数据隔离
- settings / usage_records 主键改为 (user_id, key) / (user_id, request_id),
索引一律以 user_id 打头;collect_runs / audit_log 增加 user_id
- query / collect / scheduler 全链路把 uid 作为 conn 之后的第一个位置参数且无默认值
(漏传直接 TypeError,不会退化成「返回全量」)
- 配置三级回落 个人→实例→DEFAULTS;NO_FALLBACK_KEYS={cookie,user_agent} 不回落
凭证保密
- 新增 workbuddy_portal/crypto.py:手写 ChaCha20(RFC8439 §2.3) + HMAC-SHA256
encrypt-then-MAC,零第三方依赖;主密钥 cookie_key 与 SECRET_KEY 分键位存放
- get_secret() 是取明文的唯一通道;get_settings() 把加密键置空;
secret_state() 只回 {set,chars,tail,broken};升级时自动加密历史明文
注册与验证码
- 新增 /register 与 workbuddy_portal/captcha.py(手写 PNG + 点阵字模 + 干扰线)
- 验证码答案只存服务端表、不进 session,一次性、5 分钟过期、按 purpose 隔离
- allow_register / register_max_per_ip / captcha_policy / captcha_length 四个实例级开关
- 失败限速改为 IP + 用户名双维度;停用账号每请求回查、立即失效
页面
- 新增 /profile(个人中心)与注册页;登录页加验证码与自助注册入口
- /config 增加凭证状态、cookie_broken 告警、实例级设置区;/users 增加邮箱/状态与启停
修复
- base.html 顶层 {% set me %} 覆盖子模板同名变量,导致个人中心「注册于」渲染为空
- WB_COOKIE_SECURE 未写进 compose 的 environment,在 .env 里设了不生效
- 「修改登录密码」提示写「至少 6 位」,与实际策略(≥8 位 + 两类字符)不符
- 「用户管理」删除说明写「可勾选保留」,与页面实际行为不符
- 注册页与 flash 文案里的 **强调** Markdown 字面量
验证与文档
- smoke.py 99 → 165 项断言(多用户隔离 / 凭证保密 / 注册与验证码 / 3 条防回归)
- check_live.py 56 → 83 项断言(新增注册 / 验证码 / 安全响应头一节)
- demo_data.py 造两个账号;shots.py 自动过验证码、重出 11 张截图
- README / SECURITY / ARCHITECTURE / API / DEPLOYMENT / USER-GUIDE / FAQ / CHANGELOG / CONTRIBUTING 同步
508 行
21 KiB
Python
508 行
21 KiB
Python
# -*- coding: utf-8 -*-
|
||
# SPDX-License-Identifier: MIT
|
||
# Copyright (c) 2026 Wang Chuanli
|
||
|
||
"""采集主流程:云端增量 -> SQLite(去重、断点、漂移校验、运行记录)。
|
||
|
||
与原 fetch_usage.py 的差别:
|
||
* 存档正本从 CSV 换成 SQLite,去重由 `ON CONFLICT(user_id, request_id)` 承担
|
||
* 断点由 `SELECT MAX(ts) WHERE user_id=?` 承担,不再需要全量读入内存
|
||
* 每次运行落一条 collect_runs 记录,页面据此展示任务历史与日志
|
||
* 采集互斥用文件锁,保证「单写者」——SQLite 只允许一个写进程
|
||
|
||
**多用户约定(最重要)**
|
||
所有写入函数都要求显式传入 `uid`,且 `uid` 是 `conn` 之后的第一个位置参数、
|
||
没有默认值。采集**只使用该账号自己保存的 Cookie**:
|
||
|
||
* 明文的 Cookie 从来只存在于内存里(库里是 crypto.encrypt 后的密文)
|
||
* 不再支持 `WB_COOKIE` 环境变量作为采集凭证 —— 那会让所有人共用一份凭证,
|
||
一旦生效就是「A 的采集把数据写进 B 的账号」这种串号事故。
|
||
环境变量只保留给 `manage.py import-creds` 做一次性导入。
|
||
"""
|
||
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 NotReady(Exception):
|
||
"""该账号还没配好凭证 —— 不是错误,只是「没什么可做的」。"""
|
||
|
||
|
||
# ---------------- 互斥锁 ----------------
|
||
class _Lock:
|
||
"""全局单写者锁。
|
||
|
||
刻意**不做成按用户加锁**:SQLite 同一时刻只允许一个写事务,
|
||
按用户并行反而会在 busy_timeout 上互相拖死。串行跑完所有人的采集,
|
||
总耗时与并发差别很小(每个人一天也就拉一次)。
|
||
"""
|
||
|
||
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, uid, trigger):
|
||
cur = conn.execute("INSERT INTO collect_runs(user_id,trigger,status,started_at)"
|
||
" VALUES(?,?,?,?)", (uid, trigger, "running", db.now_str()))
|
||
return cur.lastrowid
|
||
|
||
|
||
def finish_run(conn, run_id, status, **kw):
|
||
sets = ["status=?", "finished_at=?"]
|
||
vals = [status, db.now_str()]
|
||
for f in ("duration_ms", "win_from", "win_to", "fetched", "added", "dup",
|
||
"total", "conflicts", "exit_code", "message", "detail"):
|
||
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(user_id,request_id,ts,day,hour,model,client,credits,prompt,
|
||
first_seen,last_seen,cloud_ts)
|
||
VALUES(:user_id,:request_id,:ts,:day,:hour,:model,:client,:credits,:prompt,
|
||
:first_seen,:last_seen,:cloud_ts)
|
||
ON CONFLICT(user_id,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, uid):
|
||
ts = (n.get("ts") or "").strip()[:19]
|
||
hh = ts[11:13]
|
||
return {
|
||
"user_id": uid,
|
||
"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, uid, normalized, drift_tolerance=5, log=None):
|
||
"""按 (user_id, 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, uid))
|
||
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 user_id=? AND request_id IN (%s)"
|
||
% ph, [uid] + 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, uid):
|
||
return conn.execute("SELECT COUNT(*) FROM usage_records WHERE user_id=?",
|
||
(uid or 0,)).fetchone()[0]
|
||
|
||
|
||
def last_ts(conn, uid):
|
||
return conn.execute("SELECT MAX(ts) FROM usage_records WHERE user_id=?",
|
||
(uid or 0,)).fetchone()[0]
|
||
|
||
|
||
# ---------------- 凭证读取(解密) ----------------
|
||
def load_credentials(conn, uid):
|
||
"""取该账号的 (cookie, user_agent)。
|
||
|
||
Cookie 从库里读出来是密文,由 db.get_secret 解密;解不开会抛
|
||
db.SecretUnreadable(多半是 instance.json 里的 cookie_key 被换过),
|
||
这时应当明确告诉用户「重新粘贴 Cookie」,而不是当成「未配置」静默跳过。
|
||
"""
|
||
cookie = db.get_secret(conn, "cookie", uid).strip()
|
||
ua = (db.get_setting(conn, "user_agent", "", uid) or "").strip()
|
||
return cookie, ua
|
||
|
||
|
||
# ---------------- 主同步 ----------------
|
||
def sync(conn, uid, trigger="manual", from_dt=None, to_dt=None, verify_days=None,
|
||
do_write=True, log=None):
|
||
"""增量同步(仅限 uid 这个账号)。
|
||
|
||
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, uid=uid)
|
||
cookie, ua = load_credentials(conn, uid)
|
||
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, uid)
|
||
rewind = db.get_int(conn, "rewind_minutes", 2, uid)
|
||
drift = db.get_int(conn, "drift_tolerance_minutes", 5, uid)
|
||
max_prompt = db.get_int(conn, "max_prompt", 0, uid)
|
||
timeout = db.get_int(conn, "timeout", 30, uid)
|
||
ssl_verify = db.get_bool(conn, "ssl_verify", True, uid)
|
||
if verify_days is None:
|
||
verify_days = db.get_int(conn, "verify_days", 0, uid)
|
||
# 夹到合法区间,避免历史脏数据(如超大的 page_size)把云端打爆
|
||
lo, hi, _ = config.NUM_SETTINGS["page_size"]
|
||
page_size = max(lo, min(hi, page_size))
|
||
|
||
run_id = start_run(conn, uid, 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,请到「配置管理」页粘贴自己账号的 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 NotReady(msg)
|
||
|
||
total_before = record_count(conn, uid)
|
||
tail = last_ts(conn, uid)
|
||
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, uid, 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 WHERE user_id=? ORDER BY day DESC LIMIT ?",
|
||
(uid, verify_days))]
|
||
_log("完整性校验:最近 %d 天" % len(days))
|
||
for d in sorted(days):
|
||
local = conn.execute("SELECT COUNT(*) FROM usage_records WHERE user_id=? AND day=?",
|
||
(uid, 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, uid,
|
||
[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, uid)
|
||
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, "uid": uid}
|
||
|
||
|
||
def run_sync(trigger="manual", **kw):
|
||
"""带锁的同步入口(供 CLI / 调度器 / 页面手动触发共用)。
|
||
|
||
必须显式给出 `uid=...`;漏传会由 sync() 直接报 TypeError,
|
||
不会退化成「用某个默认账号去采集」。
|
||
"""
|
||
with _Lock():
|
||
conn = db.thread_conn()
|
||
return sync(conn, trigger=trigger, **kw)
|
||
|
||
|
||
# ---------------- 补全 / 导入 / 导出 ----------------
|
||
def fill_prompt(conn, uid, log=print):
|
||
"""补全该账号缺失的 User Prompt(官网导出的 xlsx 会丢约 22%,云端仍保留)。"""
|
||
s = db.get_settings(conn, uid=uid)
|
||
cookie, ua = load_credentials(conn, uid)
|
||
max_prompt = db.get_int(conn, "max_prompt", 0, uid)
|
||
if not cookie:
|
||
raise NotReady("未配置 Cookie")
|
||
todo = conn.execute("SELECT request_id, day FROM usage_records "
|
||
"WHERE user_id=? AND COALESCE(prompt,'')='' ORDER BY day",
|
||
(uid,)).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, uid),
|
||
timeout=db.get_int(conn, "timeout", 30, uid),
|
||
ssl_verify=db.get_bool(conn, "ssl_verify", True, uid))
|
||
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 user_id=? AND request_id=?",
|
||
(p, uid, row["request_id"]))
|
||
n += 1
|
||
left = conn.execute("SELECT COUNT(*) FROM usage_records WHERE user_id=?"
|
||
" AND COALESCE(prompt,'')=''", (uid,)).fetchone()[0]
|
||
log("补全 %d 条,仍为空 %d 条" % (n, left))
|
||
return n
|
||
|
||
|
||
def import_xlsx(conn, uid, path, log=print):
|
||
"""从官网「用量明细 - 导出」的 xlsx 合入(按 request_id 去重)。"""
|
||
try:
|
||
import openpyxl
|
||
except ImportError:
|
||
raise ApiError("需要 openpyxl:pip install openpyxl")
|
||
max_prompt = db.get_int(conn, "max_prompt", 0, uid)
|
||
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, uid, rows)
|
||
log("[xlsx] 读取 %d 条,去重后新增 %d 条,该账号存档共 %d 条"
|
||
% (len(rows), added, record_count(conn, uid)))
|
||
return added
|
||
|
||
|
||
def export_csv(conn, uid, path=None, username=None):
|
||
"""导出与旧存档 / 官网 xlsx 完全同构的 CSV(备份与对端交换用)。
|
||
|
||
文件名带账号名:多用户下所有人导出到同一个目录,
|
||
不带归属就会互相覆盖。
|
||
"""
|
||
if path is None:
|
||
tag = username or ("u%s" % uid)
|
||
path = os.path.join(config.EXPORT_DIR, "usage_records_%s.csv" % tag)
|
||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||
rows = conn.execute("SELECT request_id,credits,prompt,model,client,ts FROM usage_records "
|
||
"WHERE user_id=? ORDER BY ts, request_id", (uid,))
|
||
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, uid, path, log=print):
|
||
"""把旧版 data/usage_records.csv 全量导入 SQLite(幂等,可重复执行)。
|
||
|
||
导入的数据归属 `uid` 指定的账号 —— 老存档是单用户的,
|
||
必须由调用方明确「这份数据算谁的」。
|
||
"""
|
||
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, uid)
|
||
added, dup, _ = upsert(conn, uid, recs)
|
||
log("[migrate] 源文件 %d 条 → 新增 %d / 已存在 %d,入库前 %d 条,现共 %d 条"
|
||
% (len(recs), added, dup, before, record_count(conn, uid)))
|
||
return added
|