FastAPI单元测试困惑:如何隔离测试端点与CRUD函数
问题背景
我使用FastAPI开发后端应用,采用SQLModel(结合SQLAlchemy与Pydantic)并连接PostgreSQL数据库。目前已基于预发布PG数据库完成集成测试,验证端点功能正常,但不知如何编写单元测试来隔离测试端点及被调用的函数。核心困惑是crud/items.py与routers/items.py中所有函数均依赖db_session: Session参数,不知如何实现隔离测试。同时希望得到代码优化建议。
项目简化架构如下:
app/ ├── api/ │ ├── core/ │ │ ├── config.py # 获取环境配置并分发至应用 │ │ ├── .env │ ├── crud/ │ │ ├── items.py # 路由调用的CRUD函数 │ ├── db/ │ │ ├── session.py # 处理数据库引擎的get_session函数 │ ├── models/ │ │ ├── items.py # 数据库中的SQLModel对象定义 │ ├── routers/ │ │ ├── items.py # 路由系统 │ ├── schemas/ │ │ ├── items.py # 应用中使用的Python对象定义 │ ├── main.py # 主应用 ├── tests/ # pytest测试用例 │ ├── unit_tests/ │ ├── integration_tests/ │ │ ├── test_items.py
单元测试解决方案
单元测试的核心是隔离外部依赖,不需要连接真实PostgreSQL,通过Mock模拟数据库会话(Session)的行为即可。以下是具体实现:
1. 单元测试CRUD函数(crud/items.py)
直接给CRUD函数传入Mock的Session,验证函数逻辑(如SQL查询执行、数据处理)是否正确。
创建tests/unit_tests/test_crud_items.py:
from unittest.mock import Mock, patch from sqlmodel import select from api.crud.items import get_item, create_new_item from api.models.items import Item from api.schemas.items import ItemCreate def test_get_item(): # 模拟Session和查询结果 mock_session = Mock() mock_item = Item(id=1, city_name="Paris") mock_exec = Mock() mock_exec.first.return_value = mock_item mock_session.exec.return_value = mock_exec # 调用函数 result = get_item(mock_session, item_id=1) # 验证查询逻辑 mock_session.exec.assert_called_once_with(select(Item).where(Item.id == 1)) assert result == mock_item def test_create_new_item(): mock_session = Mock() item_create = ItemCreate(city_name="London") mock_item = Item(id=2, city_name="London") # 模拟refresh方法为对象赋值id def mock_refresh(obj): obj.id = 2 mock_session.refresh.side_effect = mock_refresh # 调用函数 result = create_new_item(mock_session, obj_input=item_create) # 验证CRUD操作 mock_session.add.assert_called_once() added_item = mock_session.add.call_args[0][0] assert added_item.city_name == "London" mock_session.commit.assert_called_once() mock_session.refresh.assert_called_once_with(added_item) assert result.id == 2
2. 单元测试路由(routers/items.py)
利用FastAPI的app.dependency_overrides替换get_session为Mock的Session,结合TestClient验证HTTP逻辑(状态码、返回值、异常处理)。
创建tests/unit_tests/test_routers_items.py:
from unittest.mock import Mock, patch from fastapi.testclient import TestClient from api.main import app from api.db.session import get_session from api.models.items import Item client = TestClient(app) def test_read_item_success(): # 替换Session依赖 mock_session = Mock() def override_get_session(): yield mock_session app.dependency_overrides[get_session] = override_get_session # Mock CRUD函数返回值 with patch("api.routers.items.get_item") as mock_get_item: mock_get_item.return_value = Item(id=1, city_name="Berlin") response = client.get("/api/items/1") assert response.status_code == 200 assert response.json() == {"id": 1, "city_name": "Berlin"} mock_get_item.assert_called_once_with(db_session=mock_session, item_id=1) # 清理依赖覆盖 app.dependency_overrides.clear() def test_read_item_not_found(): mock_session = Mock() def override_get_session(): yield mock_session app.dependency_overrides[get_session] = override_get_session with patch("api.routers.items.get_item") as mock_get_item: mock_get_item.return_value = None response = client.get("/api/items/999") assert response.status_code == 404 assert response.json() == {"detail": "Item not found"} app.dependency_overrides.clear() def test_create_item(): mock_session = Mock() def override_get_session(): yield mock_session app.dependency_overrides[get_session] = override_get_session with patch("api.routers.items.create_new_item") as mock_create: mock_create.return_value = Item(id=3, city_name="Madrid") response = client.post("/api/items/", json={"city_name": "Madrid"}) assert response.status_code == 200 assert response.json() == {"id": 3, "city_name": "Madrid"} assert mock_create.call_args[1]["obj_input"].city_name == "Madrid" app.dependency_overrides.clear()
代码优化建议
1. 数据库配置与Session管理
- 调整.env位置:将
.env从api/core移到项目根目录(与app同级),避免被打包进API模块,方便不同环境(开发/测试/生产)使用独立配置文件。 - 添加测试专用Session:在
db/session.py中新增内存SQLite的Session生成函数,用于集成测试时隔离数据:
集成测试中替换依赖:# db/session.py新增 def get_test_session(): test_engine = create_engine("sqlite:///:memory:", connect_args={"check_same_thread": False}) SQLModel.metadata.create_all(test_engine) with Session(test_engine) as session: yield session# tests/integration_tests/test_items.py修改 from api.db.session import get_session, get_test_session app.dependency_overrides[get_session] = get_test_session client = TestClient(app)
2. 模型与Schema优化
- 统一字段必填性:当前Item模型中
city_name是必填项,但ItemBase中定义为Optional[str],会导致创建数据时数据库报错。建议将ItemBase中的city_name改为city_name: str,保持与模型一致。 - 简化Schema设置:SQLModel 0.4+版本支持直接在
__table_args__中指定数据库schema,无需单独使用MetaData:# models/items.py修改 class Item(SQLModel, table=True): __tablename__ = "cities" __table_args__ = {"schema": "pouetpouet"} id: int = Field(primary_key=True, default=None) city_name: str
3. CRUD函数优化
- 移除不必要的编码转换:
create_new_item中无需使用jsonable_encoder,直接通过obj_input.dict()解包即可:def create_new_item(db_session: Session, *, obj_input: ItemCreate) -> Item: db_obj = Item(**obj_input.dict()) db_session.add(db_obj) db_session.commit() db_session.refresh(db_obj) return db_obj - 拆分异常逻辑:CRUD函数可抛出自定义异常(如
ItemNotFound),路由层捕获后转换为HTTPException,职责更清晰。
4. 测试结构优化
- 添加pytest配置:在项目根目录创建
pytest.ini,指定测试路径和Python路径:[pytest] testpaths = tests pythonpath = . - 复用测试Fixtures:创建
tests/conftest.py定义通用fixture,减少重复代码:import pytest from unittest.mock import Mock from fastapi.testclient import TestClient from api.main import app from api.db.session import get_session @pytest.fixture def mock_session(): return Mock() @pytest.fixture def client_with_mock_session(mock_session): def override_get_session(): yield mock_session app.dependency_overrides[get_session] = override_get_session yield TestClient(app) app.dependency_overrides.clear()
内容的提问来源于stack exchange,提问作者FloCAD
相关产品推荐
相关产品推荐

