From b4a2a8c31ddc850b1ce458d060bcfee20c99449a Mon Sep 17 00:00:00 2001 From: linkeyu Date: Fri, 24 Jul 2026 11:14:31 +0800 Subject: [PATCH] =?UTF-8?q?=E5=8A=9F=E8=83=BD=EF=BC=9A=E9=99=90=E5=88=B6?= =?UTF-8?q?=E6=AF=8F=E4=BD=8D=E7=94=A8=E6=88=B7=E6=AF=8F=E5=A4=A9=E6=9C=80?= =?UTF-8?q?=E5=A4=9A=E5=8F=91=E8=B5=B7=20100=20=E6=AC=A1=E6=AF=94=E4=BB=B7?= =?UTF-8?q?=20(#165)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## 变更内容 - 登录用户按北京时间自然日计算比价发起次数,每人每天最多 100 次。 - 第 101 次起返回 HTTP 429,并提示“今日已比价超过100次,请明天再试”。 - 同一 trace_id 的网络重试按幂等处理,不会重复计数。 - 使用用户行锁串行化同一账号的并发请求,避免并发突破上限。 - 复用现有 comparison_record 的 running 记录,无需新增数据库迁移。 ## 本地验证 - 比价额度专项测试 11 项通过。 - 本次修改涉及文件的 Ruff 检查通过。 - 已覆盖未登录、重复 trace_id、跨自然日、第 100 次放行及第 101 次拒绝。 --------- Co-authored-by: CodexSandboxOffline <798648091@qq.com> Reviewed-on: https://gitea.shaguabijia.com/WonderableAI/shaguabijia-app-server/pulls/165 Co-authored-by: linkeyu Co-committed-by: linkeyu --- app/api/v1/compare_record.py | 37 +++++++++ app/repositories/comparison.py | 85 ++++++++++++++++++++- app/schemas/compare_record.py | 15 +++- tests/test_compare_daily_limit.py | 120 ++++++++++++++++++++++++++++++ 4 files changed, 255 insertions(+), 2 deletions(-) create mode 100644 tests/test_compare_daily_limit.py diff --git a/app/api/v1/compare_record.py b/app/api/v1/compare_record.py index f1d67d6..5afc8ac 100644 --- a/app/api/v1/compare_record.py +++ b/app/api/v1/compare_record.py @@ -20,6 +20,8 @@ from app.db.session import SessionLocal from app.models.comparison import ComparisonRecord from app.repositories import comparison as crud_compare from app.schemas.compare_record import ( + CompareStartReserveIn, + CompareStartReserveOut, CompareStatsOut, ComparisonRecordCreatedOut, ComparisonRecordDetailOut, @@ -35,6 +37,41 @@ logger = logging.getLogger("shagua.compare_record") router = APIRouter(prefix="/api/v1/compare", tags=["compare-record"]) +@router.post( + "/start", + response_model=CompareStartReserveOut, + summary="预占一次当日比价发起次数(每人每天最多100次)", +) +def reserve_compare_start( + payload: CompareStartReserveIn, + user: CurrentUser, + db: DbSession, +) -> CompareStartReserveOut: + try: + _, used = crud_compare.reserve_daily_start( + db, + user_id=user.id, + trace_id=payload.trace_id, + business_type=payload.business_type, + device_id=payload.device_id, + ) + except crud_compare.DailyCompareStartLimitExceeded: + raise HTTPException( + status_code=status.HTTP_429_TOO_MANY_REQUESTS, + detail="今日已比价超过100次,请明天再试", + ) from None + except crud_compare.ComparisonTraceOwnershipError: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail="比价任务标识冲突,请重新发起", + ) from None + return CompareStartReserveOut( + limit=crud_compare.DAILY_COMPARE_START_LIMIT, + used=used, + remaining=max(crud_compare.DAILY_COMPARE_START_LIMIT - used, 0), + ) + + @router.post( "/record", response_model=ComparisonRecordCreatedOut, diff --git a/app/repositories/comparison.py b/app/repositories/comparison.py index 7dceeab..53602de 100644 --- a/app/repositories/comparison.py +++ b/app/repositories/comparison.py @@ -5,7 +5,7 @@ """ from __future__ import annotations -from datetime import datetime +from datetime import datetime, timedelta from sqlalchemy import func, or_, select from sqlalchemy.orm import Session, defer @@ -14,8 +14,19 @@ from app.core.rewards import CN_TZ from app.models.ad_feed_reward import AdFeedRewardRecord from app.models.comparison import ComparisonRecord from app.models.savings import SavingsRecord +from app.models.user import User from app.schemas.compare_record import ComparisonRecordIn +DAILY_COMPARE_START_LIMIT = 100 + + +class DailyCompareStartLimitExceeded(Exception): + """The authenticated user has consumed today's comparison-start quota.""" + + +class ComparisonTraceOwnershipError(Exception): + """A trace id already belongs to a different authenticated user.""" + def _yuan_to_cents(yuan: float | None) -> int | None: """元(float)→ 分(int)。None 透传。""" @@ -243,6 +254,78 @@ def _get_by_trace(db: Session, trace_id: str) -> ComparisonRecord | None: ).scalar_one_or_none() +def reserve_daily_start( + db: Session, + *, + user_id: int, + trace_id: str, + business_type: str = "food", + device_id: str | None = None, + now: datetime | None = None, +) -> tuple[ComparisonRecord, int]: + """Atomically reserve one of a user's 100 Beijing-day comparison starts. + + ``trace_id`` makes client retries idempotent. Locking the user row serializes + concurrent starts for one account, so parallel requests cannot both consume + the final available slot. The reservation is the existing ``running`` + comparison row; later result reporting updates that same row. + """ + db.execute(select(User.id).where(User.id == user_id).with_for_update()).scalar_one() + + existing = _get_by_trace(db, trace_id) + if existing is not None: + if existing.user_id not in (None, user_id): + raise ComparisonTraceOwnershipError + if existing.user_id is None: + existing.user_id = user_id + if existing.device_id is None and device_id: + existing.device_id = device_id + db.commit() + db.refresh(existing) + + existing_at = existing.created_at + if existing_at.tzinfo is not None: + existing_at = existing_at.astimezone(CN_TZ).replace(tzinfo=None) + day_start = existing_at.replace(hour=0, minute=0, second=0, microsecond=0) + day_end = day_start + timedelta(days=1) + used = db.scalar( + select(func.count(ComparisonRecord.id)).where( + ComparisonRecord.user_id == user_id, + ComparisonRecord.created_at >= day_start, + ComparisonRecord.created_at < day_end, + ) + ) or 0 + return existing, int(used) + + current = now or datetime.now(CN_TZ) + if current.tzinfo is not None: + current = current.astimezone(CN_TZ).replace(tzinfo=None) + day_start = current.replace(hour=0, minute=0, second=0, microsecond=0) + day_end = day_start + timedelta(days=1) + used = db.scalar( + select(func.count(ComparisonRecord.id)).where( + ComparisonRecord.user_id == user_id, + ComparisonRecord.created_at >= day_start, + ComparisonRecord.created_at < day_end, + ) + ) or 0 + if used >= DAILY_COMPARE_START_LIMIT: + raise DailyCompareStartLimitExceeded + + rec = ComparisonRecord( + trace_id=trace_id, + user_id=user_id, + business_type=business_type or "food", + device_id=device_id, + status="running", + created_at=current, + ) + db.add(rec) + db.commit() + db.refresh(rec) + return rec, int(used) + 1 + + def harvest_running( db: Session, *, diff --git a/app/schemas/compare_record.py b/app/schemas/compare_record.py index 91ccd36..3887bf0 100644 --- a/app/schemas/compare_record.py +++ b/app/schemas/compare_record.py @@ -13,7 +13,6 @@ from datetime import datetime from pydantic import BaseModel, ConfigDict, Field, field_validator - # ===== 上报请求 ===== class ComparisonItemIn(BaseModel): @@ -198,6 +197,20 @@ class ComparisonRecordCreatedOut(BaseModel): id: int = Field(..., description="写入(或已存在)的记录 id") +class CompareStartReserveIn(BaseModel): + """Reserve one authenticated comparison start before the agent begins.""" + + trace_id: str = Field(..., min_length=1, max_length=64) + business_type: str = Field(default="food", min_length=1, max_length=16) + device_id: str | None = Field(default=None, max_length=64) + + +class CompareStartReserveOut(BaseModel): + limit: int + used: int + remaining: int + + class CompareStatsOut(BaseModel): """「我的」页省钱战绩卡(比价口径)聚合。""" diff --git a/tests/test_compare_daily_limit.py b/tests/test_compare_daily_limit.py new file mode 100644 index 0000000..a10afec --- /dev/null +++ b/tests/test_compare_daily_limit.py @@ -0,0 +1,120 @@ +from __future__ import annotations + +import time +from datetime import datetime, timedelta + +from sqlalchemy import func, select + +from app.core.rewards import CN_TZ +from app.core.security import decode_token +from app.db.session import SessionLocal +from app.models.comparison import ComparisonRecord + + +def _login(client) -> tuple[str, int]: + phone = f"137{int(time.time() * 1000) % 100000000:08d}" + sent = client.post("/api/v1/auth/sms/send", json={"phone": phone}) + assert sent.status_code == 200, sent.text + logged_in = client.post( + "/api/v1/auth/sms/login", + json={"phone": phone, "code": "123456"}, + ) + assert logged_in.status_code == 200, logged_in.text + token = logged_in.json()["access_token"] + return token, int(decode_token(token, expected_type="access")["sub"]) + + +def _headers(token: str) -> dict[str, str]: + return {"Authorization": f"Bearer {token}"} + + +def test_compare_start_requires_login(client) -> None: + response = client.post( + "/api/v1/compare/start", + json={"trace_id": "quota-no-auth", "business_type": "food"}, + ) + assert response.status_code == 401 + + +def test_compare_start_is_idempotent_by_trace_id(client) -> None: + token, user_id = _login(client) + payload = { + "trace_id": f"quota-idempotent-{user_id}", + "business_type": "ecom", + "device_id": "quota-device", + } + + first = client.post("/api/v1/compare/start", json=payload, headers=_headers(token)) + retry = client.post("/api/v1/compare/start", json=payload, headers=_headers(token)) + + assert first.status_code == 200, first.text + assert first.json() == {"limit": 100, "used": 1, "remaining": 99} + assert retry.status_code == 200, retry.text + assert retry.json() == first.json() + with SessionLocal() as db: + count = db.scalar( + select(func.count(ComparisonRecord.id)).where( + ComparisonRecord.trace_id == payload["trace_id"] + ) + ) + record = db.execute( + select(ComparisonRecord).where( + ComparisonRecord.trace_id == payload["trace_id"] + ) + ).scalar_one() + assert count == 1 + assert record.user_id == user_id + assert record.status == "running" + assert record.business_type == "ecom" + assert record.device_id == "quota-device" + + +def test_compare_start_rejects_101st_beijing_day_attempt(client) -> None: + token, user_id = _login(client) + now = datetime.now(CN_TZ).replace(tzinfo=None) + with SessionLocal() as db: + db.add_all( + [ + ComparisonRecord( + user_id=user_id, + trace_id=f"quota-full-{user_id}-{index}", + status="failed", + created_at=now, + ) + for index in range(99) + ] + ) + db.add( + ComparisonRecord( + user_id=user_id, + trace_id=f"quota-yesterday-{user_id}", + status="success", + created_at=now - timedelta(days=1), + ) + ) + db.commit() + + final_allowed_trace = f"quota-final-allowed-{user_id}" + allowed = client.post( + "/api/v1/compare/start", + json={"trace_id": final_allowed_trace, "business_type": "food"}, + headers=_headers(token), + ) + assert allowed.status_code == 200, allowed.text + assert allowed.json() == {"limit": 100, "used": 100, "remaining": 0} + + rejected_trace = f"quota-rejected-{user_id}" + response = client.post( + "/api/v1/compare/start", + json={"trace_id": rejected_trace, "business_type": "food"}, + headers=_headers(token), + ) + + assert response.status_code == 429 + assert response.json()["detail"] == "今日已比价超过100次,请明天再试" + with SessionLocal() as db: + assert db.scalar( + select(func.count(ComparisonRecord.id)).where( + ComparisonRecord.trace_id == rejected_trace + ) + ) == 0