from __future__ import annotations import os import sys from pathlib import Path from typing import AsyncIterator, Callable from urllib.parse import urlparse import pytest import pytest_asyncio import sqlalchemy as sa from alembic import command from alembic.config import Config from sqlalchemy import event from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker, create_async_engine # 确保在任何 pytest rootdir 下都能 `import app.*` SERVER_DIR = Path(__file__).resolve().parents[1] # .../server if str(SERVER_DIR) not in sys.path: sys.path.insert(0, str(SERVER_DIR)) def _read_env_kv(env_path: Path) -> dict[str, str]: """ 读取 .env 文件中的 KEY=VALUE(最小实现,避免引入额外依赖)。 """ data: dict[str, str] = {} if not env_path.exists(): return data for raw in env_path.read_text(encoding="utf-8").splitlines(): line = raw.strip() if not line or line.startswith("#"): continue if "=" not in line: continue k, v = line.split("=", 1) k = k.strip() v = v.strip().strip('"').strip("'") if k: data[k] = v return data def _get_database_url() -> str: """ 获取测试用数据库连接串。 约定(与 alembic/env.py 保持一致): - 优先读取环境变量 `DATABASE_URL` - 若未设置,则按 `APP_ENV`(默认 dev)读取 `server/.env.dev` 或 `server/.env.prod` """ env_url = (os.getenv("DATABASE_URL") or "").strip() if env_url: return env_url server_dir = Path(__file__).resolve().parents[1] # .../server app_env = (os.getenv("APP_ENV") or "dev").strip() or "dev" env_file = server_dir / (".env.prod" if app_env == "prod" else ".env.dev") kv = _read_env_kv(env_file) url = (kv.get("DATABASE_URL") or "").strip() if url: return url raise RuntimeError( "缺少 DATABASE_URL:请设置环境变量 DATABASE_URL,或在 server/.env.dev(或 .env.prod)中配置 DATABASE_URL。" ) def _assert_safe_mysql_test_db(url: str) -> None: """ 为了避免对开发库造成破坏性影响,集成测试只允许连接到“测试库”。 规则(V1): - 必须是 mysql+aiomysql://... - 为避免误连生产库:不允许数据库名为 'mindfulness'(prod 默认库名) - 建议使用独立测试库(例如 mindfulness_dev_test) """ if not url.startswith("mysql+"): raise RuntimeError(f"当前仅允许 MySQL 集成测试(mysql+aiomysql)。实际:{url!r}") parsed = urlparse(url.replace("mysql+aiomysql://", "mysql://", 1)) db_name = (parsed.path or "").lstrip("/") if db_name.lower() == "mindfulness": raise RuntimeError( "为避免误连生产库,集成测试不允许连接到数据库 'mindfulness'。" "请改用 dev 测试库(例如 mindfulness_dev_test 或 mindfulness_dev)。" ) def _run_alembic_upgrade_head() -> None: """ 使用 Alembic 将测试库升级到最新 schema。 说明: - 依赖 env.py 内部读取 DATABASE_URL - 仅在 session 级别执行一次,避免每个测试都跑迁移 """ server_dir = Path(__file__).resolve().parents[1] # .../server alembic_ini = server_dir / "alembic.ini" cfg = Config(str(alembic_ini)) # 确保脚本路径正确(alembic.ini 里一般已配置,这里兜底) cfg.set_main_option("script_location", "alembic") command.upgrade(cfg, "head") def _assert_schema_exists(url: str) -> None: """ 非破坏性检查:要求目标库已经存在所需表。 说明: - 默认不在测试中运行 Alembic(避免任何 schema 变更) - 若要自动迁移,请设置环境变量 ALLOW_SCHEMA_MIGRATION=1 """ allow_migration = (os.getenv("ALLOW_SCHEMA_MIGRATION") or "").strip() == "1" if allow_migration: _run_alembic_upgrade_head() return # 使用 PyMySQL 做同步检查,避免依赖 MySQLdb(不要求系统安装 mysqlclient) sync_url = url.replace("mysql+aiomysql://", "mysql+pymysql://", 1) engine = sa.create_engine(sync_url, future=True) try: insp = sa.inspect(engine) tables = set(insp.get_table_names()) required = {"contents", "content_profiles", "content_risk_flags"} missing = sorted(required - tables) if missing: raise RuntimeError( "集成测试检测到 schema 不完整(缺少表:" + ", ".join(missing) + ")。为避免破坏性操作,测试不会自动迁移。" "请先手动在该库执行 `alembic upgrade head`,或设置 ALLOW_SCHEMA_MIGRATION=1 允许测试自动迁移。" ) finally: engine.dispose() @pytest.fixture(scope="session", autouse=True) def _migrate_db_once() -> None: """ Session 级 schema 检查(仅在配置了数据库连接时启用)。 说明: - 纯函数单元测试不需要 MySQL;若未配置 DATABASE_URL,则跳过检查 - 集成测试(依赖 db_session/async_engine)仍会在获取 DATABASE_URL 时失败,从而提示用户配置 """ try: url = _get_database_url() except RuntimeError: return _assert_safe_mysql_test_db(url) _assert_schema_exists(url) @pytest_asyncio.fixture async def async_engine() -> AsyncIterator[AsyncEngine]: url = _get_database_url() _assert_safe_mysql_test_db(url) engine = create_async_engine(url, pool_pre_ping=True) try: yield engine finally: await engine.dispose() @pytest_asyncio.fixture async def db_session(async_engine: AsyncEngine) -> AsyncIterator[AsyncSession]: """ 提供一个干净的 AsyncSession。 清理策略:每个测试都在事务中执行,并在结束时回滚(不做 DELETE/TRUNCATE)。 说明: - 用例里严禁调用 session.commit(),只允许 flush() - 这样不会对测试库产生持久化写入,更不会影响开发库 """ SessionLocal: async_sessionmaker[AsyncSession] = async_sessionmaker( bind=async_engine, expire_on_commit=False, autoflush=False, autocommit=False, ) async with SessionLocal() as session: trans = await session.begin() try: yield session finally: await trans.rollback() @pytest.fixture def query_counter(async_engine: AsyncEngine) -> Callable[[], int]: """ 返回一个函数:调用可获得当前累计查询次数。 """ count = {"n": 0} def before_cursor_execute(*args, **kwargs): # type: ignore[no-untyped-def] count["n"] += 1 event.listen(async_engine.sync_engine, "before_cursor_execute", before_cursor_execute) def get_count() -> int: return int(count["n"]) def fin() -> None: event.remove(async_engine.sync_engine, "before_cursor_execute", before_cursor_execute) # 用 yield 确保测试后移除监听,避免重复绑定导致统计偏大 try: yield get_count # type: ignore[misc] finally: fin()