如何在pytest中Mock使用上下文管理器的类方法返回值?
问题描述
我有一个database.py模块,其中包含如下使用上下文管理器实现数据库连接管理的DatabaseClient类:
from someDatabase import theClient class DatabaseClient: def __init__(self): self.connection = None def __enter__(self): self.connection = self.database_connection() return self def __exit__(self, exc_type, exc_value, traceback): if self.connection: self.connection.close() def database_connection(self): client = theClient(<connection params>) return client def database_query(self, table: str, query: str): response = self.connection.search( body = query, table = table ) return response
在Flask应用中,我有如下路由:
from app.utils.Database import Database from app.queries import thequery from flask import Blueprint, jsonify api = Blueprint('the_api', __name__) @api.route("/api/this/route", methods=["GET"]) def get_some_stuff(**kwargs): try: # 从请求获取输入参数 with Database() as db: response = db.database_query("some_table", thequery) # 转换响应对象 return jsonify(response), 200 except Exception as e: raise e
我希望仅使用pytest,MockDatabaseClient类的database_query()方法,使其返回指定的测试样本数据来测试上述Flask接口。请问正确的patch方式是什么?
我的测试代码框架如下:
import pytest from app import app from tests.sample_data.api_responses.get_some_stuff import ( get_some_stuff_api_response ) from tests.sample_data.query_responses.get_some_stuff import ( get_some_stuff_query_response ) @pytest.fixture def client(): with app.test_client() as client: yield client def test_get_some_stuff(client): # 这里应该怎么patch,让Database.database_query返回get_some_stuff_query_response? # 调用路由 response = client.get( f"/api/this/route", ) assert response.status_code == 200 data = response.json assert data == get_some_stuff_api_response
正确的Patch方式
核心原则是:patch路由代码中实际导入Database类的位置,而非原始DatabaseClient类的定义模块。因为路由里是从app.utils.Database导入的Database(推测是DatabaseClient的别名或导出类),所以要针对该导入路径进行patch。
修改后的完整测试代码如下:
import pytest from unittest.mock import patch from app import app from app.queries import thequery # 需要导入thequery用于验证调用 from tests.sample_data.api_responses.get_some_stuff import ( get_some_stuff_api_response ) from tests.sample_data.query_responses.get_some_stuff import ( get_some_stuff_query_response ) @pytest.fixture def client(): with app.test_client() as client: yield client def test_get_some_stuff(client): # 针对路由中导入Database的路径进行patch with patch('app.utils.Database.Database.database_query') as mock_query: # 设置mock方法的返回值 mock_query.return_value = get_some_stuff_query_response # 调用测试路由 response = client.get("/api/this/route") # 断言状态码和响应数据 assert response.status_code == 200 data = response.json assert data == get_some_stuff_api_response # 可选:验证mock方法是否按预期被调用 mock_query.assert_called_once_with("some_table", thequery)
关键说明
- Python的
unittest.mock.patch作用于类被导入的位置,而非类的定义位置。路由代码中从app.utils.Database导入Database,所以必须patch该路径下的类方法,才能让路由代码使用mock后的逻辑。 - 由于
Database是上下文管理器,patch它的database_query方法后,当路由执行with Database() as db:时,生成的实例会自动使用mock后的方法,返回你指定的测试样本数据。
如果app.utils.Database中的Database是DatabaseClient的别名(比如from .database import DatabaseClient as Database),patch路径也可以写成app.utils.Database.DatabaseClient.database_query,但更推荐使用路由中实际引用的导入名称对应的路径。
内容的提问来源于stack exchange,提问作者says
相关产品推荐
相关产品推荐

