Skip to content

Commit f0b217b

Browse files
committed
refactor: 收敛重复样板(错误类型/统计 scope/DB 事务)
- UpstreamHTTPError 提到 provider/base.py:两个 provider 此前逐字重复 status/body/kind(),子类只绑 classify_status - admin_stats:五个端点重复的 `username if is_admin else principal.username` 收进 _scope() - Database.transaction() 上下文:repo/collector 里 14 处 connect→execute→commit 样板收敛,多语句写入(set_pinned)异常时回滚, 此前中途抛错会留下半完成状态;新增提交/回滚两个分支的测试 - TECHNICAL §4:可选能力是「不实现即不定义 + 调用方 getattr 探测 → 400」, 不是原文写的「trae 抛 NotImplementedError」(与实现不符) - TECHNICAL §7:写入统一走 transaction(),无应用层写锁
1 parent 68978aa commit f0b217b

9 files changed

Lines changed: 181 additions & 134 deletions

File tree

‎TECHNICAL.md‎

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -194,7 +194,7 @@ class Provider(Protocol):
194194
def import_credential(self, raw: dict) -> Credential: ...
195195
def refresh(self, cred: DecryptedCred) -> None: ...
196196

197-
# 可选能力(CB 独有;trae 抛 NotImplementedError)
197+
# 可选能力(CB 独有:多账号切换)
198198
def list_accounts(self, cred: DecryptedCred) -> list[Account]: ...
199199
def switch_account(self, cred: DecryptedCred, account_id: str) -> None: ...
200200

@@ -213,6 +213,11 @@ class Provider(Protocol):
213213
约定:
214214
- `DecryptedCred` = db 取出 `data_enc` → Fernet 解密后的 dataclass;provider 不接触 sqlite
215215
- `stream_chat` 只产出 `Event`,产出前先做 HTTP 状态码检查;`classify` 由 executor 调用
216+
- 可选能力**不实现即不定义**(不是抛 `NotImplementedError`):调用方用
217+
`getattr`/`hasattr` 探测,缺失时返回 400「该凭证不支持此操作」,而不是 500。
218+
当前可选集:`start_auth`/`poll_auth`(仅支持 poll 的 provider 才有)、
219+
`complete_callback`(仅 TRAE)、`list_accounts`/`switch_account`(仅 CodeBuddy)、
220+
`credential_from`/`checkin_scope`(刷新与签到任务的能力探测)
216221
- 两个 provider 共用 `engine/sse.py` 的帧解析器(SSE 规范层),事件语义各自映射
217222

218223
---
@@ -266,7 +271,7 @@ class Scheduler:
266271
def pin(self, credential_id: str | None) -> None: ...
267272
```
268273

269-
状态全部落 `credentials` 表(`cooling_until` / `err_count` / `health` / `disabled` / `quota_expiry_ladder`),进程重启不丢冷却状态。写路径无应用层锁:每个写方法直接走当前线程的连接,并发写靠 SQLite WAL + `busy_timeout=5000` 串行化。
274+
状态全部落 `credentials` 表(`cooling_until` / `err_count` / `health` / `disabled` / `quota_expiry_ladder`),进程重启不丢冷却状态。写路径无应用层锁:并发写靠 SQLite WAL + `busy_timeout=5000` 串行化;每个写方法走 `Database.transaction()` 上下文(正常提交、异常回滚),不再散落 `connect()/commit()` 样板。
270275

271276
到期积分只算一处:`expiring_credits()`。选号走 `Candidate.expiry_credits()`,管理台列表走 `GET /api/credentials` 的 `quota_expiring_credits`(窗口值随响应返回 `expiry_window_seconds`),两处共用同一实现,界面数字与选号顺序不会漂移;渠道无到期信息时返回 `null`(不显示),窗口关闭或确实无积分临近过期时返回 `0`(同样不显示)。
272277

@@ -322,7 +327,7 @@ PRAGMA busy_timeout = 5000;
322327
PRAGMA foreign_keys = ON; -- api_keys 之外无外键(users.txt 无表)
323328
```
324329

325-
- 连接:`threading.local()` 每线程一个 `sqlite3.Connection(row_factory=sqlite3.Row)`;引擎与 FastAPI 线程池各自持有自己的连接。没有应用层写锁,写入并发由 SQLite 自身串行化(WAL + `busy_timeout=5000` 下短写足够;确需多语句原子性时用显式事务)
330+
- 连接:`threading.local()` 每线程一个 `sqlite3.Connection(row_factory=sqlite3.Row)`;引擎与 FastAPI 线程池各自持有自己的连接。写入统一走 `Database.transaction()`(`conn.commit()` / 异常 `rollback()`),没有应用层写锁,并发由 SQLite 自身串行化(WAL + `busy_timeout=5000` 下短写足够)
326331
- 加密:`Fernet(base64.urlsafe_b64encode(sha256(APP_SECRET).digest()))`;APP_SECRET 丢失 = 凭证全部不可解,只能重录(Q13 已明示)
327332
- migration:启动时读 `schema.sql` 逐条 `CREATE TABLE IF NOT EXISTS`(只加不改,列注释可改);新增列写进 `migrate._MIGRATION_COLUMNS` 走 `ALTER TABLE ... ADD COLUMN`(重复列名忽略,老库幂等补列),删表写进 `migrate._MIGRATION_DROPS` 走 `DROP TABLE IF EXISTS`(`CREATE TABLE IF NOT EXISTS` 对老库无效,不给删会遗留死表),同时 `SCHEMA_VERSION + 1`,版本记在 `PRAGMA user_version`
328333

‎src/api/admin_stats.py‎

Lines changed: 15 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -4,46 +4,50 @@
44

55
from fastapi import APIRouter, Depends
66

7+
from ..auth.rbac import Principal
78
from .deps import Services, principal_from_request
89

910

11+
def _scope(principal: Principal, username: str | None) -> str | None:
12+
"""admin 传 username 可查全局/指定人;普通用户永远只能看自己。"""
13+
return username if principal.is_admin else principal.username
14+
15+
1016
def create_router(services: Services) -> APIRouter:
1117
router = APIRouter()
1218
stats_query = services.stats_query
1319

1420
@router.get("/api/stats/overview")
1521
async def stats_overview(principal=Depends(principal_from_request),
1622
username: str | None = None, since: int | None = None):
17-
target = username if principal.is_admin else principal.username
18-
return stats_query.overview(username=target, since=since)
23+
return stats_query.overview(username=_scope(principal, username), since=since)
1924

2025
@router.get("/api/stats/by-provider")
2126
async def stats_by_provider(principal=Depends(principal_from_request),
2227
username: str | None = None, since: int | None = None):
23-
target = username if principal.is_admin else principal.username
24-
return {"providers": stats_query.by_provider(username=target, since=since)}
28+
return {"providers": stats_query.by_provider(
29+
username=_scope(principal, username), since=since)}
2530

2631
@router.get("/api/stats/timeline")
2732
async def stats_timeline(principal=Depends(principal_from_request),
2833
username: str | None = None, since: int | None = None,
2934
metric: str = "requests"):
30-
target = username if principal.is_admin else principal.username
31-
return {"points": stats_query.timeline(username=target, since=since, metric=metric)}
35+
return {"points": stats_query.timeline(
36+
username=_scope(principal, username), since=since, metric=metric)}
3237

3338
@router.get("/api/stats/model-timeline")
3439
async def stats_model_timeline(principal=Depends(principal_from_request),
3540
username: str | None = None, since: int | None = None,
3641
metric: str = "requests"):
37-
target = username if principal.is_admin else principal.username
38-
return stats_query.model_timeline(username=target, since=since, metric=metric)
42+
return stats_query.model_timeline(
43+
username=_scope(principal, username), since=since, metric=metric)
3944

4045
@router.get("/api/stats/events")
4146
async def stats_events(principal=Depends(principal_from_request),
4247
username: str | None = None, since: int | None = None,
4348
before: int | None = None, limit: int = 50):
4449
# 明细保留 90 天;单页上限 200,防止一次拉爆
45-
target = username if principal.is_admin else principal.username
46-
return stats_query.events(username=target, since=since, before=before,
47-
limit=max(1, min(limit, 200)))
50+
return stats_query.events(username=_scope(principal, username), since=since,
51+
before=before, limit=max(1, min(limit, 200)))
4852

4953
return router

‎src/db/conn.py‎

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,8 @@
77

88
import sqlite3
99
import threading
10+
from collections.abc import Iterator
11+
from contextlib import contextmanager
1012
from pathlib import Path
1113

1214
PRAGMAS = (
@@ -42,3 +44,19 @@ def close(self) -> None:
4244
if conn is not None:
4345
conn.close()
4446
self._local.conn = None
47+
48+
@contextmanager
49+
def transaction(self) -> Iterator[sqlite3.Connection]:
50+
"""写操作事务:正常退出提交,异常回滚。
51+
52+
多语句写入必须走这个上下文(如「先清 pin 再置 pin」),
53+
否则中途抛错会留下半完成状态;单语句写入也可用,
54+
让「connect → execute → commit」的样板只出现一次。
55+
"""
56+
conn = self.connect()
57+
try:
58+
yield conn
59+
except BaseException:
60+
conn.rollback()
61+
raise
62+
conn.commit()

‎src/db/repo.py‎

Lines changed: 60 additions & 65 deletions
Original file line numberDiff line numberDiff line change
@@ -55,86 +55,81 @@ def add(self, *, provider: str, credential_data: dict, nickname: str = "",
5555
added_by: str = "", now: int | None = None) -> str:
5656
credential_id = _new_id("cred")
5757
payload = json.dumps(credential_data, ensure_ascii=False).encode("utf-8")
58-
self._db.connect().execute(
59-
"INSERT INTO credentials (id, provider, nickname, data_enc, created_at, added_by) "
60-
"VALUES (?,?,?,?,?,?)",
61-
(credential_id, provider, nickname, self._cipher.encrypt(payload),
62-
int(now if now is not None else time.time()), added_by),
63-
)
64-
self._db.connect().commit()
58+
with self._db.transaction() as conn:
59+
conn.execute(
60+
"INSERT INTO credentials (id, provider, nickname, data_enc, created_at, added_by) "
61+
"VALUES (?,?,?,?,?,?)",
62+
(credential_id, provider, nickname, self._cipher.encrypt(payload),
63+
int(now if now is not None else time.time()), added_by),
64+
)
6565
return credential_id
6666

6767
def delete(self, credential_id: str) -> bool:
68-
cursor = self._db.connect().execute("DELETE FROM credentials WHERE id = ?",
69-
(credential_id,))
70-
self._db.connect().commit()
68+
with self._db.transaction() as conn:
69+
cursor = conn.execute("DELETE FROM credentials WHERE id = ?", (credential_id,))
7170
return cursor.rowcount > 0
7271

7372
def set_enabled(self, credential_id: str, enabled: bool) -> bool:
74-
cursor = self._db.connect().execute(
75-
"UPDATE credentials SET enabled = ? WHERE id = ?", (1 if enabled else 0, credential_id))
76-
self._db.connect().commit()
73+
with self._db.transaction() as conn:
74+
cursor = conn.execute(
75+
"UPDATE credentials SET enabled = ? WHERE id = ?",
76+
(1 if enabled else 0, credential_id))
7777
return cursor.rowcount > 0
7878

7979
def set_pinned(self, credential_id: str | None) -> None:
80-
conn = self._db.connect()
81-
conn.execute("UPDATE credentials SET pinned = 0")
82-
if credential_id is not None:
83-
conn.execute("UPDATE credentials SET pinned = 1 WHERE id = ?", (credential_id,))
84-
conn.commit()
80+
with self._db.transaction() as conn:
81+
conn.execute("UPDATE credentials SET pinned = 0")
82+
if credential_id is not None:
83+
conn.execute("UPDATE credentials SET pinned = 1 WHERE id = ?", (credential_id,))
8584

8685
def revive(self, credential_id: str) -> bool:
8786
"""解除硬禁用(session 死亡)与冷却,允许凭证重新参与调度。
8887
8988
没有这个入口时,凭证一旦因 session 失效被硬禁用就只能删除重建,
9089
重新登录后也无法复用同一条记录。
9190
"""
92-
cursor = self._db.connect().execute(
93-
"UPDATE credentials SET disabled = 0, disabled_reason = NULL, cooling_until = NULL, "
94-
"err_count = 0 WHERE id = ?", (credential_id,))
95-
self._db.connect().commit()
91+
with self._db.transaction() as conn:
92+
cursor = conn.execute(
93+
"UPDATE credentials SET disabled = 0, disabled_reason = NULL, "
94+
"cooling_until = NULL, err_count = 0 WHERE id = ?", (credential_id,))
9695
return cursor.rowcount > 0
9796

9897
def save_error(self, credential_id: str, outcome: ErrorOutcome) -> None:
99-
conn = self._db.connect()
100-
if outcome.disabled:
101-
conn.execute(
102-
"UPDATE credentials SET disabled = 1, disabled_reason = ?, err_count = 0, "
103-
"cooling_until = NULL WHERE id = ?", ("session dead", credential_id))
104-
else:
105-
conn.execute(
106-
"UPDATE credentials SET cooling_until = ?, err_count = ? WHERE id = ?",
107-
(outcome.cooling_until, outcome.err_count, credential_id))
108-
conn.commit()
98+
with self._db.transaction() as conn:
99+
if outcome.disabled:
100+
conn.execute(
101+
"UPDATE credentials SET disabled = 1, disabled_reason = ?, err_count = 0, "
102+
"cooling_until = NULL WHERE id = ?", ("session dead", credential_id))
103+
else:
104+
conn.execute(
105+
"UPDATE credentials SET cooling_until = ?, err_count = ? WHERE id = ?",
106+
(outcome.cooling_until, outcome.err_count, credential_id))
109107

110108
def save_success(self, credential_id: str) -> None:
111-
conn = self._db.connect()
112-
conn.execute("UPDATE credentials SET err_count = 0 WHERE id = ?", (credential_id,))
113-
conn.commit()
109+
with self._db.transaction() as conn:
110+
conn.execute("UPDATE credentials SET err_count = 0 WHERE id = ?", (credential_id,))
114111

115112
def save_credential_data(self, credential_id: str, credential_data: dict) -> None:
116113
payload = json.dumps(credential_data, ensure_ascii=False).encode("utf-8")
117-
conn = self._db.connect()
118-
conn.execute("UPDATE credentials SET data_enc = ? WHERE id = ?",
119-
(self._cipher.encrypt(payload), credential_id))
120-
conn.commit()
114+
with self._db.transaction() as conn:
115+
conn.execute("UPDATE credentials SET data_enc = ? WHERE id = ?",
116+
(self._cipher.encrypt(payload), credential_id))
121117

122118
def save_quota(self, credential_id: str, quota: Quota) -> None:
123-
conn = self._db.connect()
124-
conn.execute(
125-
"UPDATE credentials SET quota_remaining = ?, quota_total = ?, quota_cycle_end = ?, "
126-
"quota_expiry_ladder = ?, quota_probed_at = ?, health = ? WHERE id = ?",
127-
(quota.remaining, quota.total, quota.cycle_end,
128-
_ladder_text(quota.expiry_ladder), quota.probed_at,
129-
health_score(quota), credential_id),
130-
)
131-
conn.commit()
119+
with self._db.transaction() as conn:
120+
conn.execute(
121+
"UPDATE credentials SET quota_remaining = ?, quota_total = ?, "
122+
"quota_cycle_end = ?, quota_expiry_ladder = ?, quota_probed_at = ?, health = ? "
123+
"WHERE id = ?",
124+
(quota.remaining, quota.total, quota.cycle_end,
125+
_ladder_text(quota.expiry_ladder), quota.probed_at,
126+
health_score(quota), credential_id),
127+
)
132128

133129
def mark_probe_failed(self, credential_id: str, now: int | None = None) -> None:
134-
conn = self._db.connect()
135-
conn.execute("UPDATE credentials SET quota_probed_at = ?, health = NULL WHERE id = ?",
136-
(int(now if now is not None else time.time()), credential_id))
137-
conn.commit()
130+
with self._db.transaction() as conn:
131+
conn.execute("UPDATE credentials SET quota_probed_at = ?, health = NULL WHERE id = ?",
132+
(int(now if now is not None else time.time()), credential_id))
138133

139134
# ------------------------------------------------------------- 读取
140135

@@ -210,13 +205,13 @@ def create(self, username: str, name: str = "", now: int | None = None) -> dict[
210205
plaintext = generate_api_key()
211206
key_id = _new_id("key")
212207
created_at = int(now if now is not None else time.time())
213-
self._db.connect().execute(
214-
"INSERT INTO api_keys (id, username, name, key_digest, preview, created_at) "
215-
"VALUES (?,?,?,?,?,?)",
216-
(key_id, username, name, digest_api_key(plaintext), preview_api_key(plaintext),
217-
created_at),
218-
)
219-
self._db.connect().commit()
208+
with self._db.transaction() as conn:
209+
conn.execute(
210+
"INSERT INTO api_keys (id, username, name, key_digest, preview, created_at) "
211+
"VALUES (?,?,?,?,?,?)",
212+
(key_id, username, name, digest_api_key(plaintext), preview_api_key(plaintext),
213+
created_at),
214+
)
220215
return {"id": key_id, "username": username, "name": name, "api_key": plaintext,
221216
"preview": preview_api_key(plaintext), "created_at": created_at}
222217

@@ -226,9 +221,9 @@ def verify(self, api_key: str) -> str | None:
226221
"SELECT id, username FROM api_keys WHERE key_digest = ?", (digest,)).fetchone()
227222
if row is None:
228223
return None
229-
self._db.connect().execute("UPDATE api_keys SET last_used_at = ? WHERE id = ?",
230-
(int(time.time()), row["id"]))
231-
self._db.connect().commit()
224+
with self._db.transaction() as conn:
225+
conn.execute("UPDATE api_keys SET last_used_at = ? WHERE id = ?",
226+
(int(time.time()), row["id"]))
232227
return row["username"]
233228

234229
def list_for(self, username: str) -> list[dict[str, Any]]:
@@ -238,7 +233,7 @@ def list_for(self, username: str) -> list[dict[str, Any]]:
238233
return [dict(row) for row in rows]
239234

240235
def delete(self, key_id: str, username: str) -> bool:
241-
cursor = self._db.connect().execute(
242-
"DELETE FROM api_keys WHERE id = ? AND username = ?", (key_id, username))
243-
self._db.connect().commit()
236+
with self._db.transaction() as conn:
237+
cursor = conn.execute(
238+
"DELETE FROM api_keys WHERE id = ? AND username = ?", (key_id, username))
244239
return cursor.rowcount > 0

‎src/provider/base.py‎

Lines changed: 20 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -6,9 +6,10 @@
66

77
from __future__ import annotations
88

9+
from collections.abc import Callable
910
from dataclasses import dataclass
1011
from enum import StrEnum
11-
from typing import Protocol, runtime_checkable
12+
from typing import ClassVar, Protocol, runtime_checkable
1213

1314

1415
class EventKind(StrEnum):
@@ -39,6 +40,24 @@ def body_hint(body: bytes, limit: int = 160) -> str:
3940
return text[:limit]
4041

4142

43+
class UpstreamHTTPError(Exception):
44+
"""上游非 2xx;两个 provider 共用同一形状(status + 原始 body)。
45+
46+
`kind()` 由子类绑定各自的 classify_status,因为 1005 等业务码的
47+
判定规则是上游协议私有的。
48+
"""
49+
50+
classify_status: ClassVar[Callable[[int, bytes], ErrKind]]
51+
52+
def __init__(self, status: int, body: bytes) -> None:
53+
self.status = status
54+
self.body = body
55+
super().__init__(f"upstream http {status}: {body_hint(body)}")
56+
57+
def kind(self) -> ErrKind:
58+
return self.classify_status(self.status, self.body)
59+
60+
4261
@dataclass(slots=True)
4362
class Usage:
4463
input_tokens: int | None = None

‎src/provider/codebuddy/client.py‎

Lines changed: 5 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,8 @@
1717
import httpx
1818

1919
from ...engine.sse import iter_frames
20-
from ...provider.base import ErrKind, Event, Model, Quota, body_hint
20+
from ...provider import base
21+
from ...provider.base import ErrKind, Event, Model, Quota
2122
from . import events as cb_events
2223
from .credential import CodeBuddyCredential, parse_credential
2324
from .events import UpstreamProtocolViolation
@@ -435,14 +436,10 @@ def _cached_refresh(client: CodeBuddyClient):
435436
CODEBUDDY_IDE_VERSION = "1.42.0"
436437

437438

438-
class UpstreamHTTPError(Exception):
439-
def __init__(self, status: int, body: bytes) -> None:
440-
self.status = status
441-
self.body = body
442-
super().__init__(f"upstream http {status}: {body_hint(body)}")
439+
class UpstreamHTTPError(base.UpstreamHTTPError):
440+
"""CodeBuddy 上游非 2xx;kind() 走 CB 的 1005/400 规则。"""
443441

444-
def kind(self) -> ErrKind:
445-
return cb_events.classify_status(self.status, self.body)
442+
classify_status = staticmethod(cb_events.classify_status)
446443

447444

448445
@dataclass(slots=True)

0 commit comments

Comments
 (0)