420 lines
14 KiB
Python
420 lines
14 KiB
Python
# -*- coding: utf-8 -*-
|
|
"""SQLite 存储:提供商配置 / 测试记录 / 每次采样指标 / 全量日志"""
|
|
import os
|
|
import json
|
|
import math
|
|
import sqlite3
|
|
import threading
|
|
import time
|
|
|
|
import config
|
|
|
|
_lock = threading.RLock() # RLock:允许 get_logs_after 内嵌套调用 get_last_run
|
|
|
|
SCHEMA = """
|
|
CREATE TABLE IF NOT EXISTS configs(
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
name TEXT DEFAULT '',
|
|
provider TEXT DEFAULT 'openai',
|
|
base_url TEXT DEFAULT '',
|
|
api_key TEXT DEFAULT '',
|
|
model TEXT DEFAULT '',
|
|
temperature REAL DEFAULT 0.7,
|
|
created_at TEXT DEFAULT (datetime('now','localtime'))
|
|
);
|
|
|
|
CREATE TABLE IF NOT EXISTS tests(
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
created_at TEXT DEFAULT (datetime('now','localtime')),
|
|
status TEXT DEFAULT 'running',
|
|
provider TEXT DEFAULT '',
|
|
model TEXT DEFAULT '',
|
|
name TEXT DEFAULT '',
|
|
config_json TEXT DEFAULT '{}',
|
|
gen_cfg_json TEXT DEFAULT '{}',
|
|
summary_json TEXT DEFAULT '{}',
|
|
error TEXT DEFAULT '',
|
|
started_at REAL DEFAULT 0
|
|
);
|
|
|
|
CREATE TABLE IF NOT EXISTS test_runs(
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
test_id INTEGER NOT NULL,
|
|
run_index INTEGER DEFAULT 0,
|
|
context_length INTEGER DEFAULT 0,
|
|
metrics_json TEXT DEFAULT '{}',
|
|
error TEXT DEFAULT ''
|
|
);
|
|
|
|
CREATE TABLE IF NOT EXISTS logs(
|
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
|
test_id INTEGER NOT NULL,
|
|
level TEXT DEFAULT 'INFO',
|
|
msg TEXT DEFAULT '',
|
|
rel REAL DEFAULT 0,
|
|
ts TEXT DEFAULT (datetime('now','localtime'))
|
|
);
|
|
|
|
CREATE INDEX IF NOT EXISTS idx_runs_test ON test_runs(test_id);
|
|
CREATE INDEX IF NOT EXISTS idx_logs_test ON logs(test_id);
|
|
"""
|
|
|
|
|
|
def _connect():
|
|
os.makedirs(config.DATA_DIR, exist_ok=True)
|
|
conn = sqlite3.connect(config.DB_PATH, check_same_thread=False, timeout=30)
|
|
conn.row_factory = sqlite3.Row
|
|
conn.execute("PRAGMA journal_mode=WAL")
|
|
return conn
|
|
|
|
|
|
def _migrate(conn):
|
|
"""老库升级:test_runs 补 context_length、tests 补 name"""
|
|
cur = conn.execute("PRAGMA table_info(test_runs)")
|
|
cols = [r[1] for r in cur.fetchall()]
|
|
if "context_length" not in cols:
|
|
conn.execute("ALTER TABLE test_runs ADD COLUMN context_length INTEGER DEFAULT 0")
|
|
cur = conn.execute("PRAGMA table_info(tests)")
|
|
tcols = [r[1] for r in cur.fetchall()]
|
|
if "name" not in tcols:
|
|
conn.execute("ALTER TABLE tests ADD COLUMN name TEXT DEFAULT ''")
|
|
|
|
|
|
def init_db():
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
conn.executescript(SCHEMA)
|
|
_migrate(conn)
|
|
conn.commit()
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
# ───────────────────────── 提供商配置 ─────────────────────────
|
|
|
|
def add_config(cfg: dict) -> int:
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
cur = conn.execute(
|
|
"INSERT INTO configs(name,provider,base_url,api_key,model,temperature) VALUES(?,?,?,?,?,?)",
|
|
(cfg.get("name", ""), cfg.get("provider", "openai"),
|
|
cfg.get("base_url", ""), cfg.get("api_key", ""),
|
|
cfg.get("model", ""), float(cfg.get("temperature", 0.7))))
|
|
conn.commit()
|
|
return cur.lastrowid
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def list_configs():
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
rows = conn.execute("SELECT id,name,provider,base_url,model,temperature,"
|
|
"CASE WHEN api_key<>'' THEN 1 ELSE 0 END AS has_key "
|
|
"FROM configs ORDER BY id DESC").fetchall()
|
|
return [dict(r) for r in rows]
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def get_config(cid: int):
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
r = conn.execute("SELECT * FROM configs WHERE id=?", (cid,)).fetchone()
|
|
return dict(r) if r else None
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def delete_config(cid: int):
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
conn.execute("DELETE FROM configs WHERE id=?", (cid,))
|
|
conn.commit()
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def update_config(cid: int, cfg: dict):
|
|
"""局部更新:只更新请求里出现的字段"""
|
|
allow = {"name", "provider", "base_url", "api_key", "model", "temperature"}
|
|
fields = {k: v for k, v in cfg.items() if k in allow and v is not None}
|
|
if not fields:
|
|
return
|
|
sets = ",".join("%s=?" % k for k in fields)
|
|
vals = list(fields.values()) + [cid]
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
conn.execute("UPDATE configs SET %s WHERE id=?" % sets, vals)
|
|
conn.commit()
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
# ───────────────────────── 测试记录 ─────────────────────────
|
|
|
|
def create_test(cfg: dict, gen: dict) -> int:
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
cur = conn.execute(
|
|
"INSERT INTO tests(status,provider,model,name,config_json,gen_cfg_json,started_at) "
|
|
"VALUES('running',?,?,?,?,?,?)",
|
|
(cfg.get("provider", "openai"), cfg.get("model", ""),
|
|
gen.get("name") or cfg.get("name") or "",
|
|
json.dumps(cfg, ensure_ascii=False), json.dumps(gen, ensure_ascii=False),
|
|
time.time()))
|
|
conn.commit()
|
|
return cur.lastrowid
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def update_status(tid: int, status: str, summary: dict = None, error: str = ""):
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
conn.execute("UPDATE tests SET status=?, summary_json=?, error=? WHERE id=?",
|
|
(status,
|
|
json.dumps(summary or {}, ensure_ascii=False),
|
|
error or "", tid))
|
|
conn.commit()
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def get_test(tid: int):
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
r = conn.execute("SELECT * FROM tests WHERE id=?", (tid,)).fetchone()
|
|
if not r:
|
|
return None
|
|
d = dict(r)
|
|
d["config"] = json.loads(d.pop("config_json") or "{}")
|
|
d["gen"] = json.loads(d.pop("gen_cfg_json") or "{}")
|
|
d["summary"] = json.loads(d.pop("summary_json") or "{}")
|
|
return d
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
# 历史列表排序字段 → 取值函数(数值字段取 summary 聚合值,None 恒排最后)
|
|
SORT_KEYS = {
|
|
"id": lambda t: t["id"],
|
|
"created_at": lambda t: t["created_at"],
|
|
"name": lambda t: (t.get("name") or "").lower(),
|
|
"provider": lambda t: (t.get("provider") or "").lower(),
|
|
"model": lambda t: (t.get("model") or "").lower(),
|
|
"status": lambda t: t.get("status") or "",
|
|
"samples": lambda t: (t.get("summary") or {}).get("samples_ok"),
|
|
"ttft": lambda t: (t.get("summary") or {}).get("avg_ttft_ms"),
|
|
"prefill": lambda t: (t.get("summary") or {}).get("avg_prefill_speed"),
|
|
"decode": lambda t: (t.get("summary") or {}).get("avg_decode_speed"),
|
|
"total_ms": lambda t: (t.get("summary") or {}).get("avg_total_ms"),
|
|
}
|
|
|
|
|
|
def query_tests(q="", status="", provider="", sort="id", order="desc",
|
|
page=1, page_size=20):
|
|
"""历史记录:关键词搜索 + 状态/提供商筛选 + 字段排序 + 分页。
|
|
返回 {"items": [...], "total": N, "page": p, "page_size": s, "total_pages": n}"""
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
where, args = [], []
|
|
if q:
|
|
like = "%" + q + "%"
|
|
where.append("(id LIKE ? OR name LIKE ? OR model LIKE ? OR provider LIKE ?"
|
|
" OR status LIKE ? OR created_at LIKE ?)")
|
|
args += [like] * 6
|
|
if status:
|
|
sts = [s.strip() for s in status.split(",") if s.strip()]
|
|
if sts:
|
|
where.append("status IN (%s)" % ",".join("?" * len(sts)))
|
|
args += sts
|
|
if provider:
|
|
where.append("provider = ?")
|
|
args.append(provider)
|
|
wsql = ("WHERE " + " AND ".join(where)) if where else ""
|
|
|
|
rows = conn.execute(
|
|
"SELECT id,created_at,status,provider,model,name,summary_json,gen_cfg_json,error "
|
|
"FROM tests " + wsql + " ORDER BY id DESC", args).fetchall()
|
|
out = []
|
|
for r in rows:
|
|
d = dict(r)
|
|
d["summary"] = json.loads(d.pop("summary_json") or "{}")
|
|
d["gen"] = json.loads(d.pop("gen_cfg_json") or "{}")
|
|
out.append(d)
|
|
|
|
# Python 侧排序:None 值恒排最后(无论升降序)
|
|
keyfn = SORT_KEYS.get(sort) or SORT_KEYS["id"]
|
|
present = [t for t in out if keyfn(t) is not None]
|
|
absent = [t for t in out if keyfn(t) is None]
|
|
present.sort(key=keyfn, reverse=(order == "desc"))
|
|
out = present + absent
|
|
|
|
total = len(out)
|
|
total_pages = max(1, math.ceil(total / page_size))
|
|
page = min(max(1, page), total_pages)
|
|
start = (page - 1) * page_size
|
|
return {
|
|
"items": out[start:start + page_size],
|
|
"total": total,
|
|
"page": page,
|
|
"page_size": page_size,
|
|
"total_pages": total_pages,
|
|
}
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def clear_tests():
|
|
"""一键清空全部历史(测试记录 + 采样指标 + 日志)"""
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
conn.execute("DELETE FROM test_runs")
|
|
conn.execute("DELETE FROM logs")
|
|
conn.execute("DELETE FROM tests")
|
|
conn.commit()
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def delete_test(tid: int):
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
conn.execute("DELETE FROM tests WHERE id=?", (tid,))
|
|
conn.execute("DELETE FROM test_runs WHERE test_id=?", (tid,))
|
|
conn.execute("DELETE FROM logs WHERE test_id=?", (tid,))
|
|
conn.commit()
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def export_tests():
|
|
"""导出全部历史:每条含 基本字段 + 配置(Key打码) + 生成参数 + 摘要指标,
|
|
不含逐条采样与完整日志(避免文件过大)"""
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
rows = conn.execute(
|
|
"SELECT id,created_at,status,provider,model,name,config_json,gen_cfg_json,summary_json,error "
|
|
"FROM tests ORDER BY id").fetchall()
|
|
out = []
|
|
for r in rows:
|
|
d = dict(r)
|
|
cfg = json.loads(d.pop("config_json") or "{}")
|
|
if cfg.get("api_key"):
|
|
cfg["api_key"] = "******"
|
|
d["config"] = cfg
|
|
d["gen"] = json.loads(d.pop("gen_cfg_json") or "{}")
|
|
d["summary"] = json.loads(d.pop("summary_json") or "{}")
|
|
out.append(d)
|
|
return out
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
# ───────────────────────── 采样指标 ─────────────────────────
|
|
|
|
def add_run(tid: int, run_index: int, metrics: dict, error: str = "", context_length: int = 0):
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
conn.execute(
|
|
"INSERT INTO test_runs(test_id,run_index,context_length,metrics_json,error) VALUES(?,?,?,?,?)",
|
|
(tid, run_index, context_length, json.dumps(metrics, ensure_ascii=False), error))
|
|
conn.commit()
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def get_runs(tid: int):
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
rows = conn.execute(
|
|
"SELECT run_index,context_length,metrics_json,error FROM test_runs "
|
|
"WHERE test_id=? ORDER BY run_index", (tid,)).fetchall()
|
|
out = []
|
|
for r in rows:
|
|
d = dict(r)
|
|
d["metrics"] = json.loads(d.pop("metrics_json") or "{}")
|
|
out.append(d)
|
|
return out
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def get_last_run(tid: int):
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
r = conn.execute(
|
|
"SELECT metrics_json FROM test_runs WHERE test_id=? "
|
|
"ORDER BY run_index DESC LIMIT 1", (tid,)).fetchone()
|
|
return json.loads(r["metrics_json"] or "{}") if r else None
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
# ───────────────────────── 日志 ─────────────────────────
|
|
|
|
def add_log(tid: int, level: str, msg: str, rel: float = None):
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
if rel is None:
|
|
r = conn.execute("SELECT started_at FROM tests WHERE id=?", (tid,)).fetchone()
|
|
rel = (time.time() - (r["started_at"] or time.time())) if r else 0.0
|
|
conn.execute("INSERT INTO logs(test_id,level,msg,rel) VALUES(?,?,?,?)",
|
|
(tid, level, msg, round(rel, 3)))
|
|
conn.commit()
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def get_logs_after(tid: int, after_id: int = 0):
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
t = conn.execute("SELECT status,summary_json,error FROM tests WHERE id=?", (tid,)).fetchone()
|
|
if not t:
|
|
return None
|
|
rows = conn.execute(
|
|
"SELECT id,level,msg,rel FROM logs WHERE test_id=? AND id>? ORDER BY id",
|
|
(tid, after_id)).fetchall()
|
|
logs = [dict(r) for r in rows]
|
|
last = logs[-1]["id"] if logs else after_id
|
|
return {
|
|
"status": t["status"],
|
|
"error": t["error"],
|
|
"summary": json.loads(t["summary_json"] or "{}"),
|
|
"last_run": get_last_run(tid),
|
|
"logs": logs,
|
|
"after": last,
|
|
}
|
|
finally:
|
|
conn.close()
|
|
|
|
|
|
def get_logs(tid: int):
|
|
with _lock:
|
|
conn = _connect()
|
|
try:
|
|
rows = conn.execute(
|
|
"SELECT id,level,msg,rel,ts FROM logs WHERE test_id=? ORDER BY id", (tid,)).fetchall()
|
|
return [dict(r) for r in rows]
|
|
finally:
|
|
conn.close()
|