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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 15:10:00