新功能:个性化推荐算法

This commit is contained in:
吕新雨
2026-02-02 16:47:37 +08:00
parent 936094211b
commit 6dc4e2b943
119 changed files with 7427 additions and 357 deletions

223
server/tests/conftest.py Normal file
View File

@@ -0,0 +1,223 @@
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()