文件
workbuddy-portal/workbuddy_portal/collect.py
T
wangchuanli 86631ae7ab chore: 项目定名为 workbuddy-portal,容器化并补齐文档体系
## 项目定名
- 目录 wb_usage_portal → workbuddy-portal
- Python 包 wb_usage → workbuddy_portal(含 session cookie 名)
- 界面品牌统一为 WorkBuddy Portal;项目标识收敛到 config 单一来源

## 容器化
- Dockerfile:多阶段构建,依赖层与源码解耦;非 root(uid 1000);内置健康检查
- docker-compose.yml:单服务 + 绑定挂载 data/logs + 日志轮转 + TZ
- docker/entrypoint.sh:幂等初始化 → exec serve(LF 行尾,已由 .gitattributes 锁定)
- docker/healthcheck.py:纯标准库探活 /login(slim 镜像无 curl)
- .dockerignore / .env.example;数据目录可用 WB_DATA_DIR 等环境变量覆盖

## 文档
- docs/USER-GUIDE.md    用户使用手册(含 9 张真实界面截图)
- docs/DEPLOYMENT.md    部署运维(Docker / 裸机 / 反代 / 备份 / 推 Gitea 注册表)
- docs/ARCHITECTURE.md  架构与设计说明(含已知坑与红线、验证体系)
- docs/API.md           接口参考(路径 / 参数 / 返回结构 / 错误码)
- docs/FAQ.md           常见问题;docs/CHANGELOG.md 变更日志

## 修复缺陷(8)
1. /records/export 必然 500:生成器在请求上下文销毁后才迭代,改用自建连接
2. 大屏页图表全白:相对路径把 echarts.min.js 解析成 /vendor/... → 404
3. /users 500:路由已注册但模板缺失
4. 明细页日期筛选失效:视图传 f.frm、模板读 f.from
5. 配置页维护按钮全死:调用了不存在的 WBU.bindMaint()
6. 审计只能看最近 40 条:LIMIT 写死
7. 明细页多跑一条无用 SELECT:day_list() 取了没人用
8. 登录页锁定阈值未从配置注入

## 安全加固
- 新增 safe_next():拒绝 //evil.com 等协议相对 URL 的开放重定向
- 缺 CSRF 的写请求统一 400
- 默认开启云端 HTTPS 证书校验(ssl_verify=1);Cookie 是账号凭证
- 登录失败计数表加上限与 TTL
- /logout 拆分为 POST(执行) + GET(仅提示),防 <img src=/logout> 静默退出
- settings 内部簿记键 slot:* 读写两侧过滤,不再从 /api/settings 泄漏

## 内部质量与工具
- 设置项写时校验 + 读时兜底,杜绝「一个手滑的数字让采集整个跑不起来」
- 全局 ValueError → 400:手写 query string 不再暴露 500 页面
- CSV 导出改 csv.writer 流式写入(原手工拼串,字段含逗号会串列)
- bundle 明细加 20000 上限并回传 recordsTotal/recordsTruncated,不静默丢数据
- tools/smoke.py 离线回归 99 项;tools/check_live.py 真实 HTTP 56 项
- tools/shots.py Playwright 逐页截图 + JS 报错收集

## 验证
- compileall 通过;smoke 99/99;对容器实例 check_live 56/56;截图 0 JS 报错
- 容器内采集实测成功(trigger=startup 补跑:新增 11 条)
2026-09-14 14:55:50 +08:00

450 行
19 KiB
Python
原始文件 Blame 文件历史

此文件含有模棱两可的 Unicode 字符
此文件含有可能会与其他字符混淆的 Unicode 字符。 如果您是想特意这样的,可以安全地忽略该警告。 使用 Escape 按钮显示他们。
# -*- coding: utf-8 -*-
"""采集主流程:云端增量 -> SQLite(去重、断点、漂移校验、运行记录)。
与原 fetch_usage.py 的差别:
* 存档正本从 CSV 换成 SQLite,去重由 `ON CONFLICT(request_id)` 承担
* 断点由 `SELECT MAX(ts)` 承担,不再需要全量读入内存
* 每次运行落一条 collect_runs 记录,页面据此展示任务历史与日志
* 采集互斥用文件锁,保证「单写者」——SQLite 只允许一个写进程
"""
import csv
import os
import time
from datetime import datetime, timedelta
from . import client, config, db
from .client import ApiError
FIELDS = ["RequestID", "积分消耗", "User Prompt", "模型", "客户端", "时间"]
LOCK_PATH = os.path.join(config.DATA_DIR, "collect.lock")
LOCK_STALE_SECONDS = 30 * 60 # 锁超过 30 分钟视为僵尸锁,可抢占
class Busy(Exception):
"""已有采集在跑。"""
# ---------------- 互斥锁 ----------------
class _Lock:
def __init__(self, path=LOCK_PATH):
self.path = path
self.fd = None
def __enter__(self):
config.ensure_dirs()
if os.path.exists(self.path):
try:
age = time.time() - os.path.getmtime(self.path)
except OSError:
age = 0
if age < LOCK_STALE_SECONDS:
raise Busy("已有采集任务正在运行(锁文件 %s,%.0f 秒前创建)"
% (self.path, age))
os.remove(self.path) # 僵尸锁
try:
self.fd = os.open(self.path, os.O_CREAT | os.O_EXCL | os.O_WRONLY)
os.write(self.fd, ("%d %s" % (os.getpid(), db.now_str())).encode())
except FileExistsError:
raise Busy("已有采集任务正在运行")
return self
def __exit__(self, *a):
try:
if self.fd is not None:
os.close(self.fd)
except OSError:
pass
self.fd = None
try:
os.remove(self.path)
except OSError:
pass
# ---------------- 运行记录 ----------------
def start_run(conn, trigger):
cur = conn.execute("INSERT INTO collect_runs(trigger,status,started_at) VALUES(?,?,?)",
(trigger, "running", db.now_str()))
return cur.lastrowid
def finish_run(conn, run_id, status, **kw):
fields = ["finished_at", "duration_ms", "win_from", "win_to", "fetched",
"added", "dup", "total", "conflicts", "exit_code", "message", "detail"]
sets = ["status=?", "finished_at=?"]
vals = [status, db.now_str()]
for f in fields[1:]:
if f in kw:
sets.append("%s=?" % f)
vals.append(kw[f])
vals.append(run_id)
conn.execute("UPDATE collect_runs SET %s WHERE id=?" % ",".join(sets), vals)
# ---------------- 入库 ----------------
_UPSERT = """
INSERT INTO usage_records(request_id,ts,day,hour,model,client,credits,prompt,
first_seen,last_seen,cloud_ts)
VALUES(:request_id,:ts,:day,:hour,:model,:client,:credits,:prompt,:first_seen,:last_seen,:cloud_ts)
ON CONFLICT(request_id) DO UPDATE SET
last_seen = excluded.last_seen,
cloud_ts = excluded.cloud_ts,
ts = CASE WHEN excluded.ts <> '' AND excluded.ts < usage_records.ts
THEN excluded.ts ELSE usage_records.ts END,
day = CASE WHEN excluded.ts <> '' AND excluded.ts < usage_records.ts
THEN excluded.day ELSE usage_records.day END,
hour = CASE WHEN excluded.ts <> '' AND excluded.ts < usage_records.ts
THEN excluded.hour ELSE usage_records.hour END,
prompt = CASE WHEN COALESCE(usage_records.prompt,'') = ''
THEN excluded.prompt ELSE usage_records.prompt END,
model = CASE WHEN COALESCE(usage_records.model,'') IN ('','-')
THEN excluded.model ELSE usage_records.model END,
client = CASE WHEN COALESCE(usage_records.client,'') IN ('','-')
THEN excluded.client ELSE usage_records.client END
"""
def _row_dict(n):
ts = (n.get("ts") or "").strip()[:19]
hh = ts[11:13]
return {
"request_id": n["request_id"],
"ts": ts,
"day": ts[:10],
"hour": int(hh) if hh.isdigit() else 0,
"model": n.get("model") or "-",
"client": n.get("client") or "-",
"credits": n.get("credits") or 0.0,
"prompt": n.get("prompt") or "",
"first_seen": db.now_str(),
"last_seen": db.now_str(),
"cloud_ts": ts,
}
def upsert(conn, normalized, drift_tolerance=5, log=None):
"""按 request_id 去重写入。返回 (added, dup, conflicts) 与逐行警告。
每 200 条一个事务(连接是 autocommit,不显式 BEGIN 的话每条 INSERT 都要
单独 fsync)。upsert 本身幂等,所以按块提交是安全的。
"""
recs = []
for n in normalized:
if not n.get("request_id") or not n.get("ts"):
continue
recs.append(_row_dict(n))
added = dup = 0
conflicts = []
for i in range(0, len(recs), 200):
chunk = recs[i:i + 200]
ids = [r["request_id"] for r in chunk]
ph = ",".join("?" * len(ids))
own_tx = not conn.in_transaction
if own_tx:
conn.execute("BEGIN")
try:
old = {x["request_id"]: x["ts"] for x in
conn.execute("SELECT request_id,ts FROM usage_records WHERE request_id IN (%s)" % ph, ids)}
for r in chunk:
prev = old.get(r["request_id"])
if prev is None:
added += 1
else:
dup += 1
try:
delta = (datetime.strptime(prev, "%Y-%m-%d %H:%M:%S")
- datetime.strptime(r["ts"], "%Y-%m-%d %H:%M:%S")).total_seconds()
except (ValueError, TypeError):
delta = 0
if delta > drift_tolerance * 60:
conflicts.append((r["request_id"], prev, r["ts"]))
conn.executemany(_UPSERT, chunk)
if own_tx:
conn.execute("COMMIT")
except Exception:
if own_tx:
conn.execute("ROLLBACK")
raise
if conflicts and log:
log("[warn] %d 条 RequestID 相同且云端开始时间早于本地超过 %d 分钟(保留本地最早时间,未覆盖):"
% (len(conflicts), drift_tolerance))
for rid, t_old, t_new in conflicts[:10]:
log(" %s: 本地 %s / 云端 %s" % (rid, t_old, t_new))
if len(conflicts) > 10:
log(" ... 另有 %d 条" % (len(conflicts) - 10))
return added, dup, conflicts
def record_count(conn):
return conn.execute("SELECT COUNT(*) FROM usage_records").fetchone()[0]
def last_ts(conn):
return conn.execute("SELECT MAX(ts) FROM usage_records").fetchone()[0]
# ---------------- 主同步 ----------------
def sync(conn, trigger="manual", from_dt=None, to_dt=None, verify_days=None,
do_write=True, log=None):
"""增量同步。
trigger: manual | schedule | cli | startup(写进 collect_runs 便于区分来源)
返回 result dict;异常时抛出 ApiError(调用方决定如何展示)。
"""
lines = []
def _log(msg):
lines.append(msg)
if log:
log(msg)
s = db.get_settings(conn)
cookie = (s.get("cookie") or "").strip() or os.environ.get("WB_COOKIE", "").strip()
ua = (s.get("user_agent") or "").strip() or os.environ.get("WB_UA", "").strip()
api_base = s.get("api_base") or config.API_BASE
api_path = s.get("api_path") or config.API_PATH
# 一律走 db.get_int/get_float:settings 表的值由后台页面自由输入,
# 直接 int() 会让一个手滑的字符把整条采集链路打断(历史 bug)。
page_size = db.get_int(conn, "page_size", 200)
rewind = db.get_int(conn, "rewind_minutes", 2)
drift = db.get_int(conn, "drift_tolerance_minutes", 5)
max_prompt = db.get_int(conn, "max_prompt", 0)
timeout = db.get_int(conn, "timeout", 30)
ssl_verify = db.get_bool(conn, "ssl_verify", True)
if verify_days is None:
verify_days = db.get_int(conn, "verify_days", 0)
# 夹到合法区间,避免历史脏数据(如超大的 page_size)把云端打爆
lo, hi, _ = config.NUM_SETTINGS["page_size"]
page_size = max(lo, min(hi, page_size))
run_id = start_run(conn, trigger) if do_write else None
t0 = time.time()
base = {"win_from": None, "win_to": None, "fetched": 0,
"added": 0, "dup": 0, "conflicts": 0}
if not cookie:
msg = "未配置 Cookie,请到「配置管理」页粘贴,或设置环境变量 WB_COOKIE"
_log("[error] " + msg)
if run_id:
finish_run(conn, run_id, "error", exit_code=2, message=msg,
duration_ms=int((time.time() - t0) * 1000), detail="\n".join(lines), **base)
raise ApiError(msg)
total_before = record_count(conn)
tail = last_ts(conn)
now = datetime.now()
_log("存档:%d 条%s" % (total_before, (",最后记录 " + tail) if tail else "(空)"))
# 断点 = 本地最后一条记录时间(回退 rewind 分钟防边界遗漏)
if from_dt:
start = from_dt
elif tail:
start = datetime.strptime(tail, "%Y-%m-%d %H:%M:%S") - timedelta(minutes=rewind)
else:
start = now - timedelta(days=30)
_log("存档为空,默认回填最近 30 天")
end = to_dt or now
if start >= end:
start = end - timedelta(minutes=rewind)
_log("同步区间:%s ~ %s" % (start.strftime("%Y-%m-%d %H:%M:%S"),
end.strftime("%Y-%m-%d %H:%M:%S")))
try:
raw, _totals = client.fetch_range(start, end, cookie, ua, api_base, api_path,
page_size=page_size, timeout=timeout,
log=_log, ssl_verify=ssl_verify)
except ApiError as e:
msg = str(e)
_log("[error] " + msg)
if run_id:
finish_run(conn, run_id, "error", exit_code=3 if e.cookie_expired else 5,
message=msg, duration_ms=int((time.time() - t0) * 1000),
detail="\n".join(lines), win_from=start.strftime("%Y-%m-%d %H:%M:%S"),
win_to=end.strftime("%Y-%m-%d %H:%M:%S"), **{k: v for k, v in base.items()
if k not in ("win_from", "win_to")})
raise
new_rows = [client.normalize(r, max_prompt=max_prompt) for r in raw]
_log("云端返回:%d 条" % len(raw))
added, dup, conflicts = upsert(conn, new_rows, drift_tolerance=drift, log=_log)
# 整日完整性校验(默认关闭;用于排查缺记录)
if verify_days > 0:
days = [r["day"] for r in conn.execute(
"SELECT DISTINCT day FROM usage_records ORDER BY day DESC LIMIT ?", (verify_days,))]
_log("完整性校验:最近 %d 天" % len(days))
for d in sorted(days):
local = conn.execute("SELECT COUNT(*) FROM usage_records WHERE day=?", (d,)).fetchone()[0]
d0 = datetime.strptime(d, "%Y-%m-%d")
try:
raw2, t2 = client.fetch_range(d0, d0.replace(hour=23, minute=59, second=59),
cookie, ua, api_base, api_path,
page_size=page_size, timeout=timeout,
ssl_verify=ssl_verify)
except ApiError as e:
_log(" %s 校验失败:%s" % (d, e))
continue
cloud = t2.get(d, 0)
if local < cloud:
a2, _, _ = upsert(conn, [client.normalize(r, max_prompt=max_prompt) for r in raw2],
drift_tolerance=drift)
added += a2
_log(" %s:云端 %d / 本地 %d → 补入 %d 条" % (d, cloud, local, a2))
else:
_log(" %s:云端 %d / 本地 %d OK" % (d, cloud, local))
total_after = record_count(conn)
status = "warn" if conflicts else "ok"
msg = "新增 %d 条,重复 %d 条,存档共 %d 条" % (added, dup, total_after)
_log("RESULT: added=%d dup=%d total=%d" % (added, dup, total_after))
if run_id:
finish_run(conn, run_id, status,
duration_ms=int((time.time() - t0) * 1000),
win_from=start.strftime("%Y-%m-%d %H:%M:%S"),
win_to=end.strftime("%Y-%m-%d %H:%M:%S"),
fetched=len(raw), added=added, dup=dup, total=total_after,
conflicts=len(conflicts), exit_code=0, message=msg,
detail="\n".join(lines))
return {"status": status, "added": added, "dup": dup, "fetched": len(raw),
"total": total_after, "conflicts": len(conflicts),
"win_from": start.strftime("%Y-%m-%d %H:%M:%S"),
"win_to": end.strftime("%Y-%m-%d %H:%M:%S"),
"message": msg, "lines": lines, "run_id": run_id}
def run_sync(trigger="manual", **kw):
"""带锁的同步入口(供 CLI / 调度器 / 页面手动触发共用)。"""
with _Lock():
conn = db.thread_conn()
return sync(conn, trigger=trigger, **kw)
# ---------------- 补全 / 导入 / 导出 ----------------
def fill_prompt(conn, log=print):
"""补全缺失的 User Prompt(官网导出的 xlsx 会丢约 22%,云端仍保留)。"""
s = db.get_settings(conn)
cookie = (s.get("cookie") or "").strip() or os.environ.get("WB_COOKIE", "").strip()
ua = (s.get("user_agent") or "").strip()
max_prompt = db.get_int(conn, "max_prompt", 0)
if not cookie:
raise ApiError("未配置 Cookie")
todo = conn.execute("SELECT request_id, day FROM usage_records "
"WHERE COALESCE(prompt,'')='' ORDER BY day").fetchall()
if not todo:
log("没有缺失的 User Prompt")
return 0
days = sorted({r["day"] for r in todo})
log("待补全 %d 条,分布在 %d 天" % (len(todo), len(days)))
pool = {}
for d in days:
d0 = datetime.strptime(d, "%Y-%m-%d")
raw, _ = client.fetch_range(d0, d0.replace(hour=23, minute=59, second=59),
cookie, ua, s.get("api_base") or config.API_BASE,
s.get("api_path") or config.API_PATH,
page_size=db.get_int(conn, "page_size", 200),
timeout=db.get_int(conn, "timeout", 30),
ssl_verify=db.get_bool(conn, "ssl_verify", True))
for r in raw:
rid = (r.get("requestId") or "").strip()
if rid:
pool[rid] = client.normalize(r, max_prompt=max_prompt)["prompt"]
log(" %s 云端 %d 条" % (d, len(raw)))
n = 0
for row in todo:
p = pool.get(row["request_id"], "")
if p:
conn.execute("UPDATE usage_records SET prompt=? WHERE request_id=?", (p, row["request_id"]))
n += 1
left = conn.execute("SELECT COUNT(*) FROM usage_records WHERE COALESCE(prompt,'')=''").fetchone()[0]
log("补全 %d 条,仍为空 %d 条" % (n, left))
return n
def import_xlsx(conn, path, log=print):
"""从官网「用量明细 - 导出」的 xlsx 合入(按 request_id 去重)。"""
try:
import openpyxl
except ImportError:
raise ApiError("需要 openpyxl:pip install openpyxl")
s = db.get_settings(conn)
max_prompt = db.get_int(conn, "max_prompt", 0)
wb = openpyxl.load_workbook(path, read_only=True, data_only=True)
it = wb.worksheets[0].iter_rows(values_only=True)
header = [str(c).strip() if c is not None else "" for c in next(it)]
idx = {n: i for i, n in enumerate(header)}
def col(*names):
for n in names:
if n in idx:
return idx[n]
raise ApiError("xlsx 缺少列 %s,实际表头:%s" % (names, header))
i_rid, i_cr, i_px = col("RequestID"), col("积分消耗"), col("User Prompt")
i_m, i_cl, i_t = col("模型"), col("客户端"), col("时间")
rows = []
for row in it:
if row is None or row[i_t] is None:
continue
t = str(row[i_t])[:19]
try:
cr = float(row[i_cr] or 0)
except (TypeError, ValueError):
cr = 0.0
px = " ".join(str(row[i_px] or "").split())
if max_prompt and len(px) > max_prompt:
px = px[:max_prompt]
rid = str(row[i_rid] or "").strip() or ("xlsx-%s-%d" % (t, len(rows)))
rows.append({"request_id": rid, "ts": t, "credits": round(cr, 2), "prompt": px,
"model": str(row[i_m] or "-").strip() or "-",
"client": str(row[i_cl] or "-").strip() or "-"})
added, dup, _ = upsert(conn, rows)
log("[xlsx] 读取 %d 条,去重后新增 %d 条,存档共 %d 条" % (len(rows), added, record_count(conn)))
return added
def export_csv(conn, path=None):
"""导出与旧存档 / 官网 xlsx 完全同构的 CSV(备份与对端交换用)。"""
path = path or os.path.join(config.EXPORT_DIR, "usage_records.csv")
os.makedirs(os.path.dirname(path), exist_ok=True)
rows = conn.execute("SELECT request_id,credits,prompt,model,client,ts FROM usage_records "
"ORDER BY ts, request_id")
n = 0
with open(path, "w", encoding="utf-8-sig", newline="") as f:
w = csv.writer(f)
w.writerow(FIELDS)
for r in rows:
w.writerow([r["request_id"], "%.2f" % r["credits"], r["prompt"] or "",
r["model"], r["client"], r["ts"]])
n += 1
return path, n
def migrate_from_csv(conn, path, log=print):
"""把旧版 data/usage_records.csv 全量导入 SQLite(幂等,可重复执行)。"""
if not os.path.exists(path):
raise FileNotFoundError(path)
with open(path, "rb") as fb:
if fb.read(2) == b"PK":
raise ApiError("文件被 Excel 另存成了 xlsx(扩展名仍是 .csv):%s" % path)
with open(path, "r", encoding="utf-8-sig", newline="") as f:
recs = []
for r in csv.DictReader(f):
rid = (r.get("RequestID") or "").strip()
ts = (r.get("时间") or "").strip()[:19]
if not rid or len(ts) < 19:
continue
try:
cr = round(float(r.get("积分消耗") or 0), 2)
except (TypeError, ValueError):
cr = 0.0
recs.append({"request_id": rid, "ts": ts, "credits": cr,
"prompt": " ".join(str(r.get("User Prompt") or "").split()),
"model": (r.get("模型") or "-").strip() or "-",
"client": (r.get("客户端") or "-").strip() or "-"})
before = record_count(conn)
added, dup, _ = upsert(conn, recs)
log("[migrate] 源文件 %d 条 → 新增 %d / 已存在 %d,入库前 %d 条,现共 %d 条"
% (len(recs), added, dup, before, record_count(conn)))
return added