fix: 代码审查问题修复
- unit_of_work.py: 简化为纯 session 上下文管理器,消除 double-close 风险 - validation.py: 移除未使用的 import,简化 _score_to_1x2 逻辑 - repositories.py: 将 func 导入移到模块顶层 - 更新所有 get_uow() 调用点使用新接口(yield session 而非 uow 对象)
This commit is contained in:
@@ -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,
|
||||
|
||||
+2
-2
@@ -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)
|
||||
|
||||
+4
-5
@@ -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
|
||||
|
||||
+4
-7
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user