147 lines
5.1 KiB
Python
147 lines
5.1 KiB
Python
"""
|
|||
|
|
使用扣减接口(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),
|
||
|
|
)
|