diff --git a/src/api/deps.py b/src/api/deps.py index 0f095c0..6f04641 100644 --- a/src/api/deps.py +++ b/src/api/deps.py @@ -187,8 +187,13 @@ async def rate_limit_predict(request: Request) -> None: client_ip = get_client_ip(request) if not _predict_limiter.is_allowed(client_ip): - logger.warning("rate limit exceeded for %s", client_ip) + logger.warning("rate limit exceeded for %s", ip) raise HTTPException( status_code=429, detail="请求过于频繁,请稍后再试(每分钟最多 10 次)", ) + + +def get_predict_rate_limit_remaining(request: Request) -> int: + """查询当前 IP 剩余的预测配额(用于响应中提示前端)。""" + return _predict_limiter.remaining(get_client_ip(request)) diff --git a/src/api/routes/matches.py b/src/api/routes/matches.py index 1b64101..655299c 100644 --- a/src/api/routes/matches.py +++ b/src/api/routes/matches.py @@ -72,7 +72,13 @@ async def list_matches( d = datetime.strptime(date, "%Y-%m-%d") except ValueError: raise HTTPException(400, "date 格式应为 YYYY-MM-DD") - q = q.where(Match.match_date >= d, Match.match_date < d + timedelta(days=1)) + # date 是用户本地日期(默认北京 UTC+8);match_date 存 UTC,需转换: + # 本地 00:00 (UTC+8) = UTC 前一天 16:00;本地 24:00 = UTC 当天 16:00 + from datetime import timezone as tz_mod + tz_cn = tz_mod(timedelta(hours=8)) + local_start = d.replace(tzinfo=tz_cn) + local_end = local_start + timedelta(days=1) + q = q.where(Match.match_date >= local_start, Match.match_date < local_end) # 未开赛按日期正序(最近的排最前,便于预测);其余按日期倒序(最新赛果在前) if status == "scheduled": diff --git a/src/api/routes/predict.py b/src/api/routes/predict.py index 4df8811..79dd32e 100644 --- a/src/api/routes/predict.py +++ b/src/api/routes/predict.py @@ -8,11 +8,11 @@ from __future__ import annotations import logging -from fastapi import APIRouter, Depends, HTTPException, Query +from fastapi import APIRouter, Depends, HTTPException, Query, Request from sqlalchemy import select from sqlalchemy.orm import selectinload -from src.api.deps import rate_limit_predict, require_admin +from src.api.deps import get_predict_rate_limit_remaining, rate_limit_predict, require_admin from src.api.schemas import PredictOut, PredictRequest, PredictionOut from src.db.base import AsyncSession, get_db_read, short_read from src.db.models import Match, Prediction @@ -24,8 +24,8 @@ router = APIRouter(prefix="/api/v1", tags=["predict"]) @router.post("/predict", response_model=PredictOut, dependencies=[Depends(rate_limit_predict)]) -async def predict(req: PredictRequest): - """对一场比赛调 LLM 预测。mode=multi(默认,5专家+终裁)或 single。 +async def predict(req: PredictRequest, request: Request): + """对一场比赛调 LLM 预测。mode=multi(默认,5专家+终裁)、single 或 baseline。 公开接口,仅做限流保护(不要求登录)。 @@ -100,7 +100,7 @@ async def predict(req: PredictRequest): latency_ms=result.get("latency_ms", 0) if result_dict else result.latency_ms, prompt_tokens=result.get("prompt_tokens") if result_dict else getattr(result, "prompt_tokens", None), completion_tokens=result.get("completion_tokens") if result_dict else getattr(result, "completion_tokens", None), - rate_limit_remaining=_predict_limiter.remaining(get_client_ip(request)), + rate_limit_remaining=get_predict_rate_limit_remaining(request), ) diff --git a/src/llm/agents/orchestrator.py b/src/llm/agents/orchestrator.py index 63e1bb1..c88339d 100644 --- a/src/llm/agents/orchestrator.py +++ b/src/llm/agents/orchestrator.py @@ -95,10 +95,10 @@ class MultiPredictResult: pred_1x2: str | None subjective_confidence: float | None reasoning: str | None + context: str agent_outputs: list[dict] agent_weights: dict | None status: str = "success" - context: str latency_ms: int | None = None prompt_tokens: int | None = None completion_tokens: int | None = None