refactor(tools/skills_sync,memory): split skills_sync client wire/org and bundled/optional ops; extract memory_tool_store and session_search common/discover
This commit is contained in:
1256
tools/memory_tool.py
1256
tools/memory_tool.py
File diff suppressed because it is too large
Load Diff
543
tools/memory_tool_store.py
Normal file
543
tools/memory_tool_store.py
Normal file
@@ -0,0 +1,543 @@
|
||||
"""MemoryStore — bounded, file-backed curated memory (MEMORY.md / USER.md).
|
||||
Entries are joined by ``ENTRY_DELIMITER``; budgets are in chars (model-
|
||||
independent). Module state that tests monkeypatch (``get_memory_dir``,
|
||||
``fcntl``/``msvcrt``) stays in ``tools.memory_tool`` and is read lazily."""
|
||||
|
||||
import logging
|
||||
import time
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
from utils import atomic_write_text
|
||||
from tools.threat_patterns import first_threat_message as _first_threat_message
|
||||
|
||||
logger = logging.getLogger("tools.memory_tool")
|
||||
|
||||
# System-prompt block header prefixes rendered by MemoryStore._render_block.
|
||||
# agent/conversation_compression.py uses them to detect a leftover block for a
|
||||
# target whose entries have since been emptied — keep in lockstep.
|
||||
MEMORY_BLOCK_HEADERS = {
|
||||
"memory": "MEMORY (your personal notes)",
|
||||
"user": "USER PROFILE (who the user is)",
|
||||
}
|
||||
|
||||
ENTRY_DELIMITER = "\n§\n"
|
||||
|
||||
# Sentinel from ``_reload_target``: the file EXISTS but could not be read.
|
||||
# Distinct from a drift-backup path (str) and a clean reload (None); the caller
|
||||
# must abort rather than persist over an unreadable file.
|
||||
_READ_FAILED = object()
|
||||
|
||||
|
||||
def _memory_dir() -> Path:
|
||||
from tools import memory_tool
|
||||
|
||||
return memory_tool.get_memory_dir()
|
||||
|
||||
|
||||
def _scan_memory_content(content: str) -> Optional[str]:
|
||||
"""Scan memory content for injection/exfil patterns ("strict" scope: memory
|
||||
enters the system prompt as a frozen snapshot, so a poisoned entry persists
|
||||
across sessions). Returns the error string if blocked."""
|
||||
return _first_threat_message(content, scope="strict")
|
||||
|
||||
|
||||
def _drift_error(path: "Path", bak_path: str) -> Dict[str, Any]:
|
||||
"""Error dict for external drift: the on-disk file wouldn't round-trip
|
||||
through the parser/serializer, so flushing would discard content added by a
|
||||
patch tool, shell append, manual edit, or sister session."""
|
||||
return {
|
||||
"success": False,
|
||||
"error": (
|
||||
f"Refusing to write {path.name}: file on disk has content that wouldn't round-trip "
|
||||
f"through the memory tool (likely added by the patch tool, a shell append, a manual edit, "
|
||||
f"or a concurrent session). A snapshot was saved to {bak_path}. Resolve the drift first — "
|
||||
f"either rewrite the file as a clean §-delimited list of entries, or move the extra "
|
||||
f"content out — then retry. This guard exists to prevent silent data loss (issue #26045)."
|
||||
),
|
||||
"drift_backup": bak_path,
|
||||
"remediation": (
|
||||
"Open the .bak file, integrate the missing entries into the memory tool one at a time via "
|
||||
"memory(action=add, content=...), then remove or rewrite the original file to a clean state."
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def _read_failed_error(path: "Path") -> Dict[str, Any]:
|
||||
"""Error dict for an unreadable (but existing) memory file. Treating it as
|
||||
empty and saving would rewrite the whole file from ``[]`` — wiping memory."""
|
||||
return {"success": False, "error": (
|
||||
f"Refusing to write {path.name}: the file exists on disk but could not be read right now "
|
||||
f"(temporarily locked by another program, a permission change, invalid/corrupt text encoding, "
|
||||
f"or a filesystem error). Treating an unreadable file as empty and saving would wipe existing "
|
||||
f"memory, so the write is refused. Nothing was changed — retry in a moment."
|
||||
)}
|
||||
|
||||
|
||||
def _find_unique_match(entries: List[str], old_text: str) -> Tuple[Optional[int], bool]:
|
||||
"""``(index, ambiguous)`` for entries containing *old_text*. Exact-duplicate
|
||||
matches are safe (first wins); distinct matches → ``(None, True)``."""
|
||||
matches = [i for i, e in enumerate(entries) if old_text in e]
|
||||
if not matches:
|
||||
return None, False
|
||||
if len({entries[i] for i in matches}) > 1:
|
||||
return None, True
|
||||
return matches[0], False
|
||||
|
||||
|
||||
class MemoryStore:
|
||||
"""Bounded curated memory with file persistence; one instance per AIAgent.
|
||||
``_system_prompt_snapshot`` is frozen at load time (prefix-cache stable);
|
||||
``memory_entries`` / ``user_entries`` are live state persisted to disk."""
|
||||
|
||||
# After this many failed consolidation attempts (overflow / zero-match) in
|
||||
# ONE turn, return a terminal "save skipped" result instead of "retry in
|
||||
# this turn", so a fragile replace/add can't loop the turn to budget
|
||||
# exhaustion and suppress the user's reply.
|
||||
_MAX_CONSOLIDATION_FAILURES_PER_TURN = 3
|
||||
|
||||
def __init__(self, memory_char_limit: int = 2200, user_char_limit: int = 1375, *,
|
||||
memory_enabled: bool = True, user_profile_enabled: bool = True):
|
||||
self.memory_entries: List[str] = []
|
||||
self.user_entries: List[str] = []
|
||||
self.memory_char_limit = memory_char_limit
|
||||
self.user_char_limit = user_char_limit
|
||||
self.memory_enabled = memory_enabled
|
||||
self.user_profile_enabled = user_profile_enabled
|
||||
self._system_prompt_snapshot: Dict[str, str] = {"memory": "", "user": ""}
|
||||
# Per-turn counter of failed at-capacity consolidation attempts; reset
|
||||
# at each turn boundary by reset_consolidation_failures().
|
||||
self._consolidation_failures = 0
|
||||
|
||||
def target_enabled(self, target: str) -> bool:
|
||||
"""Return whether this session's selected built-in store is writable."""
|
||||
return self.user_profile_enabled if target == "user" else self.memory_enabled
|
||||
|
||||
def reset_consolidation_failures(self) -> None:
|
||||
"""Reset the per-turn consolidation-failure counter (call at turn start)."""
|
||||
self._consolidation_failures = 0
|
||||
|
||||
def _consolidation_failure(self, response: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""Count an at-capacity consolidation failure. Under the per-turn cap
|
||||
return ``response`` unchanged (it says how to retry); past it return a
|
||||
TERMINAL result so the model stops looping — a failed memory side effect
|
||||
must never block the turn's reply."""
|
||||
self._consolidation_failures += 1
|
||||
if self._consolidation_failures <= self._MAX_CONSOLIDATION_FAILURES_PER_TURN:
|
||||
return response
|
||||
return {"success": False, "done": True, "error": (
|
||||
f"Memory consolidation failed {self._consolidation_failures} times this turn. Stop retrying "
|
||||
"memory calls — leave memory unchanged for now and continue with your reply to the user. "
|
||||
"The fact can be saved in a later turn."
|
||||
)}
|
||||
|
||||
def load_from_disk(self):
|
||||
"""Load MEMORY.md / USER.md and capture the frozen system-prompt snapshot.
|
||||
Threat hits are replaced in the SNAPSHOT only by a ``[BLOCKED: …]``
|
||||
placeholder; live lists keep the raw text so the user can see and remove
|
||||
poisoned entries (dropping them silently would hide the attack).
|
||||
Scanning is deterministic from disk bytes, so the snapshot stays stable."""
|
||||
mem_dir = _memory_dir()
|
||||
mem_dir.mkdir(parents=True, exist_ok=True)
|
||||
# Deduplicate (order-preserving, first occurrence wins).
|
||||
self.memory_entries = list(dict.fromkeys(self._read_file(mem_dir / "MEMORY.md")))
|
||||
self.user_entries = list(dict.fromkeys(self._read_file(mem_dir / "USER.md")))
|
||||
self._system_prompt_snapshot = {
|
||||
target: self._render_block(target, self._sanitize_entries_for_snapshot(entries, filename))
|
||||
for target, entries, filename in (
|
||||
("memory", self.memory_entries, "MEMORY.md"),
|
||||
("user", self.user_entries, "USER.md"),
|
||||
)
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _sanitize_entries_for_snapshot(entries: List[str], filename: str) -> List[str]:
|
||||
"""Return *entries* with any threat-matching entry replaced by a
|
||||
``[BLOCKED: …]`` placeholder (strict scope, same as writes). Empty or
|
||||
already-blocked entries pass through unchanged."""
|
||||
from tools.threat_patterns import scan_for_threats
|
||||
|
||||
sanitized: List[str] = []
|
||||
for entry in entries:
|
||||
findings = scan_for_threats(entry, scope="strict") if entry and not entry.startswith("[BLOCKED:") else None
|
||||
if not findings:
|
||||
sanitized.append(entry)
|
||||
continue
|
||||
logger.warning("Memory entry from %s blocked at load time: %s", filename, ", ".join(findings))
|
||||
sanitized.append(f"[BLOCKED: {filename} entry contained threat pattern(s): {', '.join(findings)}. "
|
||||
f"Removed from system prompt; use memory(action=remove) to delete the original.]")
|
||||
return sanitized
|
||||
|
||||
@staticmethod
|
||||
@contextmanager
|
||||
def _file_lock(path: Path):
|
||||
"""Exclusive lock for read-modify-write safety, on a separate .lock file
|
||||
so the memory file itself can still be atomically replaced."""
|
||||
from tools import memory_tool as _mt # fcntl/msvcrt live (and are patched) there
|
||||
|
||||
fcntl, msvcrt = _mt.fcntl, _mt.msvcrt
|
||||
lock_path = path.with_suffix(path.suffix + ".lock")
|
||||
lock_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
if fcntl is None and msvcrt is None:
|
||||
yield
|
||||
return
|
||||
fd = open(lock_path, "a+", encoding="utf-8")
|
||||
try:
|
||||
if fcntl:
|
||||
fcntl.flock(fd, fcntl.LOCK_EX)
|
||||
else:
|
||||
fd.seek(0)
|
||||
msvcrt.locking(fd.fileno(), msvcrt.LK_LOCK, 1)
|
||||
yield
|
||||
finally:
|
||||
try:
|
||||
if fcntl:
|
||||
fcntl.flock(fd, fcntl.LOCK_UN)
|
||||
else:
|
||||
fd.seek(0)
|
||||
msvcrt.locking(fd.fileno(), msvcrt.LK_UNLCK, 1)
|
||||
except OSError:
|
||||
pass
|
||||
fd.close()
|
||||
|
||||
@staticmethod
|
||||
def _path_for(target: str) -> Path:
|
||||
return _memory_dir() / ("USER.md" if target == "user" else "MEMORY.md")
|
||||
|
||||
def _reload_target(self, target: str, *, skip_drift: bool = False):
|
||||
"""Re-read entries from disk (under file lock) before mutating.
|
||||
|
||||
Returns ``None`` on a clean reload; the backup path (str) on external
|
||||
drift (caller must abort — flushing would discard un-roundtrippable
|
||||
content); or ``_READ_FAILED`` when the file exists but could not be
|
||||
read (caller MUST abort — rewriting from an assumed-empty view would
|
||||
wipe it; even append-only ``add`` rewrites the whole file). Drift check
|
||||
and parse use the SAME raw snapshot — a failed second read used to count
|
||||
as "no drift". *skip_drift* skips the round-trip check (``add``).
|
||||
"""
|
||||
raw, read_ok = self._read_raw_checked(self._path_for(target))
|
||||
if not read_ok:
|
||||
return _READ_FAILED
|
||||
bak = None if skip_drift else self._detect_external_drift(target, raw)
|
||||
self._set_entries(target, list(dict.fromkeys(self._parse_entries(raw))))
|
||||
return bak
|
||||
|
||||
def _reload_or_error(self, target: str, *, skip_drift: bool = False) -> Optional[Dict[str, Any]]:
|
||||
"""Reload under lock; return the abort error dict, or None to proceed."""
|
||||
bak = self._reload_target(target, skip_drift=skip_drift)
|
||||
if bak is _READ_FAILED:
|
||||
return _read_failed_error(self._path_for(target))
|
||||
if bak:
|
||||
return _drift_error(self._path_for(target), bak)
|
||||
return None
|
||||
|
||||
def save_to_disk(self, target: str):
|
||||
"""Persist entries to the appropriate file. Called after every mutation."""
|
||||
_memory_dir().mkdir(parents=True, exist_ok=True)
|
||||
self._write_file(self._path_for(target), self._entries_for(target))
|
||||
|
||||
def _entries_for(self, target: str) -> List[str]:
|
||||
return self.user_entries if target == "user" else self.memory_entries
|
||||
|
||||
def _set_entries(self, target: str, entries: List[str]):
|
||||
if target == "user":
|
||||
self.user_entries = entries
|
||||
else:
|
||||
self.memory_entries = entries
|
||||
|
||||
def _char_count(self, target: str) -> int:
|
||||
return len(ENTRY_DELIMITER.join(self._entries_for(target)))
|
||||
|
||||
def _char_limit(self, target: str) -> int:
|
||||
return self.user_char_limit if target == "user" else self.memory_char_limit
|
||||
|
||||
def _usage(self, target: str) -> str:
|
||||
return f"{self._char_count(target):,}/{self._char_limit(target):,}"
|
||||
|
||||
def _failure_with_entries(self, target: str, message: str) -> Dict[str, Any]:
|
||||
"""Consolidation failure that shows the live entries so the model can
|
||||
decide what to consolidate."""
|
||||
return self._consolidation_failure({"success": False, "error": message,
|
||||
"current_entries": self._entries_for(target), "usage": self._usage(target)})
|
||||
|
||||
def _locate(self, target: str, old_text: str, verb: str):
|
||||
"""Resolve *old_text* to a unique entry index, or an error dict."""
|
||||
entries = self._entries_for(target)
|
||||
idx, ambiguous = _find_unique_match(entries, old_text)
|
||||
if ambiguous:
|
||||
return None, {"success": False, "error": f"Multiple entries matched '{old_text}'. Be more specific.",
|
||||
"matches": self._previews([e for e in entries if old_text in e])}
|
||||
if idx is None:
|
||||
return None, self._consolidation_failure({
|
||||
"success": False,
|
||||
"error": f"No entry matched '{old_text}'. Check current_entries below and retry with the exact text of the entry you want to {verb}.",
|
||||
"current_entries": entries,
|
||||
})
|
||||
return idx, None
|
||||
|
||||
def _commit(self, target: str, entries: List[str], message: str) -> Dict[str, Any]:
|
||||
self._set_entries(target, entries)
|
||||
self.save_to_disk(target)
|
||||
return self._success_response(target, message)
|
||||
|
||||
def add(self, target: str, content: str) -> Dict[str, Any]:
|
||||
"""Append a new entry. Returns error if it would exceed the char limit."""
|
||||
content = content.strip()
|
||||
if not content:
|
||||
return {"success": False, "error": "Content cannot be empty."}
|
||||
scan_error = _scan_memory_content(content)
|
||||
if scan_error:
|
||||
return {"success": False, "error": scan_error}
|
||||
with self._file_lock(self._path_for(target)):
|
||||
# Append-only: skip the drift guard (appending never clobbers
|
||||
# un-roundtrippable content), but still refuse on a failed read —
|
||||
# add rewrites the WHOLE file from the parsed entries.
|
||||
err = self._reload_or_error(target, skip_drift=True)
|
||||
if err:
|
||||
return err
|
||||
entries = self._entries_for(target)
|
||||
limit = self._char_limit(target)
|
||||
if content in entries:
|
||||
return self._success_response(target, "Entry already exists (no duplicate added).")
|
||||
if len(ENTRY_DELIMITER.join(entries + [content])) > limit:
|
||||
return self._failure_with_entries(target, (
|
||||
f"Memory at {self._char_count(target):,}/{limit:,} chars. Adding this entry "
|
||||
f"({len(content)} chars) would exceed the limit. Consolidate now: use 'replace' to merge "
|
||||
f"overlapping entries into shorter ones or 'remove' stale or less important entries (see "
|
||||
f"current_entries below), then retry this add — all in this turn."
|
||||
))
|
||||
entries.append(content)
|
||||
return self._commit(target, entries, "Entry added.")
|
||||
|
||||
def replace(self, target: str, old_text: str, new_content: str) -> Dict[str, Any]:
|
||||
"""Find entry containing old_text substring, replace it with new_content."""
|
||||
old_text = old_text.strip()
|
||||
new_content = new_content.strip()
|
||||
if not old_text:
|
||||
return {"success": False, "error": "old_text cannot be empty."}
|
||||
if not new_content:
|
||||
return {"success": False, "error": "new_content cannot be empty. Use 'remove' to delete entries."}
|
||||
scan_error = _scan_memory_content(new_content)
|
||||
if scan_error:
|
||||
return {"success": False, "error": scan_error}
|
||||
with self._file_lock(self._path_for(target)):
|
||||
err = self._reload_or_error(target)
|
||||
if err:
|
||||
return err
|
||||
idx, err = self._locate(target, old_text, "replace")
|
||||
if err:
|
||||
return err
|
||||
entries = self._entries_for(target)
|
||||
limit = self._char_limit(target)
|
||||
new_total = len(ENTRY_DELIMITER.join(entries[:idx] + [new_content] + entries[idx + 1:]))
|
||||
if new_total > limit:
|
||||
return self._failure_with_entries(target, (
|
||||
f"Replacement would put memory at {new_total:,}/{limit:,} chars. Shorten the new content, "
|
||||
f"or 'remove' other stale or less important entries to make room (see current_entries "
|
||||
f"below), then retry — all in this turn."
|
||||
))
|
||||
entries[idx] = new_content
|
||||
return self._commit(target, entries, "Entry replaced.")
|
||||
|
||||
def remove(self, target: str, old_text: str) -> Dict[str, Any]:
|
||||
"""Remove the entry containing old_text substring."""
|
||||
old_text = old_text.strip()
|
||||
if not old_text:
|
||||
return {"success": False, "error": "old_text cannot be empty."}
|
||||
with self._file_lock(self._path_for(target)):
|
||||
err = self._reload_or_error(target)
|
||||
if err:
|
||||
return err
|
||||
idx, err = self._locate(target, old_text, "remove")
|
||||
if err:
|
||||
return err
|
||||
entries = self._entries_for(target)
|
||||
entries.pop(idx)
|
||||
return self._commit(target, entries, "Entry removed.")
|
||||
|
||||
# -- Batch --
|
||||
|
||||
@staticmethod
|
||||
def _apply_batch_op(working: List[str], act: str, content: str, old_text: str, pos: str) -> Optional[str]:
|
||||
"""Apply one batch op to *working* in place; return an error message or None."""
|
||||
if act == "add":
|
||||
if not content:
|
||||
return f"{pos}: content is required."
|
||||
if content not in working: # idempotent -- skip duplicate, don't fail the batch
|
||||
working.append(content)
|
||||
return None
|
||||
if act not in ("replace", "remove"):
|
||||
return f"{pos}: unknown action. Use add, replace, or remove."
|
||||
if not old_text:
|
||||
return f"{pos}: old_text is required."
|
||||
if act == "replace" and not content:
|
||||
return f"{pos}: content is required (use action='remove' to delete)."
|
||||
idx, ambiguous = _find_unique_match(working, old_text)
|
||||
if ambiguous:
|
||||
return f"{pos}: '{old_text}' matched multiple distinct entries -- be more specific."
|
||||
if idx is None:
|
||||
return f"{pos}: no entry matched '{old_text}'."
|
||||
if act == "replace":
|
||||
working[idx] = content
|
||||
else:
|
||||
working.pop(idx)
|
||||
return None
|
||||
|
||||
def apply_batch(self, target: str, operations: List[Dict[str, Any]]) -> Dict[str, Any]:
|
||||
"""Apply a sequence of add/replace/remove ops to one target atomically.
|
||||
|
||||
Ops are validated and applied against the FINAL budget only — so a single
|
||||
call can free space (remove/replace) and add new entries without the
|
||||
multi-turn consolidate-then-retry dance. All-or-nothing: if any op is
|
||||
malformed, doesn't match, or the net result exceeds the char limit,
|
||||
NOTHING is written and the first failure plus live state is returned.
|
||||
"""
|
||||
if not operations:
|
||||
return {"success": False, "error": "operations list is empty."}
|
||||
|
||||
# Scan every add/replace content BEFORE touching disk -- one poisoned
|
||||
# op rejects the whole batch.
|
||||
for i, op in enumerate(operations):
|
||||
op = op or {}
|
||||
scan_error = op.get("action") in {"add", "replace"} and op.get("content") and _scan_memory_content(op["content"])
|
||||
if scan_error:
|
||||
return {"success": False, "error": f"Operation {i + 1}: {scan_error}"}
|
||||
|
||||
with self._file_lock(self._path_for(target)):
|
||||
err = self._reload_or_error(target)
|
||||
if err:
|
||||
return err
|
||||
# Work on a copy; only commit if the whole batch validates.
|
||||
working: List[str] = list(self._entries_for(target))
|
||||
limit = self._char_limit(target)
|
||||
for i, op in enumerate(operations):
|
||||
op = op or {}
|
||||
act = op.get("action")
|
||||
content = (op.get("content") or op.get("new_text") or "").strip()
|
||||
old_text = (op.get("old_text") or "").strip()
|
||||
pos = f"Operation {i + 1} ({act or 'unknown'})"
|
||||
msg = self._apply_batch_op(working, act, content, old_text, pos)
|
||||
if msg:
|
||||
return self._batch_error(target, msg)
|
||||
# Budget check against the FINAL state only.
|
||||
new_total = len(ENTRY_DELIMITER.join(working))
|
||||
if new_total > limit:
|
||||
return self._failure_with_entries(target, (
|
||||
f"After applying all {len(operations)} operations, memory would be at "
|
||||
f"{new_total:,}/{limit:,} chars -- over the limit. Remove or shorten more "
|
||||
f"entries in the same batch (see current_entries below), then retry."
|
||||
))
|
||||
return self._commit(target, working, f"Applied {len(operations)} operation(s).")
|
||||
|
||||
def _batch_error(self, target: str, message: str) -> Dict[str, Any]:
|
||||
"""Build a batch-abort error that reports live (uncommitted) state."""
|
||||
return self._failure_with_entries(
|
||||
target, message + " No operations were applied (batch is all-or-nothing)."
|
||||
)
|
||||
|
||||
def format_for_system_prompt(self, target: str) -> Optional[str]:
|
||||
"""Return the frozen load-time snapshot for system-prompt injection (NOT
|
||||
live state — mid-session writes don't affect it, preserving the prefix
|
||||
cache). None if the snapshot is empty."""
|
||||
return self._system_prompt_snapshot.get(target, "") or None
|
||||
|
||||
# -- Internal helpers --
|
||||
|
||||
@staticmethod
|
||||
def _previews(entries: List[str], width: int = 80) -> List[str]:
|
||||
"""Truncated one-line previews of entries for error feedback."""
|
||||
return [e[:width] + ("..." if len(e) > width else "") for e in entries]
|
||||
|
||||
def _usage_pct(self, target: str, current: int) -> str:
|
||||
"""``"<pct>% — <current>/<limit> chars"`` for the given target."""
|
||||
limit = self._char_limit(target)
|
||||
pct = min(100, int((current / limit) * 100)) if limit > 0 else 0
|
||||
return f"{pct}% — {current:,}/{limit:,} chars"
|
||||
|
||||
def _success_response(self, target: str, message: str = None) -> Dict[str, Any]:
|
||||
# A successful write means the consolidation loop made progress, so the
|
||||
# per-turn failure budget resets (the cap counts consecutive failures).
|
||||
self._consolidation_failures = 0
|
||||
# Intentionally TERMINAL and without the entries list: echoing entries
|
||||
# invites the model to "find more to fix" and re-issue the same ops.
|
||||
# Entries are only shown on error/over-budget paths.
|
||||
resp = {"success": True, "done": True, "target": target,
|
||||
"usage": self._usage_pct(target, self._char_count(target)),
|
||||
"entry_count": len(self._entries_for(target))}
|
||||
if message:
|
||||
resp["message"] = message
|
||||
resp["note"] = "Write saved. This update is complete — do not repeat it."
|
||||
return resp
|
||||
|
||||
def _render_block(self, target: str, entries: List[str]) -> str:
|
||||
"""Render a system prompt block with header and usage indicator."""
|
||||
if not entries:
|
||||
return ""
|
||||
content = ENTRY_DELIMITER.join(entries)
|
||||
title = MEMORY_BLOCK_HEADERS["user" if target == "user" else "memory"]
|
||||
separator = "═" * 46
|
||||
return f"{separator}\n{title} [{self._usage_pct(target, len(content))}]\n{separator}\n{content}"
|
||||
|
||||
@staticmethod
|
||||
def _read_raw_checked(path: Path) -> Tuple[str, bool]:
|
||||
"""Read raw text as ``(raw, read_ok)``. ``read_ok`` is False ONLY when the
|
||||
file EXISTS but can't be read (absent file → ``("", True)``). Invalid
|
||||
UTF-8 counts as unreadable; decoding stays STRICT because
|
||||
``errors="replace"`` would hand callers a lossy view that a save then
|
||||
persists over the real bytes. ``utf-8-sig`` strips a Notepad BOM that
|
||||
otherwise glues U+FEFF onto the first entry forever."""
|
||||
if not path.exists():
|
||||
return "", True
|
||||
try:
|
||||
return path.read_text(encoding="utf-8-sig"), True
|
||||
except (OSError, UnicodeDecodeError):
|
||||
return "", False
|
||||
|
||||
@staticmethod
|
||||
def _parse_entries(raw: str) -> List[str]:
|
||||
"""Split raw memory-file text into stripped, non-empty entries. Splits on
|
||||
the full ENTRY_DELIMITER so a bare "§" inside an entry is preserved."""
|
||||
return [e for e in (x.strip() for x in raw.split(ENTRY_DELIMITER)) if e]
|
||||
|
||||
@staticmethod
|
||||
def _read_entries_checked(path: Path) -> Tuple[List[str], bool]:
|
||||
"""Read + parse as ``(entries, read_ok)`` — see ``_read_raw_checked``."""
|
||||
raw, read_ok = MemoryStore._read_raw_checked(path)
|
||||
return MemoryStore._parse_entries(raw), read_ok
|
||||
|
||||
@staticmethod
|
||||
def _read_file(path: Path) -> List[str]:
|
||||
"""Read a memory file into entries (empty list on any error). Only for
|
||||
read-only callers (``load_from_disk``, learning_mutations); mutation
|
||||
paths must use ``_read_raw_checked`` so they can refuse to overwrite an
|
||||
unreadable file."""
|
||||
return MemoryStore._read_entries_checked(path)[0]
|
||||
|
||||
def _detect_external_drift(self, target: str, raw: str) -> Optional[str]:
|
||||
"""Backup-path string if *raw* (the caller's checked-read snapshot) shows
|
||||
external drift, else None. Signals: round-trip mismatch, or one parsed
|
||||
entry exceeding the whole-file char limit (no tool-written entry can —
|
||||
an external writer appended free-form content a flush would truncate).
|
||||
The file is snapshotted to ``.bak.<ts>`` so the operator can recover it."""
|
||||
if not raw.strip():
|
||||
return None
|
||||
parsed = self._parse_entries(raw)
|
||||
if raw.strip() == ENTRY_DELIMITER.join(parsed) and max(map(len, parsed), default=0) <= self._char_limit(target):
|
||||
return None
|
||||
path = self._path_for(target)
|
||||
bak_path = path.with_suffix(path.suffix + f".bak.{int(time.time())}")
|
||||
try:
|
||||
bak_path.write_text(raw, encoding="utf-8")
|
||||
except OSError:
|
||||
return str(bak_path) + " (BACKUP FAILED — file unchanged on disk)"
|
||||
return str(bak_path)
|
||||
|
||||
@staticmethod
|
||||
def _write_file(path: Path, entries: List[str]):
|
||||
"""Atomic temp-file + rename: readers see the old or the new complete
|
||||
file, never a truncated one."""
|
||||
try:
|
||||
atomic_write_text(path, ENTRY_DELIMITER.join(entries), tmp_prefix=".mem_")
|
||||
except OSError as e:
|
||||
raise RuntimeError(f"Failed to write memory file {path}: {e}")
|
||||
File diff suppressed because it is too large
Load Diff
230
tools/session_search_tool_common.py
Normal file
230
tools/session_search_tool_common.py
Normal file
@@ -0,0 +1,230 @@
|
||||
"""Shared helpers for the session_search tool: source classification, lineage
|
||||
resolution, message storage state, and response shaping. Imported by
|
||||
``tools.session_search_tool`` (which re-exports the names) and
|
||||
``tools.session_search_tool_discover``."""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
from hermes_state_common import _RESET_END_REASONS
|
||||
|
||||
# Hidden from browsing/searching: integrations (HERMES_SESSION_SOURCE=tool),
|
||||
# delegate subagent runs, kanban workers — not the user's history.
|
||||
_HIDDEN_SESSION_SOURCES = ("kanban", "subagent", "tool")
|
||||
|
||||
# Searchable but DEMOTED below interactive sessions: cron sessions' repetitive
|
||||
# vocabulary dominates bare BM25 and starves out the user's own sessions
|
||||
# ("recall blindness"). Demoting keeps them reachable when they're the only match.
|
||||
_DEMOTED_SESSION_SOURCES = ("cron",)
|
||||
|
||||
# FTS rows scanned before dedup-by-lineage — well above the handful of distinct
|
||||
# sessions a query returns, so interactive matches buried under cron hits are
|
||||
# still in hand for the demotion pass.
|
||||
_DISCOVER_SCAN_LIMIT = 300
|
||||
|
||||
# Raw FTS rows are only a discovery-plan input; the response hydrates its own
|
||||
# anchored window and bookends after lineage dedup.
|
||||
_DISCOVER_SEARCH_FIELDS = ("id", "session_id", "role", "snippet", "source", "model", "session_started")
|
||||
|
||||
# Generated context-compaction handoff summaries (agent/context_compressor.py);
|
||||
# excluded from bookends so huge compaction payloads aren't re-introduced.
|
||||
_COMPACTION_PREFIXES = ("[CONTEXT COMPACTION", "[CONTEXT SUMMARY]:")
|
||||
|
||||
# /new, /reset, idle/daily expiry and CLI /new ("new_session") end the
|
||||
# predecessor WITHOUT carrying its transcript forward — unlike compression
|
||||
# continuations and live delegation children. Derived from the gateway set so
|
||||
# this tool and the recovery fence cannot drift.
|
||||
_FRESH_RESET_END_REASONS = frozenset(_RESET_END_REASONS) | {"new_session"}
|
||||
|
||||
|
||||
def _quiet(fn, default, msg, *log_args, with_exc: bool = False):
|
||||
"""Call ``fn()``; on any exception debug-log *msg* (appending the exception
|
||||
as a final ``%s`` arg when *with_exc*) and return *default*."""
|
||||
try:
|
||||
return fn()
|
||||
except Exception as e:
|
||||
logging.debug(msg, *(log_args + (e,) if with_exc else log_args), exc_info=True)
|
||||
return default
|
||||
|
||||
|
||||
def _format_timestamp(ts: Union[int, float, str, None]) -> str:
|
||||
"""Unix timestamp (number or numeric string) or ISO string -> readable date.
|
||||
"unknown" for None; str(ts) if conversion fails."""
|
||||
if ts is None:
|
||||
return "unknown"
|
||||
try:
|
||||
value = ts
|
||||
if isinstance(ts, str):
|
||||
if not ts.replace(".", "").replace("-", "").isdigit():
|
||||
return ts
|
||||
value = float(ts)
|
||||
if isinstance(value, (int, float)):
|
||||
return datetime.fromtimestamp(value).strftime("%B %d, %Y at %I:%M %p")
|
||||
except (ValueError, OSError, OverflowError) as e:
|
||||
logging.debug("Failed to format timestamp %s: %s", ts, e, exc_info=True)
|
||||
except Exception as e:
|
||||
logging.debug("Unexpected error formatting timestamp %s: %s", ts, e, exc_info=True)
|
||||
return str(ts)
|
||||
|
||||
|
||||
def _session_meta_block(meta: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""The ``session_meta`` sub-object shared by read/scroll responses."""
|
||||
return {"when": _format_timestamp(meta.get("started_at")), "source": meta.get("source"),
|
||||
"model": meta.get("model"), "title": meta.get("title")}
|
||||
|
||||
|
||||
def _ok(**payload) -> str:
|
||||
"""Serialize a successful tool result (``success`` first, then *payload* in order)."""
|
||||
return json.dumps({"success": True, **payload}, ensure_ascii=False)
|
||||
|
||||
|
||||
def _is_compaction_summary(content: str) -> bool:
|
||||
"""Return True if *content* looks like a generated compaction handoff."""
|
||||
return bool(content) and content.lstrip().startswith(_COMPACTION_PREFIXES)
|
||||
|
||||
|
||||
def _resolve_to_parent(db, session_id: str) -> tuple[str, bool]:
|
||||
"""Walk parent_session_id to the lineage root -> ``(root_id, has_compression_hop)``.
|
||||
The flag distinguishes a compression-split lineage (parent content summarised
|
||||
away) from a delegation lineage (child content still visible to the parent).
|
||||
Falls back to ``(session_id, False)`` on errors."""
|
||||
if not session_id:
|
||||
return session_id, False
|
||||
visited: set[str] = set()
|
||||
cur, has_compression = session_id, False
|
||||
while cur and cur not in visited:
|
||||
visited.add(cur)
|
||||
s = _quiet(lambda: db.get_session(cur), None, "Error resolving parent for %s: %s", cur, with_exc=True)
|
||||
if not s:
|
||||
break
|
||||
if s.get("end_reason") == "compression":
|
||||
has_compression = True
|
||||
if not s.get("parent_session_id"):
|
||||
break
|
||||
cur = s["parent_session_id"]
|
||||
return cur, has_compression
|
||||
|
||||
|
||||
def _resolve_lineage(db, session_id: str) -> str:
|
||||
"""Return only the lineage root (ignores the compression hop)."""
|
||||
return _resolve_to_parent(db, session_id)[0]
|
||||
|
||||
|
||||
def _session_end_reason(db, session_id: str) -> Optional[str]:
|
||||
"""Return the session's ``end_reason``, or None if missing/unended/error."""
|
||||
if not session_id:
|
||||
return None
|
||||
try:
|
||||
s = db.get_session(session_id)
|
||||
return (s.get("end_reason") or None) if s else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _session_left_live_context(db, session_id: str) -> bool:
|
||||
"""True when *session_id*'s transcript is no longer in anyone's live context:
|
||||
``compression`` (summarised into the continuation child) or a fresh reset
|
||||
(child starts empty). Everything else stays excluded from same-lineage
|
||||
recall — live delegation children (``end_reason is None``) are visible to
|
||||
the parent agent, and ``branched`` parents were copied verbatim into the
|
||||
branch child, so their content IS the current context."""
|
||||
end_reason = _session_end_reason(db, session_id)
|
||||
return end_reason == "compression" or end_reason in _FRESH_RESET_END_REASONS
|
||||
|
||||
|
||||
def _get_message_storage_state(db, message_id) -> Optional[Dict[str, Any]]:
|
||||
"""Return the owning session and visibility flags for *message_id*."""
|
||||
if not message_id:
|
||||
return None
|
||||
|
||||
def _lookup():
|
||||
with db._lock:
|
||||
return db._conn.execute(
|
||||
"SELECT session_id, active, compacted FROM messages WHERE id = ?", (message_id,)
|
||||
).fetchone()
|
||||
|
||||
row = _quiet(_lookup, None, "message storage-state lookup failed for %s", message_id)
|
||||
return dict(row) if row is not None else None
|
||||
|
||||
|
||||
def _is_compacted_state(state: Optional[Dict[str, Any]]) -> bool:
|
||||
"""Compaction archives are ``active=0, compacted=1`` (content summarised
|
||||
away by archive_and_compact). Rewind/undo rows are ``active=0, compacted=0``
|
||||
and must stay hidden."""
|
||||
return state is not None and state["active"] == 0 and state["compacted"] == 1
|
||||
|
||||
|
||||
def _is_compacted_message(db, message_id) -> bool:
|
||||
"""True if *message_id* is a compaction-archived row — pre-compaction content
|
||||
no longer in live context, so it should stay discoverable even on the
|
||||
current session. False on any error (caller falls back to skipping)."""
|
||||
return _is_compacted_state(_get_message_storage_state(db, message_id))
|
||||
|
||||
|
||||
def _annotate_rebuild_status(db, payload: Dict[str, Any]) -> None:
|
||||
"""Add a rebuild-progress note while the deferred FTS backfill is running,
|
||||
so the agent can explain thin/slow results instead of treating them as
|
||||
ground truth. No-op (never raises) when no rebuild is pending."""
|
||||
try:
|
||||
status = db.fts_rebuild_status()
|
||||
except Exception:
|
||||
return
|
||||
if status is None:
|
||||
return
|
||||
payload["index_rebuild"] = {"percent": status["percent"], "note": (
|
||||
f"The search index is rebuilding in the background ({status['percent']}% done, "
|
||||
f"{status['indexed']:,} of {status['total']:,} messages). Results from older messages "
|
||||
f"may be incomplete until it finishes."
|
||||
)}
|
||||
|
||||
|
||||
def _order_for_recall(raw_results: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
|
||||
"""Stable-sort FTS rows so interactive sessions rank above automation.
|
||||
BM25 order is preserved within each class; only cross-class order changes,
|
||||
so a cron hit never displaces an interactive hit during lineage dedup."""
|
||||
return sorted(raw_results, key=lambda r: 1 if (r.get("source") or "") in _DEMOTED_SESSION_SOURCES else 0)
|
||||
|
||||
|
||||
def _shape_message(m: Dict[str, Any], anchor_id: Optional[int] = None,
|
||||
max_content_len: Optional[int] = None) -> Dict[str, Any]:
|
||||
"""Slim a message row for the tool response. Keeps content even if empty
|
||||
(absent content is meaningful — tool-call-only assistant turns). With
|
||||
*max_content_len*, content is truncated and ``content_truncated`` /
|
||||
``original_content_chars`` added."""
|
||||
content = m.get("content")
|
||||
if isinstance(content, str) and "\x1b" in content:
|
||||
# Recalled messages can carry ANSI escapes (archived terminal output).
|
||||
from tools.ansi_strip import strip_ansi
|
||||
|
||||
content = strip_ansi(content)
|
||||
original_chars = None
|
||||
if max_content_len and content and len(content) > max_content_len:
|
||||
original_chars = len(content)
|
||||
content = content[:max_content_len] + "…"
|
||||
entry = {"id": m.get("id"), "role": m.get("role"), "content": content, "timestamp": m.get("timestamp")}
|
||||
entry.update({k: m.get(k) for k in ("tool_name", "tool_calls", "tool_call_id") if m.get(k)})
|
||||
if anchor_id is not None and m.get("id") == anchor_id:
|
||||
entry["anchor"] = True
|
||||
if original_chars is not None:
|
||||
entry["content_truncated"] = True
|
||||
entry["original_content_chars"] = original_chars
|
||||
return {k: v for k, v in entry.items() if v is not None or k == "content"}
|
||||
|
||||
|
||||
def _session_link(session_id: str, profile: str = None) -> str:
|
||||
"""The reference the agent writes to point the user at a session — same
|
||||
value the desktop composer emits, so it renders as a titled link. The
|
||||
profile segment is omitted when it can't be named confidently (a bare id
|
||||
still resolves, it just can't disambiguate across profiles)."""
|
||||
name = (profile or "").strip()
|
||||
if not name:
|
||||
def _active():
|
||||
from hermes_cli.profiles import get_active_profile_name
|
||||
|
||||
resolved = get_active_profile_name()
|
||||
return "" if resolved == "custom" else resolved
|
||||
|
||||
name = _quiet(_active, "", "get_active_profile_name failed for session link")
|
||||
return f"@session:{name}/{session_id}" if name else f"@session:{session_id}"
|
||||
192
tools/session_search_tool_discover.py
Normal file
192
tools/session_search_tool_discover.py
Normal file
@@ -0,0 +1,192 @@
|
||||
"""Discovery shape of session_search: FTS5 query, title match, lineage dedup
|
||||
and adaptive/full hydration of the surviving results."""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from tools.registry import tool_error
|
||||
from tools.session_search_tool_common import (
|
||||
_DISCOVER_SCAN_LIMIT, _DISCOVER_SEARCH_FIELDS, _HIDDEN_SESSION_SOURCES, _annotate_rebuild_status,
|
||||
_format_timestamp, _is_compacted_message, _is_compaction_summary, _order_for_recall, _quiet,
|
||||
_resolve_lineage, _resolve_to_parent, _session_left_live_context, _session_link, _shape_message,
|
||||
)
|
||||
|
||||
|
||||
def _normalize_title_query(query: str) -> str:
|
||||
"""Strip common quoting the model may include around a remembered title."""
|
||||
return query.strip().strip("`'\"")
|
||||
|
||||
|
||||
def _title_match_result(db, query: str, current_lineage_root: Optional[str]) -> Optional[Dict[str, Any]]:
|
||||
"""Return a discovery-shaped result when the query matches a session title."""
|
||||
title_query = _normalize_title_query(query)
|
||||
if not title_query:
|
||||
return None
|
||||
session_id = _quiet(lambda: db.resolve_session_by_title(title_query), None,
|
||||
"resolve_session_by_title failed for %r", title_query)
|
||||
if not session_id:
|
||||
return None
|
||||
lineage_root = _resolve_lineage(db, session_id)
|
||||
# Same-lineage title hits are in-context only while the session is live;
|
||||
# /new-reset and compression-ended parents are not.
|
||||
if (
|
||||
current_lineage_root
|
||||
and lineage_root == current_lineage_root
|
||||
and not _session_left_live_context(db, session_id)
|
||||
):
|
||||
return None
|
||||
|
||||
session_meta = _quiet(lambda: db.get_session(lineage_root) or db.get_session(session_id), None,
|
||||
"get_session failed for title match %s", session_id) or {}
|
||||
if session_meta.get("source") in _HIDDEN_SESSION_SOURCES:
|
||||
return None
|
||||
messages = _quiet(lambda: db.get_messages(session_id), [], "get_messages failed for title match %s", session_id)
|
||||
anchor_id = messages[0].get("id") if messages else None
|
||||
view = {}
|
||||
if anchor_id is not None:
|
||||
view = _quiet(lambda: db.get_anchored_view(session_id, anchor_id, window=5, bookend=3), {},
|
||||
"get_anchored_view failed for title match %s/%s", session_id, anchor_id)
|
||||
entry = {
|
||||
"session_id": session_id, "when": _format_timestamp(session_meta.get("started_at")),
|
||||
"source": session_meta.get("source", "unknown"), "model": session_meta.get("model") or "unknown",
|
||||
"title": session_meta.get("title") or title_query, "matched_role": "session_title",
|
||||
"match_message_id": anchor_id,
|
||||
"snippet": f"Session title matched: {session_meta.get('title') or title_query}",
|
||||
"bookend_start": [_shape_message(m) for m in (view.get("bookend_start") or messages[:3])],
|
||||
"messages": [_shape_message(m, anchor_id=anchor_id) for m in (view.get("window") or messages[:5])],
|
||||
"bookend_end": [_shape_message(m) for m in (view.get("bookend_end") or messages[-3:])],
|
||||
"messages_before": view.get("messages_before", 0),
|
||||
"messages_after": view.get("messages_after", max(len(messages) - 5, 0)),
|
||||
"detail": "full", "_lineage_root": lineage_root,
|
||||
}
|
||||
if lineage_root and lineage_root != session_id:
|
||||
entry["parent_session_id"] = lineage_root
|
||||
return entry
|
||||
|
||||
|
||||
def _discover_payload(db, query: str, detail: str, results: list, **extra) -> str:
|
||||
payload = {"success": True, "mode": "discover", "query": query, "detail": detail,
|
||||
"results": results, "count": len(results), **extra}
|
||||
_annotate_rebuild_status(db, payload)
|
||||
return json.dumps(payload, ensure_ascii=False)
|
||||
|
||||
|
||||
def _dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root) -> None:
|
||||
"""Fill *seen_sessions* (lineage_root -> first surviving FTS row) up to *limit*.
|
||||
|
||||
The raw owning session_id stays on the row — only it pairs validly with the
|
||||
FTS match id for the anchored window. Current-lineage hits are skipped
|
||||
UNLESS the transcript left live context: compression-ended session, /new-
|
||||
reset predecessor (hiding it made gateway recall blind after every /new),
|
||||
or an in-place compacted row on the SAME session_id. A live delegation
|
||||
child has end_reason=None, so it stays excluded.
|
||||
"""
|
||||
for r in raw_results:
|
||||
if len(seen_sessions) >= limit:
|
||||
break
|
||||
raw_sid = r["session_id"]
|
||||
resolved_sid, _ = _resolve_to_parent(db, raw_sid)
|
||||
is_compacted_hit = _is_compacted_message(db, r.get("id"))
|
||||
is_ended_session = _session_left_live_context(db, raw_sid)
|
||||
if (
|
||||
current_lineage_root
|
||||
and resolved_sid == current_lineage_root
|
||||
and not (is_ended_session or is_compacted_hit)
|
||||
):
|
||||
continue
|
||||
if current_session_id and raw_sid == current_session_id and not is_compacted_hit:
|
||||
continue
|
||||
seen_sessions.setdefault(resolved_sid, {**r, "_lineage_root": resolved_sid})
|
||||
|
||||
|
||||
def _bookend(view: Dict[str, Any], key: str) -> List[Dict[str, Any]]:
|
||||
return [_shape_message(m, max_content_len=1200) for m in (view.get(key) or [])
|
||||
if not _is_compaction_summary(m.get("content", ""))]
|
||||
|
||||
|
||||
def _hydrate_hit(db, lineage_root: str, match_info: Dict[str, Any], result_detail: str) -> Optional[Dict[str, Any]]:
|
||||
"""Build one discovery result from a surviving FTS row; None if the anchored
|
||||
view can't be loaded (the hit is dropped)."""
|
||||
hit_sid = match_info.get("session_id") or lineage_root
|
||||
msg_id = match_info.get("id")
|
||||
try:
|
||||
view = db.get_anchored_view(hit_sid, msg_id, window=5, bookend=3)
|
||||
except Exception as e:
|
||||
logging.warning("get_anchored_view failed for %s/%s: %s", hit_sid, msg_id, e, exc_info=True)
|
||||
return None
|
||||
session_meta = _quiet(lambda: db.get_session(lineage_root), None, "get_session failed for %s", lineage_root) or {}
|
||||
full = result_detail == "full"
|
||||
window_messages = view.get("window") or []
|
||||
if not full:
|
||||
window_messages = [m for m in window_messages if m.get("id") == msg_id]
|
||||
entry = {
|
||||
"session_id": hit_sid,
|
||||
"when": _format_timestamp(session_meta.get("started_at") or match_info.get("session_started")),
|
||||
"source": session_meta.get("source") or match_info.get("source", "unknown"),
|
||||
"model": session_meta.get("model") or match_info.get("model") or "unknown",
|
||||
"title": session_meta.get("title") or None, "matched_role": match_info.get("role"),
|
||||
"match_message_id": msg_id, "snippet": match_info.get("snippet") or "",
|
||||
"bookend_start": _bookend(view, "bookend_start") if full else [],
|
||||
"messages": [_shape_message(m, anchor_id=msg_id, max_content_len=4000) for m in window_messages],
|
||||
"bookend_end": _bookend(view, "bookend_end") if full else [],
|
||||
"messages_before": view.get("messages_before", 0), "messages_after": view.get("messages_after", 0),
|
||||
"detail": result_detail,
|
||||
}
|
||||
if lineage_root and lineage_root != hit_sid:
|
||||
entry["parent_session_id"] = lineage_root
|
||||
return entry
|
||||
|
||||
|
||||
def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort: Optional[str],
|
||||
detail: str, current_session_id: str = None, link_profile: str = None) -> str:
|
||||
"""Discovery shape: FTS5 plus adaptive or full result hydration."""
|
||||
role_list = role_filter if role_filter else ["user", "assistant"]
|
||||
current_lineage_root = _resolve_lineage(db, current_session_id) if current_session_id else None
|
||||
title_result = _title_match_result(db, query, current_lineage_root)
|
||||
|
||||
try:
|
||||
raw_results = db.search_messages(
|
||||
query=query, role_filter=role_list, exclude_sources=list(_HIDDEN_SESSION_SOURCES),
|
||||
limit=_DISCOVER_SCAN_LIMIT, offset=0, sort=sort, fields=_DISCOVER_SEARCH_FIELDS,
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error("FTS5 search failed: %s", e, exc_info=True)
|
||||
return tool_error(f"Search failed: {e}", success=False)
|
||||
|
||||
# Demote cron rows below interactive ones BEFORE dedup so a high-volume
|
||||
# cron corpus can't starve the user's own sessions out of the top `limit`.
|
||||
raw_results = _order_for_recall(raw_results)
|
||||
|
||||
if not raw_results and not title_result:
|
||||
return _discover_payload(db, query, detail, [], message=(
|
||||
"No matching sessions found. FTS5 ANDs all terms by default — "
|
||||
"broaden with OR (`alpha OR beta`), exact-match with quoted "
|
||||
"phrases, exclude with NOT, or prefix-match with `deploy*`."
|
||||
))
|
||||
|
||||
seen_sessions: Dict[str, Dict[str, Any]] = {}
|
||||
results = []
|
||||
if title_result:
|
||||
title_lineage = title_result.pop("_lineage_root", None)
|
||||
if title_lineage:
|
||||
seen_sessions[title_lineage] = {"_title_only": True}
|
||||
results.append(title_result)
|
||||
_dedupe_by_lineage(db, raw_results, limit, seen_sessions, current_session_id, current_lineage_root)
|
||||
|
||||
for lineage_root, match_info in seen_sessions.items():
|
||||
if match_info.get("_title_only"):
|
||||
continue
|
||||
# Adaptive: only the top-ranked result is fully hydrated.
|
||||
entry = _hydrate_hit(db, lineage_root, match_info, "full" if detail == "full" or not results else "compact")
|
||||
if entry is not None:
|
||||
results.append(entry)
|
||||
for entry in results:
|
||||
entry["link"] = _session_link(entry["session_id"], link_profile)
|
||||
return _discover_payload(db, query, detail, results, sessions_searched=len(seen_sessions), link_hint=(
|
||||
"When referring the user to a session, write its `link` value "
|
||||
"verbatim inline mid-sentence (it renders as a titled link) — never "
|
||||
"as markdown, in backticks, on its own line, or next to the "
|
||||
"title/id/date. To read more around a compact result, scroll: "
|
||||
"session_search(session_id=..., around_message_id=match_message_id)."
|
||||
))
|
||||
1590
tools/skills_sync.py
1590
tools/skills_sync.py
File diff suppressed because it is too large
Load Diff
240
tools/skills_sync_bundled_ops.py
Normal file
240
tools/skills_sync_bundled_ops.py
Normal file
@@ -0,0 +1,240 @@
|
||||
"""Bundled-skill maintenance ops: reset, diff, list-modified, opt-out, remove-pristine.
|
||||
|
||||
Extracted from ``tools.skills_sync``. Profile-scoped paths and patchable helpers
|
||||
(``_get_bundled_dir``, ``sync_skills``, ...) are resolved through ``_ss()`` at
|
||||
call time so monkeypatching ``tools.skills_sync`` keeps working.
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
|
||||
def _ss():
|
||||
from tools import skills_sync
|
||||
|
||||
return skills_sync
|
||||
|
||||
|
||||
def _is_tracked_user_modification(origin_hash: str, user_hash: str) -> bool:
|
||||
"""Whether an on-disk skill is a user modification ``hermes update`` keeps. Shared by
|
||||
the sync loop and ``list_user_modified_bundled_skills`` so they never drift: needs a
|
||||
recorded origin hash (un-baselined v1 entries don't count) AND differing content."""
|
||||
return bool(origin_hash) and user_hash != origin_hash
|
||||
|
||||
|
||||
def _bundled_by_name(bundled_dir: Path) -> dict:
|
||||
return dict(_ss()._discover_bundled_skills(bundled_dir))
|
||||
|
||||
|
||||
def reset_bundled_skill(name: str, restore: bool = False) -> dict:
|
||||
"""Reset a bundled skill's manifest tracking so future syncs work normally.
|
||||
|
||||
An edited bundled skill stays ``user_modified`` forever — even after copying the
|
||||
bundled version back — because the manifest holds the OLD origin hash; clearing
|
||||
the entry breaks that loop. ``restore`` also deletes the user's copy so the next
|
||||
sync re-copies bundled. Returns ``{ok, action, message, synced}``; action is
|
||||
manifest_cleared / restored / not_in_manifest / bundled_missing / not_reset.
|
||||
"""
|
||||
ss = _ss()
|
||||
manifest = ss._read_manifest()
|
||||
bundled_dir = ss._get_bundled_dir()
|
||||
bundled_by_name = _bundled_by_name(bundled_dir)
|
||||
in_manifest = name in manifest
|
||||
is_bundled = name in bundled_by_name
|
||||
|
||||
def _fail(action: str, message: str) -> dict:
|
||||
return {"ok": False, "action": action, "message": message, "synced": None}
|
||||
|
||||
if not in_manifest and not is_bundled:
|
||||
return _fail("not_in_manifest", f"'{name}' is not a tracked bundled skill. Nothing to reset. "
|
||||
f"(Hub-installed skills use `hermes skills uninstall`.)")
|
||||
|
||||
# Delete the user's copy BEFORE touching the manifest so a failed rmtree
|
||||
# cannot leave the skill in a manifest-less limbo state.
|
||||
deleted_user_copy = False
|
||||
if restore:
|
||||
if not is_bundled:
|
||||
return _fail("bundled_missing", f"'{name}' has no bundled source — manifest entry preserved "
|
||||
f"but cannot restore from bundled (skill was removed upstream).")
|
||||
dest = ss._compute_relative_dest(bundled_by_name[name], bundled_dir)
|
||||
if dest.exists():
|
||||
try:
|
||||
ss._rmtree_writable(dest)
|
||||
except (OSError, IOError) as e:
|
||||
return _fail("not_reset", f"Could not delete user copy at {dest}: {e}. "
|
||||
f"Manifest entry preserved — nothing was changed.")
|
||||
deleted_user_copy = True
|
||||
|
||||
if in_manifest:
|
||||
del manifest[name]
|
||||
ss._write_manifest(manifest)
|
||||
synced = ss.sync_skills(quiet=True)
|
||||
|
||||
if not restore:
|
||||
action, message = "manifest_cleared", (f"Cleared manifest entry for '{name}'. Future `hermes update` runs "
|
||||
f"will re-baseline against your current copy and accept upstream changes.")
|
||||
elif deleted_user_copy:
|
||||
action, message = "restored", f"Restored '{name}' from bundled source."
|
||||
else:
|
||||
action, message = "restored", f"Restored '{name}' (no prior user copy, re-copied from bundled)."
|
||||
return {"ok": True, "action": action, "message": message, "synced": synced}
|
||||
|
||||
|
||||
def list_user_modified_bundled_skills() -> List[dict]:
|
||||
"""Bundled skills ``hermes update`` keeps because the user edited them (same test
|
||||
the sync loop uses). Name-sorted ``{"name", "dest", "bundled_src"}`` dicts."""
|
||||
ss = _ss()
|
||||
manifest = ss._read_manifest()
|
||||
if not manifest:
|
||||
return []
|
||||
bundled_dir = ss._get_bundled_dir()
|
||||
modified: List[dict] = []
|
||||
for skill_name, skill_dir in ss._discover_bundled_skills(bundled_dir):
|
||||
origin_hash = manifest.get(skill_name, "") # empty = untracked/un-baselined v1: next sync handles it
|
||||
dest = ss._compute_relative_dest(skill_dir, bundled_dir)
|
||||
if origin_hash and dest.exists() and _is_tracked_user_modification(origin_hash, ss._dir_hash(dest)):
|
||||
modified.append({"name": skill_name, "dest": dest, "bundled_src": skill_dir})
|
||||
return sorted(modified, key=lambda e: e["name"])
|
||||
|
||||
|
||||
def _read_for_diff(path: Path) -> Tuple[Optional[bytes], Optional[str]]:
|
||||
"""Read a file once for diffing: ``(raw_bytes, text)`` with ``text=None`` for
|
||||
binary content, ``(None, None)`` if unreadable."""
|
||||
try:
|
||||
data = path.read_bytes()
|
||||
return data, (None if b"\x00" in data else data.decode("utf-8"))
|
||||
except OSError:
|
||||
return None, None
|
||||
except UnicodeDecodeError:
|
||||
return data, None
|
||||
|
||||
|
||||
def diff_bundled_skill(name: str) -> dict:
|
||||
"""Diff a user's copy of a bundled skill against stock. Returns ``{ok, name, found,
|
||||
modified, message, diffs}``; each diff is ``{"path", "status", "diff"}`` with status
|
||||
modified / added (only in user copy) / removed (only in bundled) / binary."""
|
||||
import difflib
|
||||
|
||||
from tools.skills_sync_optional import _skill_file_list
|
||||
|
||||
ss = _ss()
|
||||
|
||||
def _fail(found: bool, message: str) -> dict:
|
||||
return {"ok": False, "name": name, "found": found, "modified": False, "diffs": [], "message": message}
|
||||
|
||||
bundled_dir = ss._get_bundled_dir()
|
||||
bundled_src = _bundled_by_name(bundled_dir).get(name)
|
||||
if bundled_src is None:
|
||||
return _fail(False, f"'{name}' is not a tracked bundled skill (no stock version to "
|
||||
f"diff against). Hub-installed skills use `hermes skills inspect`.")
|
||||
dest = ss._compute_relative_dest(bundled_src, bundled_dir)
|
||||
if not dest.exists():
|
||||
return _fail(True, f"No local copy of '{name}' found at {dest}.")
|
||||
|
||||
user_files = set(_skill_file_list(dest))
|
||||
stock_files = set(_skill_file_list(bundled_src))
|
||||
diffs: List[dict] = []
|
||||
for rel in sorted(user_files | stock_files):
|
||||
if rel not in stock_files:
|
||||
diffs.append({"path": rel, "status": "added", "diff": f"+ only in your copy: {rel}"})
|
||||
elif rel not in user_files:
|
||||
diffs.append({"path": rel, "status": "removed", "diff": f"- only in stock: {rel}"})
|
||||
else:
|
||||
user_bytes, user_text = _read_for_diff(dest / rel)
|
||||
stock_bytes, stock_text = _read_for_diff(bundled_src / rel)
|
||||
if user_text is None or stock_text is None:
|
||||
# At least one side is binary — report only if the bytes differ.
|
||||
if user_bytes != stock_bytes:
|
||||
diffs.append({"path": rel, "status": "binary", "diff": "<binary file differs>"})
|
||||
elif user_text != stock_text:
|
||||
text = "".join(difflib.unified_diff(
|
||||
stock_text.splitlines(keepends=True), user_text.splitlines(keepends=True),
|
||||
fromfile=f"stock/{rel}", tofile=f"yours/{rel}",
|
||||
))
|
||||
diffs.append({"path": rel, "status": "modified", "diff": text})
|
||||
|
||||
message = (f"'{name}' differs from the stock version in {len(diffs)} file(s)." if diffs
|
||||
else f"'{name}' matches the stock version.")
|
||||
return {"ok": True, "name": name, "found": True, "modified": bool(diffs), "diffs": diffs, "message": message}
|
||||
|
||||
|
||||
_OPT_OUT_MESSAGES = { # (enabled, changed) -> message
|
||||
(True, True): "Opted out of bundled skills. Future install / update / sync runs will not seed bundled skills into this profile.",
|
||||
(True, False): "Already opted out — marker was already present.",
|
||||
(False, True): "Opted back in. The next `hermes update` (or `hermes skills opt-in --sync`) will re-seed bundled skills.",
|
||||
(False, False): "Not opted out — no marker to remove.",
|
||||
}
|
||||
|
||||
|
||||
def set_bundled_skills_opt_out(enabled: bool) -> dict:
|
||||
"""Toggle the .no-bundled-skills marker: the on-disk half of ``hermes skills
|
||||
opt-out`` / ``opt-in`` that stops installer/update/sync seeding. Removing
|
||||
already-present skills is a separate step (``remove_pristine_bundled_skills``).
|
||||
Returns ``{ok, changed, marker, message}``."""
|
||||
ss = _ss()
|
||||
marker = ss._hermes_home() / ss.NO_BUNDLED_SKILLS_MARKER
|
||||
existed = marker.exists()
|
||||
try:
|
||||
if enabled:
|
||||
ss._hermes_home().mkdir(parents=True, exist_ok=True)
|
||||
marker.write_text(
|
||||
"This profile opted out of bundled-skill seeding (`hermes skills opt-out`).\n"
|
||||
"Delete this file to re-enable sync on the next `hermes update`.\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
elif existed:
|
||||
marker.unlink()
|
||||
except OSError as e:
|
||||
return {"ok": False, "changed": False, "marker": str(marker),
|
||||
"message": f"Could not update opt-out marker at {marker}: {e}"}
|
||||
changed = enabled != existed
|
||||
return {"ok": True, "changed": changed, "marker": str(marker), "message": _OPT_OUT_MESSAGES[(enabled, changed)]}
|
||||
|
||||
|
||||
def remove_pristine_bundled_skills(dry_run: bool = False) -> dict:
|
||||
"""Delete bundled skills that are present, manifest-tracked, AND unmodified.
|
||||
|
||||
Removed ONLY when in the sync manifest (genuinely bundled, not hub/hand-written),
|
||||
still in the bundled source (hash-comparable), and byte-identical to the origin
|
||||
hash; everything else lands in ``skipped``. Removed skills lose their manifest
|
||||
entry so a later opt-in re-seed treats them as new.
|
||||
Returns ``{ok, removed, skipped: [{name, reason}], dry_run, message}``.
|
||||
"""
|
||||
ss = _ss()
|
||||
manifest = ss._read_manifest()
|
||||
bundled_dir = ss._get_bundled_dir()
|
||||
bundled_by_name = _bundled_by_name(bundled_dir)
|
||||
|
||||
removed: List[str] = []
|
||||
skipped: List[dict] = []
|
||||
for name, origin_hash in sorted(manifest.items()):
|
||||
src = bundled_by_name.get(name)
|
||||
if src is None:
|
||||
skipped.append({"name": name, "reason": "no bundled source (removed upstream)"})
|
||||
continue
|
||||
dest = ss._compute_relative_dest(src, bundled_dir)
|
||||
if not dest.exists():
|
||||
# Already gone from disk; just forget the stale manifest entry.
|
||||
if not dry_run:
|
||||
manifest.pop(name, None)
|
||||
continue
|
||||
if ss._dir_hash(dest) != origin_hash:
|
||||
skipped.append({"name": name, "reason": "user-modified (kept)"})
|
||||
continue
|
||||
if not dry_run:
|
||||
try:
|
||||
ss._rmtree_writable(dest)
|
||||
except (OSError, IOError) as e:
|
||||
skipped.append({"name": name, "reason": f"delete failed: {e}"})
|
||||
continue
|
||||
manifest.pop(name, None)
|
||||
removed.append(name)
|
||||
|
||||
if not dry_run and removed:
|
||||
ss._write_manifest(manifest)
|
||||
|
||||
verb = "Would remove" if dry_run else "Removed"
|
||||
return {
|
||||
"ok": True, "removed": removed, "skipped": skipped, "dry_run": dry_run,
|
||||
"message": f"{verb} {len(removed)} pristine bundled skill(s); kept {len(skipped)}.",
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
381
tools/skills_sync_client_org.py
Normal file
381
tools/skills_sync_client_org.py
Normal file
@@ -0,0 +1,381 @@
|
||||
"""Org-shared skills: org pull + propose (``~/.hermes/skills/_org/<org_id>/``).
|
||||
|
||||
Org skills live under a DISTINCT local namespace (read-only to the runtime; a
|
||||
local edit is a personal fork until proposed). The canonical set is
|
||||
``refs/org/<org_id>/HEAD`` -- the SAME object model as personal sync.
|
||||
PERSONAL-ORG GATE: NAS stamps ``org_role`` ONLY for multi-member orgs; no
|
||||
claim => pull/propose raise SyncInertError and personal sync is untouched.
|
||||
``propose_skill`` must stay non-interactive (automation will drive it).
|
||||
Module state (``_skills_dir``, ``_org_dir``, base URL, device id) stays in
|
||||
``tools.skills_sync_client`` and is read lazily so tests can monkeypatch it.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import shutil
|
||||
from pathlib import Path, PurePosixPath
|
||||
from typing import Any, Callable, Dict, List, Optional
|
||||
|
||||
from tools.skills_sync_client_wire import (
|
||||
DEFAULT_MAX_OBJECT_BYTES, ObjectSet, SyncClient, SyncConflict, SyncError, _check_version,
|
||||
assemble_root_from_skill_trees, build_commit, build_tree, materialize_tree, read_ref_hash,
|
||||
root_tree_of_commit, skill_trees_of_root,
|
||||
)
|
||||
|
||||
logger = logging.getLogger("tools.skills_sync_client")
|
||||
|
||||
ORG_DIR_NAME = "_org"
|
||||
|
||||
# Propose re-splices onto a moved org HEAD at most this many times. Small:
|
||||
# contention means other members are actively proposing; unbounded would spin.
|
||||
_ORG_CAS_MAX_ATTEMPTS = 5
|
||||
|
||||
|
||||
def _ssc():
|
||||
from tools import skills_sync_client
|
||||
|
||||
return skills_sync_client
|
||||
|
||||
|
||||
def org_head_ref(org_id: str) -> str:
|
||||
return f"refs/org/{org_id}/HEAD"
|
||||
|
||||
|
||||
def resolve_org_identity() -> Dict[str, Any]:
|
||||
"""``resolve_identity()`` extended with ``org_id`` + ``org_role``.
|
||||
|
||||
Raises SyncInertError when the token carries no ``org_role`` claim (personal
|
||||
org / issuer predates org support): org sync is unavailable, NOT an error.
|
||||
"""
|
||||
from tools.skills_sync_client import SyncInertError, resolve_identity
|
||||
|
||||
identity = resolve_identity()
|
||||
claims = identity.get("claims") or {}
|
||||
org_id = claims.get("org_id")
|
||||
org_role = claims.get("org_role")
|
||||
if not org_id:
|
||||
raise SyncInertError("no organisation associated with this account")
|
||||
if not isinstance(org_role, str) or not org_role:
|
||||
raise SyncInertError("this account isn't a member of a shared organisation")
|
||||
identity["org_id"] = str(org_id)
|
||||
identity["org_role"] = org_role
|
||||
return identity
|
||||
|
||||
|
||||
def _org_client(identity: Optional[Dict[str, Any]], client: Optional[SyncClient]):
|
||||
"""Resolve (identity, client, caps) for an org operation; raises SyncInertError
|
||||
when the base URL is missing or the server lacks the ``org`` feature."""
|
||||
ssc = _ssc()
|
||||
identity = identity or resolve_org_identity()
|
||||
if client is None:
|
||||
base_url = ssc.resolve_sync_base_url()
|
||||
if not base_url:
|
||||
raise ssc.SyncInertError("no sync base URL configured")
|
||||
client = SyncClient(base_url, identity["api_key"])
|
||||
caps = client.capabilities()
|
||||
_check_version(caps)
|
||||
if "org" not in (caps.get("features") or []):
|
||||
raise ssc.SyncInertError("this server does not support org-shared skills")
|
||||
return identity, client, caps
|
||||
|
||||
|
||||
def _read_org_head(client: SyncClient, org_id: str) -> Optional[str]:
|
||||
"""Current org HEAD, or None. MUST read through the ORG endpoint."""
|
||||
return read_ref_hash(client, org_head_ref(org_id), org_scope=True)
|
||||
|
||||
|
||||
# Local mirror sidecars
|
||||
|
||||
def _mirror_root(org_id: str) -> Path:
|
||||
return _ssc()._org_dir() / org_id
|
||||
|
||||
|
||||
def _write_sidecar(what: str, path_fn: Callable[[], Path], text: str) -> None:
|
||||
"""Best-effort sidecar write (path resolution included); never raises."""
|
||||
try:
|
||||
path = path_fn()
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(text, encoding="utf-8")
|
||||
except Exception as e:
|
||||
logger.debug("skills_sync_client: %s write failed: %s", what, e)
|
||||
|
||||
|
||||
def _skill_dir_fingerprint(path: Path) -> str:
|
||||
"""Stable content hash of a materialized skill dir (sorted relative path +
|
||||
bytes, independent of filesystem order and mtimes). "" on read failure."""
|
||||
h = hashlib.sha256()
|
||||
try:
|
||||
for f in sorted(p for p in path.rglob("*") if p.is_file()):
|
||||
h.update(str(f.relative_to(path)).replace("\\", "/").encode("utf-8"))
|
||||
h.update(b"\0")
|
||||
h.update(f.read_bytes())
|
||||
h.update(b"\0")
|
||||
except OSError as e:
|
||||
logger.debug("skills_sync_client: fingerprint failed for %s: %s", path, e)
|
||||
return ""
|
||||
return h.hexdigest()
|
||||
|
||||
|
||||
def _sidecar_path(org_id: Optional[str], const: str) -> Path:
|
||||
"""``<mirror>/<agent.skill_utils.<const>>`` (org-level when org_id is None)."""
|
||||
import agent.skill_utils as sku
|
||||
|
||||
return (_mirror_root(org_id) if org_id else _ssc()._org_dir()) / getattr(sku, const)
|
||||
|
||||
|
||||
def _org_baseline_path(org_id: str) -> Path:
|
||||
"""Sidecar recording the upstream fingerprint of each mirrored skill."""
|
||||
return _sidecar_path(org_id, "ORG_BASELINE_FILE")
|
||||
|
||||
|
||||
def _read_org_baseline(org_id: str) -> Dict[str, Any]:
|
||||
try:
|
||||
return json.loads(_org_baseline_path(org_id).read_text(encoding="utf-8"))
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def _write_org_baseline(org_id: str, baseline: Dict[str, Any]) -> None:
|
||||
_write_sidecar(
|
||||
"baseline", lambda: _org_baseline_path(org_id), json.dumps(baseline, indent=2, sort_keys=True)
|
||||
)
|
||||
|
||||
|
||||
def _write_org_provenance(org_id: str, data: Dict[str, Any]) -> None:
|
||||
_write_sidecar(
|
||||
"org provenance", lambda: _sidecar_path(org_id, "ORG_PROVENANCE_FILE"), json.dumps(data, indent=2)
|
||||
)
|
||||
|
||||
|
||||
def _write_active_org_marker(org_id: str) -> None:
|
||||
"""Record which org's mirror may resolve (agent/skill_utils.read_active_org_id)."""
|
||||
_write_sidecar("active-org marker", lambda: _sidecar_path(None, "ORG_ACTIVE_MARKER"), org_id)
|
||||
|
||||
|
||||
def _clear_active_org_marker() -> None:
|
||||
"""Remove the active-org marker so org skills stop resolving."""
|
||||
try:
|
||||
marker = _sidecar_path(None, "ORG_ACTIVE_MARKER")
|
||||
if marker.exists():
|
||||
marker.unlink()
|
||||
logger.info(
|
||||
"skills_sync_client: cleared active-org marker "
|
||||
"(token has no org workflow); org skills no longer resolve"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug("skills_sync_client: marker clear failed: %s", e)
|
||||
|
||||
|
||||
def org_skill_is_locally_modified(skill_rel_path: str, org_id: str) -> bool:
|
||||
"""True when the local copy of an org skill differs from what upstream sent.
|
||||
No recorded baseline (pre-existing mirror) => unmodified; the next pull
|
||||
records one."""
|
||||
dest = _mirror_root(org_id) / PurePosixPath(skill_rel_path)
|
||||
if not dest.is_dir():
|
||||
return False
|
||||
entry = _read_org_baseline(org_id).get(skill_rel_path) or {}
|
||||
recorded = entry.get("fingerprint") if isinstance(entry, dict) else entry
|
||||
return bool(recorded) and _skill_dir_fingerprint(dest) != recorded
|
||||
|
||||
|
||||
def list_locally_modified_org_skills(org_id: Optional[str] = None) -> List[str]:
|
||||
"""Org skills with local edits that upstream has not seen."""
|
||||
try:
|
||||
from agent.skill_utils import read_active_org_id
|
||||
|
||||
org_id = org_id or read_active_org_id(_ssc()._skills_dir())
|
||||
if not org_id:
|
||||
return []
|
||||
return sorted(rel for rel in _read_org_baseline(org_id) if org_skill_is_locally_modified(rel, org_id))
|
||||
except Exception as e:
|
||||
logger.debug("skills_sync_client: modified-scan failed: %s", e)
|
||||
return []
|
||||
|
||||
|
||||
def list_org_skill_names() -> List[str]:
|
||||
"""Skill names present in the local org mirror (empty when none pulled)."""
|
||||
names: List[str] = []
|
||||
try:
|
||||
from agent.skill_utils import read_active_org_id
|
||||
|
||||
org_id = read_active_org_id(_ssc()._skills_dir())
|
||||
root = _mirror_root(org_id) if org_id else None
|
||||
if root and root.is_dir():
|
||||
for skill_md in root.rglob("SKILL.md"):
|
||||
rel = skill_md.parent.relative_to(root)
|
||||
if rel.parts:
|
||||
names.append(str(rel).replace("\\", "/"))
|
||||
except Exception as e:
|
||||
logger.debug("skills_sync_client: org skill listing failed: %s", e)
|
||||
return sorted(names)
|
||||
|
||||
|
||||
# Pull / propose
|
||||
|
||||
def pull_org_skills(
|
||||
client: Optional[SyncClient] = None, *, identity: Optional[Dict[str, Any]] = None
|
||||
) -> Dict[str, Any]:
|
||||
"""Pull the org canonical set into the local mirror (fast-forward only; no
|
||||
client merge on the org path). A mirrored skill with LOCAL edits is never
|
||||
clobbered: it is skipped, and reported in ``conflicted`` when upstream also
|
||||
moved; the member's change of record is ``propose_skill``.
|
||||
Returns ``{ok, org_id, head, updated, conflicted}``."""
|
||||
identity = identity or resolve_org_identity()
|
||||
if "org_id" not in identity:
|
||||
raise _ssc().SyncInertError("no organisation context available")
|
||||
identity, client, _caps = _org_client(identity, client)
|
||||
org_id = identity["org_id"]
|
||||
|
||||
head = _read_org_head(client, org_id)
|
||||
# Token-gated marker: written HERE because this runs only after the token's
|
||||
# org_id + org_role were verified. A stale mirror from a previous org stops
|
||||
# resolving the moment a pull runs under a different org.
|
||||
_write_active_org_marker(org_id)
|
||||
if not head:
|
||||
return {"ok": True, "org_id": org_id, "head": None, "updated": []}
|
||||
|
||||
head_commit = client.get_commit_json(head, org_scope=True)
|
||||
skill_trees = skill_trees_of_root(client, head_commit["tree"], org_scope=True)
|
||||
|
||||
dest_root = _mirror_root(org_id)
|
||||
updated: List[str] = []
|
||||
conflicted: List[str] = []
|
||||
baseline = _read_org_baseline(org_id)
|
||||
for rel_path, tree_hash in sorted(skill_trees.items()):
|
||||
dest = dest_root / PurePosixPath(rel_path)
|
||||
try:
|
||||
if dest.exists():
|
||||
if org_skill_is_locally_modified(rel_path, org_id):
|
||||
if (baseline.get(rel_path) or {}).get("tree") != tree_hash:
|
||||
conflicted.append(rel_path)
|
||||
continue
|
||||
shutil.rmtree(dest)
|
||||
dest.mkdir(parents=True, exist_ok=True)
|
||||
materialize_tree(client, tree_hash, dest, org_scope=True)
|
||||
baseline[rel_path] = {"fingerprint": _skill_dir_fingerprint(dest), "tree": tree_hash}
|
||||
updated.append(rel_path)
|
||||
except Exception as e:
|
||||
logger.warning("skills_sync_client: org skill materialize failed for %s: %s", rel_path, e)
|
||||
# Provenance for the skill_view header: the HEAD author is token-verified
|
||||
# by the plane at push time, so it is trustworthy to display.
|
||||
author = head_commit.get("author") or {}
|
||||
_write_org_provenance(org_id, {
|
||||
"org_id": org_id, "head": head, "author_user_id": author.get("owner", ""),
|
||||
"author_device": author.get("device", ""), "ts": head_commit.get("ts", ""), "skills": updated,
|
||||
})
|
||||
_write_org_baseline(org_id, baseline)
|
||||
if conflicted:
|
||||
logger.warning(
|
||||
"skills_sync_client: %d org skill(s) have local edits AND upstream "
|
||||
"changes; left untouched: %s", len(conflicted), ", ".join(conflicted),
|
||||
)
|
||||
return {"ok": True, "org_id": org_id, "head": head, "updated": updated, "conflicted": conflicted}
|
||||
|
||||
|
||||
def propose_skill(
|
||||
skill_name: str,
|
||||
client: Optional[SyncClient] = None,
|
||||
*,
|
||||
identity: Optional[Dict[str, Any]] = None,
|
||||
message: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Propose a local (personal) skill's content to the org canonical set.
|
||||
|
||||
Snapshots the skill dir as an org-scoped commit splicing that ONE skill
|
||||
subtree into the current org HEAD (proposals are per-skill deltas, never a
|
||||
wholesale replace), uploads with ``?scope=org``, then CAS-es the org HEAD:
|
||||
ADMIN/OWNER -> server merges -> ``{ok, merged: True}``; MEMBER -> 202
|
||||
proposal -> ``{ok, proposal_pending: True, proposal_id, ref}``, never
|
||||
presented as live.
|
||||
|
||||
If HEAD moves between read and CAS, the skill is re-spliced onto the NEW
|
||||
head (not replayed from the old root, which would drop the other member's
|
||||
skill) up to ``_ORG_CAS_MAX_ATTEMPTS`` times.
|
||||
"""
|
||||
ssc = _ssc()
|
||||
identity, client, caps = _org_client(identity, client)
|
||||
org_id = identity["org_id"]
|
||||
max_bytes = int(caps.get("max_object_bytes") or DEFAULT_MAX_OBJECT_BYTES)
|
||||
|
||||
rel = ssc._skill_rel_path(skill_name)
|
||||
if rel is None:
|
||||
raise SyncError(f"skill '{skill_name}' not found under the skills dir")
|
||||
skill_dir = ssc._skills_dir() / rel
|
||||
if not (skill_dir / "SKILL.md").exists():
|
||||
raise SyncError(f"skill '{skill_name}' has no SKILL.md")
|
||||
|
||||
objects = ObjectSet()
|
||||
skill_tree = build_tree(skill_dir, objects, max_object_bytes=max_bytes)
|
||||
|
||||
for attempt in range(1, _ORG_CAS_MAX_ATTEMPTS + 1):
|
||||
base_head = _read_org_head(client, org_id)
|
||||
skill_map = (
|
||||
skill_trees_of_root(client, root_tree_of_commit(client, base_head, org_scope=True), org_scope=True)
|
||||
if base_head
|
||||
else {}
|
||||
)
|
||||
skill_map[str(rel)] = skill_tree
|
||||
root_hash = assemble_root_from_skill_trees(skill_map, objects)
|
||||
commit_hash = build_commit(
|
||||
root_hash, [base_head] if base_head else [], owner=identity["owner"],
|
||||
device=ssc.stable_device_id(), message=message or f"propose {skill_name}", objects=objects,
|
||||
)
|
||||
client.put_objects(objects.objects, org_scope=True)
|
||||
try:
|
||||
result = client.cas_ref(org_head_ref(org_id), base_head, commit_hash)
|
||||
break
|
||||
except SyncConflict as conflict:
|
||||
if attempt >= _ORG_CAS_MAX_ATTEMPTS:
|
||||
raise SyncError(
|
||||
"the organisation's skills changed while this was being "
|
||||
f"proposed, and {attempt} attempts to catch up all lost "
|
||||
"the race — run the command again",
|
||||
status=409,
|
||||
) from conflict
|
||||
logger.debug(
|
||||
"propose_skill: org HEAD moved (actual=%r), re-splicing (attempt %d)", conflict.actual, attempt
|
||||
)
|
||||
|
||||
if result.get("proposal_pending"):
|
||||
return {
|
||||
"ok": True, "proposal_pending": True, "proposal_id": result.get("proposal_id"),
|
||||
"ref": result.get("ref"), "commit": commit_hash, "org_id": org_id,
|
||||
}
|
||||
return {
|
||||
"ok": True, "merged": True, "head": result.get("hash", commit_hash),
|
||||
"commit": commit_hash, "org_id": org_id,
|
||||
}
|
||||
|
||||
|
||||
def maybe_pull_org_skills() -> Optional[Dict[str, Any]]:
|
||||
"""Best-effort org pull if all gates hold (logged in, org_role claim,
|
||||
feature enabled, base URL). Never raises; None when inert.
|
||||
|
||||
Marker hygiene: when the token VERIFIABLY lacks the org claim (personal org
|
||||
/ left the org) the active-org marker is cleared so mirrored org skills stop
|
||||
resolving. When identity cannot be resolved at all (offline, logged out) the
|
||||
marker is left alone -- offline grace keeps pulled org skills working.
|
||||
"""
|
||||
ssc = _ssc()
|
||||
try:
|
||||
identity = resolve_org_identity()
|
||||
except ssc.SyncInertError:
|
||||
try:
|
||||
if not (ssc.resolve_identity().get("claims") or {}).get("org_role"):
|
||||
_clear_active_org_marker()
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.debug("skills_sync_client: maybe_pull_org_skills inert/failed: %s", e)
|
||||
return None
|
||||
try:
|
||||
if not ssc.sync_feature_enabled() or not ssc.resolve_sync_base_url():
|
||||
return None
|
||||
return pull_org_skills(identity=identity)
|
||||
except Exception as e:
|
||||
logger.debug("skills_sync_client: maybe_pull_org_skills inert/failed: %s", e)
|
||||
return None
|
||||
439
tools/skills_sync_client_wire.py
Normal file
439
tools/skills_sync_client_wire.py
Normal file
@@ -0,0 +1,439 @@
|
||||
"""Skill Sync wire model: content-addressed objects, the HTTP client, tree walks.
|
||||
|
||||
Everything here is independent of local skill state (no ``~/.hermes`` reads);
|
||||
``tools/skills_sync_client.py`` orchestrates push/pull on top of it and
|
||||
re-exports these names.
|
||||
|
||||
Wire contract (version 1). ``hsp_version`` / ``X-HSP-Object-Type`` are deployed
|
||||
protocol identifiers and are NOT renamed with the product name "Skill Sync".
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import stat as _stat
|
||||
from contextlib import suppress
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path, PurePosixPath
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
logger = logging.getLogger("tools.skills_sync_client")
|
||||
|
||||
WIRE_VERSION = "1"
|
||||
DEFAULT_MAX_OBJECT_BYTES = 26214400 # 25 MiB, mirrors capabilities default
|
||||
|
||||
KIND_BLOB = "blob"
|
||||
KIND_TREE = "tree"
|
||||
KIND_COMMIT = "commit"
|
||||
|
||||
MODE_FILE = "file"
|
||||
MODE_EXEC = "exec"
|
||||
MODE_DIR = "dir"
|
||||
|
||||
ARTIFACT_TYPE_SKILL = "skill"
|
||||
_EXEC_BITS = _stat.S_IXUSR | _stat.S_IXGRP | _stat.S_IXOTH
|
||||
|
||||
# `sync-manifest`: per-skill opt-in is CONTENT in the object model, not a
|
||||
# device-local flag -- a root-level blob in the tree at refs/user/<owner>/HEAD
|
||||
# recording {name, enabled}. The plane manifest is authoritative; the local
|
||||
# `.usage.json` `sync` flag is only the editable intent (reconciled FROM it on
|
||||
# pull, TO it on push). Shape MUST match gateway-gateway src/sync/manifest.ts.
|
||||
SYNC_MANIFEST_ENTRY_NAME = "sync-manifest"
|
||||
SYNC_MANIFEST_TYPE = "sync-manifest"
|
||||
SYNC_MANIFEST_VERSION = 1
|
||||
|
||||
|
||||
# Content addressing. The wire uses the FULL 64-hex sha256 -- a different
|
||||
# namespace from the truncated 16-hex local `content_hash` (skills_guard.py).
|
||||
|
||||
def wire_address(data: bytes) -> str:
|
||||
"""Return ``sha256:<64-hex>`` -- the wire address of ``data``."""
|
||||
return "sha256:" + hashlib.sha256(data).hexdigest()
|
||||
|
||||
|
||||
def canonical_json_bytes(obj: Dict[str, Any]) -> bytes:
|
||||
"""Canonical JSON for tree/commit hashing: UTF-8, sorted keys, no whitespace,
|
||||
no trailing newline. Arrays must already be in contract order (tree entries
|
||||
by name, commit parents by significance). Client and server MUST produce
|
||||
byte-identical output or a push fails ``422 hash_mismatch``."""
|
||||
return json.dumps(obj, sort_keys=True, separators=(",", ":"), ensure_ascii=False).encode("utf-8")
|
||||
|
||||
|
||||
def build_sync_manifest_bytes(skills: Dict[str, bool]) -> bytes:
|
||||
"""Serialize ``{name: enabled}`` into canonical ``sync-manifest`` bytes
|
||||
(entries sorted by name for a stable content address)."""
|
||||
return canonical_json_bytes({
|
||||
"type": SYNC_MANIFEST_TYPE, "version": SYNC_MANIFEST_VERSION,
|
||||
"skills": [{"name": name, "enabled": bool(enabled)} for name, enabled in sorted(skills.items())],
|
||||
})
|
||||
|
||||
|
||||
def parse_sync_manifest(data: bytes) -> Optional[Dict[str, bool]]:
|
||||
"""Parse ``sync-manifest`` bytes into ``{name: enabled}``, or None if malformed.
|
||||
|
||||
Strict (mirrors gateway-gateway ``parseSyncManifest``): unknown type, version
|
||||
!= 1, non-list skills, or a malformed entry all reject -- a malformed
|
||||
manifest must not be mistaken for "no skills opted in".
|
||||
"""
|
||||
try:
|
||||
value = json.loads(data.decode("utf-8"))
|
||||
except Exception:
|
||||
return None
|
||||
if (
|
||||
not isinstance(value, dict)
|
||||
or value.get("type") != SYNC_MANIFEST_TYPE
|
||||
or value.get("version") != SYNC_MANIFEST_VERSION
|
||||
or not isinstance(value.get("skills"), list)
|
||||
):
|
||||
return None
|
||||
out: Dict[str, bool] = {}
|
||||
for raw in value["skills"]:
|
||||
if not isinstance(raw, dict):
|
||||
return None
|
||||
name, enabled = raw.get("name"), raw.get("enabled")
|
||||
if not isinstance(name, str) or not name or not isinstance(enabled, bool):
|
||||
return None
|
||||
out[name] = enabled
|
||||
return out
|
||||
|
||||
|
||||
# Object building
|
||||
|
||||
class ObjectSet:
|
||||
"""Objects to push, ``hash -> (kind, bytes)``, deduped by content address."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.objects: Dict[str, Tuple[str, bytes]] = {}
|
||||
|
||||
def add(self, kind: str, data: bytes) -> str:
|
||||
addr = wire_address(data)
|
||||
self.objects.setdefault(addr, (kind, data))
|
||||
return addr
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.objects)
|
||||
|
||||
|
||||
def _entry(name: str, kind: str, hash_: str, mode: str) -> Dict[str, str]:
|
||||
return {"name": name, "kind": kind, "hash": hash_, "mode": mode}
|
||||
|
||||
|
||||
def _add_tree(entries: List[Dict[str, str]], objects: ObjectSet) -> str:
|
||||
"""Canonicalize *entries* (sorted by name, byte order) into a tree object."""
|
||||
entries.sort(key=lambda e: e["name"])
|
||||
return objects.add(KIND_TREE, canonical_json_bytes({"type": KIND_TREE, "entries": entries}))
|
||||
|
||||
|
||||
def _file_mode(path: Path) -> str:
|
||||
"""``exec`` if +x else ``file``. No symlink / other modes are emitted."""
|
||||
with suppress(OSError):
|
||||
if path.stat().st_mode & _EXEC_BITS:
|
||||
return MODE_EXEC
|
||||
return MODE_FILE
|
||||
|
||||
|
||||
def build_tree(dir_path: Path, objects: ObjectSet, *, max_object_bytes: int) -> str:
|
||||
"""Recursively build objects for *dir_path*; return the tree address. Files
|
||||
-> blobs, subdirs -> nested trees; symlinks/special files are skipped
|
||||
(contract: no symlinks). A blob over *max_object_bytes* raises ValueError
|
||||
so the caller can surface / skip the artifact (server -> 413)."""
|
||||
entries: List[Dict[str, str]] = []
|
||||
for child in sorted(dir_path.iterdir(), key=lambda p: p.name):
|
||||
if child.is_symlink():
|
||||
logger.debug("skills_sync_client: skipping symlink %s", child)
|
||||
elif child.is_dir():
|
||||
sub_hash = build_tree(child, objects, max_object_bytes=max_object_bytes)
|
||||
entries.append(_entry(child.name, KIND_TREE, sub_hash, MODE_DIR))
|
||||
elif child.is_file():
|
||||
data = child.read_bytes()
|
||||
if len(data) > max_object_bytes:
|
||||
raise ValueError(
|
||||
f"file {child} is {len(data)} bytes > max_object_bytes {max_object_bytes}"
|
||||
)
|
||||
entries.append(_entry(child.name, KIND_BLOB, objects.add(KIND_BLOB, data), _file_mode(child)))
|
||||
return _add_tree(entries, objects)
|
||||
|
||||
|
||||
def build_commit(
|
||||
tree_hash: str, parents: List[str], *, owner: str, device: str, message: str,
|
||||
objects: ObjectSet, ts: Optional[str] = None,
|
||||
) -> str:
|
||||
"""Build a commit object and return its address.
|
||||
|
||||
``parents``: 0 for the first commit, 1 for an edit, 2 for a merge (order
|
||||
significant: parents[0] = base fast-forwarded from, parents[1] = other head).
|
||||
"""
|
||||
return objects.add(KIND_COMMIT, canonical_json_bytes({
|
||||
"type": KIND_COMMIT, "tree": tree_hash, "parents": list(parents),
|
||||
"author": {"owner": owner, "device": device},
|
||||
"ts": ts or datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ"),
|
||||
"message": message, "artifact_type": ARTIFACT_TYPE_SKILL,
|
||||
}))
|
||||
|
||||
|
||||
def build_root_tree(node: Dict[str, Any], objects: ObjectSet, *, manifest_hash: Optional[str] = None) -> str:
|
||||
"""Canonicalize a nested ``{name: {"__tree__": hash} | subdict}`` root into trees.
|
||||
|
||||
``manifest_hash`` (top level only) adds the root ``sync-manifest`` BLOB entry.
|
||||
It cannot collide with a skill dir: skill entries are trees, this is a blob.
|
||||
"""
|
||||
entries: List[Dict[str, str]] = []
|
||||
for name, child in node.items():
|
||||
if isinstance(child, dict) and "__tree__" in child and len(child) == 1:
|
||||
entries.append(_entry(name, KIND_TREE, child["__tree__"], MODE_DIR))
|
||||
else:
|
||||
entries.append(_entry(name, KIND_TREE, build_root_tree(child, objects), MODE_DIR))
|
||||
if manifest_hash is not None:
|
||||
entries.append(_entry(SYNC_MANIFEST_ENTRY_NAME, KIND_BLOB, manifest_hash, MODE_FILE))
|
||||
return _add_tree(entries, objects)
|
||||
|
||||
|
||||
def nest_skill_tree(root: Dict[str, Any], rel_parts: Tuple[str, ...], tree_hash: str) -> None:
|
||||
"""Insert a skill tree leaf into the nested root structure by path parts."""
|
||||
node = root
|
||||
for part in rel_parts[:-1]:
|
||||
node = node.setdefault(part, {})
|
||||
node[rel_parts[-1]] = {"__tree__": tree_hash}
|
||||
|
||||
|
||||
def assemble_root_from_skill_trees(skill_trees: Dict[str, str], objects: ObjectSet) -> str:
|
||||
"""Build a profile-root tree from ``{posix_rel_path: tree_hash}``.
|
||||
|
||||
The skill trees are assumed already durable (they came from either side of
|
||||
a merge / the org HEAD); only the new intermediate/root trees are added.
|
||||
"""
|
||||
root: Dict[str, Any] = {}
|
||||
for path, tree_hash in skill_trees.items():
|
||||
nest_skill_tree(root, PurePosixPath(path).parts, tree_hash)
|
||||
return build_root_tree(root, objects)
|
||||
|
||||
|
||||
# HTTP client (routes under /v1/sync/)
|
||||
|
||||
class SyncError(RuntimeError):
|
||||
"""A non-recoverable wire error (4xx the client can't retry)."""
|
||||
|
||||
def __init__(self, message: str, *, status: Optional[int] = None):
|
||||
super().__init__(message)
|
||||
self.status = status
|
||||
|
||||
|
||||
class SyncConflict(RuntimeError):
|
||||
"""CAS lost (409). NOT a rejection -- pushed objects are already durable.
|
||||
|
||||
``actual`` is the current head to merge against, or None when the ref does
|
||||
not exist server-side (the server sends ""). None means "retry as a create";
|
||||
it must never be fetched as an object -- normalized here, not per call site.
|
||||
"""
|
||||
|
||||
def __init__(self, actual: Optional[str]):
|
||||
self.actual: Optional[str] = actual or None
|
||||
super().__init__(
|
||||
f"CAS conflict; actual head {self.actual}"
|
||||
if self.actual
|
||||
else "CAS conflict; the ref does not exist yet"
|
||||
)
|
||||
|
||||
|
||||
def _check_version(caps: Dict[str, Any]) -> None:
|
||||
"""Reject an incompatible server major version."""
|
||||
ver = str(caps.get("hsp_version") or "") # wire field name
|
||||
if ver.split(".", 1)[0] != WIRE_VERSION:
|
||||
raise SyncError(
|
||||
f"this server speaks sync version {ver!r}, but this Hermes speaks "
|
||||
f"{WIRE_VERSION} — update Hermes to sync with it"
|
||||
)
|
||||
|
||||
|
||||
class SyncClient:
|
||||
"""Sync client bound to a base URL + Nous bearer.
|
||||
|
||||
Org refs/objects live behind SEPARATE ``org/`` routes, not a prefix filter:
|
||||
the personal routes are hard-scoped to the token's owner and would silently
|
||||
answer an org query with personal data. Callers reading org content MUST
|
||||
pass ``org_scope=True`` on every hop.
|
||||
"""
|
||||
|
||||
def __init__(self, base_url: str, api_key: str, *, timeout: float = 30.0):
|
||||
self.base = base_url.rstrip("/")
|
||||
self.api_key = api_key
|
||||
self.timeout = timeout
|
||||
import requests # core dependency
|
||||
|
||||
self._session = requests.Session()
|
||||
self._session.headers["Authorization"] = f"Bearer {api_key}"
|
||||
|
||||
def _url(self, path: str) -> str:
|
||||
return f"{self.base}/v1/sync/{path.lstrip('/')}"
|
||||
|
||||
@staticmethod
|
||||
def _check(r, op: str, ok=(200,), errors: Optional[Dict[int, str]] = None) -> None:
|
||||
"""Raise SyncError unless the status is in *ok*; *errors* maps specific
|
||||
statuses to their message, anything else gets ``"<op> failed: <code>"``."""
|
||||
if r.status_code in ok:
|
||||
return
|
||||
msg = (errors or {}).get(r.status_code)
|
||||
raise SyncError(msg or f"{op} failed: {r.status_code}", status=r.status_code)
|
||||
|
||||
def capabilities(self) -> Dict[str, Any]:
|
||||
"""GET capabilities (no auth required)."""
|
||||
r = self._session.get(self._url("capabilities"), timeout=self.timeout)
|
||||
self._check(r, "capabilities")
|
||||
return r.json()
|
||||
|
||||
def get_refs(self, prefix: str, *, org_scope: bool = False) -> List[Dict[str, str]]:
|
||||
"""GET refs?prefix=... (or org/refs, filtered client-side by *prefix*)."""
|
||||
path = "org/refs" if org_scope else "refs"
|
||||
params = None if org_scope else {"prefix": prefix}
|
||||
r = self._session.get(self._url(path), params=params, timeout=self.timeout)
|
||||
self._check(r, "get_refs")
|
||||
refs = (r.json() or {}).get("refs", [])
|
||||
if org_scope:
|
||||
refs = [r_ for r_ in refs if str(r_.get("name", "")).startswith(prefix)]
|
||||
return refs
|
||||
|
||||
def get_object(self, obj_hash: str, *, org_scope: bool = False) -> Tuple[str, bytes]:
|
||||
"""GET objects/:hash -> ``(kind, bytes)``. Kind comes from the object-type
|
||||
header; a blob (octet-stream) is returned as ``blob``."""
|
||||
path = f"org/objects/{obj_hash}" if org_scope else f"objects/{obj_hash}"
|
||||
r = self._session.get(self._url(path), timeout=self.timeout)
|
||||
self._check(
|
||||
r, "get_object",
|
||||
errors={404: f"object {obj_hash} not found", 403: f"object {obj_hash} not readable"},
|
||||
)
|
||||
return r.headers.get("X-HSP-Object-Type") or KIND_BLOB, r.content
|
||||
|
||||
def _get_json_of_kind(self, obj_hash: str, expected: str, org_scope: bool) -> Dict[str, Any]:
|
||||
kind, data = self.get_object(obj_hash, org_scope=org_scope)
|
||||
if kind != expected:
|
||||
raise SyncError(f"{obj_hash} is {kind}, expected {expected}")
|
||||
return json.loads(data.decode("utf-8"))
|
||||
|
||||
def get_commit_json(self, commit_hash: str, *, org_scope: bool = False) -> Dict[str, Any]:
|
||||
return self._get_json_of_kind(commit_hash, KIND_COMMIT, org_scope)
|
||||
|
||||
def get_tree_json(self, tree_hash: str, *, org_scope: bool = False) -> Dict[str, Any]:
|
||||
return self._get_json_of_kind(tree_hash, KIND_TREE, org_scope)
|
||||
|
||||
def put_objects(self, objects: Dict[str, Tuple[str, bytes]], *, org_scope: bool = False) -> Dict[str, Any]:
|
||||
"""POST objects -- batch upload as multipart/form-data: field name = the
|
||||
claimed ``sha256:<hex>``, ``filename`` = object type, body = raw bytes
|
||||
(the contract requires raw bytes, not base64-in-JSON). The server
|
||||
recomputes every hash and rejects the whole batch with 422 on mismatch;
|
||||
known hashes are idempotent no-ops. ``org_scope`` adds ``?scope=org`` so
|
||||
objects land org-readable (required before an org CAS/propose)."""
|
||||
files = [(h, (kind, data, "application/octet-stream")) for h, (kind, data) in objects.items()]
|
||||
r = self._session.post(
|
||||
self._url("objects"), files=files, params={"scope": "org"} if org_scope else None, timeout=self.timeout
|
||||
)
|
||||
self._check(
|
||||
r, "put_objects", ok=(200, 201),
|
||||
errors={413: "object too large (413)", 422: f"hash_mismatch (422): {r.text}"},
|
||||
)
|
||||
return r.json() if r.content else {}
|
||||
|
||||
def cas_ref(self, name: str, from_hash: Optional[str], to_hash: str) -> Dict[str, Any]:
|
||||
"""POST refs/:name -- atomic compare-and-swap. Raises SyncConflict on 409.
|
||||
|
||||
A non-admin member's CAS on an org HEAD is converted server-side to a
|
||||
proposal (202) and surfaced as ``{"proposal_pending": True, ...}``: a
|
||||
SUCCESS-shaped outcome that must never be presented as live/merged.
|
||||
"""
|
||||
r = self._session.post(
|
||||
self._url(f"refs/{name}"), json={"from": from_hash, "to": to_hash}, timeout=self.timeout
|
||||
)
|
||||
if r.status_code == 202:
|
||||
return {"proposal_pending": True, **(r.json() if r.content else {})}
|
||||
if r.status_code == 409:
|
||||
# "" actual = the ref does not exist server-side (SyncConflict -> None).
|
||||
raise SyncConflict((r.json() or {}).get("actual", ""))
|
||||
self._check(r, "cas_ref", errors={403: "forbidden (403) -- owner/permission"})
|
||||
return r.json() if r.content else {}
|
||||
|
||||
|
||||
# Reading remote trees
|
||||
|
||||
def read_ref_hash(client: SyncClient, ref: str, *, org_scope: bool = False) -> Optional[str]:
|
||||
"""Hash of *ref* (queried with itself as prefix), or None if absent."""
|
||||
refs = client.get_refs(ref, org_scope=org_scope)
|
||||
return next((r.get("hash") for r in refs if r.get("name") == ref), None)
|
||||
|
||||
|
||||
def root_tree_of_commit(client: SyncClient, commit_hash: str, *, org_scope: bool = False) -> str:
|
||||
return client.get_commit_json(commit_hash, org_scope=org_scope)["tree"]
|
||||
|
||||
|
||||
def skill_trees_of_root(client: SyncClient, root_tree_hash: str, *, org_scope: bool = False) -> Dict[str, str]:
|
||||
"""Flatten a profile-root tree into ``{posix_rel_path: skill_tree_hash}``. A
|
||||
skill tree is any subtree containing a ``SKILL.md`` blob, keyed by its path
|
||||
so category nesting is preserved."""
|
||||
result: Dict[str, str] = {}
|
||||
|
||||
def _walk(tree_hash: str, prefix: str) -> None:
|
||||
entries = client.get_tree_json(tree_hash, org_scope=org_scope).get("entries", [])
|
||||
if prefix and any(e.get("name") == "SKILL.md" and e.get("kind") == KIND_BLOB for e in entries):
|
||||
result[prefix] = tree_hash
|
||||
return
|
||||
for e in entries:
|
||||
if e.get("kind") == KIND_TREE:
|
||||
_walk(e["hash"], f"{prefix}/{e['name']}" if prefix else e["name"])
|
||||
|
||||
_walk(root_tree_hash, "")
|
||||
return result
|
||||
|
||||
|
||||
def read_manifest_of_root(client: SyncClient, root_tree_hash: str) -> Optional[Dict[str, bool]]:
|
||||
"""``{name: enabled}`` from the root-level ``sync-manifest`` blob, or None if
|
||||
absent/malformed. This is how a device learns another device's opt-ins."""
|
||||
try:
|
||||
tree = client.get_tree_json(root_tree_hash)
|
||||
except Exception as e:
|
||||
logger.debug("skills_sync_client: manifest root read failed: %s", e)
|
||||
return None
|
||||
for e in tree.get("entries", []):
|
||||
if e.get("name") == SYNC_MANIFEST_ENTRY_NAME and e.get("kind") == KIND_BLOB:
|
||||
try:
|
||||
_kind, data = client.get_object(e["hash"])
|
||||
except Exception as ex:
|
||||
logger.debug("skills_sync_client: manifest blob fetch failed: %s", ex)
|
||||
return None
|
||||
return parse_sync_manifest(data)
|
||||
return None
|
||||
|
||||
|
||||
def materialize_tree(client: SyncClient, tree_hash: str, dest: Path, *, org_scope: bool = False) -> None:
|
||||
"""Write the tree at *tree_hash* into *dest* (created if needed): blobs ->
|
||||
files (+x restored for ``exec``), trees -> subdirectories. Does NOT delete
|
||||
files absent from the tree (caller decides). Refuses path traversal."""
|
||||
dest.mkdir(parents=True, exist_ok=True)
|
||||
for entry in client.get_tree_json(tree_hash, org_scope=org_scope).get("entries", []):
|
||||
name = entry.get("name", "")
|
||||
if not name or "/" in name or name in (".", ".."):
|
||||
logger.warning("skills_sync_client: skipping unsafe tree entry %r", name)
|
||||
continue
|
||||
target = dest / name
|
||||
kind = entry.get("kind")
|
||||
if kind == KIND_TREE:
|
||||
materialize_tree(client, entry["hash"], target, org_scope=org_scope)
|
||||
elif kind == KIND_BLOB:
|
||||
_, data = client.get_object(entry["hash"], org_scope=org_scope)
|
||||
target.write_bytes(data)
|
||||
if entry.get("mode") == MODE_EXEC:
|
||||
with suppress(OSError):
|
||||
target.chmod(target.stat().st_mode | _EXEC_BITS)
|
||||
|
||||
|
||||
def merge_skill(base: Optional[str], ours: Optional[str], theirs: Optional[str]) -> str:
|
||||
"""Three-way decision for one skill's tree hash: ``ours`` / ``theirs`` /
|
||||
``either`` / ``overlap`` / ``none``. A side "modified" the skill when its
|
||||
hash differs from the common base (same semantics as skills_sync.py's
|
||||
origin/user/incoming block)."""
|
||||
if ours == theirs:
|
||||
return "either" if ours is not None else "none"
|
||||
if theirs == base: # only we moved
|
||||
return "ours"
|
||||
if ours == base: # only they moved
|
||||
return "theirs"
|
||||
return "overlap"
|
||||
262
tools/skills_sync_optional.py
Normal file
262
tools/skills_sync_optional.py
Normal file
@@ -0,0 +1,262 @@
|
||||
"""Official optional-skill provenance: hub-lock backfill and restore.
|
||||
|
||||
Extracted from ``tools.skills_sync``. Profile-scoped paths and patchable
|
||||
helpers are resolved through ``_ss()`` at call time so tests and multi-profile
|
||||
runtimes that patch ``tools.skills_sync`` globals keep working.
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path, PurePosixPath
|
||||
from typing import Dict, Iterator, List, Optional, Set, Tuple
|
||||
|
||||
from agent.skill_utils import is_excluded_skill_path
|
||||
from utils import atomic_write_text
|
||||
|
||||
logger = logging.getLogger("tools.skills_sync")
|
||||
|
||||
|
||||
def _ss():
|
||||
from tools import skills_sync
|
||||
|
||||
return skills_sync
|
||||
|
||||
|
||||
def _content_hash(directory: Path) -> str:
|
||||
"""Same hash style the skills hub lock uses; hashing is provenance metadata
|
||||
only, so fall back to the local MD5 if guard deps are unavailable."""
|
||||
try:
|
||||
from tools.skills_guard import content_hash
|
||||
|
||||
return content_hash(directory)
|
||||
except Exception:
|
||||
return _ss()._dir_hash(directory)
|
||||
|
||||
|
||||
def _safe_rel_install_path(path: Path, base: Path) -> str:
|
||||
"""Return a normalized relative POSIX path, rejecting traversal/absolute paths."""
|
||||
posix = path.relative_to(base).as_posix()
|
||||
pure = PurePosixPath(posix)
|
||||
parts = [part for part in pure.parts if part not in {"", "."}]
|
||||
if pure.is_absolute() or not parts or ".." in parts:
|
||||
raise ValueError(f"Unsafe optional skill path: {posix}")
|
||||
return "/".join(parts)
|
||||
|
||||
|
||||
def _skill_file_list(skill_dir: Path) -> List[str]:
|
||||
"""List files inside a skill directory in lock-file format."""
|
||||
return [f.relative_to(skill_dir).as_posix() for f in sorted(skill_dir.rglob("*")) if f.is_file()]
|
||||
|
||||
|
||||
def _hub_lock_path() -> Path:
|
||||
return _ss()._skills_dir() / ".hub" / "lock.json"
|
||||
|
||||
|
||||
def _load_hub_lock() -> Optional[dict]:
|
||||
"""Parse the skills-hub lock; None when missing or unreadable."""
|
||||
try:
|
||||
return json.loads(_hub_lock_path().read_text(encoding="utf-8"))
|
||||
except (FileNotFoundError, json.JSONDecodeError, OSError):
|
||||
return None
|
||||
|
||||
|
||||
def _hub_lock_entries(data: Optional[dict]) -> List[dict]:
|
||||
return [e for e in ((data or {}).get("installed") or {}).values() if isinstance(e, dict)]
|
||||
|
||||
|
||||
def _read_hub_install_paths() -> Set[str]:
|
||||
"""Install paths recorded in the hub lock, as POSIX strings. Hub-installed skills
|
||||
are owned by the hub, never by bundled sync: rename recovery must not move them even
|
||||
when content matches a bundled origin hash, or the lock's ``install_path`` dangles."""
|
||||
return {str(e["install_path"]).strip("/") for e in _hub_lock_entries(_load_hub_lock()) if e.get("install_path")}
|
||||
|
||||
|
||||
def _write_hub_lock(lock_path: Path, data: dict) -> None:
|
||||
"""Atomic write so a crash mid-write can't wipe all provenance (the
|
||||
JSONDecodeError fallback in the reader resets ``installed`` to empty)."""
|
||||
atomic_write_text(lock_path, json.dumps(data, indent=2, ensure_ascii=False) + "\n", tmp_prefix=".lock_")
|
||||
|
||||
|
||||
def _iter_optional_skills(optional_dir: Path, *, root_relative: bool) -> Iterator[Tuple[Path, Path, str]]:
|
||||
"""Yield ``(skill_md, src, install_path)`` for every safe official optional skill."""
|
||||
for skill_md in sorted(optional_dir.rglob("SKILL.md")):
|
||||
if root_relative and is_excluded_skill_path(skill_md.relative_to(optional_dir), root=optional_dir):
|
||||
continue
|
||||
if not root_relative and is_excluded_skill_path(skill_md):
|
||||
continue
|
||||
try:
|
||||
yield skill_md, skill_md.parent, _safe_rel_install_path(skill_md.parent, optional_dir)
|
||||
except ValueError as e:
|
||||
logger.debug("Skipping optional skill with unsafe path %s: %s", skill_md.parent, e)
|
||||
|
||||
|
||||
def _optional_skill_index() -> Dict[str, Tuple[str, str, Path]]:
|
||||
"""Official optional skills keyed by BOTH folder name and frontmatter name, so callers
|
||||
may pass either the hub-lock slug or the user-facing name. Values are
|
||||
``(folder_name, install_path, source_dir)``."""
|
||||
ss = _ss()
|
||||
optional_dir = ss._get_optional_dir()
|
||||
index: Dict[str, Tuple[str, str, Path]] = {}
|
||||
if not optional_dir.exists():
|
||||
return index
|
||||
for skill_md, src, install_path in _iter_optional_skills(optional_dir, root_relative=True):
|
||||
value = (src.name, install_path, src)
|
||||
index[src.name] = value
|
||||
index[ss._read_skill_name(skill_md, src.name)] = value
|
||||
return index
|
||||
|
||||
|
||||
def _move_to_restore_backup(path: Path, backup_root: Path) -> str:
|
||||
"""Move an existing skill directory into a restore backup, preserving rel path."""
|
||||
rel = path.relative_to(_ss()._skills_dir())
|
||||
target = backup_root / rel
|
||||
suffix = 0
|
||||
while target.exists():
|
||||
suffix += 1
|
||||
target = (backup_root / rel).with_name(f"{rel.name}-{suffix}")
|
||||
_ss()._move_dir(path, target)
|
||||
return rel.as_posix()
|
||||
|
||||
|
||||
def _find_active_copies(folder_name: str, src_frontmatter: str, dest: Path) -> List[Path]:
|
||||
"""Active copies of an official skill (by frontmatter name or folder slug),
|
||||
even when the curator moved it into another category; excludes ``dest``."""
|
||||
ss = _ss()
|
||||
names = {folder_name, src_frontmatter}
|
||||
return [
|
||||
md.parent
|
||||
for md in ss._iter_active_skill_mds(sort=True)
|
||||
if md.parent != dest and (md.parent.name == folder_name or ss._read_skill_name(md, md.parent.name) in names)
|
||||
]
|
||||
|
||||
|
||||
def restore_official_optional_skill(name: str, *, restore: bool = False) -> dict:
|
||||
"""Restore one or all official optional skills from repo source. ``restore=False``
|
||||
only performs exact-match provenance backfill; ``restore=True`` repairs mutated /
|
||||
reorganized skills by backing up matching active copies and copying the official
|
||||
source into its canonical path."""
|
||||
ss = _ss()
|
||||
|
||||
def _fail(message: str) -> dict:
|
||||
return {"ok": False, "message": message, "restored": [], "backfilled": [], "backed_up": []}
|
||||
|
||||
index = _optional_skill_index()
|
||||
if not index:
|
||||
return _fail("No official optional skills directory found.")
|
||||
if name in {"all", "*"}:
|
||||
targets = sorted(set(index.values()), key=lambda item: item[1])
|
||||
elif name in index:
|
||||
targets = [index[name]]
|
||||
else:
|
||||
return _fail(f"Official optional skill not found: {name}")
|
||||
|
||||
restored: List[str] = []
|
||||
backed_up: List[str] = []
|
||||
timestamp = datetime.now(timezone.utc).strftime("%Y%m%d-%H%M%S")
|
||||
backup_root = ss._skills_dir() / ".restore-backups" / f"official-optional-{timestamp}"
|
||||
|
||||
for folder_name, install_path, src in targets if restore else []:
|
||||
dest = ss._skills_dir() / Path(*install_path.split("/"))
|
||||
canonical_ok = dest.exists() and ss._dir_hash(dest) == ss._dir_hash(src)
|
||||
src_frontmatter = ss._read_skill_name(src / "SKILL.md", folder_name)
|
||||
for match in _find_active_copies(folder_name, src_frontmatter, dest):
|
||||
if match.exists():
|
||||
backed_up.append(_move_to_restore_backup(match, backup_root))
|
||||
if dest.exists() and not canonical_ok:
|
||||
backed_up.append(_move_to_restore_backup(dest, backup_root))
|
||||
if not dest.exists():
|
||||
ss._copy_dir(src, dest)
|
||||
restored.append(folder_name)
|
||||
|
||||
return {
|
||||
"ok": True, "message": "Official optional skill repair complete.", "restored": restored,
|
||||
"backfilled": _backfill_optional_provenance(quiet=True), "backed_up": backed_up,
|
||||
"backup_dir": str(backup_root) if backed_up else "",
|
||||
}
|
||||
|
||||
|
||||
def _index_installed_skill_dirs_by_name() -> Dict[str, List[Path]]:
|
||||
"""Index installed skills by directory name with one active-tree scan,
|
||||
skipping anything that resolves outside the skills tree (symlinks/external)."""
|
||||
ss = _ss()
|
||||
index: Dict[str, List[Path]] = {}
|
||||
root = ss._skills_dir().resolve()
|
||||
for skill_md in ss._iter_active_skill_mds():
|
||||
try:
|
||||
skill_md.parent.resolve().relative_to(root)
|
||||
except (OSError, ValueError):
|
||||
continue
|
||||
index.setdefault(skill_md.parent.name, []).append(skill_md.parent)
|
||||
return index
|
||||
|
||||
|
||||
def _relocated_dest(src_name: str, index: Dict[str, List[Path]]) -> Optional[Tuple[Path, str]]:
|
||||
"""The active tree may hold a skill under a DIFFERENT category path than the
|
||||
repo (upstream reorganizes; the installed copy keeps its old location). Fall
|
||||
back to a UNIQUE same-directory-name match — an ambiguous name gives no basis
|
||||
to pick one. Returns ``(dest, install_path)`` or None."""
|
||||
candidates = index.get(src_name, [])
|
||||
if len(candidates) != 1:
|
||||
return None
|
||||
dest = candidates[0]
|
||||
try:
|
||||
return dest, _safe_rel_install_path(dest, _ss()._skills_dir())
|
||||
except ValueError as e:
|
||||
logger.debug("Skipping relocated optional skill %s: %s", dest, e)
|
||||
return None
|
||||
|
||||
|
||||
def _backfill_optional_provenance(quiet: bool = False) -> List[str]:
|
||||
"""Mark already-present official optional skills as hub-installed: skills that used
|
||||
to be bundled (or were hand-copied) and now live under optional-skills/ get official
|
||||
provenance when byte-identical to the source. Modified/local skills are left alone."""
|
||||
ss = _ss()
|
||||
optional_dir = ss._get_optional_dir()
|
||||
if not optional_dir.exists():
|
||||
return []
|
||||
|
||||
data = _load_hub_lock()
|
||||
if data is None:
|
||||
data = {"version": 1, "installed": {}}
|
||||
installed = data.setdefault("installed", {})
|
||||
existing_paths = {entry.get("install_path") for entry in _hub_lock_entries(data)}
|
||||
|
||||
backfilled: List[str] = []
|
||||
installed_dir_index: Optional[Dict[str, List[Path]]] = None
|
||||
for _skill_md, src, install_path in _iter_optional_skills(optional_dir, root_relative=False):
|
||||
lock_name = src.name
|
||||
if lock_name in installed or install_path in existing_paths:
|
||||
continue
|
||||
dest = ss._skills_dir() / Path(*install_path.split("/"))
|
||||
if not dest.is_dir():
|
||||
if installed_dir_index is None:
|
||||
installed_dir_index = _index_installed_skill_dirs_by_name()
|
||||
found = _relocated_dest(src.name, installed_dir_index)
|
||||
if found is None:
|
||||
continue
|
||||
dest, install_path = found # still requires a byte-identical hash below
|
||||
if install_path in existing_paths or ss._dir_hash(dest) != ss._dir_hash(src):
|
||||
continue
|
||||
|
||||
timestamp = datetime.now(timezone.utc).isoformat()
|
||||
installed[lock_name] = {
|
||||
"source": "official",
|
||||
"identifier": f"official/{install_path}",
|
||||
"trust_level": "builtin",
|
||||
"scan_verdict": "backfilled",
|
||||
"content_hash": _content_hash(dest),
|
||||
"install_path": install_path,
|
||||
"files": _skill_file_list(dest),
|
||||
"metadata": {"backfilled_from": "optional-skills"},
|
||||
"installed_at": timestamp,
|
||||
"updated_at": timestamp,
|
||||
}
|
||||
existing_paths.add(install_path)
|
||||
backfilled.append(lock_name)
|
||||
if not quiet:
|
||||
print(f" = {lock_name} (official optional provenance backfilled)")
|
||||
|
||||
if backfilled:
|
||||
_write_hub_lock(_hub_lock_path(), data)
|
||||
return backfilled
|
||||
Reference in New Issue
Block a user