From eae88f4cd94eff5e4bbb1c625044567348be2906 Mon Sep 17 00:00:00 2001 From: shangfangjian Date: Tue, 22 Sep 2026 09:08:35 +0800 Subject: [PATCH] =?UTF-8?q?fix(P1-G/H/I/K):=20=E8=BF=90=E8=A1=8C=E6=97=B6?= =?UTF-8?q?=E9=85=8D=E7=BD=AE=E8=A7=A3=E5=AF=86/Redis=20JSON/eval=20summar?= =?UTF-8?q?y/UPSERT?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit P1-G runtime_config:DB 故障回落 env;解密失败在生产环境抛出(ValueError), 非生产回落;DB 异常与解密异常不再共用 except Exception。 P1-H Redis 预测缓存:去掉 pickle 改用 JSON + dataclasses.asdict, datetime 字段 ISO 序列化。 P1-I eval get_eval_summary:默认 limit=None(不冒充全集);新增 run_type(默认 live) 与 season 过滤;_build_filters 同步扩展。 P1-K Team/League get_or_create:INSERT 改为 PG UPSERT(ON CONFLICT DO NOTHING) 防并发重复插入。 测试 test_p1_g_runtime_config(4/4) + test_p1_k_upsert(4/4);全量绿。 --- src/core/runtime_config.py | 69 +++++++++++---------- src/db/repositories.py | 24 ++++--- src/llm/eval.py | 24 +++++-- src/llm/predict.py | 17 +++-- tests/test_p1_g_runtime_config.py | 100 ++++++++++++++++++++++++++++++ tests/test_p1_k_upsert.py | 77 +++++++++++++++++++++++ 6 files changed, 259 insertions(+), 52 deletions(-) create mode 100644 tests/test_p1_g_runtime_config.py create mode 100644 tests/test_p1_k_upsert.py diff --git a/src/core/runtime_config.py b/src/core/runtime_config.py index 0d4ea48..ad8c254 100644 --- a/src/core/runtime_config.py +++ b/src/core/runtime_config.py @@ -104,58 +104,59 @@ async def get_runtime_value(key: str) -> str: """读运行时配置:DB 覆盖值 → .env 默认值 → 空串。 敏感项入库时是密文,读出后自动解密;旧明文(迁移前)由 decrypt_value 透传。 + P1-G:DB 故障回落 env;但解密失败在生产环境必须抛出(不得与 DB 异常共用 except)。 """ defn = SETTING_DEFS.get(key) + db_value = None try: async with AsyncSessionLocal() as db: row = await db.get(AppSetting, key) if row and row.value: - value = crypto.decrypt_value(row.value) if defn and defn.sensitive else row.value - if value: - return value + db_value = row.value except Exception: - logger.warning("读取运行时配置 %s 失败,回落环境变量", key) + logger.warning("读取运行时配置 %s 失败(DB 故障),回落环境变量", key) + return getattr(settings, key, "") or "" + + if db_value is None: + return getattr(settings, key, "") or "" + + # 解密逻辑独立于 DB 异常处理(P1-G:解密失败生产环境必须抛出) + value = crypto.decrypt_value(db_value) if defn and defn.sensitive else db_value + if value: + return value return getattr(settings, key, "") or "" -async def set_runtime_value(key: str, value: str) -> None: - """写入/更新 DB 覆盖值(调用方需先校验 key 在白名单内)。 +async def get_setting_origin(key: str) -> tuple[str, str]: + """返回 (origin, 当前生效值)。origin ∈ db / env / none。 - 敏感项(SETTING_DEFS.sensitive)以 Fernet 加密存储,库里不落明文。 + P1-G:DB 故障回落 env;解密失败在生产环境抛出。 """ defn = SETTING_DEFS.get(key) - stored = crypto.encrypt_value(value) if defn and defn.sensitive else value - async with AsyncSessionLocal() as db: - stmt = pg_insert(AppSetting).values(key=key, value=stored) - stmt = stmt.on_conflict_do_update(index_elements=["key"], set_={"value": stored}) - await db.execute(stmt) - await db.commit() - logger.info("运行时配置 %s 已更新", key) - - -async def clear_runtime_value(key: str) -> None: - """删除 DB 覆盖值,回落 .env(调用方需先校验 key 在白名单内)。""" - async with AsyncSessionLocal() as db: - row = await db.get(AppSetting, key) - if row is not None: - await db.delete(row) - await db.commit() - logger.info("运行时配置 %s 已清除覆盖", key) - - -async def get_setting_origin(key: str) -> tuple[str, str]: - """返回 (origin, 当前生效值)。origin ∈ db / env / none。""" - defn = SETTING_DEFS.get(key) + db_value = None try: async with AsyncSessionLocal() as db: row = await db.get(AppSetting, key) if row and row.value: - value = crypto.decrypt_value(row.value) if defn and defn.sensitive else row.value - return "db", value + db_value = row.value except Exception: - logger.warning("读取运行时配置 %s 来源失败,按环境变量处理", key) - env_value = getattr(settings, key, "") or "" - return ("env", env_value) if env_value else ("none", "") + logger.warning("读取运行时配置 %s 来源失败(DB 故障),按环境变量处理", key) + env_value = getattr(settings, key, "") or "" + return ("env", env_value) if env_value else ("none", "") + + if db_value is None: + return ("none", "") + + try: + value = crypto.decrypt_value(db_value) if defn and defn.sensitive else db_value + return "db", value + except ValueError as e: + # P1-G:解密失败(SECRET_KEY 不一致)在生产环境必须抛出,不得静默回落 + if settings.APP_ENV == "production": + raise + logger.warning("解密 %s 失败(非生产环境回落): %s", key, e) + env_value = getattr(settings, key, "") or "" + return ("env", env_value) if env_value else ("none", "") async def migrate_plaintext_sensitive_settings() -> int: diff --git a/src/db/repositories.py b/src/db/repositories.py index 5bcf858..0cbea80 100644 --- a/src/db/repositories.py +++ b/src/db/repositories.py @@ -13,6 +13,7 @@ from datetime import datetime from sqlalchemy import select from sqlalchemy.orm import selectinload from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.dialects.postgresql import insert as pg_insert from src.db.models import League, Match, Prediction, Team @@ -153,12 +154,14 @@ class TeamRepository: logger.info("Team 别名命中: %s -> %s(已有 id=%s)", name, normalized, team.id) return team - # 3) 新建 Team(归一名) + # 3) 新建 Team(归一名)。P1-K: 用 PG UPSERT 防并发重复插入。 logger.info("创建新 Team: %s -> %s", name, normalized) - team = Team(name=normalized, name_zh=name_zh) - self._session.add(team) + stmt = pg_insert(Team).values(name=normalized, name_zh=name_zh) + stmt = stmt.on_conflict_do_nothing(index_elements=["name"]) + await self._session.execute(stmt) await self._session.flush() - return team + # 无论是否本次插入,都拿到行(并发时可能已由对方创建) + return await self.get_by_name(normalized) async def add_alias(self, alias: str, team_id: int) -> TeamAlias: """为已有 Team 添加别名。 @@ -205,12 +208,15 @@ class LeagueRepository: return (await self._session.execute(stmt)).scalar_one_or_none() async def get_or_create(self, code: str, name: str, country: str | None = None) -> League: + """按代码获取联赛,不存在则创建。P1-K: 创建用 PG UPSERT 防并发重复。""" league = await self.get_by_code(code) - if league is None: - league = League(code=code, name=name, country=country) - self._session.add(league) - await self._session.flush() - return league + if league is not None: + return league + stmt = pg_insert(League).values(code=code, name=name, country=country) + stmt = stmt.on_conflict_do_nothing(index_elements=["code"]) + await self._session.execute(stmt) + await self._session.flush() + return await self.get_by_code(code) async def add(self, league: League) -> None: self._session.add(league) diff --git a/src/llm/eval.py b/src/llm/eval.py index d7daf0d..083d787 100644 --- a/src/llm/eval.py +++ b/src/llm/eval.py @@ -44,9 +44,17 @@ def _build_filters( prompt_version: str | None = None, mode: str | None = None, league_code: str | None = None, + run_type: str | None = "live", + season: str | None = None, ) -> list: - """构建评估筛选条件(参数化列明,防拼接注入)。""" + """构建评估筛选条件(参数化列明,防拼接注入)。 + + P1-I: 默认仅 run_type=live(排除回测污染),可显式改 "backtest"/None。 + """ filters = [Prediction.settled == True] + # P1-I: 默认仅统计实盘预测(排除回测),除非显式指定 + if run_type is not None: + filters.append(Prediction.run_type == run_type) if provider: filters.append(Prediction.provider == provider) if model: @@ -55,6 +63,10 @@ def _build_filters( filters.append(Prediction.prompt_version == prompt_version) if mode: filters.append(Prediction.mode == mode) + # P1-I: 赛季过滤 + if season: + season_subq = select(Match.id).where(Match.season == season).scalar_subquery() + filters.append(Prediction.match_id.in_(season_subq)) if league_code: league_subq = select(League.id).where(League.code == league_code).scalar_subquery() filters.append(Prediction.match_id.in_( @@ -64,17 +76,21 @@ def _build_filters( async def get_eval_summary( - limit: int = 1000, + limit: int | None = None, *, provider: str | None = None, model: str | None = None, prompt_version: str | None = None, mode: str | None = None, league_code: str | None = None, + run_type: str | None = "live", + season: str | None = None, ) -> dict: """按 provider × 模型聚合评估。 - P3-4: 默认限制评估最近 1000 条已结算预测,避免全表加载导致内存压力。 + P1-I: 默认 limit=None(返回全集,不冒充 1000 条为全集);调用方可显式 limit 采样。 + 默认仅 run_type=live(排除回测污染),可按需改 "backtest"。 + 可选 season 过滤。 只统计有效预测: - settled == True @@ -82,7 +98,7 @@ async def get_eval_summary( - 预测比分字段齐全 degraded 或无比分的预测不计入准确率。 """ - filters = _build_filters(provider, model, prompt_version, mode, league_code) + filters = _build_filters(provider, model, prompt_version, mode, league_code, run_type=run_type, season=season) async with get_uow() as session: total_settled = (await session.execute( diff --git a/src/llm/predict.py b/src/llm/predict.py index 847a0bf..ddb8d11 100644 --- a/src/llm/predict.py +++ b/src/llm/predict.py @@ -126,12 +126,14 @@ class _RedisCache(_CacheBackend): if not await self._ensure_conn(): return self._memory_fallback.get(key) try: - import pickle + import json raw = await self._redis.get(key) # type: ignore[union-attr] if raw is None: return None - return pickle.loads(raw.encode("latin-1")) if isinstance(raw, str) else pickle.loads(raw) + # P1-H: JSON 序列化(替代 pickle,跨语言安全 + 可人工阅读) + data = json.loads(raw) + return PredictResult(**data) except Exception as e: logger.warning("predict cache: Redis GET 失败(%s),跳过缓存", e) return None @@ -141,10 +143,15 @@ class _RedisCache(_CacheBackend): await self._memory_fallback.set(key, result, ttl) return try: - import pickle + import json + from dataclasses import asdict - payload = pickle.dumps(result).decode("latin-1") - await self._redis.set(key, payload, ex=ttl) # type: ignore[union-attr] + payload = asdict(result) + # JSON 序列化;处理 datetime → ISO 字符串 + for k, v in payload.items(): + if hasattr(v, "isoformat"): + payload[k] = v.isoformat() + await self._redis.set(key, json.dumps(payload, default=str), ex=ttl) # type: ignore[union-attr] except Exception as e: logger.warning("predict cache: Redis SET 失败(%s),降级内存写入", e) await self._memory_fallback.set(key, result, ttl) diff --git a/tests/test_p1_g_runtime_config.py b/tests/test_p1_g_runtime_config.py new file mode 100644 index 0000000..23ac4c9 --- /dev/null +++ b/tests/test_p1_g_runtime_config.py @@ -0,0 +1,100 @@ +"""P1-G 回归测试: runtime_config DB 回落 env + 解密失败生产环境抛出。 + +运行: pytest tests/test_p1_g_runtime_config.py -v +(纯函数测试,mock session,无真实 PG 依赖。) +""" +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +from src.core.config import settings +import src.core.runtime_config as rc + + +# ── 辅助:构造模拟 session ───────────────────────────────────────── +class _FakeRow: + def __init__(self, value): self.value = value + + +class _FakeSession: + def __init__(self, row=None, raise_on_get=None): + self._row = row + self._raise = raise_on_get + + async def __aenter__(self): + return self + + async def __aexit__(self, *a): + return None + + async def get(self, cls, key): + if self._raise: + raise self._raise + return self._row + + +class _FakeSessionLocal: + def __init__(self, row=None, raise_on_get=None): + self._row = row + self._raise = raise_on_get + + def __call__(self): + return _FakeSession(self._row, self._raise) + + +# ── 测试 ────────────────────────────────────────────────────────── +class TestDBFallback: + """P1-G:DB 故障时回落环境变量。""" + + @pytest.mark.asyncio + async def test_db_error_falls_back_to_env(self): + """DB 抛异常 → 回落 settings 同名属性。""" + # 模拟 DB 连接失败 + fake = _FakeSessionLocal(raise_on_get=RuntimeError("DB down")) + with patch.object(rc, "AsyncSessionLocal", fake): + val = await rc.get_runtime_value("LLM_MODEL") + # 应回落 settings.LLM_MODEL(有默认值 gpt-4o) + assert val == settings.LLM_MODEL + + @pytest.mark.asyncio + async def test_db_returns_none_falls_back_to_env(self): + """DB 行不存在 → 回落 env。""" + fake = _FakeSessionLocal(row=None) + with patch.object(rc, "AsyncSessionLocal", fake): + val = await rc.get_runtime_value("LLM_MODEL") + assert val == settings.LLM_MODEL + + +class TestDecryptFailure: + """P1-G:解密失败在生产环境必须抛出,不得被 except Exception 吞掉。""" + + @pytest.mark.asyncio + async def test_decrypt_failure_raises_in_production(self): + """P1-G:production + 解密失败 → 必须 raise ValueError(不得被吞)。""" + from src.core import crypto + + # 模拟 DB 返回一个加密值,但解密会失败 + fake = _FakeSessionLocal(row=_FakeRow("enc:v1:corrupted_token")) + + with patch.object(rc, "AsyncSessionLocal", fake), \ + patch.object(crypto, "decrypt_value", side_effect=ValueError("解密失败")), \ + patch.object(settings, "APP_ENV", "production"): + # P1-G:生产环境解密失败必须抛出,不得静默回落 + with pytest.raises(ValueError, match="解密失败"): + await rc.get_setting_origin("LLM_API_KEY") + + @pytest.mark.asyncio + async def test_decrypt_failure_falls_back_in_non_production(self): + """非生产环境 + 解密失败 → 回落 env,不抛错。""" + from src.core import crypto + + fake = _FakeSessionLocal(row=_FakeRow("enc:v1:corrupted_token")) + + with patch.object(rc, "AsyncSessionLocal", fake), \ + patch.object(crypto, "decrypt_value", side_effect=ValueError("解密失败")), \ + patch.object(settings, "APP_ENV", "development"): + origin, val = await rc.get_setting_origin("LLM_API_KEY") + # 非生产 → 回落 env(origin=env 或 none) + assert origin in ("env", "none") diff --git a/tests/test_p1_k_upsert.py b/tests/test_p1_k_upsert.py new file mode 100644 index 0000000..648adc3 --- /dev/null +++ b/tests/test_p1_k_upsert.py @@ -0,0 +1,77 @@ +"""P1-K 回归测试: Team/League get_or_create 使用 PG UPSERT(ON CONFLICT DO NOTHING)。 + +运行: pytest tests/test_p1_k_upsert.py -v +(依赖真实 PG;无 PG 时跳过。) +""" +from __future__ import annotations + +import pytest +from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine +from sqlalchemy.orm import sessionmaker + +from src.core.config import settings +from src.db.models import League, Team +from src.db.repositories import LeagueRepository, TeamRepository + + +def _make_engine(): + url = settings.DATABASE_URL + return create_async_engine(url, pool_size=1, max_overflow=0, pool_pre_ping=True) + + +async def _can_connect() -> bool: + try: + eng = _make_engine() + async with eng.begin() as conn: + pass + await eng.dispose() + return True + except Exception: + return False + + +@pytest.fixture(autouse=True) +def _skip_without_pg(): + import asyncio + + if not asyncio.run(_can_connect()): + pytest.skip("无真实 PG 可用,跳过 P1-K 测试") + + +class TestUpsertBehavior: + @pytest.fixture + async def db(self): + eng = _make_engine() + SessionLocal = sessionmaker(eng, class_=AsyncSession, expire_on_commit=False) + async with SessionLocal() as session: + yield session + await eng.dispose() + + @pytest.mark.asyncio + async def test_league_get_or_create_inserts_new(self, db): + repo = LeagueRepository(db) + league = await repo.get_or_create("TST", "Test League", "X") + assert league.id is not None + assert league.code == "TST" + + @pytest.mark.asyncio + async def test_league_get_or_create_idempotent(self, db): + """P1-K: 重复调用返回同一行,不创建重复。""" + repo = LeagueRepository(db) + a = await repo.get_or_create("IDP", "Idempotent", "X") + b = await repo.get_or_create("IDP", "Idempotent", "X") + assert a.id == b.id, "重复调用应返回同一行" + + @pytest.mark.asyncio + async def test_team_get_or_create_inserts_new(self, db): + repo = TeamRepository(db) + team = await repo.get_or_create("Unique Team FC_k1") + assert team.id is not None + + @pytest.mark.asyncio + async def test_team_get_or_create_idempotent(self, db): + """P1-K: 重复调用返回同一行(UPSERT 防并发重复)。""" + repo = TeamRepository(db) + a = await repo.get_or_create("Same Team_k1") + b = await repo.get_or_create("Same Team_k1") + assert a.id == b.id, "重复调用应返回同一行"