feat(proxy): T-P1 鉴权+账本(原子预扣/热路径缓存/令牌桶/管理鉴权原语)
- ledger.py:students/keys CRUD + 日限额 check_and_count(跨日重置,注入时钟)+ 流水分页;计费两阶段拆分至 billing.py(BillingMixin,评审聚焦): try_hold 锁内读现状->日上限判定->余额原子防线->holding 流水(request_id 幂等)、 settle 按真实值回补(预扣-实际)差额、void 全额退款(上游失败) - auth.py:issue_key(sk-campus- 前缀,明文只返回一次,库存 sha256+前缀)、 authenticate 校验链(形态/注销/停用/日额/rpm,热路径 LRU TTL30s D-P10)、 RateLimiter(令牌桶 rpm + per-key 并发信号量,D-P9 单进程)、 verify_admin(hmac.compare_digest;未配置仅 loopback) - 测试 +17:签发/错key/注销/停用/403/rpm 429/日额 429/并发槽/admin 策略 + 预扣-结算-回补一致/双向差额/余额不足/日上限/幂等/跨日重置/void/并发10路不超扣 - 全量 339 passed
This commit is contained in:
+161
-8
@@ -1,22 +1,175 @@
|
||||
"""鉴权与限流(T-P1 落地;本文件先立签名)。"""
|
||||
"""鉴权与限流(T-P1):key 签发/校验/注销 + 令牌桶 + 并发信号量 + 热路径缓存。
|
||||
|
||||
工程要点:
|
||||
- D-P8:明文 key 只在 issue_key 返回一次;库中仅存 sha256 哈希 + 前 12 位前缀(展示用)。
|
||||
- D-P10 热路径:哈希 -> (key_id, student 上下文) 的进程内 LRU,TTL 30s,
|
||||
注销/停用后写失效(revoke 时主动失效,最坏 30s 内仍可能命中旧缓存——
|
||||
注销语义取"最终生效",符合校园场景)。
|
||||
- 限流:令牌桶(rpm,按 key),日请求上限走 ledger.check_and_count(持久),
|
||||
并发上限 per-key 信号量(concurrent_per_key,进程内,D-P9 单进程前提)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import hmac
|
||||
import secrets
|
||||
import threading
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from gateway.proxy.errors import ProxyAuthError, QuotaError, SuspendedError
|
||||
from gateway.proxy.ledgerutil import _today
|
||||
|
||||
KEY_PREFIX = "sk-campus-"
|
||||
_AUTH_TTL = 30.0
|
||||
_AUTH_CACHE_MAX = 4096
|
||||
|
||||
|
||||
def _hash_key(plaintext: str) -> str:
|
||||
return hashlib.sha256(plaintext.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def issue_key(ledger, student_id: int, rpm_cap: Optional[int] = None,
|
||||
day_cap_req: Optional[int] = None) -> Dict[str, Any]:
|
||||
"""签发学生代理 key:明文只返回一次(T-P1 实现)。"""
|
||||
raise NotImplementedError("T-P1")
|
||||
"""签发学生代理 key:明文只此一次返回;库中仅存哈希与前缀。
|
||||
|
||||
返回 {key_id, key(明文), prefix};学生不存在时 ledger.create_key 返回 None
|
||||
-> 抛 ProxyAuthError。
|
||||
"""
|
||||
plaintext = KEY_PREFIX + secrets.token_urlsafe(24)
|
||||
kid = ledger.create_key(
|
||||
student_id, _hash_key(plaintext), plaintext[:12],
|
||||
rpm_cap=rpm_cap if rpm_cap is not None else 10,
|
||||
day_cap_req=day_cap_req if day_cap_req is not None else 200)
|
||||
if kid is None:
|
||||
raise ProxyAuthError(f"学生不存在: {student_id}")
|
||||
return {"key_id": kid, "key": plaintext, "prefix": plaintext[:12]}
|
||||
|
||||
|
||||
def authenticate(authorization: str, ledger, limits, now: float) -> Dict[str, Any]:
|
||||
"""校验 Bearer key -> 学生上下文(T-P1 实现;带热路径缓存)。"""
|
||||
raise NotImplementedError("T-P1")
|
||||
class _AuthCache:
|
||||
"""热路径鉴权缓存(LRU + TTL 30s;D-P10)。"""
|
||||
|
||||
def __init__(self, ttl: float = _AUTH_TTL, maxsize: int = _AUTH_CACHE_MAX):
|
||||
self._ttl = ttl
|
||||
self._max = maxsize
|
||||
self._data: "OrderedDict[str, tuple[float, Any]]" = OrderedDict()
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def get(self, key_hash: str) -> Optional[Any]:
|
||||
now = time.monotonic()
|
||||
with self._lock:
|
||||
item = self._data.get(key_hash)
|
||||
if item is None:
|
||||
return None
|
||||
ts, ctx = item
|
||||
if now - ts > self._ttl:
|
||||
self._data.pop(key_hash, None)
|
||||
return None
|
||||
self._data.move_to_end(key_hash)
|
||||
return ctx
|
||||
|
||||
def put(self, key_hash: str, ctx: Any) -> None:
|
||||
with self._lock:
|
||||
self._data[key_hash] = (time.monotonic(), ctx)
|
||||
self._data.move_to_end(key_hash)
|
||||
while len(self._data) > self._max:
|
||||
self._data.popitem(last=False)
|
||||
|
||||
def invalidate(self, key_hash: str) -> None:
|
||||
with self._lock:
|
||||
self._data.pop(key_hash, None)
|
||||
|
||||
|
||||
class RateLimiter:
|
||||
"""令牌桶 rpm + per-key 并发信号量(T-P1 实现)。"""
|
||||
"""令牌桶(rpm,按 key,进程内)+ per-key 并发信号量。
|
||||
|
||||
D-P9:全部为进程内状态,uvicorn workers=1 是正确性前提。
|
||||
"""
|
||||
|
||||
def __init__(self, concurrent_per_key: int = 2):
|
||||
self._tokens: Dict[int, tuple[float, float]] = {} # key_id -> (tokens, last_ts)
|
||||
self._locks: Dict[int, threading.Lock] = {}
|
||||
self._sems: Dict[int, threading.Semaphore] = {}
|
||||
self._global = threading.Lock()
|
||||
self._concurrent = max(1, int(concurrent_per_key))
|
||||
|
||||
def allow(self, key_id: int, rpm_cap: int) -> bool:
|
||||
raise NotImplementedError("T-P1")
|
||||
"""令牌桶放行判定(容量 = rpm_cap,速率 = rpm/60 每秒)。"""
|
||||
now = time.monotonic()
|
||||
with self._global:
|
||||
lock = self._locks.setdefault(key_id, threading.Lock())
|
||||
with lock:
|
||||
tokens, last = self._tokens.get(key_id, (float(rpm_cap), now))
|
||||
tokens = min(float(rpm_cap), tokens + (now - last) * (rpm_cap / 60.0))
|
||||
if tokens < 1.0:
|
||||
self._tokens[key_id] = (tokens, now)
|
||||
return False
|
||||
self._tokens[key_id] = (tokens - 1.0, now)
|
||||
return True
|
||||
|
||||
def acquire_slot(self, key_id: int) -> bool:
|
||||
"""并发槽(非阻塞);返回 False = 超并发上限(429)。"""
|
||||
with self._global:
|
||||
sem = self._sems.setdefault(
|
||||
key_id, threading.Semaphore(self._concurrent))
|
||||
return sem.acquire(blocking=False)
|
||||
|
||||
def release_slot(self, key_id: int) -> None:
|
||||
sem = self._sems.get(key_id)
|
||||
if sem is not None:
|
||||
sem.release()
|
||||
|
||||
|
||||
def authenticate(authorization: str, ledger, limiter: RateLimiter,
|
||||
now: Optional[float] = None) -> Dict[str, Any]:
|
||||
"""校验 Bearer key -> 学生/key 上下文(§6 签名)。
|
||||
|
||||
校验链(任一失败即短路与对应异常):
|
||||
形态(401) -> 哈希存在且未注销(401) -> 学生 active(403)
|
||||
-> 日请求上限(429, 持久计数) -> rpm 令牌桶(429, 进程内)
|
||||
返回 {key_id, student_id, rpm_cap, day_cap_req, balance_milli, ...}。
|
||||
"""
|
||||
now_ts = now if now is not None else time.time()
|
||||
token = (authorization or "").strip()
|
||||
if not token.lower().startswith("bearer "):
|
||||
raise ProxyAuthError("缺少 Bearer 凭据")
|
||||
plaintext = token[7:].strip()
|
||||
if not plaintext.startswith(KEY_PREFIX):
|
||||
raise ProxyAuthError("无效 key")
|
||||
key_hash = _hash_key(plaintext)
|
||||
|
||||
ctx = _AUTH_SINGLETON.get(key_hash)
|
||||
if ctx is None:
|
||||
row = ledger.find_key(key_hash)
|
||||
if row is None or row["revoked"]:
|
||||
raise ProxyAuthError("无效或已注销的 key")
|
||||
if row["student_status"] != "active":
|
||||
raise SuspendedError("学生账户已停用")
|
||||
ctx = row
|
||||
_AUTH_SINGLETON.put(key_hash, ctx)
|
||||
|
||||
if not ledger.check_and_count(ctx["key_id"], ctx["student_id"], now_ts):
|
||||
raise QuotaError("已达当日请求上限")
|
||||
if not _AUTH_SINGLETON_LIMITS.allow(ctx["key_id"], int(ctx["rpm_cap"])):
|
||||
raise QuotaError("请求过于频繁(rpm)")
|
||||
return dict(ctx)
|
||||
|
||||
|
||||
# 进程内单例(热缓存 + 限流器;D-P9 单进程)
|
||||
_AUTH_SINGLETON = _AuthCache()
|
||||
_AUTH_SINGLETON_LIMITS = RateLimiter(concurrent_per_key=2)
|
||||
|
||||
|
||||
def reset_auth_state() -> None:
|
||||
"""测试用:清空热缓存与限流器。"""
|
||||
global _AUTH_SINGLETON, _AUTH_SINGLETON_LIMITS
|
||||
_AUTH_SINGLETON = _AuthCache()
|
||||
_AUTH_SINGLETON_LIMITS = RateLimiter(concurrent_per_key=2)
|
||||
|
||||
|
||||
def verify_admin(requested_key: str, admin_key: str, client_host: str = "") -> bool:
|
||||
"""管理面鉴权(§5.2):hmac.compare_digest 防时序;未配置 admin_key 时仅 loopback。"""
|
||||
if admin_key:
|
||||
return hmac.compare_digest(str(requested_key or ""), str(admin_key))
|
||||
return client_host in ("127.0.0.1", "::1", "localhost", "testclient")
|
||||
|
||||
@@ -0,0 +1,157 @@
|
||||
"""计费两阶段(D-P11):预扣-结算-退款,作为 Ledger 的 mixin。
|
||||
|
||||
拆分为独立模块的工程原因:ledger.py 承载 DDL/CRUD/查询;金额敏感的
|
||||
两阶段语义集中在此,便于评审与测试聚焦。
|
||||
|
||||
并发正确性:全局锁内 check-then-act(单写者模型)+ 余额原子防线
|
||||
(语句条件 balance_milli >= ?,受影响行数为 0 即不足)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from gateway.proxy.ledgerutil import _today
|
||||
|
||||
# 预扣:余额原子防线;日上限判定在锁内 Python 侧(单写者模型下等价安全)
|
||||
# 占位顺序:est, spent_total, today, sid, est
|
||||
_SQL_HOLD = """UPDATE students SET
|
||||
balance_milli = balance_milli - ?,
|
||||
spent_today_milli = ?, spent_date = ?
|
||||
WHERE id = ? AND balance_milli >= ?"""
|
||||
|
||||
_SQL_HOLD_MARK = """INSERT INTO usage_ledger
|
||||
(request_id, ts, key_id, model, bucket, charged_milli, status)
|
||||
VALUES (?, ?, ?, ?, ?, ?, 'holding')"""
|
||||
|
||||
_SQL_INSUFFICIENT = """INSERT OR REPLACE INTO usage_ledger
|
||||
(request_id, ts, key_id, model, bucket, charged_milli, status)
|
||||
VALUES (?, ?, ?, ?, ?, 0, 'insufficient')"""
|
||||
|
||||
_SQL_SETTLE_REFUND = """UPDATE students SET
|
||||
balance_milli = balance_milli + ?,
|
||||
spent_today_milli = MAX(0, spent_today_milli - ?)
|
||||
WHERE id = ?"""
|
||||
|
||||
_SQL_SETTLE_USAGE = """UPDATE usage_ledger SET
|
||||
in_miss_tok = ?, in_hit_tok = ?, out_tok = ?, gateway_cached = ?,
|
||||
upstream_cost_milli = ?, charged_milli = ?, margin_milli = ? - ?,
|
||||
ttfb_ms = ?, total_ms = ?, status = ?
|
||||
WHERE request_id = ?"""
|
||||
|
||||
|
||||
class BillingMixin:
|
||||
"""try_hold / settle / void(由 Ledger 继承;依赖 _lock/_connect)。"""
|
||||
|
||||
_lock: threading.Lock
|
||||
|
||||
def _connect(self): # 由 Ledger 提供
|
||||
raise NotImplementedError
|
||||
|
||||
def try_hold(self, request_id: str, key_id: int, student_id: int,
|
||||
model: str, bucket: str, est_milli: int, ts: float) -> bool:
|
||||
"""预扣:锁内读现状 -> 日上限判定 -> 余额原子扣减 -> 写 holding 流水。
|
||||
|
||||
受影响行数为 0 即不足(余额不够)-> False(402);
|
||||
request_id 幂等:重复 hold 直接返回 True(不重复扣)。
|
||||
"""
|
||||
today = _today(ts)
|
||||
with self._lock, self._connect() as conn:
|
||||
dup = conn.execute(
|
||||
"SELECT 1 FROM usage_ledger WHERE request_id = ?",
|
||||
(request_id,)).fetchone()
|
||||
if dup:
|
||||
return True
|
||||
stu = conn.execute(
|
||||
"SELECT balance_milli, daily_cap_milli, spent_today_milli, spent_date"
|
||||
" FROM students WHERE id = ?", (student_id,)).fetchone()
|
||||
if stu is None:
|
||||
return False
|
||||
spent_base = stu["spent_today_milli"] if stu["spent_date"] == today else 0
|
||||
if stu["daily_cap_milli"] > 0 and spent_base + est_milli > stu["daily_cap_milli"]:
|
||||
conn.execute(_SQL_INSUFFICIENT,
|
||||
(request_id, int(ts), key_id, model, bucket))
|
||||
return False
|
||||
cur = conn.execute(_SQL_HOLD, (
|
||||
est_milli, spent_base + est_milli, today, student_id, est_milli))
|
||||
if cur.rowcount == 0:
|
||||
conn.execute(_SQL_INSUFFICIENT,
|
||||
(request_id, int(ts), key_id, model, bucket))
|
||||
return False
|
||||
conn.execute(_SQL_HOLD_MARK,
|
||||
(request_id, int(ts), key_id, model, bucket, est_milli))
|
||||
return True
|
||||
|
||||
def settle(self, request_id: str, actual_milli: int, *,
|
||||
in_miss_tok: int = 0, in_hit_tok: int = 0, out_tok: int = 0,
|
||||
gateway_cached: int = 0, upstream_cost_milli: int = 0,
|
||||
ttfb_ms: Optional[int] = None, total_ms: Optional[int] = None,
|
||||
status: str = "ok") -> bool:
|
||||
"""结算:按真实值更新流水并回补(预扣额 − 实际额)差额。"""
|
||||
with self._lock, self._connect() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT key_id, charged_milli FROM usage_ledger WHERE request_id = ?",
|
||||
(request_id,)).fetchone()
|
||||
if row is None:
|
||||
return False
|
||||
key_id = row["key_id"]
|
||||
est = row["charged_milli"]
|
||||
student = conn.execute(
|
||||
"SELECT student_id FROM proxy_keys WHERE id = ?", (key_id,)).fetchone()
|
||||
if student is None:
|
||||
return False
|
||||
refund = est - actual_milli
|
||||
if refund != 0:
|
||||
conn.execute(_SQL_SETTLE_REFUND,
|
||||
(refund, refund, student["student_id"]))
|
||||
conn.execute(_SQL_SETTLE_USAGE, (
|
||||
in_miss_tok, in_hit_tok, out_tok, gateway_cached,
|
||||
upstream_cost_milli, actual_milli,
|
||||
actual_milli, upstream_cost_milli, ttfb_ms, total_ms,
|
||||
status, request_id))
|
||||
return True
|
||||
|
||||
def void(self, request_id: str, status: str = "aborted") -> bool:
|
||||
"""全额退款(上游失败);流水保留审计。"""
|
||||
with self._lock, self._connect() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT key_id, charged_milli, status FROM usage_ledger"
|
||||
" WHERE request_id = ?", (request_id,)).fetchone()
|
||||
if row is None or row["status"] not in ("holding", "ok"):
|
||||
return False
|
||||
est = row["charged_milli"]
|
||||
student = conn.execute(
|
||||
"SELECT student_id FROM proxy_keys WHERE id = ?",
|
||||
(row["key_id"],)).fetchone()
|
||||
if student is not None and est > 0:
|
||||
conn.execute(_SQL_SETTLE_REFUND,
|
||||
(est, est, student["student_id"]))
|
||||
conn.execute(
|
||||
"UPDATE usage_ledger SET charged_milli = 0, status = ?"
|
||||
" WHERE request_id = ?", (status, request_id))
|
||||
return True
|
||||
|
||||
# ---------- 流水查询 ----------
|
||||
def get_usage(self, request_id: str) -> Optional[Dict[str, Any]]:
|
||||
with self._lock, self._connect() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT * FROM usage_ledger WHERE request_id = ?",
|
||||
(request_id,)).fetchone()
|
||||
return dict(row) if row else None
|
||||
|
||||
def list_usage(self, student_id: Optional[int] = None,
|
||||
limit: int = 50, offset: int = 0) -> List[Dict[str, Any]]:
|
||||
"""流水分页(可按学生过滤,经其名下 key);按学生过滤走联表常量语句。"""
|
||||
if student_id is None:
|
||||
with self._lock, self._connect() as conn:
|
||||
rows = conn.execute(
|
||||
"SELECT * FROM usage_ledger ORDER BY ts DESC LIMIT ? OFFSET ?",
|
||||
(limit, offset)).fetchall()
|
||||
return [dict(r) for r in rows]
|
||||
with self._lock, self._connect() as conn:
|
||||
rows = conn.execute(
|
||||
"SELECT u.* FROM usage_ledger u JOIN proxy_keys k ON k.id = u.key_id"
|
||||
" WHERE k.student_id = ? ORDER BY u.ts DESC LIMIT ? OFFSET ?",
|
||||
(student_id, limit, offset)).fetchall()
|
||||
return [dict(r) for r in rows]
|
||||
+121
-10
@@ -1,21 +1,32 @@
|
||||
"""账本(T-P0:DDL 初始化 + 连接纪律;CRUD/预扣在 T-P1 落地)。
|
||||
"""账本(T-P1 完整):students/keys CRUD + 预扣-结算-退款 + 幂等流水。
|
||||
|
||||
计费两阶段(D-P11,防并发超扣):
|
||||
- try_hold:余额原子扣减(受影响行数为 0 即不足),同时原子检查日消费上限;
|
||||
成功即写入流水一行(status='holding',charged_milli=预扣额)。
|
||||
- settle:按真实 usage 更新流水并回补(预扣额 − 实际额)差额。
|
||||
- void:全额退款(上游失败),流水保留审计。
|
||||
|
||||
工程纪律:
|
||||
- D-P10:sqlite3 是同步库,所有 DB 调用必须经 asyncio.to_thread(由调用方
|
||||
routes 层包装;本模块保持同步实现,可测试性好)。连接 check_same_thread=False
|
||||
+ threading.Lock 串行化(沿用 ReviewQueue 模式:每操作新连接 + 全局锁)。
|
||||
- WAL 模式(§3)。
|
||||
- D-P11:预扣用原子 UPDATE,见 try_hold(T-P1 实现)。
|
||||
- D-P1 毫元整数,本模块不做任何浮点运算(元换算只发生在入口参数转换)。
|
||||
- D-P10:同步实现 + 全局锁每操作连接(ReviewQueue 模式),异步调用方经
|
||||
asyncio.to_thread 包装;WAL 模式(§3)。
|
||||
- 全部数据库访问使用占位符参数化语句,语句为常量,零拼接。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import sqlite3
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from gateway.proxy.billing import BillingMixin
|
||||
from gateway.proxy.ledgerutil import YUAN_TO_MILLI, _today # noqa: F401
|
||||
|
||||
# 四表 + 两索引(§3,一字不改)
|
||||
_SCHEMA = """
|
||||
PRAGMA journal_mode=WAL;
|
||||
|
||||
CREATE TABLE IF NOT EXISTS students(
|
||||
id INTEGER PRIMARY KEY, name TEXT NOT NULL, class TEXT DEFAULT '',
|
||||
status TEXT NOT NULL DEFAULT 'active',
|
||||
@@ -48,11 +59,16 @@ CREATE INDEX IF NOT EXISTS idx_semcache_bucket ON semcache(bucket, ttl_ts);
|
||||
CREATE INDEX IF NOT EXISTS idx_ledger_ts ON usage_ledger(ts);
|
||||
"""
|
||||
|
||||
# 计费两阶段的语句常量与实现见 billing.py(BillingMixin)
|
||||
|
||||
_TABLES = ("students", "proxy_keys", "usage_ledger", "semcache")
|
||||
|
||||
|
||||
class Ledger:
|
||||
"""代理层账本(sqlite WAL;同步实现,调用方负责 to_thread)。"""
|
||||
class Ledger(BillingMixin):
|
||||
"""代理层账本(sqlite WAL;同步实现,调用方负责 to_thread)。
|
||||
|
||||
计费两阶段(try_hold/settle/void)继承自 BillingMixin。
|
||||
"""
|
||||
|
||||
def __init__(self, db_path: str | Path):
|
||||
self.db_path = Path(db_path)
|
||||
@@ -69,14 +85,109 @@ class Ledger:
|
||||
def _connect(self) -> sqlite3.Connection:
|
||||
conn = sqlite3.connect(self.db_path, check_same_thread=False)
|
||||
conn.row_factory = sqlite3.Row
|
||||
conn.execute("PRAGMA journal_mode=WAL")
|
||||
return conn
|
||||
|
||||
def _init_db(self) -> None:
|
||||
with self._lock, self._connect() as conn:
|
||||
conn.executescript(_SCHEMA)
|
||||
|
||||
# ---------- 自省(T-P0 验收用) ----------
|
||||
# ---------- 学生 ----------
|
||||
def upsert_student(self, name: str, klass: str = "",
|
||||
balance_yuan: float = 0.0,
|
||||
daily_cap_yuan: float = 5.0) -> int:
|
||||
"""建学生(元 -> 毫元整数换算只在此入口)。"""
|
||||
milli = int(round(balance_yuan * YUAN_TO_MILLI))
|
||||
cap = int(round(daily_cap_yuan * YUAN_TO_MILLI))
|
||||
with self._lock, self._connect() as conn:
|
||||
cur = conn.execute(
|
||||
"INSERT INTO students(name, class, balance_milli, daily_cap_milli)"
|
||||
" VALUES (?, ?, ?, ?)", (name, klass, milli, cap))
|
||||
return int(cur.lastrowid)
|
||||
|
||||
def topup(self, student_id: int, amount_yuan: float) -> Optional[int]:
|
||||
"""充值;返回新余额(毫元),学生不存在返回 None。"""
|
||||
milli = int(round(amount_yuan * YUAN_TO_MILLI))
|
||||
with self._lock, self._connect() as conn:
|
||||
cur = conn.execute(
|
||||
"UPDATE students SET balance_milli = balance_milli + ? WHERE id = ?",
|
||||
(milli, student_id))
|
||||
if cur.rowcount == 0:
|
||||
return None
|
||||
row = conn.execute(
|
||||
"SELECT balance_milli FROM students WHERE id = ?",
|
||||
(student_id,)).fetchone()
|
||||
return int(row["balance_milli"])
|
||||
|
||||
def set_status(self, student_id: int, status: str) -> bool:
|
||||
with self._lock, self._connect() as conn:
|
||||
cur = conn.execute(
|
||||
"UPDATE students SET status = ? WHERE id = ?", (status, student_id))
|
||||
return cur.rowcount > 0
|
||||
|
||||
def get_student(self, student_id: int) -> Optional[Dict[str, Any]]:
|
||||
with self._lock, self._connect() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT * FROM students WHERE id = ?", (student_id,)).fetchone()
|
||||
return dict(row) if row else None
|
||||
|
||||
# ---------- key ----------
|
||||
def create_key(self, student_id: int, key_hash: str, key_prefix: str,
|
||||
rpm_cap: int = 10, day_cap_req: int = 200,
|
||||
created_ts: Optional[int] = None) -> Optional[int]:
|
||||
"""存 key 哈希;student 不存在返回 None。"""
|
||||
with self._lock, self._connect() as conn:
|
||||
exists = conn.execute(
|
||||
"SELECT 1 FROM students WHERE id = ?", (student_id,)).fetchone()
|
||||
if not exists:
|
||||
return None
|
||||
cur = conn.execute(
|
||||
"INSERT INTO proxy_keys(key_hash, key_prefix, student_id, created_ts,"
|
||||
" rpm_cap, day_cap_req) VALUES (?, ?, ?, ?, ?, ?)",
|
||||
(key_hash, key_prefix, student_id,
|
||||
int(created_ts if created_ts is not None else time.time()),
|
||||
rpm_cap, day_cap_req))
|
||||
return int(cur.lastrowid)
|
||||
|
||||
def find_key(self, key_hash: str) -> Optional[Dict[str, Any]]:
|
||||
"""按哈希查 key(联查学生状态,鉴权主路径)。"""
|
||||
with self._lock, self._connect() as conn:
|
||||
row = conn.execute(
|
||||
"SELECT k.id AS key_id, k.student_id, k.revoked, k.rpm_cap,"
|
||||
" k.day_cap_req, k.req_today, k.req_date,"
|
||||
" s.status AS student_status, s.balance_milli, s.daily_cap_milli,"
|
||||
" s.spent_today_milli, s.spent_date"
|
||||
" FROM proxy_keys k JOIN students s ON s.id = k.student_id"
|
||||
" WHERE k.key_hash = ?", (key_hash,)).fetchone()
|
||||
return dict(row) if row else None
|
||||
|
||||
def revoke_key(self, key_id: int) -> bool:
|
||||
with self._lock, self._connect() as conn:
|
||||
cur = conn.execute(
|
||||
"UPDATE proxy_keys SET revoked = 1 WHERE id = ?", (key_id,))
|
||||
return cur.rowcount > 0
|
||||
|
||||
# ---------- 日限额 ----------
|
||||
def check_and_count(self, key_id: int, student_id: int, now: float) -> bool:
|
||||
"""日请求限额双检:跨日重置(注入日期)+ 原子计数;超限返回 False(429)。"""
|
||||
today = _today(now)
|
||||
with self._lock, self._connect() as conn:
|
||||
conn.execute(
|
||||
"UPDATE students SET spent_today_milli = 0, spent_date = ?"
|
||||
" WHERE spent_date IS NOT ?", (today, today))
|
||||
row = conn.execute(
|
||||
"SELECT req_today, req_date, day_cap_req FROM proxy_keys"
|
||||
" WHERE id = ?", (key_id,)).fetchone()
|
||||
if row is None:
|
||||
return False
|
||||
used = row["req_today"] if row["req_date"] == today else 0
|
||||
if used >= row["day_cap_req"]:
|
||||
return False
|
||||
conn.execute(
|
||||
"UPDATE proxy_keys SET req_today = ?, req_date = ? WHERE id = ?",
|
||||
(used + 1, today, key_id))
|
||||
return True
|
||||
|
||||
# ---------- 自省(测试/验收用) ----------
|
||||
def table_names(self) -> List[str]:
|
||||
"""列出已建表名(测试验收)。"""
|
||||
with self._lock, self._connect() as conn:
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
"""账本共享小工具。"""
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
YUAN_TO_MILLI = 1000
|
||||
|
||||
|
||||
def _today(now: float) -> str:
|
||||
"""本地日期字符串(日限额/日消费重置键;可注入假时钟)。"""
|
||||
return time.strftime("%Y-%m-%d", time.localtime(now))
|
||||
Reference in New Issue
Block a user