Files
mindfulness/server/tests/conftest.py
2026-02-02 16:47:37 +08:00

224 lines
7.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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()