diff --git a/.env.example b/.env.example index 1cd4407..b03ed4e 100644 --- a/.env.example +++ b/.env.example @@ -28,3 +28,5 @@ REDIS_PASSWORD=your_redis_password_here FREE_AUTH_DAYS=7 SCENE_TTL_SECONDS=300 SESSION_TTL_HOURS=24 +# 时间授权 usage_logs 节流窗口(秒) +USAGE_LOG_THROTTLE_SECONDS=60 diff --git a/auth.py b/auth.py index be7c1f8..dcc4fad 100644 --- a/auth.py +++ b/auth.py @@ -7,6 +7,8 @@ 业务函数(微信事件侧,由 wechat.py 调用): handle_scan() 处理扫码事件:建用户、发免费授权、绑定场景、签发会话 + activate_pending_authorization() 惰性激活 pending 授权(/usage/consume 也会调用) + serialize_authorization() 授权行序列化,供 /auth 与 /usage 复用 """ import logging @@ -81,26 +83,38 @@ async def get_status(scene: str = Query(..., description="create_scene 返回的 if scene_row["status"] == "expired": return {"status": "expired"} - # 已扫码授权:判断用户当前是否有可用授权 - auth_row = await _get_active_authorization(cur, scene_row["user_id"]) - if auth_row is None: - # 免费已领过且无有效授权 → 引导充值(阶段 3) - return {"status": "need_purchase"} + # 已扫码授权:需在事务内惰性激活 pending(可能切换 active 授权) + await conn.begin() + try: + async with conn.cursor(aiomysql.DictCursor) as cur: + await activate_pending_authorization(cur, scene_row["user_id"]) + # 判断用户当前是否有可用授权 + auth_row = await get_active_authorization(cur, scene_row["user_id"]) + if auth_row is None: + # 免费已领过且无有效授权 → 引导充值(阶段 3) + await conn.commit() + return {"status": "need_purchase"} - await cur.execute( - "SELECT token, (expires_at > NOW()) AS not_expired FROM sessions " - "WHERE scene_id = %s ORDER BY id DESC LIMIT 1", - (scene_row["id"],), - ) - session_row = await cur.fetchone() - if session_row is None or not session_row["not_expired"]: - return {"status": "expired"} + await cur.execute( + "SELECT token, (expires_at > NOW()) AS not_expired FROM sessions " + "WHERE scene_id = %s ORDER BY id DESC LIMIT 1", + (scene_row["id"],), + ) + session_row = await cur.fetchone() + if session_row is None or not session_row["not_expired"]: + await conn.commit() + return {"status": "expired"} - return { - "status": "authorized", - "session_token": session_row["token"], - "authorization": _serialize_authorization(auth_row), - } + result = { + "status": "authorized", + "session_token": session_row["token"], + "authorization": serialize_authorization(auth_row), + } + await conn.commit() + return result + except Exception: + await conn.rollback() + raise # --------------------------------------------------------------------------- @@ -148,8 +162,10 @@ async def handle_scan(scene_str: str, openid: str) -> str: user_id = await _find_or_create_user(cur, openid) await _grant_free_authorization(cur, user_id) + # 若旧授权已失效,先激活 pending,再判断是否有可用授权 + await activate_pending_authorization(cur, user_id) # 免费授权发完后仍无可用授权 → 需要充值(阶段 3) - has_auth = await _get_active_authorization(cur, user_id) is not None + has_auth = await get_active_authorization(cur, user_id) is not None await cur.execute( "UPDATE auth_scenes SET status = 'authorized', user_id = %s, authorized_at = NOW() " @@ -173,6 +189,69 @@ async def handle_scan(scene_str: str, openid: str) -> str: raise +async def activate_pending_authorization(cur, user_id: int) -> None: + """ + 惰性激活 pending 授权(互斥原则:同一用户同一时刻最多一条 active)。 + + 调用方必须已开启事务。步骤: + 1. 先把已失效的 active 标记为 expired / exhausted + 2. 若仍有 active,直接返回,保证互斥 + 3. 按 FIFO 激活一条 pending;时间授权按原时长从当前时刻重新起算 + (pending 等待期间 end_at 可能已过期,故重新锚定) + + 注:函数内先对 users 行加排他锁,串行化同一用户的并发激活。 + """ + await cur.execute("SELECT id FROM users WHERE id = %s FOR UPDATE", (user_id,)) + if await cur.fetchone() is None: + return + + # 1. 失效 active + await cur.execute( + "UPDATE authorizations SET status = 'expired' " + "WHERE user_id = %s AND status = 'active' AND type = 'time' " + "AND (end_at IS NULL OR end_at <= NOW())", + (user_id,), + ) + await cur.execute( + "UPDATE authorizations SET status = 'exhausted' " + "WHERE user_id = %s AND status = 'active' AND type = 'points' " + "AND remaining_points <= 0", + (user_id,), + ) + + # 2. 互斥检查:已有 active 则不激活 + await cur.execute( + "SELECT id FROM authorizations WHERE user_id = %s AND status = 'active' LIMIT 1", + (user_id,), + ) + if await cur.fetchone() is not None: + return + + # 3. FIFO 激活一条 pending + await cur.execute( + "SELECT id FROM authorizations WHERE user_id = %s AND status = 'pending' " + "ORDER BY created_at ASC, id ASC LIMIT 1", + (user_id,), + ) + pending_row = await cur.fetchone() + if pending_row is None: + return + + # end_at 的赋值必须写在 start_at 之前:MySQL 的 SET 从左到右求值, + # 否则 TIMESTAMPDIFF 会读到刚被改成 NOW() 的 start_at,时长归零。 + await cur.execute( + "UPDATE authorizations SET " + "end_at = IF(type = 'time', " + " DATE_ADD(NOW(), INTERVAL TIMESTAMPDIFF(SECOND, start_at, end_at) SECOND), " + " end_at), " + "start_at = IF(type = 'time', NOW(), start_at), " + "status = 'active', updated_at = NOW() " + "WHERE id = %s AND status = 'pending'", + (pending_row["id"],), + ) + logger.info("已激活 pending 授权 user_id=%s auth_id=%s", user_id, pending_row["id"]) + + async def _find_or_create_user(cur, openid: str) -> int: """按 openid 查找用户,不存在则创建,并刷新 last_seen_at""" await cur.execute("SELECT id FROM users WHERE openid = %s", (openid,)) @@ -223,7 +302,7 @@ async def _has_active_authorization(cur, user_id: int) -> bool: return await cur.fetchone() is not None -async def _get_active_authorization(cur, user_id: int): +async def get_active_authorization(cur, user_id: int): """取用户当前 active 授权;已失效的惰性置为 expired / exhausted 并返回 None""" await cur.execute( "SELECT id, type, end_at, remaining_points, total_points, " @@ -253,7 +332,7 @@ async def _get_active_authorization(cur, user_id: int): return row -def _serialize_authorization(row) -> dict: +def serialize_authorization(row) -> dict: return { "type": row["type"], "end_at": row["end_at"].isoformat() if row["end_at"] else None, diff --git a/config.py b/config.py index 6f12911..6c2af7b 100644 --- a/config.py +++ b/config.py @@ -30,3 +30,5 @@ REDIS_PASSWORD = os.getenv("REDIS_PASSWORD", "") FREE_AUTH_DAYS = int(os.getenv("FREE_AUTH_DAYS", "7")) # 首次关注赠送天数 SCENE_TTL_SECONDS = int(os.getenv("SCENE_TTL_SECONDS", "300")) # 二维码/scene 有效期(秒) SESSION_TTL_HOURS = int(os.getenv("SESSION_TTL_HOURS", "24")) # session_token 有效期(小时) +# 时间授权 usage_logs 节流窗口(秒):同一用户+设备在此窗口内只记一条 +USAGE_LOG_THROTTLE_SECONDS = int(os.getenv("USAGE_LOG_THROTTLE_SECONDS", "60")) diff --git a/usage.py b/usage.py new file mode 100644 index 0000000..689232c --- /dev/null +++ b/usage.py @@ -0,0 +1,147 @@ +""" +使用扣减接口(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), + ) diff --git a/wechat.py b/wechat.py index e775762..ce77280 100644 --- a/wechat.py +++ b/wechat.py @@ -4,7 +4,7 @@ GET /wechat - 微信服务器验证(签名校验 + 返回 echostr) POST /wechat - 接收微信推送的消息和事件(subscribe / SCAN 触发扫码授权) -授权接口在 auth.py 中定义,通过 include_router 挂载。 +授权接口在 auth.py、使用扣减接口在 usage.py 中定义,通过 include_router 挂载。 """ import hashlib @@ -18,6 +18,7 @@ from fastapi.responses import PlainTextResponse import auth import db +import usage from config import WECHAT_TOKEN logging.basicConfig( @@ -37,8 +38,9 @@ async def lifespan(app: FastAPI): logger.info("MySQL 连接池已关闭") -app = FastAPI(title="WeChat API", version="0.2.0", lifespan=lifespan) +app = FastAPI(title="WeChat API", version="0.3.0", lifespan=lifespan) app.include_router(auth.router) +app.include_router(usage.router) def verify_signature(signature: str, timestamp: str, nonce: str) -> bool: