Files
tzt 324c419c35 feat(sense): T-G3 标签+校准(labeler 四规则推导 + split-conformal + last-good)
- labeler.py:derive_true_tier 纯函数(规则0 人工优先 / plan_multi->T3 /
  失败->下一档(T3保持) / 成功->executed / 空结果跳过)+ derive_true_tiers
  批量回填 + 180d 留存清理
- calibrate.py:compute_thresholds(τ1/τ3 分组扫描——并列分数整组判定,
  精度 >= 1-α 的最大覆盖阈值;<min_labels -> ok=False 沿用 last-good);
  save_thresholds 工件落盘+登记;load_active 回退链 active->同kind扫描->last-good->保守值
- store:all_observations/list_artifacts 支撑方法
- 测试 +9(四规则/批量清理/覆盖达标/标签不足/工件往返与回退),全量 395 passed
2026-09-05 13:53:45 +08:00

191 lines
8.4 KiB
Python
Raw Permalink 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.
"""Sense 存储(T-G0DDL + 基础查询;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)