FastAPI集成测试:实现get_db依赖覆盖与变更检测回滚
FastAPI集成测试:事务回滚与数据变更检测
核心思路
- 每个测试用例独立运行在数据库事务中,测试结束后回滚事务,彻底避免数据污染
- 在事务回滚前,通过检查SQLAlchemy会话状态、对比表行数两种方式,确认测试是否产生数据变更
修改后的conftest.py代码
from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker, event from fastapi.testclient import TestClient import pytest from app import app, get_db # 导入项目的app和get_db依赖 from app.models import Base # 导入模型基类,用于创建测试表 # 测试数据库URL,根据实际环境调整 DB_URL = "sqlite:///./test.db" engine = create_engine(DB_URL) TestingSessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine) # 全局测试数据库初始化(仅执行一次) @pytest.fixture(scope="session") def setup_test_db(): # 创建所有表结构 Base.metadata.create_all(bind=engine) # 初始化测试基础数据 db = TestingSessionLocal() additional_db_init(db) # 你的自定义初始化逻辑 db.commit() db.close() yield # 测试全部结束后清理表 Base.metadata.drop_all(bind=engine) # 每个测试用例独立的数据库会话(带事务回滚) @pytest.fixture(scope="function") def db(setup_test_db): connection = engine.connect() # 开启顶层事务,测试结束后统一回滚 transaction = connection.begin() session = TestingSessionLocal(bind=connection) # 记录初始表行数,用于后续变更对比 initial_counts = {} # 替换为你的实际模型类 from app.models import Item, User for model in [Item, User]: initial_counts[model.__tablename__] = session.query(model).count() # 开启嵌套事务,支持测试代码内的commit操作(不会影响顶层事务) session.begin_nested() # 自动重启嵌套事务,适配测试中的commit调用 @event.listens_for(session, "after_transaction_end") def restart_savepoint(session, transaction): if transaction.nested and not transaction._parent.nested: session.begin_nested() yield session # 检测数据变更 has_changes = False # 1. 检查会话中的变更对象(新增/修改/删除) if session.new or session.dirty or session.deleted: has_changes = True print("\n测试产生数据变更:") if session.new: print(f"新增对象:{[obj.__class__.__name__ for obj in session.new]}") if session.dirty: print(f"修改对象:{[obj.__class__.__name__ for obj in session.dirty]}") if session.deleted: print(f"删除对象:{[obj.__class__.__name__ for obj in session.deleted]}") # 2. 对比表行数变化(可选,更直观) current_counts = {} for model in [Item, User]: current_counts[model.__tablename__] = session.query(model).count() if current_counts[model.__tablename__] != initial_counts[model.__tablename__]: has_changes = True print(f"表 {model.__tablename__} 行数变化:{initial_counts[model.__tablename__]} → {current_counts[model.__tablename__]}") if not has_changes: print("\n测试未产生数据变更") # 清理会话与事务 session.close() transaction.rollback() connection.close() # 每个测试用例独立的TestClient,绑定带事务的db会话 @pytest.fixture(scope="function") def test_client(db): # 覆盖get_db依赖,返回当前测试的db会话 def override_get_db(): try: yield db finally: pass # db fixture会统一处理会话关闭和事务回滚 app.dependency_overrides[get_db] = override_get_db with TestClient(app) as client: yield client # 清除依赖覆盖,避免影响其他测试 del app.dependency_overrides[get_db]
关键说明
- 事务隔离:每个测试用例的数据库操作都在独立事务中运行,测试结束后回滚顶层事务,保证所有修改不会持久化到数据库
- 嵌套事务兼容:通过事件监听自动重启嵌套事务,允许测试代码中调用
db.commit(),实际仅在嵌套事务内生效,不影响全局数据 - 变更检测:同时通过会话对象状态和表行数对比两种方式,全面检测测试过程中的数据变更,便于调试和验证测试逻辑
- fixture作用域:
setup_test_db为session级(仅执行一次),负责表创建和全局数据初始化;db和test_client为function级,确保每个测试环境独立干净
测试用例示例
def test_add_item(test_client): # 发送创建item的请求 response = test_client.post("/item", json={"name": "Test Item", "price": 19.99}) assert response.status_code == 200 # 验证返回数据正确性 data = response.json() assert data["name"] == "Test Item"
测试结束后,控制台会输出该测试是否产生数据变更,所有修改会被自动回滚,下一个测试启动时数据库状态与初始完全一致。
内容的提问来源于stack exchange,提问作者Ruuza
相关产品推荐
相关产品推荐

