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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 21:05:09