- 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
204 lines
9.1 KiB
Python
204 lines
9.1 KiB
Python
"""账本(T-P1 完整):students/keys CRUD + 预扣-结算-退款 + 幂等流水。
|
||
|
||
计费两阶段(D-P11,防并发超扣):
|
||
- try_hold:余额原子扣减(受影响行数为 0 即不足),同时原子检查日消费上限;
|
||
成功即写入流水一行(status='holding',charged_milli=预扣额)。
|
||
- settle:按真实 usage 更新流水并回补(预扣额 − 实际额)差额。
|
||
- void:全额退款(上游失败),流水保留审计。
|
||
|
||
工程纪律:
|
||
- 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',
|
||
balance_milli INTEGER NOT NULL DEFAULT 0,
|
||
daily_cap_milli INTEGER NOT NULL DEFAULT 5000,
|
||
spent_today_milli INTEGER NOT NULL DEFAULT 0, spent_date TEXT DEFAULT '');
|
||
|
||
CREATE TABLE IF NOT EXISTS proxy_keys(
|
||
id INTEGER PRIMARY KEY, key_hash TEXT UNIQUE NOT NULL, key_prefix TEXT NOT NULL,
|
||
student_id INTEGER NOT NULL REFERENCES students(id),
|
||
created_ts INTEGER NOT NULL, revoked INTEGER NOT NULL DEFAULT 0,
|
||
rpm_cap INTEGER NOT NULL DEFAULT 10, day_cap_req INTEGER NOT NULL DEFAULT 200,
|
||
req_today INTEGER NOT NULL DEFAULT 0, req_date TEXT DEFAULT '');
|
||
|
||
CREATE TABLE IF NOT EXISTS usage_ledger(
|
||
request_id TEXT PRIMARY KEY, ts INTEGER NOT NULL, key_id INTEGER NOT NULL,
|
||
model TEXT NOT NULL, bucket TEXT NOT NULL DEFAULT 'default',
|
||
in_miss_tok INTEGER NOT NULL DEFAULT 0, in_hit_tok INTEGER NOT NULL DEFAULT 0,
|
||
out_tok INTEGER NOT NULL DEFAULT 0, gateway_cached INTEGER NOT NULL DEFAULT 0,
|
||
upstream_cost_milli INTEGER NOT NULL DEFAULT 0, charged_milli INTEGER NOT NULL DEFAULT 0,
|
||
margin_milli INTEGER NOT NULL DEFAULT 0, ttfb_ms INTEGER, total_ms INTEGER,
|
||
status TEXT NOT NULL);
|
||
|
||
CREATE TABLE IF NOT EXISTS semcache(
|
||
cache_key TEXT PRIMARY KEY,
|
||
bucket TEXT NOT NULL, q_norm TEXT NOT NULL, answer TEXT NOT NULL, model TEXT NOT NULL,
|
||
created_ts INTEGER NOT NULL, ttl_ts INTEGER NOT NULL,
|
||
doc_version INTEGER NOT NULL DEFAULT 1, hits INTEGER NOT NULL DEFAULT 0);
|
||
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(BillingMixin):
|
||
"""代理层账本(sqlite WAL;同步实现,调用方负责 to_thread)。
|
||
|
||
计费两阶段(try_hold/settle/void)继承自 BillingMixin。
|
||
"""
|
||
|
||
def __init__(self, db_path: str | Path):
|
||
self.db_path = Path(db_path)
|
||
self.db_path.parent.mkdir(parents=True, exist_ok=True)
|
||
self._lock = threading.Lock()
|
||
self._init_db()
|
||
|
||
# ---------- 初始化 ----------
|
||
@classmethod
|
||
def init_db(cls, db_path: str | Path) -> "Ledger":
|
||
"""工厂(§6 签名):建库建表(幂等)。"""
|
||
return cls(db_path)
|
||
|
||
def _connect(self) -> sqlite3.Connection:
|
||
conn = sqlite3.connect(self.db_path, check_same_thread=False)
|
||
conn.row_factory = sqlite3.Row
|
||
return conn
|
||
|
||
def _init_db(self) -> None:
|
||
with self._lock, self._connect() as conn:
|
||
conn.executescript(_SCHEMA)
|
||
|
||
# ---------- 学生 ----------
|
||
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:
|
||
rows = conn.execute(
|
||
"SELECT name FROM sqlite_master WHERE type='table'").fetchall()
|
||
return [r["name"] for r in rows]
|
||
|
||
def index_names(self) -> List[str]:
|
||
with self._lock, self._connect() as conn:
|
||
rows = conn.execute(
|
||
"SELECT name FROM sqlite_master WHERE type='index' AND name LIKE 'idx_%'"
|
||
).fetchall()
|
||
return [r["name"] for r in rows]
|