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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 07:37:13