diff --git a/src/api/routes/ingest.py b/src/api/routes/ingest.py index 5b44cf5..61f0186 100644 --- a/src/api/routes/ingest.py +++ b/src/api/routes/ingest.py @@ -16,15 +16,14 @@ async def ingest_bzzoiro_route(req: IngestBzzoiroRequest): """触发 bzzoiro 采集。""" source = get_source("bzzoiro") try: - async with get_uow() as uow: + async with get_uow() as session: result = await source.ingest( - uow.session, + session, leagues=req.leagues, date_from=req.date_from, date_to=req.date_to, status=req.status, ) - await uow.commit() return IngestResponse(**result) except Exception as e: raise HTTPException(500, str(e)) @@ -35,9 +34,8 @@ async def ingest_understat_route(req: IngestUnderstatRequest): """触发 understat xG 回填。""" source = get_source("understat") try: - async with get_uow() as uow: - result = await source.ingest(uow.session, league=req.league, season=req.season) - await uow.commit() + async with get_uow() as session: + result = await source.ingest(session, league=req.league, season=req.season) return IngestSimpleResponse(**result) except Exception as e: raise HTTPException(500, str(e)) @@ -47,9 +45,8 @@ async def ingest_understat_route(req: IngestUnderstatRequest): async def ingest_injuries_route(req: IngestInjuriesRequest): """触发伤停采集。""" try: - async with get_uow() as uow: - result = await ingest_injuries(uow.session, date=req.date) - await uow.commit() + async with get_uow() as session: + result = await ingest_injuries(session, date=req.date) return IngestSimpleResponse(**result) except Exception as e: raise HTTPException(500, str(e)) diff --git a/src/db/repositories.py b/src/db/repositories.py index 3dcd144..9d08e16 100644 --- a/src/db/repositories.py +++ b/src/db/repositories.py @@ -5,7 +5,7 @@ Repository 只负责查询,不负责事务提交。 """ from __future__ import annotations -from sqlalchemy import select +from sqlalchemy import func, select from sqlalchemy.orm import selectinload from sqlalchemy.ext.asyncio import AsyncSession @@ -38,8 +38,6 @@ class MatchRepository: self, league_id: int, home_team_id: int, away_team_id: int, date ) -> Match | None: """按联赛+主队+客队+日期查找比赛(天级匹配)。""" - from sqlalchemy import func - if hasattr(date, "date"): date = date.date() diff --git a/src/db/unit_of_work.py b/src/db/unit_of_work.py index 1939205..12775ce 100644 --- a/src/db/unit_of_work.py +++ b/src/db/unit_of_work.py @@ -1,63 +1,33 @@ """工作单元(Unit of Work):统一事务边界。 使用方式: - async with UnitOfWork(db) as uow: - await uow.matches.get_by_id(1) - await uow.matches.add(new_match) - # 退出时自动 commit,异常时 rollback + async with get_uow() as uow: + await uow.session.get(Match, 1) + await uow.commit() """ from __future__ import annotations from collections.abc import AsyncIterator from contextlib import asynccontextmanager -from typing import AsyncGenerator - -from sqlalchemy.ext.asyncio import AsyncSession from src.db.base import AsyncSessionLocal -class UnitOfWork: - """工作单元:封装事务边界。""" - - def __init__(self, session: AsyncSession) -> None: - self._session = session - self.committed = False - - @property - def session(self) -> AsyncSession: - return self._session - - async def commit(self) -> None: - await self._session.commit() - self.committed = True - - async def rollback(self) -> None: - await self._session.rollback() - - async def close(self) -> None: - await self._session.close() - - async def __aenter__(self) -> "UnitOfWork": - return self - - async def __aexit__(self, exc_type, exc_val, exc_tb) -> None: - if exc_type is not None: - await self.rollback() - await self.close() - - @asynccontextmanager -async def get_uow() -> AsyncGenerator[UnitOfWork, None]: - """创建新的工作单元(用于非路由上下文)。""" +async def get_uow() -> AsyncIterator[AsyncSessionLocal]: + """创建新的工作单元(用于非路由上下文)。 + + 用法: + async with get_uow() as session: + await session.get(...) + # 退出时自动 commit(无异常) 或 rollback(有异常) + """ session = AsyncSessionLocal() - uow = UnitOfWork(session) try: - yield uow - if not uow.committed: - await uow.commit() + yield session + await session.commit() except Exception: - await uow.rollback() + await session.rollback() raise finally: - await uow.close() + await session.close() diff --git a/src/llm/agents/orchestrator.py b/src/llm/agents/orchestrator.py index f261d89..cffdc5c 100644 --- a/src/llm/agents/orchestrator.py +++ b/src/llm/agents/orchestrator.py @@ -186,9 +186,8 @@ async def predict_match_multi( ).hexdigest() # 4. 存库(使用 UnitOfWork) - async with get_uow() as uow: - db = uow.session - m = await db.get(Match, match_id) + async with get_uow() as session: + m = await session.get(Match, match_id) if m is None: raise ValueError(f"match {match_id} not found") @@ -219,9 +218,8 @@ async def predict_match_multi( cutoff_at=cutoff_at, input_hash=input_hash, ) - db.add(pred) - await uow.commit() - await db.refresh(pred) + session.add(pred) + await session.refresh(pred) return MultiPredictResult( prediction_id=pred.id, diff --git a/src/llm/backtest.py b/src/llm/backtest.py index aa62a35..25fb45e 100644 --- a/src/llm/backtest.py +++ b/src/llm/backtest.py @@ -108,9 +108,9 @@ async def run_backtest( Returns: BacktestSummary 含逐场结果 + 汇总统计 """ - async with get_uow() as uow: + async with get_uow() as session: matches = await _get_historical_matches( - uow.session, league_id=league_id, date_from=date_from, date_to=date_to, limit=limit + session, league_id=league_id, date_from=date_from, date_to=date_to, limit=limit ) summary = BacktestSummary(total=len(matches), scored=0) diff --git a/src/llm/eval.py b/src/llm/eval.py index 1f90522..8f8b63b 100644 --- a/src/llm/eval.py +++ b/src/llm/eval.py @@ -13,14 +13,13 @@ logger = logging.getLogger(__name__) async def settle_prediction(prediction_id: int, home_goals: int, away_goals: int) -> Prediction: """回填实际结果。""" - async with get_uow() as uow: - pred = await uow.session.get(Prediction, prediction_id) + async with get_uow() as session: + pred = await session.get(Prediction, prediction_id) if pred is None: raise ValueError(f"prediction {prediction_id} not found") pred.actual_home_goals = home_goals pred.actual_away_goals = away_goals pred.settled = True - await uow.commit() return pred @@ -35,12 +34,12 @@ def _actual_1x2(home: int, away: int) -> str: async def get_eval_summary() -> dict: """按 provider × 模型聚合评估。""" - async with get_uow() as uow: + async with get_uow() as session: stmt = ( select(Prediction) .where(Prediction.settled == True) ) - result = await uow.session.execute(stmt) + result = await session.execute(stmt) rows = list(result.scalars().all()) from collections import defaultdict diff --git a/src/llm/predict.py b/src/llm/predict.py index 1f41c70..c950e70 100644 --- a/src/llm/predict.py +++ b/src/llm/predict.py @@ -145,10 +145,9 @@ async def _predict_single( raise RuntimeError(f"LLM 输出校验失败: {e}") # 4. 存预测(使用 UnitOfWork 统一事务) - async with get_uow() as uow: - db = uow.session + async with get_uow() as session: # 验证 match 存在 - m = await db.get(Match, match_id) + m = await session.get(Match, match_id) if m is None: raise ValueError(f"match {match_id} not found") @@ -169,9 +168,8 @@ async def _predict_single( cutoff_at=cutoff_at, input_hash=input_hash, ) - db.add(pred) - await db.commit() - await db.refresh(pred) + session.add(pred) + await session.refresh(pred) result = PredictResult( prediction_id=pred.id, @@ -190,5 +188,4 @@ async def _predict_single( # 5. 写入缓存 _set_cached(match_id, settings.LLM_PROVIDER, provider.model, version, result) - await uow.commit() return result diff --git a/src/llm/validation.py b/src/llm/validation.py index 1a10539..7472571 100644 --- a/src/llm/validation.py +++ b/src/llm/validation.py @@ -6,8 +6,6 @@ from __future__ import annotations from pydantic import BaseModel, Field, field_validator, model_validator -from src.db.models import Prediction - class AgentReportSchema(BaseModel): """单个专家 Agent 输出的校验 schema。""" @@ -62,23 +60,20 @@ class PredictionOutputSchema(BaseModel): @model_validator(mode="after") def check_consistency(self) -> "PredictionOutputSchema": - """验证比分与胜平负一致。""" + """验证比分与胜平负一致,不一致则自动修正。""" expected = _score_to_1x2(self.pred_home_goals, self.pred_away_goals) - if expected and self.pred_1x2 != expected: - # 自动修正而非拒绝(LLM 常见小错误) + if self.pred_1x2 != expected: self.pred_1x2 = expected return self -def _score_to_1x2(home: float, away: float) -> str | None: +def _score_to_1x2(home: float, away: float) -> str: """从比分推导胜平负。""" if home > away: return "1" - if home == away: - return "X" if home < away: return "2" - return None + return "X" def validate_agent_output(raw: dict) -> AgentReportSchema: