- usage.py: POST /usage/consume,事务内校验会话(严格 device_id)、 激活 pending、按 time/points 分支扣减;失败返回 ok:false + reason - 积分扣减用条件 UPDATE 原子完成,status 赋值写在自减之前 (MySQL SET 从左到右求值,否则 IF 读到已减 1 的值,判空差 1) - 时间授权 usage_logs 按 (user_id, device_id) 60 秒节流;积分每次必写 - auth.py: 新增 activate_pending_authorization(),失效 active 后 FIFO 激活 pending,时间授权按原时长从当前时刻重新锚定 - get_status/handle_scan 接入惰性激活;get_status 改事务包裹 - _get_active_authorization/_serialize_authorization 改公开供 usage 复用 - config/.env.example: 新增 USAGE_LOG_THROTTLE_SECONDS
341 lines
13 KiB
Python
341 lines
13 KiB
Python
"""
|
||
授权接口与核心业务逻辑
|
||
|
||
路由(MFC 侧):
|
||
POST /auth/create_scene 生成 scene_str、创建微信临时二维码并落库
|
||
GET /auth/status 轮询扫码授权结果
|
||
|
||
业务函数(微信事件侧,由 wechat.py 调用):
|
||
handle_scan() 处理扫码事件:建用户、发免费授权、绑定场景、签发会话
|
||
activate_pending_authorization() 惰性激活 pending 授权(/usage/consume 也会调用)
|
||
serialize_authorization() 授权行序列化,供 /auth 与 /usage 复用
|
||
"""
|
||
|
||
import logging
|
||
import secrets
|
||
|
||
import aiomysql
|
||
from fastapi import APIRouter, HTTPException, Query
|
||
from pydantic import BaseModel
|
||
|
||
import db
|
||
from config import FREE_AUTH_DAYS, SCENE_TTL_SECONDS, SESSION_TTL_HOURS
|
||
from wechat_api import create_temp_qrcode
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
router = APIRouter(prefix="/auth", tags=["auth"])
|
||
|
||
SCENE_PREFIX = "pc_"
|
||
|
||
|
||
class CreateSceneRequest(BaseModel):
|
||
device_id: str
|
||
|
||
|
||
@router.post("/create_scene")
|
||
async def create_scene(payload: CreateSceneRequest):
|
||
"""生成唯一 scene_str,调用微信接口创建临时二维码,写入 auth_scenes"""
|
||
scene_str = SCENE_PREFIX + secrets.token_hex(16)
|
||
try:
|
||
qr_url = await create_temp_qrcode(scene_str, SCENE_TTL_SECONDS)
|
||
except Exception as exc:
|
||
# 微信接口不可用(appid/secret 未配置、网络异常、token 失效等)
|
||
logger.exception("创建二维码失败 scene_str=%s", scene_str)
|
||
raise HTTPException(status_code=502, detail=f"创建微信二维码失败: {exc}")
|
||
|
||
async with db.acquire() as conn:
|
||
async with conn.cursor() as cur:
|
||
await cur.execute(
|
||
"INSERT INTO auth_scenes (scene_str, device_id, status, created_at, expires_at) "
|
||
"VALUES (%s, %s, 'pending', NOW(), DATE_ADD(NOW(), INTERVAL %s SECOND))",
|
||
(scene_str, payload.device_id, SCENE_TTL_SECONDS),
|
||
)
|
||
|
||
return {"scene_str": scene_str, "qr_url": qr_url, "expires_in": SCENE_TTL_SECONDS}
|
||
|
||
|
||
@router.get("/status")
|
||
async def get_status(scene: str = Query(..., description="create_scene 返回的 scene_str")):
|
||
"""查询扫码授权状态:pending / authorized / expired / need_purchase"""
|
||
async with db.acquire() as conn:
|
||
async with conn.cursor(aiomysql.DictCursor) as cur:
|
||
await cur.execute(
|
||
"SELECT id, status, user_id, (expires_at > NOW()) AS not_expired "
|
||
"FROM auth_scenes WHERE scene_str = %s",
|
||
(scene,),
|
||
)
|
||
scene_row = await cur.fetchone()
|
||
if scene_row is None:
|
||
raise HTTPException(status_code=404, detail="scene 不存在")
|
||
|
||
# 尚未扫码:过期则惰性置为 expired
|
||
if scene_row["status"] in ("pending", "scanned"):
|
||
if not scene_row["not_expired"]:
|
||
await cur.execute(
|
||
"UPDATE auth_scenes SET status = 'expired' "
|
||
"WHERE id = %s AND status IN ('pending', 'scanned')",
|
||
(scene_row["id"],),
|
||
)
|
||
return {"status": "expired"}
|
||
return {"status": "pending"}
|
||
|
||
if scene_row["status"] == "expired":
|
||
return {"status": "expired"}
|
||
|
||
# 已扫码授权:需在事务内惰性激活 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"]:
|
||
await conn.commit()
|
||
return {"status": "expired"}
|
||
|
||
result = {
|
||
"status": "authorized",
|
||
"session_token": session_row["token"],
|
||
"authorization": serialize_authorization(auth_row),
|
||
}
|
||
await conn.commit()
|
||
return result
|
||
except Exception:
|
||
await conn.rollback()
|
||
raise
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# 业务逻辑
|
||
# ---------------------------------------------------------------------------
|
||
|
||
async def handle_scan(scene_str: str, openid: str) -> str:
|
||
"""
|
||
处理扫码事件,返回结果(取值与 /auth/status 的 status 一致):
|
||
|
||
- "authorized" 场景绑定成功,且用户有可用授权
|
||
- "need_purchase" 场景绑定成功,但用户没有可用授权(需充值)
|
||
- "invalid" 场景不存在、已过期,或已被处理过(重复扫码)
|
||
|
||
整个流程在事务内完成,并对 scene 行加排他锁,保证同一 scene 只被处理一次。
|
||
"""
|
||
async with db.acquire() as conn:
|
||
await conn.begin()
|
||
try:
|
||
async with conn.cursor(aiomysql.DictCursor) as cur:
|
||
await cur.execute(
|
||
"SELECT id, status, device_id, (expires_at > NOW()) AS not_expired "
|
||
"FROM auth_scenes WHERE scene_str = %s FOR UPDATE",
|
||
(scene_str,),
|
||
)
|
||
scene_row = await cur.fetchone()
|
||
|
||
if scene_row is None:
|
||
await conn.rollback()
|
||
return "invalid"
|
||
|
||
if scene_row["status"] == "authorized":
|
||
# 重复扫码:只处理第一次,后续忽略
|
||
await conn.rollback()
|
||
return "invalid"
|
||
|
||
if not scene_row["not_expired"]:
|
||
await cur.execute(
|
||
"UPDATE auth_scenes SET status = 'expired' "
|
||
"WHERE id = %s AND status IN ('pending', 'scanned')",
|
||
(scene_row["id"],),
|
||
)
|
||
await conn.commit()
|
||
return "invalid"
|
||
|
||
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
|
||
|
||
await cur.execute(
|
||
"UPDATE auth_scenes SET status = 'authorized', user_id = %s, authorized_at = NOW() "
|
||
"WHERE id = %s",
|
||
(user_id, scene_row["id"]),
|
||
)
|
||
|
||
token = secrets.token_urlsafe(32)
|
||
await cur.execute(
|
||
"INSERT INTO sessions (token, user_id, device_id, scene_id, created_at, expires_at) "
|
||
"VALUES (%s, %s, %s, %s, NOW(), DATE_ADD(NOW(), INTERVAL %s HOUR))",
|
||
(token, user_id, scene_row["device_id"], scene_row["id"], SESSION_TTL_HOURS),
|
||
)
|
||
|
||
await conn.commit()
|
||
result = "authorized" if has_auth else "need_purchase"
|
||
logger.info("扫码完成 scene_str=%s openid=%s 结果=%s", scene_str, openid, result)
|
||
return result
|
||
except Exception:
|
||
await conn.rollback()
|
||
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,))
|
||
row = await cur.fetchone()
|
||
if row is None:
|
||
# INSERT IGNORE + 重查:并发扫码时避免唯一键冲突报错
|
||
await cur.execute(
|
||
"INSERT IGNORE INTO users (openid, created_at, last_seen_at, has_claimed_free) "
|
||
"VALUES (%s, NOW(), NOW(), 0)",
|
||
(openid,),
|
||
)
|
||
await cur.execute("SELECT id FROM users WHERE openid = %s", (openid,))
|
||
row = await cur.fetchone()
|
||
|
||
await cur.execute("UPDATE users SET last_seen_at = NOW() WHERE id = %s", (row["id"],))
|
||
return row["id"]
|
||
|
||
|
||
async def _grant_free_authorization(cur, user_id: int) -> None:
|
||
"""
|
||
首次关注赠送 7 天时间授权。
|
||
|
||
以 has_claimed_free 的条件更新作为幂等闸门:只有把 0 改成 1 的那一次才真正发授权。
|
||
若用户已有 active 授权(互斥原则),新授权以 pending 保存。
|
||
"""
|
||
await cur.execute(
|
||
"UPDATE users SET has_claimed_free = 1 WHERE id = %s AND has_claimed_free = 0",
|
||
(user_id,),
|
||
)
|
||
if cur.rowcount != 1:
|
||
return
|
||
|
||
status = "pending" if await _has_active_authorization(cur, user_id) else "active"
|
||
await cur.execute(
|
||
"INSERT INTO authorizations "
|
||
"(user_id, type, start_at, end_at, remaining_points, total_points, source, status, created_at, updated_at) "
|
||
"VALUES (%s, 'time', NOW(), DATE_ADD(NOW(), INTERVAL %s DAY), 0, 0, 'free', %s, NOW(), NOW())",
|
||
(user_id, FREE_AUTH_DAYS, status),
|
||
)
|
||
logger.info("已发放免费授权 user_id=%s days=%s status=%s", user_id, FREE_AUTH_DAYS, status)
|
||
|
||
|
||
async def _has_active_authorization(cur, user_id: int) -> bool:
|
||
await cur.execute(
|
||
"SELECT 1 FROM authorizations WHERE user_id = %s AND status = 'active' LIMIT 1",
|
||
(user_id,),
|
||
)
|
||
return await cur.fetchone() is not None
|
||
|
||
|
||
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, "
|
||
"(end_at IS NOT NULL AND end_at > NOW()) AS time_valid "
|
||
"FROM authorizations WHERE user_id = %s AND status = 'active' ORDER BY id DESC LIMIT 1",
|
||
(user_id,),
|
||
)
|
||
row = await cur.fetchone()
|
||
if row is None:
|
||
return None
|
||
|
||
if row["type"] == "time":
|
||
if not row["time_valid"]:
|
||
await cur.execute(
|
||
"UPDATE authorizations SET status = 'expired' WHERE id = %s AND status = 'active'",
|
||
(row["id"],),
|
||
)
|
||
return None
|
||
return row
|
||
|
||
if row["remaining_points"] <= 0:
|
||
await cur.execute(
|
||
"UPDATE authorizations SET status = 'exhausted' WHERE id = %s AND status = 'active'",
|
||
(row["id"],),
|
||
)
|
||
return None
|
||
return row
|
||
|
||
|
||
def serialize_authorization(row) -> dict:
|
||
return {
|
||
"type": row["type"],
|
||
"end_at": row["end_at"].isoformat() if row["end_at"] else None,
|
||
"remaining_points": row["remaining_points"],
|
||
}
|