fix(P1-G/H/I/K): 运行时配置解密/Redis JSON/eval summary/UPSERT

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);全量绿。
This commit is contained in:
shangfangjian
2026-09-22 09:08:35 +08:00
parent 2b52478b8f
commit eae88f4cd9
6 changed files with 259 additions and 52 deletions
+35 -34
View File
@@ -104,58 +104,59 @@ async def get_runtime_value(key: str) -> str:
"""读运行时配置:DB 覆盖值 → .env 默认值 → 空串。 """读运行时配置:DB 覆盖值 → .env 默认值 → 空串。
敏感项入库时是密文,读出后自动解密;旧明文(迁移前)由 decrypt_value 透传。 敏感项入库时是密文,读出后自动解密;旧明文(迁移前)由 decrypt_value 透传。
P1-G:DB 故障回落 env;但解密失败在生产环境必须抛出(不得与 DB 异常共用 except)。
""" """
defn = SETTING_DEFS.get(key) defn = SETTING_DEFS.get(key)
db_value = None
try: try:
async with AsyncSessionLocal() as db: async with AsyncSessionLocal() as db:
row = await db.get(AppSetting, key) row = await db.get(AppSetting, key)
if row and row.value: if row and row.value:
value = crypto.decrypt_value(row.value) if defn and defn.sensitive else row.value db_value = row.value
if value:
return value
except Exception: 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 "" return getattr(settings, key, "") or ""
async def set_runtime_value(key: str, value: str) -> None: async def get_setting_origin(key: str) -> tuple[str, str]:
"""写入/更新 DB 覆盖值(调用方需先校验 key 在白名单内) """返回 (origin, 当前生效值)。origin ∈ db / env / none
敏感项(SETTING_DEFS.sensitive)以 Fernet 加密存储,库里不落明文 P1-G:DB 故障回落 env;解密失败在生产环境抛出
""" """
defn = SETTING_DEFS.get(key) defn = SETTING_DEFS.get(key)
stored = crypto.encrypt_value(value) if defn and defn.sensitive else value db_value = None
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)
try: try:
async with AsyncSessionLocal() as db: async with AsyncSessionLocal() as db:
row = await db.get(AppSetting, key) row = await db.get(AppSetting, key)
if row and row.value: if row and row.value:
value = crypto.decrypt_value(row.value) if defn and defn.sensitive else row.value db_value = row.value
return "db", value
except Exception: except Exception:
logger.warning("读取运行时配置 %s 来源失败,按环境变量处理", key) logger.warning("读取运行时配置 %s 来源失败(DB 故障),按环境变量处理", key)
env_value = getattr(settings, key, "") or "" env_value = getattr(settings, key, "") or ""
return ("env", env_value) if env_value else ("none", "") 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: async def migrate_plaintext_sensitive_settings() -> int:
+15 -9
View File
@@ -13,6 +13,7 @@ from datetime import datetime
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.orm import selectinload from sqlalchemy.orm import selectinload
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.dialects.postgresql import insert as pg_insert
from src.db.models import League, Match, Prediction, Team 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) logger.info("Team 别名命中: %s -> %s(已有 id=%s)", name, normalized, team.id)
return team return team
# 3) 新建 Team(归一名) # 3) 新建 Team(归一名)。P1-K: 用 PG UPSERT 防并发重复插入。
logger.info("创建新 Team: %s -> %s", name, normalized) logger.info("创建新 Team: %s -> %s", name, normalized)
team = Team(name=normalized, name_zh=name_zh) stmt = pg_insert(Team).values(name=normalized, name_zh=name_zh)
self._session.add(team) stmt = stmt.on_conflict_do_nothing(index_elements=["name"])
await self._session.execute(stmt)
await self._session.flush() await self._session.flush()
return team # 无论是否本次插入,都拿到行(并发时可能已由对方创建)
return await self.get_by_name(normalized)
async def add_alias(self, alias: str, team_id: int) -> TeamAlias: async def add_alias(self, alias: str, team_id: int) -> TeamAlias:
"""为已有 Team 添加别名。 """为已有 Team 添加别名。
@@ -205,12 +208,15 @@ class LeagueRepository:
return (await self._session.execute(stmt)).scalar_one_or_none() 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: 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) league = await self.get_by_code(code)
if league is None: if league is not None:
league = League(code=code, name=name, country=country) return league
self._session.add(league) stmt = pg_insert(League).values(code=code, name=name, country=country)
await self._session.flush() stmt = stmt.on_conflict_do_nothing(index_elements=["code"])
return league await self._session.execute(stmt)
await self._session.flush()
return await self.get_by_code(code)
async def add(self, league: League) -> None: async def add(self, league: League) -> None:
self._session.add(league) self._session.add(league)
+20 -4
View File
@@ -44,9 +44,17 @@ def _build_filters(
prompt_version: str | None = None, prompt_version: str | None = None,
mode: str | None = None, mode: str | None = None,
league_code: str | None = None, league_code: str | None = None,
run_type: str | None = "live",
season: str | None = None,
) -> list: ) -> list:
"""构建评估筛选条件(参数化列明,防拼接注入)。""" """构建评估筛选条件(参数化列明,防拼接注入)。
P1-I: 默认仅 run_type=live(排除回测污染),可显式改 "backtest"/None。
"""
filters = [Prediction.settled == True] filters = [Prediction.settled == True]
# P1-I: 默认仅统计实盘预测(排除回测),除非显式指定
if run_type is not None:
filters.append(Prediction.run_type == run_type)
if provider: if provider:
filters.append(Prediction.provider == provider) filters.append(Prediction.provider == provider)
if model: if model:
@@ -55,6 +63,10 @@ def _build_filters(
filters.append(Prediction.prompt_version == prompt_version) filters.append(Prediction.prompt_version == prompt_version)
if mode: if mode:
filters.append(Prediction.mode == 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: if league_code:
league_subq = select(League.id).where(League.code == league_code).scalar_subquery() league_subq = select(League.id).where(League.code == league_code).scalar_subquery()
filters.append(Prediction.match_id.in_( filters.append(Prediction.match_id.in_(
@@ -64,17 +76,21 @@ def _build_filters(
async def get_eval_summary( async def get_eval_summary(
limit: int = 1000, limit: int | None = None,
*, *,
provider: str | None = None, provider: str | None = None,
model: str | None = None, model: str | None = None,
prompt_version: str | None = None, prompt_version: str | None = None,
mode: str | None = None, mode: str | None = None,
league_code: str | None = None, league_code: str | None = None,
run_type: str | None = "live",
season: str | None = None,
) -> dict: ) -> dict:
"""按 provider × 模型聚合评估。 """按 provider × 模型聚合评估。
P3-4: 默认限制评估最近 1000 条已结算预测,避免全表加载导致内存压力 P1-I: 默认 limit=None(返回全集,不冒充 1000 条为全集);调用方可显式 limit 采样
默认仅 run_type=live(排除回测污染),可按需改 "backtest"
可选 season 过滤。
只统计有效预测: 只统计有效预测:
- settled == True - settled == True
@@ -82,7 +98,7 @@ async def get_eval_summary(
- 预测比分字段齐全 - 预测比分字段齐全
degraded 或无比分的预测不计入准确率。 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: async with get_uow() as session:
total_settled = (await session.execute( total_settled = (await session.execute(
+12 -5
View File
@@ -126,12 +126,14 @@ class _RedisCache(_CacheBackend):
if not await self._ensure_conn(): if not await self._ensure_conn():
return self._memory_fallback.get(key) return self._memory_fallback.get(key)
try: try:
import pickle import json
raw = await self._redis.get(key) # type: ignore[union-attr] raw = await self._redis.get(key) # type: ignore[union-attr]
if raw is None: if raw is None:
return 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: except Exception as e:
logger.warning("predict cache: Redis GET 失败(%s),跳过缓存", e) logger.warning("predict cache: Redis GET 失败(%s),跳过缓存", e)
return None return None
@@ -141,10 +143,15 @@ class _RedisCache(_CacheBackend):
await self._memory_fallback.set(key, result, ttl) await self._memory_fallback.set(key, result, ttl)
return return
try: try:
import pickle import json
from dataclasses import asdict
payload = pickle.dumps(result).decode("latin-1") payload = asdict(result)
await self._redis.set(key, payload, ex=ttl) # type: ignore[union-attr] # 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: except Exception as e:
logger.warning("predict cache: Redis SET 失败(%s),降级内存写入", e) logger.warning("predict cache: Redis SET 失败(%s),降级内存写入", e)
await self._memory_fallback.set(key, result, ttl) await self._memory_fallback.set(key, result, ttl)
+100
View File
@@ -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")
+77
View File
@@ -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, "重复调用应返回同一行"