"""Sense 存储(T-G0:DDL + 基础查询;labeler/grader 的查询在后续任务落地)。 §4 数据模型:tier_observations(观察-标签闭环主表)+ sense_artifacts(工件登记)。 WAL;全部访问经全局锁 + 每操作新连接(ReviewQueue 模式),异步调用方 to_thread。 隐私(D-G6):不落 query 原文(只存哈希+特征+int8 embedding BLOB);留存 180 天。 """ from __future__ import annotations import sqlite3 import threading import time from pathlib import Path from typing import Any, Dict, List, Optional _SCHEMA = """ CREATE TABLE IF NOT EXISTS tier_observations( id INTEGER PRIMARY KEY, ts INTEGER NOT NULL, request_id TEXT NOT NULL, consumer TEXT NOT NULL, bucket TEXT NOT NULL DEFAULT 'default', domain TEXT DEFAULT '', decided_tier TEXT NOT NULL, executed_tier TEXT NOT NULL, probs TEXT NOT NULL, policy_version TEXT NOT NULL, features TEXT NOT NULL, embedding BLOB, outcome TEXT DEFAULT '', true_tier TEXT DEFAULT '', human_override TEXT DEFAULT ''); CREATE INDEX IF NOT EXISTS idx_obs_ts ON tier_observations(ts); CREATE INDEX IF NOT EXISTS idx_obs_policy ON tier_observations(policy_version); CREATE TABLE IF NOT EXISTS sense_artifacts( version TEXT PRIMARY KEY, kind TEXT NOT NULL, path TEXT NOT NULL, metrics TEXT NOT NULL, created_ts INTEGER NOT NULL, active INTEGER DEFAULT 0); """ class SenseStore: """sense.sqlite3 访问(同步实现;异步调用方 to_thread)。""" def __init__(self, db_path: str | Path): self.db_path = Path(db_path) self.db_path.parent.mkdir(parents=True, exist_ok=True) self._lock = threading.Lock() self._init_db() @classmethod def init_db(cls, db_path: str | Path) -> "SenseStore": return cls(db_path) def _connect(self) -> sqlite3.Connection: conn = sqlite3.connect(self.db_path, check_same_thread=False) conn.row_factory = sqlite3.Row return conn def _init_db(self) -> None: with self._lock, self._connect() as conn: conn.executescript(_SCHEMA) # ---------- 观察 ---------- def insert_observation(self, obs: Dict[str, Any]) -> None: """写入一条观察(T-G2 批量缓冲的落盘终点)。""" with self._lock, self._connect() as conn: conn.execute( """INSERT OR IGNORE INTO tier_observations (ts, request_id, consumer, bucket, domain, decided_tier, executed_tier, probs, policy_version, features, embedding, outcome, true_tier, human_override) VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?)""", (int(obs.get("ts", time.time())), obs["request_id"], obs["consumer"], obs.get("bucket", "default"), obs.get("domain", ""), obs["decided_tier"], obs["executed_tier"], obs.get("probs", "{}"), obs.get("policy_version", ""), obs.get("features", "{}"), obs.get("embedding"), obs.get("outcome", ""), obs.get("true_tier", ""), obs.get("human_override", ""))) def update_outcome(self, request_id: str, outcome: str, executed_tier: Optional[str] = None) -> bool: """升级阶梯回写(outcome=escalated 等)。""" with self._lock, self._connect() as conn: if executed_tier: cur = conn.execute( "UPDATE tier_observations SET outcome = ?, executed_tier = ?" " WHERE request_id = ?", (outcome, executed_tier, request_id)) else: cur = conn.execute( "UPDATE tier_observations SET outcome = ? WHERE request_id = ?", (outcome, request_id)) return cur.rowcount > 0 def set_true_tier(self, request_id: str, true_tier: str) -> bool: """labeler 回填(T-G3)。""" with self._lock, self._connect() as conn: cur = conn.execute( "UPDATE tier_observations SET true_tier = ? WHERE request_id = ?", (true_tier, request_id)) return cur.rowcount > 0 def labeled_rows(self, policy_version: str = "", limit: int = 100000) -> List[Dict[str, Any]]: """true_tier 非空的行(训练/校准输入)。""" with self._lock, self._connect() as conn: if policy_version: rows = conn.execute( "SELECT * FROM tier_observations WHERE true_tier != ''" " AND policy_version = ? ORDER BY ts LIMIT ?", (policy_version, limit)).fetchall() else: rows = conn.execute( "SELECT * FROM tier_observations WHERE true_tier != ''" " ORDER BY ts LIMIT ?", (limit,)).fetchall() return [dict(r) for r in rows] def count_labeled(self) -> int: """已标签条数(min_labels 晋升门检查)。""" with self._lock, self._connect() as conn: row = conn.execute( "SELECT COUNT(*) AS c FROM tier_observations" " WHERE true_tier != ''").fetchone() return int(row["c"]) def all_observations(self, limit: int = 200000) -> List[Dict[str, Any]]: """全量遍历(labeler 夜间推导输入;量大时可分页)。""" with self._lock, self._connect() as conn: rows = conn.execute( "SELECT * FROM tier_observations ORDER BY id LIMIT ?", (limit,)).fetchall() return [dict(r) for r in rows] def purge_older_than(self, ts: float) -> int: """留存清理(D-G6:默认 180 天,夜间任务顺带)。""" with self._lock, self._connect() as conn: cur = conn.execute( "DELETE FROM tier_observations WHERE ts < ?", (int(ts),)) return cur.rowcount # ---------- 工件 ---------- def register_artifact(self, version: str, kind: str, path: str, metrics: Dict[str, Any], active: bool = False, created_ts: Optional[int] = None) -> None: """登记模型/阈值工件(默认 active=0,人工 promote 切换)。""" with self._lock, self._connect() as conn: conn.execute( "INSERT OR REPLACE INTO sense_artifacts" "(version, kind, path, metrics, created_ts, active)" " VALUES (?,?,?,?,?,?)", (version, kind, path, json_dumps(metrics), int(created_ts if created_ts is not None else time.time()), 1 if active else 0)) def activate_artifact(self, version: str, kind: str) -> bool: """切换 active(同 kind 互斥)。""" with self._lock, self._connect() as conn: conn.execute("UPDATE sense_artifacts SET active = 0 WHERE kind = ?", (kind,)) cur = conn.execute( "UPDATE sense_artifacts SET active = 1 WHERE version = ? AND kind = ?", (version, kind)) return cur.rowcount > 0 def active_artifact(self, kind: str) -> Optional[Dict[str, Any]]: with self._lock, self._connect() as conn: row = conn.execute( "SELECT * FROM sense_artifacts WHERE kind = ? AND active = 1" " ORDER BY created_ts DESC LIMIT 1", (kind,)).fetchone() return dict(row) if row else None def list_artifacts(self, kind: str) -> List[Dict[str, Any]]: """同 kind 工件按时间倒序(load_active 的 last-good 扫描用)。""" with self._lock, self._connect() as conn: rows = conn.execute( "SELECT * FROM sense_artifacts WHERE kind = ?" " ORDER BY created_ts DESC", (kind,)).fetchall() return [dict(r) for r in rows] # ---------- 自省 ---------- def table_names(self) -> List[str]: with self._lock, self._connect() as conn: rows = conn.execute( "SELECT name FROM sqlite_master WHERE type='table'").fetchall() return [r["name"] for r in rows] def index_names(self) -> List[str]: with self._lock, self._connect() as conn: rows = conn.execute( "SELECT name FROM sqlite_master WHERE type='index'" " AND name LIKE 'idx_%'").fetchall() return [r["name"] for r in rows] def json_dumps(obj: Any) -> str: import json return json.dumps(obj, ensure_ascii=False)