fix: 数据库与数据管线审查问题修复 (6 P1 + 5 P2) #3

Merged
shangfangjian merged 1 commits from fix/db-pipeline-review-issues into main 2026-09-16 03:14:37 +08:00
11 changed files with 458 additions and 123 deletions
Showing only changes of commit ff0045ad93 - Show all commits
@@ -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",
)
+19 -5
View File
@@ -195,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:
# 球队: 内存查找 + 按需创建 # 球队: 内存查找 + 按需创建
+97 -40
View File
@@ -104,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
@@ -125,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 {}
@@ -145,7 +152,7 @@ async def ingest_injuries(db, *, date: str | None = None) -> dict:
except (ValueError, AttributeError): except (ValueError, AttributeError):
pass pass
# P2-1: 强制 int 转换,API 可能返回字符串 # 强制 int 转换,API 可能返回字符串
player_id = player.get("id") player_id = player.get("id")
try: try:
player_id = int(player_id) if player_id is not None else None player_id = int(player_id) if player_id is not None else None
@@ -157,40 +164,92 @@ async def ingest_injuries(db, *, date: str | None = None) -> dict:
except (ValueError, TypeError): except (ValueError, TypeError):
fixture_id = None 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]:
"""查询某场比赛前某队的伤停名单(比赛日仍缺阵的)。 """查询某场比赛前某队的伤停名单(比赛日仍缺阵的)。
@@ -198,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())
+11 -6
View File
@@ -1,9 +1,9 @@
"""数据规范化:任意数据源原始记录 → NormalizedMatch。 """数据规范化:任意数据源原始记录 → NormalizedMatch。
迁移自旧项目 app/data/normalize.py,简化: 迁移自旧项目 app/data/normalize.py,简化:
- 去掉 XGBackfill 双轨(不再需要独立回填) - 去掉 XGBackoff 双轨(不再需要独立回填)
- 去掉 PIT 时间契约(无训练集要防泄漏) - 去掉 PIT 时间契约(无训练集要防泄漏)
- 保留核心清洗契约(队名归一、日期解析、数值范围) - 保留核心清洗契约(队名归一、日期解析、数值范围)
""" """
from __future__ import annotations from __future__ import annotations
@@ -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
View File
@@ -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
+6
View File
@@ -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
View File
@@ -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()
+5 -1
View File
@@ -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]:
"""创建新的工作单元(用于非路由上下文)。 """创建新的工作单元(用于非路由上下文)。
用法: 用法:
+100 -48
View File
@@ -1,16 +1,22 @@
"""上下文构建器:数据切片 + 拼接。 """上下文构建器:数据切片 + 拼接。
架构: 架构:
- match_header: 比赛基础信息(对阵双方/联赛/时间) - match_header: 比赛基础信息(对阵双方/联赛/时间)
- 切片函数: 每个领域 agent 一个数据切片(h2h / form / standings / injuries / xg) - 切片函数: 每个领域 agent 一个数据切片(h2h / form / standings / injuries / xg)
- 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
@@ -298,41 +347,44 @@ async def build_context(match_id: int, *, form_last: int = 5, h2h_last: int = 5,
不再靠文案子串匹配(见审查报告 P2-1)。 不再靠文案子串匹配(见审查报告 P2-1)。
P2-6: backtest=True 时 cutoff = match_date - 1天,确保只用赛前数据。 P2-6: backtest=True 时 cutoff = match_date - 1天,确保只用赛前数据。
P1-1: 使用单个共享 session 贯穿所有切片查询,避免连接池耗尽。
""" """
header = await load_match_header(match_id) async with AsyncSessionLocal() as db:
# P2-6: 回测模式下 cutoff 提前 1 天,防止比赛日数据泄漏 header = await load_match_header(match_id, db=db)
cutoff = header.match_dt # P2-6: 回测模式下 cutoff 提前 1 天,防止比赛日数据泄漏
if backtest and header.match_dt: cutoff = header.match_dt
from datetime import timedelta if backtest and header.match_dt:
cutoff = header.match_dt - timedelta(days=1) from datetime import timedelta
parts = [header_text(header), ""] cutoff = header.match_dt - timedelta(days=1)
parts = [header_text(header), ""]
form_res = await form_slice(header, limit=form_last, before=cutoff) 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=cutoff) 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=cutoff) 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=cutoff) 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=cutoff) injuries_res = await injuries_slice(header, before=cutoff, db=db)
parts.append(injuries_res.text) parts.append(injuries_res.text)
return MatchContext( return MatchContext(
match_id=match_id, match_id=match_id,
text="\n".join(parts), text="\n".join(parts),
has_stats=form_res.has_data or stats_res.has_data, has_stats=form_res.has_data or stats_res.has_data,
has_injuries=injuries_res.has_data, has_injuries=injuries_res.has_data,
match_dt=header.match_dt, match_dt=header.match_dt,
) )
# ============================================================ # ============================================================
+11 -9
View File
@@ -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,20 +39,21 @@ 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)
def clear_prompt_cache() -> None: def clear_prompt_cache() -> None: