You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Django单元测试Mock数据库失效:实际数据库遭修改求助

问题:Django单元测试未Mock Firebase Firestore,直接修改真实数据库

我在为POST方法编写Django单元测试时,发现测试没有使用Mock数据库,反而持续修改实际的Firebase Firestore数据库。以下是相关代码:

测试方法

def test_post_contest(self):
    mock_collection = self.mock_db.collection.return_value
    mock_document = mock_collection.document.return_value

    mock_document.set.return_value = None

    response = self.client.post(
        reverse('contest-list'),
        data={
            'ONI2024F2': {
                'name': 'ONI2024F2', 
                'problems': ['problem1', 'problem2'],
                'administratorId': '0987654321'
            }
        },
        format='json'
    )
    self.assertEqual(response.status_code, 201)
    self.assertEqual(response.json(), {'id': 'ONI2024F2'})

setUp方法

def setUp(self):
    self.client = APIClient()
    self.mock_db = patch('firebase_config.db').start()
    self.addCleanup(patch.stopall)

firebase_config.py文件

import os
import firebase_admin
from firebase_admin import credentials, firestore

#Path to Firebase credentials
cred_path = os.path.join(os.path.dirname(__file__), 'credentials','firebase_credentials.json')

#Initialize Firebase
cred = credentials.Certificate(cred_path)
firebase_admin.initialize_app(cred)

#Firestore client
db = firestore.client()

完整测试文件

from django.test import TestCase
from django.urls import reverse
from unittest.mock import patch, MagicMock
from rest_framework.test import APIClient

class ContestViewTest(TestCase):
    
    def setUp(self):
        self.client = APIClient()
        self.mock_db = patch('firebase_config.db').start()
        self.addCleanup(patch.stopall)

    @patch('OIECApp.views.ContestView.referenceToJson')
    def test_get_contest_list(self, mock_referenceToJson):
        mock_referenceToJson.return_value = {'name': 'ONI2024F1'}

        mock_collection = self.mock_db.collection.return_value
        mock_doc = MagicMock()
        mock_doc.id = 'ONI2024F1'
        mock_collection.stream.return_value = [mock_doc]

        response = self.client.get(reverse('contest-list'))

        self.assertEqual(response.status_code, 200)
        self.assertEqual(response.json(), {'IOI2023': {'name': 'ONI2024F1'}, 'ONI2024F1': {'name': 'ONI2024F1'}})

    def test_post_contest(self):
        mock_collection = self.mock_db.collection.return_value
        mock_document = mock_collection.document.return_value

        mock_document.set.return_value = None

        response = self.client.post(
            reverse('contest-list'),
            data={
                'ONI2024F2': {
                    'name': 'ONI2024F2', 
                    'problems': ['problem1', 'problem2'],
                    'administratorId': '0987654321'
                }
            },
            format='json'
        )
        self.assertEqual(response.status_code, 201)
        self.assertEqual(response.json(), {'id': 'ONI2024F2'})

    @patch('OIECApp.views.ContestView.referenceToJson')
    def test_get_contest_detail(self, mock_referenceToJson):
        mock_referenceToJson.return_value = {'name': 'ONI2024F1'}

        mock_collection = self.mock_db.collection.return_value
        mock_doc = MagicMock()
        mock_doc.id = 'ONI2024F1'
        mock_collection.stream.return_value = [mock_doc]

        response = self.client.get(reverse('contest-detail', args=['ONI2024F1']))

        self.assertEqual(response.status_code, 200)
        self.assertEqual(response.json(), {'ONI2024F1': {'name': 'ONI2024F1'}})

    def test_put_contest(self):
        mock_collection = self.mock_db.collection.return_value
        mock_document = mock_collection.document.return_value

        # Mock the update operation
        mock_document.update.return_value = None

        response = self.client.put(
            reverse('contest-detail', args=['ONI2025F1']),
            data={
                'ONI2025F1': {
                    'name': 'ONI2025F1',
                    'problems': ['problem1', 'problem3'],
                    'administratorId': '0987654321'
                }
            },
            format='json'
        )
        self.assertEqual(response.status_code, 201)
        self.assertEqual(response.json(), {'id': 'ONI2025F1'})

    def test_delete_contest(self):
        mock_collection = self.mock_db.collection.return_value
        mock_document = mock_collection.document.return_value

        # Mock the delete operation
        mock_document.delete.return_value = None

        response = self.client.delete(reverse('contest-detail', args=['ONI2024F2']))
        self.assertEqual(response.status_code, 204)

HTTP方法实现文件

from django.http import Http404, JsonResponse
from rest_framework import status
from rest_framework.response import Response
from rest_framework.views import APIView
from firebase_config import db
from firebase_admin import firestore

collection_name = "Contest"

def referenceToJson(element):
    doc_dict = element.to_dict()
    for x in doc_dict:
        tmp = doc_dict[x]
        if isinstance(tmp, list):
            counter = 0
            for y in tmp:
                if isinstance(y, firestore.firestore.DocumentReference):
                    doc_ref = y.get().id
                    tmp[counter]=doc_ref
                counter += 1
        else:
            if isinstance(tmp, firestore.firestore.DocumentReference):
                doc_ref = tmp.get().id
                doc_dict[x]=doc_ref
    return doc_dict

class ContestList(APIView):

    def get(self, request):
        contest_ref = db.collection(collection_name)
        contests = dict ()
        
        for doc in contest_ref.stream():
            idContest = doc.id
            doc_dict = referenceToJson(doc)
            contests[idContest]= doc_dict
        return JsonResponse(contests)
    
    def post(self, request):
        try: 
            datos = request.data
            clave = list(datos.keys())[0]
            valores = datos[clave]
            for c in valores.keys():
                if(c == 'Problems') or (c == 'problems'):
                    if isinstance(valores[c], list):
                        contador = 0
                        for element in valores[c]:
                            print(element)
                            ingreso = db.collection('Problems').document(valores[c][contador])
                            valores[c][contador]=ingreso
                            contador += 1
                    else:
                        ingreso = db.collection('Problems').document(valores[c])
                        valores[c] = ingreso
                elif(c == 'administratorId'):
                    ingreso = db.collection('Administrator').document(valores[c])
                    valores[c] = ingreso
            db.collection(collection_name).document(clave).set(valores)
            return JsonResponse({'id':clave}, status=201)
        except Exception as e:  
            return JsonResponse({'error':str(e)}, status=500)

class ContestDetail(APIView):

    def get(self, request, id):
        contest_ref = db.collection(collection_name)
        contest = dict ()
        
        for doc in contest_ref.stream():
            if(doc.id == id):
                idContest = doc.id
                doc_dict = referenceToJson(doc)
                contest[idContest] = doc_dict

        return JsonResponse(contest)

    def put(self,request,id):
        try: 
            datos = request.data
            clave = list(datos.keys())[0]
            valores = datos[clave]
            for c in valores.keys():
                if(c == 'Problems') or (c == 'problems'):
                    if isinstance(valores[c], list):
                        list_ref = []
                        for element in valores[c]:
                            ingreso = db.collection('Problems').document(element)
                            list_ref.append(ingreso)
                        valores[c] = list_ref
                    else:
                        ingreso = db.collection('Problems').document(valores[c])
                        valores[c] = ingreso
                elif(c == 'administratorId'):
                    ingreso = db.collection('Administrator').document(valores[c])
                    valores[c] = ingreso
            db.collection(collection_name).document(clave).update(valores)
            return JsonResponse({'id':clave}, status=201)
        except Exception as e:  
            return JsonResponse({'error':str(e)}, status=500)

    def delete(self, request, id):
        try:
            contest_ref = db.collection(collection_name).document(id)
            contest_ref.delete()
            return JsonResponse({'message': 'Deleted'}, status=204)
        except Exception as e:
            return JsonResponse({'error': str(e)}, status=500)

排查与解决方案

问题根源

你当前在测试中patch('firebase_config.db'),但在视图文件OIECApp.views中,是直接导入了db对象——视图加载时就已经引用了真实的db实例。Python的unittest.mock.patch需要作用于被使用的位置,而非定义的位置,所以当前Mock并未生效,测试仍在调用真实数据库。

修复步骤

  1. 修正测试中的patch目标
    将测试中patch('firebase_config.db')改为patch('OIECApp.views.db'),因为视图文件里直接使用了导入的db对象,必须Mock这个引用位置:

    def setUp(self):
        self.client = APIClient()
        # 修改patch路径为视图中使用db的位置
        self.mock_db = patch('OIECApp.views.db').start()
        self.addCleanup(patch.stopall)
    
  2. 覆盖所有数据库操作的Mock
    POST方法中不仅操作了Contest集合,还调用了Problems和Administrator集合,需要确保这些调用都被Mock处理:

    def test_post_contest(self):
        # 创建通用的mock集合和文档对象
        mock_collection = MagicMock()
        mock_document = MagicMock()
        mock_document.set.return_value = None
    
        # 让collection()无论传入什么参数都返回mock集合
        self.mock_db.collection.return_value = mock_collection
        # 让document()返回mock文档
        mock_collection.document.return_value = mock_document
    
        response = self.client.post(
            reverse('contest-list'),
            data={
                'ONI2024F2': {
                    'name': 'ONI2024F2', 
                    'problems': ['problem1', 'problem2'],
                    'administratorId': '0987654321'
                }
            },
            format='json'
        )
        self.assertEqual(response.status_code, 201)
        self.assertEqual(response.json(), {'id': 'ONI2024F2'})
    
        # 验证是否调用了正确的集合和文档操作
        self.mock_db.collection.assert_any_call('Problems')
        self.mock_db.collection.assert_any_call('Administrator')
        self.mock_db.collection.assert_any_call('Contest')
        mock_collection.document.assert_called_with('ONI2024F2')
        mock_document.set.assert_called_once()
    
  3. 可选:避免测试时初始化真实Firebase
    修改firebase_config.py,仅在非测试环境初始化Firebase:

    import os
    import firebase_admin
    from firebase_admin import credentials, firestore
    
    db = None
    
    def initialize_firebase():
        global db
        if not firebase_admin._apps:
            cred_path = os.path.join(os.path.dirname(__file__), 'credentials','firebase_credentials.json')
            cred = credentials.Certificate(cred_path)
            firebase_admin.initialize_app(cred)
            db = firestore.client()
    
    # 非测试环境自动初始化
    if not os.environ.get('TESTING'):
        initialize_firebase()
    

    然后在测试的setUp中设置环境变量:

    def setUp(self):
        os.environ['TESTING'] = 'True'
        self.client = APIClient()
        self.mock_db = patch('OIECApp.views.db').start()
        self.addCleanup(patch.stopall)
        self.addCleanup(lambda: os.environ.pop('TESTING', None))
    

验证修复

修改后重新运行测试,所有数据库操作将指向Mock对象,不会再修改真实的Firestore数据库。


内容的提问来源于stack exchange,提问作者Fausto Briones

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.19 09:35:54