FastAPI+pytest:如何覆写带参数的FastAPI依赖?
如何正确覆写FastAPI中带参数的依赖函数?
问题场景
我有一个符合MCVE规范的FastAPI应用,运行正常:
# run.py import uvicorn from fastapi import FastAPI, Depends app = FastAPI() # 带参数的依赖生成函数 def dep_with_arg(inp): def sub_dep(): return inp return sub_dep @app.get("/a") def a(v: str = Depends(dep_with_arg("a"))): return v # 无参数的依赖函数 def dep_without_arg(): return "b" @app.get("/b") def b(v: str = Depends(dep_without_arg)): return v def main(): uvicorn.run( "run:app", host="0.0.0.0", reload=True, port=8000, workers=1 ) if __name__ == "__main__": main()
注意/a和/b端点的依赖差异:/a使用了接受参数的依赖生成函数,调用dep_with_arg("a")会返回内部的sub_dep函数作为实际依赖;而/b直接使用无参数的依赖函数。两个端点运行均正常。
我尝试在测试中覆写依赖,但test_b通过,所有针对/a的测试都失败:
测试代码:
from run import app, dep_without_arg, dep_with_arg from starlette.testclient import TestClient def test_b(): def dep_override_for_dep_without_arg(): return "bbb" test_client = TestClient(app=app) test_client.app.dependency_overrides[dep_without_arg] = dep_override_for_dep_without_arg resp = test_client.get("/b") assert resp.json() == "bbb" # 以下测试均失败 def test_a_method_1(): def dep_override_for_dep_with_arg(inp): def sub_dep(): return "aaa" return sub_dep test_client = TestClient(app=app) test_client.app.dependency_overrides[dep_with_arg] = dep_override_for_dep_with_arg resp = test_client.get("/a") assert resp.json() == "aaa" def test_a_method_2(): def dep_override_for_dep_with_arg(inp): return "aaa" test_client = TestClient(app=app) test_client.app.dependency_overrides[dep_with_arg] = dep_override_for_dep_with_arg resp = test_client.get("/a") assert resp.json() == "aaa" def test_a_method_3(): def dep_override_for_dep_with_arg(inp): return "aaa" test_client = TestClient(app=app) test_client.app.dependency_overrides[dep_with_arg] = dep_override_for_dep_with_arg("aaa") resp = test_client.get("/a") assert resp.json() == "aaa" def test_a_method_4(): def dep_override_for_dep_with_arg(): return "aaa" test_client = TestClient(app=app) test_client.app.dependency_overrides[dep_with_arg] = dep_override_for_dep_with_arg resp = test_client.get("/a") assert resp.json() == "aaa"
失败结果:
FAILED test.py::test_a_method_1 - AssertionError: assert 'a' == 'aaa' FAILED test.py::test_a_method_2 - AssertionError: assert 'a' == 'aaa' FAILED test.py::test_a_method_3 - AssertionError: assert 'a' == 'aaa' FAILED test.py::test_a_method_4 - AssertionError: assert 'a' == 'aaa'
核心原因
FastAPI在处理Depends(dep_with_arg("a"))时,会先执行dep_with_arg("a")得到内部的sub_dep函数,之后实际注册到端点的依赖是这个sub_dep,而非外部的dep_with_arg函数。之前的测试代码都是直接覆写dep_with_arg,根本不会影响到已经生成的sub_dep实例,所以覆写无效。
正确解决方案
方法1:暴露生成后的依赖实例(推荐)
修改原代码,将dep_with_arg("a")生成的sub_dep保存为变量,直接在端点中引用这个变量:
# run.py 修改部分 def dep_with_arg(inp): def sub_dep(): return inp return sub_dep # 提前生成依赖实例并保存 sub_dep_a = dep_with_arg("a") @app.get("/a") def a(v: str = Depends(sub_dep_a)): return v
测试时直接覆写这个sub_dep_a实例:
def test_a_correct(): def dep_override(): return "aaa" test_client = TestClient(app=app) from run import sub_dep_a test_client.app.dependency_overrides[sub_dep_a] = dep_override resp = test_client.get("/a") assert resp.json() == "aaa"
这种方式清晰直观,代码可读性高,是最推荐的方案。
方法2:不修改原代码,动态查找并替换依赖实例
如果不想改动原代码,可以通过遍历FastAPI的路由列表,找到/a端点对应的依赖实例,再进行覆写:
def test_a_without_modifying_source(): def dep_override(): return "aaa" test_client = TestClient(app=app) # 找到/a对应的路由 for route in test_client.app.routes: if route.path == "/a" and route.name == "a": # 遍历端点的依赖项 for dep in route.dependencies: # 确认是dep_with_arg生成的sub_dep if hasattr(dep.dependency, '__closure__'): test_client.app.dependency_overrides[dep.dependency] = dep_override break break resp = test_client.get("/a") assert resp.json() == "aaa"
这种方式不需要修改原代码,但逻辑相对复杂,适合无法改动源码的场景。
方法3:重构依赖设计(长期方案)
如果需要频繁对这类带参数的依赖进行测试,可以重构依赖的设计方式,比如使用类依赖或者将参数通过其他方式注入(比如请求对象、配置等),让依赖本身更容易被覆写。例如:
# 重构后的依赖 class DepWithArg: def __init__(self, inp): self.inp = inp def __call__(self): return self.inp # 端点使用 @app.get("/a") def a(v: str = Depends(DepWithArg("a"))): return v
测试时可以直接覆写DepWithArg类的实例:
def test_a_refactored(): def dep_override(): return "aaa" test_client = TestClient(app=app) # 找到对应的DepWithArg实例并覆写 for route in test_client.app.routes: if route.path == "/a": for dep in route.dependencies: if isinstance(dep.dependency, DepWithArg): test_client.app.dependency_overrides[dep.dependency] = dep_override break break resp = test_client.get("/a") assert resp.json() == "aaa"
内容的提问来源于stack exchange,提问作者Amin Ba
相关产品推荐
相关产品推荐

