224 lines
7.1 KiB
Python
224 lines
7.1 KiB
Python
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()
|
||
|