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
+32 -31
View File
@@ -104,56 +104,57 @@ 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
db_value = row.value
except Exception:
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
except Exception:
logger.warning("读取运行时配置 %s 失败,回落环境变量", key)
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)
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", "")
+14 -8
View File
@@ -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()
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)
+20 -4
View File
@@ -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(
+12 -5
View File
@@ -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)
+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, "重复调用应返回同一行"