Files
Video-Tree-TRM5/adapters/baseline_diagnosis_store.py
T

155 lines
5.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""SqliteDiagnosisSignalStorebaseline 逐题诊断信号的 SQLite 持久化适配器。
实现 core/evolution/protocols.py::DiagnosisSignalStore 端口。信号行以
(question_id, baseline_run_id, diag_fingerprint) 为主键,INSERT OR REPLACE
保证逐题幂等 upsert;bool 字段以 0/1 存储,可空字段以 NULL 存储。表建在
harness.db,供离线诊断编排写入、视频级切分选择器读取。
"""
from __future__ import annotations
import sqlite3
from pathlib import Path
from core.evolution.types import DiagnosisSignalRow
# 表列顺序即 DiagnosisSignalRow 字段顺序(question_id..session_id),
# upsert 写入与 load 还原共用,避免手写列名两处漂移。
_COLUMNS: tuple[str, ...] = (
"question_id",
"video_id",
"baseline_run_id",
"diag_fingerprint",
"task_type",
"error_type",
"cause_category",
"tier",
"evolution_target",
"degraded",
"infra",
"session_id",
)
class SqliteDiagnosisSignalStore:
"""逐题诊断信号的 SQLite 存储实现。
构造时按需建表(CREATE TABLE IF NOT EXISTS),主键
(question_id, baseline_run_id, diag_fingerprint) 保证同键覆盖。
每次写操作单事务 commit,保证原子落盘与断点续跑。
参数:
db_path: SQLite 数据库文件路径(通常为 harness.db)。
"""
def __init__(self, db_path: str) -> None:
Path(db_path).parent.mkdir(parents=True, exist_ok=True)
self._conn = sqlite3.connect(db_path)
self._conn.row_factory = sqlite3.Row
self._init_table()
def _init_table(self) -> None:
"""创建 baseline_diagnosis 表(若不存在)。"""
self._conn.execute(
"""
CREATE TABLE IF NOT EXISTS baseline_diagnosis (
question_id TEXT NOT NULL,
video_id TEXT NOT NULL,
baseline_run_id TEXT NOT NULL,
diag_fingerprint TEXT NOT NULL,
task_type TEXT NOT NULL,
error_type TEXT,
cause_category TEXT,
tier TEXT NOT NULL,
evolution_target TEXT,
degraded INTEGER NOT NULL,
infra INTEGER NOT NULL,
session_id TEXT,
PRIMARY KEY (question_id, baseline_run_id, diag_fingerprint)
)
"""
)
self._conn.commit()
def upsert(self, row: DiagnosisSignalRow) -> None:
"""写入或覆盖单题诊断信号(按主键幂等,单事务)。
参数:
row: 待持久化的诊断信号行。bool 字段转 0/1None 存 NULL。
关键实现:
用 INSERT OR REPLACE 按主键覆盖,避免重复行;commit 保证原子。
"""
placeholders = ", ".join("?" for _ in _COLUMNS)
col_names = ", ".join(_COLUMNS)
values = (
row.question_id,
row.video_id,
row.baseline_run_id,
row.diag_fingerprint,
row.task_type,
row.error_type,
row.cause_category,
row.tier,
row.evolution_target,
int(row.degraded),
int(row.infra),
row.session_id,
)
self._conn.execute(
f"INSERT OR REPLACE INTO baseline_diagnosis ({col_names}) VALUES ({placeholders})",
values,
)
self._conn.commit()
def done_question_ids(self, baseline_run_id: str, diag_fingerprint: str) -> set[str]:
"""查询指定 run 与诊断指纹下已完成的 question_id 集合。
参数:
baseline_run_id: baseline run 标识。
diag_fingerprint: 诊断口径指纹。
返回:
已落盘信号的 question_id 集合;无匹配时为空集,供断点续跑跳过。
"""
cursor = self._conn.execute(
"SELECT DISTINCT question_id FROM baseline_diagnosis"
" WHERE baseline_run_id = ? AND diag_fingerprint = ?",
(baseline_run_id, diag_fingerprint),
)
return {r["question_id"] for r in cursor.fetchall()}
def load(self, baseline_run_id: str, diag_fingerprint: str) -> list[DiagnosisSignalRow]:
"""加载指定 run 与诊断指纹下的全部诊断信号行。
参数:
baseline_run_id: baseline run 标识。
diag_fingerprint: 诊断口径指纹。
返回:
还原后的 DiagnosisSignalRow 列表(0/1→boolNULL→None)。
"""
col_names = ", ".join(_COLUMNS)
cursor = self._conn.execute(
f"SELECT {col_names} FROM baseline_diagnosis"
" WHERE baseline_run_id = ? AND diag_fingerprint = ?",
(baseline_run_id, diag_fingerprint),
)
return [
DiagnosisSignalRow(
question_id=r["question_id"],
video_id=r["video_id"],
baseline_run_id=r["baseline_run_id"],
diag_fingerprint=r["diag_fingerprint"],
task_type=r["task_type"],
error_type=r["error_type"],
cause_category=r["cause_category"],
tier=r["tier"],
evolution_target=r["evolution_target"],
degraded=bool(r["degraded"]),
infra=bool(r["infra"]),
session_id=r["session_id"],
)
for r in cursor.fetchall()
]