feat: 阶段 2 - /usage/consume、pending 授权惰性激活
- 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
This commit is contained in:
1 parent
a798fbad60
commit
cfb2a36c87
5 files changed
+255
-23
No files matched your search
@@ -28,3 +28,5 @@ REDIS_PASSWORD=your_redis_password_here
|
|||||||
FREE_AUTH_DAYS=7
|
FREE_AUTH_DAYS=7
|
||||||
SCENE_TTL_SECONDS=300
|
SCENE_TTL_SECONDS=300
|
||||||
SESSION_TTL_HOURS=24
|
SESSION_TTL_HOURS=24
|
||||||
|
# 时间授权 usage_logs 节流窗口(秒)
|
||||||
|
USAGE_LOG_THROTTLE_SECONDS=60
|
||||||
@@ -7,6 +7,8 @@
|
|||||||
|
|
||||||
业务函数(微信事件侧,由 wechat.py 调用):
|
业务函数(微信事件侧,由 wechat.py 调用):
|
||||||
handle_scan() 处理扫码事件:建用户、发免费授权、绑定场景、签发会话
|
handle_scan() 处理扫码事件:建用户、发免费授权、绑定场景、签发会话
|
||||||
|
activate_pending_authorization() 惰性激活 pending 授权(/usage/consume 也会调用)
|
||||||
|
serialize_authorization() 授权行序列化,供 /auth 与 /usage 复用
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
@@ -81,26 +83,38 @@ async def get_status(scene: str = Query(..., description="create_scene 返回的
|
|||||||
if scene_row["status"] == "expired":
|
if scene_row["status"] == "expired":
|
||||||
return {"status": "expired"}
|
return {"status": "expired"}
|
||||||
|
|
||||||
# 已扫码授权:判断用户当前是否有可用授权
|
# 已扫码授权:需在事务内惰性激活 pending(可能切换 active 授权)
|
||||||
auth_row = await _get_active_authorization(cur, scene_row["user_id"])
|
await conn.begin()
|
||||||
if auth_row is None:
|
try:
|
||||||
# 免费已领过且无有效授权 → 引导充值(阶段 3)
|
async with conn.cursor(aiomysql.DictCursor) as cur:
|
||||||
return {"status": "need_purchase"}
|
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(
|
await cur.execute(
|
||||||
"SELECT token, (expires_at > NOW()) AS not_expired FROM sessions "
|
"SELECT token, (expires_at > NOW()) AS not_expired FROM sessions "
|
||||||
"WHERE scene_id = %s ORDER BY id DESC LIMIT 1",
|
"WHERE scene_id = %s ORDER BY id DESC LIMIT 1",
|
||||||
(scene_row["id"],),
|
(scene_row["id"],),
|
||||||
)
|
)
|
||||||
session_row = await cur.fetchone()
|
session_row = await cur.fetchone()
|
||||||
if session_row is None or not session_row["not_expired"]:
|
if session_row is None or not session_row["not_expired"]:
|
||||||
return {"status": "expired"}
|
await conn.commit()
|
||||||
|
return {"status": "expired"}
|
||||||
|
|
||||||
return {
|
result = {
|
||||||
"status": "authorized",
|
"status": "authorized",
|
||||||
"session_token": session_row["token"],
|
"session_token": session_row["token"],
|
||||||
"authorization": _serialize_authorization(auth_row),
|
"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)
|
user_id = await _find_or_create_user(cur, openid)
|
||||||
await _grant_free_authorization(cur, user_id)
|
await _grant_free_authorization(cur, user_id)
|
||||||
|
# 若旧授权已失效,先激活 pending,再判断是否有可用授权
|
||||||
|
await activate_pending_authorization(cur, user_id)
|
||||||
# 免费授权发完后仍无可用授权 → 需要充值(阶段 3)
|
# 免费授权发完后仍无可用授权 → 需要充值(阶段 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(
|
await cur.execute(
|
||||||
"UPDATE auth_scenes SET status = 'authorized', user_id = %s, authorized_at = NOW() "
|
"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
|
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:
|
async def _find_or_create_user(cur, openid: str) -> int:
|
||||||
"""按 openid 查找用户,不存在则创建,并刷新 last_seen_at"""
|
"""按 openid 查找用户,不存在则创建,并刷新 last_seen_at"""
|
||||||
await cur.execute("SELECT id FROM users WHERE openid = %s", (openid,))
|
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
|
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"""
|
"""取用户当前 active 授权;已失效的惰性置为 expired / exhausted 并返回 None"""
|
||||||
await cur.execute(
|
await cur.execute(
|
||||||
"SELECT id, type, end_at, remaining_points, total_points, "
|
"SELECT id, type, end_at, remaining_points, total_points, "
|
||||||
@@ -253,7 +332,7 @@ async def _get_active_authorization(cur, user_id: int):
|
|||||||
return row
|
return row
|
||||||
|
|
||||||
|
|
||||||
def _serialize_authorization(row) -> dict:
|
def serialize_authorization(row) -> dict:
|
||||||
return {
|
return {
|
||||||
"type": row["type"],
|
"type": row["type"],
|
||||||
"end_at": row["end_at"].isoformat() if row["end_at"] else None,
|
"end_at": row["end_at"].isoformat() if row["end_at"] else None,
|
||||||
|
|||||||
@@ -30,3 +30,5 @@ REDIS_PASSWORD = os.getenv("REDIS_PASSWORD", "")
|
|||||||
FREE_AUTH_DAYS = int(os.getenv("FREE_AUTH_DAYS", "7")) # 首次关注赠送天数
|
FREE_AUTH_DAYS = int(os.getenv("FREE_AUTH_DAYS", "7")) # 首次关注赠送天数
|
||||||
SCENE_TTL_SECONDS = int(os.getenv("SCENE_TTL_SECONDS", "300")) # 二维码/scene 有效期(秒)
|
SCENE_TTL_SECONDS = int(os.getenv("SCENE_TTL_SECONDS", "300")) # 二维码/scene 有效期(秒)
|
||||||
SESSION_TTL_HOURS = int(os.getenv("SESSION_TTL_HOURS", "24")) # session_token 有效期(小时)
|
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"))
|
||||||
@@ -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),
|
||||||
|
)
|
||||||
@@ -4,7 +4,7 @@
|
|||||||
GET /wechat - 微信服务器验证(签名校验 + 返回 echostr)
|
GET /wechat - 微信服务器验证(签名校验 + 返回 echostr)
|
||||||
POST /wechat - 接收微信推送的消息和事件(subscribe / SCAN 触发扫码授权)
|
POST /wechat - 接收微信推送的消息和事件(subscribe / SCAN 触发扫码授权)
|
||||||
|
|
||||||
授权接口在 auth.py 中定义,通过 include_router 挂载。
|
授权接口在 auth.py、使用扣减接口在 usage.py 中定义,通过 include_router 挂载。
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import hashlib
|
import hashlib
|
||||||
@@ -18,6 +18,7 @@ from fastapi.responses import PlainTextResponse
|
|||||||
|
|
||||||
import auth
|
import auth
|
||||||
import db
|
import db
|
||||||
|
import usage
|
||||||
from config import WECHAT_TOKEN
|
from config import WECHAT_TOKEN
|
||||||
|
|
||||||
logging.basicConfig(
|
logging.basicConfig(
|
||||||
@@ -37,8 +38,9 @@ async def lifespan(app: FastAPI):
|
|||||||
logger.info("MySQL 连接池已关闭")
|
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(auth.router)
|
||||||
|
app.include_router(usage.router)
|
||||||
|
|
||||||
|
|
||||||
def verify_signature(signature: str, timestamp: str, nonce: str) -> bool:
|
def verify_signature(signature: str, timestamp: str, nonce: str) -> bool:
|
||||||
|
|||||||
Reference in new issue
Block a user