From 51daec63d23f483edd7a8f814a0f416ec1e64cd3 Mon Sep 17 00:00:00 2001 From: ngfchl Date: Tue, 7 Jul 2026 16:53:39 +0800 Subject: [PATCH] fix: harden game concurrency and rule writes --- db/__init__.py | 4 + db/models.py | 1 + services/game_service.py | 165 ++++++++++++++++++++++----------------- web/api.py | 23 +++--- 4 files changed, 111 insertions(+), 82 deletions(-) diff --git a/db/__init__.py b/db/__init__.py index 1bd08a4..c7885ce 100644 --- a/db/__init__.py +++ b/db/__init__.py @@ -82,6 +82,10 @@ async def ensure_runtime_schema(): UPDATE lottery_events SET prize_count = 1 WHERE prize_count IS NULL OR prize_count <= 0; UPDATE lottery_events SET remaining_count = CASE WHEN status = 'ended' THEN 0 ELSE GREATEST(1, prize_count) END WHERE remaining_count IS NULL OR remaining_count < 0; ALTER TABLE lottery_participants ADD COLUMN IF NOT EXISTS is_winner BOOLEAN NOT NULL DEFAULT FALSE; + DELETE FROM lottery_participants a + USING lottery_participants b + WHERE a.id > b.id AND a.event_id = b.event_id AND a.user_id = b.user_id; + CREATE UNIQUE INDEX IF NOT EXISTS uq_lottery_participants_event_user ON lottery_participants(event_id, user_id); CREATE TABLE IF NOT EXISTS lottery_events ( id BIGSERIAL PRIMARY KEY, created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), diff --git a/db/models.py b/db/models.py index 6c51131..a2c765d 100644 --- a/db/models.py +++ b/db/models.py @@ -214,3 +214,4 @@ class LotteryParticipant(TimestampMixin, Model): class Meta: table = "lottery_participants" indexes = (("event_id", "user_id"),) + unique_together = (("event_id", "user_id"),) diff --git a/services/game_service.py b/services/game_service.py index 89c7d7f..cb1e00d 100644 --- a/services/game_service.py +++ b/services/game_service.py @@ -4,6 +4,8 @@ import logging from datetime import datetime, timedelta from db.models import GameScore, LotteryEvent, LotteryParticipant, ShopItem, ShopTransaction +from tortoise.transactions import in_transaction +from tortoise.exceptions import IntegrityError logger = logging.getLogger("spam_guard") @@ -80,21 +82,27 @@ async def get_scores_summary(user_id: int, chat_id: int = 0) -> dict: return {"scores": scores, "total": total, "by_type": by_type} -async def deduct_points(chat_id: int, user_id: int, cost: int) -> bool: - summary = await get_scores_summary(user_id, chat_id) - if summary["total"] < cost: +async def _deduct_points_locked(conn, chat_id: int, user_id: int, cost: int) -> bool: + scores = await GameScore.filter(chat_id=chat_id, user_id=user_id).using_db(conn).select_for_update().all() + total = sum((s.score or 0) for s in scores) + if total < cost: return False remaining = cost - for score in sorted(summary["scores"], key=lambda x: x.score or 0, reverse=True): + for score in sorted(scores, key=lambda x: x.score or 0, reverse=True): if remaining <= 0: break - take = min(score.score, remaining) + take = min(score.score or 0, remaining) score.score -= take remaining -= take - await score.save() + await score.save(using_db=conn) return True +async def deduct_points(chat_id: int, user_id: int, cost: int) -> bool: + async with in_transaction() as conn: + return await _deduct_points_locked(conn, chat_id, user_id, cost) + + async def lottery_draw(chat_id: int, user_id: int, username: str) -> dict: cost = 10 if not await deduct_points(chat_id, user_id, cost): @@ -154,28 +162,34 @@ async def list_shop_transactions(chat_id: int, limit: int = 100) -> list[dict]: async def buy_shop_item(chat_id: int, user_id: int, item_id: str, username: str = "") -> dict: - items = await list_shop_items(chat_id) - item = next((x for x in items if str(x.get("item_key")) == str(item_id) and x.get("enabled", True)), None) - if not item: - return {"ok": False, "reason": "商品不存在"} - db_item = await ShopItem.filter(chat_id=int(item.get("chat_id") or chat_id), item_key=str(item.get("item_key"))).first() - if db_item and db_item.stock == 0: - return {"ok": False, "reason": "库存不足", "item": item} - if not await deduct_points(chat_id, user_id, int(item["cost"])): - return {"ok": False, "reason": "积分不足", "item": item} - if db_item and db_item.stock > 0: - db_item.stock -= 1 - await db_item.save(update_fields=["stock", "updated_at"]) - item["stock"] = db_item.stock - await ShopTransaction.create( - chat_id=chat_id, - user_id=user_id, - username=username, - item_key=str(item.get("item_key")), - item_name=item.get("name") or "", - cost=int(item.get("cost") or 0), - status="success", - ) + async with in_transaction() as conn: + db_item = await ShopItem.filter(chat_id=chat_id, item_key=str(item_id), enabled=True).using_db(conn).select_for_update().first() + if not db_item: + # 兼容默认内置商品:没有数据库记录时视为无限库存。 + builtin = SHOP_ITEMS.get(str(item_id)) + if not builtin: + return {"ok": False, "reason": "商品不存在"} + item = {"item_key": str(item_id), "name": builtin["name"], "cost": builtin["cost"], "stock": -1, "desc": builtin["desc"], "enabled": True} + else: + item = {"id": db_item.id, "chat_id": db_item.chat_id, "item_key": db_item.item_key, "name": db_item.name, "cost": db_item.cost, "stock": db_item.stock, "desc": db_item.desc, "enabled": db_item.enabled} + if db_item.stock == 0: + return {"ok": False, "reason": "库存不足", "item": item} + if not await _deduct_points_locked(conn, chat_id, user_id, int(item["cost"])): + return {"ok": False, "reason": "积分不足", "item": item} + if db_item and db_item.stock > 0: + db_item.stock -= 1 + await db_item.save(using_db=conn, update_fields=["stock", "updated_at"]) + item["stock"] = db_item.stock + await ShopTransaction.create( + chat_id=chat_id, + user_id=user_id, + username=username, + item_key=str(item.get("item_key")), + item_name=item.get("name") or "", + cost=int(item.get("cost") or 0), + status="success", + using_db=conn, + ) summary = await get_scores_summary(user_id, chat_id) return {"ok": True, "item": item, "total": summary["total"]} @@ -285,38 +299,41 @@ async def check_lottery_eligible(event: LotteryEvent, user_id: int) -> tuple[boo async def join_lottery_event(event_id: int, user_id: int, username: str = "", first_name: str = "") -> dict: - event = await LotteryEvent.filter(id=event_id).first() - if not event or event.status != "active": - return {"ok": False, "reason": "抽奖不存在或已结束"} - ok, reason = await check_lottery_eligible(event, user_id) - if not ok: - return {"ok": False, "reason": reason} - exists = await LotteryParticipant.filter(event_id=event_id, user_id=user_id).exists() - if exists: - return {"ok": False, "reason": "你已经参与过了"} + async with in_transaction() as conn: + event = await LotteryEvent.filter(id=event_id).using_db(conn).select_for_update().first() + if not event or event.status != "active": + return {"ok": False, "reason": "抽奖不存在或已结束"} + ok, reason = await check_lottery_eligible(event, user_id) + if not ok: + return {"ok": False, "reason": reason} + exists = await LotteryParticipant.filter(event_id=event_id, user_id=user_id).using_db(conn).exists() + if exists: + return {"ok": False, "reason": "你已经参与过了"} - is_instant = event.draw_type == "instant" or event.end_type == "instant" - if is_instant and int(event.remaining_count or 0) <= 0: - event.status = "ended" - event.ended_at = datetime.now() - await event.save() - return {"ok": False, "reason": "奖品已抽完,活动已结束"} - - participant = await LotteryParticipant.create(event_id=event_id, user_id=user_id, username=username, first_name=first_name, is_winner=is_instant) - count = await count_lottery_participants(event_id) - won = False - ended = False - if is_instant: - won = True - event.remaining_count = max(0, int(event.remaining_count or 0) - 1) - event.winner_id = user_id - event.winner_name = username or first_name or str(user_id) - if event.remaining_count <= 0: + is_instant = event.draw_type == "instant" or event.end_type == "instant" + if is_instant and int(event.remaining_count or 0) <= 0: event.status = "ended" event.ended_at = datetime.now() - ended = True - await event.save() - return {"ok": True, "event": event, "participant": participant, "count": count, "won": won, "ended": ended, "should_draw": (not is_instant) and event.end_type == "people" and count >= event.end_value} + await event.save(using_db=conn) + return {"ok": False, "reason": "奖品已抽完,活动已结束"} + try: + participant = await LotteryParticipant.create(event_id=event_id, user_id=user_id, username=username, first_name=first_name, is_winner=is_instant, using_db=conn) + except IntegrityError: + return {"ok": False, "reason": "你已经参与过了"} + count = await LotteryParticipant.filter(event_id=event_id).using_db(conn).count() + won = False + ended = False + if is_instant: + won = True + event.remaining_count = max(0, int(event.remaining_count or 0) - 1) + event.winner_id = user_id + event.winner_name = username or first_name or str(user_id) + if event.remaining_count <= 0: + event.status = "ended" + event.ended_at = datetime.now() + ended = True + await event.save(using_db=conn) + return {"ok": True, "event": event, "participant": participant, "count": count, "won": won, "ended": ended, "should_draw": (not is_instant) and event.end_type == "people" and count >= event.end_value} def lottery_event_payload(event: LotteryEvent | None) -> dict | None: @@ -350,7 +367,6 @@ def lottery_participant_payload(p: LotteryParticipant | None) -> dict | None: return { "id": p.id, "event_id": p.event_id, - "chat_id": p.chat_id, "user_id": p.user_id, "username": p.username, "first_name": p.first_name, @@ -361,19 +377,22 @@ def lottery_participant_payload(p: LotteryParticipant | None) -> dict | None: async def draw_lottery_event(event_id: int) -> dict: - event = await LotteryEvent.filter(id=event_id).first() - if not event or event.status != "active": - return {"ok": False, "reason": "抽奖不存在或已结束"} - participants = await LotteryParticipant.filter(event_id=event_id).all() - if not participants: - event.status = "cancelled" + async with in_transaction() as conn: + event = await LotteryEvent.filter(id=event_id).using_db(conn).select_for_update().first() + if not event or event.status != "active": + return {"ok": False, "reason": "抽奖不存在或已结束"} + participants = await LotteryParticipant.filter(event_id=event_id).using_db(conn).all() + if not participants: + event.status = "cancelled" + event.ended_at = datetime.now() + await event.save(using_db=conn) + return {"ok": False, "reason": "无人参与,抽奖已取消", "event": lottery_event_payload(event)} + winner = random.choice(participants) + winner.is_winner = True + await winner.save(using_db=conn) + event.status = "ended" + event.winner_id = winner.user_id + event.winner_name = winner.username or winner.first_name or str(winner.user_id) event.ended_at = datetime.now() - await event.save() - return {"ok": False, "reason": "无人参与,抽奖已取消", "event": lottery_event_payload(event)} - winner = random.choice(participants) - event.status = "ended" - event.winner_id = winner.user_id - event.winner_name = winner.username or winner.first_name or str(winner.user_id) - event.ended_at = datetime.now() - await event.save() - return {"ok": True, "event": lottery_event_payload(event), "winner": lottery_participant_payload(winner), "count": len(participants)} + await event.save(using_db=conn) + return {"ok": True, "event": lottery_event_payload(event), "winner": lottery_participant_payload(winner), "count": len(participants)} diff --git a/web/api.py b/web/api.py index 1ed54c8..27af550 100644 --- a/web/api.py +++ b/web/api.py @@ -36,6 +36,7 @@ _log_buffer: list[dict] = [] MAX_LOG_BUFFER = 300 MAX_LOG_MESSAGE_CHARS = 1200 _restart_bot_callback = None +_rule_write_lock = asyncio.Lock() def set_restart_bot_callback(callback): @@ -615,15 +616,19 @@ async def api_add_spam_rule(request: Request): return JSONResponse({"ok": False, "error": str(e)}, status_code=400) path = spam_detector.RULES_DIR / ("spam_keywords.txt" if rule_type == "keyword" else "spam_regex.txt") line = f"{stored_pattern}|{score}" - existing = path.read_text(encoding="utf-8") if path.exists() else "" - existing_patterns = {raw.rsplit("|", 1)[0].strip() for raw in existing.splitlines() if raw.strip() and not raw.strip().startswith("#")} - if stored_pattern in existing_patterns: - return JSONResponse({"ok": False, "error": "规则已存在"}, status_code=409) - with path.open("a", encoding="utf-8") as f: - if existing and not existing.endswith("\n"): - f.write("\n") - f.write(f"\n# Web 手动添加规则 {format_dt(now_tz())}\n{line}\n") - counts = spam_detector.reload_external_rules() + async with _rule_write_lock: + existing = path.read_text(encoding="utf-8") if path.exists() else "" + existing_patterns = {raw.rsplit("|", 1)[0].strip() for raw in existing.splitlines() if raw.strip() and not raw.strip().startswith("#")} + if stored_pattern in existing_patterns: + return JSONResponse({"ok": False, "error": "规则已存在"}, status_code=409) + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("a", encoding="utf-8") as f: + if existing and not existing.endswith("\n"): + f.write("\n") + note = f" 原始: {original_pattern}" if original_pattern else "" + f.write(f"\n# Web 手动添加规则 {format_dt(now_tz())}{note}\n{line}\n") + counts = spam_detector.reload_external_rules() + await Action.create(action="SPAM_RULE_ADD", details={"type": rule_type, "pattern": stored_pattern, "score": score, "original": original_pattern}) return JSONResponse({"ok": True, "type": rule_type, "pattern": stored_pattern, "original_pattern": original_pattern, "score": score, "rule_counts": counts})