Pytest测试FastAPI/SQLAlchemy应用:无法切换至测试数据库
FastAPI/SQLAlchemy测试依赖覆盖失效,仍连接生产数据库问题排查
我正在为FastAPI/SQLAlchemy应用编写测试,想要使用独立的空测试数据库。已经在conftest.py中添加了依赖覆盖,但override_get_db()函数从未被调用,导致测试仍在生产数据库上运行,请求排查代码问题。
main.py
from fastapi import FastAPI from fastapi.middleware.cors import CORSMiddleware from routes.address import router as address_router app = FastAPI() app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) app.include_router(address_router)
database.py
from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker from models.base import Base from config import Config from sqlalchemy.orm import Session engine = create_engine( Config.DATABASE_URI, echo=True, ) SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine) def get_db(): print(f"Connecting to database: {Config.DATABASE_URI}") Base.metadata.create_all(engine) db: Session = SessionLocal() try: yield db finally: db.close()
routes/address.py
from fastapi import APIRouter, Depends from sqlalchemy.orm import Session from crud.address import get, get_all, create, update, delete from database.database import get_db from schemas.address import AddressCreate router = APIRouter() @router.get("/address/{address_id}") async def get_address(address_id: int, db: Session = Depends(get_db)): return get(db, address_id) @router.get("/address/") async def get_all_addresss(db: Session = Depends(get_db)): return get_all(db) @router.post("/address/") async def create_address(address: AddressCreate, db: Session = Depends(get_db)): return create(db, address) @router.put("/address/{address_id}") async def update_address( address_id: int, address: AddressCreate, db: Session = Depends(get_db) ): return update(db, address_id, address) @router.delete("/address/{address_id}") async def delete_address(address_id: int, db: Session = Depends(get_db)): return delete(db, address_id)
conftest.py
import pytest from fastapi.testclient import TestClient from sqlalchemy import Engine, StaticPool, create_engine from sqlalchemy.orm import sessionmaker from main import app from config import Config from src.database.database import get_db from src.models.base import Base print("Loading conftest.py") TEST_DATABASE_URI = "sqlite:///:memory:" @pytest.fixture(scope="session") def engine() -> Engine: print(f"Using database URI: {Config.TEST_DATABASE_URI}") return create_engine( Config.TEST_DATABASE_URI, connect_args={"check_same_thread": False}, poolclass=StaticPool, echo=True, ) @pytest.fixture(scope="function") def test_db(engine): Base.metadata.create_all(bind=engine) TestingSessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine) db = TestingSessionLocal() try: yield db finally: db.close() Base.metadata.drop_all(bind=engine) @pytest.fixture(scope="function") def override_get_db(): def _override_get_db(): print("Using test database") try: yield test_db finally: test_db.close() return _override_get_db @pytest.fixture(scope="function") def test_app(override_get_db): print("Applying dependency override") app.dependency_overrides[get_db] = override_get_db yield app print("Clearing dependency override") app.dependency_overrides.clear() @pytest.fixture(scope="function") def client(test_app): return TestClient(test_app)
test_address.py
def test_create_address(client): response = client.post( "/address/", json={ "city": "Springfield", "country": "USA", }, ) assert response.status_code == 200 response_data = response.json() assert response_data["city"] == "Springfield" assert response_data["country"] == "USA" assert "id" in response_data
问题排查与修复
1. get_db导入路径不匹配
FastAPI的依赖覆盖基于对象引用,conftest.py中从src.database.database导入get_db,但实际应用的routes/address.py是从database.database导入的,两者是不同对象,导致覆盖失效。
修复:
统一导入路径,修改conftest.py的导入语句:
# 替换原有导入 from database.database import get_db from models.base import Base
2. override_get_db错误引用test_db fixture
override_get_db内部直接引用test_db fixture,未通过依赖注入获取,导致无法正确拿到数据库会话实例,而且test_db的关闭逻辑已经在自身fixture中处理,无需重复调用。
修复:
让override_get_db依赖test_db,调整代码如下:
@pytest.fixture(scope="function") def override_get_db(test_db): # 注入test_db fixture def _override_get_db(): print("Using test database") try: yield test_db finally: pass # test_db的关闭已在自身fixture中处理 return _override_get_db
3. 额外优化:移除get_db中的表创建逻辑
database.py的get_db函数每次请求都执行Base.metadata.create_all(engine),会影响性能,且测试环境的表创建已经在test_db fixture中处理。建议将表创建逻辑移到应用启动脚本或初始化步骤:
# 修改后的database.py get_db函数 def get_db(): print(f"Connecting to database: {Config.DATABASE_URI}") db: Session = SessionLocal() try: yield db finally: db.close()
验证修复
运行测试后,控制台会打印Using test database,测试将使用内存中的独立数据库,不会操作生产数据库,测试完成后数据自动销毁。
内容的提问来源于stack exchange,提问作者MatGood
相关产品推荐
相关产品推荐

