feat: 核心模块增强 — 加密 + 运行时配置 + 日志缓冲

- crypto.py: API Key 加密/解密工具
- runtime_config.py: 运行时动态配置管理
- log_buffer.py: 内存日志缓冲区
- config.py: 新增加密配置项
- http_client.py: 增强重试和错误处理
This commit is contained in:
shangfangjian
2026-09-19 11:58:03 +08:00
parent b3e2c52b49
commit 786f10aa11
57 changed files with 3178 additions and 488 deletions
+16
View File
@@ -15,17 +15,29 @@ from src.core.config import settings
async def lifespan(app: FastAPI) -> AsyncIterator[None]:
from src.db.base import init_db
from src.core.http_client import close_client
from src.core.runtime_config import (
ensure_admin_password_hashed,
migrate_plaintext_sensitive_settings,
)
await init_db() # 验证连接,不建表
await migrate_plaintext_sensitive_settings() # 明文敏感配置 → 加密(幂等)
await ensure_admin_password_hashed() # .env 明文密码 → scrypt 哈希(幂等)
yield
await close_client()
def create_app() -> FastAPI:
from src.core.log_buffer import setup_memory_logging
setup_memory_logging(settings.LOG_LEVEL)
# 生产环境不暴露 OpenAPI 文档(避免向访客泄露接口结构)
openapi_url = "/openapi.json" if settings.APP_ENV != "production" else None
app = FastAPI(
title="Profeto API",
description="足球数据 + LLM 预测服务",
version="0.1.0",
lifespan=lifespan,
openapi_url=openapi_url,
)
origins = [o.strip() for o in settings.CORS_ORIGINS.split(",") if o.strip()]
@@ -44,12 +56,16 @@ def create_app() -> FastAPI:
from src.api.routes.ingest import router as ingest_router
from src.api.routes.eval import router as eval_router
from src.api.routes.backtest import router as backtest_router
from src.api.routes.auth import router as auth_router
from src.api.routes.admin_settings import router as admin_settings_router
app.include_router(matches_router)
app.include_router(predict_router)
app.include_router(ingest_router)
app.include_router(eval_router)
app.include_router(backtest_router)
app.include_router(auth_router)
app.include_router(admin_settings_router)
@app.get("/health")
async def health():
+77 -16
View File
@@ -1,37 +1,98 @@
"""API 依赖:鉴权等横切关注点。
审查报告 P2-7:ingest / backtest / settle 这类「写入型或高成本」接口此前
完全无鉴权 —— 任何能访问到服务的人都可触发采集、或直接烧掉 LLM 额度。
策略(渐进式,不破坏本地开发):
- `ADMIN_API_KEY` 未配置 → 直接放行,并打一次 warning
这样本地 `docker compose up` 无需额外配置即可用。
- 已配置 → 必须带匹配的 `X-API-Key` 请求头,否则 401
- 管理员密码:库中 scrypt 哈希优先,回落 .env 初始值;后台可在线修改
- 密码已配置(哈希或 .env)→ 管理后台可用密码登录,登录后颁发 HttpOnly
Cookie 会话;受保护接口接受 Cookie 会话或 X-API-Key
- 仅配置 `ADMIN_API_KEY` → 受保护接口只接受 `X-API-Key` 请求头(机器/脚本调用)。
- 两者都未配置 → 直接放行,并打一次 warning(本地开发模式)。
"""
from __future__ import annotations
import hmac
import logging
import secrets
import time
from fastapi import Header, HTTPException
from fastapi import Header, HTTPException, Request
from src.core.config import settings
from src.core.runtime_config import get_admin_password_hash
logger = logging.getLogger(__name__)
SESSION_COOKIE = "profeto_session"
async def require_admin_key(x_api_key: str | None = Header(None, alias="X-API-Key")) -> None:
"""保护「写入型 / 高成本」接口的依赖。
用法: `@router.post("/ingest/bzzoiro", dependencies=[Depends(require_admin_key)])`
def _sign(exp_ts: int, secret: bytes) -> str:
msg = f"profeto-admin:{exp_ts}".encode()
return hmac.new(secret, msg, "sha256").hexdigest()
async def get_session_secret() -> bytes:
"""会话签名密钥 = HMAC(SECRET_KEY, 管理员凭证指纹)。
指纹来自密码哈希(密码本身永不参与签名):密码变更 → 指纹变化
→ 全部旧会话失效,无需额外吊销机制。
"""
expected = settings.ADMIN_API_KEY
if not expected:
from src.core.runtime_config import get_admin_credential_fingerprint
fingerprint = await get_admin_credential_fingerprint()
key = settings.SECRET_KEY or f"fallback:{settings.DATABASE_URL}"
return hmac.new(key.encode(), b"session:" + fingerprint.encode(), "sha256").digest()
def create_session_token(secret: bytes) -> str:
exp_ts = int(time.time()) + settings.ADMIN_SESSION_TTL_HOURS * 3600
return f"{exp_ts}.{_sign(exp_ts, secret)}"
def verify_session_token(token: str, secret: bytes) -> bool:
try:
exp_raw, sig = token.split(".", 1)
exp_ts = int(exp_raw)
if exp_ts < int(time.time()):
return False
return secrets.compare_digest(sig, _sign(exp_ts, secret))
except (ValueError, TypeError):
return False
async def auth_configured() -> bool:
"""是否已启用鉴权(密码哈希/.env 密码/API Key 任一)。"""
return bool(
await get_admin_password_hash()
or settings.ADMIN_PASSWORD
or settings.ADMIN_API_KEY
)
async def require_admin(
request: Request,
x_api_key: str | None = Header(None, alias="X-API-Key"),
) -> None:
"""统一保护管理接口:接受 Cookie 会话(密码登录)或 X-API-Key。
用法: `@router.get("/leagues", dependencies=[Depends(require_admin)])`
"""
if not await auth_configured():
logger.warning(
"ADMIN_API_KEY 未设置,采集/回测接口当前【无鉴权】。"
"生产环境请设置该环境变量"
"管理员密码 / ADMIN_API_KEY 未设置,管理接口当前【无鉴权】。"
"生产环境请至少设置其中一项"
)
return
if not x_api_key or not secrets.compare_digest(x_api_key, expected):
raise HTTPException(status_code=401, detail="无效或缺失的 X-API-Key")
# 1) Cookie 会话(密码登录颁发)
token = request.cookies.get(SESSION_COOKIE)
if token and verify_session_token(token, await get_session_secret()):
return
# 2) X-API-Key(机器/脚本调用;key 明文只在内存中,与第三方交互必需)
if (
settings.ADMIN_API_KEY
and x_api_key
and secrets.compare_digest(x_api_key, settings.ADMIN_API_KEY)
):
return
raise HTTPException(status_code=401, detail="未登录或凭证无效")
+319
View File
@@ -0,0 +1,319 @@
"""后台管理路由:数据源配置的查看、修改与连通性测试。
所有接口需管理员鉴权(require_admin)。配置项白名单见
src/core/runtime_config.py SETTING_DEFS,之外的 key 一律拒绝。
"""
from __future__ import annotations
import logging
import time
from datetime import date, datetime, timezone
from fastapi import APIRouter, Depends, HTTPException, Query
from pydantic import BaseModel
from sqlalchemy import func, select
import httpx
from src.api.deps import require_admin
from src.core.config import settings
from src.core.http_client import get_client
from src.core.log_buffer import get_entries
from src.core.runtime_config import (
AGENT_META,
SETTING_DEFS,
clear_runtime_value,
get_runtime_value,
get_setting_origin,
mask_value,
set_runtime_value,
)
from src.db.base import AsyncSession, get_db_read
from src.db.models import Injury, MatchStats
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/v1/admin", tags=["admin"], dependencies=[Depends(require_admin)])
# ── 数据源元数据 ────────────────────────────────────────────────
_SOURCES: list[dict] = [
{
"name": "bzzoiro",
"label": "Bzzoiro",
"description": "历史赛程与比分数据,覆盖全球主要联赛",
"setting_keys": ["BZZOIRO_KEY", "BZZOIRO_BASE"],
},
{
"name": "understat",
"label": "Understat",
"description": "xG(预期进球)进阶数据,无需 API Key,网页抓取",
"setting_keys": [],
},
{
"name": "injuries",
"label": "Injuries (API-Football)",
"description": "球员伤停信息,用于预测时考虑阵容完整性",
"setting_keys": ["API_FOOTBALL_KEY"],
},
]
class SettingUpdateIn(BaseModel):
value: str
async def _last_ingestion(db: AsyncSession, source: str) -> datetime | None:
"""各源最近一次采集时间(取自数据血缘字段,无记录返回 None)。"""
if source == "injuries":
return (await db.execute(select(func.max(Injury.retrieved_at)))).scalar()
return (
await db.execute(
select(func.max(MatchStats.retrieved_at)).where(MatchStats.source == source)
)
).scalar()
@router.get("/datasources")
async def list_datasources(db: AsyncSession = Depends(get_db_read)):
"""数据源列表:各配置项的脱敏值、来源(db/env/none)与最近采集时间。"""
result = []
for src in _SOURCES:
settings_out = []
for key in src["setting_keys"]:
origin, value = await get_setting_origin(key)
defn = SETTING_DEFS[key]
settings_out.append(
{
"key": key,
"label": defn.label,
"description": defn.description,
"sensitive": defn.sensitive,
"configured": origin != "none",
"masked": mask_value(value, defn.sensitive),
"origin": origin,
}
)
key_configured = all(s["configured"] for s in settings_out) if settings_out else True
last = await _last_ingestion(db, src["name"])
result.append(
{
"name": src["name"],
"label": src["label"],
"description": src["description"],
"key_configured": key_configured,
"last_ingestion": last.isoformat() if last else None,
"settings": settings_out,
}
)
return result
@router.get("/settings")
async def list_settings():
"""全部可配置项(脱敏),供后台各配置页渲染。"""
out = []
for key, defn in SETTING_DEFS.items():
origin, value = await get_setting_origin(key)
out.append(
{
"key": key,
"label": defn.label,
"description": defn.description,
"sensitive": defn.sensitive,
"configured": origin != "none",
"masked": mask_value(value, defn.sensitive),
"origin": origin,
}
)
return out
# ── LLM 可用模型检测 ────────────────────────────────────────────
@router.get("/logs")
async def read_logs(
level: str | None = Query(None, description="最低级别: DEBUG/INFO/WARNING/ERROR"),
keyword: str | None = Query(None, description="消息或 logger 关键字"),
limit: int = Query(200, ge=1, le=1000),
):
"""查询应用运行日志(内存环形缓冲,最新在前;进程重启后清零)。"""
entries = get_entries(level, keyword, limit)
return {"entries": entries, "count": len(entries)}
@router.get("/llm/agents")
async def list_llm_agents():
"""各专家/终裁的独立 LLM 配置状态(含当前生效模型的解析结果)。"""
out = []
for agent in AGENT_META:
aid = agent["id"].upper()
pfx = f"AGENT_{aid}_"
fields = {}
for suffix in ("MODEL", "BASE_URL", "API_KEY"):
origin, value = await get_setting_origin(f"{pfx}{suffix}")
defn = SETTING_DEFS[f"{pfx}{suffix}"]
fields[suffix.lower()] = {
"configured": origin != "none",
"masked": mask_value(value, defn.sensitive),
"origin": origin,
}
# 生效模型 = 覆盖 → 层级默认(专家/终裁 env) → 全局 LLM_MODEL
tier_default = (
settings.LLM_AGGREGATOR_MODEL if agent["id"] == "aggregator" else settings.LLM_SPECIALIST_MODEL
)
effective_model = (
fields["model"]["masked"]
if fields["model"]["configured"]
else (tier_default or await get_runtime_value("LLM_MODEL"))
)
out.append(
{
"id": agent["id"],
"label": agent["label"],
"fields": fields,
"effective_model": effective_model,
}
)
return out
@router.get("/llm/models")
async def list_llm_models():
"""探测当前 LLM 服务可用的模型列表(OpenAI 兼容 GET /models)。
只读探测,不产生费用;配置缺失或服务不可达时返回 ok=false 与原因。
"""
base_url = (await get_runtime_value("LLM_BASE_URL")).rstrip("/")
api_key = await get_runtime_value("LLM_API_KEY")
if not base_url or not api_key:
return {"ok": False, "models": [], "detail": "LLM_BASE_URL 或 LLM_API_KEY 未配置"}
client = get_client()
start = time.monotonic()
try:
resp = await client.get(
f"{base_url}/models",
headers={"Authorization": f"Bearer {api_key}"},
timeout=httpx.Timeout(connect=10.0, read=20.0, write=10.0, pool=10.0),
)
except Exception as e:
return {
"ok": False,
"models": [],
"latency_ms": int((time.monotonic() - start) * 1000),
"detail": f"无法连接 LLM 服务: {e}",
}
latency = int((time.monotonic() - start) * 1000)
if resp.status_code in (401, 403):
return {"ok": False, "models": [], "latency_ms": latency, "detail": "密钥无效或无权限(HTTP 401/403)"}
if resp.status_code != 200:
return {"ok": False, "models": [], "latency_ms": latency, "detail": f"服务返回 HTTP {resp.status_code}"}
try:
data = resp.json()
except Exception:
return {"ok": False, "models": [], "latency_ms": latency, "detail": "响应不是合法 JSON"}
models: list[str] = []
items = data.get("data") if isinstance(data, dict) else None
if isinstance(items, list):
models = sorted(
str(m.get("id")) for m in items if isinstance(m, dict) and m.get("id")
)
if not models:
return {"ok": False, "models": [], "latency_ms": latency, "detail": "服务未返回模型列表"}
return {"ok": True, "models": models, "latency_ms": latency, "detail": f"{len(models)} 个可用模型"}
@router.put("/settings/{key}")
async def update_setting(key: str, body: SettingUpdateIn):
"""更新配置项(写入 app_settings 覆盖 .env)。传空值请改用 DELETE。"""
if key not in SETTING_DEFS:
raise HTTPException(404, f"不支持的配置项: {key}")
value = body.value.strip()
if not value:
raise HTTPException(400, "值不能为空;如需回落 .env 请调用清除接口")
await set_runtime_value(key, value)
defn = SETTING_DEFS[key]
return {"key": key, "masked": mask_value(value, defn.sensitive), "origin": "db"}
@router.delete("/settings/{key}")
async def clear_setting(key: str):
"""清除 DB 覆盖值,回落 .env 默认。"""
if key not in SETTING_DEFS:
raise HTTPException(404, f"不支持的配置项: {key}")
await clear_runtime_value(key)
origin, value = await get_setting_origin(key)
defn = SETTING_DEFS[key]
return {
"key": key,
"masked": mask_value(value, defn.sensitive),
"origin": origin,
}
# ── 连通性测试 ──────────────────────────────────────────────────
_TEST_TIMEOUT = 15
async def _probe(url: str, headers: dict | None = None, params: dict | None = None) -> dict:
"""单次 HTTP 探测,返回 (ok, status, latency_ms, detail)。不重试。"""
client = get_client()
start = time.monotonic()
try:
resp = await client.get(url, headers=headers, params=params, timeout=_TEST_TIMEOUT)
except Exception as e:
return {
"ok": False,
"status": None,
"latency_ms": int((time.monotonic() - start) * 1000),
"detail": f"无法连接: {e}",
}
latency = int((time.monotonic() - start) * 1000)
status = resp.status_code
if status == 200:
detail = "连接成功"
elif status in (401, 403):
detail = "服务可达,但密钥无效或无权限"
else:
detail = f"服务返回 HTTP {status}"
return {"ok": status == 200, "status": status, "latency_ms": latency, "detail": detail}
@router.post("/datasources/{name}/test")
async def test_datasource(name: str):
"""轻量连通性测试:真实请求上游一次,不触发任何入库。"""
src = next((s for s in _SOURCES if s["name"] == name), None)
if src is None:
raise HTTPException(404, f"未知数据源: {name}")
if name == "bzzoiro":
key = await get_runtime_value("BZZOIRO_KEY")
if not key:
return {"ok": False, "status": None, "latency_ms": 0, "detail": "BZZOIRO_KEY 未配置"}
base = (await get_runtime_value("BZZOIRO_BASE")).rstrip("/")
today = date.today().isoformat()
return await _probe(
f"{base}/events/",
headers={"Authorization": f"Token {key}", "Accept": "application/json"},
params={"date_from": today, "date_to": today},
)
if name == "understat":
return await _probe(
"https://understat.com/league/EPL/2025",
headers={"User-Agent": "Mozilla/5.0", "Accept": "text/html"},
)
# injuries (api-football)
api_key = await get_runtime_value("API_FOOTBALL_KEY")
if not api_key:
return {"ok": False, "status": None, "latency_ms": 0, "detail": "API_FOOTBALL_KEY 未配置"}
return await _probe(
"https://v3.football.api-sports.io/status",
headers={"x-apisports-key": api_key},
)
+136
View File
@@ -0,0 +1,136 @@
"""管理后台认证路由:密码登录 → HttpOnly Cookie 会话;支持在线修改密码。
管理员密码以 scrypt 哈希存于数据库(.env 明文仅作初始值,启动时自动迁移为哈希)。
修改密码会改变会话签名密钥,所有已登录会话随之失效,需重新登录。
"""
from __future__ import annotations
import logging
import secrets
import time
from collections import defaultdict, deque
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from pydantic import BaseModel
from src.api.deps import (
SESSION_COOKIE,
auth_configured,
create_session_token,
get_session_secret,
require_admin,
verify_session_token,
)
from src.core import crypto
from src.core.config import settings
from src.core.runtime_config import (
get_admin_password_hash,
get_setting_origin,
set_admin_password_hash,
verify_admin_password,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/v1/auth", tags=["auth"])
# 简易防爆破:10 分钟窗口内同一 IP 连续失败 5 次即锁定 10 分钟(内存态,重启清零)
_MAX_FAILS = 5
_WINDOW_SECONDS = 600
_fail_times: dict[str, deque[float]] = defaultdict(deque)
# 新密码强度要求
_MIN_PASSWORD_LEN = 8
_MAX_PASSWORD_LEN = 128
class LoginIn(BaseModel):
password: str
class PasswordChangeIn(BaseModel):
current_password: str
new_password: str
def _client_ip(request: Request) -> str:
return request.client.host if request.client else "unknown"
def _is_locked(ip: str) -> bool:
dq = _fail_times.get(ip)
if not dq:
return False
now = time.time()
while dq and now - dq[0] > _WINDOW_SECONDS:
dq.popleft()
return len(dq) >= _MAX_FAILS
@router.post("/login")
async def login(body: LoginIn, request: Request, response: Response):
ip = _client_ip(request)
if not (await get_admin_password_hash() or settings.ADMIN_PASSWORD):
raise HTTPException(status_code=503, detail="服务器未配置管理员密码,登录不可用")
if _is_locked(ip):
logger.warning("管理员登录尝试过于频繁 (ip=%s)", ip)
raise HTTPException(status_code=429, detail="失败次数过多,请 10 分钟后再试")
if not await verify_admin_password(body.password):
_fail_times[ip].append(time.time())
logger.warning("管理员登录失败 (ip=%s)", ip)
raise HTTPException(status_code=401, detail="密码错误")
_fail_times.pop(ip, None)
response.set_cookie(
key=SESSION_COOKIE,
value=create_session_token(await get_session_secret()),
max_age=settings.ADMIN_SESSION_TTL_HOURS * 3600,
httponly=True,
samesite="lax",
path="/",
)
logger.info("管理员登录成功 (ip=%s)", ip)
return {"ok": True, "expires_in_hours": settings.ADMIN_SESSION_TTL_HOURS}
@router.post("/logout")
async def logout(response: Response):
response.delete_cookie(key=SESSION_COOKIE, path="/")
return {"ok": True}
@router.get("/me")
async def me(request: Request):
"""前端登录门禁探测。未启用鉴权时视为已登录(本地开发模式)。"""
token = request.cookies.get(SESSION_COOKIE)
authenticated = not await auth_configured() or bool(
token and verify_session_token(token, await get_session_secret())
)
has_hash = bool(await get_admin_password_hash())
return {
"authenticated": authenticated,
"enabled": await auth_configured(),
"password_origin": "db" if has_hash else ("env" if settings.ADMIN_PASSWORD else "none"),
}
@router.post("/change-password", dependencies=[Depends(require_admin)])
async def change_password(body: PasswordChangeIn, request: Request, response: Response):
"""修改管理员密码:验证当前密码 → 写运行时覆盖 → 清除会话(全端登出)。"""
if not await auth_configured():
raise HTTPException(status_code=503, detail="服务器未配置管理员密码,无法修改")
if not await verify_admin_password(body.current_password):
logger.warning("修改密码失败:当前密码错误 (ip=%s)", _client_ip(request))
raise HTTPException(status_code=401, detail="当前密码错误")
new = body.new_password
if not (_MIN_PASSWORD_LEN <= len(new) <= _MAX_PASSWORD_LEN):
raise HTTPException(status_code=400, detail=f"新密码长度需在 {_MIN_PASSWORD_LEN}-{_MAX_PASSWORD_LEN} 位之间")
if await verify_admin_password(new):
raise HTTPException(status_code=400, detail="新密码不能与当前密码相同")
await set_admin_password_hash(crypto.hash_password(new))
# 密码即会话签名密钥,修改后所有旧会话失效;主动清除当前 Cookie 要求重新登录
response.delete_cookie(key=SESSION_COOKIE, path="/")
logger.info("管理员密码已修改 (ip=%s),所有会话已失效", _client_ip(request))
return {"ok": True, "message": "密码已修改,请用新密码重新登录"}
+4 -2
View File
@@ -6,7 +6,7 @@ import logging
from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel, Field
from src.api.deps import require_admin_key
from src.api.deps import require_admin
from src.llm.backtest import run_backtest
logger = logging.getLogger(__name__)
@@ -23,7 +23,7 @@ class BacktestRequest(BaseModel):
model: str | None = Field(None, description="指定模型 (空=默认)")
@router.post("/backtest", dependencies=[Depends(require_admin_key)])
@router.post("/backtest", dependencies=[Depends(require_admin)])
async def backtest(req: BacktestRequest):
"""对历史比赛运行回测。
@@ -64,6 +64,8 @@ async def backtest(req: BacktestRequest):
"league_code": r.league_code,
"home_team": r.home_team,
"away_team": r.away_team,
"home_team_zh": r.home_team_zh,
"away_team_zh": r.away_team_zh,
"match_date": r.match_date,
"actual_score": f"{r.actual_home}-{r.actual_away}",
"actual_1x2": r.actual_1x2,
+3 -3
View File
@@ -5,7 +5,7 @@ import logging
from fastapi import APIRouter, Depends, HTTPException
from src.api.deps import require_admin_key
from src.api.deps import require_admin
from src.api.schemas import EvalSummaryOut, SettleRequest
from src.db.base import AsyncSession, get_db, get_db_read
from src.llm.eval import get_eval_summary, settle_prediction
@@ -15,7 +15,7 @@ logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/v1", tags=["eval"])
@router.post("/eval/settle", dependencies=[Depends(require_admin_key)])
@router.post("/eval/settle", dependencies=[Depends(require_admin)])
async def settle(req: SettleRequest, db: AsyncSession = Depends(get_db)):
"""回填实际结果。"""
try:
@@ -29,7 +29,7 @@ async def settle(req: SettleRequest, db: AsyncSession = Depends(get_db)):
raise HTTPException(500, "回填失败,请查看服务器日志")
@router.get("/eval/summary", response_model=EvalSummaryOut)
@router.get("/eval/summary", response_model=EvalSummaryOut, dependencies=[Depends(require_admin)])
async def eval_summary():
"""提供商/模型准确率对比。"""
return await get_eval_summary()
+91 -27
View File
@@ -1,12 +1,14 @@
"""采集路由。"""
from __future__ import annotations
import asyncio
import logging
from fastapi import APIRouter, Depends, HTTPException
from src.api.deps import require_admin_key
from src.api.deps import require_admin
from src.api.schemas import IngestBzzoiroRequest, IngestResponse, IngestUnderstatRequest, IngestInjuriesRequest, IngestSimpleResponse
from src.data.config import BZZOIRO_LEAGUE_IDS, FDCO_TO_UNDERSTAT
from src.data.sources import get_source
from src.data.injuries import ingest_injuries
from src.db.unit_of_work import get_uow
@@ -15,46 +17,108 @@ logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/v1", tags=["ingest"])
# 后台采集任务注册表:持强引用防止被 GC
_background_tasks: set[asyncio.Task] = set()
@router.post("/ingest/bzzoiro", response_model=IngestResponse, dependencies=[Depends(require_admin_key)])
def _spawn(coro) -> None:
"""启动后台采集任务;异常已在任务内记录到系统日志。"""
task = asyncio.create_task(coro)
_background_tasks.add(task)
task.add_done_callback(_background_tasks.discard)
@router.post("/ingest/bzzoiro", dependencies=[Depends(require_admin)])
async def ingest_bzzoiro_route(req: IngestBzzoiroRequest):
"""触发 bzzoiro 采集。"""
source = get_source("bzzoiro")
# 未指定联赛 = 采集全部已知联赛;未指定状态 = 已完赛 + 未开赛都采集
leagues = req.leagues or list(BZZOIRO_LEAGUE_IDS.keys())
statuses = [req.status] if req.status else ["finished", "scheduled"]
_spawn(_run_bzzoiro(leagues, req.date_from, req.date_to, statuses))
return {
"ok": True,
"message": f"采集任务已启动(后台执行,状态: {', '.join(statuses)}),请在「系统日志」查看进度与结果",
}
async def _run_bzzoiro(leagues: list[str], date_from: str | None, date_to: str | None, statuses: list[str]) -> None:
"""后台执行 bzzoiro 采集:上游限速时单次可能耗时数分钟,必须脱离请求生命周期。"""
try:
source = get_source("bzzoiro")
merged: dict = {"leagues": {}, "total_inserted": 0, "total_updated": 0, "errors": []}
async with get_uow() as session:
result = await source.ingest(
session,
leagues=req.leagues,
date_from=req.date_from,
date_to=req.date_to,
status=req.status,
)
return IngestResponse(**result)
except Exception as e:
logger.exception("bzzoiro ingest failed")
raise HTTPException(500, "数据采集失败,请查看服务器日志")
for st in statuses:
r = await source.ingest(
session,
leagues=leagues,
date_from=date_from,
date_to=date_to,
status=st,
)
merged["total_inserted"] += r.get("total_inserted", 0)
merged["total_updated"] += r.get("total_updated", 0)
merged["errors"].extend(r.get("errors", []))
for code, stat in r.get("leagues", {}).items():
acc = merged["leagues"].setdefault(code, {"inserted": 0, "updated": 0, "errors": []})
acc["inserted"] += stat.get("inserted", 0)
acc["updated"] += stat.get("updated", 0)
acc["errors"].extend(stat.get("errors", []))
league_errors = {c: stat["errors"] for c, stat in merged["leagues"].items() if stat.get("errors")}
logger.info(
"bzzoiro 采集完成: 新增 %d, 更新 %d, 联赛 %d 个, 状态 %s",
merged["total_inserted"], merged["total_updated"], len(merged["leagues"]), statuses,
)
if league_errors:
sample = {c: errs[:1] for c, errs in list(league_errors.items())[:3]}
logger.warning("bzzoiro 部分联赛存在错误: %s", sample)
if merged["errors"]:
logger.warning("bzzoiro 采集错误 %d 条: %s", len(merged["errors"]), merged["errors"][:3])
except Exception:
logger.exception("bzzoiro 采集任务失败")
@router.post("/ingest/understat", response_model=IngestSimpleResponse, dependencies=[Depends(require_admin_key)])
@router.post("/ingest/understat", dependencies=[Depends(require_admin)])
async def ingest_understat_route(req: IngestUnderstatRequest):
"""触发 understat xG 回填。"""
source = get_source("understat")
leagues_to_run = [req.league] if req.league else list(FDCO_TO_UNDERSTAT.keys())
_spawn(_run_understat(leagues_to_run, req.season))
return {"ok": True, "message": "xG 回填任务已启动(后台执行),请在「系统日志」查看结果"}
async def _run_understat(leagues_to_run: list[str], season: int) -> None:
try:
source = get_source("understat")
merged: dict = {"count": 0, "updated": 0, "skipped": 0, "unmatched": 0, "errors": []}
async with get_uow() as session:
result = await source.ingest(session, league=req.league, season=req.season)
return IngestSimpleResponse(**result)
except Exception as e:
logger.exception("understat ingest failed")
raise HTTPException(500, "xG 回填失败,请查看服务器日志")
for league in leagues_to_run:
r = await source.ingest(session, league=league, season=season)
for k in ("count", "updated", "skipped", "unmatched"):
merged[k] += r.get(k, 0)
merged["errors"].extend(r.get("errors", []))
logger.info(
"understat 回填完成: 联赛 %d 个, 更新 %d, 未匹配 %d, 错误 %d",
len(leagues_to_run), merged["updated"], merged["unmatched"], len(merged["errors"]),
)
except Exception:
logger.exception("understat 回填任务失败")
@router.post("/ingest/injuries", response_model=IngestSimpleResponse, dependencies=[Depends(require_admin_key)])
@router.post("/ingest/injuries", dependencies=[Depends(require_admin)])
async def ingest_injuries_route(req: IngestInjuriesRequest):
"""触发伤停采集。"""
_spawn(_run_injuries(req.date))
return {"ok": True, "message": "伤停采集任务已启动(后台执行),请在「系统日志」查看结果"}
async def _run_injuries(date: str | None) -> None:
try:
async with get_uow() as session:
result = await ingest_injuries(session, date=req.date)
return IngestSimpleResponse(**result)
except Exception as e:
logger.exception("injuries ingest failed")
raise HTTPException(500, "伤停采集失败,请查看服务器日志")
result = await ingest_injuries(session, date=date)
logger.info(
"injuries 采集完成: 新增 %d, 更新 %d, 错误 %d",
result.get("count", 0), result.get("updated", 0), len(result.get("errors", [])),
)
if result.get("errors"):
logger.warning("injuries 采集错误: %s", result["errors"][:3])
except Exception:
logger.exception("injuries 采集任务失败")
+8 -2
View File
@@ -7,6 +7,7 @@ from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy import select
from sqlalchemy.orm import selectinload
from src.api.deps import require_admin
from src.api.schemas import MatchListOut, MatchOut
from src.db.base import AsyncSession, get_db_read
from src.db.models import League, Match
@@ -14,7 +15,7 @@ from src.db.models import League, Match
router = APIRouter(prefix="/api/v1", tags=["data"])
@router.get("/leagues", response_model=list[dict])
@router.get("/leagues", response_model=list[dict], dependencies=[Depends(require_admin)])
async def list_leagues(db: AsyncSession = Depends(get_db_read)):
stmt = select(League).order_by(League.name)
result = await db.execute(stmt)
@@ -64,7 +65,12 @@ async def list_matches(
raise HTTPException(400, "date 格式应为 YYYY-MM-DD")
q = q.where(Match.match_date >= d, Match.match_date < d + timedelta(days=1))
rows = (await db.execute(q.order_by(Match.match_date.desc(), Match.id.desc()).limit(limit + 1))).scalars().all()
# 未开赛按日期正序(最近的排最前,便于预测);其余按日期倒序(最新赛果在前)
if status == "scheduled":
order = (Match.match_date.asc(), Match.id.asc())
else:
order = (Match.match_date.desc(), Match.id.desc())
rows = (await db.execute(q.order_by(*order).limit(limit + 1))).scalars().all()
has_more = len(rows) > limit
rows = rows[:limit]
+19 -3
View File
@@ -7,9 +7,10 @@ from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy import select
from sqlalchemy.orm import selectinload
from src.api.deps import require_admin
from src.api.schemas import PredictOut, PredictRequest, PredictionOut
from src.db.base import AsyncSession, get_db, get_db_read
from src.db.models import Prediction
from src.db.models import Match, Prediction
from src.llm.predict import predict_match, PredictResult
logger = logging.getLogger(__name__)
@@ -20,6 +21,12 @@ router = APIRouter(prefix="/api/v1", tags=["predict"])
@router.post("/predict", response_model=PredictOut)
async def predict(req: PredictRequest, db: AsyncSession = Depends(get_db)):
"""对一场比赛调 LLM 预测。mode=multi(默认,5专家+终裁)或 single。"""
# 已完赛比赛不再支持预测(回测走服务层直调,不受此限)
match = await db.get(Match, req.match_id)
if match is None:
raise HTTPException(404, "match not found")
if match.match_status == "finished":
raise HTTPException(400, "该比赛已完赛,不再支持预测")
try:
result = await predict_match(
req.match_id,
@@ -28,6 +35,9 @@ async def predict(req: PredictRequest, db: AsyncSession = Depends(get_db)):
mode=req.mode,
)
except ValueError as e:
msg = str(e)
if "已结算" in msg:
raise HTTPException(409, msg)
logger.warning("predict validation error: %s", e)
raise HTTPException(404, "比赛不存在")
except RuntimeError as e:
@@ -46,6 +56,8 @@ async def predict(req: PredictRequest, db: AsyncSession = Depends(get_db)):
mode=getattr(result, "mode", "single"),
pred_home_goals=result.pred_home_goals,
pred_away_goals=result.pred_away_goals,
alt_pred_home_goals=result.alt_pred_home_goals,
alt_pred_away_goals=result.alt_pred_away_goals,
pred_1x2=result.pred_1x2,
subjective_confidence=result.subjective_confidence,
reasoning=result.reasoning,
@@ -54,9 +66,13 @@ async def predict(req: PredictRequest, db: AsyncSession = Depends(get_db)):
context=result.context,
latency_ms=result.latency_ms,
)
logger.info(
"预测完成 match=%s mode=%s pred=%s:%s (%s)",
req.match_id, req.mode, result.pred_home_goals, result.pred_away_goals, result.pred_1x2,
)
@router.get("/predictions", response_model=list[PredictionOut])
@router.get("/predictions", response_model=list[PredictionOut], dependencies=[Depends(require_admin)])
async def list_predictions(
match_id: int | None = None,
limit: int = Query(50, ge=1, le=200),
@@ -90,7 +106,7 @@ async def list_predictions(
]
@router.get("/predictions/{prediction_id}", response_model=PredictionOut)
@router.get("/predictions/{prediction_id}", response_model=PredictionOut, dependencies=[Depends(require_admin)])
async def get_prediction(prediction_id: int, db: AsyncSession = Depends(get_db_read)):
p = await db.get(Prediction, prediction_id)
if p is None:
+9 -5
View File
@@ -1,7 +1,7 @@
"""Pydantic schemas。"""
from __future__ import annotations
from datetime import datetime
from datetime import date, datetime
from typing import Any
from pydantic import BaseModel, Field
@@ -53,6 +53,8 @@ class PredictOut(BaseModel):
mode: str = "single"
pred_home_goals: float | None
pred_away_goals: float | None
alt_pred_home_goals: int | None = None
alt_pred_away_goals: int | None = None
pred_1x2: str | None
subjective_confidence: float | None
reasoning: str | None
@@ -71,6 +73,8 @@ class PredictionOut(BaseModel):
mode: str = "single"
pred_home_goals: float | None
pred_away_goals: float | None
alt_pred_home_goals: int | None = None
alt_pred_away_goals: int | None = None
pred_1x2: str | None
subjective_confidence: float | None
reasoning: str | None
@@ -82,10 +86,10 @@ class PredictionOut(BaseModel):
class IngestBzzoiroRequest(BaseModel):
leagues: list[str] = Field(..., description="联赛代码列表,如 ['E0','SP1']")
leagues: list[str] = Field(default_factory=list, description="联赛代码列表,如 ['E0','SP1'];空 = 全部已知联赛")
date_from: str | None = None
date_to: str | None = None
status: str = "finished"
status: str | None = Field(None, description="finished/scheduled;空 = 两者都采集")
class IngestResponse(BaseModel):
@@ -96,8 +100,8 @@ class IngestResponse(BaseModel):
class IngestUnderstatRequest(BaseModel):
league: str = Field(..., description="联赛代码,如 'E0'")
season: int = Field(..., description="赛季起始年,如 2025 表示 2025-2026 赛季")
league: str | None = Field(None, description="联赛代码,如 'E0';空 = 全部已知联赛")
season: int = Field(default_factory=lambda: date.today().year, description="赛季起始年,如 2025 表示 2025-2026 赛季")
class IngestInjuriesRequest(BaseModel):
+14 -3
View File
@@ -45,10 +45,21 @@ class Settings(BaseSettings):
DB_POOL_RECYCLE: int = 1800
# --- 管理接口鉴权 ---
# 采集 / 回测等高成本或写入型接口需要此 Key(请求头 X-API-Key)
# 留空表示「未启用鉴权」(本地开发默认),生产环境必须设置
# 见审查报告 P2-7:ingest/backtest 无鉴权可被任意调用并烧掉 LLM 额度
# 管理后台登录密码(POST /api/v1/auth/login),登录后颁发 HttpOnly Cookie 会话
# 采集 / 回测等高成本或写入型接口同样需要此密码或下方 API Key
# 两者均留空表示「未启用鉴权」(本地开发默认),生产环境必须至少设置一项
# 注意:.env 中的 ADMIN_PASSWORD 是初始值;后台修改密码后以数据库中的
# scrypt 哈希为准,建议随后删除此明文项。
ADMIN_PASSWORD: str = ""
ADMIN_API_KEY: str = ""
# 管理后台会话有效期(小时)
ADMIN_SESSION_TTL_HOURS: int = 168
# --- 加密主密钥 ---
# 敏感配置(数据源/LLM 的 API Key)入库加密、会话签名都由它派生。
# 只存于部署机 .env,切勿入库或提交代码。生成: openssl rand -base64 32
# 变更后已加密配置将无法解密(需在后台重新保存)。
SECRET_KEY: str = ""
settings = Settings()
+104
View File
@@ -0,0 +1,104 @@
"""安全原语:对称加密(Fernet/AES)与密码哈希(scrypt)。
- API Key 等需要原文调用的敏感值:入库前用 SECRET_KEY 派生的 Fernet 密钥加密,
存储格式 `enc:v1:<token>`;读取时解密。SECRET_KEY 只存于部署机 .env,不入库。
- 管理员密码:只存 scrypt 哈希(单向,不可逆),验证用,永远不需要还原原文。
`enc:v1:` 前缀 + 透传设计使旧明文数据无需停机即可共存,由启动迁移一次性加密。
"""
from __future__ import annotations
import base64
import hashlib
import hmac as _hmac
import logging
import secrets
from cryptography.fernet import Fernet, InvalidToken
from src.core.config import settings
logger = logging.getLogger(__name__)
_ENC_PREFIX = "enc:v1:"
# scrypt 参数(OWASP 推荐: n=2^17 更强,取 n=2^15 平衡 NAS CPU)
_SCRYPT_N = 2**15
_SCRYPT_R = 8
_SCRYPT_P = 1
# OpenSSL 默认 maxmem 限制约 32MB,显式放宽到 128MB
_SCRYPT_MAXMEM = 128 * 1024 * 1024
def _fernet() -> Fernet:
"""由 SECRET_KEY 确定性派生 Fernet 密钥(任意字符串输入均可)。
SECRET_KEY 未配置时回落派生自 DATABASE_URL(仅为不让开发环境崩溃;
生产必须显式配置,否则加密强度受限 —— 启动时会打 warning)。
"""
raw = settings.SECRET_KEY
if not raw:
logger.warning(
"SECRET_KEY 未设置,加密密钥回落派生自 DATABASE_URL。"
"请在 .env 配置强随机 SECRET_KEY(openssl rand -base64 32)。"
)
raw = f"fallback:{settings.DATABASE_URL}"
digest = hashlib.sha256(raw.encode()).digest()
return Fernet(base64.urlsafe_b64encode(digest))
def encrypt_value(plaintext: str) -> str:
"""加密敏感值,带版本前缀;空值原样返回。"""
if not plaintext:
return plaintext
token = _fernet().encrypt(plaintext.encode()).decode()
return f"{_ENC_PREFIX}{token}"
def decrypt_value(stored: str) -> str:
"""解密 `enc:v1:` 前缀的值;无前缀(旧明文)原样返回,便于平滑迁移。"""
if not stored or not stored.startswith(_ENC_PREFIX):
return stored
token = stored[len(_ENC_PREFIX):]
try:
return _fernet().decrypt(token.encode()).decode()
except InvalidToken:
# 密钥不匹配(通常是 SECRET_KEY 变了):报错而非静默返回错误数据
raise ValueError(
"敏感配置解密失败:SECRET_KEY 与加密时不一致。"
"恢复原 SECRET_KEY 或在后台重新保存对应配置项。"
) from None
def is_encrypted(stored: str) -> bool:
return bool(stored) and stored.startswith(_ENC_PREFIX)
def hash_password(password: str) -> str:
"""scrypt 哈希,存储格式 scrypt$N$r$p$salt_hex$dk_hex。"""
salt = secrets.token_bytes(16)
dk = hashlib.scrypt(
password.encode(), salt=salt, n=_SCRYPT_N, r=_SCRYPT_R, p=_SCRYPT_P,
dklen=32, maxmem=_SCRYPT_MAXMEM,
)
return f"scrypt${_SCRYPT_N}${_SCRYPT_R}${_SCRYPT_P}${salt.hex()}${dk.hex()}"
def verify_password(password: str, stored: str) -> bool:
"""校验密码与存储的 scrypt 哈希是否匹配。"""
try:
algo, n, r, p, salt_hex, dk_hex = stored.split("$")
if algo != "scrypt":
return False
dk = hashlib.scrypt(
password.encode(),
salt=bytes.fromhex(salt_hex),
n=int(n),
r=int(r),
p=int(p),
dklen=len(bytes.fromhex(dk_hex)),
maxmem=_SCRYPT_MAXMEM,
)
return _hmac.compare_digest(dk, bytes.fromhex(dk_hex))
except (ValueError, TypeError):
return False
+82
View File
@@ -0,0 +1,82 @@
"""内存日志缓冲:供后台「系统日志」页查看应用运行日志。
把应用日志(stdout)同时捕获到进程内环形缓冲(deque),提供级别/关键字/条数
过滤查询。缓冲在进程重启后清零;需要持久化的审计请另行落库。
"""
from __future__ import annotations
import logging
import threading
from collections import deque
_BUFFER: deque[dict] = deque(maxlen=2000)
_LOCK = threading.Lock()
_LEVEL_ORDER = {"DEBUG": 10, "INFO": 20, "WARNING": 30, "ERROR": 40, "CRITICAL": 50}
class MemoryLogHandler(logging.Handler):
"""把日志记录写入内存环形缓冲。"""
def __init__(self) -> None:
super().__init__()
# format() 会在有 exc_info 时自动附带异常堆栈文本
self.setFormatter(logging.Formatter("%(message)s"))
def emit(self, record: logging.LogRecord) -> None:
try:
entry = {
"ts": record.created,
"level": record.levelname,
"logger": record.name,
"message": self.format(record),
}
with _LOCK:
_BUFFER.append(entry)
except Exception: # noqa: BLE001 日志采集绝不影响业务
self.handleError(record)
class _SQLNoiseFilter(logging.Filter):
"""过滤 SQLAlchemy 的 DEBUG/INFO 回显(只留警告以上)。"""
def filter(self, record: logging.LogRecord) -> bool:
return not (
record.name.startswith("sqlalchemy.") and record.levelno < logging.WARNING
)
def get_entries(
min_level: str | None = None,
keyword: str | None = None,
limit: int = 200,
) -> list[dict]:
"""按条件查询缓冲日志,最新在前。"""
min_no = _LEVEL_ORDER.get((min_level or "").upper(), 0)
kw = (keyword or "").strip().lower()
with _LOCK:
items = list(_BUFFER)
items.reverse()
out: list[dict] = []
for e in items:
if _LEVEL_ORDER.get(e["level"], 0) < min_no:
continue
if kw and kw not in e["message"].lower() and kw not in e["logger"].lower():
continue
out.append(e)
if len(out) >= limit:
break
return out
def setup_memory_logging(level: str = "INFO") -> None:
"""挂载内存 handler 到 root logger(幂等),并确保 root 级别不低于 INFO。"""
root = logging.getLogger()
if any(isinstance(h, MemoryLogHandler) for h in root.handlers):
return
handler = MemoryLogHandler()
handler.setLevel(logging.INFO)
handler.addFilter(_SQLNoiseFilter())
root.addHandler(handler)
if root.level == logging.NOTSET or root.level > logging.INFO:
root.setLevel(getattr(logging, level.upper(), logging.INFO))
+259
View File
@@ -0,0 +1,259 @@
"""运行时配置:数据库优先,回落 .env。
后台「数据源」页可在线修改的配置项存 app_settings 表;
读取时 DB 有值用 DB,否则回落同名环境变量(pydantic settings)。
DB 读取失败时也回落环境变量,保证采集不因管理表故障而中断。
"""
from __future__ import annotations
import logging
import secrets
from dataclasses import dataclass
from sqlalchemy import select
from sqlalchemy.dialects.postgresql import insert as pg_insert
from src.core import crypto
from src.core.config import settings
from src.db.base import AsyncSessionLocal
from src.db.models import AppSetting
# 管理员密码哈希在 app_settings 中的键(不进 SETTING_DEFS 白名单:
# 只能走专门的改密接口 —— 需验证当前密码,不能被通用配置接口绕过)
ADMIN_PASSWORD_HASH_KEY = "ADMIN_PASSWORD_HASH"
# .env 明文密码的键名(仅作为初始值;后台改密后以哈希为准)
_ADMIN_ENV_KEY = "ADMIN_PASSWORD"
logger = logging.getLogger(__name__)
@dataclass(frozen=True)
class SettingDef:
key: str
label: str
description: str
sensitive: bool
# 允许在后台查看/修改的配置项白名单(之外的 key 一律拒绝读写)
SETTING_DEFS: dict[str, SettingDef] = {
"BZZOIRO_KEY": SettingDef(
"BZZOIRO_KEY", "Bzzoiro API Key", "比赛赛程 / 比分数据源凭证", sensitive=True,
),
"BZZOIRO_BASE": SettingDef(
"BZZOIRO_BASE", "Bzzoiro API 地址", "Bzzoiro 接口基础地址", sensitive=False,
),
"API_FOOTBALL_KEY": SettingDef(
"API_FOOTBALL_KEY", "API-Football Key", "伤停数据源凭证(api-sports)", sensitive=True,
),
"LLM_API_KEY": SettingDef(
"LLM_API_KEY", "LLM API Key", "大模型服务凭证(OpenAI 兼容接口)", sensitive=True,
),
"LLM_BASE_URL": SettingDef(
"LLM_BASE_URL", "LLM 接口地址", "如 https://api.deepseek.com/v1", sensitive=False,
),
"LLM_MODEL": SettingDef(
"LLM_MODEL", "LLM 模型", "如 deepseek-chat / gpt-4o", sensitive=False,
),
}
# ── 按角色独立配置 LLM 的键(5 专家 + 终裁) ──
# 每个角色可独立覆盖 模型 / 接口地址 / API Key;留空继承分层默认(见 orchestrator._agent_provider)。
AGENT_META: list[dict] = [
{"id": "form", "label": "近期状态分析专家"},
{"id": "stats", "label": "攻防数据分析专家"},
{"id": "home_away", "label": "主客因素分析专家"},
{"id": "injuries", "label": "阵容完整性分析专家"},
{"id": "h2h", "label": "历史交锋分析专家"},
{"id": "aggregator", "label": "终裁分析专家"},
]
for _agent in AGENT_META:
_u = _agent["id"].upper()
SETTING_DEFS[f"AGENT_{_u}_MODEL"] = SettingDef(
f"AGENT_{_u}_MODEL", f"{_agent['label']} 模型", "留空继承默认(专家层/全局)", sensitive=False,
)
SETTING_DEFS[f"AGENT_{_u}_BASE_URL"] = SettingDef(
f"AGENT_{_u}_BASE_URL", f"{_agent['label']} 接口地址", "留空继承全局 LLM_BASE_URL", sensitive=False,
)
SETTING_DEFS[f"AGENT_{_u}_API_KEY"] = SettingDef(
f"AGENT_{_u}_API_KEY", f"{_agent['label']} API Key", "留空继承全局 LLM_API_KEY", sensitive=True,
)
def mask_value(value: str, sensitive: bool) -> str:
"""脱敏展示:敏感值只留末 4 位;非敏感值原样返回。"""
if not value:
return ""
if not sensitive:
return value
return f"****{value[-4:]}" if len(value) >= 8 else "****"
async def get_runtime_value(key: str) -> str:
"""读运行时配置:DB 覆盖值 → .env 默认值 → 空串。
敏感项入库时是密文,读出后自动解密;旧明文(迁移前)由 decrypt_value 透传。
"""
defn = SETTING_DEFS.get(key)
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
except Exception:
logger.warning("读取运行时配置 %s 失败,回落环境变量", key)
return getattr(settings, key, "") or ""
async def set_runtime_value(key: str, value: str) -> None:
"""写入/更新 DB 覆盖值(调用方需先校验 key 在白名单内)。
敏感项(SETTING_DEFS.sensitive)以 Fernet 加密存储,库里不落明文。
"""
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)
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
except Exception:
logger.warning("读取运行时配置 %s 来源失败,按环境变量处理", key)
env_value = getattr(settings, key, "") or ""
return ("env", env_value) if env_value else ("none", "")
async def migrate_plaintext_sensitive_settings() -> int:
"""一次性迁移:把库中仍是明文的敏感项加密(幂等,启动时执行)。
返回加密的条数。
"""
migrated = 0
sensitive_keys = {k for k, d in SETTING_DEFS.items() if d.sensitive}
async with AsyncSessionLocal() as db:
rows = (await db.execute(select(AppSetting))).scalars().all()
for row in rows:
if row.key not in sensitive_keys or crypto.is_encrypted(row.value):
continue
row.value = crypto.encrypt_value(row.value)
migrated += 1
await db.commit()
if migrated:
logger.info("已加密迁移 %d 条明文敏感配置", migrated)
return migrated
# ── 管理员密码:只存 scrypt 哈希,永不存明文 ──────────────────────
async def get_admin_password_hash() -> str:
"""库中管理员密码哈希;无则空串。"""
try:
async with AsyncSessionLocal() as db:
row = await db.get(AppSetting, ADMIN_PASSWORD_HASH_KEY)
return row.value if row else ""
except Exception:
logger.warning("读取管理员密码哈希失败")
return ""
async def set_admin_password_hash(hash_str: str) -> None:
async with AsyncSessionLocal() as db:
stmt = pg_insert(AppSetting).values(key=ADMIN_PASSWORD_HASH_KEY, value=hash_str)
stmt = stmt.on_conflict_do_update(index_elements=["key"], set_={"value": hash_str})
await db.execute(stmt)
await db.commit()
logger.info("管理员密码哈希已更新")
async def verify_admin_password(candidate: str) -> bool:
"""校验管理员密码:优先哈希;哈希不存在时回落 .env 明文(未迁移的旧部署)。"""
stored_hash = await get_admin_password_hash()
if stored_hash:
return crypto.verify_password(candidate, stored_hash)
env_pw = getattr(settings, _ADMIN_ENV_KEY, "") or ""
return bool(env_pw) and secrets.compare_digest(candidate, env_pw)
async def get_admin_credential_fingerprint() -> str:
"""管理员凭证指纹(作为会话签名密钥的输入)。
用密码哈希而非密码本身:凭证变化 → 指纹变化 → 全部会话失效。
"""
stored_hash = await get_admin_password_hash()
if stored_hash:
return f"hash:{stored_hash}"
env_pw = getattr(settings, _ADMIN_ENV_KEY, "") or ""
return f"env:{env_pw}" if env_pw else ""
async def _get_raw_setting(key: str) -> str:
async with AsyncSessionLocal() as db:
row = await db.get(AppSetting, key)
return row.value if row else ""
async def _delete_setting(key: str) -> None:
async with AsyncSessionLocal() as db:
row = await db.get(AppSetting, key)
if row is not None:
await db.delete(row)
await db.commit()
async def ensure_admin_password_hashed() -> bool:
"""启动迁移:确保管理员密码只以 scrypt 哈希存在(幂等)。
迁移来源优先级:
1. 库中旧版明文 ADMIN_PASSWORD 行(旧代码写入的当前密码,迁移后删除该明文行)
2. .env 的 ADMIN_PASSWORD 初始值
"""
if await get_admin_password_hash():
# 哈希已存在:清除旧版可能残留的明文行
if await _get_raw_setting(_ADMIN_ENV_KEY):
await _delete_setting(_ADMIN_ENV_KEY)
logger.info("已删除遗留的明文 ADMIN_PASSWORD 行(哈希已存在)")
return False
legacy_plain = await _get_raw_setting(_ADMIN_ENV_KEY)
if legacy_plain:
await set_admin_password_hash(crypto.hash_password(legacy_plain))
await _delete_setting(_ADMIN_ENV_KEY)
logger.info("已将库中明文管理员密码迁移为 scrypt 哈希,明文行已删除")
return True
env_pw = getattr(settings, _ADMIN_ENV_KEY, "") or ""
if not env_pw:
return False
await set_admin_password_hash(crypto.hash_password(env_pw))
logger.info(
"已将 .env 中的明文管理员密码迁移为 scrypt 哈希。"
"建议现在从 .env 中删除 ADMIN_PASSWORD 明文行。"
)
return True
+20 -9
View File
@@ -14,10 +14,13 @@ from datetime import datetime, timezone
from sqlalchemy import select
from src.core.config import settings
import httpx
from src.core.runtime_config import get_runtime_value
from src.core.http_client import get_client
from src.data.config import BZZOIRO_LEAGUE_IDS, LEAGUE_COUNTRIES, LEAGUE_NAMES, REQUEST_INTERVAL
from src.data.normalize import normalize_bzzoiro
from src.data.team_names_zh import zh_name
from src.data.sources import register
from src.db.models import League, Match, MatchStats, Team
@@ -46,9 +49,9 @@ def _match_key(home_team_id: int, away_team_id: int, match_date) -> tuple[int, i
async def _fetch_json_async(path: str, params: dict | None = None, max_retries: int = 3) -> dict | list:
"""异步 HTTP(bzzoiro 使用 httpx,不再阻塞事件循环线程池)。"""
base = settings.BZZOIRO_BASE.rstrip("/")
base = (await get_runtime_value("BZZOIRO_BASE")).rstrip("/")
url = f"{base}/{path.lstrip('/')}"
key = settings.BZZOIRO_KEY
key = await get_runtime_value("BZZOIRO_KEY")
if not key:
raise RuntimeError("BZZOIRO_KEY 未设置")
@@ -61,7 +64,14 @@ async def _fetch_json_async(path: str, params: dict | None = None, max_retries:
for attempt in range(max_retries):
try:
client = get_client()
resp = await client.get(url, headers=headers, params=params, timeout=30)
# 整请求兜底: httpx 无 total 超时,用 wait_for 防「滴水式」限速挂死
resp = await asyncio.wait_for(
client.get(
url, headers=headers, params=params,
timeout=httpx.Timeout(connect=10.0, read=30.0, write=10.0, pool=10.0),
),
timeout=60.0,
)
resp.raise_for_status()
return resp.json()
except Exception as e:
@@ -199,7 +209,8 @@ class BzzoiroSource:
# 避免加载联赛全部历史比赛到内存(多赛季采集时内存溢出)
if normalized_matches:
from datetime import timedelta
dates = [nm.date for nm in normalized_matches if nm.date is not None]
# normalized_matches 存的是 (nm, raw) 元组,遍历需解包
dates = [nm.date for nm, _raw in normalized_matches if nm.date is not None]
if dates:
min_dt = min(dates) - timedelta(days=30)
max_dt = max(dates) + timedelta(days=30)
@@ -219,7 +230,7 @@ class BzzoiroSource:
# 球队: 内存查找 + 按需创建
home_team_id = team_name_to_id.get(nm.home_team)
if home_team_id is None:
home = Team(name=nm.home_team)
home = Team(name=nm.home_team, name_zh=zh_name(nm.home_team))
db.add(home)
await db.flush()
home_team_id = home.id
@@ -227,7 +238,7 @@ class BzzoiroSource:
away_team_id = team_name_to_id.get(nm.away_team)
if away_team_id is None:
away = Team(name=nm.away_team)
away = Team(name=nm.away_team, name_zh=zh_name(nm.away_team))
db.add(away)
await db.flush()
away_team_id = away.id
@@ -273,7 +284,7 @@ class BzzoiroSource:
home_red_cards=nm.home_red_cards,
away_red_cards=nm.away_red_cards,
source="bzzoiro",
source_event_id=str(raw.get("id", "")),
source_record_id=str(raw.get("id", "")),
retrieved_at=now,
available_at=now,
)
@@ -299,7 +310,7 @@ class BzzoiroSource:
existing_match.stats = MatchStats(
match_id=existing_match.id,
source="bzzoiro",
source_event_id=str(raw.get("id", "")),
source_record_id=str(raw.get("id", "")),
retrieved_at=now,
available_at=now,
)
+1 -1
View File
@@ -43,4 +43,4 @@ LEAGUE_COUNTRIES: dict[str, str] = {
"EL": "Europe",
}
REQUEST_INTERVAL = 1.2 # bzzoiro 限速(秒)
REQUEST_INTERVAL = 2.0 # bzzoiro 限速(秒);上游限速严厉时宁可慢一点
+9 -3
View File
@@ -16,7 +16,7 @@ from typing import Any
import httpx
from src.core.config import settings
from src.core.runtime_config import get_runtime_value
from src.core.http_client import get_client
logger = logging.getLogger(__name__)
@@ -39,7 +39,7 @@ async def fetch_injuries(*, date: str | None = None, fixture_id: int | None = No
Returns:
伤停记录列表
"""
api_key = settings.API_FOOTBALL_KEY
api_key = await get_runtime_value("API_FOOTBALL_KEY")
if not api_key:
raise RuntimeError("API_FOOTBALL_KEY 未设置")
@@ -77,7 +77,13 @@ async def fetch_injuries(*, date: str | None = None, fixture_id: int | None = No
for attempt in range(3):
try:
client = get_client()
resp = await client.get(url, headers=headers, params=params, timeout=30)
resp = await asyncio.wait_for(
client.get(
url, headers=headers, params=params,
timeout=httpx.Timeout(connect=10.0, read=30.0, write=10.0, pool=10.0),
),
timeout=60.0,
)
resp.raise_for_status()
break
except Exception as e:
+3 -3
View File
@@ -18,7 +18,7 @@ VALID_STATUS = {"finished", "scheduled", "in_play", "paused", "postponed", "canc
STATUS_MAP = {
"finished": "finished", "completed": "finished", "done": "finished", "awarded": "finished",
"scheduled": "scheduled", "upcoming": "scheduled",
"scheduled": "scheduled", "upcoming": "scheduled", "notstarted": "scheduled", "not_started": "scheduled",
"in_play": "in_play", "live": "in_play",
"paused": "paused", "postponed": "postponed",
"cancelled": "cancelled", "canceled": "cancelled", "abandoned": "cancelled",
@@ -200,8 +200,8 @@ def normalize_understat(raw: dict, league_type: str) -> NormalizedMatch | None:
away = normalize_name(away_name)
if not home or not away or home == away:
return None
home_xg = raw.get("xG", {}).get("h") if isinstance(raw.get("xG"), dict) else None
away_xg = raw.get("xG", {}).get("a") if isinstance(raw.get("xG"), dict) else None
home_xg = _to_float(raw["xG"].get("h")) if isinstance(raw.get("xG"), dict) else None
away_xg = _to_float(raw["xG"].get("a")) if isinstance(raw.get("xG"), dict) else None
return NormalizedMatch(
league_type=league_type,
date=dt,
+7 -3
View File
@@ -29,9 +29,13 @@ class DataSource(Protocol):
_SOURCES: dict[str, DataSource] = {}
def register(source: DataSource) -> DataSource:
"""装饰器:将数据源注册到全局注册表。"""
_SOURCES[source.name] = source
def register(source):
"""装饰器:将数据源注册到全局注册表。
兼容类注册与实例注册:类会被实例化后存入(保证 get_source 返回实例)。
"""
obj = source() if isinstance(source, type) else source
_SOURCES[obj.name] = obj
return source
+168
View File
@@ -0,0 +1,168 @@
"""球队中文译名表: 规范化英文名 → 中文。
来源: 手工整理(五大联赛全部 + 欧战常客)。
未收录的球队保持英文显示(前端回落),新队采集入库时自动查此表。
"""
TEAM_NAME_ZH: dict[str, str] = {
# ── 英格兰 ──
"Arsenal": "阿森纳", "Aston Villa": "阿斯顿维拉", "Chelsea": "切尔西",
"Liverpool FC": "利物浦", "Liverpool": "利物浦",
"Manchester City": "曼城", "Manchester United": "曼联",
"Tottenham Hotspur": "托特纳姆热刺", "Newcastle United": "纽卡斯尔联",
"West Ham United": "西汉姆联", "Everton": "埃弗顿", "Fulham": "富勒姆",
"Crystal Palace": "水晶宫", "Brentford": "布伦特福德",
"Brighton & Hove Albion": "布莱顿", "Brighton and Hove Albion": "布莱顿",
"Wolverhampton": "狼队", "Wolverhampton Wanderers": "狼队",
"Nottingham Forest": "诺丁汉森林", "AFC Bournemouth": "伯恩茅斯",
"Leeds United": "利兹联", "Leicester City": "莱斯特城",
"Ipswich Town": "伊普斯维奇", "Southampton": "南安普顿",
"Norwich City": "诺维奇城", "Sheffield United": "谢菲尔德联",
"Sheffield Wednesday": "谢周三", "Stoke City": "斯托克城",
"Sunderland": "桑德兰", "Burnley": "伯恩利", "Watford": "沃特福德",
"Hull City": "赫尔城", "Huddersfield Town": "哈德斯菲尔德",
"Luton Town": "卢顿", "Cardiff City": "加的夫城", "Swansea City": "斯旺西",
"West Bromwich Albion": "西布罗姆维奇", "Birmingham City": "伯明翰",
"Blackburn Rovers": "布莱克本", "Bolton Wanderers": "博尔顿",
"Barnsley": "巴恩斯利", "Blackpool": "布莱克浦",
"Bradford City": "布拉德福德", "Charlton Athletic": "查尔顿竞技",
"Coventry City": "考文垂", "Derby County": "德比郡",
"Middlesbrough": "米德尔斯堡", "Milton Keynes Dons": "米尔顿凯恩斯",
"Oldham Athletic": "奥尔德姆竞技", "Portsmouth": "朴茨茅斯",
"Queens Park Rangers": "女王公园巡游者", "Reading": "雷丁",
"Swindon Town": "斯温登", "Wigan Athletic": "维冈竞技",
# ── 苏格兰/爱尔兰 ──
"Celtic": "凯尔特人", "Rangers": "流浪者", "Aberdeen": "阿伯丁",
"Heart of Midlothian": "哈茨", "Hibernian": "希伯尼安",
"Derry City": "德里城", "Shelbourne": "谢尔本", "Larne FC": "拉恩",
"Linfield FC": "林斯菲尔德", "Shamrock Rovers": "沙姆罗克流浪者",
# ── 西班牙 ──
"Real Madrid": "皇家马德里", "FC Barcelona": "巴塞罗那",
"Atlético Madrid": "马德里竞技", "Athletic Club": "毕尔巴鄂竞技",
"Real Sociedad": "皇家社会", "Villarreal": "比利亚雷亚尔",
"Real Betis": "皇家贝蒂斯", "Sevilla": "塞维利亚", "Valencia": "瓦伦西亚",
"Celta Vigo": "塞尔塔", "Osasuna": "奥萨苏纳", "Getafe": "赫塔菲",
"Rayo Vallecano": "巴列卡诺", "Mallorca": "马略卡", "Girona FC": "赫罗纳",
"Girona": "赫罗纳", "Espanyol": "西班牙人", "UD Las Palmas": "拉斯帕尔马斯",
"Las Palmas": "拉斯帕尔马斯", "Deportivo Alavés": "阿拉维斯",
"Leganés": "莱加内斯", "Elche": "埃尔切", "Levante UD": "莱万特",
"Malaga CF": "马拉加", "Deportivo de A Coruna": "拉科鲁尼亚",
"Real Oviedo": "皇家奥维耶多", "Real Racing Club": "桑坦德竞技",
"Real Valladolid": "巴利亚多利德",
# ── 意大利 ──
"Juventus": "尤文图斯", "AC Milan": "AC米兰", "Inter Milan": "国际米兰",
"SSC Napoli": "那不勒斯", "AS Roma": "罗马", "Lazio": "拉齐奥",
"Atalanta": "亚特兰大", "ACF Fiorentina": "佛罗伦萨", "Bologna": "博洛尼亚",
"Torino": "都灵", "Udinese": "乌迪内斯", "Genoa": "热那亚",
"Cagliari": "卡利亚里", "Hellas Verona": "维罗纳", "Lecce": "莱切",
"Empoli": "恩波利", "Parma": "帕尔马", "Como": "科莫", "Venezia": "威尼斯",
"Pisa": "比萨", "Cremonese": "克雷莫纳", "AC Monza": "蒙扎",
"Frosinone": "弗罗西诺内", "Sassuolo": "萨索洛",
# ── 德国 ──
"FC Bayern Munchen": "拜仁慕尼黑", "Borussia Dortmund": "多特蒙德",
"Bayer 04 Leverkusen": "勒沃库森", "RB Leipzig": "莱比锡红牛",
"Borussia Mönchengladbach": "门兴格拉德巴赫", "VfB Stuttgart": "斯图加特",
"Eintracht Frankfurt": "法兰克福", "VfL Wolfsburg": "沃尔夫斯堡",
"SC Freiburg": "弗赖堡", "TSG Hoffenheim": "霍芬海姆",
"1. FC Union Berlin": "柏林联合", "1. FC Koln": "科隆",
"1. FSV Mainz 05": "美因茨", "FC Augsburg": "奥格斯堡",
"SV Werder Bremen": "云达不来梅", "VfL Bochum 1848": "波鸿",
"1. FC Heidenheim": "海登海姆", "FC St. Pauli": "圣保利",
"Holstein Kiel": "荷尔斯泰因基尔", "FC Schalke 04": "沙尔克04",
"Hamburger SV": "汉堡", "SC Paderborn 07": "帕德博恩",
"SV 07 Elversberg": "埃弗斯贝格",
# ── 法国 ──
"Paris Saint-Germain": "巴黎圣日耳曼", "Olympique de Marseille": "马赛",
"Olympique Lyonnais": "里昂", "AS Monaco": "摩纳哥", "Lille OSC": "里尔",
"OGC Nice": "尼斯", "RC Lens": "朗斯", "Stade Rennais": "雷恩",
"RC Strasbourg": "斯特拉斯堡", "Stade Brestois": "布雷斯特",
"Stade de Reims": "兰斯", "FC Nantes": "南特", "Toulouse FC": "图卢兹",
"Montpellier HSC": "蒙彼利埃", "AS Saint-Étienne": "圣埃蒂安",
"AJ Auxerre": "欧塞尔", "Le Havre AC": "勒阿弗尔", "FC Lorient": "洛里昂",
"Metz": "梅斯", "Angers SCO": "昂热", "Guingamp": "甘冈", "Troyes": "特鲁瓦",
"Le Mans": "勒芒", "Paris FC": "巴黎FC", "USL Dunkerque": "敦刻尔克",
"Rodez AF": "罗德兹", "Red Star FC": "巴黎红星",
# ── 荷兰/比利时 ──
"AFC Ajax": "阿贾克斯", "PSV Eindhoven": "埃因霍温", "Feyenoord": "费耶诺德",
"AZ Alkmaar": "阿尔克马尔", "FC Twente": "特温特", "FC Utrecht": "乌得勒支",
"NEC Nijmegen": "奈梅亨", "Go Ahead Eagles": "前进之鹰",
"Club Brugge KV": "布鲁日", "RSC Anderlecht": "安德莱赫特",
"KRC Genk": "亨克", "Royale Union Saint-Gilloise": "圣吉罗斯联合",
"Sint-Truidense VV": "圣特鲁伊登",
# ── 葡萄牙 ──
"FC Porto": "波尔图", "Benfica": "本菲卡", "Sporting CP": "里斯本竞技",
"Sporting Braga": "布拉加", "Torreense": "托雷恩塞",
# ── 土耳其 ──
"Galatasaray": "加拉塔萨雷", "Fenerbahce": "费内巴切",
"Besiktas JK": "贝西克塔斯", "Trabzonspor": "特拉布宗体育",
"Samsunspor": "萨姆松体育",
# ── 北欧 ──
"Bodø/Glimt": "博多闪耀", "Viking FK": "维京", "Tromsø IL": "特罗姆瑟",
"SK Brann": "布兰", "Lillestrøm SK": "利勒斯特罗姆",
"Malmo FF": "马尔默", "IF Elfsborg": "埃尔夫斯堡", "BK Hacken": "哈肯",
"Hammarby IF": "哈马比", "Fredrikstad FK": "腓特烈斯塔",
"Mjallby AIF": "米亚尔比", "AGF": "奥胡斯", "FC Midtjylland": "中日德兰",
"FC København": "哥本哈根", "Klaksvikar Itrottarfelag": "克拉克斯维克",
"Vikingur Gøta": "戈塔维京人", "Vikingur Reykjavik": "雷克雅未克维京人",
"Breidablik Kopavogur": "布雷达布利克", "IF Vestri": "韦斯特里",
"Kuopion Palloseura": "库奥皮奥", "Ilves": "伊尔维斯",
# ── 瑞士/奥地利 ──
"Basel": "巴塞尔", "BSC Young Boys": "伯尔尼年轻人", "FC Lugano": "卢加诺",
"Servette FC": "塞尔维特", "FC Thun": "图恩",
"FC St. Gallen 1879": "圣加仑", "LASK": "林茨", "SK Sturm Graz": "格拉茨风暴",
"Red Bull Salzburg": "萨尔茨堡红牛", "Wolfsberger AC": "沃尔夫斯贝格",
# ── 中东欧 ──
"Shakhtar Donetsk": "顿涅茨克矿工", "Dynamo Kyiv": "基辅迪纳摩",
"Dinamo Minsk": "明斯克迪纳摩", "ML Vitebsk": "维捷布斯克",
"Legia Warszawa": "华沙莱吉亚", "Lech Poznan": "波兹南莱赫",
"Jagiellonia Białystok": "比亚韦斯托克亚盖隆尼亚",
"Gornik Zabrze": "扎布热矿工", "MSK Zilina": "日利纳",
"SK Slovan Bratislava": "布拉迪斯拉发斯洛万",
"FC Spartak Trnava": "特尔纳瓦斯巴达", "SK Slavia Praha": "布拉格斯拉维亚",
"AC Sparta Praha": "布拉格斯巴达", "FC Viktoria Plzen": "比尔森胜利",
"SK Sigma Olomouc": "奥洛莫茨西格玛", "FC Hradec Kralove": "赫拉德茨克拉洛韦",
"Banik Ostrava": "俄斯特拉发矿工", "Ferencvaros TC": "费伦茨瓦罗斯",
"Paksi FC": "帕克斯", "ETO FC Gyor": "杰尔", "CFR 1907 Cluj": "克卢日",
"FCSB": "布加勒斯特星", "FC Universitatea Cluj": "克卢日大学",
"Universitatea Craiova": "克拉约瓦大学", "Ludogorets": "卢多戈雷茨",
"Levski Sofia": "索非亚列夫斯基", "CSKA Sofia": "索非亚中央陆军",
"GNK Dinamo Zagreb": "萨格勒布迪纳摩", "HNK Hajduk Split": "斯普利特海杜克",
"HNK Rijeka": "里耶卡", "NK Olimpija Ljubljana": "卢布尔雅那奥林匹亚",
"NK Celje": "采列", "NK Aluminij Kidricevo": "阿卢米尼",
"FK Partizan": "贝尔格莱德游击队", "FK Vojvodina": "伏伊伏丁那",
"FK Crvena Zvezda": "贝尔格莱德红星", "FK Crvena zvezda": "贝尔格莱德红星",
"FK Borac Banja Luka": "巴尼亚卢卡战士", "HSK Zrinjski Mostar": "莫斯塔尔兹林斯基",
"FK Buducnost Podgorica": "波德戈里察未来",
"FK Sutjeska Niksic": "苏捷斯卡", "Sheriff Tiraspol": "谢里夫",
"FC Petrocub Hincesti": "佩特罗库布", "FC Milsami Orhei": "米尔萨米",
# ── 希腊/塞浦路斯/以色列 ──
"Olympiacos FC": "奥林匹亚科斯", "Panathinaikos FC": "帕纳辛纳科斯",
"PAOK": "塞萨洛尼基PAOK", "AEK Athens": "雅典AEK", "OFI Crete": "克里特OFI",
"Omonia Nicosia": "尼科西亚奥莫尼亚", "AEK Larnaca": "拉纳卡AEK",
"Pafos FC": "帕福斯", "Hapoel Be'er Sheva": "贝尔谢巴夏普尔",
"Maccabi Tel Aviv": "特拉维夫马卡比",
# ── 东南欧/高加索/中亚 ──
"Qarabag FK": "卡拉巴赫", "Sabah FK": "萨巴赫",
"FC Ararat-Armenia": "亚美尼亚阿拉拉特", "FC Noah": "诺亚",
"FK Aktobe": "阿克托别", "Kairat Almaty": "阿拉木图凯拉特",
"FC Kairat Almaty": "阿拉木图凯拉特",
"FK Vardar Skopje": "瓦尔达尔斯科普里", "KF Shkendija": "什肯迪贾",
"KF Egnatia": "埃格纳蒂亚", "FK Zalgiris": "萨尔吉里斯",
"FK Kauno Zalgiris": "考那斯萨尔吉里斯", "Riga FC": "里加",
"RFS": "里加足球学校", "FCI Levadia Tallinn": "塔林列瓦迪亚",
"Flora Tallinn": "塔林弗洛拉",
# ── 小联赛/外围 ──
"The New Saints": "新圣徒", "Lincoln Red Imps": "林肯红魔",
"Ħamrun Spartans FC": "哈姆伦斯巴达", "Floriana FC": "弗洛里亚纳",
"SP Tre Fiori": "特雷菲奥里", "SS Virtus": "维尔图斯",
"Inter Club d'Escaldes": "埃斯卡尔德斯", "Differdange FC 03": "迪费尔当",
"Atert Bissen": "比森", "FC Drita": "德里塔", "FC Prishtina": "普里什蒂纳",
"FC Iberia 1999": "伊比利亚1999", "FK Buducnost": "波德戈里察未来",
}
def zh_name(name: str | None) -> str | None:
"""查中文译名;未收录返回 None(由调用方回落英文)。"""
if not name:
return None
return TEAM_NAME_ZH.get(name.strip())
+17 -4
View File
@@ -52,7 +52,13 @@ async def fetch_understat(league_code: str, season: int) -> list[dict]:
for attempt in range(3):
try:
client = get_client()
resp = await client.get(url, headers=headers, timeout=30)
resp = await asyncio.wait_for(
client.get(
url, headers=headers,
timeout=httpx.Timeout(connect=10.0, read=30.0, write=10.0, pool=10.0),
),
timeout=60.0,
)
resp.raise_for_status()
break
except Exception as e:
@@ -65,9 +71,16 @@ async def fetch_understat(league_code: str, season: int) -> list[dict]:
else:
raise RuntimeError(f"understat fetch failed: {last_exc}")
# understat 返回 JS 对象,需要提取 JSON
# 优先按 JSON 响应解析(getLeagueData 接口返回 {teams, players, dates})
try:
data = resp.json()
except Exception:
data = None
if isinstance(data, dict) and isinstance(data.get("dates"), list):
return data["dates"]
# 兼容旧版联赛页面:内嵌 var datesData = JSON.parse('...')
text = resp.text
# 匹配 var datesData = JSON.parse('...');
match = re.search(r"var\s+datesData\s*=\s*JSON\.parse\('([^']+)'\)", text)
if not match:
logger.warning("understat 响应格式不符: %s...", text[:200])
@@ -187,7 +200,7 @@ class UnderstatSource:
existing.stats = MatchStats(
match_id=existing.id,
source="understat",
source_event_id=str(raw.get("id", "")),
source_record_id=str(raw.get("id", "")),
retrieved_at=now,
available_at=now,
)
+12
View File
@@ -177,6 +177,9 @@ class Prediction(Base):
latency_ms: Mapped[int | None] = mapped_column(Integer)
pred_home_goals: Mapped[float | None] = mapped_column(Float)
pred_away_goals: Mapped[float | None] = mapped_column(Float)
# 备选比分(次可能比分,可空)
alt_pred_home_goals: Mapped[int | None] = mapped_column(Integer)
alt_pred_away_goals: Mapped[int | None] = mapped_column(Integer)
pred_1x2: Mapped[str | None] = mapped_column(String(3))
subjective_confidence: Mapped[float | None] = mapped_column(Float) # LLM 主观置信度,非概率
reasoning: Mapped[str | None] = mapped_column(Text)
@@ -216,3 +219,12 @@ class Prediction(Base):
CheckConstraint("mode IN ('single', 'multi')", name="ck_mode_enum"),
CheckConstraint("status IN ('success', 'failed', 'degraded')", name="ck_status_enum"),
)
class AppSetting(Base):
"""后台管理的运行时设置(如数据源 API Key),读取时优先于 .env 默认值。"""
__tablename__ = "app_settings"
key: Mapped[str] = mapped_column(String(100), primary_key=True)
value: Mapped[str] = mapped_column(Text, nullable=False)
updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), default=_utcnow, onupdate=_utcnow)
+1 -1
View File
@@ -183,7 +183,7 @@ async def run_agent(
user=user_prompt,
json_mode=True,
temperature=0.2,
max_tokens=600,
max_tokens=4096, # 推理模型需要更大余量
)
if resp.error:
logger.warning("agent %s LLM failed: %s", spec.name, resp.error)
+75 -41
View File
@@ -13,6 +13,7 @@ from src.core.config import settings
from src.db.base import AsyncSessionLocal
from src.db.models import Match, Prediction
from src.db.unit_of_work import get_uow
from src.llm.predict import _upsert_prediction
from src.llm.agents.base import AgentReport, AgentSpec, load_agent_prompt
from src.llm.context_builder import (
MatchHeader,
@@ -24,6 +25,7 @@ from src.llm.context_builder import (
load_match_header,
stats_slice,
)
from src.core.runtime_config import get_runtime_value
from src.llm.provider import LLMProvider, get_default_provider
logger = logging.getLogger(__name__)
@@ -58,7 +60,20 @@ SPECIALIST_SPECS: list[AgentSpec] = [
),
]
AGGREGATOR_SYSTEM = "你是足球预测终裁专家。综合各领域报告输出最终预测。只输出 JSON。"
AGGREGATOR_SYSTEM = (
"你是足球预测终裁专家。综合各领域专家报告输出最终预测。"
"引用专家时必须使用报告中的专家全名(如「攻防数据分析专家」),禁止使用英文代码。"
"只输出 JSON。"
)
# 专家代码 → 终裁/展示统一称呼
AGENT_LABELS_ZH: dict[str, str] = {
"form": "近期状态分析专家",
"stats": "攻防数据分析专家",
"home_away": "主客因素分析专家",
"injuries": "阵容完整性分析专家",
"h2h": "历史交锋分析专家",
}
@dataclass
@@ -70,6 +85,8 @@ class MultiPredictResult:
mode: str
pred_home_goals: float | None
pred_away_goals: float | None
alt_pred_home_goals: int | None
alt_pred_away_goals: int | None
pred_1x2: str | None
subjective_confidence: float | None
reasoning: str | None
@@ -80,31 +97,38 @@ class MultiPredictResult:
raw: dict | None
def _get_specialist_provider() -> LLMProvider:
"""专家模型: LLM_SPECIALIST_MODEL 回落 LLM_MODEL。"""
p = get_default_provider()
if settings.LLM_SPECIALIST_MODEL:
p.model = settings.LLM_SPECIALIST_MODEL
return p
async def _agent_provider(agent_id: str, *, tier: str) -> LLMProvider:
"""构造某 agent 专属 provider。
def _get_aggregator_provider() -> LLMProvider:
"""终裁模型: LLM_AGGREGATOR_MODEL 回落 LLM_MODEL。"""
p = get_default_provider()
if settings.LLM_AGGREGATOR_MODEL:
p.model = settings.LLM_AGGREGATOR_MODEL
覆盖优先级:
模型: AGENT_MODEL_{ID}(运行时) → 层级默认(LLM_SPECIALIST/AGGREGATOR_MODEL) → 全局 LLM_MODEL
地址/密钥: AGENT_BASE_URL_{ID} / AGENT_API_KEY_{ID}(运行时) → 全局 LLM_BASE_URL / LLM_API_KEY
"""
pfx = f"AGENT_{agent_id.upper()}_"
p = await get_default_provider()
tier_model = settings.LLM_SPECIALIST_MODEL if tier == "specialist" else settings.LLM_AGGREGATOR_MODEL
if tier_model:
p.model = tier_model
model = await get_runtime_value(f"{pfx}MODEL")
if model:
p.model = model
base = await get_runtime_value(f"{pfx}BASE_URL")
if base:
p.base_url = base
key = await get_runtime_value(f"{pfx}API_KEY")
if key:
p.api_key = key
return p
async def run_specialists(
header: MatchHeader,
*,
provider: LLMProvider,
version: str = "v1",
) -> list[AgentReport]:
"""并行执行 5 个专家 agent。fail-open: 单个失败不影响其他。"""
tasks = [
_run_one(spec, header, provider, version=version)
_run_one(spec, header, await _agent_provider(spec.name, tier="specialist"), version=version)
for spec in SPECIALIST_SPECS
]
results = await asyncio.gather(*tasks, return_exceptions=True)
@@ -125,7 +149,13 @@ async def _run_one(spec, header, provider, *, version) -> AgentReport:
def _reports_to_json(reports: list[AgentReport]) -> str:
return json.dumps([r.to_dict() for r in reports], ensure_ascii=False, indent=1)
"""报告序列化: agent 字段直接用中文专家全名,引导终裁用统一称呼引用。"""
out = []
for r in reports:
d = r.to_dict()
d["agent"] = AGENT_LABELS_ZH.get(d.get("agent", ""), d.get("agent"))
out.append(d)
return json.dumps(out, ensure_ascii=False, indent=1)
async def run_aggregator(
@@ -147,7 +177,7 @@ async def run_aggregator(
user=user_prompt,
json_mode=True,
temperature=0.2,
max_tokens=1000,
max_tokens=4096, # 推理模型需要更大余量
)
if resp.error:
raise RuntimeError(f"aggregator LLM error: {resp.error}")
@@ -171,12 +201,11 @@ async def predict_match_multi(
prediction_cutoff_at = header.match_dt # 默认:比赛时间作为数据截止
now = datetime.now(timezone.utc)
# 2. 并行专家
specialist_provider = _get_specialist_provider()
reports = await run_specialists(header, provider=specialist_provider, version=version)
# 2. 并行专家(各自独立配置)
reports = await run_specialists(header, version=version)
# 3. 终裁
aggregator_provider = _get_aggregator_provider()
aggregator_provider = await _agent_provider("aggregator", tier="aggregator")
final, agg_prompt_tokens, agg_completion_tokens = await run_aggregator(
header, reports, provider=aggregator_provider, version=version
)
@@ -203,30 +232,33 @@ async def predict_match_multi(
# agent_weights 同样必须过校验(旧实现直接取 raw 值落库,未做任何检查)
agent_weights = validate_agent_weights(final.get("agent_weights"))
pred = Prediction(
pred = await _upsert_prediction(
session,
match_id=match_id,
provider=settings.LLM_PROVIDER,
provider_name=settings.LLM_PROVIDER,
model=aggregator_provider.model,
prompt_version=f"multi_{version}",
mode="multi",
prompt_tokens=sum(r.prompt_tokens or 0 for r in reports) + agg_prompt_tokens,
completion_tokens=sum(r.completion_tokens or 0 for r in reports) + agg_completion_tokens,
latency_ms=latency_ms,
pred_home_goals=validated.pred_home_goals,
pred_away_goals=validated.pred_away_goals,
pred_1x2=validated.pred_1x2,
subjective_confidence=validated.subjective_confidence,
reasoning=validated.reasoning,
raw_response=final,
agent_outputs=[r.to_dict() for r in reports],
status="success",
match_kickoff_at=match_kickoff_at,
prediction_cutoff_at=prediction_cutoff_at,
prediction_created_at=now,
input_hash=input_hash,
values={
"prompt_version": f"multi_{version}",
"prompt_tokens": sum(r.prompt_tokens or 0 for r in reports) + agg_prompt_tokens,
"completion_tokens": sum(r.completion_tokens or 0 for r in reports) + agg_completion_tokens,
"latency_ms": latency_ms,
"pred_home_goals": validated.pred_home_goals,
"pred_away_goals": validated.pred_away_goals,
"alt_pred_home_goals": validated.alt_pred_home_goals,
"alt_pred_away_goals": validated.alt_pred_away_goals,
"pred_1x2": validated.pred_1x2,
"subjective_confidence": validated.subjective_confidence,
"reasoning": validated.reasoning,
"raw_response": final,
"agent_outputs": [r.to_dict() for r in reports],
"status": "success",
"match_kickoff_at": match_kickoff_at,
"prediction_cutoff_at": prediction_cutoff_at,
"prediction_created_at": now,
"input_hash": input_hash,
},
)
session.add(pred)
await session.refresh(pred)
return MultiPredictResult(
prediction_id=pred.id,
@@ -236,6 +268,8 @@ async def predict_match_multi(
mode="multi",
pred_home_goals=pred.pred_home_goals,
pred_away_goals=pred.pred_away_goals,
alt_pred_home_goals=pred.alt_pred_home_goals,
alt_pred_away_goals=pred.alt_pred_away_goals,
pred_1x2=pred.pred_1x2,
subjective_confidence=pred.subjective_confidence,
reasoning=pred.reasoning,
+9
View File
@@ -20,6 +20,7 @@ from src.db.unit_of_work import get_uow
from src.llm.eval import settle_prediction
from src.llm.predict import predict_match
from src.llm.utils import actual_1x2
from src.data.team_names_zh import zh_name
logger = logging.getLogger(__name__)
@@ -31,6 +32,8 @@ class BacktestMatchResult:
league_code: str | None
home_team: str
away_team: str
home_team_zh: str | None
away_team_zh: str | None
match_date: str
actual_home: int
actual_away: int
@@ -54,6 +57,8 @@ class BacktestCandidate:
league_code: str | None
home_team: str
away_team: str
home_team_zh: str | None
away_team_zh: str | None
match_date: datetime
home_goals: int
away_goals: int
@@ -113,6 +118,8 @@ async def _get_historical_matches(
league_code=m.league.code if m.league else None,
home_team=m.home_team.name if m.home_team else "?",
away_team=m.away_team.name if m.away_team else "?",
home_team_zh=zh_name(m.home_team.name) if m.home_team else None,
away_team_zh=zh_name(m.away_team.name) if m.away_team else None,
match_date=m.match_date,
home_goals=m.home_goals,
away_goals=m.away_goals,
@@ -164,6 +171,8 @@ async def run_backtest(
league_code=c.league_code,
home_team=c.home_team,
away_team=c.away_team,
home_team_zh=c.home_team_zh,
away_team_zh=c.away_team_zh,
match_date=c.match_date.strftime("%Y-%m-%d") if c.match_date else "?",
actual_home=c.home_goals,
actual_away=c.away_goals,
+68 -21
View File
@@ -11,6 +11,8 @@ from pathlib import Path
from threading import Lock
from src.core.config import settings
from sqlalchemy import select
from src.db.base import AsyncSessionLocal
from src.db.models import Match, Prediction
from src.db.unit_of_work import get_uow
@@ -89,6 +91,8 @@ class PredictResult:
prompt_version: str
pred_home_goals: float | None
pred_away_goals: float | None
alt_pred_home_goals: int | None
alt_pred_away_goals: int | None
pred_1x2: str | None
subjective_confidence: float | None
reasoning: str | None
@@ -97,6 +101,43 @@ class PredictResult:
raw: dict | None
async def _upsert_prediction(
session,
*,
match_id: int,
provider_name: str,
model: str,
mode: str,
values: dict,
) -> Prediction:
"""按 (match, provider, model) 唯一约束写入预测。
已存在且未结算 → 覆盖更新(重新预测语义);已结算 → 拒绝(保护评估数据)。
"""
existing = (
await session.execute(
select(Prediction).where(
Prediction.match_id == match_id,
Prediction.provider == provider_name,
Prediction.model == model,
)
)
).scalar_one_or_none()
if existing is not None and existing.settled:
raise ValueError("该比赛已有已结算的预测,不能重新预测")
pred = existing if existing is not None else Prediction(
match_id=match_id, provider=provider_name, model=model,
)
pred.mode = mode
for k, v in values.items():
setattr(pred, k, v)
if existing is None:
session.add(pred)
await session.flush() # 拿到自增 id;事务由 UnitOfWork 退出时提交
return pred
async def predict_match(
match_id: int,
*,
@@ -140,7 +181,7 @@ async def _predict_single(
) -> PredictResult:
"""单次调用路径(原有实现)。"""
if provider is None:
provider = get_default_provider()
provider = await get_default_provider()
if model:
provider.model = model
version = prompt_version or "v1"
@@ -172,7 +213,7 @@ async def _predict_single(
user=user_prompt,
json_mode=True,
temperature=0.3,
max_tokens=800,
max_tokens=4096, # 推理模型的 reasoning 也计入输出 token,需留足余量
)
if resp.error:
@@ -194,28 +235,32 @@ async def _predict_single(
if m is None:
raise ValueError(f"match {match_id} not found")
pred = Prediction(
pred = await _upsert_prediction(
session,
match_id=match_id,
provider=settings.LLM_PROVIDER,
provider_name=settings.LLM_PROVIDER,
model=provider.model,
prompt_version=version,
prompt_tokens=resp.prompt_tokens,
completion_tokens=resp.completion_tokens,
latency_ms=resp.latency_ms,
pred_home_goals=validated.pred_home_goals,
pred_away_goals=validated.pred_away_goals,
pred_1x2=validated.pred_1x2,
subjective_confidence=validated.subjective_confidence,
reasoning=validated.reasoning,
raw_response=resp.raw,
status="success",
match_kickoff_at=match_kickoff_at,
prediction_cutoff_at=prediction_cutoff_at,
prediction_created_at=now,
input_hash=input_hash,
mode="single",
values={
"prompt_version": version,
"prompt_tokens": resp.prompt_tokens,
"completion_tokens": resp.completion_tokens,
"latency_ms": resp.latency_ms,
"pred_home_goals": validated.pred_home_goals,
"pred_away_goals": validated.pred_away_goals,
"alt_pred_home_goals": validated.alt_pred_home_goals,
"alt_pred_away_goals": validated.alt_pred_away_goals,
"pred_1x2": validated.pred_1x2,
"subjective_confidence": validated.subjective_confidence,
"reasoning": validated.reasoning,
"raw_response": resp.raw,
"status": "success",
"match_kickoff_at": match_kickoff_at,
"prediction_cutoff_at": prediction_cutoff_at,
"prediction_created_at": now,
"input_hash": input_hash,
},
)
session.add(pred)
await session.refresh(pred)
result = PredictResult(
prediction_id=pred.id,
@@ -224,6 +269,8 @@ async def _predict_single(
prompt_version=version,
pred_home_goals=pred.pred_home_goals,
pred_away_goals=pred.pred_away_goals,
alt_pred_home_goals=pred.alt_pred_home_goals,
alt_pred_away_goals=pred.alt_pred_away_goals,
pred_1x2=pred.pred_1x2,
subjective_confidence=pred.subjective_confidence,
reasoning=pred.reasoning,
+6 -3
View File
@@ -9,7 +9,8 @@
裁决规则:
- 各报告的 confidence 和 data_sufficiency 是采信依据: no_data/error 状态的报告必须忽略,不得编造
- 5 专家维度: form(近期状态) / stats(攻防数据) / home_away(主客因素) / injuries(阵容完整性) / h2h(历史交锋)
- 5 专家: 近期状态分析专家 / 攻防数据分析专家 / 主客因素分析专家 / 阵容完整性分析专家 / 历史交锋分析专家
- 引用专家意见时使用上述全称,不要使用英文代码(form/stats/h2h 等)
- home_edge 是各专家的方向性判断(-1~1),冲突时给出你的权衡理由
- agent_weights 体现你对各报告的采信度(0-1,总和无须为 1)
- reasoning 需引用具体报告的证据
@@ -17,8 +18,10 @@
严格按此 JSON 输出,不要其他内容:
```json
{
"pred_home_goals": <float, >,
"pred_away_goals": <float, >,
"pred_home_goals": <int 0-10, ,>,
"pred_away_goals": <int 0-10, ,>,
"alt_pred_home_goals": <int 0-10, (,)>,
"alt_pred_away_goals": <int 0-10, >,
"1x2": "<'1'|'X'|'2'>",
"confidence": <0.0-1.0>,
"reasoning": "<250 字内推理,引用各报告证据>",
+4 -2
View File
@@ -5,8 +5,10 @@
严格按此 JSON 输出:
```json
{
"pred_home_goals": "<float, 预测主队进球>",
"pred_away_goals": "<float, 预测客队进球>",
"pred_home_goals": "<int 0-10, 预测主队进球,必须是整数>",
"pred_away_goals": "<int 0-10, 预测客队进球,必须是整数>",
"alt_pred_home_goals": "<int 0-10, 备选比分主队进球(第二可能的比分,必须与主选不同)>",
"alt_pred_away_goals": "<int 0-10, 备选比分客队进球>",
"1x2": "<'1'|'X'|'2'>",
"confidence": "<0.0-1.0>",
"score_probable": {"home": "<int>", "away": "<int>", "prob": "<float>"},
+4 -2
View File
@@ -11,8 +11,10 @@
严格按此 JSON 输出,不要其他内容:
```json
{
"pred_home_goals": "<float, 预测主队进球>",
"pred_away_goals": "<float, 预测客队进球>",
"pred_home_goals": "<int 0-10, 预测主队进球,必须是整数>",
"pred_away_goals": "<int 0-10, 预测客队进球,必须是整数>",
"alt_pred_home_goals": "<int 0-10, 备选比分主队进球(第二可能的比分,必须与主选不同)>",
"alt_pred_away_goals": "<int 0-10, 备选比分客队进球>",
"1x2": "<'1'|'X'|'2'>",
"confidence": "<0.0-1.0>",
"score_probable": {"home": "<int>", "away": "<int>", "prob": "<float>"},
+19 -6
View File
@@ -10,8 +10,11 @@ import time
from dataclasses import dataclass, field
from typing import Any
import httpx
from src.core.config import settings
from src.core.http_client import get_client
from src.core.runtime_config import get_runtime_value
logger = logging.getLogger(__name__)
@@ -67,17 +70,26 @@ class LLMProvider:
start = time.perf_counter()
try:
client = get_client()
# 连接与读取分离:端点不可达时 10s 内快速失败,
# 避免每个 agent 各挂满 LLM_TIMEOUT 导致整次预测长时间无响应
resp = await client.post(
f"{self.base_url}/chat/completions",
headers=headers,
json=payload,
timeout=self.timeout,
timeout=httpx.Timeout(connect=10.0, read=float(self.timeout), write=float(self.timeout), pool=10.0),
)
resp.raise_for_status()
data = resp.json()
latency = int((time.perf_counter() - start) * 1000)
usage = data.get("usage", {})
content = data["choices"][0]["message"]["content"]
message = data["choices"][0]["message"]
content = message.get("content") or ""
if not content:
# 推理模型可能把 token 全花在 reasoning_content 上
raise RuntimeError(
"模型未返回文本内容"
+ ("(token 花在推理上,请增大 max_tokens)" if message.get("reasoning_content") else "")
)
parsed = None
if json_mode:
try:
@@ -105,10 +117,11 @@ class LLMProvider:
return LLMResponse(content="", error=str(e), latency_ms=latency)
def get_default_provider() -> LLMProvider:
async def get_default_provider() -> LLMProvider:
"""构造默认 provider:运行时配置(DB)优先,回落 .env。"""
return LLMProvider(
api_key=settings.LLM_API_KEY,
base_url=settings.LLM_BASE_URL,
model=settings.LLM_MODEL,
api_key=await get_runtime_value("LLM_API_KEY"),
base_url=await get_runtime_value("LLM_BASE_URL"),
model=await get_runtime_value("LLM_MODEL"),
timeout=settings.LLM_TIMEOUT,
)
+32 -4
View File
@@ -5,6 +5,7 @@
from __future__ import annotations
import logging
from decimal import ROUND_HALF_UP, Decimal
from pydantic import BaseModel, Field, field_validator, model_validator
@@ -54,12 +55,28 @@ class AgentReportSchema(BaseModel):
class PredictionOutputSchema(BaseModel):
"""最终预测输出的校验 schema。"""
pred_home_goals: float = Field(ge=0.0, le=10.0)
pred_away_goals: float = Field(ge=0.0, le=10.0)
pred_home_goals: int = Field(ge=0, le=10)
pred_away_goals: int = Field(ge=0, le=10)
# 备选比分(次可能比分);缺失/无效/与主选相同 → None
alt_pred_home_goals: int | None = Field(default=None, ge=0, le=10)
alt_pred_away_goals: int | None = Field(default=None, ge=0, le=10)
pred_1x2: str
subjective_confidence: float = Field(ge=0.0, le=1.0)
reasoning: str = ""
@model_validator(mode="after")
def check_alt_score(self) -> "PredictionOutputSchema":
"""备选比分与主选相同则丢弃(备选必须是不同比分)。"""
if (
self.alt_pred_home_goals is not None
and self.alt_pred_away_goals is not None
and self.alt_pred_home_goals == self.pred_home_goals
and self.alt_pred_away_goals == self.pred_away_goals
):
self.alt_pred_home_goals = None
self.alt_pred_away_goals = None
return self""
@field_validator("pred_1x2")
@classmethod
def validate_1x2(cls, v: str) -> str:
@@ -172,9 +189,20 @@ def validate_prediction_output(raw: dict) -> PredictionOutputSchema:
logger.warning("Deprecated field 'confidence' used, prefer 'subjective_confidence'")
conf = raw["confidence"]
def _alt(side: str):
v = raw.get(f"alt_pred_{side}_goals")
if v is None:
return None
try:
return int(Decimal(str(v)).quantize(Decimal("1"), rounding=ROUND_HALF_UP))
except Exception:
return None
return PredictionOutputSchema(
pred_home_goals=float(raw.get("pred_home_goals", 0)),
pred_away_goals=float(raw.get("pred_away_goals", 0)),
pred_home_goals=int(Decimal(str(raw.get("pred_home_goals", 0))).quantize(Decimal("1"), rounding=ROUND_HALF_UP)),
pred_away_goals=int(Decimal(str(raw.get("pred_away_goals", 0))).quantize(Decimal("1"), rounding=ROUND_HALF_UP)),
alt_pred_home_goals=_alt("home"),
alt_pred_away_goals=_alt("away"),
pred_1x2=raw.get("1x2") or raw.get("pred_1x2", "X"),
subjective_confidence=float(conf if conf is not None else 0.5),
reasoning=str(raw.get("reasoning", ""))[:1000],