fastapi:测试: 测试带数据库连接的异步函数
一,说明:
1,测试异步函数时,你必须在函数前加上 await,
并且测试用例本身也要是 async def,同时需要打上 @pytest.mark.asyncio 标签
2,在 FastAPI 中,很多 service 层或 CRUD 函数都需要传入一个数据库 AsyncSession。
测试这类函数时,我们可以结合 db_session 固件(Fixture)来实现。
二,代码:
数据库异步函数:
# 根据用户名查询得到用户信息的一条记录
async def get_user_by_username(db: AsyncSession, username: str):
"""从数据库中异步获取用户信息"""
result = await db.execute(select(User).filter(User.username == username))
return result.scalars().first()
测试配置
# tests/conftest.py (API 专有)
import asyncio
import pytest
import pytest_asyncio
from httpx import AsyncClient, ASGITransport
from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker, AsyncSession
from app.api.main import api_app
from app.core.database import Base, get_db
from main import app
# 使用独立的测试数据库
# _test
# TEST_DATABASE_URL = "mysql+aiomysql://root:rootpassword@127.0.0.1:3306/test_db"
TEST_DATABASE_URL = "mysql+aiomysql://root:rootpassword@localhost:3306/media_bank"
test_engine = create_async_engine(TEST_DATABASE_URL, echo=False)
TestingSessionLocal = async_sessionmaker(bind=test_engine, expire_on_commit=False)
# 初始化表结构(通常在测试套件开始时运行一次)
@pytest_asyncio.fixture(scope="session", autouse=True)
async def init_test_database():
pass
'''
async with test_engine.begin() as conn:
# 测试前:清空并重新创建所有表
await conn.run_sync(Base.metadata.drop_all)
await conn.run_sync(Base.metadata.create_all)
yield
async with test_engine.begin() as conn:
# 测试结束后:可选清理
await conn.run_sync(Base.metadata.drop_all)
'''
# 核心:每次测试独立的 AsyncSession,并利用外部事务自动回滚
@pytest_asyncio.fixture(scope="session")
async def db_session() -> AsyncSession:
async with test_engine.connect() as connection:
# 开启一个根事务
transaction = await connection.begin()
# 将 session 绑定到这个连接上
async with TestingSessionLocal(bind=connection) as session:
yield session
# 测试完成后,无条件回滚!数据库不会留下任何痕迹
await transaction.rollback()
# 核心:异步 HTTP 客户端
@pytest_asyncio.fixture(scope="function")
async def async_client(db_session: AsyncSession):
# 重写 FastAPI 的依赖项,注入带有自动回滚功能的 db_session
async def _get_test_db():
yield db_session
api_app.dependency_overrides[get_db] = _get_test_db
# 使用 httpx.AsyncClient 替代 TestClient
async with AsyncClient(transport=ASGITransport(app=api_app), base_url="http://127.0.0.1:8000") as client:
yield client
api_app.dependency_overrides.clear()
测试函数:
@pytest.mark.asyncio(scope="session")
async def test_get_user_by_username(db_session):
'''
# 1. 准备数据:利用测试 session 往数据库塞入一条测试数据
mock_user = User(id=1, username="admin")
db_session.add(mock_user)
await db_session.commit()
'''
# 2. 调用我们要测试的独立函数,将测试 session 作为参数传进去
user = await get_user_by_username(db=db_session, username='admin')
print('user:', user)
# 3. 断言结果
assert user is not None
assert user.username == "admin"
浙公网安备 33010602011771号