""" 使用扣减接口(MFC 侧) 路由: POST /usage/consume 校验会话令牌并扣减一次使用 授权状态机复用 auth.py:激活 pending、取 active 授权、序列化。 """ import logging import aiomysql from fastapi import APIRouter from pydantic import BaseModel import auth import db from config import USAGE_LOG_THROTTLE_SECONDS logger = logging.getLogger(__name__) router = APIRouter(prefix="/usage", tags=["usage"]) class ConsumeRequest(BaseModel): session_token: str device_id: str @router.post("/consume") async def consume(payload: ConsumeRequest): """ 校验 session_token 并扣减一次使用。 成功:{"ok": true, "authorization": {...}} 失败:{"ok": false, "reason": "expired | exhausted | invalid_token"} """ async with db.acquire() as conn: await conn.begin() try: result = await _do_consume(conn, payload) except Exception: await conn.rollback() raise await conn.commit() return result async def _do_consume(conn, payload: ConsumeRequest) -> dict: async with conn.cursor(aiomysql.DictCursor) as cur: # 1. 校验会话:token 存在、未过期,且设备与签发时一致 await cur.execute( "SELECT user_id, device_id, (expires_at > NOW()) AS valid " "FROM sessions WHERE token = %s", (payload.session_token,), ) session_row = await cur.fetchone() if ( session_row is None or not session_row["valid"] or not session_row["device_id"] or session_row["device_id"] != payload.device_id ): return {"ok": False, "reason": "invalid_token"} user_id = session_row["user_id"] await cur.execute("UPDATE users SET last_seen_at = NOW() WHERE id = %s", (user_id,)) # 2. 惰性激活 pending 授权(内部对 users 行加锁,串行化同一用户并发) await auth.activate_pending_authorization(cur, user_id) # 3. 取当前 active 授权 auth_row = await auth.get_active_authorization(cur, user_id) if auth_row is None: return {"ok": False, "reason": await _reason_without_active(cur, user_id)} if auth_row["type"] == "time": await _log_time_usage(cur, user_id, payload.device_id, auth_row["id"]) return {"ok": True, "authorization": auth.serialize_authorization(auth_row)} # 4. 积分授权:条件 UPDATE 原子扣减,禁止先读后写 # 注意 SET 求值顺序:status 必须写在 remaining_points 自减之前, # 否则 IF 读到的是已减 1 的值,判空差 1。 await cur.execute( "UPDATE authorizations SET " "status = IF(remaining_points <= 1, 'exhausted', 'active'), " "remaining_points = remaining_points - 1, " "updated_at = NOW() " "WHERE id = %s AND status = 'active' AND remaining_points > 0", (auth_row["id"],), ) if cur.rowcount != 1: # 并发下已被其他请求扣完 return {"ok": False, "reason": "exhausted"} await cur.execute( "INSERT INTO usage_logs " "(user_id, device_id, authorization_id, cost_type, cost_points, used_at) " "VALUES (%s, %s, %s, 'points', 1, NOW())", (user_id, payload.device_id, auth_row["id"]), ) # 重查该行,返回扣减后的余额 await cur.execute( "SELECT id, type, end_at, remaining_points, total_points " "FROM authorizations WHERE id = %s", (auth_row["id"],), ) updated_row = await cur.fetchone() logger.info( "积分扣减 user_id=%s auth_id=%s 剩余=%s", user_id, auth_row["id"], updated_row["remaining_points"], ) return {"ok": True, "authorization": auth.serialize_authorization(updated_row)} async def _reason_without_active(cur, user_id: int) -> str: """无可用授权时,按最近一条失效授权的类型判定失败原因""" await cur.execute( "SELECT type FROM authorizations " "WHERE user_id = %s AND status IN ('expired', 'exhausted') " "ORDER BY updated_at DESC, id DESC LIMIT 1", (user_id,), ) row = await cur.fetchone() if row is not None and row["type"] == "points": return "exhausted" return "expired" async def _log_time_usage(cur, user_id: int, device_id: str, authorization_id: int) -> None: """时间授权写使用日志;同一用户+设备在节流窗口内只记一条""" await cur.execute( "SELECT 1 FROM usage_logs " "WHERE user_id = %s AND device_id = %s AND cost_type = 'time' " "AND used_at > NOW() - INTERVAL %s SECOND LIMIT 1", (user_id, device_id, USAGE_LOG_THROTTLE_SECONDS), ) if await cur.fetchone() is not None: return await cur.execute( "INSERT INTO usage_logs " "(user_id, device_id, authorization_id, cost_type, cost_points, used_at) " "VALUES (%s, %s, %s, 'time', 0, NOW())", (user_id, device_id, authorization_id), )