如何为FastAPI异步数据库连接创建测试数据库以开展测试?
为FastAPI+SQLAlchemy应用创建独立测试数据库的解决方案
1. 调整配置与数据库连接代码
首先让应用支持区分生产和测试环境的数据库,修改配置逻辑,从主数据库URL派生测试库地址:
# settings.py from pydantic_settings import BaseSettings class Settings(BaseSettings): DB_CONNECT: str @property def TEST_DB_CONNECT(self): # 根据数据库类型生成测试库URL if "postgresql" in self.DB_CONNECT: base_url, db_name = self.DB_CONNECT.rsplit("/", 1) return f"{base_url}/{db_name}_test" elif "mysql" in self.DB_CONNECT: base_url, db_name = self.DB_CONNECT.rsplit("/", 1) return f"{base_url}/{db_name}_test" elif "sqlite" in self.DB_CONNECT: return "sqlite:///./test_temp.db" else: raise ValueError("不支持的数据库类型") settings = Settings()
接着修改数据库连接代码,使其支持传入自定义URL,方便测试时替换:
# your_app/database.py import databases from sqlalchemy import create_engine from sqlalchemy.orm import declarative_base from settings import settings def get_database(url: str = settings.DB_CONNECT): return databases.Database(url) def get_engine(url: str = settings.DB_CONNECT): connect_args = {"check_same_thread": False} if "sqlite" in url else {} return create_engine(url, connect_args=connect_args) engine = get_engine() database = get_database() Base = declarative_base()
2. 编写测试专用pytest夹具
实现创建测试库→建表→执行测试→销毁资源的完整流程:
# tests/conftest.py import pytest_asyncio from sqlalchemy import text from settings import settings from your_app.database import get_engine, get_database, Base @pytest_asyncio.fixture(scope="session") async def test_db_url(): return settings.TEST_DB_CONNECT @pytest_asyncio.fixture(scope="session") async def test_engine(test_db_url): return get_engine(test_db_url) @pytest_asyncio.fixture(scope="session") async def test_database(test_db_url): db = get_database(test_db_url) await db.connect() yield db await db.disconnect() @pytest_asyncio.fixture(scope="session", autouse=True) async def setup_test_database(test_engine, test_db_url): # 创建测试数据库(需数据库用户有建库权限) if "postgresql" in test_db_url: default_engine = get_engine(test_db_url.rsplit("/", 1)[0] + "/postgres") with default_engine.connect() as conn: conn.execute(text(f"DROP DATABASE IF EXISTS {test_db_url.rsplit('/',1)[1]}")) conn.execute(text(f"CREATE DATABASE {test_db_url.rsplit('/',1)[1]}")) conn.commit() elif "mysql" in test_db_url: default_engine = get_engine(test_db_url.rsplit("/", 1)[0] + "/mysql") with default_engine.connect() as conn: conn.execute(text(f"DROP DATABASE IF EXISTS {test_db_url.rsplit('/',1)[1]}")) conn.execute(text(f"CREATE DATABASE {test_db_url.rsplit('/',1)[1]}")) conn.commit() # 创建所有表结构 async with test_engine.begin() as conn: await conn.run_sync(Base.metadata.create_all) yield # 清理资源 async with test_engine.begin() as conn: await conn.run_sync(Base.metadata.drop_all) if "postgresql" in test_db_url or "mysql" in test_db_url: with default_engine.connect() as conn: conn.execute(text(f"DROP DATABASE IF EXISTS {test_db_url.rsplit('/',1)[1]}")) conn.commit()
3. 在测试用例中使用测试数据库
通过依赖替换,让测试用的FastAPI客户端连接到测试库:
# tests/test_endpoints.py from your_app.main import app from fastapi.testclient import TestClient from your_app.database import get_database import pytest_asyncio @pytest_asyncio.fixture async def client(test_database): app.dependency_overrides[get_database] = lambda: test_database yield TestClient(app) app.dependency_overrides.clear() async def test_sample_endpoint(client): response = client.get("/sample") assert response.status_code == 200
关键注意事项
- 确保数据库用户拥有创建/删除数据库的权限,否则建库步骤会失败。
- SQLite无需手动建库,直接使用临时文件或内存库即可。
- 若需要每个测试用例都重建数据库,可将夹具的
scope="session"改为scope="function"。
内容的提问来源于stack exchange,提问作者Tema Vedernikov
相关产品推荐
相关产品推荐

