feat: 核心模块增强 — 加密 + 运行时配置 + 日志缓冲
- crypto.py: API Key 加密/解密工具 - runtime_config.py: 运行时动态配置管理 - log_buffer.py: 内存日志缓冲区 - config.py: 新增加密配置项 - http_client.py: 增强重试和错误处理
This commit is contained in:
@@ -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
@@ -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="未登录或凭证无效")
|
||||
|
||||
@@ -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},
|
||||
)
|
||||
@@ -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": "密码已修改,请用新密码重新登录"}
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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 采集任务失败")
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
@@ -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
|
||||
@@ -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))
|
||||
@@ -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
@@ -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
@@ -43,4 +43,4 @@ LEAGUE_COUNTRIES: dict[str, str] = {
|
||||
"EL": "Europe",
|
||||
}
|
||||
|
||||
REQUEST_INTERVAL = 1.2 # bzzoiro 限速(秒)
|
||||
REQUEST_INTERVAL = 2.0 # bzzoiro 限速(秒);上游限速严厉时宁可慢一点
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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 字内推理,引用各报告证据>",
|
||||
|
||||
@@ -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>"},
|
||||
|
||||
@@ -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
@@ -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
@@ -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],
|
||||
|
||||
Reference in New Issue
Block a user