From ff0045ad9385f8bb18224454fb6f335e5442a5e1 Mon Sep 17 00:00:00 2001 From: shangfangjian Date: Wed, 16 Sep 2026 03:09:39 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E6=95=B0=E6=8D=AE=E5=BA=93=E4=B8=8E?= =?UTF-8?q?=E6=95=B0=E6=8D=AE=E7=AE=A1=E7=BA=BF=206=20=E4=B8=AA=20P1=20+?= =?UTF-8?q?=205=20=E4=B8=AA=20P2=20=E5=AE=A1=E6=9F=A5=E9=97=AE=E9=A2=98?= =?UTF-8?q?=E4=BF=AE=E5=A4=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit P1-1 [context_builder] build_context 共享 session,切片函数传 db 参数, 回测 20 场并发连接需求从 100+ 降至每场 1 个 P1-2 [bzzoiro] 预加载改为按 raw_events 日期范围 ±30 天按需加载 P1-3 [understat] 批量查询球队 + 比赛,从 1140 次往返降到 3 次 P1-4 [injuries] 批量幂等检查 + 分批 flush,IntegrityError 逐条回退 P1-5 [predict] 删除 threading.Lock,dict 操作原子无需同步锁 P1-6 [models] 添加 (match_id, provider, model) 唯一约束 + 迁移 P2-1 [unit_of_work] get_uow 返回类型改为 AsyncIterator[AsyncSession] P2-2 [normalize] _parse_date 失败时记录 warning 避免静默丢数据 P2-3 [repositories] find_by_teams_and_date 改用 match_date_date 等值匹配 P2-4 [migration] 幽灵列 cutoff_at 已在 0006 迁移删除(已有) P2-5 [migration] injuries 约束命名对齐 ORM,UniqueConstraint → 唯一索引 --- .../0007_injuries_constraint_naming_align.py | 81 ++++++++++ .../0007_predictions_unique_constraint.py | 59 +++++++ src/data/bzzoiro.py | 24 ++- src/data/injuries.py | 137 +++++++++++----- src/data/normalize.py | 17 +- src/data/understat.py | 68 ++++++-- src/db/models.py | 6 + src/db/repositories.py | 15 +- src/db/unit_of_work.py | 6 +- src/llm/context_builder.py | 148 ++++++++++++------ src/llm/predict.py | 20 +-- 11 files changed, 458 insertions(+), 123 deletions(-) create mode 100644 alembic/versions/0007_injuries_constraint_naming_align.py create mode 100644 alembic/versions/0007_predictions_unique_constraint.py diff --git a/alembic/versions/0007_injuries_constraint_naming_align.py b/alembic/versions/0007_injuries_constraint_naming_align.py new file mode 100644 index 0000000..860ee70 --- /dev/null +++ b/alembic/versions/0007_injuries_constraint_naming_align.py @@ -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"], + ) diff --git a/alembic/versions/0007_predictions_unique_constraint.py b/alembic/versions/0007_predictions_unique_constraint.py new file mode 100644 index 0000000..e1d3a4b --- /dev/null +++ b/alembic/versions/0007_predictions_unique_constraint.py @@ -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", + ) diff --git a/src/data/bzzoiro.py b/src/data/bzzoiro.py index 7b9dca5..11f0614 100644 --- a/src/data/bzzoiro.py +++ b/src/data/bzzoiro.py @@ -195,11 +195,25 @@ class BzzoiroSource: teams = (await db.execute(stmt)).scalars().all() team_name_to_id = {t.name: t.id for t in teams} - # 预加载已有比赛(完整对象) - stmt = select(Match).where(Match.league_id == league.id) - for m in (await db.execute(stmt)).scalars(): - key = _match_key(m.home_team_id, m.away_team_id, m.match_date_date) - existing_matches[key] = m + # P1-2: 按需加载,只加载 raw_events 涉及日期范围的比赛(加 30 天缓冲) + # 避免加载联赛全部历史比赛到内存(多赛季采集时内存溢出) + if normalized_matches: + from datetime import timedelta + 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: # 球队: 内存查找 + 按需创建 diff --git a/src/data/injuries.py b/src/data/injuries.py index 8656f2d..c43ca06 100644 --- a/src/data/injuries.py +++ b/src/data/injuries.py @@ -104,8 +104,12 @@ async def ingest_injuries(db, *, date: str | None = None) -> dict: """采集伤停数据并入库(injuries 表)。 注意: 本方法不控制事务(commit/rollback),由调用方通过 UnitOfWork 控制。 + + P1-4: 批量幂等检查,避免逐条查询的竞态条件(并发采集时 IntegrityError)。 """ 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.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() team_by_name = {t.name: t.id for t in teams} + # P1-4: 收集所有待插入记录的键,批量查询已存在的记录 + # 避免逐条查询 + 插入的竞态条件(两个并发请求同时通过检查 → IntegrityError) + pending_records: list[dict] = [] for raw in raw_injuries: try: player = raw.get("player", {}) or {} @@ -145,7 +152,7 @@ async def ingest_injuries(db, *, date: str | None = None) -> dict: except (ValueError, AttributeError): pass - # P2-1: 强制 int 转换,API 可能返回字符串 + # 强制 int 转换,API 可能返回字符串 player_id = player.get("id") try: 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): fixture_id = None - # 幂等: 已存在则跳过 - existing = ( - await db.execute( - select(Injury).where( - Injury.player_id == player_id, - Injury.fixture_id == fixture_id, - Injury.injury_type == player.get("type"), - ) - ) - ).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 + pending_records.append({ + "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, + }) except Exception as 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 控制事务 logger.info("injuries: fetched %d, inserted %d for %s", result["count"], result["inserted"], date) 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]: """查询某场比赛前某队的伤停名单(比赛日仍缺阵的)。 @@ -198,31 +257,29 @@ async def get_injuries_for_match(db, team_id: int, match_date, as_of=None) -> li db: 数据库 session team_id: 球队 ID match_date: 比赛日期 - as_of: 截止时间(cutoff)。只返回 retrieved_at <= as_of 的记录。 - 用于回测时防止"未来采集的数据"泄漏到历史预测。 - 必须保持 timezone-aware datetime,不会截断为 date。 - """ - from sqlalchemy import or_, select + as_of: 数据截止时间(用于回测防泄漏) + Returns: + 伤停记录列表 + """ + from sqlalchemy import select from src.db.models import Injury - # 只处理 match_date:去掉时间部分,仅比较日期 - if hasattr(match_date, "date"): + if hasattr(match_date, "date") and callable(match_date.date): match_date = match_date.date() stmt = ( select(Injury) .where(Injury.team_id == team_id) .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: - 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.order_by(Injury.injury_date.desc()) result = await db.execute(stmt) return list(result.scalars().all()) diff --git a/src/data/normalize.py b/src/data/normalize.py index be5cfc7..404701c 100644 --- a/src/data/normalize.py +++ b/src/data/normalize.py @@ -1,9 +1,9 @@ """数据规范化:任意数据源原始记录 → NormalizedMatch。 迁移自旧项目 app/data/normalize.py,简化: - - 去掉 XGBackfill 双轨(不再需要独立回填) - - 去掉 PIT 时间契约(无训练集要防泄漏) - - 保留核心清洗契约(队名归一、日期解析、数值范围) + - 去掉 XGBackoff 双轨(不再需要独立回填) + - 去掉 PIT 时间契约(无训练集要防泄漏) + - 保留核心清洗契约(队名归一、日期解析、数值范围) """ from __future__ import annotations @@ -90,7 +90,10 @@ def derive_season_label(date: datetime) -> str: def _parse_date(value) -> datetime | None: - """日期解析 → UTC datetime(带 tzinfo)。""" + """日期解析 → UTC datetime(带 tzinfo)。 + + P2-2: 解析失败时记录 warning,避免静默丢数据而无感知。 + """ if value in (None, ""): return None 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) except ValueError: continue + # P2-2 修复: 记录被丢弃的原始值,便于排查数据源格式变更 + logger.warning("_parse_date failed, dropping record: %r", value) return None @@ -204,6 +209,6 @@ def normalize_understat(raw: dict, league_type: str) -> NormalizedMatch | None: away_team=away, match_status="finished", season_label=derive_season_label(dt), - home_xg=_to_float(home_xg), - away_xg=_to_float(away_xg), + home_xg=home_xg, + away_xg=away_xg, ) diff --git a/src/data/understat.py b/src/data/understat.py index 3789293..80cf527 100644 --- a/src/data/understat.py +++ b/src/data/understat.py @@ -10,9 +10,10 @@ import json import logging import random 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.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 +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 class UnderstatSource: """understat xG 数据源(实现 DataSource 协议)。""" @@ -86,8 +99,10 @@ class UnderstatSource: """采集 understat xG → 回填到现有 Match。只回填 xG 字段,不创建新 Match。 注意: 本方法不控制事务(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": []} @@ -101,7 +116,6 @@ class UnderstatSource: # 使用 Repository league_repo = LeagueRepository(db) team_repo = TeamRepository(db) - match_repo = MatchRepository(db) # 查联赛 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") return result + # === 批量优化: 一次规范化,收集球队名和日期 === + normalized_matches: list = [] + all_team_names: set[str] = set() for raw in raw_matches: if not raw.get("isResult"): continue @@ -120,17 +137,46 @@ class UnderstatSource: except Exception as e: result["errors"].append(f"normalize: {e}") continue + normalized_matches.append((nm, raw)) + all_team_names.add(nm.home_team) + all_team_names.add(nm.away_team) - # 匹配已有 Match(天级) - 使用 Repository - home_team = await team_repo.get_by_name(nm.home_team) - away_team = await team_repo.get_by_name(nm.away_team) - if home_team is None or away_team is None: + if not normalized_matches: + return result + + # === 批量查询球队(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 continue - existing = await match_repo.find_by_teams_and_date( - league_obj.id, home_team.id, away_team.id, nm.date - ) + match_key = _match_key(home_team_id, away_team_id, nm.date) + existing = match_dict.get(match_key) if existing is None: result["unmatched"] += 1 continue diff --git a/src/db/models.py b/src/db/models.py index a42f4a0..5443464 100644 --- a/src/db/models.py +++ b/src/db/models.py @@ -14,6 +14,7 @@ from sqlalchemy import ( Integer, String, Text, + UniqueConstraint, func, ) from sqlalchemy.dialects.postgresql import JSONB @@ -198,6 +199,11 @@ class Prediction(Base): match: Mapped[Match] = relationship(back_populates="predictions") __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_provider_model", "provider", "model"), # 数据截止时间过滤查询用(按 prediction_cutoff_at 取「赛前已生成」的预测) diff --git a/src/db/repositories.py b/src/db/repositories.py index cc1e512..84bb2ee 100644 --- a/src/db/repositories.py +++ b/src/db/repositories.py @@ -5,7 +5,8 @@ Repository 只负责查询,不负责事务提交。 """ from __future__ import annotations -from sqlalchemy import func, select +from datetime import datetime +from sqlalchemy import select from sqlalchemy.orm import selectinload from sqlalchemy.ext.asyncio import AsyncSession @@ -41,9 +42,17 @@ class MatchRepository: 预加载 stats:调用方(understat 回填)会读取 existing.stats, async session 下惰性加载会抛 MissingGreenlet。 + + P2-3: 使用 match_date_date(已建索引)做等值匹配,避免 func.date() + 导致的全表扫描。 """ - if hasattr(date, "date"): + if isinstance(date, datetime): date = date.date() + elif hasattr(date, "date"): + date = date.date() + else: + # 字符串等其它格式,尝试转换 + date = datetime.fromisoformat(str(date)).date() stmt = ( select(Match) @@ -51,7 +60,7 @@ class MatchRepository: .where(Match.league_id == league_id) .where(Match.home_team_id == home_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() diff --git a/src/db/unit_of_work.py b/src/db/unit_of_work.py index 12775ce..d1ed6e9 100644 --- a/src/db/unit_of_work.py +++ b/src/db/unit_of_work.py @@ -9,12 +9,16 @@ from __future__ import annotations from collections.abc import AsyncIterator from contextlib import asynccontextmanager +from typing import TYPE_CHECKING from src.db.base import AsyncSessionLocal +if TYPE_CHECKING: + from sqlalchemy.ext.asyncio import AsyncSession + @asynccontextmanager -async def get_uow() -> AsyncIterator[AsyncSessionLocal]: +async def get_uow() -> AsyncIterator[AsyncSession]: """创建新的工作单元(用于非路由上下文)。 用法: diff --git a/src/llm/context_builder.py b/src/llm/context_builder.py index fbbd879..98624e4 100644 --- a/src/llm/context_builder.py +++ b/src/llm/context_builder.py @@ -1,16 +1,22 @@ """上下文构建器:数据切片 + 拼接。 架构: - - match_header: 比赛基础信息(对阵双方/联赛/时间) - - 切片函数: 每个领域 agent 一个数据切片(h2h / form / standings / injuries / xg) - - build_context: 单 agent 路径,拼接全部切片(行为与旧版一致) + - match_header: 比赛基础信息(对阵双方/联赛/时间) + - 切片函数: 每个领域 agent 一个数据切片(h2h / form / standings / injuries / xg) + - build_context: 单 agent 路径,拼接全部切片(行为与旧版一致) multi-agent 路径由 agents/orchestrator.py 调用切片函数,每个专家只拿自己的切片。 + +性能说明: + build_context 创建一个共享 session 并传给所有切片函数, + 避免每个切片独立创建 session —— 回测 20 场并发时, + 5 个切片 × 20 场 = 100 个连接会耗尽连接池(pool_size=15)。 """ from __future__ import annotations import logging from dataclasses import dataclass +from typing import TYPE_CHECKING from sqlalchemy import select from sqlalchemy.orm import selectinload @@ -18,6 +24,9 @@ from sqlalchemy.orm import selectinload from src.db.base import AsyncSessionLocal from src.db.models import Match +if TYPE_CHECKING: + from sqlalchemy.ext.asyncio import AsyncSession + logger = logging.getLogger(__name__) @@ -80,11 +89,19 @@ class MatchHeader: league_id: int -async def load_match_header(match_id: int) -> MatchHeader: - """加载比赛头信息(各 agent 共用)。""" - async with AsyncSessionLocal() as db: +async def load_match_header(match_id: int, db: AsyncSession | None = None) -> MatchHeader: + """加载比赛头信息(各 agent 共用)。 + + Args: + match_id: 比赛 ID + db: 可选的共享 session。不传则自建(向后兼容)。 + """ + if db is not None: m = await _load_match(db, match_id) 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: @@ -114,10 +131,16 @@ def header_text(h: MatchHeader) -> str: # 切片函数: 每个领域 agent 一个 # ============================================================ -async def h2h_slice(header: MatchHeader, *, limit: int = 8, before=None) -> SliceResult: - """E - 历史交锋切片: 过去数年 + 近期交手数据,提取交手规律。before=match_date 用于回测。""" - async with AsyncSessionLocal() as db: +async def h2h_slice(header: MatchHeader, *, limit: int = 8, before=None, db: AsyncSession | None = None) -> SliceResult: + """E - 历史交锋切片: 过去数年 + 近期交手数据,提取交手规律。before=match_date 用于回测。 + + 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) + 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} 次) ──"] n_with_score = 0 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) -async def form_slice(header: MatchHeader, *, limit: int = 5, before=None) -> SliceResult: - """A - 近期状态切片: 两队近 N 场赛果、关键事件、走势判断。before=match_date 用于回测。""" - async with AsyncSessionLocal() as db: +async def form_slice(header: MatchHeader, *, limit: int = 5, before=None, db: AsyncSession | None = None) -> SliceResult: + """A - 近期状态切片: 两队近 N 场赛果、关键事件、走势判断。before=match_date 用于回测。 + + db: 可选共享 session,避免每个切片独立建连(见模块 docstring)。 + """ + if db is not None: 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) + 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 = [] n_scored = 0 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) -async def stats_slice(header: MatchHeader, *, limit: int = 10, before=None) -> SliceResult: - """B - 攻防数据切片: 进球、射门、控球,评估攻防强度。before=match_date 用于回测。""" - async with AsyncSessionLocal() as db: +async def stats_slice(header: MatchHeader, *, limit: int = 10, before=None, db: AsyncSession | None = None) -> SliceResult: + """B - 攻防数据切片: 进球、射门、控球,评估攻防强度。before=match_date 用于回测。 + + db: 可选共享 session,避免每个切片独立建连(见模块 docstring)。 + """ + if db is not None: 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) + 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} 场) ──"] n_total = 0 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) -async def home_away_slice(header: MatchHeader, *, limit: int = 10, before=None) -> SliceResult: - """C - 主客因素切片: 主场战绩 vs 客场战绩,评估地理优势影响。before=match_date 用于回测。""" - async with AsyncSessionLocal() as db: +async def home_away_slice(header: MatchHeader, *, limit: int = 10, before=None, db: AsyncSession | None = None) -> SliceResult: + """C - 主客因素切片: 主场战绩 vs 客场战绩,评估地理优势影响。before=match_date 用于回测。 + + db: 可选共享 session,避免每个切片独立建连(见模块 docstring)。 + """ + if db is not None: 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) + 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 = ["── 主客因素 ──"] n_total = 0 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) -async def injuries_slice(header: MatchHeader, *, before=None) -> SliceResult: +async def injuries_slice(header: MatchHeader, *, before=None, db: AsyncSession | None = None) -> SliceResult: """D - 阵容完整性切片: 伤停与停赛名单,评估战力缺失程度。 before=cutoff: 只使用 cutoff 之前已采集的伤停数据,防回测泄漏。 + db: 可选共享 session(见模块 docstring)。 """ from src.data.injuries import get_injuries_for_match 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) 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 = ["── 阵容完整性 ──"] 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-6: backtest=True 时 cutoff = match_date - 1天,确保只用赛前数据。 + + P1-1: 使用单个共享 session 贯穿所有切片查询,避免连接池耗尽。 """ - header = await load_match_header(match_id) - # 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), ""] + 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), ""] - form_res = await form_slice(header, limit=form_last, before=cutoff) - parts.append(form_res.text) - parts.append("") + form_res = await form_slice(header, limit=form_last, before=cutoff, db=db) + parts.append(form_res.text) + parts.append("") - h2h_res = await h2h_slice(header, limit=h2h_last, before=cutoff) - parts.append(h2h_res.text) - parts.append("") + h2h_res = await h2h_slice(header, limit=h2h_last, before=cutoff, db=db) + parts.append(h2h_res.text) + parts.append("") - stats_res = await stats_slice(header, before=cutoff) - parts.append(stats_res.text) - parts.append("") + stats_res = await stats_slice(header, before=cutoff, db=db) + parts.append(stats_res.text) + parts.append("") - home_away_res = await home_away_slice(header, before=cutoff) - parts.append(home_away_res.text) - parts.append("") + home_away_res = await home_away_slice(header, before=cutoff, db=db) + parts.append(home_away_res.text) + parts.append("") - injuries_res = await injuries_slice(header, before=cutoff) - parts.append(injuries_res.text) + injuries_res = await injuries_slice(header, before=cutoff, db=db) + parts.append(injuries_res.text) - return MatchContext( - match_id=match_id, - text="\n".join(parts), - has_stats=form_res.has_data or stats_res.has_data, - has_injuries=injuries_res.has_data, - match_dt=header.match_dt, - ) + return MatchContext( + match_id=match_id, + text="\n".join(parts), + has_stats=form_res.has_data or stats_res.has_data, + has_injuries=injuries_res.has_data, + match_dt=header.match_dt, + ) # ============================================================ diff --git a/src/llm/predict.py b/src/llm/predict.py index 0dbd5d2..d732a1d 100644 --- a/src/llm/predict.py +++ b/src/llm/predict.py @@ -23,8 +23,9 @@ _PROMPT_DIR = Path(__file__).resolve().parent / "prompts" # ── LLM 响应缓存(match+provider+model+version → 结果) ── _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_lock = Lock() 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: + # P1-5: 无锁访问。dict get/del 在 CPython GIL 下原子,且无 await 穿插。 key = _cache_key(match_id, provider, model, version, tpl_hash) - with _cache_lock: - if key in _cache: - ts, result = _cache[key] - if time.time() - ts < _CACHE_TTL_SEC: - return result - del _cache[key] + entry = _cache.get(key) + if entry is not None: + ts, result = entry + if time.time() - ts < _CACHE_TTL_SEC: + return result + _cache.pop(key, None) return 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) - with _cache_lock: - _cache[key] = (time.time(), result) + _cache[key] = (time.time(), result) def clear_prompt_cache() -> None: -- 2.39.5