如何在FastAPI单元测试中正确Mock第三方Qdrant客户端库?
如何在FastAPI单元测试中正确Mock第三方Qdrant客户端库?
我来帮你梳理下问题出在哪,你现在的Mock没生效主要是两个原因:patch的目标路径不对,再加上你的import语句有点混淆,导致Mock没覆盖到代码里实际在使用的Qdrant客户端实例。
第一步:先修正main.py里的import语句(很关键)
你当前main.py里的import写法:
import qdrant_client as QdrantClient
其实是把整个qdrant_client模块导入并别名为QdrantClient,之后你用QdrantClient(url=...)实例化客户端——这种写法很容易混淆模块和类,也会给后续Mock带来麻烦。推荐改成标准的类导入方式:
# 推荐写法:直接导入QdrantClient类 from qdrant_client import QdrantClient qdrant_client = QdrantClient(url=...)
或者如果你想保留模块引用,就明确调用类:
import qdrant_client # 明确使用模块下的QdrantClient类 qdrant_client = qdrant_client.QdrantClient(url=...)
第二步:用正确的姿势patch覆盖客户端
不管你用哪种import方式,核心原则只有一个:patch的目标必须是你的代码中实际使用那个对象的位置,也就是app.main模块里的那个qdrant_client全局实例,或者你导入的类引用。
方案一:直接patch全局的客户端实例(最省心)
因为你在main.py里初始化了一个全局的qdrant_client实例,所有接口都会用到它,直接patch这个全局变量是最简单的方案:
from unittest import TestCase, patch, MagicMock from fastapi.testclient import TestClient from app.main import app class TestMyEndpoint(TestCase): @patch('app.main.qdrant_client') def test_my_endpoint(self, mock_qdrant): # Mock你需要用到的方法返回值,根据业务逻辑调整 # 如果启动时会调用create_collection,先mock它 mock_qdrant.create_collection.return_value = None # 再mock query_points的返回结果 mock_qdrant.query_points.return_value = { "result": [{"id": 1, "payload": {"content": "测试内容"}}] } with TestClient(app) as client: resp = client.get("/my_endpoint", params={"query": "你好"}) # 断言响应状态码符合预期 self.assertEqual(resp.status_code, 200) # 验证query_points是否被正确调用 mock_qdrant.query_points.assert_called_once() # 也可以更精确地验证调用参数 # mock_qdrant.query_points.assert_called_once_with(collection_name="my_col", query_vector=...)
方案二:从类层面patch QdrantClient
如果你想从类的角度Mock,确保patch的是main.py中导入的那个类引用:
比如你用了from qdrant_client import QdrantClient,那patch目标就是'app.main.QdrantClient',然后在测试中Mock它的实例方法:
from unittest import TestCase, patch, MagicMock from fastapi.testclient import TestClient from app.main import app class TestMyEndpoint(TestCase): @patch('app.main.QdrantClient') def test_my_endpoint(self, mock_qdrant_class): # 创建一个Mock实例,让类返回这个实例 mock_qdrant_instance = MagicMock() mock_qdrant_class.return_value = mock_qdrant_instance # 给实例的方法设置返回值 mock_qdrant_instance.create_collection.return_value = None mock_qdrant_instance.query_points.return_value = {"result": []} with TestClient(app) as client: resp = client.get("/my_endpoint", params={"query": "测试查询"}) self.assertEqual(resp.status_code, 200) # 验证QdrantClient类是否被正确实例化 mock_qdrant_class.assert_called_once_with(url="你配置的url") # 验证query_points方法是否被调用 mock_qdrant_instance.query_points.assert_called_once()
为什么你原来的代码没生效?
你原来的写法@patch('app.main.QdrantClient', qdrant_client=MagicMock())有两个明显的问题:
- patch的参数用法错了,
@patch的第二个参数是new(用来指定Mock对象),不能用关键字参数qdrant_client=...,正确写法是@patch('app.main.QdrantClient', new=MagicMock()),或者直接写@patch('app.main.QdrantClient')让patch自动创建MagicMock。 - 更核心的是,你patch的是
QdrantClient类(或者模块),但代码里实际在使用的是已经初始化好的全局qdrant_client实例——这个实例是在patch执行前就创建的,所以Mock的类根本影响不到它,代码还是会调用真实的客户端方法。
备注:内容来源于stack exchange,提问作者Clovis
相关产品推荐
相关产品推荐

