Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e89ab1a0c9 | ||
|
|
ff0045ad93 | ||
|
|
983b620659 | ||
|
|
fa69795d69 | ||
|
|
c5f92c9a54 |
@@ -3,6 +3,10 @@ APP_ENV=development
|
|||||||
LOG_LEVEL=INFO
|
LOG_LEVEL=INFO
|
||||||
|
|
||||||
# ---- 数据库 ----
|
# ---- 数据库 ----
|
||||||
|
POSTGRES_USER=football
|
||||||
|
POSTGRES_PASSWORD=football
|
||||||
|
POSTGRES_DB=football
|
||||||
|
POSTGRES_PORT=5432
|
||||||
DATABASE_URL=postgresql+asyncpg://football:football@localhost:5432/football
|
DATABASE_URL=postgresql+asyncpg://football:football@localhost:5432/football
|
||||||
|
|
||||||
# ---- LLM (OpenAI-compatible,必填一个) ----
|
# ---- LLM (OpenAI-compatible,必填一个) ----
|
||||||
|
|||||||
@@ -0,0 +1,81 @@
|
|||||||
|
"""修复 injuries 表约束命名与 ORM 声明不一致
|
||||||
|
|
||||||
|
Revision ID: 0007_injuries_constraint_naming_align
|
||||||
|
Revises: 0006_schema_model_drift_cleanup
|
||||||
|
Create Date: 2026-09-16
|
||||||
|
|
||||||
|
背景(见代码审查报告 P2-5):
|
||||||
|
0003 迁移使用 sa.UniqueConstraint 创建唯一约束,
|
||||||
|
而 ORM models.py 中声明为 Index(..., unique=True)。
|
||||||
|
虽然 PostgreSQL 中两者效果相同(都保证唯一性),
|
||||||
|
但 pg_catalog 中表示不同,会导致:
|
||||||
|
- alembic autogenerate 持续报告漂移
|
||||||
|
- 约束命名约定不一致(uc_ 前缀 vs ix_ 前缀)
|
||||||
|
|
||||||
|
本迁移将 UniqueConstraint 替换为唯一索引,与 ORM 声明对齐。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Sequence, Union
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
import sqlalchemy as sa
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision: str = '0007_injuries_constraint_naming_align'
|
||||||
|
down_revision: Union[str, None] = '0006_schema_model_drift_cleanup'
|
||||||
|
branch_labels: Union[str, Sequence[str], None] = None
|
||||||
|
depends_on: Union[str, Sequence[str], None] = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
bind = op.get_bind()
|
||||||
|
inspector = sa.inspect(bind)
|
||||||
|
|
||||||
|
# 检查当前约束类型
|
||||||
|
constraints = {
|
||||||
|
c["name"]: c
|
||||||
|
for c in inspector.get_unique_constraints("injuries")
|
||||||
|
}
|
||||||
|
indexes = {
|
||||||
|
i["name"]: i
|
||||||
|
for i in inspector.get_indexes("injuries")
|
||||||
|
}
|
||||||
|
|
||||||
|
# 如果存在 UniqueConstraint 形式的 ix_injuries_player_fixture,替换为唯一索引
|
||||||
|
if "ix_injuries_player_fixture" in constraints:
|
||||||
|
# 删除唯一约束
|
||||||
|
op.drop_constraint("ix_injuries_player_fixture", "injuries", type_="unique")
|
||||||
|
|
||||||
|
# 如果不存在同名唯一索引,创建它(与 ORM 声明一致)
|
||||||
|
if "ix_injuries_player_fixture" not in indexes:
|
||||||
|
op.create_index(
|
||||||
|
"ix_injuries_player_fixture",
|
||||||
|
"injuries",
|
||||||
|
["player_id", "fixture_id", "injury_type"],
|
||||||
|
unique=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
bind = op.get_bind()
|
||||||
|
inspector = sa.inspect(bind)
|
||||||
|
|
||||||
|
indexes = {
|
||||||
|
i["name"]: i
|
||||||
|
for i in inspector.get_indexes("injuries")
|
||||||
|
}
|
||||||
|
constraints = {
|
||||||
|
c["name"]: c
|
||||||
|
for c in inspector.get_unique_constraints("injuries")
|
||||||
|
}
|
||||||
|
|
||||||
|
# 恢复为 UniqueConstraint 形式
|
||||||
|
if "ix_injuries_player_fixture" in indexes:
|
||||||
|
op.drop_index("ix_injuries_player_fixture", table_name="injuries")
|
||||||
|
|
||||||
|
if "ix_injuries_player_fixture" not in constraints:
|
||||||
|
op.create_unique_constraint(
|
||||||
|
"ix_injuries_player_fixture",
|
||||||
|
"injuries",
|
||||||
|
["player_id", "fixture_id", "injury_type"],
|
||||||
|
)
|
||||||
@@ -0,0 +1,59 @@
|
|||||||
|
"""为 predictions 表添加 match_id+provider+model 唯一约束
|
||||||
|
|
||||||
|
Revision ID: 0007_predictions_unique_constraint
|
||||||
|
Revises: 0007_injuries_constraint_naming_align
|
||||||
|
Create Date: 2026-09-16
|
||||||
|
|
||||||
|
背景(见代码审查报告 P1-6):
|
||||||
|
同一 match_id + provider + model 组合不应产生重复预测。
|
||||||
|
当前缺少数据库级唯一约束,回测多次运行或并发采集可能产生重复记录,
|
||||||
|
导致统计偏差。
|
||||||
|
|
||||||
|
先清理已存在的重复记录(保留最早创建的那条),再添加唯一约束。
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Sequence, Union
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
import sqlalchemy as sa
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision: str = '0007_predictions_unique_constraint'
|
||||||
|
down_revision: Union[str, None] = '0007_injuries_constraint_naming_align'
|
||||||
|
branch_labels: Union[str, Sequence[str], None] = None
|
||||||
|
depends_on: Union[str, Sequence[str], None] = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
# 1. 清理已存在的重复记录(保留 id 最小的)
|
||||||
|
op.execute(
|
||||||
|
"""
|
||||||
|
DELETE FROM predictions
|
||||||
|
WHERE id NOT IN (
|
||||||
|
SELECT MIN(id)
|
||||||
|
FROM predictions
|
||||||
|
GROUP BY match_id, provider, model
|
||||||
|
)
|
||||||
|
AND match_id IN (
|
||||||
|
SELECT match_id
|
||||||
|
FROM predictions
|
||||||
|
GROUP BY match_id, provider, model
|
||||||
|
HAVING COUNT(*) > 1
|
||||||
|
)
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
|
||||||
|
# 2. 添加唯一约束
|
||||||
|
op.create_unique_constraint(
|
||||||
|
"uq_predictions_match_provider_model",
|
||||||
|
"predictions",
|
||||||
|
["match_id", "provider", "model"],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_constraint(
|
||||||
|
"uq_predictions_match_provider_model",
|
||||||
|
"predictions",
|
||||||
|
type_="unique",
|
||||||
|
)
|
||||||
+6
-6
@@ -2,15 +2,15 @@ services:
|
|||||||
postgres:
|
postgres:
|
||||||
image: postgres:16-alpine
|
image: postgres:16-alpine
|
||||||
environment:
|
environment:
|
||||||
POSTGRES_USER: football
|
POSTGRES_USER: ${POSTGRES_USER:?POSTGRES_USER 未设置}
|
||||||
POSTGRES_PASSWORD: football
|
POSTGRES_PASSWORD: ${POSTGRES_PASSWORD:?POSTGRES_PASSWORD 未设置}
|
||||||
POSTGRES_DB: football
|
POSTGRES_DB: ${POSTGRES_DB:-football}
|
||||||
ports:
|
ports:
|
||||||
- "5432:5432"
|
- "${POSTGRES_PORT:-5432}:5432"
|
||||||
volumes:
|
volumes:
|
||||||
- pgdata:/var/lib/postgresql/data
|
- pgdata:/var/lib/postgresql/data
|
||||||
healthcheck:
|
healthcheck:
|
||||||
test: ["CMD-SHELL", "pg_isready -U football"]
|
test: ["CMD-SHELL", "pg_isready -U ${POSTGRES_USER:?POSTGRES_USER 未设置}"]
|
||||||
interval: 5s
|
interval: 5s
|
||||||
timeout: 5s
|
timeout: 5s
|
||||||
retries: 5
|
retries: 5
|
||||||
@@ -19,7 +19,7 @@ services:
|
|||||||
build: .
|
build: .
|
||||||
command: uvicorn src.api.app:app --host 0.0.0.0 --port 8000 --reload
|
command: uvicorn src.api.app:app --host 0.0.0.0 --port 8000 --reload
|
||||||
ports:
|
ports:
|
||||||
- "8000:8000"
|
- "${API_PORT:-8000}:8000"
|
||||||
env_file: .env
|
env_file: .env
|
||||||
depends_on:
|
depends_on:
|
||||||
postgres:
|
postgres:
|
||||||
|
|||||||
@@ -80,7 +80,7 @@
|
|||||||
|
|
||||||
/* ── 版面切换文字标签(联赛/状态/模式) ── */
|
/* ── 版面切换文字标签(联赛/状态/模式) ── */
|
||||||
.tab {
|
.tab {
|
||||||
@apply whitespace-nowrap px-0.5 py-1 text-sm text-ink-500 transition-colors hover:text-ink-900;
|
@apply relative whitespace-nowrap px-0.5 py-1 text-sm text-ink-500 transition-colors hover:text-ink-900;
|
||||||
}
|
}
|
||||||
.tab-on {
|
.tab-on {
|
||||||
@apply font-medium text-press;
|
@apply font-medium text-press;
|
||||||
|
|||||||
+4
-2
@@ -29,12 +29,14 @@ def create_app() -> FastAPI:
|
|||||||
)
|
)
|
||||||
|
|
||||||
origins = [o.strip() for o in settings.CORS_ORIGINS.split(",") if o.strip()]
|
origins = [o.strip() for o in settings.CORS_ORIGINS.split(",") if o.strip()]
|
||||||
|
methods = [m.strip() for m in settings.CORS_METHODS.split(",") if m.strip()]
|
||||||
|
headers = [h.strip() for h in settings.CORS_HEADERS.split(",") if h.strip()]
|
||||||
app.add_middleware(
|
app.add_middleware(
|
||||||
CORSMiddleware,
|
CORSMiddleware,
|
||||||
allow_origins=origins,
|
allow_origins=origins,
|
||||||
allow_credentials=True,
|
allow_credentials=True,
|
||||||
allow_methods=["*"],
|
allow_methods=methods,
|
||||||
allow_headers=["*"],
|
allow_headers=headers,
|
||||||
)
|
)
|
||||||
|
|
||||||
from src.api.routes.matches import router as matches_router
|
from src.api.routes.matches import router as matches_router
|
||||||
|
|||||||
@@ -19,24 +19,18 @@ from src.core.config import settings
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
_warned_unset = False
|
|
||||||
|
|
||||||
|
|
||||||
async def require_admin_key(x_api_key: str | None = Header(None, alias="X-API-Key")) -> None:
|
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)])`
|
用法: `@router.post("/ingest/bzzoiro", dependencies=[Depends(require_admin_key)])`
|
||||||
"""
|
"""
|
||||||
global _warned_unset
|
|
||||||
|
|
||||||
expected = settings.ADMIN_API_KEY
|
expected = settings.ADMIN_API_KEY
|
||||||
if not expected:
|
if not expected:
|
||||||
if not _warned_unset:
|
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"ADMIN_API_KEY 未设置,采集/回测接口当前【无鉴权】。"
|
"ADMIN_API_KEY 未设置,采集/回测接口当前【无鉴权】。"
|
||||||
"生产环境请设置该环境变量。"
|
"生产环境请设置该环境变量。"
|
||||||
)
|
)
|
||||||
_warned_unset = True
|
|
||||||
return
|
return
|
||||||
|
|
||||||
if not x_api_key or not secrets.compare_digest(x_api_key, expected):
|
if not x_api_key or not secrets.compare_digest(x_api_key, expected):
|
||||||
|
|||||||
@@ -32,6 +32,17 @@ class Settings(BaseSettings):
|
|||||||
|
|
||||||
# --- CORS ---
|
# --- CORS ---
|
||||||
CORS_ORIGINS: str = "http://localhost:5173,http://localhost:3000"
|
CORS_ORIGINS: str = "http://localhost:5173,http://localhost:3000"
|
||||||
|
CORS_METHODS: str = "GET,POST,PUT,DELETE,OPTIONS"
|
||||||
|
CORS_HEADERS: str = "Authorization,Content-Type,X-API-Key,Accept"
|
||||||
|
|
||||||
|
# --- HTTP ---
|
||||||
|
HTTP_DEFAULT_TIMEOUT: int = 30
|
||||||
|
|
||||||
|
# --- database pool ---
|
||||||
|
DB_POOL_SIZE: int = 5
|
||||||
|
DB_MAX_OVERFLOW: int = 10
|
||||||
|
DB_POOL_TIMEOUT: int = 30
|
||||||
|
DB_POOL_RECYCLE: int = 1800
|
||||||
|
|
||||||
# --- 管理接口鉴权 ---
|
# --- 管理接口鉴权 ---
|
||||||
# 采集 / 回测等高成本或写入型接口需要此 Key(请求头 X-API-Key)。
|
# 采集 / 回测等高成本或写入型接口需要此 Key(请求头 X-API-Key)。
|
||||||
|
|||||||
@@ -2,24 +2,27 @@
|
|||||||
|
|
||||||
使用方:
|
使用方:
|
||||||
- src/llm/provider.py: LLM 调用
|
- src/llm/provider.py: LLM 调用
|
||||||
|
- src/data/bzzoiro.py: bzzoiro 比赛数据
|
||||||
- src/data/understat.py: xG 抓取
|
- src/data/understat.py: xG 抓取
|
||||||
- src/data/injuries.py: 伤停抓取
|
- src/data/injuries.py: 伤停抓取
|
||||||
|
|
||||||
生命周期由 FastAPI lifespan 管理(关闭时 aclose)。
|
生命周期由 FastAPI lifespan 管理(关闭时 aclose)。
|
||||||
|
调用方可通过 `timeout` 参数覆盖 per-request 超时。
|
||||||
"""
|
"""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
|
from src.core.config import settings
|
||||||
|
|
||||||
_shared_client: httpx.AsyncClient | None = None
|
_shared_client: httpx.AsyncClient | None = None
|
||||||
_default_timeout = 30
|
|
||||||
|
|
||||||
|
|
||||||
def get_client() -> httpx.AsyncClient:
|
def get_client() -> httpx.AsyncClient:
|
||||||
"""获取共享客户端(懒初始化)。"""
|
"""获取共享客户端(懒初始化)。"""
|
||||||
global _shared_client
|
global _shared_client
|
||||||
if _shared_client is None or _shared_client.is_closed:
|
if _shared_client is None or _shared_client.is_closed:
|
||||||
_shared_client = httpx.AsyncClient(timeout=_default_timeout)
|
_shared_client = httpx.AsyncClient(timeout=settings.HTTP_DEFAULT_TIMEOUT)
|
||||||
return _shared_client
|
return _shared_client
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,89 +0,0 @@
|
|||||||
"""重试工具:带指数退避的瞬态错误重试。
|
|
||||||
|
|
||||||
NOTE(审查报告 P3):当前全项目**无调用点** —— bzzoiro 在 `_fetch_json_sync`
|
|
||||||
里自带了一套重试逻辑,understat/injuries 各自也有。这里保留是作为后续统一
|
|
||||||
重试策略的落点,但请勿误以为它已在生效。
|
|
||||||
|
|
||||||
如果决定不引入统一重试,建议删除本文件以避免"看起来有重试、实际没有"的误判。
|
|
||||||
"""
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import asyncio
|
|
||||||
import functools
|
|
||||||
import logging
|
|
||||||
import random
|
|
||||||
import time
|
|
||||||
from typing import Callable, Iterable, TypeVar
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
T = TypeVar("T")
|
|
||||||
|
|
||||||
|
|
||||||
def with_retry(
|
|
||||||
*,
|
|
||||||
max_retries: int = 3,
|
|
||||||
base_delay: float = 1.0,
|
|
||||||
max_delay: float = 30.0,
|
|
||||||
retryable_exceptions: Iterable[type[BaseException]] = (Exception,),
|
|
||||||
on_retry: Callable[[Exception, int], None] | None = None,
|
|
||||||
) -> Callable:
|
|
||||||
"""重试装饰器(同步/异步通用,指数退避 + 抖动)。
|
|
||||||
|
|
||||||
Args:
|
|
||||||
max_retries: 最大重试次数
|
|
||||||
base_delay: 基础延迟(秒)
|
|
||||||
max_delay: 最大延迟(秒)
|
|
||||||
retryable_exceptions: 触发重试的异常类型
|
|
||||||
on_retry: 重试回调(exception, attempt)
|
|
||||||
"""
|
|
||||||
retryable = tuple(retryable_exceptions)
|
|
||||||
|
|
||||||
def decorator(func: Callable) -> Callable:
|
|
||||||
@functools.wraps(func)
|
|
||||||
async def async_wrapper(*args, **kwargs):
|
|
||||||
last_exc: Exception | None = None
|
|
||||||
for attempt in range(max_retries + 1):
|
|
||||||
try:
|
|
||||||
return await func(*args, **kwargs)
|
|
||||||
except retryable as e:
|
|
||||||
last_exc = e
|
|
||||||
if attempt == max_retries:
|
|
||||||
break
|
|
||||||
delay = min(base_delay * (2 ** attempt), max_delay)
|
|
||||||
delay += random.uniform(0, delay * 0.1) # 抖动
|
|
||||||
logger.warning(
|
|
||||||
"%s failed (attempt %d/%d), retry in %.1fs: %s",
|
|
||||||
func.__name__, attempt + 1, max_retries, delay, e,
|
|
||||||
)
|
|
||||||
if on_retry:
|
|
||||||
on_retry(e, attempt + 1)
|
|
||||||
await asyncio.sleep(delay)
|
|
||||||
raise last_exc # type: ignore[misc]
|
|
||||||
|
|
||||||
@functools.wraps(func)
|
|
||||||
def sync_wrapper(*args, **kwargs):
|
|
||||||
last_exc: Exception | None = None
|
|
||||||
for attempt in range(max_retries + 1):
|
|
||||||
try:
|
|
||||||
return func(*args, **kwargs)
|
|
||||||
except retryable as e:
|
|
||||||
last_exc = e
|
|
||||||
if attempt == max_retries:
|
|
||||||
break
|
|
||||||
delay = min(base_delay * (2 ** attempt), max_delay)
|
|
||||||
delay += random.uniform(0, delay * 0.1)
|
|
||||||
logger.warning(
|
|
||||||
"%s failed (attempt %d/%d), retry in %.1fs: %s",
|
|
||||||
func.__name__, attempt + 1, max_retries, delay, e,
|
|
||||||
)
|
|
||||||
if on_retry:
|
|
||||||
on_retry(e, attempt + 1)
|
|
||||||
time.sleep(delay)
|
|
||||||
raise last_exc # type: ignore[misc]
|
|
||||||
|
|
||||||
if asyncio.iscoroutinefunction(func):
|
|
||||||
return async_wrapper
|
|
||||||
return sync_wrapper
|
|
||||||
|
|
||||||
return decorator
|
|
||||||
+47
-36
@@ -9,16 +9,13 @@ import asyncio
|
|||||||
import json as _json
|
import json as _json
|
||||||
import logging
|
import logging
|
||||||
import random
|
import random
|
||||||
import time as _time
|
|
||||||
import urllib.error
|
|
||||||
import urllib.parse
|
|
||||||
import urllib.request
|
|
||||||
from collections.abc import Iterable
|
from collections.abc import Iterable
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
|
|
||||||
from src.core.config import settings
|
from src.core.config import settings
|
||||||
|
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.config import BZZOIRO_LEAGUE_IDS, LEAGUE_COUNTRIES, LEAGUE_NAMES, REQUEST_INTERVAL
|
||||||
from src.data.normalize import normalize_bzzoiro
|
from src.data.normalize import normalize_bzzoiro
|
||||||
from src.data.sources import register
|
from src.data.sources import register
|
||||||
@@ -47,43 +44,46 @@ def _match_key(home_team_id: int, away_team_id: int, match_date) -> tuple[int, i
|
|||||||
return (home_team_id, away_team_id, d.isoformat() if d is not None else "")
|
return (home_team_id, away_team_id, d.isoformat() if d is not None else "")
|
||||||
|
|
||||||
|
|
||||||
def _fetch_json_sync(path: str, params: dict | None = None, max_retries: int = 3) -> dict | list:
|
async def _fetch_json_async(path: str, params: dict | None = None, max_retries: int = 3) -> dict | list:
|
||||||
"""同步 HTTP(bzzoiro 客户端保持同步,在 async 函数里 run_in_executor)。"""
|
"""异步 HTTP(bzzoiro 使用 httpx,不再阻塞事件循环线程池)。"""
|
||||||
base = settings.BZZOIRO_BASE.rstrip("/")
|
base = settings.BZZOIRO_BASE.rstrip("/")
|
||||||
url = f"{base}/{path.lstrip('/')}"
|
url = f"{base}/{path.lstrip('/')}"
|
||||||
if params:
|
|
||||||
url += "?" + urllib.parse.urlencode(params)
|
|
||||||
key = settings.BZZOIRO_KEY
|
key = settings.BZZOIRO_KEY
|
||||||
if not key:
|
if not key:
|
||||||
raise RuntimeError("BZZOIRO_KEY 未设置")
|
raise RuntimeError("BZZOIRO_KEY 未设置")
|
||||||
|
|
||||||
|
headers = {
|
||||||
|
"Authorization": f"Token {key}",
|
||||||
|
"Accept": "application/json",
|
||||||
|
}
|
||||||
|
|
||||||
last_exc: Exception | None = None
|
last_exc: Exception | None = None
|
||||||
for attempt in range(max_retries):
|
for attempt in range(max_retries):
|
||||||
try:
|
try:
|
||||||
req = urllib.request.Request(url)
|
client = get_client()
|
||||||
req.add_header("Authorization", f"Token {key}")
|
resp = await client.get(url, headers=headers, params=params, timeout=30)
|
||||||
req.add_header("Accept", "application/json")
|
resp.raise_for_status()
|
||||||
with urllib.request.urlopen(req, timeout=30) as resp:
|
return resp.json()
|
||||||
return _json.loads(resp.read().decode("utf-8"))
|
except Exception as e:
|
||||||
except urllib.error.HTTPError as e:
|
|
||||||
last_exc = e
|
last_exc = e
|
||||||
if e.code == 429:
|
status = getattr(getattr(e, "response", None), "status_code", None)
|
||||||
# 指数退避: 429 通常意味着限速
|
if status == 429:
|
||||||
delay = min(2 ** attempt, 16) + random.uniform(0, 1)
|
delay = min(2 ** attempt, 16) + random.uniform(0, 1)
|
||||||
logger.warning("bzzoiro 429, retry %d in %.1fs", attempt + 1, delay)
|
logger.warning("bzzoiro 429, retry %d in %.1fs", attempt + 1, delay)
|
||||||
_time.sleep(delay)
|
await asyncio.sleep(delay)
|
||||||
continue
|
continue
|
||||||
if 500 <= e.code < 600:
|
if 500 <= (status or 0) < 600:
|
||||||
delay = min(2 ** attempt, 16) + random.uniform(0, 1)
|
delay = min(2 ** attempt, 16) + random.uniform(0, 1)
|
||||||
logger.warning("bzzoiro %d, retry %d in %.1fs", e.code, attempt + 1, delay)
|
logger.warning("bzzoiro %d, retry %d in %.1fs", status, attempt + 1, delay)
|
||||||
_time.sleep(delay)
|
await asyncio.sleep(delay)
|
||||||
continue
|
continue
|
||||||
raise # 4xx 直接抛
|
# 网络错误(连接失败/超时)也退避重试
|
||||||
except (urllib.error.URLError, TimeoutError, ConnectionError) as e:
|
if isinstance(e, (TimeoutError, ConnectionError, OSError)):
|
||||||
last_exc = e
|
|
||||||
delay = min(2 ** attempt, 16) + random.uniform(0, 1)
|
delay = min(2 ** attempt, 16) + random.uniform(0, 1)
|
||||||
logger.warning("bzzoiro network error, retry %d in %.1fs: %s", attempt + 1, delay, e)
|
logger.warning("bzzoiro network error, retry %d in %.1fs: %s", attempt + 1, delay, e)
|
||||||
_time.sleep(delay)
|
await asyncio.sleep(delay)
|
||||||
|
continue
|
||||||
|
raise
|
||||||
raise RuntimeError(f"bzzoiro request failed after {max_retries} attempts: {last_exc}")
|
raise RuntimeError(f"bzzoiro request failed after {max_retries} attempts: {last_exc}")
|
||||||
|
|
||||||
|
|
||||||
@@ -95,15 +95,13 @@ async def fetch_bzzoiro_events(
|
|||||||
date_to: str | None = None,
|
date_to: str | None = None,
|
||||||
limit: int = 200,
|
limit: int = 200,
|
||||||
) -> list[dict]:
|
) -> list[dict]:
|
||||||
"""抓取 bzzoiro 原始事件(异步包装)。"""
|
"""抓取 bzzoiro 原始事件(纯异步,无需 run_in_executor)。"""
|
||||||
league_id = BZZOIRO_LEAGUE_IDS.get(league_code)
|
league_id = BZZOIRO_LEAGUE_IDS.get(league_code)
|
||||||
if league_id is None:
|
if league_id is None:
|
||||||
raise ValueError(f"未知联赛代码: {league_code}")
|
raise ValueError(f"未知联赛代码: {league_code}")
|
||||||
|
|
||||||
loop = asyncio.get_event_loop()
|
|
||||||
rows: list[dict] = []
|
rows: list[dict] = []
|
||||||
offset = 0
|
offset = 0
|
||||||
payload: dict | list = {}
|
|
||||||
while True:
|
while True:
|
||||||
params: dict = {
|
params: dict = {
|
||||||
"league_id": league_id,
|
"league_id": league_id,
|
||||||
@@ -115,8 +113,7 @@ async def fetch_bzzoiro_events(
|
|||||||
params["date_from"] = str(date_from)[:10]
|
params["date_from"] = str(date_from)[:10]
|
||||||
if date_to:
|
if date_to:
|
||||||
params["date_to"] = str(date_to)[:10]
|
params["date_to"] = str(date_to)[:10]
|
||||||
# 显式位置参数,避免 lambda 闭包捕获循环变量
|
payload = await _fetch_json_async("/events/", params)
|
||||||
payload = await loop.run_in_executor(None, _fetch_json_sync, "/events/", params)
|
|
||||||
batch = payload.get("results") or []
|
batch = payload.get("results") or []
|
||||||
if not batch:
|
if not batch:
|
||||||
break
|
break
|
||||||
@@ -186,8 +183,8 @@ class BzzoiroSource:
|
|||||||
try:
|
try:
|
||||||
nm.validate()
|
nm.validate()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.debug("normalize skip: %s", e)
|
# P1-3: 统一使用 warning,不追加到 errors(仅运行时错误入 errors)
|
||||||
league_r["errors"].append(f"normalize: {e}")
|
logger.warning("normalize skip: %s", e)
|
||||||
continue
|
continue
|
||||||
normalized_matches.append((nm, raw))
|
normalized_matches.append((nm, raw))
|
||||||
all_team_names.add(nm.home_team)
|
all_team_names.add(nm.home_team)
|
||||||
@@ -198,11 +195,25 @@ class BzzoiroSource:
|
|||||||
teams = (await db.execute(stmt)).scalars().all()
|
teams = (await db.execute(stmt)).scalars().all()
|
||||||
team_name_to_id = {t.name: t.id for t in teams}
|
team_name_to_id = {t.name: t.id for t in teams}
|
||||||
|
|
||||||
# 预加载已有比赛(完整对象)
|
# P1-2: 按需加载,只加载 raw_events 涉及日期范围的比赛(加 30 天缓冲)
|
||||||
stmt = select(Match).where(Match.league_id == league.id)
|
# 避免加载联赛全部历史比赛到内存(多赛季采集时内存溢出)
|
||||||
for m in (await db.execute(stmt)).scalars():
|
if normalized_matches:
|
||||||
key = _match_key(m.home_team_id, m.away_team_id, m.match_date_date)
|
from datetime import timedelta
|
||||||
existing_matches[key] = m
|
dates = [nm.date for nm 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)
|
||||||
|
stmt = (
|
||||||
|
select(Match)
|
||||||
|
.where(Match.league_id == league.id)
|
||||||
|
.where(Match.match_date >= min_dt)
|
||||||
|
.where(Match.match_date <= max_dt)
|
||||||
|
)
|
||||||
|
existing_matches = {
|
||||||
|
_match_key(m.home_team_id, m.away_team_id, m.match_date_date): m
|
||||||
|
for m in (await db.execute(stmt)).scalars()
|
||||||
|
}
|
||||||
|
# else: existing_matches 保持空 dict(全量新比赛)
|
||||||
|
|
||||||
for nm, raw in normalized_matches:
|
for nm, raw in normalized_matches:
|
||||||
# 球队: 内存查找 + 按需创建
|
# 球队: 内存查找 + 按需创建
|
||||||
|
|||||||
+108
-41
@@ -8,6 +8,7 @@ import asyncio
|
|||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import random
|
import random
|
||||||
|
import tempfile
|
||||||
import time
|
import time
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -23,8 +24,8 @@ logger = logging.getLogger(__name__)
|
|||||||
API_BASE = "https://v3.football.api-sports.io"
|
API_BASE = "https://v3.football.api-sports.io"
|
||||||
DEFAULT_HOST = "v3.football.api-sports.io"
|
DEFAULT_HOST = "v3.football.api-sports.io"
|
||||||
|
|
||||||
# 缓存目录
|
# P2-3: 缓存目录改用系统临时目录,避免源码树内写入
|
||||||
_CACHE_DIR = Path(__file__).resolve().parent.parent.parent / "data" / "injuries_cache"
|
_CACHE_DIR = Path(tempfile.gettempdir()) / "profeto_injuries"
|
||||||
|
|
||||||
|
|
||||||
async def fetch_injuries(*, date: str | None = None, fixture_id: int | None = None, league_id: int | None = None) -> list[dict]:
|
async def fetch_injuries(*, date: str | None = None, fixture_id: int | None = None, league_id: int | None = None) -> list[dict]:
|
||||||
@@ -103,8 +104,12 @@ async def ingest_injuries(db, *, date: str | None = None) -> dict:
|
|||||||
"""采集伤停数据并入库(injuries 表)。
|
"""采集伤停数据并入库(injuries 表)。
|
||||||
|
|
||||||
注意: 本方法不控制事务(commit/rollback),由调用方通过 UnitOfWork 控制。
|
注意: 本方法不控制事务(commit/rollback),由调用方通过 UnitOfWork 控制。
|
||||||
|
|
||||||
|
P1-4: 批量幂等检查,避免逐条查询的竞态条件(并发采集时 IntegrityError)。
|
||||||
"""
|
"""
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.exc import IntegrityError
|
||||||
|
from sqlalchemy.orm import selectinload
|
||||||
|
|
||||||
from src.data.team_names import normalize as normalize_name
|
from src.data.team_names import normalize as normalize_name
|
||||||
from src.db.models import Injury, Team
|
from src.db.models import Injury, Team
|
||||||
@@ -124,6 +129,9 @@ async def ingest_injuries(db, *, date: str | None = None) -> dict:
|
|||||||
teams = (await db.execute(select(Team))).scalars().all()
|
teams = (await db.execute(select(Team))).scalars().all()
|
||||||
team_by_name = {t.name: t.id for t in teams}
|
team_by_name = {t.name: t.id for t in teams}
|
||||||
|
|
||||||
|
# P1-4: 收集所有待插入记录的键,批量查询已存在的记录
|
||||||
|
# 避免逐条查询 + 插入的竞态条件(两个并发请求同时通过检查 → IntegrityError)
|
||||||
|
pending_records: list[dict] = []
|
||||||
for raw in raw_injuries:
|
for raw in raw_injuries:
|
||||||
try:
|
try:
|
||||||
player = raw.get("player", {}) or {}
|
player = raw.get("player", {}) or {}
|
||||||
@@ -144,43 +152,104 @@ async def ingest_injuries(db, *, date: str | None = None) -> dict:
|
|||||||
except (ValueError, AttributeError):
|
except (ValueError, AttributeError):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
# 强制 int 转换,API 可能返回字符串
|
||||||
player_id = player.get("id")
|
player_id = player.get("id")
|
||||||
|
try:
|
||||||
|
player_id = int(player_id) if player_id is not None else None
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
player_id = None
|
||||||
fixture_id = fixture.get("id")
|
fixture_id = fixture.get("id")
|
||||||
|
try:
|
||||||
|
fixture_id = int(fixture_id) if fixture_id is not None else None
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
fixture_id = None
|
||||||
|
|
||||||
# 幂等: 已存在则跳过
|
pending_records.append({
|
||||||
existing = (
|
"player_id": player_id,
|
||||||
await db.execute(
|
"player_name": player_name,
|
||||||
select(Injury).where(
|
"team_id": team_id,
|
||||||
Injury.player_id == player_id,
|
"fixture_id": fixture_id,
|
||||||
Injury.fixture_id == fixture_id,
|
"league_id": (raw.get("league") or {}).get("id"),
|
||||||
Injury.injury_type == player.get("type"),
|
"injury_type": player.get("type"),
|
||||||
)
|
"reason": player.get("reason"),
|
||||||
)
|
"injury_date": injury_date,
|
||||||
).scalar_one_or_none()
|
})
|
||||||
|
|
||||||
if existing is not None:
|
|
||||||
continue
|
|
||||||
|
|
||||||
injury = Injury(
|
|
||||||
player_id=player_id,
|
|
||||||
player_name=player_name,
|
|
||||||
team_id=team_id,
|
|
||||||
fixture_id=fixture_id,
|
|
||||||
league_id=(raw.get("league") or {}).get("id"),
|
|
||||||
injury_type=player.get("type"),
|
|
||||||
reason=player.get("reason"),
|
|
||||||
injury_date=injury_date,
|
|
||||||
)
|
|
||||||
db.add(injury)
|
|
||||||
result["inserted"] += 1
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
result["errors"].append(f"parse error: {e}")
|
result["errors"].append(f"parse error: {e}")
|
||||||
|
|
||||||
|
# P1-4: 批量查询已存在的记录(1 次 DB 往返)
|
||||||
|
existing_keys: set[tuple] = set()
|
||||||
|
if pending_records:
|
||||||
|
# 构造查询条件:所有 (player_id, fixture_id, injury_type) 组合
|
||||||
|
# 使用 OR 条件批量查询
|
||||||
|
conditions = []
|
||||||
|
for rec in pending_records:
|
||||||
|
conditions.append(
|
||||||
|
(Injury.player_id == rec["player_id"])
|
||||||
|
& (Injury.fixture_id == rec["fixture_id"])
|
||||||
|
& (Injury.injury_type == rec["injury_type"])
|
||||||
|
)
|
||||||
|
if conditions:
|
||||||
|
from sqlalchemy import or_
|
||||||
|
stmt = select(Injury.player_id, Injury.fixture_id, Injury.injury_type).where(or_(*conditions))
|
||||||
|
rows = (await db.execute(stmt)).all()
|
||||||
|
existing_keys = {(r[0], r[1], r[2]) for r in rows}
|
||||||
|
|
||||||
|
# P1-4: 批量插入(跳过已存在的)
|
||||||
|
for rec in pending_records:
|
||||||
|
key = (rec["player_id"], rec["fixture_id"], rec["injury_type"])
|
||||||
|
if key in existing_keys:
|
||||||
|
continue
|
||||||
|
|
||||||
|
injury = Injury(**rec)
|
||||||
|
db.add(injury)
|
||||||
|
result["inserted"] += 1
|
||||||
|
|
||||||
|
# 每 50 条 flush 一次,减少内存压力,同时捕获 IntegrityError
|
||||||
|
if result["inserted"] % 50 == 0:
|
||||||
|
try:
|
||||||
|
await db.flush()
|
||||||
|
except IntegrityError:
|
||||||
|
# P1-4: 并发采集时可能仍有竞态,回退到逐条插入
|
||||||
|
await db.rollback()
|
||||||
|
logger.warning("injuries batch IntegrityError, falling back to per-record insert")
|
||||||
|
return await _ingest_injuries_fallback(db, pending_records, result)
|
||||||
|
|
||||||
|
# 最终 flush
|
||||||
|
try:
|
||||||
|
await db.flush()
|
||||||
|
except IntegrityError:
|
||||||
|
await db.rollback()
|
||||||
|
logger.warning("injuries final flush IntegrityError, falling back to per-record insert")
|
||||||
|
return await _ingest_injuries_fallback(db, pending_records, result)
|
||||||
|
|
||||||
# 注意: 不在此处 commit,由调用方 UnitOfWork 控制事务
|
# 注意: 不在此处 commit,由调用方 UnitOfWork 控制事务
|
||||||
logger.info("injuries: fetched %d, inserted %d for %s", result["count"], result["inserted"], date)
|
logger.info("injuries: fetched %d, inserted %d for %s", result["count"], result["inserted"], date)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
async def _ingest_injuries_fallback(db, pending_records: list[dict], result: dict) -> dict:
|
||||||
|
"""P1-4: 逐条插入回退,捕获每条 IntegrityError 避免整批回滚。"""
|
||||||
|
from sqlalchemy.exc import IntegrityError
|
||||||
|
from src.db.models import Injury
|
||||||
|
|
||||||
|
inserted = 0
|
||||||
|
for rec in pending_records:
|
||||||
|
injury = Injury(**rec)
|
||||||
|
db.add(injury)
|
||||||
|
try:
|
||||||
|
await db.flush()
|
||||||
|
inserted += 1
|
||||||
|
except IntegrityError:
|
||||||
|
await db.rollback()
|
||||||
|
# 已存在或其他冲突,跳过
|
||||||
|
continue
|
||||||
|
|
||||||
|
result["inserted"] = inserted
|
||||||
|
logger.info("injuries fallback: inserted %d records", inserted)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
async def get_injuries_for_match(db, team_id: int, match_date, as_of=None) -> list[Injury]:
|
async def get_injuries_for_match(db, team_id: int, match_date, as_of=None) -> list[Injury]:
|
||||||
"""查询某场比赛前某队的伤停名单(比赛日仍缺阵的)。
|
"""查询某场比赛前某队的伤停名单(比赛日仍缺阵的)。
|
||||||
|
|
||||||
@@ -188,31 +257,29 @@ async def get_injuries_for_match(db, team_id: int, match_date, as_of=None) -> li
|
|||||||
db: 数据库 session
|
db: 数据库 session
|
||||||
team_id: 球队 ID
|
team_id: 球队 ID
|
||||||
match_date: 比赛日期
|
match_date: 比赛日期
|
||||||
as_of: 截止时间(cutoff)。只返回 retrieved_at <= as_of 的记录。
|
as_of: 数据截止时间(用于回测防泄漏)
|
||||||
用于回测时防止"未来采集的数据"泄漏到历史预测。
|
|
||||||
必须保持 timezone-aware datetime,不会截断为 date。
|
|
||||||
"""
|
|
||||||
from sqlalchemy import or_, select
|
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
伤停记录列表
|
||||||
|
"""
|
||||||
|
from sqlalchemy import select
|
||||||
from src.db.models import Injury
|
from src.db.models import Injury
|
||||||
|
|
||||||
# 只处理 match_date:去掉时间部分,仅比较日期
|
if hasattr(match_date, "date") and callable(match_date.date):
|
||||||
if hasattr(match_date, "date"):
|
|
||||||
match_date = match_date.date()
|
match_date = match_date.date()
|
||||||
|
|
||||||
stmt = (
|
stmt = (
|
||||||
select(Injury)
|
select(Injury)
|
||||||
.where(Injury.team_id == team_id)
|
.where(Injury.team_id == team_id)
|
||||||
.where(Injury.injury_date <= match_date)
|
.where(Injury.injury_date <= match_date)
|
||||||
.where(or_(Injury.return_date.is_(None), Injury.return_date >= match_date))
|
.where(
|
||||||
|
(Injury.return_date.is_(None)) | (Injury.return_date >= match_date)
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
# 回测防泄漏: 只使用 as_of 时间点之前已采集的数据
|
|
||||||
# 注意: as_of 保持 datetime,不截断为 date,避免错误排除同日合法数据
|
|
||||||
if as_of is not None:
|
if as_of is not None:
|
||||||
stmt = stmt.where(Injury.retrieved_at.is_not(None))
|
if hasattr(as_of, "date") and callable(as_of.date):
|
||||||
|
as_of = as_of.date()
|
||||||
stmt = stmt.where(Injury.retrieved_at <= as_of)
|
stmt = stmt.where(Injury.retrieved_at <= as_of)
|
||||||
|
|
||||||
stmt = stmt.order_by(Injury.injury_date.desc())
|
|
||||||
result = await db.execute(stmt)
|
result = await db.execute(stmt)
|
||||||
return list(result.scalars().all())
|
return list(result.scalars().all())
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
"""数据规范化:任意数据源原始记录 → NormalizedMatch。
|
"""数据规范化:任意数据源原始记录 → NormalizedMatch。
|
||||||
|
|
||||||
迁移自旧项目 app/data/normalize.py,简化:
|
迁移自旧项目 app/data/normalize.py,简化:
|
||||||
- 去掉 XGBackfill 双轨(不再需要独立回填)
|
- 去掉 XGBackoff 双轨(不再需要独立回填)
|
||||||
- 去掉 PIT 时间契约(无训练集要防泄漏)
|
- 去掉 PIT 时间契约(无训练集要防泄漏)
|
||||||
- 保留核心清洗契约(队名归一、日期解析、数值范围)
|
- 保留核心清洗契约(队名归一、日期解析、数值范围)
|
||||||
"""
|
"""
|
||||||
@@ -90,7 +90,10 @@ def derive_season_label(date: datetime) -> str:
|
|||||||
|
|
||||||
|
|
||||||
def _parse_date(value) -> datetime | None:
|
def _parse_date(value) -> datetime | None:
|
||||||
"""日期解析 → UTC datetime(带 tzinfo)。"""
|
"""日期解析 → UTC datetime(带 tzinfo)。
|
||||||
|
|
||||||
|
P2-2: 解析失败时记录 warning,避免静默丢数据而无感知。
|
||||||
|
"""
|
||||||
if value in (None, ""):
|
if value in (None, ""):
|
||||||
return None
|
return None
|
||||||
if isinstance(value, (int, float)):
|
if isinstance(value, (int, float)):
|
||||||
@@ -111,6 +114,8 @@ def _parse_date(value) -> datetime | None:
|
|||||||
return datetime.strptime(s[:19], fmt).replace(tzinfo=timezone.utc)
|
return datetime.strptime(s[:19], fmt).replace(tzinfo=timezone.utc)
|
||||||
except ValueError:
|
except ValueError:
|
||||||
continue
|
continue
|
||||||
|
# P2-2 修复: 记录被丢弃的原始值,便于排查数据源格式变更
|
||||||
|
logger.warning("_parse_date failed, dropping record: %r", value)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
@@ -204,6 +209,6 @@ def normalize_understat(raw: dict, league_type: str) -> NormalizedMatch | None:
|
|||||||
away_team=away,
|
away_team=away,
|
||||||
match_status="finished",
|
match_status="finished",
|
||||||
season_label=derive_season_label(dt),
|
season_label=derive_season_label(dt),
|
||||||
home_xg=_to_float(home_xg),
|
home_xg=home_xg,
|
||||||
away_xg=_to_float(away_xg),
|
away_xg=away_xg,
|
||||||
)
|
)
|
||||||
|
|||||||
+57
-11
@@ -10,9 +10,10 @@ import json
|
|||||||
import logging
|
import logging
|
||||||
import random
|
import random
|
||||||
import re
|
import re
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
|
|
||||||
from sqlalchemy import func, select
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.orm import selectinload
|
||||||
|
|
||||||
from src.core.http_client import get_client
|
from src.core.http_client import get_client
|
||||||
from src.data.config import FDCO_TO_UNDERSTAT, LEAGUE_NAMES
|
from src.data.config import FDCO_TO_UNDERSTAT, LEAGUE_NAMES
|
||||||
@@ -76,6 +77,18 @@ async def fetch_understat(league_code: str, season: int) -> list[dict]:
|
|||||||
return data
|
return data
|
||||||
|
|
||||||
|
|
||||||
|
def _match_key(home_team_id: int, away_team_id: int, match_date) -> tuple[int, int, str]:
|
||||||
|
"""比赛去重键:(主队, 客队, 天级日期 ISO 字符串)。
|
||||||
|
|
||||||
|
统一在这里构造,避免"预加载时用 str(date)、写入时用 isoformat()"这类
|
||||||
|
隐式格式依赖 —— 两者当前恰好相等,但一旦有人改动其一就会静默失配,
|
||||||
|
导致所有比赛被判为不存在而重复插入。
|
||||||
|
"""
|
||||||
|
if hasattr(match_date, "date") and callable(match_date.date):
|
||||||
|
match_date = match_date.date()
|
||||||
|
return (home_team_id, away_team_id, match_date.isoformat() if match_date is not None else "")
|
||||||
|
|
||||||
|
|
||||||
@register
|
@register
|
||||||
class UnderstatSource:
|
class UnderstatSource:
|
||||||
"""understat xG 数据源(实现 DataSource 协议)。"""
|
"""understat xG 数据源(实现 DataSource 协议)。"""
|
||||||
@@ -86,8 +99,10 @@ class UnderstatSource:
|
|||||||
"""采集 understat xG → 回填到现有 Match。只回填 xG 字段,不创建新 Match。
|
"""采集 understat xG → 回填到现有 Match。只回填 xG 字段,不创建新 Match。
|
||||||
|
|
||||||
注意: 本方法不控制事务(commit/rollback),由调用方通过 UnitOfWork 控制。
|
注意: 本方法不控制事务(commit/rollback),由调用方通过 UnitOfWork 控制。
|
||||||
|
|
||||||
|
P1-3: 批量查询优化,将单赛季 380 场 × 3 次 DB 往返降为 3 次查询。
|
||||||
"""
|
"""
|
||||||
from src.db.repositories import LeagueRepository, MatchRepository, TeamRepository
|
from src.db.repositories import LeagueRepository, TeamRepository
|
||||||
|
|
||||||
result = {"updated": 0, "skipped": 0, "unmatched": 0, "errors": []}
|
result = {"updated": 0, "skipped": 0, "unmatched": 0, "errors": []}
|
||||||
|
|
||||||
@@ -101,7 +116,6 @@ class UnderstatSource:
|
|||||||
# 使用 Repository
|
# 使用 Repository
|
||||||
league_repo = LeagueRepository(db)
|
league_repo = LeagueRepository(db)
|
||||||
team_repo = TeamRepository(db)
|
team_repo = TeamRepository(db)
|
||||||
match_repo = MatchRepository(db)
|
|
||||||
|
|
||||||
# 查联赛
|
# 查联赛
|
||||||
league_obj = await league_repo.get_by_code(league)
|
league_obj = await league_repo.get_by_code(league)
|
||||||
@@ -109,6 +123,9 @@ class UnderstatSource:
|
|||||||
result["errors"].append(f"league {league} not found in DB")
|
result["errors"].append(f"league {league} not found in DB")
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
# === 批量优化: 一次规范化,收集球队名和日期 ===
|
||||||
|
normalized_matches: list = []
|
||||||
|
all_team_names: set[str] = set()
|
||||||
for raw in raw_matches:
|
for raw in raw_matches:
|
||||||
if not raw.get("isResult"):
|
if not raw.get("isResult"):
|
||||||
continue
|
continue
|
||||||
@@ -120,17 +137,46 @@ class UnderstatSource:
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
result["errors"].append(f"normalize: {e}")
|
result["errors"].append(f"normalize: {e}")
|
||||||
continue
|
continue
|
||||||
|
normalized_matches.append((nm, raw))
|
||||||
|
all_team_names.add(nm.home_team)
|
||||||
|
all_team_names.add(nm.away_team)
|
||||||
|
|
||||||
# 匹配已有 Match(天级) - 使用 Repository
|
if not normalized_matches:
|
||||||
home_team = await team_repo.get_by_name(nm.home_team)
|
return result
|
||||||
away_team = await team_repo.get_by_name(nm.away_team)
|
|
||||||
if home_team is None or away_team is None:
|
# === 批量查询球队(1 次 DB 往返) ===
|
||||||
|
team_name_to_id = {}
|
||||||
|
if all_team_names:
|
||||||
|
teams = await team_repo.get_all_by_names(list(all_team_names))
|
||||||
|
team_name_to_id = {name: team.id for name, team in teams.items()}
|
||||||
|
|
||||||
|
# === 批量查询已有比赛(1 次 DB 往返,按日期范围) ===
|
||||||
|
match_dict: dict[tuple, Match] = {}
|
||||||
|
dates = [nm.date for nm, _ 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)
|
||||||
|
stmt = (
|
||||||
|
select(Match)
|
||||||
|
.options(selectinload(Match.stats))
|
||||||
|
.where(Match.league_id == league_obj.id)
|
||||||
|
.where(Match.match_date >= min_dt)
|
||||||
|
.where(Match.match_date <= max_dt)
|
||||||
|
)
|
||||||
|
for m in (await db.execute(stmt)).scalars():
|
||||||
|
key = _match_key(m.home_team_id, m.away_team_id, m.match_date_date)
|
||||||
|
match_dict[key] = m
|
||||||
|
|
||||||
|
# === 内存匹配 + 回填 xG ===
|
||||||
|
for nm, raw in normalized_matches:
|
||||||
|
home_team_id = team_name_to_id.get(nm.home_team)
|
||||||
|
away_team_id = team_name_to_id.get(nm.away_team)
|
||||||
|
if home_team_id is None or away_team_id is None:
|
||||||
result["unmatched"] += 1
|
result["unmatched"] += 1
|
||||||
continue
|
continue
|
||||||
|
|
||||||
existing = await match_repo.find_by_teams_and_date(
|
match_key = _match_key(home_team_id, away_team_id, nm.date)
|
||||||
league_obj.id, home_team.id, away_team.id, nm.date
|
existing = match_dict.get(match_key)
|
||||||
)
|
|
||||||
if existing is None:
|
if existing is None:
|
||||||
result["unmatched"] += 1
|
result["unmatched"] += 1
|
||||||
continue
|
continue
|
||||||
|
|||||||
+4
-2
@@ -17,8 +17,10 @@ engine = create_async_engine(
|
|||||||
settings.DATABASE_URL,
|
settings.DATABASE_URL,
|
||||||
echo=False,
|
echo=False,
|
||||||
pool_pre_ping=True,
|
pool_pre_ping=True,
|
||||||
pool_size=10,
|
pool_size=settings.DB_POOL_SIZE,
|
||||||
max_overflow=20,
|
max_overflow=settings.DB_MAX_OVERFLOW,
|
||||||
|
pool_timeout=settings.DB_POOL_TIMEOUT,
|
||||||
|
pool_recycle=settings.DB_POOL_RECYCLE,
|
||||||
)
|
)
|
||||||
|
|
||||||
AsyncSessionLocal = async_sessionmaker(
|
AsyncSessionLocal = async_sessionmaker(
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from sqlalchemy import (
|
|||||||
Integer,
|
Integer,
|
||||||
String,
|
String,
|
||||||
Text,
|
Text,
|
||||||
|
UniqueConstraint,
|
||||||
func,
|
func,
|
||||||
)
|
)
|
||||||
from sqlalchemy.dialects.postgresql import JSONB
|
from sqlalchemy.dialects.postgresql import JSONB
|
||||||
@@ -198,6 +199,11 @@ class Prediction(Base):
|
|||||||
match: Mapped[Match] = relationship(back_populates="predictions")
|
match: Mapped[Match] = relationship(back_populates="predictions")
|
||||||
|
|
||||||
__table_args__ = (
|
__table_args__ = (
|
||||||
|
# P1-6: 数据库级唯一约束,防止同一 match+provider+model 产生重复预测
|
||||||
|
UniqueConstraint(
|
||||||
|
"match_id", "provider", "model",
|
||||||
|
name="uq_predictions_match_provider_model",
|
||||||
|
),
|
||||||
Index("ix_predictions_match", "match_id"),
|
Index("ix_predictions_match", "match_id"),
|
||||||
Index("ix_predictions_provider_model", "provider", "model"),
|
Index("ix_predictions_provider_model", "provider", "model"),
|
||||||
# 数据截止时间过滤查询用(按 prediction_cutoff_at 取「赛前已生成」的预测)
|
# 数据截止时间过滤查询用(按 prediction_cutoff_at 取「赛前已生成」的预测)
|
||||||
|
|||||||
+12
-3
@@ -5,7 +5,8 @@ Repository 只负责查询,不负责事务提交。
|
|||||||
"""
|
"""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from sqlalchemy import func, select
|
from datetime import datetime
|
||||||
|
from sqlalchemy import select
|
||||||
from sqlalchemy.orm import selectinload
|
from sqlalchemy.orm import selectinload
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
@@ -41,9 +42,17 @@ class MatchRepository:
|
|||||||
|
|
||||||
预加载 stats:调用方(understat 回填)会读取 existing.stats,
|
预加载 stats:调用方(understat 回填)会读取 existing.stats,
|
||||||
async session 下惰性加载会抛 MissingGreenlet。
|
async session 下惰性加载会抛 MissingGreenlet。
|
||||||
|
|
||||||
|
P2-3: 使用 match_date_date(已建索引)做等值匹配,避免 func.date()
|
||||||
|
导致的全表扫描。
|
||||||
"""
|
"""
|
||||||
if hasattr(date, "date"):
|
if isinstance(date, datetime):
|
||||||
date = date.date()
|
date = date.date()
|
||||||
|
elif hasattr(date, "date"):
|
||||||
|
date = date.date()
|
||||||
|
else:
|
||||||
|
# 字符串等其它格式,尝试转换
|
||||||
|
date = datetime.fromisoformat(str(date)).date()
|
||||||
|
|
||||||
stmt = (
|
stmt = (
|
||||||
select(Match)
|
select(Match)
|
||||||
@@ -51,7 +60,7 @@ class MatchRepository:
|
|||||||
.where(Match.league_id == league_id)
|
.where(Match.league_id == league_id)
|
||||||
.where(Match.home_team_id == home_team_id)
|
.where(Match.home_team_id == home_team_id)
|
||||||
.where(Match.away_team_id == away_team_id)
|
.where(Match.away_team_id == away_team_id)
|
||||||
.where(func.date(Match.match_date) == date)
|
.where(Match.match_date_date == date)
|
||||||
)
|
)
|
||||||
return (await self._session.execute(stmt)).scalar_one_or_none()
|
return (await self._session.execute(stmt)).scalar_one_or_none()
|
||||||
|
|
||||||
|
|||||||
@@ -9,12 +9,16 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from collections.abc import AsyncIterator
|
from collections.abc import AsyncIterator
|
||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from src.db.base import AsyncSessionLocal
|
from src.db.base import AsyncSessionLocal
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
async def get_uow() -> AsyncIterator[AsyncSessionLocal]:
|
async def get_uow() -> AsyncIterator[AsyncSession]:
|
||||||
"""创建新的工作单元(用于非路由上下文)。
|
"""创建新的工作单元(用于非路由上下文)。
|
||||||
|
|
||||||
用法:
|
用法:
|
||||||
|
|||||||
+20
-28
@@ -3,9 +3,11 @@
|
|||||||
核心机制:
|
核心机制:
|
||||||
- build_context 已内置 before=match_date,天然防未来信息泄漏
|
- build_context 已内置 before=match_date,天然防未来信息泄漏
|
||||||
- 对历史比赛跑预测 → 用实际比分 settle → 统计准确率
|
- 对历史比赛跑预测 → 用实际比分 settle → 统计准确率
|
||||||
|
- 并发控制: asyncio.Semaphore 限制同时 LLM 调用数
|
||||||
"""
|
"""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
@@ -17,6 +19,7 @@ from src.db.models import Match
|
|||||||
from src.db.unit_of_work import get_uow
|
from src.db.unit_of_work import get_uow
|
||||||
from src.llm.eval import settle_prediction
|
from src.llm.eval import settle_prediction
|
||||||
from src.llm.predict import predict_match
|
from src.llm.predict import predict_match
|
||||||
|
from src.llm.utils import actual_1x2
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -69,15 +72,6 @@ class BacktestSummary:
|
|||||||
results: list[BacktestMatchResult] = field(default_factory=list)
|
results: list[BacktestMatchResult] = field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
def _actual_1x2(home: int, away: int) -> str:
|
|
||||||
"""实际比分 → 胜平负。"""
|
|
||||||
if home > away:
|
|
||||||
return "1"
|
|
||||||
if home == away:
|
|
||||||
return "X"
|
|
||||||
return "2"
|
|
||||||
|
|
||||||
|
|
||||||
async def _get_historical_matches(
|
async def _get_historical_matches(
|
||||||
db,
|
db,
|
||||||
*,
|
*,
|
||||||
@@ -156,21 +150,16 @@ async def run_backtest(
|
|||||||
|
|
||||||
summary = BacktestSummary(total=len(candidates), scored=0)
|
summary = BacktestSummary(total=len(candidates), scored=0)
|
||||||
|
|
||||||
for c in candidates:
|
# P1-6: 并发控制,同时最多 8 场预测(避免 LLM API 限流)
|
||||||
|
sem = asyncio.Semaphore(8)
|
||||||
|
|
||||||
|
async def _one(c: BacktestCandidate) -> BacktestMatchResult | None:
|
||||||
|
async with sem:
|
||||||
try:
|
try:
|
||||||
# 预测 (build_context 内部已用 before=match_date 防泄漏,
|
result = await predict_match(c.match_id, mode=mode, model=model, use_cache=False, backtest=True)
|
||||||
# injuries_slice 也使用 as_of=match_date 过滤 retrieved_at)
|
|
||||||
# 回测必须禁用结果缓存: 否则命中缓存会复用同一 prediction_id,
|
|
||||||
# 导致 settle 反复覆盖同一条记录(见 P1-3)。
|
|
||||||
result = await predict_match(c.match_id, mode=mode, model=model, use_cache=False)
|
|
||||||
|
|
||||||
# 用实际比分 settle
|
|
||||||
await settle_prediction(result.prediction_id, c.home_goals, c.away_goals)
|
await settle_prediction(result.prediction_id, c.home_goals, c.away_goals)
|
||||||
|
actual = actual_1x2(c.home_goals, c.away_goals)
|
||||||
actual = _actual_1x2(c.home_goals, c.away_goals)
|
return BacktestMatchResult(
|
||||||
correct = result.pred_1x2 == actual
|
|
||||||
|
|
||||||
bt = BacktestMatchResult(
|
|
||||||
match_id=c.match_id,
|
match_id=c.match_id,
|
||||||
league_code=c.league_code,
|
league_code=c.league_code,
|
||||||
home_team=c.home_team,
|
home_team=c.home_team,
|
||||||
@@ -183,16 +172,19 @@ async def run_backtest(
|
|||||||
pred_away=result.pred_away_goals,
|
pred_away=result.pred_away_goals,
|
||||||
pred_1x2=result.pred_1x2,
|
pred_1x2=result.pred_1x2,
|
||||||
subjective_confidence=result.subjective_confidence,
|
subjective_confidence=result.subjective_confidence,
|
||||||
correct_1x2=correct,
|
correct_1x2=result.pred_1x2 == actual,
|
||||||
prediction_id=result.prediction_id,
|
prediction_id=result.prediction_id,
|
||||||
)
|
)
|
||||||
summary.results.append(bt)
|
|
||||||
summary.scored += 1
|
|
||||||
|
|
||||||
except Exception:
|
except Exception:
|
||||||
# 用 exception 而非 warning:保留堆栈,否则集成层缺陷(如惰性加载
|
|
||||||
# 在 session 外触发)会只剩一行无堆栈的 warning,极难定位。
|
|
||||||
logger.exception("backtest match %s failed", c.match_id)
|
logger.exception("backtest match %s failed", c.match_id)
|
||||||
|
return None
|
||||||
|
|
||||||
|
# 并行执行,保持结果顺序
|
||||||
|
results = await asyncio.gather(*[_one(c) for c in candidates])
|
||||||
|
for r in results:
|
||||||
|
if r is not None:
|
||||||
|
summary.results.append(r)
|
||||||
|
summary.scored += 1
|
||||||
|
|
||||||
# 汇总统计
|
# 汇总统计
|
||||||
if summary.scored > 0:
|
if summary.scored > 0:
|
||||||
|
|||||||
+83
-24
@@ -6,11 +6,17 @@
|
|||||||
- build_context: 单 agent 路径,拼接全部切片(行为与旧版一致)
|
- build_context: 单 agent 路径,拼接全部切片(行为与旧版一致)
|
||||||
|
|
||||||
multi-agent 路径由 agents/orchestrator.py 调用切片函数,每个专家只拿自己的切片。
|
multi-agent 路径由 agents/orchestrator.py 调用切片函数,每个专家只拿自己的切片。
|
||||||
|
|
||||||
|
性能说明:
|
||||||
|
build_context 创建一个共享 session 并传给所有切片函数,
|
||||||
|
避免每个切片独立创建 session —— 回测 20 场并发时,
|
||||||
|
5 个切片 × 20 场 = 100 个连接会耗尽连接池(pool_size=15)。
|
||||||
"""
|
"""
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.orm import selectinload
|
from sqlalchemy.orm import selectinload
|
||||||
@@ -18,6 +24,9 @@ from sqlalchemy.orm import selectinload
|
|||||||
from src.db.base import AsyncSessionLocal
|
from src.db.base import AsyncSessionLocal
|
||||||
from src.db.models import Match
|
from src.db.models import Match
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@@ -80,11 +89,19 @@ class MatchHeader:
|
|||||||
league_id: int
|
league_id: int
|
||||||
|
|
||||||
|
|
||||||
async def load_match_header(match_id: int) -> MatchHeader:
|
async def load_match_header(match_id: int, db: AsyncSession | None = None) -> MatchHeader:
|
||||||
"""加载比赛头信息(各 agent 共用)。"""
|
"""加载比赛头信息(各 agent 共用)。
|
||||||
async with AsyncSessionLocal() as db:
|
|
||||||
|
Args:
|
||||||
|
match_id: 比赛 ID
|
||||||
|
db: 可选的共享 session。不传则自建(向后兼容)。
|
||||||
|
"""
|
||||||
|
if db is not None:
|
||||||
m = await _load_match(db, match_id)
|
m = await _load_match(db, match_id)
|
||||||
return _to_header(m)
|
return _to_header(m)
|
||||||
|
async with AsyncSessionLocal() as new_db:
|
||||||
|
m = await _load_match(new_db, match_id)
|
||||||
|
return _to_header(m)
|
||||||
|
|
||||||
|
|
||||||
def _to_header(m: Match) -> MatchHeader:
|
def _to_header(m: Match) -> MatchHeader:
|
||||||
@@ -114,10 +131,16 @@ def header_text(h: MatchHeader) -> str:
|
|||||||
# 切片函数: 每个领域 agent 一个
|
# 切片函数: 每个领域 agent 一个
|
||||||
# ============================================================
|
# ============================================================
|
||||||
|
|
||||||
async def h2h_slice(header: MatchHeader, *, limit: int = 8, before=None) -> SliceResult:
|
async def h2h_slice(header: MatchHeader, *, limit: int = 8, before=None, db: AsyncSession | None = None) -> SliceResult:
|
||||||
"""E - 历史交锋切片: 过去数年 + 近期交手数据,提取交手规律。before=match_date 用于回测。"""
|
"""E - 历史交锋切片: 过去数年 + 近期交手数据,提取交手规律。before=match_date 用于回测。
|
||||||
async with AsyncSessionLocal() as db:
|
|
||||||
|
db: 可选共享 session,避免每个切片独立建连(见模块 docstring)。
|
||||||
|
"""
|
||||||
|
if db is not None:
|
||||||
h2h = await _get_h2h(db, header.home_team_id, header.away_team_id, before=before, limit=limit)
|
h2h = await _get_h2h(db, header.home_team_id, header.away_team_id, before=before, limit=limit)
|
||||||
|
else:
|
||||||
|
async with AsyncSessionLocal() as new_db:
|
||||||
|
h2h = await _get_h2h(new_db, header.home_team_id, header.away_team_id, before=before, limit=limit)
|
||||||
lines = [f"── 历史交锋(近 {limit} 次) ──"]
|
lines = [f"── 历史交锋(近 {limit} 次) ──"]
|
||||||
n_with_score = 0
|
n_with_score = 0
|
||||||
if h2h:
|
if h2h:
|
||||||
@@ -141,11 +164,18 @@ async def h2h_slice(header: MatchHeader, *, limit: int = 8, before=None) -> Slic
|
|||||||
return SliceResult(text="\n".join(lines), has_data=n_with_score > 0, n_records=n_with_score)
|
return SliceResult(text="\n".join(lines), has_data=n_with_score > 0, n_records=n_with_score)
|
||||||
|
|
||||||
|
|
||||||
async def form_slice(header: MatchHeader, *, limit: int = 5, before=None) -> SliceResult:
|
async def form_slice(header: MatchHeader, *, limit: int = 5, before=None, db: AsyncSession | None = None) -> SliceResult:
|
||||||
"""A - 近期状态切片: 两队近 N 场赛果、关键事件、走势判断。before=match_date 用于回测。"""
|
"""A - 近期状态切片: 两队近 N 场赛果、关键事件、走势判断。before=match_date 用于回测。
|
||||||
async with AsyncSessionLocal() as db:
|
|
||||||
|
db: 可选共享 session,避免每个切片独立建连(见模块 docstring)。
|
||||||
|
"""
|
||||||
|
if db is not None:
|
||||||
home_form = await _get_form(db, header.home_team_id, before=before, limit=limit)
|
home_form = await _get_form(db, header.home_team_id, before=before, limit=limit)
|
||||||
away_form = await _get_form(db, header.away_team_id, before=before, limit=limit)
|
away_form = await _get_form(db, header.away_team_id, before=before, limit=limit)
|
||||||
|
else:
|
||||||
|
async with AsyncSessionLocal() as new_db:
|
||||||
|
home_form = await _get_form(new_db, header.home_team_id, before=before, limit=limit)
|
||||||
|
away_form = await _get_form(new_db, header.away_team_id, before=before, limit=limit)
|
||||||
lines = []
|
lines = []
|
||||||
n_scored = 0
|
n_scored = 0
|
||||||
for label, name, form, side in (
|
for label, name, form, side in (
|
||||||
@@ -175,11 +205,18 @@ async def form_slice(header: MatchHeader, *, limit: int = 5, before=None) -> Sli
|
|||||||
return SliceResult(text="\n".join(lines), has_data=n_scored > 0, n_records=n_scored)
|
return SliceResult(text="\n".join(lines), has_data=n_scored > 0, n_records=n_scored)
|
||||||
|
|
||||||
|
|
||||||
async def stats_slice(header: MatchHeader, *, limit: int = 10, before=None) -> SliceResult:
|
async def stats_slice(header: MatchHeader, *, limit: int = 10, before=None, db: AsyncSession | None = None) -> SliceResult:
|
||||||
"""B - 攻防数据切片: 进球、射门、控球,评估攻防强度。before=match_date 用于回测。"""
|
"""B - 攻防数据切片: 进球、射门、控球,评估攻防强度。before=match_date 用于回测。
|
||||||
async with AsyncSessionLocal() as db:
|
|
||||||
|
db: 可选共享 session,避免每个切片独立建连(见模块 docstring)。
|
||||||
|
"""
|
||||||
|
if db is not None:
|
||||||
home_form = await _get_form(db, header.home_team_id, before=before, limit=limit)
|
home_form = await _get_form(db, header.home_team_id, before=before, limit=limit)
|
||||||
away_form = await _get_form(db, header.away_team_id, before=before, limit=limit)
|
away_form = await _get_form(db, header.away_team_id, before=before, limit=limit)
|
||||||
|
else:
|
||||||
|
async with AsyncSessionLocal() as new_db:
|
||||||
|
home_form = await _get_form(new_db, header.home_team_id, before=before, limit=limit)
|
||||||
|
away_form = await _get_form(new_db, header.away_team_id, before=before, limit=limit)
|
||||||
lines = [f"── 攻防数据(近 {limit} 场) ──"]
|
lines = [f"── 攻防数据(近 {limit} 场) ──"]
|
||||||
n_total = 0
|
n_total = 0
|
||||||
for label, name, form, side in (
|
for label, name, form, side in (
|
||||||
@@ -221,11 +258,18 @@ async def stats_slice(header: MatchHeader, *, limit: int = 10, before=None) -> S
|
|||||||
return SliceResult(text="\n".join(lines), has_data=n_total > 0, n_records=n_total)
|
return SliceResult(text="\n".join(lines), has_data=n_total > 0, n_records=n_total)
|
||||||
|
|
||||||
|
|
||||||
async def home_away_slice(header: MatchHeader, *, limit: int = 10, before=None) -> SliceResult:
|
async def home_away_slice(header: MatchHeader, *, limit: int = 10, before=None, db: AsyncSession | None = None) -> SliceResult:
|
||||||
"""C - 主客因素切片: 主场战绩 vs 客场战绩,评估地理优势影响。before=match_date 用于回测。"""
|
"""C - 主客因素切片: 主场战绩 vs 客场战绩,评估地理优势影响。before=match_date 用于回测。
|
||||||
async with AsyncSessionLocal() as db:
|
|
||||||
|
db: 可选共享 session,避免每个切片独立建连(见模块 docstring)。
|
||||||
|
"""
|
||||||
|
if db is not None:
|
||||||
home_home = await _get_home_away(db, header.home_team_id, "home", before=before, limit=limit)
|
home_home = await _get_home_away(db, header.home_team_id, "home", before=before, limit=limit)
|
||||||
away_away = await _get_home_away(db, header.away_team_id, "away", before=before, limit=limit)
|
away_away = await _get_home_away(db, header.away_team_id, "away", before=before, limit=limit)
|
||||||
|
else:
|
||||||
|
async with AsyncSessionLocal() as new_db:
|
||||||
|
home_home = await _get_home_away(new_db, header.home_team_id, "home", before=before, limit=limit)
|
||||||
|
away_away = await _get_home_away(new_db, header.away_team_id, "away", before=before, limit=limit)
|
||||||
lines = ["── 主客因素 ──"]
|
lines = ["── 主客因素 ──"]
|
||||||
n_total = 0
|
n_total = 0
|
||||||
for label, name, matches, side in (
|
for label, name, matches, side in (
|
||||||
@@ -255,17 +299,22 @@ async def home_away_slice(header: MatchHeader, *, limit: int = 10, before=None)
|
|||||||
return SliceResult(text="\n".join(lines), has_data=n_total > 0, n_records=n_total)
|
return SliceResult(text="\n".join(lines), has_data=n_total > 0, n_records=n_total)
|
||||||
|
|
||||||
|
|
||||||
async def injuries_slice(header: MatchHeader, *, before=None) -> SliceResult:
|
async def injuries_slice(header: MatchHeader, *, before=None, db: AsyncSession | None = None) -> SliceResult:
|
||||||
"""D - 阵容完整性切片: 伤停与停赛名单,评估战力缺失程度。
|
"""D - 阵容完整性切片: 伤停与停赛名单,评估战力缺失程度。
|
||||||
|
|
||||||
before=cutoff: 只使用 cutoff 之前已采集的伤停数据,防回测泄漏。
|
before=cutoff: 只使用 cutoff 之前已采集的伤停数据,防回测泄漏。
|
||||||
|
db: 可选共享 session(见模块 docstring)。
|
||||||
"""
|
"""
|
||||||
from src.data.injuries import get_injuries_for_match
|
from src.data.injuries import get_injuries_for_match
|
||||||
|
|
||||||
cutoff = before or header.match_dt
|
cutoff = before or header.match_dt
|
||||||
async with AsyncSessionLocal() as db:
|
if db is not None:
|
||||||
home_injuries = await get_injuries_for_match(db, header.home_team_id, cutoff, as_of=cutoff)
|
home_injuries = await get_injuries_for_match(db, header.home_team_id, cutoff, as_of=cutoff)
|
||||||
away_injuries = await get_injuries_for_match(db, header.away_team_id, cutoff, as_of=cutoff)
|
away_injuries = await get_injuries_for_match(db, header.away_team_id, cutoff, as_of=cutoff)
|
||||||
|
else:
|
||||||
|
async with AsyncSessionLocal() as new_db:
|
||||||
|
home_injuries = await get_injuries_for_match(new_db, header.home_team_id, cutoff, as_of=cutoff)
|
||||||
|
away_injuries = await get_injuries_for_match(new_db, header.away_team_id, cutoff, as_of=cutoff)
|
||||||
|
|
||||||
lines = ["── 阵容完整性 ──"]
|
lines = ["── 阵容完整性 ──"]
|
||||||
n_records = 0
|
n_records = 0
|
||||||
@@ -291,32 +340,42 @@ async def injuries_slice(header: MatchHeader, *, before=None) -> SliceResult:
|
|||||||
# 单 agent 路径: 拼接全部切片(行为与旧版一致)
|
# 单 agent 路径: 拼接全部切片(行为与旧版一致)
|
||||||
# ============================================================
|
# ============================================================
|
||||||
|
|
||||||
async def build_context(match_id: int, *, form_last: int = 5, h2h_last: int = 5) -> MatchContext:
|
async def build_context(match_id: int, *, form_last: int = 5, h2h_last: int = 5, backtest: bool = False) -> MatchContext:
|
||||||
"""单 agent 路径的完整上下文: 拼接全部切片(before=比赛时间,防未来信息)。
|
"""单 agent 路径的完整上下文: 拼接全部切片(before=比赛时间,防未来信息)。
|
||||||
|
|
||||||
has_stats / has_injuries 直接取切片显式声明的 has_data,
|
has_stats / has_injuries 直接取切片显式声明的 has_data,
|
||||||
不再靠文案子串匹配(见审查报告 P2-1)。
|
不再靠文案子串匹配(见审查报告 P2-1)。
|
||||||
|
|
||||||
|
P2-6: backtest=True 时 cutoff = match_date - 1天,确保只用赛前数据。
|
||||||
|
|
||||||
|
P1-1: 使用单个共享 session 贯穿所有切片查询,避免连接池耗尽。
|
||||||
"""
|
"""
|
||||||
header = await load_match_header(match_id)
|
async with AsyncSessionLocal() as db:
|
||||||
|
header = await load_match_header(match_id, db=db)
|
||||||
|
# P2-6: 回测模式下 cutoff 提前 1 天,防止比赛日数据泄漏
|
||||||
|
cutoff = header.match_dt
|
||||||
|
if backtest and header.match_dt:
|
||||||
|
from datetime import timedelta
|
||||||
|
cutoff = header.match_dt - timedelta(days=1)
|
||||||
parts = [header_text(header), ""]
|
parts = [header_text(header), ""]
|
||||||
|
|
||||||
form_res = await form_slice(header, limit=form_last, before=header.match_dt)
|
form_res = await form_slice(header, limit=form_last, before=cutoff, db=db)
|
||||||
parts.append(form_res.text)
|
parts.append(form_res.text)
|
||||||
parts.append("")
|
parts.append("")
|
||||||
|
|
||||||
h2h_res = await h2h_slice(header, limit=h2h_last, before=header.match_dt)
|
h2h_res = await h2h_slice(header, limit=h2h_last, before=cutoff, db=db)
|
||||||
parts.append(h2h_res.text)
|
parts.append(h2h_res.text)
|
||||||
parts.append("")
|
parts.append("")
|
||||||
|
|
||||||
stats_res = await stats_slice(header, before=header.match_dt)
|
stats_res = await stats_slice(header, before=cutoff, db=db)
|
||||||
parts.append(stats_res.text)
|
parts.append(stats_res.text)
|
||||||
parts.append("")
|
parts.append("")
|
||||||
|
|
||||||
home_away_res = await home_away_slice(header, before=header.match_dt)
|
home_away_res = await home_away_slice(header, before=cutoff, db=db)
|
||||||
parts.append(home_away_res.text)
|
parts.append(home_away_res.text)
|
||||||
parts.append("")
|
parts.append("")
|
||||||
|
|
||||||
injuries_res = await injuries_slice(header, before=header.match_dt)
|
injuries_res = await injuries_slice(header, before=cutoff, db=db)
|
||||||
parts.append(injuries_res.text)
|
parts.append(injuries_res.text)
|
||||||
|
|
||||||
return MatchContext(
|
return MatchContext(
|
||||||
|
|||||||
+14
-8
@@ -23,8 +23,9 @@ _PROMPT_DIR = Path(__file__).resolve().parent / "prompts"
|
|||||||
|
|
||||||
# ── LLM 响应缓存(match+provider+model+version → 结果) ──
|
# ── LLM 响应缓存(match+provider+model+version → 结果) ──
|
||||||
_CACHE_TTL_SEC = 300 # 5 分钟
|
_CACHE_TTL_SEC = 300 # 5 分钟
|
||||||
|
# P1-5: 缓存仅在 asyncio 协程内同步访问(dict 操作 GIL 原子),无需 threading.Lock。
|
||||||
|
# 删除 _cache_lock,避免同步锁阻塞事件循环;dict 的 get/set 在 CPython 下原子。
|
||||||
_cache: dict[str, tuple[float, PredictResult]] = {}
|
_cache: dict[str, tuple[float, PredictResult]] = {}
|
||||||
_cache_lock = Lock()
|
|
||||||
|
|
||||||
|
|
||||||
def _cache_key(match_id: int, provider: str, model: str, version: str, tpl_hash: str) -> str:
|
def _cache_key(match_id: int, provider: str, model: str, version: str, tpl_hash: str) -> str:
|
||||||
@@ -38,19 +39,20 @@ def _cache_key(match_id: int, provider: str, model: str, version: str, tpl_hash:
|
|||||||
|
|
||||||
|
|
||||||
def _get_cached(match_id: int, provider: str, model: str, version: str, tpl_hash: str) -> PredictResult | None:
|
def _get_cached(match_id: int, provider: str, model: str, version: str, tpl_hash: str) -> PredictResult | None:
|
||||||
|
# P1-5: 无锁访问。dict get/del 在 CPython GIL 下原子,且无 await 穿插。
|
||||||
key = _cache_key(match_id, provider, model, version, tpl_hash)
|
key = _cache_key(match_id, provider, model, version, tpl_hash)
|
||||||
with _cache_lock:
|
entry = _cache.get(key)
|
||||||
if key in _cache:
|
if entry is not None:
|
||||||
ts, result = _cache[key]
|
ts, result = entry
|
||||||
if time.time() - ts < _CACHE_TTL_SEC:
|
if time.time() - ts < _CACHE_TTL_SEC:
|
||||||
return result
|
return result
|
||||||
del _cache[key]
|
_cache.pop(key, None)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def _set_cached(match_id: int, provider: str, model: str, version: str, tpl_hash: str, result: PredictResult) -> None:
|
def _set_cached(match_id: int, provider: str, model: str, version: str, tpl_hash: str, result: PredictResult) -> None:
|
||||||
|
# P1-5: 无锁写入。同上,dict set 原子。
|
||||||
key = _cache_key(match_id, provider, model, version, tpl_hash)
|
key = _cache_key(match_id, provider, model, version, tpl_hash)
|
||||||
with _cache_lock:
|
|
||||||
_cache[key] = (time.time(), result)
|
_cache[key] = (time.time(), result)
|
||||||
|
|
||||||
|
|
||||||
@@ -103,6 +105,7 @@ async def predict_match(
|
|||||||
prompt_version: str | None = None,
|
prompt_version: str | None = None,
|
||||||
mode: str = "multi",
|
mode: str = "multi",
|
||||||
use_cache: bool = True,
|
use_cache: bool = True,
|
||||||
|
backtest: bool = False,
|
||||||
) -> "PredictResult | MultiPredictResult":
|
) -> "PredictResult | MultiPredictResult":
|
||||||
"""预测入口。mode=multi(默认)走多 agent;mode=single 走单次调用。
|
"""预测入口。mode=multi(默认)走多 agent;mode=single 走单次调用。
|
||||||
|
|
||||||
@@ -110,6 +113,7 @@ async def predict_match(
|
|||||||
use_cache: 是否允许返回进程内缓存结果。回测必须传 False——
|
use_cache: 是否允许返回进程内缓存结果。回测必须传 False——
|
||||||
缓存命中不会新建 prediction 行,调用方会对同一个 prediction_id
|
缓存命中不会新建 prediction 行,调用方会对同一个 prediction_id
|
||||||
反复 settle,把不同比赛的真实比分覆盖到同一条记录上。
|
反复 settle,把不同比赛的真实比分覆盖到同一条记录上。
|
||||||
|
backtest: 是否回测模式。True 时 build_context 使用 match_date-1天 作为 cutoff。
|
||||||
"""
|
"""
|
||||||
if mode == "single":
|
if mode == "single":
|
||||||
return await _predict_single(
|
return await _predict_single(
|
||||||
@@ -118,6 +122,7 @@ async def predict_match(
|
|||||||
model=model,
|
model=model,
|
||||||
prompt_version=prompt_version,
|
prompt_version=prompt_version,
|
||||||
use_cache=use_cache,
|
use_cache=use_cache,
|
||||||
|
backtest=backtest,
|
||||||
)
|
)
|
||||||
from src.llm.agents.orchestrator import predict_match_multi
|
from src.llm.agents.orchestrator import predict_match_multi
|
||||||
|
|
||||||
@@ -131,6 +136,7 @@ async def _predict_single(
|
|||||||
model: str | None = None,
|
model: str | None = None,
|
||||||
prompt_version: str | None = None,
|
prompt_version: str | None = None,
|
||||||
use_cache: bool = True,
|
use_cache: bool = True,
|
||||||
|
backtest: bool = False,
|
||||||
) -> PredictResult:
|
) -> PredictResult:
|
||||||
"""单次调用路径(原有实现)。"""
|
"""单次调用路径(原有实现)。"""
|
||||||
if provider is None:
|
if provider is None:
|
||||||
@@ -147,8 +153,8 @@ async def _predict_single(
|
|||||||
logger.debug("predict cache hit match=%s", match_id)
|
logger.debug("predict cache hit match=%s", match_id)
|
||||||
return cached
|
return cached
|
||||||
|
|
||||||
# 1. 拼上下文
|
# 1. 拼上下文(P2-6: backtest 时使用 match_date-1天 作为 cutoff)
|
||||||
ctx = await build_context(match_id)
|
ctx = await build_context(match_id, backtest=backtest)
|
||||||
|
|
||||||
# 1.5 计算快照元数据(用于可复现性)
|
# 1.5 计算快照元数据(用于可复现性)
|
||||||
now = datetime.now(timezone.utc)
|
now = datetime.now(timezone.utc)
|
||||||
|
|||||||
@@ -0,0 +1,19 @@
|
|||||||
|
"""LLM 模块共享工具函数。"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
|
||||||
|
def actual_1x2(home: int, away: int) -> str:
|
||||||
|
"""实际比分 → 胜平负。
|
||||||
|
|
||||||
|
单一权威源: backtest.py 和 eval.py 共用,避免重复定义。
|
||||||
|
"""
|
||||||
|
if home > away:
|
||||||
|
return "1"
|
||||||
|
if home == away:
|
||||||
|
return "X"
|
||||||
|
return "2"
|
||||||
|
|
||||||
|
|
||||||
|
def is_correct_1x2(pred: str | None, actual: str) -> bool:
|
||||||
|
"""预测是否命中胜平负。"""
|
||||||
|
return pred == actual
|
||||||
@@ -10,8 +10,10 @@ from pydantic import BaseModel, Field, field_validator, model_validator
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# 已知的 5 个专家 agent 名(与 orchestrator.SPECIALIST_SPECS 保持一致)
|
# 单一权威源:从 orchestrator.SPECIALIST_SPECS 派生,避免两端独立定义导致静默偏离
|
||||||
KNOWN_AGENT_NAMES: tuple[str, ...] = ("form", "stats", "home_away", "injuries", "h2h")
|
from src.llm.agents.orchestrator import SPECIALIST_SPECS
|
||||||
|
|
||||||
|
KNOWN_AGENT_NAMES: tuple[str, ...] = tuple(spec.name for spec in SPECIALIST_SPECS)
|
||||||
|
|
||||||
|
|
||||||
class AgentReportSchema(BaseModel):
|
class AgentReportSchema(BaseModel):
|
||||||
@@ -74,7 +76,7 @@ class PredictionOutputSchema(BaseModel):
|
|||||||
"""
|
"""
|
||||||
expected = _score_to_1x2(self.pred_home_goals, self.pred_away_goals)
|
expected = _score_to_1x2(self.pred_home_goals, self.pred_away_goals)
|
||||||
if self.pred_1x2 != expected:
|
if self.pred_1x2 != expected:
|
||||||
logger.warning(
|
logger.debug(
|
||||||
"1x2 与比分不一致: 比分 %.1f-%.1f 推出 '%s',但 LLM 给出 '%s';以比分修正",
|
"1x2 与比分不一致: 比分 %.1f-%.1f 推出 '%s',但 LLM 给出 '%s';以比分修正",
|
||||||
self.pred_home_goals, self.pred_away_goals, expected, self.pred_1x2,
|
self.pred_home_goals, self.pred_away_goals, expected, self.pred_1x2,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user