480 lines
16 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.

"""本地 SQLite(data/muse.db):WAL、版本化迁移、run/event/review/card。"""
from __future__ import annotations
import hashlib
import json
import math
import os
import sqlite3
import struct
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Iterable, Mapping, Sequence
PROJECT_ROOT = next(
parent
for parent in (Path(__file__).resolve().parent, *Path(__file__).resolve().parents)
if (parent / "AGENTS.md").is_file() and (parent / ".git").exists()
)
MIGRATIONS_DIR = PROJECT_ROOT / "data" / "migrations"
_DEFAULT_DB = PROJECT_ROOT / "data" / "muse.db"
def default_db_path() -> Path:
override = os.environ.get("MUSE_DB")
return Path(override).expanduser() if override else _DEFAULT_DB
def _now() -> str:
return datetime.now(timezone.utc).replace(microsecond=0).isoformat()
def pack_embedding(vector: Sequence[float]) -> bytes:
return struct.pack(f"<{len(vector)}f", *[float(item) for item in vector])
def unpack_embedding(blob: bytes) -> list[float]:
if not blob:
return []
count = len(blob) // 4
return list(struct.unpack(f"<{count}f", blob))
def cosine(left: Sequence[float], right: Sequence[float]) -> float:
if not left or not right or len(left) != len(right):
return 0.0
dot = sum(a * b for a, b in zip(left, right))
norm_l = math.sqrt(sum(a * a for a in left))
norm_r = math.sqrt(sum(b * b for b in right))
if norm_l == 0.0 or norm_r == 0.0:
return 0.0
return dot / (norm_l * norm_r)
def skill_set_hash(root: Path | None = None) -> str:
base = root or PROJECT_ROOT
digest = hashlib.sha256()
skills = base / ".agent" / "skills"
if skills.is_dir():
for path in sorted(skills.rglob("SKILL.md")):
digest.update(path.as_posix().encode("utf-8"))
digest.update(path.read_bytes())
return "sha256:" + digest.hexdigest()
def baseline_hash(path: Path | None = None) -> str | None:
db_path = path or default_db_path()
if not db_path.is_file():
return None
digest = hashlib.sha256()
with db_path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
return "sha256:" + digest.hexdigest()
def connect(path: str | Path | None = None) -> sqlite3.Connection:
db_path = Path(path) if path is not None else default_db_path()
db_path.parent.mkdir(parents=True, exist_ok=True)
conn = sqlite3.connect(str(db_path))
conn.row_factory = sqlite3.Row
conn.execute("PRAGMA foreign_keys=ON")
mode = conn.execute("PRAGMA journal_mode=WAL").fetchone()[0]
if str(mode).lower() != "wal":
conn.close()
raise RuntimeError(f"无法开启 WAL:journal_mode={mode!r}")
_apply_migrations(conn)
return conn
def _apply_migrations(conn: sqlite3.Connection) -> None:
conn.execute(
"""CREATE TABLE IF NOT EXISTS schema_migrations (
version TEXT PRIMARY KEY,
applied_at TEXT NOT NULL
)"""
)
applied = {
row["version"]
for row in conn.execute("SELECT version FROM schema_migrations")
}
for script in sorted(MIGRATIONS_DIR.glob("*.sql")):
version = script.stem
if version in applied:
continue
conn.executescript(script.read_text(encoding="utf-8"))
conn.execute(
"INSERT INTO schema_migrations(version, applied_at) VALUES (?, ?)",
(version, _now()),
)
conn.commit()
def upsert_card(
*,
card_id: str,
kind: str,
title: str | None,
payload: Mapping[str, Any],
embedding: Sequence[float] | None = None,
work_id: int | None = None,
source_path: str | None = None,
content_hash: str | None = None,
path: str | Path | None = None,
) -> None:
blob = pack_embedding(embedding) if embedding else None
with connect(path) as conn:
conn.execute(
"""INSERT INTO cards(
id, kind, title, payload_json, embedding, work_id,
source_path, content_hash, created_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(id) DO UPDATE SET
kind=excluded.kind,
title=excluded.title,
payload_json=excluded.payload_json,
embedding=excluded.embedding,
work_id=excluded.work_id,
source_path=excluded.source_path,
content_hash=excluded.content_hash
""",
(
card_id,
kind,
title,
json.dumps(payload, ensure_ascii=False),
blob,
work_id,
source_path,
content_hash,
_now(),
),
)
conn.commit()
def _ai_context_visible(ai_rule: Any, purpose: str) -> bool:
"""aiContext 判定,与 检索知识 PG 面同语义:
true 全用途可见;false 不可见;[用途] 仅列出的可见;无规则默认可见。"""
if ai_rule is None:
return True
if isinstance(ai_rule, bool):
return ai_rule
return purpose in ai_rule
def search_card_vectors(
intent_vector: Sequence[float],
*,
kind: str | None = None,
work_id: int | None = None,
top: int = 5,
path: str | Path | None = None,
purpose: str = "generation",
ai_rules: Mapping[str, Mapping[str, Any]] | None = None,
) -> list[dict[str, Any]]:
"""本地 SQLite 向量召回(PG 检索面的降级镜像)。
与 PG 面同口径应用 aiContext 字段裁剪;本地卡没有绑定/状态面,
bindingStatus 与 productionRetrievalEligible 一律失败关闭(None/False),
不得伪造 active 资格——生产消费方(如 ProductionCardIndexRepository)
会因此拒绝本地镜像结果,而不是静默放行。
"""
if purpose not in {"generation", "planning", "detection", "extraction"}:
raise ValueError("purpose 非法")
rules_by_type: Mapping[str, Mapping[str, Any]] = ai_rules or {}
sql = "SELECT id, kind, title, payload_json, embedding, work_id, source_path, content_hash FROM cards WHERE embedding IS NOT NULL"
args: list[Any] = []
if kind:
sql += " AND kind = ?"
args.append(kind)
if work_id is not None:
sql += " AND work_id = ?"
args.append(work_id)
with connect(path) as conn:
rows = list(conn.execute(sql, args))
scored: list[tuple[float, sqlite3.Row]] = []
for row in rows:
score = cosine(intent_vector, unpack_embedding(row["embedding"]))
scored.append((score, row))
scored.sort(key=lambda item: item[0], reverse=True)
results: list[dict[str, Any]] = []
for score, row in scored[:top]:
payload = json.loads(row["payload_json"] or "{}")
fields = payload.get("字段") if isinstance(payload.get("字段"), dict) else {}
# 无内容指针行(如 PG 迁移的 embedding 指针)不产出检索结果:
# 它们没有名称/字段/摘要,召回它们只会产出空壳卡片。
if not (payload.get("名称") or payload.get("一句话摘要") or fields):
continue
card_type = payload.get("型") or row["kind"]
type_rules = rules_by_type.get(card_type, {})
visible_fields = {
key: item
for key, item in fields.items()
if _ai_context_visible(type_rules.get(key), purpose)
}
results.append(
{
"cardId": row["id"],
"type": card_type,
"name": payload.get("名称") or row["title"],
"score": float(score),
"summary": payload.get("一句话摘要"),
"visibleFields": visible_fields,
"omittedFields": sorted(set(fields) - set(visible_fields)),
"sourceId": f"sqlite-card:{row['id']}",
"sourceVersion": f"hash:{row['content_hash'] or 'none'}",
"sourceOffset": 0,
"sourceRefs": payload.get("sourceRefs") if isinstance(payload.get("sourceRefs"), list) else [],
"milestones": [],
"sourceKind": "local_card",
"sourceStatus": None,
"bindingStatus": None,
"retrievalScope": "work" if work_id is not None else "admin",
"productionRetrievalEligible": False,
}
)
return results
class FanoutSink:
"""把同一条事件转给多个 sink;runtime 不认识本类型。"""
def __init__(self, *sinks: Any) -> None:
self._sinks = sinks
def emit(self, event_type: str, **kwargs: Any) -> Any:
result = None
for sink in self._sinks:
result = sink.emit(event_type, **kwargs)
return result
class SqliteRecorder:
"""缓冲会话事件,结束时一次写入 run + events。"""
def __init__(
self,
run_id: str,
*,
kind: str,
input_obj: Mapping[str, Any],
skill_set_hash_value: str,
role: str | None = None,
path: str | Path | None = None,
) -> None:
self.run_id = run_id
self.kind = kind
self.input_obj = dict(input_obj)
self.skill_set_hash_value = skill_set_hash_value
self.role = role
self.path = path
self.events: list[dict[str, Any]] = []
self._committed = False
def emit(self, event_type: str, **kwargs: Any) -> int:
seq = len(self.events) + 1
self.events.append(
{
"seq": seq,
"kind": event_type,
"role": self.role,
"tool_name": kwargs.get("tool_name"),
"payload": {key: value for key, value in kwargs.items() if key != "tool_name"},
"created_at": _now(),
}
)
return seq
def commit(self, output_text: str, meta: Mapping[str, Any] | None = None) -> None:
if self._committed:
return
self._committed = True
db_path = Path(self.path) if self.path is not None else default_db_path()
prior = baseline_hash(db_path)
payload = json.dumps(self.input_obj, ensure_ascii=False, sort_keys=True, default=str)
meta_json = json.dumps(dict(meta or {}), ensure_ascii=False, sort_keys=True, default=str)
with connect(db_path) as conn:
try:
conn.execute(
"""INSERT INTO runs(
id, created_at, kind, input_json, output_text,
meta_json, skill_set_hash, baseline_hash
) VALUES (?, ?, ?, ?, ?, ?, ?, ?)""",
(
self.run_id,
_now(),
self.kind,
payload,
output_text or "",
meta_json,
self.skill_set_hash_value,
prior,
),
)
except sqlite3.IntegrityError:
return
for event in self.events:
conn.execute(
"""INSERT INTO events(
run_id, seq, kind, role, tool_name, payload_json, created_at
) VALUES (?, ?, ?, ?, ?, ?, ?)""",
(
self.run_id,
event["seq"],
event["kind"],
event["role"],
event["tool_name"],
json.dumps(event["payload"], ensure_ascii=False, default=str),
event["created_at"],
),
)
conn.commit()
def add_review(
*,
run_id: str,
target: str,
action: str,
reviewer: str,
reason: str | None = None,
before_text: str | None = None,
after_text: str | None = None,
kind: str = "prose",
path: str | Path | None = None,
) -> int:
if not reviewer or not str(reviewer).strip():
raise ValueError("reviewer 必填")
if action == "revise" and (before_text is None or after_text is None):
raise ValueError("revise 必须提供 before_text 与 after_text")
with connect(path) as conn:
cursor = conn.execute(
"""INSERT INTO reviews(run_id, target, action, reviewer, reason, created_at)
VALUES (?, ?, ?, ?, ?, ?)""",
(run_id, target, action, reviewer, reason, _now()),
)
review_id = int(cursor.lastrowid)
if action == "revise":
conn.execute(
"""INSERT INTO revisions(review_id, kind, before_text, after_text, created_at)
VALUES (?, ?, ?, ?, ?)""",
(review_id, kind, before_text or "", after_text or "", _now()),
)
conn.commit()
return review_id
def register_lesson(
*,
lesson_id: str,
source_run_ids: Sequence[str],
source_review_ids: Sequence[str],
kind: str,
title: str,
content: str,
target_ref: str,
lessons_dir: str | Path | None = None,
path: str | Path | None = None,
) -> Path:
if not source_run_ids or not source_review_ids:
raise ValueError("lesson 必须绑定 source_run_ids 与 source_review_ids")
root = PROJECT_ROOT / "lessons" / "pending"
dest_dir = Path(lessons_dir) if lessons_dir is not None else root
dest_dir.mkdir(parents=True, exist_ok=True)
content_path = dest_dir / f"{lesson_id}.md"
content_path.write_text(content, encoding="utf-8")
with connect(path) as conn:
conn.execute(
"""INSERT INTO lessons(
id, source_run_ids, source_review_ids, kind, title,
content_path, status, target_ref
) VALUES (?, ?, ?, ?, ?, ?, 'proposed', ?)""",
(
lesson_id,
json.dumps(list(source_run_ids)),
json.dumps(list(source_review_ids)),
kind,
title,
str(content_path),
target_ref,
),
)
conn.commit()
return content_path
def approve_and_merge_lesson(
lesson_id: str,
*,
decided_by: str,
rationale: str,
references_dir: str | Path,
path: str | Path | None = None,
) -> Path:
with connect(path) as conn:
row = conn.execute("SELECT * FROM lessons WHERE id=?", (lesson_id,)).fetchone()
if row is None:
raise ValueError(f"lesson 不存在: {lesson_id}")
source = Path(row["content_path"])
text = source.read_text(encoding="utf-8")
dest_dir = Path(references_dir)
dest_dir.mkdir(parents=True, exist_ok=True)
dest = dest_dir / f"lesson-{lesson_id}.md"
dest.write_text(text, encoding="utf-8")
conn.execute(
"""UPDATE lessons
SET status='promoted', decided_by=?, decided_at=?, rationale=?
WHERE id=?""",
(decided_by, _now(), rationale, lesson_id),
)
conn.commit()
add_review(
run_id=json.loads(row["source_run_ids"])[0],
target="lesson",
action="adopt",
reviewer=decided_by,
reason=rationale,
path=path,
)
return dest
def get_run(run_id: str, path: str | Path | None = None) -> dict[str, Any] | None:
with connect(path) as conn:
row = conn.execute("SELECT * FROM runs WHERE id = ?", (run_id,)).fetchone()
if row is None:
return None
return dict(row)
def list_events(run_id: str, path: str | Path | None = None) -> list[dict[str, Any]]:
with connect(path) as conn:
rows = conn.execute(
"SELECT * FROM events WHERE run_id = ? ORDER BY seq ASC", (run_id,)
).fetchall()
return [dict(row) for row in rows]
__all__ = [
"FanoutSink",
"SqliteRecorder",
"add_review",
"baseline_hash",
"connect",
"cosine",
"default_db_path",
"approve_and_merge_lesson",
"get_run",
"list_events",
"register_lesson",
"pack_embedding",
"search_card_vectors",
"skill_set_hash",
"unpack_embedding",
"upsert_card",
]