* LICENSE —— MIT
* THIRD-PARTY-NOTICES —— 依赖清单、再分发合规说明(含随仓库分发的
Apache ECharts 5.6.0 / Apache-2.0)与自查清单
* CONTRIBUTING.md —— 开发环境、五层验证、必须遵守的不变量、提交规范
* SECURITY.md —— 漏洞私有报告渠道、已有措施、已知非目标
* CODE_OF_CONDUCT.md —— 改编自 Contributor Covenant 2.1
* .github/ —— Bug 报告 / 功能建议表单 + PR 模板
* .editorconfig —— 与 .gitattributes 保持一致
* 全部 Python / Shell 源文件加 SPDX-License-Identifier: MIT 头
* README 增加「开源与许可」章节与许可标识
453 行
19 KiB
Python
453 行
19 KiB
Python
# -*- coding: utf-8 -*-
|
||
# SPDX-License-Identifier: MIT
|
||
# Copyright (c) 2026 Wang Chuanli
|
||
|
||
"""采集主流程:云端增量 -> 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
|