Merge simp/r3-33 (late tail) into hermes/simplify-codebase

This commit is contained in:
Teknium
2026-09-03 05:09:00 -07:00
27 changed files with 2046 additions and 3225 deletions
+54 -83
View File
@@ -1,12 +1,9 @@
"""Central manager for per-server MCP OAuth state (one instance per process).
Holds per-server providers and coordinates cross-process token reload (mtime-based disk watch,
so tokens refreshed by cron/another CLI are picked up without a restart), 401 deduplication
(N concurrent 401s with the same access_token trigger one recovery) and reconnect signalling
(``MCPServerTask`` drives the reconnect; the manager decides when). The ONLY place that
instantiates the SDK's ``OAuthClientProvider`` for runtime use. We rely on the SDK's lazy
refresh: one ``stat()`` per tool call is cheaper than an await + refresh round-trip.
"""
"""Central manager for per-server MCP OAuth state (one instance per process): per-server providers, cross-process
token reload (mtime-based disk watch so tokens refreshed by cron/another CLI are picked up without a restart), 401
deduplication (N concurrent 401s with the same access_token trigger one recovery) and reconnect signalling
(``MCPServerTask`` drives the reconnect; the manager decides when). The ONLY place that instantiates the SDK's
``OAuthClientProvider`` for runtime use; refresh stays lazy in the SDK — one ``stat()`` per tool call is cheaper
than an await + refresh round-trip."""
from __future__ import annotations
@@ -53,34 +50,30 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
def __init__(self, *args: Any, server_name: str = "", preregistered: bool = False, **kwargs: Any):
super().__init__(*args, **kwargs)
# mcp 2.0 uses a task-owned anyio.Lock held across the yielded resource request (a
# session-long GET blocks every POST; HTTPX may close the generator from another task).
# A binary semaphore keeps mutual exclusion without task ownership.
# mcp 2.0 uses a task-owned anyio.Lock held across the yielded resource request (a session-long GET blocks
# every POST; HTTPX may close the generator from another task). A binary semaphore drops task ownership.
import anyio
self.context.lock = anyio.Semaphore(1, max_value=1)
self._hermes_server_name = server_name
self._hermes_home = ""
# A config-supplied client_id rejected as invalid_client means the *config* is wrong —
# re-registration can't help, so only dynamically-registered clients auto-heal.
# A config-supplied client_id rejected as invalid_client means the *config* is wrong — only DCR clients auto-heal.
self._hermes_preregistered = preregistered
def _hermes_storage(self):
"""The context storage when it is a ``HermesTokenStorage``, else None."""
from tools.mcp_oauth import HermesTokenStorage
storage = self.context.storage
return storage if isinstance(storage, HermesTokenStorage) else None
return self.context.storage if isinstance(self.context.storage, HermesTokenStorage) else None
def _log_nonfatal(self, what: str, exc: BaseException) -> None:
logger.debug("MCP OAuth '%s': %s failed (non-fatal): %s", self._hermes_server_name, what, exc)
async def _initialize(self) -> None:
"""Load stored state, seed ``token_expiry_time``, restore/prefetch metadata. The SDK's
``_initialize`` never calls ``update_token_expiry``, so ``is_token_valid()`` is True for
any loaded token regardless of age and a restarted process ships stale Bearer tokens;
seeding the expiry (``HermesTokenStorage`` persists absolute ``expires_at``) makes the SDK
refresh first. Metadata is restored from disk, else discovered pre-flight when we hold
tokens but no metadata: otherwise ``_refresh_token`` guesses ``{server_url}/token``
(wrong for split-origin providers), 404s, and we fall through to browser reauth."""
``_initialize`` never calls ``update_token_expiry``, so a restarted process would ship stale
Bearer tokens as "valid"; seeding the expiry (``HermesTokenStorage`` persists absolute
``expires_at``) makes the SDK refresh first. Metadata is restored from disk, else discovered
pre-flight when we hold tokens but no metadata: otherwise ``_refresh_token`` guesses
``{server_url}/token`` (wrong for split-origin providers), 404s, and we fall to browser reauth."""
await super()._initialize()
tokens = self.context.current_tokens
if tokens is not None and tokens.expires_in is not None:
@@ -99,10 +92,9 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
self._log_nonfatal("pre-flight metadata discovery", exc)
async def _prefetch_oauth_metadata(self) -> None:
"""Fetch PRM + ASM from the well-known endpoints before the first request, using the
SDK's own URL builders/response handlers so we track whatever the pinned SDK expects."""
# The SDK's httpx flavour, not Hermes' — mcp 2.0 builds on httpx2 and
# `create_oauth_metadata_request` returns *its* Request objects.
"""Fetch PRM + ASM from the well-known endpoints before the first request, via the SDK's own URL
builders/response handlers so we track whatever the pinned SDK expects."""
# The SDK's httpx flavour, not Hermes': `create_oauth_metadata_request` returns *its* (httpx2) Request objects.
from tools.mcp_tool import sdk_httpx
httpx = sdk_httpx()
if httpx is None: # pragma: no cover — SDK import would have failed
@@ -119,7 +111,6 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
except httpx.HTTPError as exc:
logger.debug("MCP OAuth '%s': %s discovery to %s failed: %s", self._hermes_server_name, label, url, exc)
return None
async with httpx.AsyncClient(timeout=10.0) as client:
# PRM discovery to learn the authorization_server URL.
for url in build_protected_resource_metadata_discovery_urls(None, server_url):
@@ -160,8 +151,7 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
async def _is_invalid_client_at_token_endpoint(self, response: Any) -> bool:
"""True when *response* is the token endpoint (same scheme/host/path, query ignored)
rejecting our client_id with ``invalid_client`` — whole word, so RFC 7591's
``invalid_client_metadata`` does not trip it. The body is read only after the endpoint
matches."""
``invalid_client_metadata`` does not trip it. The body is read only after the endpoint matches."""
from urllib.parse import urlsplit
token_endpoint = getattr(getattr(self.context, "oauth_metadata", None), "token_endpoint", None)
req = getattr(response, "request", None)
@@ -171,29 +161,24 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
pa, pb = urlsplit(str(req.url)), urlsplit(str(token_endpoint))
except ValueError: # pragma: no cover — malformed URL
return False
if not (pa.scheme == pb.scheme and pa.netloc.lower() == pb.netloc.lower()
and pa.path.rstrip("/") == pb.path.rstrip("/")):
if (pa.scheme, pa.netloc.lower(), pa.path.rstrip("/")) != (pb.scheme, pb.netloc.lower(), pb.path.rstrip("/")):
return False
body = await response.aread()
return re.search(rb"\binvalid_client\b", body.lower()) is not None
return re.search(rb"\binvalid_client\b", (await response.aread()).lower()) is not None
async def _maybe_flag_poisoned_client(self, response: Any) -> None:
"""An ``invalid_client`` rejection of our ``client_id`` at the token endpoint proves the
cached registration is dead server-side: delete ``client.json`` (+ stale metadata) so the
SDK re-runs DCR next flow. Conservative: acts ONLY on 400/401 at the discovered
``token_endpoint`` (the only request carrying our ``client_id``) with ``invalid_client``
in the body; pre-registered clients are never poisoned; any failure is swallowed. The
browser-side "Redirect URI Mismatch" case has no HTTP signal (``hermes mcp reauth``)."""
"""An ``invalid_client`` rejection of our ``client_id`` at the token endpoint proves the cached registration
is dead server-side: delete ``client.json`` (+ stale metadata) so the SDK re-runs DCR next flow.
Conservative: acts ONLY on 400/401 at the discovered ``token_endpoint`` (the only request carrying our
``client_id``) with ``invalid_client`` in the body; pre-registered clients are never poisoned; any failure
is swallowed. The browser-side "Redirect URI Mismatch" case has no HTTP signal (``hermes mcp reauth``)."""
try:
if self._hermes_preregistered or getattr(response, "status_code", None) not in (400, 401):
return
if not await self._is_invalid_client_at_token_endpoint(response):
if (self._hermes_preregistered or getattr(response, "status_code", None) not in (400, 401)
or not await self._is_invalid_client_at_token_endpoint(response)):
return
storage = self._hermes_storage()
# If the rejected client_id was our CIMD URL, re-presenting it would loop (the
# server already fetched and refused it). Drop the URL so the retry takes DCR, and
# mark it on disk so the next process doesn't walk back into the same refusal
# (`hermes mcp login` clears the marker).
# A rejected CIMD URL would loop if re-presented (the server already fetched and refused
# it): drop it so the retry takes DCR, and mark it on disk so the next process doesn't walk
# back into the same refusal (`hermes mcp login` clears the marker).
cimd_url = getattr(self.context, "client_metadata_url", None)
if cimd_url and getattr(self.context.client_info, "client_id", None) == cimd_url:
logger.warning("MCP OAuth '%s': authorization server rejected our Client ID Metadata Document (%s) "
@@ -211,18 +196,15 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
self._log_nonfatal("invalid_client detection", exc)
async def async_auth_flow(self, request): # type: ignore[override]
# Pre-flow hook: reload from disk if it changed (non-fatal on error).
try:
try: # pre-flow hook: reload from disk if it changed (non-fatal on error)
await get_manager().invalidate_if_disk_changed(self._hermes_server_name, hermes_home=self._hermes_home)
except Exception as exc: # pragma: no cover — defensive
self._log_nonfatal("pre-flow disk-watch", exc)
# Bridge the bidirectional generator by hand: a naive ``async for item in inner: yield
# item`` DISCARDS the responses httpx sends back via ``asend``, and the SDK crashes on None.
inner = super().async_auth_flow(request)
resource_lock_released = False
resource_lock_released = retry_after_concurrent_auth = False
sent_access_token = None
retry_after_concurrent_auth = False
try:
outgoing = await inner.__anext__()
while True:
@@ -252,12 +234,11 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
self._persist_oauth_metadata_if_changed() # metadata discovered lazily in the 401 branch
finally:
if resource_lock_released:
# Balance the SDK's surrounding ``async with`` even when HTTPX cancels/closes
# the flow mid-request. Shield only this local bookkeeping.
# Balance the SDK's surrounding ``async with`` even when HTTPX cancels/closes the
# flow mid-request; shield only this local bookkeeping.
import anyio
with anyio.CancelScope(shield=True):
await self.context.lock.acquire()
if retry_after_concurrent_auth:
yield request
self._persist_oauth_metadata_if_changed()
@@ -274,13 +255,12 @@ class MCPOAuthManager:
def __init__(self) -> None:
self._entries: dict[tuple[str, str], _ProviderEntry] = {}
self._entries_lock = threading.Lock()
# Strong refs to in-flight 401 tasks so the loop's weak bookkeeping cannot GC them
# mid-run and leave `await pending` hanging forever.
# Strong refs to in-flight 401 tasks so the loop's weak bookkeeping cannot GC them mid-run.
self._inflight_tasks: set[asyncio.Task] = set()
def get_or_build_provider(self, server_name: str, server_url: str, oauth_config: Optional[dict]) -> Optional[Any]:
"""Cached OAuth provider for ``server_name``, built on first use (rebuilt when
``server_url`` changes). None if the MCP SDK's OAuth support is unavailable."""
"""Cached OAuth provider for ``server_name``, built on first use (rebuilt when ``server_url`` changes);
None if the MCP SDK's OAuth support is unavailable."""
key = self._key(server_name)
with self._entries_lock:
entry = self._entries.get(key)
@@ -306,8 +286,7 @@ class MCPOAuthManager:
if _HERMES_PROVIDER_CLS is None:
logger.warning("MCP OAuth '%s': SDK auth module unavailable", server_name)
return None
# Local imports avoid circular deps at module import time.
from tools.mcp_dashboard_oauth import get_dashboard_oauth_flow
from tools.mcp_dashboard_oauth import get_dashboard_oauth_flow # lazy: circular at import time
from tools.mcp_oauth import _OAUTH_AVAILABLE, OAuthNonInteractiveError, _is_interactive
from tools.mcp_oauth_provider import build_provider_kwargs, prepare_oauth_config
if not _OAUTH_AVAILABLE:
@@ -322,8 +301,7 @@ class MCPOAuthManager:
**build_provider_kwargs(cfg, storage, ssh_proxy_hint=False))
def remove(self, server_name: str, *, hermes_home: str | Path | None = None) -> _ProviderEntry | None:
"""Evict the provider from cache AND delete tokens from disk (``hermes mcp remove``,
and ``hermes mcp login`` during forced re-auth)."""
"""Evict the provider from cache AND delete tokens from disk (``hermes mcp remove`` / forced re-auth)."""
entry = self.evict(server_name, hermes_home=hermes_home)
from tools.mcp_oauth import remove_oauth_tokens
remove_oauth_tokens(server_name, hermes_home=hermes_home)
@@ -343,23 +321,20 @@ class MCPOAuthManager:
return self._entries.pop(self._key(server_name, hermes_home), None)
async def invalidate_if_disk_changed(self, server_name: str, *, hermes_home: str | Path | None = None) -> bool:
"""Force the SDK provider to reload when the tokens file mtime changed; True if
invalidated. A cron job writes fresh tokens and the next tool call picks them up."""
"""Force the SDK provider to reload when the tokens file mtime changed (e.g. a cron refresh); True if so."""
from tools.mcp_oauth import _get_token_dir, _safe_filename
entry = self._entries.get(self._key(server_name, hermes_home))
if entry is None or entry.provider is None:
return False
async with entry.lock:
tokens_path = _get_token_dir(hermes_home) / f"{_safe_filename(server_name)}.json"
try:
mtime_ns = tokens_path.stat().st_mtime_ns
except (FileNotFoundError, OSError):
mtime_ns = (_get_token_dir(hermes_home) / f"{_safe_filename(server_name)}.json").stat().st_mtime_ns
except OSError:
return False
if mtime_ns == entry.last_mtime_ns:
return False
old, entry.last_mtime_ns = entry.last_mtime_ns, mtime_ns
# `_initialized` is private SDK API but stable across the versions we pin
# (>=1.26.0); resetting it forces a reload.
# `_initialized` is private SDK API but stable across the pinned versions (>=1.26.0).
if hasattr(entry.provider, "_initialized"):
entry.provider._initialized = False # noqa: SLF001
logger.info("MCP OAuth '%s': tokens file changed (mtime %d -> %d), forcing reload", server_name, old, mtime_ns)
@@ -368,37 +343,33 @@ class MCPOAuthManager:
async def _recover_401(self, server_name: str, entry: _ProviderEntry, key: str, pending: asyncio.Future) -> None:
"""Single recovery attempt behind *pending*; always clears the dedup slot."""
try:
# Disk changed (external refresh)? Else: if the SDK can refresh in place, let the
# caller retry (the httpx.Auth flow refreshes on the next request).
can_refresh = True
if not await self.invalidate_if_disk_changed(server_name):
# Disk changed (external refresh)? Else: if the SDK can refresh in place, let the caller retry.
can_refresh = await self.invalidate_if_disk_changed(server_name)
if not can_refresh:
try:
can_refresh = bool(entry.provider.context.can_refresh_token())
except Exception: # no context / not callable / probe failed
can_refresh = False
if not pending.done():
pending.set_result(can_refresh)
except Exception as exc: # pragma: no cover — defensive
logger.warning("MCP OAuth '%s': 401 handler failed: %s", server_name, exc)
if not pending.done():
pending.set_result(False)
can_refresh = False
finally:
entry.pending_401.pop(key, None)
if not pending.done():
pending.set_result(can_refresh)
async def handle_401(self, server_name: str, failed_access_token: Optional[str] = None) -> bool:
"""Handle a 401 from a tool call. True: a (possibly new) token is available — reconnect
and retry. False: no recovery path — surface ``needs_reauth`` so the model stops
hallucinating manual refreshes. Concurrent 401s with the same ``failed_access_token``
fire one recovery attempt; the rest await its future."""
"""Handle a 401 from a tool call. True: a (possibly new) token is available — reconnect and retry. False: no
recovery path — surface ``needs_reauth`` so the model stops hallucinating manual refreshes. Concurrent 401s
with the same ``failed_access_token`` fire one recovery attempt; the rest await its future."""
entry = self._entries.get(self._key(server_name))
if entry is None or entry.provider is None:
return False
key = failed_access_token or "<unknown>"
loop = asyncio.get_running_loop()
async with entry.lock:
pending = entry.pending_401.get(key)
if pending is None:
pending = entry.pending_401[key] = loop.create_future()
pending = entry.pending_401[key] = asyncio.get_running_loop().create_future()
task = asyncio.create_task(self._recover_401(server_name, entry, key, pending))
self._inflight_tasks.add(task)
task.add_done_callback(self._inflight_tasks.discard)
+46 -75
View File
@@ -20,27 +20,25 @@ _mcp_stderr_log_lock = threading.Lock()
def _get_mcp_stderr_log() -> Any:
"""Shared append-mode handle for MCP subprocess stderr, opened once per process. Must
expose a real fd (``fileno()``) because asyncio wires the child's stderr directly to it.
Falls back to ``/dev/null``, then real stderr."""
"""Shared append-mode handle for MCP subprocess stderr, opened once per process. Must expose a
real fd (asyncio wires the child's stderr to it); falls back to ``/dev/null``, then real stderr."""
global _mcp_stderr_log_fh
with _mcp_stderr_log_lock:
if _mcp_stderr_log_fh is not None:
return _mcp_stderr_log_fh
try:
from hermes_constants import get_hermes_home
log_dir = get_hermes_home() / "logs"
log_dir.mkdir(parents=True, exist_ok=True)
# Line-buffered so output lands promptly; errors="replace" tolerates garbled binary.
fh = open(log_dir / "mcp-stderr.log", "a", encoding="utf-8", errors="replace", buffering=1)
fh.fileno() # confirm a real fd before committing
_mcp_stderr_log_fh = fh
except Exception as exc: # pragma: no cover — best-effort fallback
logger.debug("Failed to open MCP stderr log, using devnull: %s", exc)
if _mcp_stderr_log_fh is None:
try:
_mcp_stderr_log_fh = open(os.devnull, "w", encoding="utf-8")
except Exception:
_mcp_stderr_log_fh = sys.stderr
from hermes_constants import get_hermes_home
log_dir = get_hermes_home() / "logs"
log_dir.mkdir(parents=True, exist_ok=True)
# Line-buffered so output lands promptly; errors="replace" tolerates garbled binary.
fh = open(log_dir / "mcp-stderr.log", "a", encoding="utf-8", errors="replace", buffering=1)
fh.fileno() # confirm a real fd before committing
_mcp_stderr_log_fh = fh
except Exception as exc: # pragma: no cover — best-effort fallback
logger.debug("Failed to open MCP stderr log, using devnull: %s", exc)
try:
_mcp_stderr_log_fh = open(os.devnull, "w", encoding="utf-8")
except Exception:
_mcp_stderr_log_fh = sys.stderr
return _mcp_stderr_log_fh
@@ -49,8 +47,7 @@ def _write_stderr_log_header(server_name: str) -> None:
(per-line prefixes would need a pipe + reader thread)."""
fh = _core._get_mcp_stderr_log()
try:
ts = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
fh.write(f"\n===== [{ts}] starting MCP server '{server_name}' =====\n")
fh.write(f"\n===== [{datetime.now():%Y-%m-%d %H:%M:%S}] starting MCP server '{server_name}' =====\n")
fh.flush()
except Exception:
pass
@@ -59,16 +56,14 @@ def _write_stderr_log_header(server_name: str) -> None:
# Env vars safe to pass to stdio subprocesses (no secrets).
_SAFE_ENV_KEYS = frozenset({"PATH", "HOME", "USER", "LANG", "LC_ALL", "TERM", "SHELL", "TMPDIR"})
# Windows process/location vars needed by launcher-style tools (e.g. Docker
# Desktop's MCP plugin discovery); none carry secrets.
# Windows process/location vars needed by launcher-style tools (e.g. Docker Desktop's MCP plugin discovery).
_SAFE_ENV_KEYS_CASE_INSENSITIVE = frozenset({
"ALLUSERSPROFILE", "APPDATA", "COMMONPROGRAMFILES", "COMMONPROGRAMFILES(X86)",
"COMMONPROGRAMW6432", "COMPUTERNAME", "COMSPEC", "HOMEDRIVE", "HOMEPATH",
"LOCALAPPDATA", "NUMBER_OF_PROCESSORS", "OS", "PATHEXT", "PROCESSOR_ARCHITECTURE",
"PROGRAMDATA", "PROGRAMFILES", "PROGRAMFILES(X86)", "PROGRAMW6432", "PUBLIC",
"SYSTEMDRIVE", "SYSTEMROOT", "TEMP", "TMP", "USERDOMAIN", "USERNAME",
"USERPROFILE", "WINDIR",
})
"USERPROFILE", "WINDIR"})
# ${VAR_NAME} interpolation; any non-} chars allowed so MY-VAR / my.var work.
_ENV_VAR_PATTERN = re.compile(r"\$\{([^}]+)\}")
@@ -80,11 +75,9 @@ def _workspace_folder() -> str:
try:
from tools.file_tools import _authoritative_workspace_root
root = _authoritative_workspace_root()
if root:
return root
except Exception:
pass
return os.getcwd()
root = None
return root or os.getcwd()
def _workspace_basename() -> str:
@@ -94,12 +87,8 @@ def _workspace_basename() -> str:
# Cursor's case-sensitive context vars -> resolver.
_CONTEXT_VAR_RESOLVERS = {
"userHome": lambda: os.path.expanduser("~"),
"workspaceFolder": lambda: _core._workspace_folder(),
"workspaceFolderBasename": _workspace_basename,
"pathSeparator": lambda: os.sep,
"/": lambda: os.sep,
}
"userHome": lambda: os.path.expanduser("~"), "workspaceFolder": lambda: _core._workspace_folder(),
"workspaceFolderBasename": _workspace_basename, "pathSeparator": lambda: os.sep, "/": lambda: os.sep}
def _build_safe_env(user_env: Optional[dict]) -> dict:
@@ -120,8 +109,7 @@ def _build_safe_env(user_env: Optional[dict]) -> dict:
def _which_with_config_pathext(command: str, path_arg, env: dict):
"""``shutil.which`` retried under the config env's PATHEXT (Windows only):
``which(path=...)`` uses the PARENT's PATHEXT, not the config env's."""
"""``shutil.which`` retried under the config env's PATHEXT (Windows only; ``which`` uses the PARENT's)."""
cfg_pathext = next((v for k, v in env.items() if k.upper() == "PATHEXT" and isinstance(v, str) and v.strip()), None)
if not cfg_pathext or cfg_pathext == os.environ.get("PATHEXT"):
return None
@@ -137,28 +125,20 @@ def _which_with_config_pathext(command: str, path_arg, env: dict):
def _node_fallback(command: str) -> str:
"""Well-known Node install locations for bare ``npx``/``npm``/``node`` when PATH lookup
failed; *command* unchanged when none is executable."""
"""Well-known Node install locations for bare ``npx``/``npm``/``node``; *command* unchanged when none exists."""
home = os.path.expanduser("~")
hermes_home = os.path.expanduser(os.getenv("HERMES_HOME", os.path.join(home, ".hermes")))
candidates = [
os.path.join(hermes_home, "node", "bin", command),
os.path.join(home, ".local", "bin", command),
# Canonical Node location (from-source Linux, Hermes Docker image, Intel Homebrew). Needed
# when a hand-authored env.PATH omits it: npx's shebang re-execs /usr/bin/env node.
os.path.join(os.sep, "usr", "local", "bin", command)]
for candidate in candidates:
if os.path.isfile(candidate) and os.access(candidate, os.X_OK):
return candidate
return command
# /usr/local/bin: canonical Node location (from-source Linux, Hermes Docker image, Intel Homebrew),
# needed when a hand-authored env.PATH omits it — npx's shebang re-execs /usr/bin/env node.
candidates = (os.path.join(hermes_home, "node", "bin", command), os.path.join(home, ".local", "bin", command),
os.path.join(os.sep, "usr", "local", "bin", command))
return next((c for c in candidates if os.path.isfile(c) and os.access(c, os.X_OK)), command)
def _resolve_stdio_command(command: str, env: dict) -> tuple[str, dict]:
"""Resolve a stdio command against the exact subprocess env, mainly so bare
``npx``/``npm``/``node`` work under a filtered PATH."""
"""Resolve a stdio command against the exact subprocess env (bare ``npx``/``npm``/``node`` under a filtered PATH)."""
resolved_command = os.path.expanduser(str(command).strip())
resolved_env = dict(env or {})
if os.sep not in resolved_command:
path_arg = resolved_env.get("PATH")
which_hit = shutil.which(resolved_command, path=path_arg)
@@ -168,7 +148,6 @@ def _resolve_stdio_command(command: str, env: dict) -> tuple[str, dict]:
resolved_command = which_hit
elif resolved_command in {"npx", "npm", "node"}:
resolved_command = _node_fallback(resolved_command)
command_dir = os.path.dirname(resolved_command)
if command_dir:
resolved_env = _prepend_path(resolved_env, command_dir)
@@ -269,15 +248,14 @@ def _interpolate_env_vars(value):
return value
# (server_name, dotted key path) pairs already warned about; config loads
# happen on every discovery pass, so warn once per process.
# (server_name, dotted key path) pairs already warned about: config loads repeat per discovery pass.
_whitespace_warned: Set[Tuple[str, str]] = set()
def _warn_hidden_whitespace(server_name: str, config: dict) -> List[str]:
"""Warn once per (server, key path) about string values with leading/trailing whitespace (a
pasted newline causes opaque auth failures, invisible in config.yaml). Advisory only: values
are never mutated (could be intentional) nor logged (often secrets). Returns flagged paths."""
pasted newline causes opaque auth failures). Advisory only: values are never mutated (could be
intentional) nor logged (often secrets). Returns flagged paths."""
flagged: List[str] = []
def _walk(value: Any, path: str) -> None:
@@ -289,18 +267,14 @@ def _warn_hidden_whitespace(server_name: str, config: dict) -> List[str]:
elif isinstance(value, list):
for i, v in enumerate(value):
_walk(v, f"{path}[{i}]")
_walk(config, "")
for key_path in flagged:
if (server_name, key_path) in _whitespace_warned:
continue
_whitespace_warned.add((server_name, key_path))
logger.warning(
"MCP server '%s': config value '%s' has hidden leading or "
"trailing whitespace — this often causes authentication or "
"connection failures. Check for stray spaces/newlines in "
"config.yaml (or the referenced env var).",
server_name, key_path)
if (server_name, key_path) not in _whitespace_warned:
_whitespace_warned.add((server_name, key_path))
logger.warning(
"MCP server '%s': config value '%s' has hidden leading or trailing whitespace — this often "
"causes authentication or connection failures. Check for stray spaces/newlines in config.yaml "
"(or the referenced env var).", server_name, key_path)
return flagged
@@ -310,20 +284,18 @@ def _filter_suspicious_mcp_servers(servers: Dict[str, dict]) -> Dict[str, dict]:
from hermes_cli.mcp_security import validate_mcp_server_entry
except Exception:
return servers
safe_servers = {}
for name, cfg in servers.items():
issues = validate_mcp_server_entry(name, cfg) if isinstance(cfg, dict) else None
if issues:
logger.warning("Skipping suspicious MCP server '%s': %s", name, "; ".join(issues))
continue
safe_servers[name] = cfg
else:
safe_servers[name] = cfg
return safe_servers
def _portable_mcp_servers(safe_servers: Dict[str, dict]) -> None:
"""Merge plugin-provided (portable) MCP servers into *safe_servers*; native config wins
on a name clash. Never raises."""
"""Merge plugin-provided (portable) MCP servers into *safe_servers*; native config wins on a clash. Never raises."""
try:
from hermes_cli.plugins import discover_plugins, get_plugin_manager
discover_plugins()
@@ -331,15 +303,14 @@ def _portable_mcp_servers(safe_servers: Dict[str, dict]) -> None:
for name, cfg in _core._filter_suspicious_mcp_servers(portable).items():
if name in safe_servers:
logger.warning("Portable MCP server '%s' conflicts with native config; skipping", name)
continue
safe_servers[name] = dict(cfg)
else:
safe_servers[name] = dict(cfg)
except Exception:
logger.debug("Failed to load portable MCP servers", exc_info=True)
def _load_mcp_config() -> Dict[str, dict]:
"""Read ``mcp_servers`` from config.yaml as ``{name: config}`` (empty on error or in safe
mode); ``${VAR}`` placeholders are interpolated after ``.env`` is loaded."""
"""``mcp_servers`` from config.yaml as ``{name: config}`` (empty on error / safe mode), ``${VAR}`` interpolated."""
try:
from hermes_cli.config import load_config
from utils import env_var_enabled as _env_enabled
+39 -61
View File
@@ -14,29 +14,23 @@ from tools.mcp_tool_common import _sanitize_error, _core
logger = logging.getLogger("tools.mcp_tool")
# Stateless (2026-07-28) servers reject a legacy ``initialize`` with
# UnsupportedProtocolVersion (-32022) or plain method-not-found.
# Stateless (2026-07-28) servers reject a legacy ``initialize`` with this or plain method-not-found.
_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION = -32022
def _jsonrpc_code(exc: BaseException):
"""Structural ``MCPError.error.code`` (None when absent)."""
return getattr(getattr(exc, "error", None), "code", None)
def _jsonrpc_matches(exc: BaseException, code, codes: tuple, markers: tuple) -> bool:
"""Structural *code* in *codes*, else any *marker* in ``str(exc).lower()``. Never ``isinstance``
on SDK exception types: they arrive wrapped in ExceptionGroups and drift across generations."""
def _jsonrpc_matches(exc: BaseException, codes: tuple, markers: tuple, code=None) -> bool:
"""Structural ``MCPError.error.code`` (or *code*) in *codes*, else any *marker* in ``str(exc).lower()``. Never
``isinstance`` on SDK exception types: they arrive wrapped in ExceptionGroups and drift across generations."""
code = getattr(getattr(exc, "error", None), "code", None) or code
return code in codes or any(marker in str(exc).lower() for marker in markers)
def _handshake_rejected_as_modern(exc: BaseException) -> bool:
"""True when a failed ``initialize`` signals a stateless-only (2026-07-28) server."""
return _jsonrpc_matches(
exc, _jsonrpc_code(exc) or getattr(exc, "code", None),
(_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _core._JSONRPC_METHOD_NOT_FOUND),
exc, (_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _core._JSONRPC_METHOD_NOT_FOUND),
("unsupported protocol version", str(_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION)),
) or _is_method_not_found_error(exc)
code=getattr(exc, "code", None)) or _is_method_not_found_error(exc)
def _is_method_not_found_error(exc: BaseException) -> bool:
@@ -44,13 +38,12 @@ def _is_method_not_found_error(exc: BaseException) -> bool:
substring fallback includes "Unknown method: <name>" — without it the ping→list_tools keepalive
fallback never latches and reconnect-loops."""
return _jsonrpc_matches(
exc, _jsonrpc_code(exc), (_core._JSONRPC_METHOD_NOT_FOUND,),
exc, (_core._JSONRPC_METHOD_NOT_FOUND,),
(str(_core._JSONRPC_METHOD_NOT_FOUND), "method not found", "unknown method", "not found: ping"))
class InvalidMcpUrlError(ValueError):
"""A remote MCP server's ``url`` is not parseable http(s):// — validated once at startup so we
fail fast instead of burning the reconnect-backoff loop."""
"""A remote MCP server's ``url`` is not parseable http(s):// — validated once at startup to fail fast."""
class NonMcpEndpointError(ConnectionError):
@@ -93,11 +86,10 @@ def _classify_mcp_failure(exc: BaseException) -> str:
def _validate_remote_mcp_url(server_name: str, url: Any) -> str:
"""The stripped URL if it is a valid http(s) URL; else InvalidMcpUrlError naming the server
(non-string, other scheme — stdio servers use ``command`` — or empty host)."""
"""The stripped URL if valid http(s); else InvalidMcpUrlError naming the server (non-string, other scheme —
stdio servers use ``command`` — or empty host)."""
def _bad(detail: str) -> InvalidMcpUrlError:
return InvalidMcpUrlError(f"Invalid MCP URL for '{server_name}': {detail}")
if not isinstance(url, str):
raise _bad(f"expected a string, got {type(url).__name__}")
stripped = url.strip()
@@ -133,13 +125,11 @@ def _resolve_client_cert(server_name: str, config: dict):
if not os.path.isfile(expanded):
raise FileNotFoundError(f"{prefix}{label} not found at {expanded!r}")
return expanded
if not isinstance(raw_cert, (list, tuple)):
cert_path = _expand(raw_cert, "client_cert")
return (cert_path, _expand(raw_key, "client_key")) if raw_key is not None else cert_path # combined PEM
if raw_key is not None:
raise ValueError(f"{prefix}specify either client_cert as a list [cert, key] OR "
f"client_cert + client_key, not both")
raise ValueError(f"{prefix}specify either client_cert as a list [cert, key] OR client_cert + client_key, not both")
if len(raw_cert) not in (2, 3):
raise ValueError(f"{prefix}client_cert list form must have 2 or 3 elements (got {len(raw_cert)})")
pair = (_expand(raw_cert[0], "client_cert[0]"), _expand(raw_cert[1], "client_cert[1]"))
@@ -161,7 +151,6 @@ def _resolve_identity_header(server_name: str, config: dict):
def _ignore(detail: str, *args):
logger.warning("MCP server '%s': identity_header " + detail + " — ignoring", server_name, *args)
return None
if not isinstance(raw, dict):
return _ignore("must be a mapping with 'name' and 'value'/'value_from' keys (got %s)", type(raw).__name__)
name = raw.get("name")
@@ -210,7 +199,6 @@ def _make_redirect_header_stripper(original_url, *, strict: bool = False,
for _name in configured_header_names if strict else ():
while _name in headers:
del headers[_name]
return _strip_on_cross_origin_redirect
@@ -236,7 +224,6 @@ def _format_connect_error(exc: BaseException) -> str:
text = "" if getattr(current, "exceptions", None) else str(current).strip()
messages = ([text] if text else []) + [m for child in _exc_children(current) for m in _flatten_messages(child)]
return messages or [current.__class__.__name__]
missing = _find_missing(exc)
if not missing:
return _sanitize_error("; ".join(list(dict.fromkeys(_flatten_messages(exc)))[:3]))
@@ -248,11 +235,6 @@ def _format_connect_error(exc: BaseException) -> str:
return _sanitize_error(message)
# Lazily-built caches so this module imports without the SDK OAuth module.
_AUTH_ERROR_TYPES: tuple = ()
_HTTP_STATUS_ERROR_TYPES: Optional[tuple] = None
def _optional_types(module: str, *names: str) -> list:
"""``[module.name, ...]`` or ``[]`` when the module/attribute is unavailable."""
try:
@@ -262,35 +244,33 @@ def _optional_types(module: str, *names: str) -> list:
return []
def _http_status_error_types() -> tuple:
"""``HTTPStatusError`` from both httpx flavours: a 401 may come from the SDK's own stack
(``httpx2`` on mcp >= 2.0) or Hermes' pinned ``httpx``; the classes are unrelated."""
global _HTTP_STATUS_ERROR_TYPES
if _HTTP_STATUS_ERROR_TYPES is None:
sdk_mod = _core.sdk_httpx()
_HTTP_STATUS_ERROR_TYPES = tuple(dict.fromkeys(
([sdk_mod.HTTPStatusError] if sdk_mod is not None else []) + _optional_types("httpx", "HTTPStatusError")))
return _HTTP_STATUS_ERROR_TYPES
# Lazily-built ``(auth_types, http_status_types)`` so this module imports without the SDK OAuth module.
_AUTH_ERROR_TYPES: Optional[tuple] = None
def _get_auth_error_types() -> tuple:
"""Cached MCP OAuth failure types: SDK ``OAuthFlowError``/``OAuthTokenError`` (+ legacy
``UnauthorizedError``), our ``OAuthNonInteractiveError``, and both ``HTTPStatusError`` flavours
(which still need the 401 check in :func:`_is_auth_error`)."""
"""Cached ``(auth_types, http_status_types)``: SDK ``OAuthFlowError``/``OAuthTokenError`` (+ legacy
``UnauthorizedError``), our ``OAuthNonInteractiveError``, and ``HTTPStatusError`` from both httpx
flavours — a 401 may come from the SDK's own stack (``httpx2`` on mcp >= 2.0) or Hermes' pinned
``httpx``; the classes are unrelated and still need the 401 check in :func:`_is_auth_error`."""
global _AUTH_ERROR_TYPES
if not _AUTH_ERROR_TYPES:
_AUTH_ERROR_TYPES = (*_optional_types("mcp.client.auth", "OAuthFlowError", "OAuthTokenError"),
*_optional_types("mcp.client.auth", "UnauthorizedError"), # older SDKs
*_optional_types("tools.mcp_oauth", "OAuthNonInteractiveError"),
*_http_status_error_types())
if not (_AUTH_ERROR_TYPES and _AUTH_ERROR_TYPES[0]): # retry while empty (SDK may import later)
sdk_mod = _core.sdk_httpx()
http_types = tuple(dict.fromkeys(
([sdk_mod.HTTPStatusError] if sdk_mod is not None else []) + _optional_types("httpx", "HTTPStatusError")))
auth_types = (*_optional_types("mcp.client.auth", "OAuthFlowError", "OAuthTokenError"),
*_optional_types("mcp.client.auth", "UnauthorizedError"), # older SDKs
*_optional_types("tools.mcp_oauth", "OAuthNonInteractiveError"), *http_types)
_AUTH_ERROR_TYPES = (auth_types, http_types)
return _AUTH_ERROR_TYPES
def _is_auth_error(exc: BaseException) -> bool:
"""True if ``exc`` indicates an MCP OAuth failure; ``HTTPStatusError`` counts only with status 401."""
if not isinstance(exc, _get_auth_error_types()):
auth_types, http_types = _get_auth_error_types()
if not isinstance(exc, auth_types):
return False
return getattr(exc.response, "status_code", None) == 401 if isinstance(exc, _http_status_error_types()) else True
return getattr(exc.response, "status_code", None) == 401 if isinstance(exc, http_types) else True
# Lower-cased substrings meaning the transport session expired / was GC'd (OAuth token still valid).
@@ -299,18 +279,17 @@ _SESSION_EXPIRED_MARKERS: tuple = (
"unknown session", "session terminated", "closedresourceerror", "closed resource",
"transport is closed", "connection closed", "broken pipe", "end of file")
# Node budget for ``_is_session_expired_error`` (the visited set breaks cycles; this bounds acyclic
# blow-ups). Well above ``sys.getrecursionlimit()`` so deep task-group nesting is fully scanned.
# Node budget for ``_is_session_expired_error`` (the visited set breaks cycles; this bounds acyclic blow-ups).
# Well above ``sys.getrecursionlimit()`` so deep task-group nesting is fully scanned.
_EXC_TRAVERSAL_MAX_NODES = 10_000
def _is_session_expired_error(exc: BaseException) -> bool:
"""True if ``exc`` looks like a transport session expiry (Streamable-HTTP servers GC session
state on idle TTL / restart / pod rotation while the OAuth token stays valid) — the fix is a
transport reconnect, not an OAuth refresh. Iterative walk over ``exceptions`` / ``__cause__`` /
``__context__`` with a visited set AND a node budget; every reachable node is inspected so an
InterruptedError anywhere overrides transport markers, and the chain walk matters because SDK
wrappers raise a generic RuntimeError *from* a message-less ClosedResourceError."""
"""True if ``exc`` looks like a transport session expiry (Streamable-HTTP servers GC session state on idle TTL /
restart / pod rotation while the OAuth token stays valid) — the fix is a transport reconnect, not an OAuth
refresh. Iterative walk over ``exceptions`` / ``__cause__`` / ``__context__`` with a visited set AND a node
budget; every reachable node is inspected so an InterruptedError anywhere overrides transport markers, and the
chain walk matters because SDK wrappers raise a generic RuntimeError *from* a message-less ClosedResourceError."""
# AnyIO stream exceptions are often message-less, so type checks complement marker matching.
transport_error_types = tuple(_optional_types("anyio", "BrokenResourceError", "ClosedResourceError", "EndOfStream"))
stack: "list[BaseException | None]" = [exc]
@@ -325,10 +304,9 @@ def _is_session_expired_error(exc: BaseException) -> bool:
budget -= 1
if isinstance(current, InterruptedError):
return False
# Messages vary across SDK versions/servers: a narrow allow-list of stable substrings avoids
# false positives.
# Messages vary across SDK versions/servers: a narrow allow-list of stable substrings avoids false positives.
msg = str(current).lower()
found = found or isinstance(current, transport_error_types) or any(m in msg for m in _SESSION_EXPIRED_MARKERS)
stack.extend(getattr(current, "exceptions", ()))
stack.extend((getattr(current, "__cause__", None), getattr(current, "__context__", None)))
stack.extend((*getattr(current, "exceptions", ()), getattr(current, "__cause__", None),
getattr(current, "__context__", None)))
return found
+98 -157
View File
@@ -1,6 +1,5 @@
"""Registry-facing sync handlers for MCP tools and utility tools (resources/prompts), plus
the per-call recovery ladder: trust gating, circuit breaker, auth (401) refresh,
session-expired reconnect and dead-stdio respawn retry."""
"""Registry-facing sync handlers for MCP tools and utility tools (resources/prompts), plus the per-call recovery
ladder: trust gating, circuit breaker, auth (401) refresh, session-expired reconnect and dead-stdio respawn retry."""
import logging
import asyncio
@@ -21,18 +20,27 @@ from tools.mcp_tool_content import (
from tools.mcp_tool_errors import _is_session_expired_error
logger = logging.getLogger("tools.mcp_tool")
_MISSING = object()
_NEEDS_REAUTH_MSG = (
"MCP server '{s}' requires re-authentication. Run `hermes mcp login {s}` (or delete the tokens file under "
"~/.hermes/mcp-tokens/ and restart). Do NOT retry this tool — ask the user to re-authenticate.")
_STDIO_NO_RESPAWN_MSG = (
"MCP server '{s}' stdio subprocess had exited (this is not a timeout — the call never reached the server). A "
"respawn was requested but no fresh session came back within {t:.0f}s. Wait a few seconds before retrying; if it "
"keeps failing the server is not starting and needs the user.")
_STDIO_DIED_AGAIN_MSG = (
"MCP server '{s}' respawned its stdio subprocess and it exited again immediately. The server is not starting "
"cleanly — do NOT retry this tool; ask the user to check the server's command and its stderr log.")
# --------------------------------------------------------------- pre-call gates
def _trust_gate_check(server_name: str, tool_name: str) -> Optional[str]:
"""Approval gate for write-capable tools on ``trust: untrusted`` servers. None to proceed,
else a ``tool_error``. Fail-closed: approval-system errors block."""
trust = _core._server_trust_levels.get(server_name, _core._TRUST_FULL)
if trust != _core._TRUST_UNTRUSTED or _core._tool_read_only_hints.get(server_name, {}).get(tool_name) is True:
if (_core._server_trust_levels.get(server_name, _core._TRUST_FULL) != _core._TRUST_UNTRUSTED
or _core._tool_read_only_hints.get(server_name, {}).get(tool_name) is True):
return None
# Lazy import: tools.approval routes the prompt to whichever surface owns the session.
try:
try: # lazy: tools.approval routes the prompt to whichever surface owns the session
from tools.approval import request_elicitation_consent
answer = request_elicitation_consent(
f"MCP tool '{tool_name}' on UNTRUSTED server '{server_name}' wants to run. This tool is write-capable "
@@ -56,15 +64,12 @@ def _check_circuit_breaker(server_name: str) -> Optional[str]:
"""Open-breaker error, or None when calls may proceed. After the cooldown the breaker is
half-open: the next call probes; success resets, failure re-bumps and re-arms the cooldown."""
failures = _core._server_error_counts.get(server_name, 0)
if failures < _core._CIRCUIT_BREAKER_THRESHOLD:
return None
age = time.monotonic() - _core._server_breaker_opened_at.get(server_name, 0.0)
if age >= _core._CIRCUIT_BREAKER_COOLDOWN_SEC:
if failures < _core._CIRCUIT_BREAKER_THRESHOLD or age >= _core._CIRCUIT_BREAKER_COOLDOWN_SEC:
return None
remaining = max(1, int(_core._CIRCUIT_BREAKER_COOLDOWN_SEC - age))
return tool_error(f"MCP server '{server_name}' is unreachable after {failures} consecutive failures. "
f"Auto-retry available in ~{remaining}s. Do NOT retry this tool yet — use alternative "
f"approaches or ask the user to check the MCP server.")
f"Auto-retry available in ~{max(1, int(_core._CIRCUIT_BREAKER_COOLDOWN_SEC - age))}s. Do NOT retry "
f"this tool yet — use alternative approaches or ask the user to check the MCP server.")
def _acquire_call_server(server_name: str, tool_timeout: float):
@@ -73,20 +78,16 @@ def _acquire_call_server(server_name: str, tool_timeout: float):
server task to rebuild (probing a dead transport would re-arm the breaker forever)."""
not_connected = tool_error(f"MCP server '{server_name}' is not connected")
server = _core._get_connected_server_for_call(server_name)
if not server:
_core._bump_server_error(server_name)
return None, not_connected
if server.session or _core._wait_for_server_session_ready(server, timeout=min(5.0, float(tool_timeout or 5.0))):
wait = min(5.0, float(tool_timeout or 5.0))
if server and (server.session or _core._wait_for_server_session_ready(server, timeout=wait)):
return server, None
_core._bump_server_error(server_name)
if _core._signal_reconnect(server):
if server and _core._signal_reconnect(server):
return None, tool_error(f"MCP server '{server_name}' transport is down; reconnect requested. Do NOT retry this "
f"tool immediately — give it a few seconds to come back.")
return None, not_connected
# ------------------------------------------------------------ breaker bookkeeping
def _result_is_error(result) -> bool:
"""True only for a JSON payload carrying an ``error`` key (non-JSON = success)."""
try:
@@ -97,10 +98,7 @@ def _result_is_error(result) -> bool:
def _record_call_outcome(server_name: str, result) -> Any:
"""Breaker bookkeeping: an error payload from the tool itself still counts as a strike."""
if _result_is_error(result):
_core._bump_server_error(server_name)
else:
_core._reset_server_error(server_name)
(_core._bump_server_error if _result_is_error(result) else _core._reset_server_error)(server_name)
return result
@@ -110,18 +108,17 @@ def _strike(server_name: str, message: str, **extra) -> str:
return tool_error(message, **extra)
def _mcp_loop_running() -> bool:
return _core._mcp_loop is not None and _core._mcp_loop.is_running()
def _lookup_reconnectable_server(server_name: str, require_loop: bool = False):
"""The registered server object when it can be signalled to reconnect, else None.
With *require_loop*, also None unless the MCP loop is running (nothing to wait on)."""
with _core._lock:
srv = _core._servers.get(server_name)
if srv is None or not hasattr(srv, "_reconnect_event") or (require_loop and not _mcp_loop_running()):
return None
return srv
def _mcp_loop_running() -> bool:
return _core._mcp_loop is not None and _core._mcp_loop.is_running()
ok = srv is not None and hasattr(srv, "_reconnect_event") and (_mcp_loop_running() or not require_loop)
return srv if ok else None
def _retry_once(server_name: str, retry_call, op_description: str, what: str):
@@ -138,8 +135,6 @@ def _retry_once(server_name: str, retry_call, op_description: str, what: str):
return result
# --------------------------------------------------------------- recovery ladder
def _handle_auth_error_and_retry(server_name: str, exc: BaseException, retry_call, op_description: str):
"""OAuth recovery + one retry; None when *exc* is not an auth error. ``handle_401`` decides
viability; if viable, signal a reconnect (fresh credentials), wait ready, retry once. Any
@@ -147,36 +142,28 @@ def _handle_auth_error_and_retry(server_name: str, exc: BaseException, retry_cal
if not _core._is_auth_error(exc):
return None
from tools.mcp_oauth_manager import get_manager
manager = get_manager()
try:
recovered = _core._run_on_mcp_loop(lambda: manager.handle_401(server_name, None), timeout=10)
recovered = _core._run_on_mcp_loop(lambda: get_manager().handle_401(server_name, None), timeout=10)
except Exception as rec_exc:
logger.warning("MCP OAuth '%s': recovery attempt failed: %s", server_name, rec_exc)
recovered = False
if recovered:
srv = _lookup_reconnectable_server(server_name)
# Recovery + reconnect is independent evidence of viability: close the breaker here, not
# only on retry success (else a failing retry pins it open forever).
# Recovery + reconnect is independent evidence of viability: close the breaker here, not only on
# retry success (else a failing retry pins it open forever).
if srv is not None and _core._signal_reconnect_and_wait(
server_name, srv, op_description=f"{op_description} after OAuth recovery", timeout=15):
_core._reset_server_error(server_name)
result = _retry_once(server_name, retry_call, op_description, "auth recovery")
if result is not None:
return result
return _strike(
server_name,
f"MCP server '{server_name}' requires re-authentication. Run `hermes mcp login "
f"{server_name}` (or delete the tokens file under ~/.hermes/mcp-tokens/ and restart). Do "
f"NOT retry this tool — ask the user to re-authenticate.",
needs_reauth=True, server=server_name)
return _strike(server_name, _NEEDS_REAUTH_MSG.format(s=server_name), needs_reauth=True, server=server_name)
def _handle_session_expired_and_retry(server_name: str, exc: BaseException, retry_call, op_description: str):
"""Transport reconnect + one retry on session expiry; None to fall through. Skips
``handle_401``: the token is valid, only the server-side session is stale."""
if not _is_session_expired_error(exc):
return None
srv = _lookup_reconnectable_server(server_name, require_loop=True)
srv = _lookup_reconnectable_server(server_name, require_loop=True) if _is_session_expired_error(exc) else None
if srv is None:
return None
logger.info("MCP server '%s': %s failed with session-expired error (%s); signalling transport reconnect "
@@ -206,27 +193,17 @@ def _handle_stdio_child_exited_and_retry(server_name: str, exc: Exception, retry
if _mcp_loop_running():
reconnected = _core._signal_reconnect_and_wait(
server_name, srv, op_description=op_description, timeout=_core._STDIO_RESPAWN_WAIT_SEC)
else:
# No MCP loop to wait on (non-async adapters, tests): still request the respawn.
else: # No MCP loop to wait on (non-async adapters, tests): still request the respawn.
_core._signal_reconnect(srv)
if not reconnected:
return _strike(
server_name,
f"MCP server '{server_name}' stdio subprocess had exited (this is not a timeout — the "
f"call never reached the server). A respawn was requested but no fresh session came "
f"back within {_core._STDIO_RESPAWN_WAIT_SEC:.0f}s. Wait a few seconds before retrying; "
f"if it keeps failing the server is not starting and needs the user.")
return _strike(server_name, _STDIO_NO_RESPAWN_MSG.format(s=server_name, t=_core._STDIO_RESPAWN_WAIT_SEC))
try:
return _record_call_outcome(server_name, retry_call())
except _StdioChildExited as retry_exc:
# Died again right after respawn: broken server; run()'s budget takes it to the park.
logger.warning("MCP server '%s': %s stdio subprocess exited again right after respawn (%s); not retrying "
"further.", server_name, op_description, retry_exc)
return _strike(
server_name,
f"MCP server '{server_name}' respawned its stdio subprocess and it exited again "
f"immediately. The server is not starting cleanly — do NOT retry this tool; ask the "
f"user to check the server's command and its stderr log.")
return _strike(server_name, _STDIO_DIED_AGAIN_MSG.format(s=server_name))
except Exception as retry_exc:
logger.warning("MCP %s/%s retry after stdio respawn failed: %s", server_name, op_description, retry_exc)
return _strike(server_name, _sanitize_error(
@@ -234,13 +211,18 @@ def _handle_stdio_child_exited_and_retry(server_name: str, exc: Exception, retry
f"{type(retry_exc).__name__}: {_exc_str(retry_exc)}"))
def _invoke_with_recovery(server_name: str, call_once: Callable[[], str], op: str,
recoverers, on_final_failure: Callable[[BaseException], None],
record_outcome: bool = False) -> str:
"""Run ``call_once``, walking ``recoverers`` (``(server_name, exc, retry_call, op) ->
Optional[str]``, None = not its kind; order matters) on failure. Unrecovered exceptions go
through ``on_final_failure`` and become the generic call-failed error. ``record_outcome``
applies breaker bookkeeping to the FIRST attempt only; retries own theirs."""
def _dispatch(server_name: str, server: Any, op: str, call, tool_timeout: float, recoverers,
on_final_failure: Callable[[BaseException], None], record_outcome: bool = False) -> str:
"""Mark the call started on *server* (doubles may lack ``mark_tool_call``), run coroutine function *call*
on the MCP loop and, on failure, walk ``recoverers`` (``(server_name, exc, retry_call, op) -> Optional[str]``,
None = not its kind; order matters). Unrecovered exceptions go through ``on_final_failure`` and become the
generic call-failed error. ``record_outcome`` applies breaker bookkeeping to the FIRST attempt only."""
if callable(getattr(server, "mark_tool_call", None)):
server.mark_tool_call()
def call_once():
return _core._run_on_mcp_loop(call, timeout=tool_timeout)
try:
result = call_once()
return _record_call_outcome(server_name, result) if record_outcome else result
@@ -255,14 +237,6 @@ def _invoke_with_recovery(server_name: str, call_once: Callable[[], str], op: st
return tool_error(_sanitize_error(f"MCP call failed: {type(exc).__name__}: {_exc_str(exc)}"))
# ------------------------------------------------------------- the RPC itself
def _mark_server_call_started(server: Any) -> None:
"""Record a user-visible MCP operation when the server supports it."""
if callable(getattr(server, "mark_tool_call", None)):
server.mark_tool_call()
@asynccontextmanager
async def _track_inflight_rpc(server: Any, server_name: str, op: str):
"""Register the running RPC so teardown can fail it fast. A deliberate teardown
@@ -276,8 +250,8 @@ async def _track_inflight_rpc(server: Any, server_name: str, op: str):
yield
except asyncio.CancelledError:
if getattr(server, "_reconnecting", False):
raise RuntimeError(f"MCP {op} on '{server_name}' was aborted by a reconnect "
f"teardown; retry the request on the rebuilt session") from None
raise RuntimeError(f"MCP {op} on '{server_name}' was aborted by a reconnect teardown; retry the "
f"request on the rebuilt session") from None
raise
finally:
if tracked:
@@ -286,16 +260,15 @@ async def _track_inflight_rpc(server: Any, server_name: str, op: str):
async def _call_tool_racing_stdio_death(server, server_name: str, tool_name: str, args: dict):
"""``session.call_tool`` that fails fast when the stdio child is/gets dead: pre-call (a dead
child must not hold the slot for the full timeout; ``server.session`` is stale) and mid-call
(race against ``_watch_stdio_children``). Both raise :class:`_StdioChildExited` for the
respawn path, which owns the reconnect signal. callable()/``is True`` because MagicMock
attributes are truthy."""
child must not hold the slot for the full timeout) and mid-call (race against
``_watch_stdio_children``). Both raise :class:`_StdioChildExited` for the respawn path, which
owns the reconnect signal. callable()/``is True`` because MagicMock attributes are truthy."""
_stdio_dead = getattr(server, "_stdio_children_dead", None)
if callable(_stdio_dead) and _stdio_dead() is True:
raise _StdioChildExited(f"MCP stdio subprocess for '{server_name}' had already exited when the call was dispatched")
_call_coro = server.session.call_tool(tool_name, arguments=args)
_watch_children = getattr(server, "_watch_stdio_children", None)
if not (_watch_children is not None and inspect.iscoroutinefunction(_watch_children) and asyncio.iscoroutine(_call_coro)):
if not (inspect.iscoroutinefunction(_watch_children) and asyncio.iscoroutine(_call_coro)):
# Stubbed sessions return a non-awaitable, or there is no child-watcher to race: plain await.
return await _call_coro if asyncio.iscoroutine(_call_coro) else _call_coro
rpc_task = asyncio.ensure_future(_call_coro)
@@ -340,9 +313,8 @@ def _render_content_blocks(result, server_name: str) -> Tuple[str, int]:
parts.append(rendered)
usable_parts += 1
continue
# Benign empty renders log at debug; warn only for unknown shapes.
block_type = getattr(block, "type", None) or type(block).__name__
if block_type in {"text", "resource", "audio", "image"}:
if block_type in {"text", "resource", "audio", "image"}: # benign empty render
logger.debug("MCP %s: content block type %r rendered empty", server_name, block_type)
else:
logger.warning("MCP %s: dropping unsupported content block type %r", server_name, block_type)
@@ -380,8 +352,7 @@ def _render_call_tool_result(result, server_name: str) -> str:
structured = None # drop notices do not count as usable content
if structured is None and meta is None:
return json.dumps({"result": text_result}, ensure_ascii=False)
# Key order is part of the output: "result" leads when there is text, otherwise "_meta"
# precedes the (empty) "result".
# Key order is part of the output: "result" leads when there is text, otherwise "_meta" precedes it.
payload: Dict[str, Any] = {"result": text_result} if text_result else {}
if structured is not None:
payload["structuredContent" if text_result else "result"] = structured
@@ -390,20 +361,16 @@ def _render_call_tool_result(result, server_name: str) -> str:
payload.setdefault("result", text_result)
try:
return json.dumps(payload, ensure_ascii=False)
except (TypeError, ValueError):
# Non-serializable metadata: drop the extras, keep the call.
except (TypeError, ValueError): # Non-serializable metadata: drop the extras, keep the call.
return json.dumps({"result": text_result}, ensure_ascii=False)
# ------------------------------------------------------------------- handlers
def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float):
"""Sync registry handler (``handler(args_dict, **kwargs) -> str``) calling an MCP tool via the background loop."""
op = f"tools/call {tool_name}"
def _handler(args: dict, **kwargs) -> str:
# Security boundary: untrusted-server write tools need approval before ANY transport work
# (including the lazy first-use spawn).
# Security boundary: untrusted-server write tools need approval before ANY transport work (incl. lazy spawn).
error = _trust_gate_check(server_name, tool_name) or _check_circuit_breaker(server_name)
if error is not None:
return error
@@ -412,67 +379,57 @@ def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float):
return error
async def _call():
_mark_server_call_started(server)
async with server._rpc_lock, _track_inflight_rpc(server, server_name, op):
# Snapshot contextvars for the elicitation callback (MCP recv loop doesn't inherit them).
server._pending_call_context = contextvars.copy_context()
server._pending_call_context = contextvars.copy_context() # for the elicitation callback
try:
result = await _call_tool_racing_stdio_death(server, server_name, tool_name, args)
finally:
server._pending_call_context = None
# Round-trip completed: transport is healthy even if the tool returned isError.
if getattr(server, "_mark_session_proven", None) is not None:
if getattr(server, "_mark_session_proven", None) is not None: # round-trip done: transport healthy
server._mark_session_proven()
return _render_call_tool_result(result, server_name)
def _on_failure(exc):
_core._bump_server_error(server_name)
logger.error("MCP tool %s/%s call failed: %s", server_name, tool_name, exc)
return _invoke_with_recovery(
server_name, lambda: _core._run_on_mcp_loop(_call, timeout=tool_timeout), op,
return _dispatch(
server_name, server, op, _call, tool_timeout,
(_handle_stdio_child_exited_and_retry, _handle_auth_error_and_retry, _handle_session_expired_and_retry),
_on_failure, record_outcome=True)
return _handler
def _make_utility_handler(server_name: str, tool_timeout: float, op: str, log_label: str,
rpc, render, required: Optional[str] = None):
"""Shared shape of the four utility handlers: ``rpc(session, args, server_name)`` awaited
under ``_rpc_lock``, ``render(result, server_name)`` -> JSON-able payload, ``required``
validated before any transport work; owns the connected check and recovery ladder."""
def _make_utility_handler(op: str, log_label: str, rpc, render, required: Optional[str] = None):
"""``(server_name, tool_timeout) -> sync handler`` for one utility tool: ``rpc(session, args,
server_name)`` awaited under ``_rpc_lock``, ``render(result, server_name)`` -> JSON-able
payload, ``required`` validated before any transport work."""
def _factory(server_name: str, tool_timeout: float):
def _handler(args: dict, **kwargs) -> str:
server = _core._get_connected_server_for_call(server_name)
if not server or not server.session:
return tool_error(f"MCP server '{server_name}' is not connected")
if required and not args.get(required):
return tool_error(f"Missing required parameter '{required}'")
def _handler(args: dict, **kwargs) -> str:
server = _core._get_connected_server_for_call(server_name)
if not server or not server.session:
return tool_error(f"MCP server '{server_name}' is not connected")
if required and not args.get(required):
return tool_error(f"Missing required parameter '{required}'")
async def _call():
_mark_server_call_started(server)
async with server._rpc_lock:
result = await rpc(server.session, args, server_name)
return json.dumps(render(result, server_name), ensure_ascii=False)
return _invoke_with_recovery(
server_name, lambda: _core._run_on_mcp_loop(_call, timeout=tool_timeout), op,
(_handle_auth_error_and_retry, _handle_session_expired_and_retry),
lambda exc: logger.error("MCP %s/%s failed: %s", server_name, log_label, exc))
return _handler
async def _call():
async with server._rpc_lock:
result = await rpc(server.session, args, server_name)
return json.dumps(render(result, server_name), ensure_ascii=False)
return _dispatch(
server_name, server, op, _call, tool_timeout,
(_handle_auth_error_and_retry, _handle_session_expired_and_retry),
lambda exc: logger.error("MCP %s/%s failed: %s", server_name, log_label, exc))
return _handler
return _factory
def _pick(obj, *specs) -> dict:
"""``{out_key: value}`` for each ``(out_key, attr[, truthy])`` present on *obj* (``hasattr``
so SDK models and stubs behave alike; ``truthy`` also skips falsy). Key order = spec order."""
"""``{out_key: value}`` for each ``(out_key, attr[, truthy])`` present on *obj* (presence check so SDK models
and stubs behave alike; ``truthy`` also skips falsy). Key order = spec order."""
entry = {}
for out_key, attr, *truthy in specs:
if not hasattr(obj, attr):
continue
value = getattr(obj, attr)
if value or not (truthy and truthy[0]):
value = getattr(obj, attr, _MISSING)
if value is not _MISSING and (value or not (truthy and truthy[0])):
entry[out_key] = value
return entry
@@ -495,8 +452,7 @@ def _render_read_resource(result, server_name: str) -> dict:
for block in getattr(result, "contents", []):
if getattr(block, "text", None) is not None:
parts.append(strip_unicode_tags(block.text))
elif getattr(block, "blob", None) is not None:
# Binary contents go to the document cache (same contract as EmbeddedResource blocks).
elif getattr(block, "blob", None) is not None: # binary -> document cache, like EmbeddedResource blocks
rendered = _render_mcp_resource_block(SimpleNamespace(type="resource", resource=block), server_name)
parts.append(rendered or f"[binary data, {len(block.blob)} bytes]")
return {"result": "\n".join(parts)}
@@ -523,41 +479,26 @@ def _render_get_prompt(result, server_name: str) -> dict:
return {"messages": messages, **_pick(result, ("description", "description", True))}
def _utility_factory(op: str, log_label: str, rpc, render, required: Optional[str] = None):
"""``(server_name, tool_timeout) -> sync handler`` for one utility tool."""
def _factory(server_name: str, tool_timeout: float):
return _make_utility_handler(server_name, tool_timeout, op, log_label, rpc, render, required)
return _factory
_make_list_resources_handler = _utility_factory(
_make_list_resources_handler = _make_utility_handler(
"resources/list", "list_resources",
lambda session, args, sn: _core._paginate_full_list(session.list_resources, "resources", sn),
_render_resource_list)
_make_read_resource_handler = _utility_factory(
lambda session, args, sn: _core._paginate_full_list(session.list_resources, "resources", sn), _render_resource_list)
_make_read_resource_handler = _make_utility_handler(
"resources/read", "read_resource",
lambda session, args, sn: session.read_resource(args["uri"]), _render_read_resource, required="uri")
_make_list_prompts_handler = _utility_factory(
_make_list_prompts_handler = _make_utility_handler(
"prompts/list", "list_prompts",
lambda session, args, sn: _core._paginate_full_list(session.list_prompts, "prompts", sn),
_render_prompt_list)
_make_get_prompt_handler = _utility_factory(
lambda session, args, sn: _core._paginate_full_list(session.list_prompts, "prompts", sn), _render_prompt_list)
_make_get_prompt_handler = _make_utility_handler(
"prompts/get", "get_prompt",
lambda session, args, sn: session.get_prompt(args["name"], arguments=args.get("arguments", {})),
_render_get_prompt, required="name")
def _make_check_fn(server_name: str):
"""Check function that verifies the MCP connection is alive."""
"""Connection-alive check; lazy (schema-cache registered) servers count as available."""
def _check() -> bool:
with _core._lock:
server = _core._servers.get(server_name)
if server is not None and (server.session is not None or server._is_recycled_stdio()):
return True
# Lazy (schema-cache registered) servers count as available: the first real
# call spawns/connects them.
return server_name in _core._lazy_server_configs
return ((server is not None and (server.session is not None or server._is_recycled_stdio()))
or server_name in _core._lazy_server_configs)
return _check
+19 -30
View File
@@ -1,6 +1,6 @@
"""Session health for MCPServerTask: dynamic tool refresh on list_changed notifications, server
log forwarding, keepalive probes, suspect-mark / lazy-verify, in-flight call fail-fast, stdio
child liveness and stdio idle/lifetime recycling. Split from tools/mcp_tool.py."""
child liveness and stdio idle/lifetime recycling."""
import asyncio
import json
@@ -38,8 +38,7 @@ class MCPServerHealthMixin:
self._recycled_reason = None
def _stdio_recycle_deadlines(self):
"""``[(deadline, reason), ...]`` for the configured lifetime/idle limits; empty for HTTP
servers or while an RPC holds the lock."""
"""``[(deadline, reason), ...]`` for the lifetime/idle limits; empty for HTTP or while an RPC holds the lock."""
if self._is_http() or self._rpc_lock.locked():
return []
limits = ((self._lifecycle_started_at, self._max_lifetime_seconds, "max_lifetime_seconds"),
@@ -52,8 +51,7 @@ class MCPServerHealthMixin:
return next((reason for deadline, reason in self._stdio_recycle_deadlines() if now >= deadline), None)
def _next_stdio_recycle_deadline(self) -> Optional[float]:
deadlines = self._stdio_recycle_deadlines()
return min(d for d, _ in deadlines) if deadlines else None
return min((d for d, _ in self._stdio_recycle_deadlines()), default=None)
def _mark_stdio_recycled(self, reason: str) -> None:
"""Mark a stdio session dormant before its transport finishes closing."""
@@ -67,15 +65,13 @@ class MCPServerHealthMixin:
await self._refresh_tools()
except Exception:
logger.exception("MCP server '%s': dynamic tool refresh failed", self.name)
task = asyncio.create_task(_run())
self._pending_refresh_tasks.add(task)
task.add_done_callback(self._pending_refresh_tasks.discard)
return task
def _make_logging_callback(self):
"""``logging_callback`` forwarding server ``notifications/message`` into Hermes logging
tagged with the server name (the SDK default drops them)."""
"""``logging_callback`` forwarding server ``notifications/message`` into Hermes logging (SDK default drops them)."""
async def _on_log(params):
try:
level = _core._MCP_LOG_LEVEL_MAP.get(str(getattr(params, "level", "info")).lower(), logging.INFO)
@@ -95,8 +91,7 @@ class MCPServerHealthMixin:
return _on_log
def _make_message_handler(self):
"""``message_handler`` for ``ClientSession``: only ``ToolListChangedNotification`` triggers
a refresh; prompt/resource changes are logged."""
"""``message_handler``: only ``ToolListChangedNotification`` triggers a refresh; prompt/resource changes log."""
async def _handler(message):
try:
if isinstance(message, Exception):
@@ -122,20 +117,16 @@ class MCPServerHealthMixin:
return _handler
def _deregister_owned(self, tool_names: Iterable[str]) -> None:
"""Deregister *tool_names* this server's toolset still owns. Never removes a colliding
name currently owned by another server."""
"""Deregister *tool_names* this server's toolset still owns (never a colliding name owned by another server)."""
from tools.registry import registry
toolset_name = f"mcp-{self.name}"
for tool_name in tool_names:
if registry.get_toolset_for_tool(tool_name) != toolset_name:
continue
registry.deregister(tool_name, scope=_core._server_registry_scope(self.name))
_forget_mcp_tool_server(tool_name)
if registry.get_toolset_for_tool(tool_name) == f"mcp-{self.name}":
registry.deregister(tool_name, scope=_core._server_registry_scope(self.name))
_forget_mcp_tool_server(tool_name)
async def _refresh_tools(self):
"""Re-fetch tools on ``tools/list_changed`` and update the registry. The lock serializes
rapid-fire notifications; after the list_tools ``await`` all mutations are synchronous —
atomic on the event loop."""
"""Re-fetch tools on ``tools/list_changed`` and update the registry. The lock serializes rapid-fire
notifications; after the list_tools ``await`` all mutations are synchronous — atomic on the event loop."""
if not self._advertises_tools():
return # tools/list would raise MCPError(-32601)
async with self._refresh_lock:
@@ -165,6 +156,8 @@ class MCPServerHealthMixin:
"""Exercise the session; raise on a genuine connection failure. ``ping`` first (cheap,
OPTIONAL); on -32601 latch ``_ping_unsupported`` (reset per transport connection) and fall
back to ``list_tools`` when the server advertises tools, else the -32601 propagates."""
async def list_tools():
await asyncio.wait_for(self.session.list_tools(), timeout=_KEEPALIVE_RPC_TIMEOUT)
if not self._ping_unsupported:
try:
await asyncio.wait_for(self.session.send_ping(), timeout=_KEEPALIVE_RPC_TIMEOUT)
@@ -180,7 +173,7 @@ class MCPServerHealthMixin:
# A server that silently drops ping looks like a dead transport: confirm with
# list_tools before declaring it dead, else propagate the original failure.
try:
await asyncio.wait_for(self.session.list_tools(), timeout=_KEEPALIVE_RPC_TIMEOUT)
await list_tools()
except Exception:
raise exc from None
self._ping_unsupported = True # latch so later keepalives skip the 30s wait
@@ -189,7 +182,7 @@ class MCPServerHealthMixin:
return
else:
raise # closed transport, expired session, etc. — real failure
await asyncio.wait_for(self.session.list_tools(), timeout=_KEEPALIVE_RPC_TIMEOUT)
await list_tools()
def _mark_session_proven(self) -> None:
"""Record that the session demonstrated real health (keepalive or tool-call success).
@@ -204,12 +197,10 @@ class MCPServerHealthMixin:
logger.warning("MCP server '%s': revived — session healthy again after "
"parking (state: parked → connected)", self.name)
# A proven fresh transport clears the one-time permanent-failure grace and any race bookkeeping.
self._permanent_grace_used = False
self._teardown_race = False
self._permanent_grace_used = self._teardown_race = False
def mark_suspect(self, reason: str) -> None:
"""Latch a suspicion (no I/O). The NEXT call verifies via :meth:`ensure_healthy` and
recycles the transport if the probe fails."""
"""Latch a suspicion (no I/O); the NEXT call verifies via :meth:`ensure_healthy` and recycles on failure."""
if self._suspect_reason is None and reason:
logger.warning("MCP server '%s': connection marked suspect (%s); next call will health-check it",
self.name, reason)
@@ -253,8 +244,7 @@ class MCPServerHealthMixin:
victims = [t for t in self._inflight_tasks if not t.done()]
if not victims:
return
self._reconnecting = True
self._teardown_race = True
self._reconnecting = self._teardown_race = True
self.mark_suspect(f"{reason} tore down {len(victims)} in-flight call(s)")
for task in victims:
task.cancel()
@@ -272,7 +262,6 @@ class MCPServerHealthMixin:
return False
async def _watch_stdio_children(self) -> None:
"""Poll child liveness while a stdio RPC is in flight; resolves when a tracked child dies
so the caller cancels the RPC instead of waiting out the timeout."""
"""Poll child liveness during a stdio RPC; resolves when a tracked child dies so the caller cancels the RPC."""
while not self._stdio_children_dead():
await asyncio.sleep(0.25)
+52 -85
View File
@@ -1,11 +1,11 @@
"""Registering a connected (or schema-cached) MCP server's tools into the tool registry:
include/exclude filtering, trust-tier metadata capture, utility-tool selection,
name-collision resolution and the schema-cache write-through. Both entry points
(``_register_server_tools`` live, ``_register_from_cache_sync`` lazy) build ``_Candidate``
records and feed the single ``_register_candidates`` loop."""
include/exclude filtering, trust-tier metadata capture, utility-tool selection, name-collision
resolution and the schema-cache write-through. Both entry points (``_register_server_tools``
live, ``_register_from_cache_sync`` lazy) build ``_Candidate`` records for ``_register_candidates``."""
import logging
from dataclasses import dataclass
from types import SimpleNamespace
from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Optional
from tools.mcp_tool_common import _parse_boolish, _core, _resolve_tool_timeout
from tools.mcp_tool_handlers import (
@@ -27,8 +27,7 @@ _UTILITY_HANDLER_FACTORIES = {
def _normalize_server_trust(value: Any) -> str:
"""Config ``trust`` -> tier. None -> ``full`` (backward-compatible default); an
unrecognized string -> ``untrusted`` so a misspelled tier fails closed."""
"""Config ``trust`` -> tier. None -> ``full`` (compat default); unrecognized -> ``untrusted`` (fail closed)."""
if value is None:
return _core._TRUST_FULL
text = str(value).strip().lower()
@@ -40,25 +39,19 @@ def _normalize_server_trust(value: Any) -> str:
def _annotation_read_only_hint(mcp_tool: Any) -> bool:
"""True only when annotations (SDK object or schema-cache dict) carry ``readOnlyHint is
True``; unknown metadata means write-capable."""
"""True only when annotations (SDK object or cache dict) carry ``readOnlyHint is True``; unknown = write-capable."""
annotations = getattr(mcp_tool, "annotations", None)
if isinstance(annotations, dict):
return annotations.get("readOnlyHint") is True
return getattr(annotations, "readOnlyHint", None) is True
hint = annotations.get("readOnlyHint") if isinstance(annotations, dict) else getattr(annotations, "readOnlyHint", None)
return hint is True
def _record_tool_trust_metadata(server_name: str, config: dict, tools: List[Any]) -> None:
"""Capture per-server trust and per-tool readOnlyHint at discovery — the security
boundary: the call-time gate classifies from data we control, never re-read
server-supplied state."""
"""Capture per-server trust and per-tool readOnlyHint at discovery — the security boundary: the call-time gate
classifies from data we control, never re-read server-supplied state."""
with _core._lock:
_core._server_trust_levels[server_name] = _normalize_server_trust((config or {}).get("trust"))
hints = _core._tool_read_only_hints.setdefault(server_name, {})
for tool in tools:
name = getattr(tool, "name", None)
if name:
hints[name] = _annotation_read_only_hint(tool)
hints.update({t.name: _annotation_read_only_hint(t) for t in tools if getattr(t, "name", None)})
def _track_mcp_tool_server(tool_name: str, server_name: str) -> None:
@@ -80,8 +73,7 @@ def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: d
filters anything since ClientSession defines all four methods."""
tools_filter = config.get("tools") or {}
enabled = {f: _parse_boolish(tools_filter.get(f), default=True) for f in ("resources", "prompts")}
init_result = getattr(server, "initialize_result", None)
advertised = getattr(init_result, "capabilities", None) if init_result is not None else None
advertised = getattr(getattr(server, "initialize_result", None), "capabilities", None)
def _skip_reason(handler_key: str) -> Optional[str]:
family = _UTILITY_CAPABILITY_ATTRS[handler_key]
@@ -93,7 +85,6 @@ def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: d
return None
# Legacy gate (no initialize_result): the ClientSession method shares the handler key.
return None if hasattr(server.session, handler_key) else f"session lacks {handler_key}"
selected: List[dict] = []
for entry in _build_utility_schemas(server_name):
reason = _skip_reason(entry["handler_key"])
@@ -105,8 +96,7 @@ def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: d
def _existing_tool_names() -> List[str]:
"""Tool names for all currently connected servers plus lazy (cache-registered) servers,
whose tools live only in the registry."""
"""Tool names for all connected servers plus lazy (cache-registered) servers, whose tools live only in the registry."""
names: List[str] = []
for server in _core._servers.values():
names.extend(server._registered_tool_names if hasattr(server, "_registered_tool_names")
@@ -118,9 +108,8 @@ def _existing_tool_names() -> List[str]:
def _make_tool_filter(name: str, config: dict) -> Callable[[str], bool]:
"""Include/exclude predicate for a server's tool names: ``tools.include`` is a whitelist
(``[]`` = register nothing), ``tools.exclude`` a blacklist; entries are exact names or
fnmatch globs; include wins over exclude."""
"""Include/exclude predicate for a server's tool names: ``tools.include`` is a whitelist (``[]`` = register
nothing), ``tools.exclude`` a blacklist; entries are exact names or fnmatch globs; include wins over exclude."""
tools_filter = config.get("tools") or {}
include_raw = tools_filter.get("include")
include_set = _normalize_name_filter(include_raw, f"mcp_servers.{name}.tools.include")
@@ -130,30 +119,18 @@ def _make_tool_filter(name: str, config: dict) -> Callable[[str], bool]:
return lambda tool_name: not (exclude_set and matches_name_filter(tool_name, exclude_set))
class _CachedMCPTool:
"""Stand-in for MCP Tool objects loaded from the schema cache. Missing or non-dict
``annotations`` (older cache files) fail closed to write-capable."""
__slots__ = ("name", "description", "inputSchema", "annotations")
def __init__(self, name: str, description: str, inputSchema: dict, annotations: Optional[dict] = None):
self.name = name
self.description = description
self.inputSchema = inputSchema or {}
self.annotations = annotations if isinstance(annotations, dict) else None
@classmethod
def from_cache_dicts(cls, raws: Iterable[Any]) -> List["_CachedMCPTool"]:
"""Cached rows -> stand-ins; rows that are not dicts or lack a name are dropped."""
return [cls(raw["name"], raw.get("description") or "",
raw["inputSchema"] if isinstance(raw.get("inputSchema"), dict) else {}, raw.get("annotations"))
for raw in raws if isinstance(raw, dict) and raw.get("name")]
def _cached_tools(raws: Iterable[Any]) -> List[SimpleNamespace]:
"""Schema-cache rows -> stand-ins for MCP Tool objects; rows that are not dicts or lack a name
are dropped. Missing or non-dict ``annotations`` (older cache files) fail closed to write-capable."""
return [SimpleNamespace(name=raw["name"], description=raw.get("description") or "",
inputSchema=raw["inputSchema"] if isinstance(raw.get("inputSchema"), dict) else {},
annotations=raw["annotations"] if isinstance(raw.get("annotations"), dict) else None)
for raw in raws if isinstance(raw, dict) and raw.get("name")]
@dataclass
class _Candidate:
"""One registration attempt: a native tool or a generated utility. ``origin`` is the
provenance text used in collision diagnostics."""
"""One registration attempt (native tool or generated utility); ``origin`` is the provenance text in diagnostics."""
registry_name: str
origin: str
@@ -167,8 +144,8 @@ class _Candidate:
def _tool_candidates(name: str, tools: Iterable[Any], should_register: Callable[[str], bool],
tool_timeout) -> List[_Candidate]:
"""Native tools (live SDK objects or ``_CachedMCPTool``) -> candidates. The injection scan
runs on BOTH paths: the cache file is user-writable JSON."""
"""Native tools (live SDK objects or cache stand-ins) -> candidates. The injection scan runs on
BOTH paths: the cache file is user-writable JSON."""
out: List[_Candidate] = []
for t in tools:
if not should_register(t.name):
@@ -187,8 +164,8 @@ def _utility_candidates(name: str, entries: Iterable[Any], tool_timeout) -> List
for raw in entries:
schema, key = (raw.get("schema"), raw.get("handler_key")) if isinstance(raw, dict) else (None, None)
if isinstance(schema, dict) and key in _UTILITY_HANDLER_FACTORIES and schema.get("name"):
handler = _UTILITY_HANDLER_FACTORIES[key](name, tool_timeout)
out.append(_Candidate(schema["name"], f"{_UTILITY_ORIGIN_PREFIX}{key!r}", schema, handler))
out.append(_Candidate(schema["name"], f"{_UTILITY_ORIGIN_PREFIX}{key!r}", schema,
_UTILITY_HANDLER_FACTORIES[key](name, tool_timeout)))
return out
@@ -197,16 +174,15 @@ def _resolve_name_collisions(name: str, candidates: List[_Candidate]) -> List[_C
a native tool's name is shadowed (native wins); any other multi-origin collision skips every
colliding entry (fail closed). Returns survivors in order."""
unique: List[_Candidate] = []
seen: set[tuple[str, str]] = set()
origins_by_name: Dict[str, set[str]] = {}
for c in candidates:
if (c.registry_name, c.origin) in seen:
origins = origins_by_name.setdefault(c.registry_name, set())
if c.origin in origins:
logger.debug("MCP server '%s': duplicate registration candidate %s for '%s'; keeping one",
name, c.origin, c.registry_name)
continue
seen.add((c.registry_name, c.origin))
origins.add(c.origin)
unique.append(c)
origins_by_name.setdefault(c.registry_name, set()).add(c.origin)
ambiguous: Dict[str, List[str]] = {}
shadowed: set[tuple[str, str]] = set()
for registry_name, origins in origins_by_name.items():
@@ -220,29 +196,14 @@ def _resolve_name_collisions(name: str, candidates: List[_Candidate]) -> List[_C
"MCP server '%s': generated utility %s normalizes onto server-native %s — keeping the native tool "
"and dropping the utility (the utility only applies when the server has no such tool of its own)",
name, ", ".join(utility_origins), native_origins[0])
continue
ambiguous[registry_name] = sorted(origins)
else:
ambiguous[registry_name] = sorted(origins)
for registry_name, origins in sorted(ambiguous.items()):
logger.error("MCP server '%s': name normalization collision for '%s' from %s; skipping every colliding "
"entry instead of choosing an arbitrary handler", name, registry_name, ", ".join(origins))
return [c for c in unique if c.registry_name not in ambiguous and (c.registry_name, c.origin) not in shadowed]
def _log_foreign_owner(name: str, c: _Candidate, existing_toolset: str, lazy: bool) -> None:
"""Diagnostics for a name already owned by another toolset (skipped to preserve the owner)."""
if lazy:
if not c.is_utility:
logger.warning("MCP server '%s' (lazy): cached tool '%s' collides with toolset '%s' — skipping",
name, c.registry_name, existing_toolset)
return
if existing_toolset.startswith("mcp-"):
logger.error("MCP server '%s': %s normalizes to '%s', already owned by MCP toolset '%s' — skipping to "
"preserve the existing owner", name, c.origin, c.registry_name, existing_toolset)
else:
logger.warning("MCP server '%s': %s (→ '%s') collides with built-in tool in toolset '%s' — skipping to "
"preserve built-in", name, c.origin, c.registry_name, existing_toolset)
def _register_candidates(name: str, candidates: List[_Candidate], *, check_fn: Callable,
scope: Callable[[], Optional[str]], lazy: bool) -> List[str]:
"""Register candidates under toolset ``mcp-{name}``; returns the names that landed. The
@@ -253,34 +214,40 @@ def _register_candidates(name: str, candidates: List[_Candidate], *, check_fn: C
registered: List[str] = []
for c in candidates:
existing_toolset = registry.get_toolset_for_tool(c.registry_name)
if existing_toolset and existing_toolset != toolset_name:
_log_foreign_owner(name, c, existing_toolset, lazy)
if existing_toolset and existing_toolset != toolset_name: # foreign owner: skip, preserve it
if lazy:
if not c.is_utility:
logger.warning("MCP server '%s' (lazy): cached tool '%s' collides with toolset '%s' — skipping",
name, c.registry_name, existing_toolset)
elif existing_toolset.startswith("mcp-"):
logger.error("MCP server '%s': %s normalizes to '%s', already owned by MCP toolset '%s' — skipping to "
"preserve the existing owner", name, c.origin, c.registry_name, existing_toolset)
else:
logger.warning("MCP server '%s': %s (→ '%s') collides with built-in tool in toolset '%s' — skipping to "
"preserve built-in", name, c.origin, c.registry_name, existing_toolset)
continue
registry.register(
name=c.registry_name, toolset=toolset_name, schema=c.schema, handler=c.handler, check_fn=check_fn,
is_async=False, description=c.schema.get("description") or "", scope=scope())
if registry.get_toolset_for_tool(c.registry_name) != toolset_name:
if not lazy:
logger.error("MCP server '%s': registration of %s as '%s' was rejected by the registry; "
"skipping provenance/count updates", name, c.origin, c.registry_name)
continue
_core._track_mcp_tool_server(c.registry_name, name)
registered.append(c.registry_name)
if registry.get_toolset_for_tool(c.registry_name) == toolset_name:
_core._track_mcp_tool_server(c.registry_name, name)
registered.append(c.registry_name)
elif not lazy:
logger.error("MCP server '%s': registration of %s as '%s' was rejected by the registry; "
"skipping provenance/count updates", name, c.origin, c.registry_name)
if registered:
registry.register_toolset_alias(name, toolset_name)
return registered
def _write_schema_cache(name: str, server: "MCPServerTask", config: dict, should_register) -> None:
"""Write-through: persist the manifest so the next startup can register this server
lazily without spawning it. Never raises."""
"""Write-through: persist the manifest so the next startup registers this server lazily (no spawn). Never raises."""
try:
from tools.mcp_schema_cache import config_fingerprint, write_cache_entry
tools_payload = [{
"name": t.name, "description": t.description or "",
"inputSchema": t.inputSchema if isinstance(getattr(t, "inputSchema", None), dict) else {},
# Persisted so the lazy path trust-gates identically next startup.
"annotations": {"readOnlyHint": _annotation_read_only_hint(t)},
"annotations": {"readOnlyHint": _annotation_read_only_hint(t)}, # lazy path trust-gates identically
} for t in server._tools if should_register(t.name)]
utility_payload = [{"schema": e["schema"], "handler_key": e["handler_key"]}
for e in _select_utility_schemas(name, server, config)]
@@ -313,7 +280,7 @@ def _register_from_cache_sync(name: str, config: dict, entry: dict) -> List[str]
call-time gate is identical for live and cached registrations."""
from tools.mcp_schema_cache import config_fingerprint, tools_from_cache_entry, utility_tools_from_cache_entry
tool_timeout = _resolve_tool_timeout(config)
cached_tools = _CachedMCPTool.from_cache_dicts(tools_from_cache_entry(entry))
cached_tools = _cached_tools(tools_from_cache_entry(entry))
_record_tool_trust_metadata(name, config, cached_tools)
candidates = _tool_candidates(name, cached_tools, _make_tool_filter(name, config), tool_timeout)
candidates += _utility_candidates(name, utility_tools_from_cache_entry(entry), tool_timeout)
+81 -105
View File
@@ -13,8 +13,7 @@ from tools.mcp_tool_common import _core
logger = logging.getLogger("tools.mcp_tool")
# JSON-RPC ``initialize`` body used by the content-type preflight POST.
_PROBE_INITIALIZE_BODY = (
_PROBE_INITIALIZE_BODY = ( # JSON-RPC ``initialize`` body for the content-type preflight POST
'{"jsonrpc":"2.0","id":"_probe","method":"initialize","params":{"protocolVersion":"2025-03-26",'
'"capabilities":{},"clientInfo":{"name":"hermes-probe","version":"0.1"}}}')
@@ -28,15 +27,17 @@ def _is_2xx(resp) -> bool:
return 200 <= resp.status_code < 300
def _present(**kwargs) -> dict:
"""*kwargs* minus the ``None`` values (optional httpx client arguments)."""
return {k: v for k, v in kwargs.items() if v is not None}
def _pgroup_alive(pgid: Optional[int]) -> bool:
"""Signal 0 to the group succeeds iff any member is alive (POSIX only)."""
_killpg = getattr(os, "killpg", None)
if pgid is None or _killpg is None:
return False
try:
_killpg(pgid, 0)
os.killpg(pgid, 0)
return True
except (ProcessLookupError, PermissionError, OSError):
except (AttributeError, TypeError, OSError): # non-POSIX / pgid None / gone
return False
@@ -46,16 +47,17 @@ class MCPServerTransportMixin:
__slots__ = ()
def _advertises_tools(self) -> bool:
"""Whether the server advertises ``tools`` (prompt-/resource-only servers omit it and
``tools/list`` raises -32601). True when no capability info was captured (legacy fallback)."""
"""False only when captured capabilities omit ``tools`` (prompt-/resource-only servers,
where ``tools/list`` raises -32601); True without capability info (legacy fallback)."""
caps = getattr(self.initialize_result, "capabilities", None)
return caps is None or getattr(caps, "tools", None) is not None
def _session_kwargs(self) -> dict:
"""ClientSession kwargs: sampling, elicitation, notification + logging callbacks."""
kwargs = self._sampling.session_kwargs() if self._sampling else {}
if self._elicitation:
kwargs.update(self._elicitation.session_kwargs())
kwargs = {}
for handler in (self._sampling, self._elicitation):
if handler:
kwargs.update(handler.session_kwargs())
if _core._MCP_NOTIFICATION_TYPES and _core._MCP_MESSAGE_HANDLER_SUPPORTED:
kwargs["message_handler"] = self._make_message_handler()
if _core._MCP_LOGGING_CALLBACK_SUPPORTED:
@@ -63,12 +65,11 @@ class MCPServerTransportMixin:
return kwargs
async def _negotiate_session(self, session, connect_timeout: float):
"""Negotiate the protocol era (``initialize`` vs ``server/discover``); both results expose
``.capabilities``. ``protocol: auto`` (default) tries the legacy handshake FIRST and falls back
to discover only on a modern-only signal (-32022 / initialize -32601) — deliberately the
reverse of the SDK's discover-first mode: zero extra round-trips for the handshake-era
servers that dominate today. ``stateless`` probes discover first (one legacy retry on any
error); ``legacy`` is handshake only. A handshake TIMEOUT never falls back — it propagates."""
"""Negotiate the protocol era (``initialize`` vs ``server/discover``; both expose
``.capabilities``). ``auto`` tries the legacy handshake FIRST, falling back to discover only
on a modern-only signal (-32022 / initialize -32601) — the reverse of the SDK's discover-first
mode, so handshake-era servers pay zero extra round-trips. ``stateless`` probes discover first
(one legacy retry on any error); ``legacy`` is handshake only. A TIMEOUT never falls back."""
def call(method: str):
return asyncio.wait_for(getattr(session, method)(), timeout=connect_timeout)
@@ -80,7 +81,6 @@ class MCPServerTransportMixin:
raise
logger.info(log_fmt, self.name, exc, *log_extra)
return await call(fallback)
mode = str((self._config or {}).get("protocol", "auto")).lower().strip()
if mode in ("stateless", "modern", "2026-07-28"):
return await attempt("discover", "initialize", lambda exc: True,
@@ -94,14 +94,13 @@ class MCPServerTransportMixin:
# mcp 1.x has no server/discover client — nothing to fall back to.
return await attempt(
"initialize", "discover", lambda exc: _handshake_rejected_as_modern(exc) and hasattr(session, "discover"),
"MCP server '%s': legacy handshake rejected (%s) — "
"retrying via server/discover (2026-07-28 stateless server)")
"MCP server '%s': legacy handshake rejected (%s) — retrying via server/discover (2026-07-28 stateless server)")
async def _serve_session(self, session, connect_timeout: float,
label: str = "", mark_lifecycle: bool = False) -> str:
"""Handshake, discover, publish readiness, then serve until a lifecycle event. Clears stale
breaker state but leaves the session UNPROVEN: flapping transports handshake fine and drop
moments later, so only keepalive or tool-call success clears the reconnect budget."""
moments later, so only keepalive/tool-call success clears the reconnect budget."""
self.initialize_result = await self._negotiate_session(session, connect_timeout)
self.session = session
if mark_lifecycle:
@@ -118,8 +117,7 @@ class MCPServerTransportMixin:
async def _serve_transport(self, transport_cm, label: str, connect_timeout: float) -> str:
"""Open *transport_cm*, wrap its streams in a ClientSession and serve it. Streams are indexed,
not unpacked: mcp 1.x yields ``(read, write, get_session_id)``, 2.x ``(read, write)``.
A transport TaskGroup drop maps to ``"reconnect"`` instead of backoff/park."""
not unpacked (mcp 1.x yields a 3-tuple, 2.x a pair); a TaskGroup drop maps to ``"reconnect"``."""
try:
async with transport_cm as _streams:
async with _core.ClientSession(_streams[0], _streams[1], **self._session_kwargs()) as session:
@@ -131,7 +129,7 @@ class MCPServerTransportMixin:
def _track_spawned_children(self, new_pids: Set[int]) -> None:
"""Ledger the freshly spawned stdio children (pids, pgids, machine spawn ledger). pgids are
captured while alive (getpgid fails once it exits; the sweep needs it for reparented descendants)."""
captured while alive (getpgid fails after exit; the sweep needs them for reparented descendants)."""
new_pgids: Dict[int, int] = {}
for pid in new_pids:
try:
@@ -171,8 +169,7 @@ class MCPServerTransportMixin:
with _core._lock:
for pid in new_pids:
_stdio_pids.pop(pid, None)
# ``os.kill(pid, 0)`` is NOT a no-op on Windows; the child may be gone while
# descendants remain in its pgroup.
# Windows-safe pid probe; the child may be gone while descendants remain in its pgroup.
if _pid_exists(pid) or _pgroup_alive(_stdio_pgids.get(pid)):
_orphan_stdio_pids.add(pid)
_orphan_stdio_pid_servers[pid] = self.name
@@ -184,8 +181,7 @@ class MCPServerTransportMixin:
async def _run_stdio(self, config: dict):
"""Run the server using stdio transport."""
if config.get("identity_header") is not None:
# No headers on stdio — warn so a copy-pasted HTTP block doesn't mislead.
if config.get("identity_header") is not None: # copy-pasted HTTP block: warn, don't mislead
logger.warning("MCP server '%s': identity_header is only supported on "
"HTTP/SSE transports — ignored for stdio servers", self.name)
if not _core._ensure_mcp_sdk():
@@ -201,9 +197,8 @@ class MCPServerTransportMixin:
command=command, args=args, env=safe_env or None, cwd=config.get("cwd"),
# Windows pipes can split non-UTF-8 bytes at chunk boundaries; substitute, don't raise.
encoding_error_handler="replace")
session_kwargs = self._session_kwargs()
# Reap orphans of prior attempts first, else each retry piles up zombie pairs. Unscoped on
# purpose (also reaps servers that never reconnect). Off-loop: the reaper blocks up to 2s.
# Reap orphans of prior attempts first (else retries pile up zombie pairs); unscoped on purpose;
# off-loop because the reaper blocks up to 2s.
await asyncio.to_thread(_core._kill_orphaned_mcp_children)
pids_before = _core._snapshot_child_pids() # so the new child can be identified after spawn
new_pids: set = set()
@@ -212,21 +207,18 @@ class MCPServerTransportMixin:
try:
errlog = _core._get_mcp_stderr_log()
async with _core.stdio_client(server_params, errlog=errlog) as (read_stream, write_stream):
# New PIDs for force-kill cleanup, minus non-MCP children (slash_worker, LSP) that
# race into the window: they share the TUI's pgid, so leaking them into _stdio_pgids
# would make the shutdown killpg() kill the TUI itself.
# New PIDs for force-kill cleanup, minus non-MCP children (slash_worker, LSP) racing
# into the window: they share the TUI's pgid — leaking them would killpg() the TUI.
new_pids = _filter_mcp_children(_core._snapshot_child_pids() - pids_before)
if new_pids:
self._track_spawned_children(new_pids)
self._stdio_child_pids = set(new_pids) # so in-flight calls fail fast when the child dies
async with _core.ClientSession(read_stream, write_stream, **session_kwargs) as session:
# Bound the handshake here: ``connect_timeout`` only bounds the caller's ``.result()``.
# A server that never answers ``initialize`` would otherwise hang forever, skip the
# ``finally`` and leak child + pipes on every retry until EMFILE.
async with _core.ClientSession(read_stream, write_stream, **self._session_kwargs()) as session:
# Bound the handshake here (``connect_timeout`` only bounds the caller's ``.result()``):
# a server that never answers ``initialize`` would leak child + pipes per retry until EMFILE.
connect_timeout = float(config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT))
return await self._serve_session(session, connect_timeout, mark_lifecycle=True)
finally:
# Runs on clean exit, exceptions AND cancellation.
finally: # clean exit, exceptions AND cancellation
if new_pids:
self._release_spawned_children(new_pids)
@@ -234,29 +226,30 @@ class MCPServerTransportMixin:
async def _preflight_content_type(self, url: str, *, headers: Optional[dict] = None,
ssl_verify: bool = True, client_cert=None, timeout: float = 5.0) -> None:
"""Probe *url* before the SDK connects: a plain web page makes the SDK sit out the full
"""Probe *url* before the SDK connects: a plain web page would make the SDK sit out the full
``connect_timeout`` before an opaque ``CancelledError``; this raises NonMcpEndpointError within
``timeout`` instead. Allow-list based: only a 2xx with a definite non-MCP content type is
rejected, and only after a JSON-RPC ``initialize`` POST also fails to look like MCP (some
servers serve a UI on GET but speak MCP via POST). Missing content type, non-2xx or transport
errors pass silently — the real handshake stays the source of truth. Own httpx client, OUTSIDE
the SDK's anyio task group, so the error isn't wrapped in an ExceptionGroup."""
``timeout``. Allow-list based: only a 2xx with a definite non-MCP content type is rejected, and
only after a JSON-RPC ``initialize`` POST also fails to look like MCP (some servers serve a UI
on GET but speak MCP via POST). Anything else passes — the handshake stays the source of truth.
Own httpx client, OUTSIDE the SDK's anyio task group, so the error isn't group-wrapped."""
try:
import httpx as _httpx
except ImportError:
return # No httpx → skip probe; SDK import would have failed first.
def _non_mcp_2xx(resp) -> bool:
# Only judge 2xx (4xx/5xx may be an auth challenge); no content type advertised → trust the SDK.
ct = _content_type_base(resp)
return _is_2xx(resp) and bool(ct) and ct not in self._MCP_CONTENT_TYPES
probe_headers = dict(headers) if headers else {}
try:
async with _httpx.AsyncClient(verify=ssl_verify, follow_redirects=True, timeout=_httpx.Timeout(timeout),
**({"cert": client_cert} if client_cert is not None else {})) as client:
# HEAD is cheapest; fall back to GET on 405/501.
resp = await client.head(url, headers=probe_headers)
**_present(cert=client_cert)) as client:
resp = await client.head(url, headers=probe_headers) # cheapest; GET on 405/501
if resp.status_code in (405, 501):
resp = await client.get(url, headers=probe_headers)
# Non-MCP content type on HEAD/GET: try a JSON-RPC POST so POST-only servers pass.
ct = _content_type_base(resp)
if ct and ct not in self._MCP_CONTENT_TYPES and _is_2xx(resp):
if _non_mcp_2xx(resp):
post_resp = await client.post(
url, content=_PROBE_INITIALIZE_BODY,
headers={**probe_headers, "Content-Type": "application/json",
@@ -265,25 +258,20 @@ class MCPServerTransportMixin:
resp = post_resp
except _httpx.HTTPError:
return # DNS/connect/timeout/transport error — let the SDK try.
# Only judge 2xx (4xx/5xx may be an auth challenge the handshake handles); no content type
# advertised → don't second-guess the SDK.
ct_base = _content_type_base(resp)
if not _is_2xx(resp) or not ct_base or ct_base in self._MCP_CONTENT_TYPES:
if not _non_mcp_2xx(resp):
return
raise NonMcpEndpointError(
f"MCP server '{self.name}' at {url} returned Content-Type '{ct_base}', not an MCP "
ct_base = _content_type_base(resp)
raise NonMcpEndpointError(f"MCP server '{self.name}' at {url} returned Content-Type '{ct_base}', not an MCP "
f"response (expected one of: {', '.join(self._MCP_CONTENT_TYPES)}). The URL most likely "
"points at a web page rather than an MCP endpoint — check it resolves to a Streamable "
"HTTP / SSE endpoint (e.g. https://host/mcp, not https://host/).")
def _reconnect_or_reraise_group(self, eg: BaseExceptionGroup) -> str:
"""Map an SDK transport TaskGroup failure to a clean ``"reconnect"``: HTTP/SSE stream pumps
run in an anyio TaskGroup, so a transient drop escapes as a ``BaseExceptionGroup`` that would
otherwise back off and park the server for 300s over a sub-second glitch. Re-raise when it is
not a transient drop: shutdown in progress (``_shutdown_event`` is set before cancel), the
group carries KeyboardInterrupt/SystemExit or a real CancelledError, or no live session was
reached this attempt (``_ready`` unset — connect failures must back off, not hot-loop)."""
"""Map an SDK transport TaskGroup failure to a clean ``"reconnect"``: HTTP/SSE stream pumps run in an anyio
TaskGroup, so a transient drop escapes as a ``BaseExceptionGroup`` that would otherwise park the server for
300s over a sub-second glitch. Re-raise when it is not one: shutdown in progress (``_shutdown_event`` is
set before cancel), KeyboardInterrupt/SystemExit or a real CancelledError in the group, or no live session
this attempt (``_ready`` unset — connect failures must back off, not hot-loop)."""
if (self._shutdown_event.is_set()
or eg.split((KeyboardInterrupt, SystemExit))[0] is not None
or eg.split(asyncio.CancelledError)[0] is not None
@@ -295,8 +283,7 @@ class MCPServerTransportMixin:
def _build_oauth_auth(self, url: str, config: dict):
"""OAuth 2.1 PKCE via the central MCPOAuthManager (one provider reused across reconnects and
CLI paths). Setup failures (e.g. non-interactive without cached tokens) re-raise so only this
server is reported failed."""
CLI paths). Setup failures re-raise (after a warning) so only this server is reported failed."""
if self._auth_type != "oauth":
return None
try:
@@ -309,28 +296,25 @@ class MCPServerTransportMixin:
def _sse_transport(self, url: str, headers: dict, connect_timeout: float,
ssl_verify, client_cert, oauth_auth, strict_cfg_headers: bool):
"""``sse_client`` context manager for ``transport: sse`` entries."""
if strict_cfg_headers:
# Fail closed: SSE cannot enforce the redirect boundary.
if strict_cfg_headers: # fail closed: SSE cannot enforce the redirect boundary
raise ValueError(f"MCP server '{self.name}': strict_redirect_headers is "
"not supported on the SSE transport.")
if _core.sse_client is None:
raise ImportError(f"MCP server '{self.name}' requires SSE transport but "
"mcp.client.sse.sse_client is not available. "
"Upgrade the mcp package to get SSE support.")
# sse_read_timeout bounds the gap between events: SSE servers idle for minutes, so 300s
# (matching the Streamable HTTP read timeout), not tool_timeout. ``auth`` must be forwarded
# or OAuth SSE servers 401 silently.
# sse_read_timeout bounds the gap between events: SSE servers idle for minutes, so 300s (the
# Streamable HTTP read timeout), not tool_timeout. ``auth`` must be forwarded or OAuth SSE 401s silently.
sse_kwargs: dict = {"url": url, "headers": headers or None, "timeout": float(connect_timeout),
"sse_read_timeout": 300.0, **({"auth": oauth_auth} if oauth_auth is not None else {})}
"sse_read_timeout": 300.0, **_present(auth=oauth_auth)}
if client_cert is not None or ssl_verify is not True:
# sse_client has no verify/cert kwargs: an httpx_client_factory forwards the SDK's
# (headers, auth, timeout) and layers TLS on top. The client MUST come from the SDK's
# own httpx module (httpx2 on mcp >= 2.0) — see sdk_httpx().
# sse_client has no verify/cert kwargs: an httpx_client_factory forwards the SDK's (headers,
# auth, timeout) and layers TLS on top. Client MUST come from the SDK's httpx (httpx2 on mcp >= 2.0).
_httpx_mod = _core.sdk_httpx()
sse_kwargs["httpx_client_factory"] = lambda headers=None, timeout=None, auth=None: _httpx_mod.AsyncClient(
follow_redirects=True, verify=ssl_verify,
timeout=timeout if timeout is not None else _httpx_mod.Timeout(30.0, read=300.0),
**{k: v for k, v in (("headers", headers), ("auth", auth), ("cert", client_cert)) if v is not None})
**_present(headers=headers, auth=auth, cert=client_cert))
return _core.sse_client(**sse_kwargs)
def _streamable_http_transport(self, url: str, headers: dict, connect_timeout: float,
@@ -339,30 +323,27 @@ class MCPServerTransportMixin:
"""Streamable HTTP context manager: mcp >= 1.24.0 gets a caller-owned httpx client; on the
deprecated API (mcp < 1.24.0) the SDK owns the client."""
if not _core._MCP_NEW_HTTP:
if strict_cfg_headers:
# Fail closed: without an owned client redirects can't be hooked.
if strict_cfg_headers: # fail closed: without an owned client redirects can't be hooked
raise ImportError(f"MCP server '{self.name}' requires mcp >= 1.24.0 to "
"enforce the portable redirect-header boundary "
"(strict_redirect_headers). Upgrade the mcp package.")
return _core.streamablehttp_client(url, headers=headers, timeout=float(connect_timeout), verify=ssl_verify,
**({"auth": oauth_auth} if oauth_auth is not None else {}))
# Explicit AsyncClient matching the SDK's create_mcp_http_client defaults; MUST come from
# the SDK's httpx module (httpx2 on mcp >= 2.0) since the SDK sends its own Requests through it.
**_present(auth=oauth_auth))
# Explicit AsyncClient matching the SDK's create_mcp_http_client defaults; MUST come from the
# SDK's httpx (httpx2 on mcp >= 2.0) since the SDK sends its own Requests through it.
httpx = _core.sdk_httpx()
_strip_auth_on_cross_origin_redirect = _make_redirect_header_stripper(
httpx.URL(url), strict=strict_cfg_headers, configured_header_names=configured_header_names)
client_kwargs: dict = {"follow_redirects": True, "timeout": httpx.Timeout(float(connect_timeout), read=300.0),
"verify": ssl_verify, **({"headers": headers} if headers else {}),
"event_hooks": {"response": [_strip_auth_on_cross_origin_redirect]},
**{k: v for k, v in (("auth", oauth_auth), ("cert", client_cert)) if v is not None}}
**_present(auth=oauth_auth, cert=client_cert)}
@asynccontextmanager
async def _owned_client_streams():
# Caller owns the client lifecycle — the SDK skips cleanup when http_client is provided.
async def _owned_client_streams(): # the SDK skips cleanup when http_client is provided
async with httpx.AsyncClient(**client_kwargs) as http_client:
async with _core.streamable_http_client(url, http_client=http_client) as streams:
yield streams
return _owned_client_streams()
async def _run_http(self, config: dict):
@@ -375,13 +356,11 @@ class MCPServerTransportMixin:
url = config["url"]
headers = dict(config.get("headers") or {})
# Agent Plugins v1 strict_redirect_headers: configured headers MUST NOT follow a cross-origin
# redirect. Capture their names BEFORE client-generated headers are merged in.
# redirect — capture their names BEFORE client-generated headers are merged in.
configured_header_names = {key.lower() for key in headers}
# Optional per-user identity header; explicit headers of the same name win.
headers = _apply_identity_header(self.name, config, headers)
# Some servers require MCP-Protocol-Version on the first request; seed it (user override
# wins) from the HANDSHAKE version, not the latest: a 2026-07-28 header would route the
# handshake-era ``initialize()`` body onto the per-request-envelope ladder, which rejects it.
headers = _apply_identity_header(self.name, config, headers) # explicit same-name headers win
# Seed MCP-Protocol-Version (user override wins) from the HANDSHAKE version, not the latest: a
# 2026-07-28 header routes the handshake-era ``initialize()`` onto the envelope ladder, which rejects it.
if not any(key.lower() == "mcp-protocol-version" for key in headers):
headers["mcp-protocol-version"] = _core.LATEST_HANDSHAKE_VERSION
connect_timeout = config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT)
@@ -399,8 +378,7 @@ class MCPServerTransportMixin:
async def _discover_tools(self):
"""Discover tools from the connected session. Capability-gated: prompt-/resource-only
servers raise ``MCPError(-32601)`` on ``tools/list``, which would abort the connection."""
# Fresh transport: re-probe ``ping`` in case the server gained support across the reconnect.
self._ping_unsupported = False
self._ping_unsupported = False # fresh transport: re-probe ``ping`` across the reconnect
if self.session is None:
return
if not self._advertises_tools():
@@ -415,19 +393,17 @@ class MCPServerTransportMixin:
self._register_discovered_tools_if_needed()
def _register_discovered_tools_if_needed(self) -> None:
"""Publish freshly discovered tools for a registry-owned server if none are registered
(initial registration normally happens in ``_discover_and_register_server``). On reconnect,
outage handling may clear ``_ready`` and deregister stale tools; ownership via ``_servers``
authorizes publishing before readiness is restored so a revival never comes back with zero
tools — likewise a server retained after a recoverable initial failure."""
"""Publish freshly discovered tools when none are registered (initial registration normally happens in
``_discover_and_register_server``). Outage handling may clear ``_ready`` and deregister stale tools;
ownership via ``_servers`` authorizes publishing before readiness is restored so a revival (or a server
retained after a recoverable initial failure) never comes back with zero tools."""
if self._registered_tool_names:
return
if not self._ready.is_set():
with _core._lock:
if _core._servers.get(self.name) is not self:
return
self._registered_tool_names = _core._register_server_tools(self.name, self, self._config)
# A retained initial-failure server that just published tools has recovered.
with _core._lock:
owned = _core._servers.get(self.name) is self
if not owned and not self._ready.is_set():
return
self._registered_tool_names = _core._register_server_tools(self.name, self, self._config)
with _core._lock: # a retained initial-failure server that just published tools has recovered
if _core._servers.get(self.name) is self:
_core._server_connect_errors.pop(self.name, None)
+39 -59
View File
@@ -23,7 +23,7 @@ try:
except ImportError:
fcntl = None
try:
import msvcrt
import msvcrt # noqa: F401
except ImportError:
pass
@@ -32,8 +32,7 @@ logger = logging.getLogger(__name__)
# One tool-definition pass must use ONE config decision for availability and the
# dynamic target schema: the check_fn result flows to the immediately following
# dynamic_schema_overrides call; ContextVar isolates concurrent profile builds.
_memory_surface_flags: ContextVar[Optional[Tuple[bool, bool]]] = ContextVar(
"memory_surface_flags", default=None)
_memory_surface_flags: ContextVar[Optional[Tuple[bool, bool]]] = ContextVar("memory_surface_flags", default=None)
def get_memory_dir() -> Path:
@@ -46,26 +45,22 @@ from tools.memory_tool_store import ( # noqa: E402,F401 (re-exports)
def load_on_disk_store() -> "MemoryStore":
"""Fresh on-disk MemoryStore with configured limits/flags for contexts with no
live agent (gateway, Desktop, ``/memory``) so approvals enforce the SAME caps
as ``agent_init``. Falls back to defaults if config can't load; never raises."""
"""Fresh on-disk MemoryStore with configured limits/flags for contexts with no live
agent (gateway, Desktop, ``/memory``) so approvals enforce the SAME caps as
``agent_init``. Falls back to defaults if config can't load; never raises."""
try:
from hermes_cli.config import load_config
config = load_config() or {}
mem_cfg = get_builtin_memory_config(config)
memory_enabled, user_profile_enabled = get_builtin_memory_store_flags(config)
kwargs = {"memory_char_limit": int(mem_cfg.get("memory_char_limit", 2200)),
"user_char_limit": int(mem_cfg.get("user_char_limit", 1375)),
"memory_enabled": memory_enabled, "user_profile_enabled": user_profile_enabled}
store = MemoryStore(int(mem_cfg.get("memory_char_limit", 2200)), int(mem_cfg.get("user_char_limit", 1375)),
memory_enabled=memory_enabled, user_profile_enabled=user_profile_enabled)
except Exception:
kwargs: Dict[str, Any] = {} # config optional — fall back to defaults rather than break /memory
store = MemoryStore(**kwargs)
store = MemoryStore() # config optional — fall back to defaults rather than break /memory
store.load_from_disk()
return store
# -- Write-approval gate --
def _gate_or_stage(summary: str, detail: str, payload: Dict[str, Any]) -> Optional[str]:
"""JSON tool-result string when the write must NOT proceed (blocked or staged
for approval), None to proceed. Fails open if the gate module can't load."""
@@ -93,35 +88,30 @@ _STORE_ACTIONS = {
lambda label, content, old_text: (f"remove from {label}", old_text or ""))}
def _apply_write_gate(action: str, target: str, content: Optional[str], old_text: Optional[str]) -> Optional[str]:
"""Gate a single mutating op (add/replace/remove)."""
summary, detail = _STORE_ACTIONS[action][1]("user profile" if target == "user" else "memory", content, old_text)
return _gate_or_stage(summary, detail,
def _batch_op_line(op: Dict[str, Any]) -> str:
op = op or {}
act, content, old = op.get("action", "?"), op.get("content") or op.get("new_text") or "", op.get("old_text", "")
if act == "remove":
return f"- remove: {old}"
return f"- replace: {old} -> {content}" if act == "replace" else f"- {act}: {content}"
def _apply_write_gate(action: str, target: str, content: Optional[str], old_text: Optional[str],
operations: Optional[List[Dict[str, Any]]] = None) -> Optional[str]:
"""Gate one mutating op, or (``operations`` set) a whole batch as a single unit."""
label = "user profile" if target == "user" else "memory"
if operations is not None:
return _gate_or_stage(f"apply {len(operations)} op(s) to {label}",
"\n".join(_batch_op_line(op) for op in operations),
{"action": "batch", "target": target, "operations": operations})
return _gate_or_stage(*_STORE_ACTIONS[action][1](label, content, old_text),
{"action": action, "target": target, "content": content, "old_text": old_text})
def _apply_batch_write_gate(target: str, operations: List[Dict[str, Any]]) -> Optional[str]:
"""Gate a whole batch as a single unit."""
summary = f"apply {len(operations)} op(s) to {'user profile' if target == 'user' else 'memory'}"
detail_lines = []
for op in operations:
op = op or {}
act = op.get("action", "?")
content = op.get("content") or op.get("new_text") or ""
detail_lines.append(f"- remove: {op.get('old_text', '')}" if act == "remove"
else f"- replace: {op.get('old_text', '')} -> {content}" if act == "replace"
else f"- {act}: {content}")
return _gate_or_stage(summary, "\n".join(detail_lines),
{"action": "batch", "target": target, "operations": operations})
# -- Tool entry point --
def _validate_single_op(store, action, target, content, old_text) -> Optional[str]:
"""Validate BEFORE the gate so an invalid write is rejected now, not at approve
time. Missing ``old_text`` is recoverable (it can't be schema-required — needs a
combinator the Codex backend rejects — and some clients omit it): return the
current inventory plus a retry instruction instead of a dead-end."""
"""Validate BEFORE the gate so an invalid write is rejected now, not at approve time.
Missing ``old_text`` is recoverable (it can't be schema-required — needs a combinator
the Codex backend rejects): return the inventory plus a retry instruction."""
if action == "add" and not content:
return tool_error("Content is required for 'add' action.", success=False)
if action in ("replace", "remove") and not old_text:
@@ -154,26 +144,23 @@ def memory_tool(action: str = None, target: str = "memory", content: str = None,
if operations:
if not isinstance(operations, list):
return tool_error("operations must be a list of {action, content?, old_text?} objects.", success=False)
gate_result = _apply_batch_write_gate(target, operations)
# Approval gate: stages (background/gateway) or prompts inline (CLI); off by default.
gate_result = _apply_write_gate("batch", target, None, None, operations)
if gate_result is not None:
return gate_result
return json.dumps(store.apply_batch(target, operations), ensure_ascii=False)
if action not in _STORE_ACTIONS:
return tool_error(f"Unknown action '{action}'. Use: add, replace, remove", success=False)
invalid = _validate_single_op(store, action, target, content, old_text)
invalid = (_validate_single_op(store, action, target, content, old_text)
or _apply_write_gate(action, target, content, old_text))
if invalid is not None:
return invalid
# Approval gate: stages (background/gateway) or prompts inline (CLI); off by default.
gate_result = _apply_write_gate(action, target, content, old_text)
if gate_result is not None:
return gate_result
return json.dumps(_STORE_ACTIONS[action][0](store, target, content, old_text), ensure_ascii=False)
def get_builtin_memory_config(config: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Normalized ``memory`` config section ({} when missing/malformed → flags default
to enabled). ``agent_init`` reads the same section so availability and store
construction cannot diverge."""
"""Normalized ``memory`` config section ({} when missing/malformed → flags default to
enabled). ``agent_init`` reads the same section so availability and store cannot diverge."""
if config is None:
try:
from hermes_cli.config import load_config_readonly
@@ -214,8 +201,7 @@ def _memory_target_error(store: "MemoryStore", target: str) -> Optional[Dict[str
def apply_memory_pending(payload: Dict[str, Any], store: "MemoryStore") -> Dict[str, Any]:
"""Replay a staged write against the store, bypassing the gate (/memory approve)."""
action = payload.get("action")
target = payload.get("target", "memory")
action, target = payload.get("action"), payload.get("target", "memory")
target_error = _memory_target_error(store, target)
if target_error is not None:
return target_error
@@ -226,8 +212,6 @@ def apply_memory_pending(payload: Dict[str, Any], store: "MemoryStore") -> Dict[
return _STORE_ACTIONS[action][0](store, target, payload.get("content") or "", payload.get("old_text") or "")
# -- OpenAI Function-Calling Schema --
MEMORY_SCHEMA = {
"name": "memory",
"description": (
@@ -318,21 +302,17 @@ def _build_memory_schema_overrides() -> Dict[str, Any]:
_memory_surface_flags.set(None)
targets = [t for t, on in zip(("memory", "user"), flags) if on]
parameters = copy.deepcopy(MEMORY_SCHEMA["parameters"])
target_schema = parameters["properties"]["target"]
target_schema, description = parameters["properties"]["target"], MEMORY_SCHEMA["description"]
target_schema["enum"] = targets
description = MEMORY_SCHEMA["description"]
narrowed = _SINGLE_TARGET_TEXT.get(tuple(targets))
if narrowed:
if narrowed := _SINGLE_TARGET_TEXT.get(tuple(targets)):
target_schema["description"], replacement = narrowed
description = description.replace(
"TARGETS: 'user' = who the user is (name, role, preferences, style). 'memory' = your "
"notes (environment, conventions, tool quirks, lessons).",
replacement)
"notes (environment, conventions, tool quirks, lessons).", replacement)
return {"description": description, "parameters": parameters}
# --- Registry ---
from tools.registry import registry, tool_error
from tools.registry import registry, tool_error # noqa: E402 (registration at import time)
registry.register(
name="memory",
+105 -154
View File
@@ -5,7 +5,7 @@ in ``tools.memory_tool`` and is read lazily."""
import logging
import time
from contextlib import contextmanager
from contextlib import contextmanager, suppress
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
@@ -22,37 +22,36 @@ MEMORY_BLOCK_HEADERS = {
ENTRY_DELIMITER = "\n§\n"
def _memory_dir() -> Path:
from tools import memory_tool
return memory_tool.get_memory_dir()
def _scan_memory_content(content: str) -> Optional[str]:
"""Error string if *content* matches injection/exfil patterns. Strict scope:
memory enters the system prompt, so a poisoned entry persists across sessions."""
return _first_threat_message(content, scope="strict")
def _drift_error(path: "Path", bak_path: str) -> Dict[str, Any]:
def _error(message: str, **extra) -> Dict[str, Any]:
return {"success": False, "error": message, **extra}
def _drift_error(path: Path, bak_path: str) -> Dict[str, Any]:
"""External drift: the file wouldn't round-trip, so flushing would discard content."""
return {"success": False, "error": (
return _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": (
), 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.")}
"memory(action=add, content=...), then remove or rewrite the original file to a clean state."))
def _read_failed_error(path: "Path") -> Dict[str, Any]:
def _read_failed_error(path: Path) -> Dict[str, Any]:
"""Existing-but-unreadable file: saving from an assumed-empty view would wipe it."""
return {"success": False, "error": (
return _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.")}
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]:
@@ -78,19 +77,16 @@ class MemoryStore:
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.memory_char_limit, self.user_char_limit = memory_char_limit, user_char_limit
self.memory_enabled, self.user_profile_enabled = memory_enabled, user_profile_enabled
self._system_prompt_snapshot: Dict[str, str] = {"memory": "", "user": ""}
self._consolidation_failures = 0 # per turn; reset by reset_consolidation_failures()
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)."""
"""Call at turn start."""
self._consolidation_failures = 0
def _consolidation_failure(self, response: Dict[str, Any]) -> Dict[str, Any]:
@@ -109,22 +105,10 @@ class MemoryStore:
Threat hits are replaced by a ``[BLOCKED: …]`` placeholder in the SNAPSHOT only;
live lists keep the raw text so the user can see and remove poisoned entries
(dropping them silently would hide the attack)."""
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 = {
"memory": self._render_block("memory", self._sanitize_entries_for_snapshot(self.memory_entries, "MEMORY.md")),
"user": self._render_block("user", self._sanitize_entries_for_snapshot(self.user_entries, "USER.md"))}
@staticmethod
def _sanitize_entries_for_snapshot(entries: List[str], filename: str) -> List[str]:
"""*entries* with threat matches replaced by a ``[BLOCKED: …]`` placeholder
(strict scope, same as writes); empty / already-blocked entries pass through."""
from tools.threat_patterns import scan_for_threats
def _one(entry):
def _sanitize(entry, filename):
# Strict scope, same as writes; empty / already-blocked entries pass through.
findings = scan_for_threats(entry, scope="strict") if entry and not entry.startswith("[BLOCKED:") else None
if not findings:
return entry
@@ -132,7 +116,13 @@ class MemoryStore:
return (f"[BLOCKED: {filename} entry contained threat pattern(s): {', '.join(findings)}. "
f"Removed from system prompt; use memory(action=remove) to delete the original.]")
return [_one(e) for e in entries]
for target in ("memory", "user"):
path = self._path_for(target)
path.parent.mkdir(parents=True, exist_ok=True)
# Deduplicate (order-preserving, first occurrence wins).
entries = list(dict.fromkeys(self._read_file(path)))
self._set_entries(target, entries)
self._system_prompt_snapshot[target] = self._render_block(target, [_sanitize(e, path.name) for e in entries])
@staticmethod
@contextmanager
@@ -157,33 +147,13 @@ class MemoryStore:
try:
yield
finally:
try:
with suppress(OSError):
_flock(True)
except OSError:
pass
@staticmethod
def _path_for(target: str) -> Path:
return _memory_dir() / ("USER.md" if target == "user" else "MEMORY.md")
def _reload_or_error(self, target: str, *, skip_drift: bool = False) -> Optional[Dict[str, Any]]:
"""Re-read entries from disk (under lock) before mutating; return the abort
error dict or None. Aborts on external drift (flushing would discard
un-roundtrippable content) and on an existing-but-unreadable file (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"."""
path = self._path_for(target)
raw, read_ok = self._read_raw_checked(path)
if not read_ok:
return _read_failed_error(path)
bak = None if skip_drift else self._detect_external_drift(target, raw)
self._set_entries(target, list(dict.fromkeys(self._parse_entries(raw))))
return _drift_error(path, bak) if bak else 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))
from tools import memory_tool # get_memory_dir is monkeypatched there
return memory_tool.get_memory_dir() / ("USER.md" if target == "user" else "MEMORY.md")
def _entries_for(self, target: str) -> List[str]:
return self.user_entries if target == "user" else self.memory_entries
@@ -201,53 +171,45 @@ class MemoryStore:
return f"{self._char_count(target):,}/{self._char_limit(target):,}"
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"
return f"{min(100, int((current / limit) * 100)) if limit > 0 else 0}% — {current:,}/{limit:,} chars"
def _failure_with_entries(self, target: str, message: str) -> Dict[str, Any]:
"""Consolidation failure carrying the live entries so the model can 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": [e[:80] + ("..." if len(e) > 80 else "") 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
return self._consolidation_failure(
_error(message, current_entries=self._entries_for(target), usage=self._usage(target)))
def _mutate(self, target: str, mutate, *, skip_drift: bool = False) -> Dict[str, Any]:
"""Lock, reload, run ``mutate(entries, limit)`` -> ``(new_entries, message)`` or an
error dict, then persist and return the success response."""
with self._file_lock(self._path_for(target)):
err = self._reload_or_error(target, skip_drift=skip_drift)
if err:
return err
"""Lock, re-read from disk, run ``mutate(entries, limit)`` -> ``(new_entries, message)``
or an error dict, then persist and return the success response. The reload aborts
on an existing-but-unreadable file (even append-only ``add`` rewrites the whole
file) and, unless *skip_drift*, on external drift (flushing would discard
un-roundtrippable content). Drift check and parse use the SAME raw snapshot —
a failed second read used to count as "no drift"."""
path = self._path_for(target)
with self._file_lock(path):
raw, read_ok = self._read_raw_checked(path)
if not read_ok:
return _read_failed_error(path)
bak = None if skip_drift else self._detect_external_drift(target, raw)
self._set_entries(target, list(dict.fromkeys(self._parse_entries(raw))))
if bak:
return _drift_error(path, bak)
result = mutate(self._entries_for(target), self._char_limit(target))
if isinstance(result, dict):
return result
entries, message = result
self._set_entries(target, entries)
self.save_to_disk(target)
return self._success_response(target, message)
self._set_entries(target, result[0])
path.parent.mkdir(parents=True, exist_ok=True)
self._write_file(path, result[0])
return self._success_response(target, result[1])
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}
return _error("Content cannot be empty.")
if scan_error := _scan_memory_content(content):
return _error(scan_error)
def _add(entries, limit):
if content in entries:
@@ -259,7 +221,6 @@ class MemoryStore:
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."))
return entries + [content], "Entry added."
# Append-only: skip the drift guard (appending never clobbers foreign
# content) but still refuse a failed read — add rewrites the WHOLE file.
return self._mutate(target, _add, skip_drift=True)
@@ -268,29 +229,33 @@ class MemoryStore:
"""Find entry containing old_text substring, replace it with new_content."""
new_content = new_content.strip()
if not old_text.strip():
return {"success": False, "error": "old_text cannot be empty."}
return _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}
return _error("new_content cannot be empty. Use 'remove' to delete entries.")
if scan_error := _scan_memory_content(new_content):
return _error(scan_error)
return self._edit(target, old_text.strip(), new_content)
def remove(self, target: str, old_text: str) -> Dict[str, Any]:
"""Remove the entry containing old_text substring."""
if not old_text.strip():
return {"success": False, "error": "old_text cannot be empty."}
return _error("old_text cannot be empty.")
return self._edit(target, old_text.strip(), None)
def _edit(self, target: str, old_text: str, new_content: Optional[str]) -> Dict[str, Any]:
"""Locked replace (``new_content`` set) or remove (None) of the unique entry matching *old_text*."""
"""Locked replace (``new_content`` set) or remove (None) of the entry matching *old_text*."""
def _apply(entries, limit):
idx, err = self._locate(target, old_text, "replace" if new_content else "remove")
if err:
return err
idx, ambiguous = _find_unique_match(entries, old_text)
if ambiguous:
return _error(f"Multiple entries matched '{old_text}'. Be more specific.",
matches=[e[:80] + ("..." if len(e) > 80 else "") for e in entries if old_text in e])
if idx is None:
return self._consolidation_failure(_error(
f"No entry matched '{old_text}'. Check current_entries below and retry with the exact text "
f"of the entry you want to {'replace' if new_content else 'remove'}.", current_entries=entries))
replaced = entries[:idx] + ([] if new_content is None else [new_content]) + entries[idx + 1:]
if new_content is None:
return entries[:idx] + entries[idx + 1:], "Entry removed."
replaced = entries[:idx] + [new_content] + entries[idx + 1:]
return replaced, "Entry removed."
new_total = len(ENTRY_DELIMITER.join(replaced))
if new_total > limit:
return self._failure_with_entries(target, (
@@ -298,7 +263,6 @@ class MemoryStore:
f"or 'remove' other stale or less important entries to make room (see current_entries "
f"below), then retry — all in this turn."))
return replaced, "Entry replaced."
return self._mutate(target, _apply)
@staticmethod
@@ -321,79 +285,67 @@ class MemoryStore:
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)
working[idx:idx + 1] = [content] if act == "replace" else []
return None
def apply_batch(self, target: str, operations: List[Dict[str, Any]]) -> Dict[str, Any]:
"""Apply add/replace/remove ops to one target atomically against the FINAL
budget, so one call can free space and add entries. All-or-nothing: any
malformed / unmatched op or an over-limit result writes NOTHING and returns
the first failure plus live state."""
"""Apply add/replace/remove ops atomically against the FINAL budget, so one call
can free space and add entries. All-or-nothing: any malformed / unmatched op or
an over-limit result writes NOTHING and returns the first failure plus live state."""
if not operations:
return {"success": False, "error": "operations list is empty."}
return _error("operations list is empty.")
ops = [op or {} for op in operations]
# Scan every add/replace content BEFORE touching disk -- one poisoned op rejects the batch.
for i, op in enumerate(ops):
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}"}
return _error(f"Operation {i + 1}: {scan_error}")
def _apply(entries, limit):
working = list(entries) # only committed if the whole batch validates
for i, op in enumerate(ops):
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)
msg = self._apply_batch_op(working, act, (op.get("content") or op.get("new_text") or "").strip(),
(op.get("old_text") or "").strip(), f"Operation {i + 1} ({act or 'unknown'})")
if msg:
return self._failure_with_entries(
target, msg + " No operations were applied (batch is all-or-nothing).")
# Budget check against the FINAL state only.
new_total = len(ENTRY_DELIMITER.join(working))
return self._failure_with_entries(target, msg + " No operations were applied (batch is all-or-nothing).")
new_total = len(ENTRY_DELIMITER.join(working)) # budget check against the FINAL state only
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 working, f"Applied {len(operations)} operation(s)."
return self._mutate(target, _apply)
def format_for_system_prompt(self, target: str) -> Optional[str]:
"""Frozen load-time snapshot for the system prompt (NOT live state — mid-session
writes don't touch it, preserving the prefix cache); None if empty."""
"""Frozen load-time snapshot (NOT live state — mid-session writes don't touch
it, preserving the prefix cache); None if empty."""
return self._system_prompt_snapshot.get(target, "") or None
def _success_response(self, target: str, message: str = None) -> Dict[str, Any]:
# A successful write is progress: reset the per-turn (consecutive) failure budget.
"""TERMINAL and WITHOUT the entries list: echoing entries invites the model to
"find more to fix" and re-issue the same ops. A successful write resets the
per-turn failure budget."""
self._consolidation_failures = 0
# TERMINAL and WITHOUT the entries list: echoing entries invites the model to
# "find more to fix" and re-issue the same ops. Entries only appear on errors.
return {"success": True, "done": True, "target": target,
"usage": self._usage_pct(target, self._char_count(target)),
"entry_count": len(self._entries_for(target)), **({"message": message} if message else {}),
"note": "Write saved. This update is complete — do not repeat it."}
def _render_block(self, target: str, entries: List[str]) -> str:
"""Render a system prompt block with header and usage indicator."""
"""System prompt block: header + usage indicator + entries ("" when empty)."""
if not entries:
return ""
content = ENTRY_DELIMITER.join(entries)
content, sep = ENTRY_DELIMITER.join(entries), "═" * 46
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}"
return f"{sep}\n{title} [{self._usage_pct(target, len(content))}]\n{sep}\n{content}"
@staticmethod
def _read_raw_checked(path: Path) -> Tuple[str, bool]:
"""``(raw, read_ok)``; ``read_ok`` is False ONLY when the file EXISTS but can't
be read (absent → ``("", True)``). Decoding stays STRICT: ``errors="replace"``
would hand callers a lossy view that a save then persists. ``utf-8-sig`` strips
a Notepad BOM that otherwise glues U+FEFF onto the first entry forever."""
"""``(raw, read_ok)``; ``read_ok`` is False ONLY when the file EXISTS but can't be
read. Decoding stays STRICT (``errors="replace"`` would hand callers a lossy view
a save then persists); ``utf-8-sig`` strips a Notepad BOM off the first entry."""
if not path.exists():
return "", True
try:
@@ -408,19 +360,26 @@ class MemoryStore:
@staticmethod
def _read_file(path: Path) -> List[str]:
"""Entries of a memory file ([] on any error). Read-only callers only
(``load_from_disk``, learning_mutations); mutation paths must use
``_read_raw_checked`` so they can refuse to overwrite an unreadable file."""
"""Entries of a memory file ([] on any error). Read-only callers only; mutation
paths use ``_read_raw_checked`` so they can refuse to overwrite an unreadable file."""
return MemoryStore._parse_entries(MemoryStore._read_raw_checked(path)[0])
@staticmethod
def _write_file(path: Path, entries: List[str]):
"""Atomic temp-file + rename: readers never see a truncated file. Also used by
agent/learning_mutations.py."""
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}")
def _detect_external_drift(self, target: str, raw: str) -> Optional[str]:
"""Backup path if *raw* shows external drift, else None. Signals: round-trip
mismatch, or one entry over the whole-file limit (no tool-written entry can —
an external writer appended free-form text). Snapshots to ``.bak.<ts>``."""
if not raw.strip():
return None
"""``.bak.<ts>`` snapshot path if *raw* shows external drift, else None. Signals:
round-trip mismatch, or one entry over the whole-file limit (no tool-written
entry can be — an external writer appended free-form text)."""
parsed = self._parse_entries(raw)
if raw.strip() == ENTRY_DELIMITER.join(parsed) and max(map(len, parsed), default=0) <= self._char_limit(target):
if not raw.strip() or (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())}")
@@ -429,11 +388,3 @@ class MemoryStore:
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 never see a truncated file."""
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}")
+7 -13
View File
@@ -53,13 +53,10 @@ class GraphCredentials:
@property
def token_url(self) -> str:
tenant = self.tenant_id.strip().strip("/")
return f"{self.authority_url.rstrip('/')}/{tenant}/oauth2/v2.0/token"
return f"{self.authority_url.rstrip('/')}/{self.tenant_id.strip().strip('/')}/oauth2/v2.0/token"
@classmethod
def from_env(
cls, environ: dict[str, str] | None = None, *, required: bool = True
) -> "GraphCredentials | None":
def from_env(cls, environ: dict[str, str] | None = None, *, required: bool = True) -> "GraphCredentials | None":
env = environ if environ is not None else os.environ
values = [(env.get(name) or "").strip() for name in _REQUIRED_ENV]
missing = [name for name, value in zip(_REQUIRED_ENV, values) if not value]
@@ -90,10 +87,9 @@ class CachedAccessToken:
class MicrosoftGraphTokenProvider:
"""Acquire and cache Microsoft Graph app-only access tokens."""
def __init__(
self, credentials: GraphCredentials, *, timeout: float = 20.0,
skew_seconds: int = DEFAULT_TOKEN_SKEW_SECONDS, transport: httpx.AsyncBaseTransport | None = None,
) -> None:
def __init__(self, credentials: GraphCredentials, *, timeout: float = 20.0,
skew_seconds: int = DEFAULT_TOKEN_SKEW_SECONDS,
transport: httpx.AsyncBaseTransport | None = None) -> None:
self.credentials, self.timeout, self.skew_seconds = credentials, timeout, max(0, int(skew_seconds))
self._transport = transport
self._cached_token: CachedAccessToken | None = None
@@ -150,10 +146,8 @@ class MicrosoftGraphTokenProvider:
except (TypeError, ValueError) as exc:
raise MicrosoftGraphTokenError(
"Microsoft Graph token response did not include a valid expires_in.") from exc
return CachedAccessToken(
access_token=access_token,
token_type=str(payload.get("token_type") or "Bearer").strip() or "Bearer",
expires_at=time.time() + max(0, expires_in_seconds))
return CachedAccessToken(access_token, time.time() + max(0, expires_in_seconds),
str(payload.get("token_type") or "Bearer").strip() or "Bearer")
def _extract_error_detail(response: httpx.Response) -> str:
+34 -57
View File
@@ -10,7 +10,7 @@ from typing import Any, Awaitable, Callable
import httpx
from agent.retry_utils import parse_retry_after_seconds
from tools.microsoft_graph_auth import GraphCredentials, MicrosoftGraphTokenProvider, format_graph_error
from tools.microsoft_graph_auth import MicrosoftGraphTokenProvider, format_graph_error
DEFAULT_GRAPH_BASE_URL = "https://graph.microsoft.com/v1.0"
@@ -26,9 +26,8 @@ class MicrosoftGraphClientError(RuntimeError):
class MicrosoftGraphAPIError(MicrosoftGraphClientError):
"""Raised when a Graph API request fails."""
def __init__(
self, status_code: int, method: str, url: str, message: str, *,
retry_after_seconds: float | None = None, payload: Any = None) -> None:
def __init__(self, status_code: int, method: str, url: str, message: str, *,
retry_after_seconds: float | None = None, payload: Any = None) -> None:
self.status_code, self.method, self.url = status_code, method, url
self.retry_after_seconds, self.payload = retry_after_seconds, payload
super().__init__(f"Microsoft Graph API error {status_code} for {method} {url}: {message}")
@@ -39,20 +38,15 @@ class MicrosoftGraphClient:
alike): transport errors back off exponentially; 401 clears the token cache and
refetches; 429/5xx honor ``Retry-After``. Each attempt uses a fresh ``AsyncClient``."""
def __init__(
self, token_provider: MicrosoftGraphTokenProvider, *,
base_url: str = DEFAULT_GRAPH_BASE_URL, timeout: float = 60.0, max_retries: int = 3,
transport: httpx.AsyncBaseTransport | None = None,
sleep: Callable[[float], Awaitable[None]] | None = None,
user_agent: str = "Hermes-Agent/graph-client") -> None:
def __init__(self, token_provider: MicrosoftGraphTokenProvider, *,
base_url: str = DEFAULT_GRAPH_BASE_URL, timeout: float = 60.0, max_retries: int = 3,
transport: httpx.AsyncBaseTransport | None = None,
sleep: Callable[[float], Awaitable[None]] | None = None,
user_agent: str = "Hermes-Agent/graph-client") -> None:
self.token_provider, self.base_url, self.timeout = token_provider, base_url.rstrip("/"), timeout
self.max_retries, self.user_agent = max(0, int(max_retries)), user_agent
self._transport, self._sleep = transport, sleep or asyncio.sleep
@classmethod
def from_env(cls, **kwargs: Any) -> "MicrosoftGraphClient":
return cls(MicrosoftGraphTokenProvider(GraphCredentials.from_env()), **kwargs)
async def get_json(self, path: str, *, params: Params = None, headers: Headers = None) -> Any:
return self._decode_json(await self._request("GET", path, params=params, headers=headers))
@@ -60,27 +54,24 @@ class MicrosoftGraphClient:
return self._decode_json(await self._request("POST", path, json_body=json_body, headers=headers))
async def patch_json(self, path: str, *, json_body: Any | None = None, headers: Headers = None) -> Any:
"""Decoded body, or ``{}`` for a 204 / bodiless response."""
response = await self._request("PATCH", path, json_body=json_body, headers=headers)
return self._decode_json_or(response, {})
return self._decode_json(response) if response.status_code != 204 and response.content else {}
async def delete(self, path: str, *, headers: Headers = None) -> dict[str, Any]:
"""Decoded body, or ``{"deleted": True, "status_code"}`` for a 204 / bodiless response."""
response = await self._request("DELETE", path, headers=headers)
return self._decode_json_or(response, {"deleted": True, "status_code": response.status_code})
if response.status_code != 204 and response.content:
return self._decode_json(response)
return {"deleted": True, "status_code": response.status_code}
def _decode_json_or(self, response: httpx.Response, empty: Any) -> Any:
"""*empty* for a 204 / bodiless response, else the decoded JSON body."""
return empty if response.status_code == 204 or not response.content else self._decode_json(response)
async def collect_paginated(
self, path: str, *, params: Params = None, headers: Headers = None) -> list[Any]:
async def collect_paginated(self, path: str, *, params: Params = None, headers: Headers = None) -> list[Any]:
"""Follow ``@odata.nextLink`` and concatenate every page's ``value`` list."""
items: list[Any] = []
# Query params go on the first request only; @odata.nextLink already embeds them.
next_url: str | None = self._resolve_url(path)
next_params = dict(params or {})
next_url, next_params = self._resolve_url(path), dict(params or {})
while next_url:
response = await self._request("GET", next_url, params=next_params or None, headers=headers)
payload = self._decode_json(response)
payload = self._decode_json(await self._request("GET", next_url, params=next_params or None, headers=headers))
if not isinstance(payload, dict):
raise MicrosoftGraphClientError(
f"Expected paginated Graph response dict, got {type(payload).__name__}.")
@@ -89,13 +80,11 @@ class MicrosoftGraphClient:
next_url, next_params = payload.get("@odata.nextLink"), {}
return items
async def download_to_file(
self, path: str, destination: str | Path, *, headers: Headers = None, chunk_size: int = 65536
) -> dict[str, Any]:
async def download_to_file(self, path: str, destination: str | Path, *, headers: Headers = None,
chunk_size: int = 65536) -> dict[str, Any]:
"""Stream a Graph resource to disk chunk-by-chunk (large recordings never
fit in memory); written to ``.part`` and renamed into place only on success."""
url = self._resolve_url(path)
target = Path(destination)
url, target = self._resolve_url(path), Path(destination)
target.parent.mkdir(parents=True, exist_ok=True)
tmp_target = target.with_suffix(target.suffix + ".part")
@@ -103,8 +92,7 @@ class MicrosoftGraphClient:
try:
async with client.stream("GET", url, headers=request_headers) as response:
if response.status_code >= 400:
# Materialize the (small) error body so the message is meaningful.
await response.aread()
await response.aread() # small error body -> meaningful message
return response, None
with tmp_target.open("wb") as handle:
async for chunk in response.aiter_bytes(chunk_size=chunk_size):
@@ -119,9 +107,8 @@ class MicrosoftGraphClient:
os.replace(tmp_target, target)
return {"path": str(target), "size_bytes": target.stat().st_size, "content_type": content_type}
async def _request(
self, method: str, path_or_url: str, *,
params: Params = None, json_body: Any | None = None, headers: Headers = None) -> httpx.Response:
async def _request(self, method: str, path_or_url: str, *, params: Params = None,
json_body: Any | None = None, headers: Headers = None) -> httpx.Response:
url = self._resolve_url(path_or_url)
async def perform(client: httpx.AsyncClient, request_headers: dict[str, str]):
@@ -137,47 +124,37 @@ class MicrosoftGraphClient:
"""Run ``perform`` (-> ``(response, result)``) under the retry policy. ``kind``
only labels transport-failure messages. Raises ``MicrosoftGraphAPIError`` once
retries are exhausted or the status is not retryable; only 401 forces a token refresh."""
attempt = 0
last_error: Exception | None = None
while attempt <= self.max_retries:
for attempt in range(self.max_retries + 1):
token = await self.token_provider.get_access_token(
force_refresh=attempt > 0
and isinstance(last_error, MicrosoftGraphAPIError)
and last_error.status_code == 401)
request_headers = {"Authorization": f"Bearer {token}", "Accept": accept, "User-Agent": self.user_agent}
if json_body is not None:
request_headers["Content-Type"] = "application/json"
if headers:
request_headers.update(headers)
force_refresh=isinstance(last_error, MicrosoftGraphAPIError) and last_error.status_code == 401)
request_headers = {"Authorization": f"Bearer {token}", "Accept": accept, "User-Agent": self.user_agent,
**({"Content-Type": "application/json"} if json_body is not None else {}),
**(headers or {})}
exhausted = attempt >= self.max_retries
try:
async with httpx.AsyncClient(timeout=httpx.Timeout(self.timeout), transport=self._transport) as client:
response, result = await perform(client, request_headers)
except httpx.HTTPError as exc:
last_error, response = exc, None
if attempt >= self.max_retries:
if exhausted:
raise MicrosoftGraphClientError(
f"Microsoft Graph {kind} failed for {method} {url}: {exc}") from exc
else:
if response.status_code < 400:
return result
last_error = self._build_api_error(method, url, response)
status = response.status_code
if attempt >= self.max_retries or not (status in (401, 429) or 500 <= status < 600):
last_error, status = self._build_api_error(method, url, response), response.status_code
if exhausted or not (status in (401, 429) or 500 <= status < 600):
raise last_error
if status == 401:
self.token_provider.clear_cache()
await self._sleep(self._retry_delay(response, attempt))
attempt += 1
raise MicrosoftGraphClientError(f"Microsoft Graph {kind} exhausted retries for {method} {url}.")
def _resolve_url(self, path_or_url: str) -> str:
if path_or_url.startswith(("http://", "https://")):
return path_or_url
path = path_or_url if path_or_url.startswith("/") else f"/{path_or_url}"
return f"{self.base_url}{path}"
return f"{self.base_url}{path_or_url if path_or_url.startswith('/') else '/' + path_or_url}"
@staticmethod
def _decode_json(response: httpx.Response) -> Any:
@@ -201,5 +178,5 @@ class MicrosoftGraphClient:
payload = None
detail = format_graph_error(payload.get("error")) if isinstance(payload, dict) else None
return MicrosoftGraphAPIError(
response.status_code, method, url, detail if detail is not None else (response.text.strip() or "unknown error"),
response.status_code, method, url, response.text.strip() or "unknown error" if detail is None else detail,
retry_after_seconds=parse_retry_after_seconds(response.headers), payload=payload)
+121 -171
View File
@@ -1,14 +1,9 @@
#!/usr/bin/env python3
"""V4A patch format parser and applier (format used by codex, cline, etc.).
*** Begin Patch / *** End Patch wrap the operations:
*** Update File: p.py then hunks: ``@@ hint @@``, `` ctx``, ``-old``, ``+new``
*** Add File: n.py then ``+`` lines; *** Delete File: o.py; *** Move File: a -> b
operations, error = parse_v4a_patch(patch_content)
result = apply_v4a_operations(operations, file_ops)
"""
"""V4A patch parser/applier (codex, cline). ``*** Begin Patch``/``*** End Patch`` wrap ops:
``*** Update File: p`` + hunks (``@@ hint @@``, `` ctx``, ``-old``, ``+new``); ``*** Add File: n``
+ ``+`` lines; ``*** Delete File: o``; ``*** Move File: a -> b``. Entry points:
``parse_v4a_patch(text) -> (ops, error)`` and ``apply_v4a_operations(ops, file_ops)``."""
import contextlib
import difflib
import inspect
import re
@@ -42,7 +37,6 @@ class PatchOperation:
file_path: str
new_path: Optional[str] = None # MOVE only
hunks: List[Hunk] = field(default_factory=list)
content: Optional[str] = None # ADD only
# Markers must occupy the whole line at column 0 so content lines that merely
@@ -58,12 +52,9 @@ _HINT_RE = re.compile(r'@@\s*(.+?)\s*@@')
def parse_v4a_patch(patch_content: str) -> Tuple[List[PatchOperation], Optional[str]]:
"""Parse a V4A patch -> ``(operations, None)`` (``[]`` for an empty patch is not an
error) or ``([], "Parse error: ...")`` for malformed operations."""
# Tolerate CRLF bodies: a stray ``\r`` would otherwise end up in every
# HunkLine.content and defeat the anchored Begin/End markers.
"""-> ``(operations, None)`` (empty patch = ``[]``, no error) or ``([], "Parse error: …")``."""
# Tolerate CRLF: a stray ``\r`` would land in every HunkLine.content and defeat the markers.
lines = [ln[:-1] if ln.endswith('\r') else ln for ln in patch_content.split('\n')]
start_idx = -1 # parse from the top when no Begin marker is present
end_idx = len(lines)
for i, line in enumerate(lines):
@@ -72,7 +63,6 @@ def parse_v4a_patch(patch_content: str) -> Tuple[List[PatchOperation], Optional[
elif _END_MARKER.match(line):
end_idx = i
break
operations: List[PatchOperation] = []
current_op: Optional[PatchOperation] = None
current_hunk: Optional[Hunk] = None
@@ -87,8 +77,7 @@ def parse_v4a_patch(patch_content: str) -> Tuple[List[PatchOperation], Optional[
operations.append(current_op)
for line in lines[start_idx + 1:end_idx]:
op_match = next(
((kind, m) for kind, rx in _OP_MARKERS if (m := rx.match(line))), None)
op_match = next(((kind, m) for kind, rx in _OP_MARKERS if (m := rx.match(line))), None)
if op_match:
kind, m = op_match
_flush()
@@ -96,8 +85,8 @@ def parse_v4a_patch(patch_content: str) -> Tuple[List[PatchOperation], Optional[
operation=kind,
file_path=m.group(1).strip(),
new_path=m.group(2).strip() if kind is OperationType.MOVE else None)
# UPDATE hunks start lazily (at '@@' or the first hunk line); ADD
# collects all '+' lines into one hunk; DELETE/MOVE are complete.
# UPDATE hunks start lazily ('@@' or first hunk line); ADD collects all '+' lines
# into one hunk; DELETE/MOVE are complete.
current_hunk = Hunk() if kind is OperationType.ADD else None
if kind in (OperationType.DELETE, OperationType.MOVE):
operations.append(current_op)
@@ -115,18 +104,16 @@ def parse_v4a_patch(patch_content: str) -> Tuple[List[PatchOperation], Optional[
elif line[0] != '\\': # "\ No newline at end of file" marker is skipped
current_hunk.lines.append(HunkLine(' ', line)) # implicit context line
_flush()
parse_errors: List[str] = []
for op in operations:
if not op.file_path:
parse_errors.append("Operation with empty file path")
if op.operation == OperationType.UPDATE and not op.hunks:
if op.operation is OperationType.UPDATE and not op.hunks:
parse_errors.append(f"UPDATE {op.file_path!r}: no hunks found")
if op.operation == OperationType.MOVE and not op.new_path:
parse_errors.append(f"MOVE {op.file_path!r}: missing destination path (expected 'src -> dst')")
if parse_errors:
return [], "Parse error: " + "; ".join(parse_errors)
return operations, None
if op.operation is OperationType.MOVE and not op.new_path:
parse_errors.append(
f"MOVE {op.file_path!r}: missing destination path (expected 'src -> dst')")
return ([], "Parse error: " + "; ".join(parse_errors)) if parse_errors else (operations, None)
def _count_occurrences(text: str, pattern: str) -> int:
@@ -136,66 +123,57 @@ def _count_occurrences(text: str, pattern: str) -> int:
def _split_hunk(hunk: Hunk) -> Tuple[List[str], List[str]]:
"""``(search_lines, replace_lines)``: context+removed vs context+added."""
search = [l.content for l in hunk.lines if l.prefix in {' ', '-'}]
replace = [l.content for l in hunk.lines if l.prefix in {' ', '+'}]
return search, replace
return ([l.content for l in hunk.lines if l.prefix != '+'],
[l.content for l in hunk.lines if l.prefix != '-'])
def _no_match_hint(error: Optional[str], search_pattern: str, content: str) -> str:
"""Best-effort 'Did you mean...' suffix; never lets a hint failure mask the real error."""
try:
with contextlib.suppress(Exception):
from tools.fuzzy_match import format_no_match_hint
return format_no_match_hint(error, 0, search_pattern, content)
except Exception:
return ""
return ""
def _hint_ambiguity(content: str, hint: str, tail: str = "") -> Tuple[int, str]:
"""(occurrences, error) for an addition-only hunk's context hint; error is '' when unique."""
occurrences = _count_occurrences(content, hint)
if occurrences > 1:
return occurrences, (f"context hint '{hint}' is ambiguous "
f"({occurrences} occurrences){tail}")
return occurrences, ""
n = _count_occurrences(content, hint)
return n, f"context hint '{hint}' is ambiguous ({n} occurrences){tail}" if n > 1 else ""
def _validate_operations(operations: List[PatchOperation], file_ops: Any) -> List[str]:
"""Dry-run every operation; return error strings (empty list = safe to apply). UPDATE
hunks are simulated in order so later hunks see post-earlier-hunk content, as apply will."""
"""Dry-run every operation -> error strings (empty = safe). UPDATE hunks are simulated in
order so later hunks see post-earlier-hunk content, exactly as apply will."""
from tools.fuzzy_match import fuzzy_find_and_replace, is_already_applied
errors: List[str] = []
real_change_count = 0
# Virtual overlay so inter-op state validates (e.g. a MOVE creating the destination
# a later UPDATE targets): path -> content from an earlier op; paths MOVE/DELETE removed.
# Overlay so inter-op state validates (a MOVE creating the path a later UPDATE targets).
pending_content: dict = {}
removed_paths: set = set()
def _read(path: str):
if path in removed_paths and path not in pending_content:
return None, "file not found"
def _read(path: str) -> Tuple[Optional[str], Optional[str]]:
if path in pending_content:
return pending_content[path], None
if path in removed_paths:
return None, "file not found"
r = file_ops.read_file_raw(path)
return (None, r.error) if r.error else (r.content, None)
def _validate_update(op: PatchOperation) -> None:
nonlocal real_change_count
content, read_err = _read(op.file_path)
simulated, read_err = _read(op.file_path)
if read_err:
errors.append(f"{op.file_path}: {read_err}")
return
simulated = content
for hunk_index, hunk in enumerate(op.hunks, start=1):
search_lines, replace_lines = _split_hunk(hunk)
if not any(l.prefix in '-+' for l in hunk.lines):
# Inert anchor hunk (context only) — models emit these
# between real changes; ignore without failing the patch.
if search_lines == replace_lines:
# Context-only anchor hunks (models emit these between changes) are inert; identical
# -/+ lines are skipped by apply as a no-op — neither may fail validation.
real_change_count += any(l.prefix in '-+' for l in hunk.lines)
continue
real_change_count += 1
if not search_lines:
# Addition-only hunk: the context hint must be unique.
if not search_lines: # addition-only: the context hint must be unique
if hunk.context_hint:
occurrences, ambiguous = _hint_ambiguity(simulated, hunk.context_hint)
if occurrences == 0:
@@ -204,19 +182,13 @@ def _validate_operations(operations: List[PatchOperation], file_ops: Any) -> Lis
elif ambiguous:
errors.append(f"{op.file_path}: addition-only hunk {ambiguous}")
continue
search_pattern = '\n'.join(search_lines)
replacement = '\n'.join(replace_lines)
if search_lines == replace_lines:
# Identical -/+ lines: apply skips it as a no-op, so
# validation must not reject it with the identical-strings error.
continue
search_pattern, replacement = '\n'.join(search_lines), '\n'.join(replace_lines)
new_simulated, count, _strategy, match_error = fuzzy_find_and_replace(
simulated, search_pattern, replacement, replace_all=False)
if count:
simulated = new_simulated
elif not is_already_applied(simulated or "", search_pattern, replacement):
# Already-applied hunks (edit landed in a prior call) are no-ops so
# multi-hunk patches don't fail wholesale; apply performs the same skip.
# Already-applied hunks are no-ops (apply performs the same skip).
label = f"'{hunk.context_hint}'" if hunk.context_hint else "(no hint)"
errors.append(
f"{op.file_path}: hunk {hunk_index} {label} not found"
@@ -224,18 +196,20 @@ def _validate_operations(operations: List[PatchOperation], file_ops: Any) -> Lis
+ _no_match_hint(match_error, search_pattern, simulated))
pending_content[op.file_path] = simulated
def _remove(path: str) -> None:
removed_paths.add(path)
pending_content.pop(path, None)
for op in operations:
if op.operation == OperationType.UPDATE:
_validate_update(op)
continue
real_change_count += 1
if op.operation == OperationType.DELETE:
_content, read_err = _read(op.file_path)
if read_err:
if _read(op.file_path)[1]:
errors.append(f"{op.file_path}: file not found for deletion")
else:
removed_paths.add(op.file_path)
pending_content.pop(op.file_path, None)
_remove(op.file_path)
elif op.operation == OperationType.MOVE:
if not op.new_path:
errors.append(f"{op.file_path}: MOVE operation missing destination path")
@@ -243,16 +217,12 @@ def _validate_operations(operations: List[PatchOperation], file_ops: Any) -> Lis
src_content, src_err = _read(op.file_path)
if src_err:
errors.append(f"{op.file_path}: source file not found for move")
_dst, dst_err = _read(op.new_path)
if not dst_err:
if not _read(op.new_path)[1]:
errors.append(f"{op.new_path}: destination already exists — move would overwrite")
# Only a cleanly-validated move updates the overlay.
if not src_err and dst_err:
elif not src_err: # only a cleanly-validated move updates the overlay
pending_content[op.new_path] = src_content if src_content is not None else ""
pending_content.pop(op.file_path, None)
removed_paths.add(op.file_path)
# ADD: parent directory creation handled by write_file; no pre-check needed.
_remove(op.file_path)
# ADD: write_file creates parent directories; no pre-check needed.
if not errors and real_change_count == 0:
errors.append("Patch contains no changes (only context lines were provided)")
return errors
@@ -262,104 +232,101 @@ def _validate_operations(operations: List[PatchOperation], file_ops: Any) -> Lis
ApplyResult = Tuple[bool, str, Optional[str], Optional[dict]]
def _fail(error: str) -> ApplyResult:
return False, error, None, None
def _written(result: Any, diff: str) -> ApplyResult:
"""Outcome of a write: its error, else success with LSP/lint propagated from the WriteResult."""
if result.error:
return _fail(result.error)
return True, diff, getattr(result, "lsp_diagnostics", None), getattr(result, "lint", None)
def _unified_diff(path: str, old: str, new: Optional[str]) -> str:
"""Unified diff ``a/path`` -> ``b/path`` (``new=None`` = deletion, ``/dev/null``)."""
return ''.join(difflib.unified_diff(
old.splitlines(keepends=True), [] if new is None else new.splitlines(keepends=True),
fromfile=f"a/{path}", tofile="/dev/null" if new is None else f"b/{path}"))
def apply_v4a_operations(operations: List[PatchOperation], file_ops: Any) -> 'PatchResult':
"""Validate all operations, then apply them (two-phase, atomic on validation failure).
A phase-2 failure (validate/apply race) is reported with a ``git diff`` note since state
may be inconsistent. ``file_ops`` needs read_file_raw/write_file/delete_file/move_file."""
"""Two-phase: validate everything, then apply (atomic on validation failure). A phase-2
failure (validate/apply race) carries a ``git diff`` note since state may be inconsistent.
``file_ops`` needs read_file_raw/write_file/delete_file/move_file."""
from tools.file_operations import PatchResult # avoid circular import
validation_errors = _validate_operations(operations, file_ops)
if validation_errors:
def _bullets(errs: List[str]) -> str:
return "\n".join(f" • {e}" for e in errs)
if errors := _validate_operations(operations, file_ops):
return PatchResult(
success=False,
error="Patch validation failed (no files were modified):\n"
+ "\n".join(f" • {e}" for e in validation_errors))
error="Patch validation failed (no files were modified):\n" + _bullets(errors))
files: Dict[str, List[str]] = {"created": [], "deleted": [], "modified": []}
all_diffs: List[str] = []
# V4A bypasses the WriteResult/PatchResult plumbing that write_file uses,
# so LSP diagnostics and lint must be propagated explicitly per file.
# V4A bypasses write_file's WriteResult plumbing: LSP diagnostics and lint propagate per file.
lsp_blocks: List[str] = []
errors: List[str] = []
lint_results: Dict[str, dict] = {}
for op in operations:
handler, verb, bucket = _APPLY_DISPATCH[op.operation]
try:
handler, verb, bucket = _APPLY_DISPATCH[op.operation]
ok, payload, lsp, lint = handler(op, file_ops)
if not ok:
errors.append(f"Failed to {verb} {op.file_path}: {payload}")
continue
label = op.file_path
if op.operation is OperationType.MOVE:
label = f"{op.file_path} -> {op.new_path}"
files[bucket].append(label)
all_diffs.append(payload)
if lsp:
lsp_blocks.append(lsp)
if lint:
lint_results[op.file_path] = lint
except Exception as e:
errors.append(f"Error processing {op.file_path}: {str(e)}")
# Each LSP block carries its own <diagnostics file="..."> header, so plain
# concatenation keeps per-file attribution.
result_kwargs = dict(
ok, payload = None, str(e)
if not ok:
prefix = f"Failed to {verb}" if ok is False else "Error processing"
errors.append(f"{prefix} {op.file_path}: {payload}")
continue
is_move = op.operation is OperationType.MOVE
files[bucket].append(f"{op.file_path} -> {op.new_path}" if is_move else op.file_path)
all_diffs.append(payload)
if lsp:
lsp_blocks.append(lsp)
if lint:
lint_results[op.file_path] = lint
# Each LSP block carries its own <diagnostics file="..."> header; joining keeps attribution.
return PatchResult(
success=not errors,
error=("Apply phase failed (state may be inconsistent — run `git diff` to assess):\n"
+ _bullets(errors)) if errors else None,
diff='\n'.join(all_diffs),
files_modified=files["modified"], files_created=files["created"], files_deleted=files["deleted"],
lint=lint_results if lint_results else None,
lsp_diagnostics="\n\n".join(lsp_blocks) if lsp_blocks else None)
if errors:
return PatchResult(
success=False,
error="Apply phase failed (state may be inconsistent — run `git diff` to assess):\n"
+ "\n".join(f" • {e}" for e in errors),
**result_kwargs)
return PatchResult(success=True, **result_kwargs)
lint=lint_results or None, lsp_diagnostics="\n\n".join(lsp_blocks) or None)
def _write_file_accepts_pre_content(file_ops: Any) -> bool:
"""True when ``file_ops.write_file`` accepts ``pre_content``. Decided from the signature,
not by catching TypeError around the call, so a TypeError raised *inside* a capable
write_file propagates instead of triggering a duplicate write."""
try:
"""Whether ``file_ops.write_file`` accepts ``pre_content`` — read from the signature, not by
catching TypeError around the call, so a TypeError raised *inside* it can't double-write."""
with contextlib.suppress(TypeError, ValueError):
params = inspect.signature(file_ops.write_file).parameters
except (TypeError, ValueError):
return False
return "pre_content" in params or any(
p.kind is inspect.Parameter.VAR_KEYWORD for p in params.values())
return "pre_content" in params or any(
p.kind is inspect.Parameter.VAR_KEYWORD for p in params.values())
return False
def _apply_add(op: PatchOperation, file_ops: Any) -> ApplyResult:
"""Create a file from the hunks' '+' lines."""
content_lines = [line.content for hunk in op.hunks for line in hunk.lines if line.prefix == '+']
result = file_ops.write_file(op.file_path, '\n'.join(content_lines))
if result.error:
return False, result.error, None, None
diff = f"--- /dev/null\n+++ b/{op.file_path}\n" + '\n'.join(f"+{line}" for line in content_lines)
return True, diff, getattr(result, "lsp_diagnostics", None), getattr(result, "lint", None)
return _written(result, diff)
def _apply_delete(op: PatchOperation, file_ops: Any) -> ApplyResult:
"""Delete a file, producing a real unified diff of the removed content."""
# Validation already confirmed existence; the re-read guards against races.
read_result = file_ops.read_file_raw(op.file_path)
read_result = file_ops.read_file_raw(op.file_path) # re-read guards validate/apply races
if read_result.error:
return False, f"Cannot delete {op.file_path}: file not found", None, None
return _fail(f"Cannot delete {op.file_path}: file not found")
result = file_ops.delete_file(op.file_path)
if result.error:
return False, result.error, None, None
diff = ''.join(difflib.unified_diff(
read_result.content.splitlines(keepends=True), [],
fromfile=f"a/{op.file_path}", tofile="/dev/null"))
return True, diff or f"# Deleted: {op.file_path}", None, None
diff = _unified_diff(op.file_path, read_result.content, None) or f"# Deleted: {op.file_path}"
return _fail(result.error) if result.error else (True, diff, None, None)
def _apply_move(op: PatchOperation, file_ops: Any) -> ApplyResult:
result = file_ops.move_file(op.file_path, op.new_path)
if result.error:
return False, result.error, None, None
return True, f"# Moved: {op.file_path} -> {op.new_path}", None, None
return _fail(result.error) if result.error else (
True, f"# Moved: {op.file_path} -> {op.new_path}", None, None)
def _insert_addition_only(new_content: str, hunk: Hunk, insert_text: str) -> Tuple[Optional[str], Optional[str]]:
@@ -371,71 +338,54 @@ def _insert_addition_only(new_content: str, hunk: Hunk, insert_text: str) -> Tup
return None, f"Addition-only hunk: {ambiguous}"
if occurrences == 1:
eol = new_content.find('\n', new_content.find(hunk.context_hint))
if eol != -1:
return new_content[:eol + 1] + insert_text + '\n' + new_content[eol + 1:], None
return new_content + '\n' + insert_text, None
# Hint not found — append at end as a safe fallback.
if eol == -1:
return new_content + '\n' + insert_text, None
return new_content[:eol + 1] + insert_text + '\n' + new_content[eol + 1:], None
# No hint / hint not found — append at end as a safe fallback.
return new_content.rstrip('\n') + '\n' + insert_text + '\n', None
def _apply_update(op: PatchOperation, file_ops: Any) -> ApplyResult:
"""Apply each hunk via fuzzy replace, then write once."""
from tools.fuzzy_match import fuzzy_find_and_replace, is_already_applied
# Raw read: no line-number prefixes or per-line truncation.
read_result = file_ops.read_file_raw(op.file_path)
read_result = file_ops.read_file_raw(op.file_path) # raw: no line numbers / truncation
if read_result.error:
return False, f"Cannot read file: {read_result.error}", None, None
current_content = read_result.content
new_content = current_content
return _fail(f"Cannot read file: {read_result.error}")
current_content = new_content = read_result.content
for hunk in op.hunks:
search_lines, replace_lines = _split_hunk(hunk)
if search_lines and search_lines == replace_lines:
continue
search_pattern, replacement = '\n'.join(search_lines), '\n'.join(replace_lines)
if not search_lines:
new_content, err = _insert_addition_only(new_content, hunk, '\n'.join(replace_lines))
new_content, err = _insert_addition_only(new_content, hunk, replacement)
if err:
return False, err, None, None
return _fail(err)
continue
search_pattern = '\n'.join(search_lines)
replacement = '\n'.join(replace_lines)
new_content, count, _strategy, error = fuzzy_find_and_replace(
new_content, search_pattern, replacement, replace_all=False)
if not (error and count == 0):
continue
# Retry inside a window around the context hint, if any.
hint_pos = new_content.find(hunk.context_hint) if hunk.context_hint else -1
if hint_pos != -1:
window_start = max(0, hint_pos - 500)
window_end = min(len(new_content), hint_pos + 2000)
window_new, count, _strategy, error = fuzzy_find_and_replace(
new_content[window_start:window_end], search_pattern, replacement, replace_all=False
)
new_content[window_start:window_end], search_pattern, replacement, replace_all=False)
if count > 0:
new_content = new_content[:window_start] + window_new + new_content[window_end:]
error = None
if error:
# Mirror the validation-phase already-applied skip, or the two
# phases disagree and the whole patch fails here.
# Mirror validation's already-applied skip, else the two phases disagree and fail here.
if is_already_applied(new_content, search_pattern, replacement):
continue
return False, f"Could not apply hunk: {error}" + _no_match_hint(error, search_pattern, new_content), None, None
hint = _no_match_hint(error, search_pattern, new_content)
return _fail(f"Could not apply hunk: {error}" + hint)
# Pass pre_content to skip a redundant re-read inside write_file when supported.
if _write_file_accepts_pre_content(file_ops):
write_result = file_ops.write_file(op.file_path, new_content, pre_content=current_content)
else:
write_result = file_ops.write_file(op.file_path, new_content)
if write_result.error:
return False, write_result.error, None, None
diff = ''.join(difflib.unified_diff(
current_content.splitlines(keepends=True), new_content.splitlines(keepends=True),
fromfile=f"a/{op.file_path}", tofile=f"b/{op.file_path}"))
return True, diff, getattr(write_result, "lsp_diagnostics", None), getattr(write_result, "lint", None)
extra = {"pre_content": current_content} if _write_file_accepts_pre_content(file_ops) else {}
write_result = file_ops.write_file(op.file_path, new_content, **extra)
return _written(write_result, _unified_diff(op.file_path, current_content, new_content))
# operation -> (handler, verb for error text, files_* bucket)
+29 -48
View File
@@ -363,8 +363,7 @@ class ProcessRegistry:
the gateway asyncio loop (watchers, reset checks) and the cleanup thread."""
_SHELL_NOISE_SUBSTRINGS = (
"no job control in this shell",
"cannot set terminal process group",
"no job control in this shell", "cannot set terminal process group",
"tcsetattr: Inappropriate ioctl for device")
def __init__(self):
@@ -394,10 +393,8 @@ class ProcessRegistry:
self._poll_observed: set = set()
# Global watch-match circuit breaker across all sessions.
self._global_watch_lock = threading.Lock()
self._global_watch_window_start: float = 0.0
self._global_watch_window_hits: int = 0
self._global_watch_tripped_until: float = 0.0
self._global_watch_suppressed_during_trip: int = 0
self._global_watch_window_start = self._global_watch_tripped_until = 0.0
self._global_watch_window_hits = self._global_watch_suppressed_during_trip = 0
# Driver-installed sinks (desktop gateway): on_output(session, chunk) streams
# live output from reader threads; on_close(session_or_none, process_id) drops
# a read-only terminal tab without killing the process.
@@ -432,16 +429,13 @@ class ProcessRegistry:
# avoids stale notifications minutes after the process ended.
if session.exited:
return
matched_lines = []
matched_pattern = None
for line in new_text.splitlines():
pat = next((p for p in session.watch_patterns if p in line), None) # one match per line
if pat is not None:
matched_lines.append(line.rstrip())
if matched_pattern is None:
matched_pattern = pat
if not matched_lines:
hits = [ # (first matching pattern, line) — one match per line
(next(p for p in session.watch_patterns if p in line), line.rstrip())
for line in new_text.splitlines() if any(p in line for p in session.watch_patterns)]
if not hits:
return
matched_pattern = hits[0][0]
matched_lines = [line for _, line in hits]
now = time.time()
with session._lock:
if session._watch_cooldown_until and now < session._watch_cooldown_until:
@@ -535,45 +529,41 @@ class ProcessRegistry:
In cooldown: drop and count. Otherwise slide the rolling window; exceeding
the cap trips the breaker for WATCH_GLOBAL_COOLDOWN_SECONDS with ONE
"tripped" summary, and the cooldown's end emits ONE "released" summary."""
release_msg = trip_msg = None
events = [] # summary events, queued outside the lock
with self._global_watch_lock:
# Handle cooldown expiry first so we can emit the release summary.
if self._global_watch_tripped_until and now >= self._global_watch_tripped_until:
suppressed = self._global_watch_suppressed_during_trip
self._global_watch_tripped_until = 0.0
self._global_watch_suppressed_during_trip = 0
self._global_watch_window_start = now
self._global_watch_window_hits = 0
self._global_watch_window_start, self._global_watch_window_hits = now, 0
if suppressed > 0:
release_msg = self._global_watch_event(
events.append(self._global_watch_event(
"watch_overflow_released",
f"Watch-pattern notifications resumed. "
f"{suppressed} match event(s) were suppressed during the flood.",
suppressed=suppressed)
suppressed=suppressed))
if self._global_watch_tripped_until and now < self._global_watch_tripped_until:
# Still in cooldown — drop and count.
self._global_watch_suppressed_during_trip += 1
admit = False
else:
if now - self._global_watch_window_start >= WATCH_GLOBAL_WINDOW_SECONDS:
self._global_watch_window_start = now
self._global_watch_window_hits = 0
self._global_watch_window_start, self._global_watch_window_hits = now, 0
admit = self._global_watch_window_hits < WATCH_GLOBAL_MAX_PER_WINDOW
if admit:
self._global_watch_window_hits += 1
else:
self._global_watch_tripped_until = now + WATCH_GLOBAL_COOLDOWN_SECONDS
self._global_watch_suppressed_during_trip += 1
trip_msg = self._global_watch_event(
events.append(self._global_watch_event(
"watch_overflow_tripped",
f"Watch-pattern overflow: >{WATCH_GLOBAL_MAX_PER_WINDOW} "
f"notifications in {WATCH_GLOBAL_WINDOW_SECONDS}s across all processes. "
f"Suppressing further watch_match events for "
f"{WATCH_GLOBAL_COOLDOWN_SECONDS}s.")
# Queue summary events outside the lock.
for msg in (release_msg, trip_msg):
if msg is not None:
self.completion_queue.put(msg)
f"{WATCH_GLOBAL_COOLDOWN_SECONDS}s."))
for msg in events:
self.completion_queue.put(msg)
return admit
@staticmethod
@@ -589,11 +579,9 @@ class ProcessRegistry:
@staticmethod
def _safe_host_start_time(pid: Optional[int]) -> Optional[int]:
"""Kernel start ticks for a host PID, or None when unavailable."""
if not pid:
return None
try:
from gateway.status import get_process_start_time
return get_process_start_time(pid)
return get_process_start_time(pid) if pid else None
except Exception:
return None
@@ -604,9 +592,8 @@ class ProcessRegistry:
process (seen in the wild: a browser's session leader tree-killed). The kernel
start time captured at spawn must match the live one; with no baseline
(legacy checkpoints, no ``/proc``) degrade to a bare liveness check."""
if not cls._is_host_pid_alive(pid):
return False
return expected_start is None or cls._safe_host_start_time(pid) == expected_start
return cls._is_host_pid_alive(pid) and (
expected_start is None or cls._safe_host_start_time(pid) == expected_start)
def _refresh_detached_session(self, session: Optional[ProcessSession]) -> Optional[ProcessSession]:
"""Update recovered host-PID sessions when the underlying process has exited."""
@@ -619,9 +606,8 @@ class ProcessRegistry:
with session._lock:
if session.exited:
return session
session.exited = True
# No waitable handle survives recovery, so the real exit code is unknown.
session.exit_code = None
session.exited, session.exit_code = True, None
self._move_to_finished(session)
return session
@@ -630,9 +616,7 @@ class ProcessRegistry:
"""True if a psutil.Process is running and not a zombie (already dead, just unreaped)."""
try:
import psutil
if not proc.is_running():
return False
return proc.status() != psutil.STATUS_ZOMBIE
return proc.is_running() and proc.status() != psutil.STATUS_ZOMBIE
except Exception:
return False
@@ -921,8 +905,8 @@ class ProcessRegistry:
after the direct child exits (mirrors ``environments/base.py::_wait_for_process``).
Windows pipes lack select(), so the lazy ``_reconcile_local_exit`` is the net."""
first_chunk = True
# A multibyte UTF-8 char split across read1() chunks would become U+FFFD with
# stateless decoding; the incremental decoder holds the partial sequence.
# A split multibyte UTF-8 char would become U+FFFD with stateless decoding; the
# incremental decoder holds the partial sequence until the rest arrives.
decoder = codecs.getincrementaldecoder("utf-8")(errors="replace")
def _append_chunk(chunk: str):
@@ -1648,20 +1632,18 @@ class ProcessRegistry:
child's blocking line read (``readline()``, Go ``bufio.Scanner``) never returns
and the process hangs looking healthy. ``\\r\\n`` gives it both; POSIX keeps ``\\n``."""
session = self.get(session_id)
is_windows_pty = bool(_IS_WINDOWS and session is not None and session._pty)
return self.write_stdin(session_id, data + ("\r\n" if is_windows_pty else "\n"))
return self.write_stdin(session_id, data + ("\r\n" if _IS_WINDOWS and session and session._pty else "\n"))
def request_close_terminal(self, session_id: str) -> dict:
"""Ask the desktop GUI to close this process's read-only terminal tab. Does NOT
kill the process — output keeps buffering and the tab can be reopened from the
status stack. Errors when no UI close sink is wired."""
sink = self.on_close
if sink is None:
if self.on_close is None:
return {"status": "error", "error": "close_terminal is only available in the Hermes desktop app."}
# The session may already be finished (or pruned) — the tab can still
# linger and be closed, so a missing session is not an error here.
try:
sink(self.get(session_id), session_id)
self.on_close(self.get(session_id), session_id)
except Exception as e:
return {"status": "error", "error": str(e)}
return {
@@ -1836,10 +1818,9 @@ class ProcessRegistry:
recovered = 0
unresolved_scope_entries: List[Dict[str, Any]] = []
for entry in entries:
pid = entry.get("pid")
pid, pid_scope = entry.get("pid"), entry.get("pid_scope", "host")
if not pid:
continue
pid_scope = entry.get("pid_scope", "host")
if pid_scope != "host": # in-sandbox PIDs mean nothing once the env handle is gone
logger.info(
"Skipping recovery for non-host process: %s (pid=%s, scope=%s)",
+79 -133
View File
@@ -1,14 +1,13 @@
"""Human-readable rendering of background-process notification events.
Events come off ``ProcessRegistry.completion_queue`` (completion, watch_match,
watch_disabled, watch_overflow_*, async_delegation) and are turned into the
``[IMPORTANT: ...]`` / ``[ASYNC DELEGATION ...]`` text injected into the agent
conversation by the CLI drain loop, the gateway, and the TUI.
"""
"""Human-readable rendering of ``ProcessRegistry.completion_queue`` events (completion,
watch_match, watch_disabled, watch_overflow_*, async_delegation) into the
``[IMPORTANT: ...]`` / ``[ASYNC DELEGATION ...]`` text the CLI drain loop, gateway and
TUI inject into the agent conversation."""
import time
from contextlib import suppress
_DONE = ("completed", "success")
def _format_age(seconds: float) -> str:
"""Human-friendly elapsed string ('18m', '2h3m', '45s')."""
@@ -26,34 +25,28 @@ def _format_age(seconds: float) -> str:
def _model_not_found_patterns() -> "list[str]":
"""Model-not-found phrases from ``agent.error_classifier`` (same classification
the failover path uses, no hand-copied list to drift); a minimal built-in set
if the import fails so per-task blocks are never hidden."""
"""Model-not-found phrases from ``agent.error_classifier`` (the failover path's
own list, so nothing drifts); a minimal built-in set if the import fails."""
try:
from agent.error_classifier import _MODEL_NOT_FOUND_PATTERNS
return list(_MODEL_NOT_FOUND_PATTERNS)
except Exception:
return ["is not a valid model", "model not found", "model_not_found"]
def _delegation_config() -> dict:
"""Active delegation config (model/provider/fallbacks); ``{}`` on any error.
Lazy ``tools.delegate_tool._load_config`` so the renderer sees the dispatcher's
model/provider without importing the heavy delegation module at import time."""
"""Active delegation config; ``{}`` on any error. Lazy: delegate_tool is heavy."""
try:
from tools.delegate_tool import _load_config as _cfg
return _cfg() or {}
except Exception:
return {}
def _delegation_model_not_found(results, config) -> bool:
"""True when a result reflects a config-level model_not_found rejection.
Requires both a model-not-found phrase AND the currently-configured model
name in the same error/summary text, so a stale task failing on a
different (removed) model is not mis-attributed to the config."""
"""True when a result reflects a config-level model_not_found rejection: needs a
model-not-found phrase AND the currently-configured model name in the same text,
so a stale task failing on a removed model is not mis-attributed to the config."""
model = str((config or {}).get("model") or "").lower()
if not model:
return False
@@ -74,11 +67,9 @@ def _delegation_model_not_found_notice(results) -> "list[str] | None":
f'"{model}" was rejected by provider "{provider}" '
"(HTTP 400: not a valid model ID).",
"Every task in this batch failed for this reason before doing any work.",
"Check Settings → Advanced → Subagent Model (or: hermes config get delegation.model).",
]
"Check Settings → Advanced → Subagent Model (or: hermes config get delegation.model)."]
with suppress(Exception):
from hermes_cli.fallback_config import get_fallback_chain
if not get_fallback_chain(config):
lines.append("No fallback chain is configured, so no failover was attempted.")
return lines
@@ -94,71 +85,59 @@ def _is_truncated(entry: dict) -> bool:
return bool(entry.get("truncated") or entry.get("exit_reason") == "max_iterations")
def _header_lines(evt: dict, title: str, intro: str, completed_at: float) -> "list[str]":
"""Shared preamble: title, intro, blank, dispatch time and task-source lines."""
def _notice_lines(results) -> "list[str]":
"""Blank + model_not_found notice block, or [] when the notice does not apply."""
notice = _delegation_model_not_found_notice(results)
return ["", *notice] if notice else []
def _preamble(evt: dict, title: str, intro: str, completed_at: float, *, with_goal: bool) -> "list[str]":
"""Shared preamble: title, intro, blank, dispatch time, [goal], context/toolsets, role+model."""
lines = [title, intro, ""]
dispatched_at = evt.get("dispatched_at")
if isinstance(dispatched_at, (int, float)):
ts = time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(dispatched_at))
lines.append(f"Dispatched: {ts} ({_format_age(completed_at - dispatched_at)} ago)")
return lines
def _task_source_lines(evt: dict) -> "list[str]":
lines = []
if with_goal:
lines.append(f"Original goal: {evt.get('goal', '') or ''}")
if evt.get("context"):
lines.append(f"Context you provided: {evt['context']}")
if evt.get("toolsets"):
lines.append(f"Toolsets: {', '.join(evt['toolsets'])}")
lines.append(f"Role: {evt.get('role') or 'leaf'} Model: {evt.get('model') or '?'}")
return lines
def _role_model(evt: dict) -> str:
return f"Role: {evt.get('role') or 'leaf'} Model: {evt.get('model') or '?'}"
def _format_batch_delegation(evt: dict, deleg_id: str, completed_at: float) -> str:
"""Consolidated block for a delegate_task fan-out that finished as one unit."""
results = evt.get("results") or []
goals = evt.get("goals") or []
results, goals = evt.get("results") or [], evt.get("goals") or []
n = len(results) if results else len(goals)
total_dur = evt.get("total_duration_seconds", evt.get("duration_seconds", "?"))
error = evt.get("error")
lines = _header_lines(
lines = _preamble(
evt,
f"[ASYNC DELEGATION BATCH COMPLETE — {deleg_id}]",
f"A background fan-out of {n} subagent(s) you dispatched earlier "
"has finished. All ran in parallel and waited on each other; their "
"consolidated results are below. You may have moved on since "
"dispatching — act on these or re-dispatch if things have changed.",
completed_at)
lines.extend(_task_source_lines(evt))
lines.append(f"{_role_model(evt)} Total duration: {total_dur}s")
if error and not results:
lines += ["--- ERROR ---", f"The batch did not complete successfully: {error}"]
completed_at, with_goal=False)
lines[-1] += f" Total duration: {evt.get('total_duration_seconds', evt.get('duration_seconds', '?'))}s"
if evt.get("error") and not results:
lines += ["--- ERROR ---", f"The batch did not complete successfully: {evt['error']}"]
return "\n".join(lines)
# Config-level rejection notice BEFORE the per-task wall — a rejected
# delegation model fails every task identically and must not stay buried.
_notice = _delegation_model_not_found_notice(results)
if _notice:
lines += ["", *_notice]
lines += _notice_lines(results)
for r in sorted(results, key=lambda x: x.get("task_index", 0)):
idx = r.get("task_index", 0)
r_status = r.get("status", "?")
r_summary = r.get("summary")
r_error = r.get("error")
idx, r_truncated = r.get("task_index", 0), _is_truncated(r)
r_status, r_summary, r_error = r.get("status", "?"), r.get("summary"), r.get("error")
r_goal = goals[idx] if idx < len(goals) else r.get("goal", "")
r_truncated = _is_truncated(r)
icon = "⚠" if r_truncated else ("✓" if r_status in ("completed", "success") else "✗")
header = f"--- {icon} TASK {idx + 1}/{n}" + (f": {r_goal}" if r_goal else "") + f" (status={r_status}"
if r.get("api_calls"):
header += f", api_calls={r['api_calls']}"
if r.get("duration_seconds") is not None:
header += f", {r['duration_seconds']}s"
if r_truncated:
header += ", TRUNCATED: hit max_iterations — work may be incomplete"
icon = "⚠" if r_truncated else ("✓" if r_status in _DONE else "✗")
header = (f"--- {icon} TASK {idx + 1}/{n}" + (f": {r_goal}" if r_goal else "") + f" (status={r_status}"
+ (f", api_calls={r['api_calls']}" if r.get("api_calls") else "")
+ (f", {r['duration_seconds']}s" if r.get("duration_seconds") is not None else "")
+ (", TRUNCATED: hit max_iterations — work may be incomplete" if r_truncated else ""))
lines += ["", header + ") ---"]
if r_status in ("completed", "success") and r_summary:
if r_status in _DONE and r_summary:
if r_truncated:
lines.append(_TRUNCATED_SUMMARY_NOTE)
lines.append(r_summary)
@@ -174,39 +153,27 @@ def _format_batch_delegation(evt: dict, deleg_id: str, completed_at: float) -> s
def _format_async_delegation(evt: dict) -> str:
"""Format an async-delegation completion into a self-contained re-injection.
Carries the FULL original task source (goal, context, toolsets, role, model) plus
dispatch time, status, and the complete result summary: when this re-enters the
conversation the agent may be deep in unrelated context and must be able to use
the result OR re-dispatch without remembering why the subagent existed."""
"""Self-contained re-injection for an async-delegation completion: the FULL
original task source (goal, context, toolsets, role, model), dispatch time, status
and result, so an agent deep in unrelated context can act on it or re-dispatch."""
deleg_id = evt.get("delegation_id", "unknown")
completed_at = evt.get("completed_at") or time.time()
if evt.get("is_batch") or isinstance(evt.get("results"), list):
return _format_batch_delegation(evt, deleg_id, completed_at)
status = evt.get("status") or "completed"
summary = evt.get("summary")
error = evt.get("error")
status, summary, error = evt.get("status") or "completed", evt.get("summary"), evt.get("error")
truncated = _is_truncated(evt)
lines = _header_lines(
lines = _preamble(
evt,
f"[ASYNC DELEGATION COMPLETE — {deleg_id}]",
"A background subagent you dispatched earlier has finished. You may "
"have moved on since dispatching it; the full task source is below so "
"you can act on the result or re-dispatch if things have changed.",
completed_at)
lines.append(f"Original goal: {evt.get('goal', '') or ''}")
lines.extend(_task_source_lines(evt))
lines.append(_role_model(evt))
_notice = _delegation_model_not_found_notice([evt])
if _notice:
lines += ["", *_notice]
_trunc = " [TRUNCATED: hit max_iterations — work may be incomplete]" if truncated else ""
lines += [
f"Status: {status} API calls: {evt.get('api_calls', 0)} "
f"Duration: {evt.get('duration_seconds', '?')}s{_trunc}",
"--- RESULT ---",
]
if status in ("completed", "success") and summary:
completed_at, with_goal=True)
lines += _notice_lines([evt]) + [
f"Status: {status} API calls: {evt.get('api_calls', 0)} Duration: {evt.get('duration_seconds', '?')}s"
+ (" [TRUNCATED: hit max_iterations — work may be incomplete]" if truncated else ""),
"--- RESULT ---"]
if status in _DONE and summary:
if truncated:
lines.append(_TRUNCATED_SUMMARY_NOTE)
lines.append(summary)
@@ -215,92 +182,71 @@ def _format_async_delegation(evt: dict) -> str:
lines.append("The subagent was interrupted before completing" + (f": {error}" if error else "."))
else: # error / timeout / failed
lines.append(
f"The subagent did not complete successfully (status={status})." + (f"\n{error}" if error else "")
)
f"The subagent did not complete successfully (status={status})." + (f"\n{error}" if error else ""))
if summary:
lines += ["Partial output:", summary]
return "\n".join(lines)
def _delegation_attribution_line(evt: dict) -> "str | None":
"""One-line provenance for a subagent-owned process event, else None.
A background process a subagent started outlives the child and is routed to
the PARENT conversation, which otherwise sees an anonymous raw output wall.
Judged on ``owner_task_id`` (the raw spawning id) — ``task_id`` is the
container key and may be collapsed to the session key."""
"""One-line provenance for a subagent-owned process event, else None. Such a process
outlives the child and lands in the PARENT conversation, which would otherwise see an
anonymous output wall. Keyed on ``owner_task_id`` — ``task_id`` may be the session key."""
task_id = str(evt.get("owner_task_id") or evt.get("task_id") or "")
if not task_id.startswith("sa-"):
return None
info = None
with suppress(Exception):
from tools.delegate_tool import get_subagent_attribution
info = get_subagent_attribution(task_id)
if not info:
# Registry entry aged out — still attribute generically, not anonymously.
return f"Started by subagent {task_id} (delegate_task)."
goal = str(info.get("goal") or "").strip()
if len(goal) > 120:
goal = goal[:117] + "..."
deleg = info.get("delegation_id")
line = f"Started by subagent {task_id}" + (f" of delegation {deleg}" if deleg else "") + "."
if goal:
line += f' Task: "{goal}"'
return line
goal, deleg = str(info.get("goal") or "").strip(), info.get("delegation_id")
goal = goal[:117] + "..." if len(goal) > 120 else goal
return (f"Started by subagent {task_id}" + (f" of delegation {deleg}" if deleg else "") + "."
+ (f' Task: "{goal}"' if goal else ""))
_REASON_STATUS = {
"lost": "marked lost because the process backend disappeared",
"failed_start": "failed to start",
}
_REASON_STATUS = {"lost": "marked lost because the process backend disappeared", "failed_start": "failed to start"}
def _completion_status(evt: dict) -> str:
reason = evt.get("completion_reason") or "exited"
if reason == "killed":
return f"terminated by {evt.get('termination_source') or 'Hermes'}"
if reason in _REASON_STATUS:
return _REASON_STATUS[reason]
return "completed normally" if evt.get("exit_code", "?") == 0 else "exited"
return _REASON_STATUS.get(reason) or ("completed normally" if evt.get("exit_code", "?") == 0 else "exited")
def format_process_notification(evt: dict) -> "str | None":
"""Format a completion_queue event into an ``[IMPORTANT: ...]`` message."""
evt_type = evt.get("type", "completion")
_sid = evt.get("session_id", "unknown")
_cmd = evt.get("command", "unknown")
_attribution = _delegation_attribution_line(evt)
# watch_disabled and overflow events carry their own human-readable `message`;
# otherwise overflow events would fall through to the completion formatter as a
# phantom "process exited (exit code ?)".
if evt_type in ("watch_disabled", "watch_overflow_tripped", "watch_overflow_released"):
return f"[IMPORTANT: {evt.get('message', '')}]"
if evt_type == "watch_match":
_sup = evt.get("suppressed", 0)
text = f"[IMPORTANT: Background process {_sid} matched watch pattern \"{evt.get('pattern', '?')}\".\n"
if _attribution:
text += f"{_attribution}\n"
text += f"Command: {_cmd}\nMatched output:\n{evt.get('output', '')}"
if _sup:
text += f"\n({_sup} earlier matches were suppressed by rate limit)"
return text + "]"
if evt_type == "async_delegation":
return _format_async_delegation(evt)
_sid, _cmd = evt.get("session_id", "unknown"), evt.get("command", "unknown")
_attribution = _delegation_attribution_line(evt)
attribution = f"{_attribution}\n" if _attribution else ""
if evt_type == "watch_match":
_sup = evt.get("suppressed", 0)
return (
f"[IMPORTANT: Background process {_sid} matched watch pattern \"{evt.get('pattern', '?')}\".\n"
f"{attribution}Command: {_cmd}\nMatched output:\n{evt.get('output', '')}"
+ (f"\n({_sup} earlier matches were suppressed by rate limit)" if _sup else "") + "]")
_exit = evt.get("exit_code", "?")
_out = evt.get("output", "")
_signal = ", SIGTERM" if _exit in {-15, 143, "-15", "143"} else ""
text = f"[IMPORTANT: Background process {_sid} {_completion_status(evt)} (exit code {_exit}{_signal}).\n"
if _attribution:
text += f"{_attribution}\n"
# A subagent-owned process's full output belongs in the child's transcript,
# not as a raw wall in the parent — trim hard but keep enough tail to
# recognise failures.
if isinstance(_out, str) and len(_out) > 600:
_out = (
"...(output trimmed — subagent-owned process; see the "
"delegation's live transcript for full output)\n"
+ _out[-600:])
text += f"Command: {_cmd}\nOutput:\n{_out}]"
return text
# A subagent-owned process's full output belongs in the child's transcript, not as
# a raw wall in the parent — trim hard but keep enough tail to recognise failures.
if _attribution and isinstance(_out, str) and len(_out) > 600:
_out = (
"...(output trimmed — subagent-owned process; see the "
"delegation's live transcript for full output)\n"
+ _out[-600:])
return (
f"[IMPORTANT: Background process {_sid} {_completion_status(evt)} (exit code {_exit}{_signal}).\n"
f"{attribution}Command: {_cmd}\nOutput:\n{_out}]")
+23 -37
View File
@@ -4,6 +4,7 @@
costs nothing elsewhere (adapters expose reactions via ``send_message(action="react")``);
defaults to the triggering message and emits ``message.reaction`` for live painting."""
import contextlib
import json
from gateway.session_context import get_session_env
@@ -20,56 +21,41 @@ def _open_session_db():
return None
def _react(emoji: str, message_row_id, messages_back, *, db, session_key: str) -> str:
row_id = message_row_id
target_role = "user"
if row_id is None:
# Default: the latest user message; `messages_back` steps to earlier user turns
# (ids aren't visible to the model; "two messages ago" is how a person thinks).
back = max(0, int(messages_back or 0))
row_id = db.latest_message_row_id(session_key, role="user", offset=back)
if row_id is None:
return tool_error(
f"No user message found {back} back." if back else "No user message to react to yet.")
else:
target_role = db.get_message_role(session_key, int(row_id)) or "user"
try:
reactions = db.set_message_reaction(session_key, int(row_id), emoji or None, author="agent")
except Exception as exc:
return tool_error(f"Failed to set the reaction: {exc}")
if reactions is None:
return tool_error(f"Message {row_id} is not part of this conversation.")
# Paint it live; a missing bridge (non-desktop) is not an error — the reaction is
# persisted. `role` lets the renderer match a live message without a durable row id.
try:
desktop_ui.emit(
"message.reaction", {"row_id": int(row_id), "reactions": reactions, "role": target_role})
except Exception:
pass
return json.dumps({"success": True, "row_id": int(row_id), "reactions": reactions}, ensure_ascii=False)
def react_to_message_tool(emoji: str, message_row_id=None, messages_back=None) -> str:
"""Attach (or with an empty ``emoji`` retract) the agent's reaction."""
emoji = (emoji or "").strip()
session_key = get_session_env("HERMES_SESSION_KEY", "") or get_session_env("HERMES_SESSION_ID", "")
if not session_key:
return tool_error("No active session — reactions need a persisted conversation.")
db = _open_session_db()
if db is None:
return tool_error("Session storage is unavailable.")
try:
return _react(emoji, message_row_id, messages_back, db=db, session_key=session_key)
finally:
row_id, target_role = message_row_id, "user"
if row_id is None:
# Default: the latest user message; `messages_back` steps to earlier user turns
# (ids aren't visible to the model; "two messages ago" is how a person thinks).
back = max(0, int(messages_back or 0))
row_id = db.latest_message_row_id(session_key, role="user", offset=back)
if row_id is None:
return tool_error(f"No user message found {back} back." if back else "No user message to react to yet.")
else:
target_role = db.get_message_role(session_key, int(row_id)) or "user"
try:
reactions = db.set_message_reaction(session_key, int(row_id), emoji or None, author="agent")
except Exception as exc:
return tool_error(f"Failed to set the reaction: {exc}")
if reactions is None:
return tool_error(f"Message {row_id} is not part of this conversation.")
# Paint it live; a missing bridge (non-desktop) is not an error — the reaction is
# persisted. `role` lets the renderer match a live message without a durable row id.
with contextlib.suppress(Exception):
desktop_ui.emit("message.reaction", {"row_id": int(row_id), "reactions": reactions, "role": target_role})
return json.dumps({"success": True, "row_id": int(row_id), "reactions": reactions}, ensure_ascii=False)
finally:
with contextlib.suppress(Exception):
from hermes_state import release_or_close
release_or_close(db)
except Exception:
pass
def check_react_requirements() -> bool:
+123 -194
View File
@@ -1,16 +1,14 @@
"""Stdlib document-to-text extraction for ``read_file``.
Jupyter, DOCX and XLSX need no dependencies. The optional ``firecrawl-anydoc`` package
(imports as ``anydoc``) widens coverage to legacy Office, OpenDocument, RTF, EPUB and PDF.
The stdlib extractors stay authoritative for their three formats so behavior is identical
with or without anydoc. Malformed documents raise :class:`ExtractionError`; callers then
fall back to normal text/binary handling.
"""
"""Document-to-text extraction for ``read_file``: stdlib Jupyter/DOCX/XLSX (always
authoritative for those three), plus legacy Office/OpenDocument/RTF/EPUB/PDF when the
optional ``firecrawl-anydoc`` package (imports as ``anydoc``) is installed. Malformed
documents raise :class:`ExtractionError`; callers fall back to text/binary handling."""
from __future__ import annotations
import contextlib
import functools
import importlib
import itertools
import json
import os
import posixpath
@@ -29,12 +27,10 @@ __all__ = ["EXTRACTABLE_EXTENSIONS", "ExtractionError", "extract_document_bytes"
"extract_document_text", "is_extractable_document"]
EXTRACTABLE_EXTENSIONS = frozenset({".ipynb", ".docx", ".xlsx"})
# Formats handled only when the optional anydoc converter is installed.
ANYDOC_EXTENSIONS = frozenset({
".doc", ".docm", ".ppt", ".pps", ".pot", ".pptx", ".pptm", ".ppsx", ".ppsm",
".xls", ".xlsm", ".xlsb", ".odt", ".ods", ".odp", ".rtf", ".epub", ".pdf"})
# anydoc loads whole files with no streaming and the read_file char budget only applies
# after conversion — cap the input size.
# anydoc loads whole files (no streaming); read_file's char budget applies only post-conversion.
MAX_ANYDOC_BYTES = 50 * 1024 * 1024
MAX_DOCUMENT_BYTES = 50 * 1024 * 1024
_MAX_XLSX_ROWS_PER_SHEET = 5000
@@ -59,15 +55,15 @@ def _extension(path: str) -> str:
_ANYDOC_UNSET = object()
_anydoc_module: Any = _ANYDOC_UNSET
_anydoc_lock = threading.Lock()
# After a failed load, wait this long before retrying: the attempt can shell out
# to pip, so retrying every call would hammer the network where install can't succeed.
# Cooldown after a failed load: the attempt can shell out to pip, so retrying every call would
# hammer the network where install can't succeed.
ANYDOC_RETRY_SECONDS = 300.0
_anydoc_failed_at: Optional[float] = None
def _anydoc() -> Optional[Any]:
"""Lazily import the optional anydoc converter; None when unavailable. A failed load is
retried after ANYDOC_RETRY_SECONDS so one transient pip/network blip does not stick."""
"""Lazily import the optional anydoc converter (None when unavailable; failures retried after
ANYDOC_RETRY_SECONDS so one transient pip/network blip does not stick)."""
global _anydoc_module, _anydoc_failed_at
if _anydoc_module is not _ANYDOC_UNSET:
return _anydoc_module
@@ -79,8 +75,7 @@ def _anydoc() -> Optional[Any]:
return None
try:
from tools.lazy_deps import ensure as _lazy_ensure
# prompt=False: read_file must never block on an install prompt.
_lazy_ensure("tool.doc_extract", prompt=False)
_lazy_ensure("tool.doc_extract", prompt=False) # read_file must never block on a prompt
_anydoc_module = importlib.import_module("anydoc")
except Exception: # install failure, ImportError or a broken native binding
_anydoc_failed_at = time.monotonic()
@@ -101,23 +96,19 @@ def _check_size(size: int, limit: int) -> None:
@contextlib.contextmanager
def _temp_copy(data: bytes, suffix: str) -> Iterator[str]:
"""Materialize backend bytes in a private host temp file; removed even when parsing fails."""
temp_path = ""
with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as fh:
fh.write(data)
try:
with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as fh:
fh.write(data)
temp_path = fh.name
yield temp_path
yield fh.name
finally:
if temp_path:
with contextlib.suppress(OSError):
os.unlink(temp_path)
with contextlib.suppress(OSError):
os.unlink(fh.name)
def extract_document_text(path: str) -> str:
ext = _extension(path)
extractor = _STDLIB_EXTRACTORS.get(ext)
if extractor is not None:
return extractor(path)
if ext in _STDLIB_EXTRACTORS:
return _STDLIB_EXTRACTORS[ext](path)
if ext in ANYDOC_EXTENSIONS:
return _extract_anydoc(path)
raise ExtractionError(f"Unsupported document type: {path!r}")
@@ -131,14 +122,12 @@ def extract_document_bytes(data: bytes, path: str) -> str:
return _extract_anydoc_bytes(data, path)
if ext not in EXTRACTABLE_EXTENSIONS:
raise ExtractionError(f"Unsupported document type: {path!r}")
# The stdlib extractors are path-oriented.
with _temp_copy(data, ext) as temp_path:
return extract_document_text(temp_path)
with _temp_copy(data, ext) as temp_path: # the stdlib extractors are path-oriented
return _STDLIB_EXTRACTORS[ext](temp_path)
def _anydoc_missing_error(path: str) -> str:
"""Teaching error for anydoc-gated formats (deliberately absent from the schema so only
sessions that hit one pay for it)."""
"""Teaching text for anydoc-gated formats (not in the schema: only sessions hitting one pay)."""
return (
f"Cannot convert {path!r}: this format needs the optional anydoc "
"converter, which is not installed (install blocked or first "
@@ -149,10 +138,10 @@ def _anydoc_missing_error(path: str) -> str:
def _hosted_ocr_config() -> tuple:
"""Resolve hosted-OCR settings: (enabled, api_key, api_url). Never raises; no network.
Maintainer decision: the ONLY route is a direct ``FIRECRAWL_API_KEY`` (anydoc defaults the
api_url); the Nous gateway is NOT used — its Parse proxy live-probed broken (revisit when
it grows Parse support). ``file_tools.hosted_ocr: false`` disables even with a key."""
"""(enabled, api_key, api_url); never raises, no network. Maintainer decision: the ONLY route
is a direct ``FIRECRAWL_API_KEY`` (anydoc defaults api_url); the Nous gateway's Parse proxy
live-probed broken, so it is NOT used. ``file_tools.hosted_ocr: false`` disables even with a
key."""
api_key = os.environ.get("FIRECRAWL_API_KEY") or None
enabled = api_key is not None
with contextlib.suppress(Exception):
@@ -164,35 +153,30 @@ def _hosted_ocr_config() -> tuple:
def hosted_ocr_available() -> bool:
"""Public probe for read_file's schema line; a key failing at conversion time lands in
the NEEDS-OCR warning instead."""
"""Probe for read_file's schema line; a key failing at conversion time surfaces in NEEDS-OCR."""
return _hosted_ocr_config()[0]
def _needs_ocr_warning(path: str, pages, hosted_error: str = "") -> str:
"""Result text when anydoc raises NeedsOcrError and hosted OCR is off/failed. Hints at
CHECKING for an OCR skill (never names one) and never advertises the hosted_ocr knob."""
"""NeedsOcrError result when hosted OCR is off/failed; hints at CHECKING for an OCR skill
(never names one) and never advertises the hosted_ocr knob."""
page_list = ", ".join(str(p) for p in pages) if pages else "unknown"
msg = (
hosted = f"Hosted OCR was attempted and failed ({hosted_error}). " if hosted_error else ""
return (
f"[NEEDS OCR: pages {page_list} of this PDF are scanned images "
"with no text layer — their content is MISSING below. ")
if hosted_error:
msg += f"Hosted OCR was attempted and failed ({hosted_error}). "
msg += (
f"with no text layer — their content is MISSING below. {hosted}"
"If the missing pages matter: render just those pages with "
f"`pdftoppm -jpeg -r 150 -f <first> -l <last> '{path}' /tmp/page` "
"and inspect via vision_analyze, or check whether an OCR skill is "
"available (skills_list).")
return msg + "]\n"
"available (skills_list).]\n")
def _finalize_anydoc_text(text: Any, path: str, pdf_note: Callable[[], str]) -> str:
"""Normalize converter output and, for PDFs, PREPEND the coverage note (read_file
paginates: a footer may never be fetched). Covers PARTIAL gaps without NeedsOcrError."""
"""Normalize converter output; PDFs get the coverage note PREPENDED (read_file paginates, so a
footer may never be fetched) — this covers PARTIAL scan gaps that raise no NeedsOcrError."""
if not isinstance(text, str) or not text.strip():
raise ExtractionError("Document contains no extractable text")
note = pdf_note() if Path(path).suffix.lower() == ".pdf" else ""
return (note or "") + text.rstrip("\n") + "\n"
return (pdf_note() if Path(path).suffix.lower() == ".pdf" else "") + text.rstrip("\n") + "\n"
def _ocr_scanned_pdf(mod: Any, path: str, exc: BaseException) -> str:
@@ -206,40 +190,33 @@ def _ocr_scanned_pdf(mod: Any, path: str, exc: BaseException) -> str:
return mod.to_markdown(path, ocr="hosted", **extra).rstrip("\n") + "\n"
except Exception as hosted_exc: # noqa: BLE001
hosted_error = f"{type(hosted_exc).__name__}: {hosted_exc}"
# No route / disabled / hosted failed: whole doc is scans — the warning IS the result.
return _needs_ocr_warning(path, pages, hosted_error)
def _require_anydoc(path: str) -> Any:
mod = _anydoc()
if mod is None:
raise ExtractionError(_anydoc_missing_error(path))
return mod
return _needs_ocr_warning(path, pages, hosted_error) # whole doc is scans: the warning IS it
def _extract_anydoc(path: str) -> str:
mod = _require_anydoc(path)
try:
size = os.path.getsize(path)
except OSError as exc:
raise ExtractionError(str(exc)) from exc
_check_size(size, MAX_ANYDOC_BYTES)
mod = _anydoc()
if mod is None:
raise ExtractionError(_anydoc_missing_error(path))
try:
_check_size(os.path.getsize(path), MAX_ANYDOC_BYTES)
text = mod.to_markdown(path)
except ExtractionError:
raise
except OSError as exc:
raise ExtractionError(str(exc)) from exc
except Exception as exc:
needs_ocr = getattr(mod, "NeedsOcrError", None)
if needs_ocr is not None and isinstance(exc, needs_ocr):
return _ocr_scanned_pdf(mod, path, exc)
# anydoc raises one ConvertError subclass per failure mode (Unsupported, Malformed,
# Encrypted, ResourceLimit, MissingPart); all mean "no meaningful text".
# Any ConvertError subclass (Unsupported/Malformed/Encrypted/...) = "no meaningful text".
raise ExtractionError(f"{type(exc).__name__}: {exc}") from exc
return _finalize_anydoc_text(text, path, lambda: _pdf_coverage_note(path))
def _extract_anydoc_bytes(data: bytes, path: str) -> str:
mod = _require_anydoc(path)
mod = _anydoc()
if mod is None:
raise ExtractionError(_anydoc_missing_error(path))
_check_size(len(data), MAX_ANYDOC_BYTES)
try:
text = mod.to_markdown_bytes(data)
@@ -248,17 +225,13 @@ def _extract_anydoc_bytes(data: bytes, path: str) -> str:
return _finalize_anydoc_text(text, path, lambda: _pdf_coverage_note_from_bytes(data, path))
# ── Scanned-PDF coverage detection: text-layer extractors return nothing for scanned
# pages, so a mostly-scanned PDF converts "successfully" into headers with empty bodies —
# silent data loss. Count per-page text via pdftotext (form-feed separated) and warn.
# ── Scanned-PDF coverage: text-layer extractors return nothing for scanned pages, so a mostly
# scanned PDF converts "successfully" into silent data loss. Count per-page text via pdftotext.
PDF_EMPTY_PAGE_CHARS = 20 # fewer extracted chars than this = empty page
# Warn when empty pages reach both MIN_EMPTY and MIN_RATIO, or ABSOLUTE_EMPTY alone.
PDF_COVERAGE_MIN_EMPTY = 2
PDF_COVERAGE_MIN_RATIO = 0.2
PDF_COVERAGE_ABSOLUTE_EMPTY = 10
PDF_COVERAGE_MIN_EMPTY, PDF_COVERAGE_MIN_RATIO, PDF_COVERAGE_ABSOLUTE_EMPTY = 2, 0.2, 10
PDF_PAGE_SCAN_TIMEOUT = 20.0
# Cap the per-gap breakdown so alternating text/scan pages can't balloon the warning.
PDF_GAP_MAP_MAX_ENTRIES = 20
PDF_GAP_MAP_MAX_ENTRIES = 20 # cap so alternating text/scan pages can't balloon the warning
_GAP_CONTEXT_CHARS = 60
@@ -279,14 +252,10 @@ def _pdf_page_texts(path: str) -> Optional[list[str]]:
def _gap_map(counts: list[int], texts: list[str], empty: list[int]) -> str:
"""Per-gap breakdown, each empty range labeled with the last text seen before
it (usually a section header), so the agent can pick WHICH gaps to OCR."""
ranges: list[list[int]] = [] # sorted 1-based page numbers -> [start, end] runs
for p in empty:
if ranges and p == ranges[-1][1] + 1:
ranges[-1][1] = p
else:
ranges.append([p, p])
"""Per-gap breakdown labeled with the text before each gap, so the agent picks which to OCR."""
# Sorted 1-based page numbers -> (start, end) runs; consecutive pages share ``page - index``.
runs = [list(g) for _k, g in itertools.groupby(enumerate(empty), lambda e: e[1] - e[0])]
ranges = [(run[0][1], run[-1][1]) for run in runs]
lines: list[str] = []
for a, b in ranges[:PDF_GAP_MAP_MAX_ENTRIES]:
label = ""
@@ -305,17 +274,17 @@ def _gap_map(counts: list[int], texts: list[str], empty: list[int]) -> str:
def _pdf_coverage_note(path: str, display_path: Optional[str] = None) -> str:
"""Warning header when many PDF pages produced no text, else ''. ``path`` is scanned
(may be a host temp file); ``display_path`` is what the recovery command shows."""
"""Warning header when many pages yielded no text, else ''. ``display_path`` (default ``path``,
which may be a host temp file) is what the recovery command shows."""
texts = _pdf_page_texts(path)
if not texts or len(texts) < 2:
return ""
counts = [len(page.strip()) for page in texts]
empty = [i + 1 for i, n in enumerate(counts) if n < PDF_EMPTY_PAGE_CHARS]
total = len(counts)
if len(empty) < PDF_COVERAGE_MIN_EMPTY or (
len(empty) / total < PDF_COVERAGE_MIN_RATIO and len(empty) < PDF_COVERAGE_ABSOLUTE_EMPTY
):
n_empty = len(empty)
enough = n_empty / total >= PDF_COVERAGE_MIN_RATIO or n_empty >= PDF_COVERAGE_ABSOLUTE_EMPTY
if n_empty < PDF_COVERAGE_MIN_EMPTY or not enough:
return ""
shown = display_path or path
return (
@@ -335,13 +304,10 @@ def _pdf_coverage_note(path: str, display_path: Optional[str] = None) -> str:
def _pdf_coverage_note_from_bytes(data: bytes, display_path: str) -> str:
"""Coverage note for backend-transferred PDF bytes via a host temp copy (pdftotext is
path-oriented); the recovery command still names ``display_path``."""
try:
with _temp_copy(data, ".pdf") as temp_path:
return _pdf_coverage_note(temp_path, display_path=display_path)
except OSError:
return ""
"""Coverage note for backend PDF bytes via a host temp copy (pdftotext needs a path)."""
with contextlib.suppress(OSError), _temp_copy(data, ".pdf") as temp_path:
return _pdf_coverage_note(temp_path, display_path=display_path)
return ""
def _joined(lines: list[str], empty_error: str) -> str:
@@ -352,11 +318,10 @@ def _joined(lines: list[str], empty_error: str) -> str:
def _source_text(source) -> str:
if isinstance(source, str):
return source
"""Notebook source/text fields are a str or a list of str fragments."""
if isinstance(source, list):
return "".join(item for item in source if isinstance(item, str))
return ""
source = "".join(item for item in source if isinstance(item, str))
return source if isinstance(source, str) else ""
def _human_size(n_bytes: int) -> str:
@@ -366,30 +331,24 @@ def _human_size(n_bytes: int) -> str:
def _base64_bytes(payload: str) -> int:
"""Approximate decoded size of a base64 payload (whitespace ignored)."""
clean = re.sub(r"[^0-9+/=A-Za-z]", "", payload)
padding = min(2, len(clean) - len(clean.rstrip("=")))
return max(0, (len(clean) * 3) // 4 - padding)
return max(0, (len(clean) * 3) // 4 - min(2, len(clean) - len(clean.rstrip("="))))
def _clean_stream_text(text: str) -> str:
"""Strip ANSI escapes; keep only the final ``\\r`` frame of each line (tqdm redraws)."""
from tools.ansi_strip import strip_ansi
lines = []
for line in strip_ansi(text).replace("\r\n", "\n").split("\n"):
frames = [frame for frame in line.split("\r") if frame]
lines.append(frames[-1] if frames else "")
return "\n".join(lines)
return "\n".join(([f for f in line.split("\r") if f] or [""])[-1]
for line in strip_ansi(text).replace("\r\n", "\n").split("\n"))
# Per-output-block truncation so one runaway training log cannot flood the extraction.
_MAX_OUTPUT_CHARS = 20_000
_MAX_OUTPUT_CHARS = 20_000 # per code cell, so one runaway training log cannot flood the extraction
# nbformat v3 stores mime data flat on the output dict under these keys.
_V3_MIME_KEYS = (("png", "image/png"), ("jpeg", "image/jpeg"), ("svg", "image/svg+xml"), ("html", "text/html"))
def _notebook_output_text(output: Any) -> str:
"""Render one notebook output as compact text: stream text, tracebacks and textual
results kept; token-heavy payloads (images, HTML, widgets) become sized placeholders.
Handles nbformat v4 and legacy v3 (``pyout``/``pyerr``) shapes."""
"""One notebook output as compact text: stream/traceback/textual results kept; token-heavy
payloads (images, HTML, widgets) become sized placeholders. Handles v4 and legacy v3 shapes."""
if not isinstance(output, dict):
return ""
otype = output.get("output_type")
@@ -397,52 +356,41 @@ def _notebook_output_text(output: Any) -> str:
body = _clean_stream_text(_source_text(output.get("text", "")))
return body if body.strip() else ""
if otype in {"error", "pyerr"}:
traceback = output.get("traceback")
tb_text = ""
if isinstance(traceback, list):
tb_text = _clean_stream_text(
"\n".join(line for line in traceback if isinstance(line, str)))
tb = output.get("traceback")
tb_text = _clean_stream_text("\n".join(filter(lambda l: isinstance(l, str), tb))
if isinstance(tb, list) else "")
header = f"Error: {output.get('ename', '')}: {output.get('evalue', '')}".rstrip(": ")
return f"{header}\n{tb_text}".rstrip()
if otype not in {"execute_result", "display_data", "pyout"}:
return ""
data = output.get("data")
if not isinstance(data, dict): # legacy v3: mime payloads sit flat on the output dict
data = {"text/plain": output["text"]} if isinstance(output.get("text"), (str, list)) else {}
data.update((mime, output[k]) for k, mime in _V3_MIME_KEYS if k in output)
if "application/vnd.jupyter.widget-view+json" in data:
return "[interactive widget — omitted]"
# Prefer readable text: models consume text/plain far better than markup.
for mime in ("text/plain", "text/markdown"):
if mime in data:
body = _clean_stream_text(_source_text(data[mime]))
if body.strip():
return body
for mime in ("text/plain", "text/markdown"): # models consume text far better than markup
body = _clean_stream_text(_source_text(data[mime])) if mime in data else ""
if body.strip():
return body
for mime, value in data.items():
if isinstance(mime, str) and mime.startswith("image/"):
size = _base64_bytes(_source_text(value))
return f"[{mime} output — {_human_size(size)}, omitted]"
return f"[{mime} output — {_human_size(_base64_bytes(_source_text(value)))}, omitted]"
if "text/html" in data:
html = _source_text(data["text/html"])
return f"[text/html output — {len(html):,} chars, omitted]"
mimes = ", ".join(str(m) for m in data) or "unknown"
return f"[{mimes} output — omitted]"
return f"[text/html output — {len(_source_text(data['text/html'])):,} chars, omitted]"
return f"[{', '.join(str(m) for m in data) or 'unknown'} output — omitted]"
def _notebook_outputs(cell: dict, jq_pointer: str = "", filename: str = "") -> str:
outputs = cell.get("outputs")
if not isinstance(outputs, list):
return ""
blocks = [text for text in (_notebook_output_text(o) for o in outputs) if text]
if not blocks:
return ""
joined = "\n".join(blocks)
if len(joined) > _MAX_OUTPUT_CHARS:
omitted = len(joined) - _MAX_OUTPUT_CHARS
hint = f" — full output: jq -r '{jq_pointer}' {filename}" if jq_pointer and filename else ""
joined = joined[:_MAX_OUTPUT_CHARS] + f"\n… [{omitted:,} output chars truncated{hint}]"
return joined
joined = "\n".join(filter(None, map(_notebook_output_text, outputs)))
if len(joined) <= _MAX_OUTPUT_CHARS:
return joined
hint = f" — full output: jq -r '{jq_pointer}' {filename}" if jq_pointer and filename else ""
omitted = len(joined) - _MAX_OUTPUT_CHARS
return joined[:_MAX_OUTPUT_CHARS] + f"\n… [{omitted:,} output chars truncated{hint}]"
_CELL_LABELS = {"markdown": "Markdown", "code": "Code", "raw": "Raw"}
@@ -459,7 +407,7 @@ def _extract_notebook(path: str) -> str:
raw_cells = nb.get("cells")
if isinstance(raw_cells, list):
cells = [(f".cells[{i}].outputs", cell) for i, cell in enumerate(raw_cells)]
else:
else: # nbformat v3: cells live under worksheets
cells = [
(f".worksheets[{wi}].cells[{ci}].outputs", cell)
for wi, ws in enumerate(nb.get("worksheets", [])) if isinstance(ws, dict)
@@ -470,18 +418,16 @@ def _extract_notebook(path: str) -> str:
counts = dict.fromkeys(_CELL_LABELS, 0)
out: list[str] = []
for jq_pointer, cell in cells:
if not isinstance(cell, dict):
continue
typ = cell.get("cell_type")
typ = cell.get("cell_type") if isinstance(cell, dict) else None
if typ not in _CELL_LABELS:
continue
counts[typ] += 1
suffix = f" {counts[typ]}" if typ != "raw" else ""
out.extend((f"# ── {_CELL_LABELS[typ]} cell{suffix} ──", _source_text(cell.get("source", "")).rstrip("\n"), ""))
if typ == "code":
rendered = _notebook_outputs(cell, jq_pointer, nb_name)
if rendered:
out.extend((f"# ── Output (cell {counts[typ]}) ──", rendered.rstrip("\n"), ""))
source = _source_text(cell.get("source", "")).rstrip("\n")
out += [f"# ── {_CELL_LABELS[typ]} cell{suffix} ──", source, ""]
rendered = _notebook_outputs(cell, jq_pointer, nb_name) if typ == "code" else ""
if rendered:
out += [f"# ── Output (cell {counts[typ]}) ──", rendered.rstrip("\n"), ""]
return _joined(out, "Notebook contains no readable cells")
@@ -491,24 +437,21 @@ def _open_zip(path: str, kind: str) -> Iterator[zipfile.ZipFile]:
try:
with zipfile.ZipFile(path) as zf:
yield zf
except zipfile.BadZipFile as exc:
raise ExtractionError(f"Not a valid {kind}: {exc}") from exc
except OSError as exc:
raise ExtractionError(str(exc)) from exc
except (zipfile.BadZipFile, OSError) as exc:
bad_zip = isinstance(exc, zipfile.BadZipFile)
raise ExtractionError(f"Not a valid {kind}: {exc}" if bad_zip else str(exc)) from exc
def _zip_xml(zf: zipfile.ZipFile, name: str, optional: bool = False) -> Any:
"""Parse a package part; ``optional`` parts yield None when absent or malformed."""
"""Parse a package part; ``optional`` parts yield an empty element when absent or malformed."""
try:
return ET.fromstring(zf.read(name))
except KeyError as exc:
except (KeyError, ET.ParseError) as exc:
if optional:
return None
raise ExtractionError(f"Missing {name}") from exc
except ET.ParseError as exc:
if optional:
return None
raise ExtractionError(f"Malformed XML in {name}: {exc}") from exc
return ET.Element("missing")
raise ExtractionError(
f"Missing {name}" if isinstance(exc, KeyError) else f"Malformed XML in {name}: {exc}"
) from exc
def _extract_docx(path: str) -> str:
@@ -518,9 +461,9 @@ def _extract_docx(path: str) -> str:
breaks = {f"{w}tab": "\t", f"{w}br": "\n", f"{w}cr": "\n"}
lines: list[str] = []
for para in root.iter(f"{w}p"):
buf = [(node.text or "") if node.tag == f"{w}t" else breaks.get(node.tag, "")
for node in para.iter()]
lines.extend("".join(buf).split("\n"))
text = "".join(
(n.text or "") if n.tag == f"{w}t" else breaks.get(n.tag, "") for n in para.iter())
lines.extend(text.split("\n"))
return _joined(lines, "DOCX contains no extractable text")
@@ -529,38 +472,27 @@ def _extract_xlsx(path: str) -> str:
with _open_zip(path, "XLSX") as zf:
names = set(zf.namelist())
sst = _zip_xml(zf, "xl/sharedStrings.xml", optional=True)
shared = [] if sst is None else [
"".join(t.text or "" for t in item.iter(f"{s}t")) for item in sst.iter(f"{s}si")]
shared = ["".join(t.text or "" for t in item.iter(f"{s}t")) for item in sst.iter(f"{s}si")]
rels_root = _zip_xml(zf, "xl/_rels/workbook.xml.rels", optional=True)
rels = {} if rels_root is None else {
rel.get("Id", ""): rel.get("Target", "")
for rel in rels_root.iter(f"{pr}Relationship") if rel.get("Id")}
rels = {rel.get("Id", ""): rel.get("Target", "")
for rel in rels_root.iter(f"{pr}Relationship") if rel.get("Id")}
out: list[str] = []
for sheet in _zip_xml(zf, "xl/workbook.xml").iter(f"{s}sheet"):
if sheet.get("state", "visible") in {"hidden", "veryHidden"}:
continue
target = rels.get(sheet.get(f"{r}id", ""), "").lstrip("/")
part = posixpath.normpath(target if target.startswith("xl/") else f"xl/{target}")
if part not in names:
if sheet.get("state", "visible") in {"hidden", "veryHidden"} or part not in names:
continue
try:
with contextlib.suppress(ET.ParseError):
rows = _sheet_rows(zf.read(part), shared)
except ET.ParseError:
continue
out.append(f"# ── Sheet: {sheet.get('name', 'Sheet')} ──")
out.extend("\t".join(row) for row in rows)
if not rows:
out.append("(empty)")
out.append("")
out += [f"# ── Sheet: {sheet.get('name', 'Sheet')} ──",
*(["\t".join(row) for row in rows] or ["(empty)"]), ""]
return _joined(out, "XLSX has no visible sheets with content")
def _col_index(ref: str) -> int:
idx = 0
for ch in ref:
if not ch.isalpha():
break
idx = idx * 26 + ord(ch.upper()) - ord("A") + 1
"""0-based column of a cell ref: ``A1`` -> 0, ``AB7`` -> 27 (bijective base-26 letters)."""
idx = functools.reduce(lambda acc, ch: acc * 26 + ord(ch.upper()) - ord("A") + 1,
itertools.takewhile(str.isalpha, ref), 0)
return max(idx - 1, 0)
@@ -568,18 +500,15 @@ def _sheet_rows(xml_bytes: bytes, shared: list[str]) -> list[list[str]]:
root = ET.fromstring(xml_bytes)
s = f"{{{_NS_S}}}"
rows: list[list[str]] = []
for row in root.iter(f"{s}row"):
if len(rows) >= _MAX_XLSX_ROWS_PER_SHEET:
break
for row in itertools.islice(root.iter(f"{s}row"), _MAX_XLSX_ROWS_PER_SHEET):
cells: dict[int, str] = {}
max_col = -1
for cell in row.iter(f"{s}c"):
col = _col_index(cell.get("r", "")) if cell.get("r") else max_col + 1
if col >= _MAX_XLSX_COLS:
continue
cells[col] = _cell_value(cell, shared, s)
max_col = max(max_col, col)
rows.append([cells.get(i, "") for i in range(max_col + 1)] if max_col >= 0 else [])
if col < _MAX_XLSX_COLS:
cells[col] = _cell_value(cell, shared, s)
max_col = max(max_col, col)
rows.append([cells.get(i, "") for i in range(max_col + 1)])
while rows and not any(value.strip() for value in rows[-1]):
rows.pop()
return rows
+97 -143
View File
@@ -1,11 +1,7 @@
"""Sanitize tool JSON schemas for broad LLM-backend compatibility.
Strict backends reject shapes OpenAI/Anthropic accept: llama.cpp's grammar converter fails
on ``{"type": "object"}`` without ``properties``, bare-string schemas and ``type`` arrays;
Anthropic rejects nullable ``anyOf`` at the top of ``input_schema``; Fireworks rejects
``default`` beside ``$ref``; Codex rejects top-level combinators. This module walks the
final schema tree on a deep copy and fixes only those shapes.
"""
"""Sanitize tool JSON schemas for strict LLM backends. llama.cpp's grammar converter fails on
``{"type": "object"}`` without ``properties``, bare-string schemas and ``type`` arrays; Anthropic
rejects nullable ``anyOf`` at the top of ``input_schema``; Fireworks rejects ``default`` beside
``$ref``; Codex rejects top-level combinators. Walks a deep copy and fixes only those shapes."""
from __future__ import annotations
@@ -16,15 +12,12 @@ from typing import Any, Callable
logger = logging.getLogger(__name__)
# Anthropic (and Bedrock/Vertex/Azure fronting it) reject property keys not matching this;
# one bad key anywhere in the tools array 400s the request (Cloudflare's MCP ships 61).
# Anthropic (and Bedrock/Vertex/Azure fronting it) reject property keys not matching this; one bad
# key anywhere in the tools array 400s the request (Cloudflare's MCP ships 61).
_PROP_KEY_RE = re.compile(r"^[a-zA-Z0-9_.-]{1,64}$")
_PROP_KEY_BAD_CHARS = re.compile(r"[^a-zA-Z0-9_.-]")
_UNION_KEYS = ("anyOf", "oneOf")
# Outer-node metadata carried onto a union's replacement node.
_UNION_META_KEYS = ("title", "description", "default", "examples")
_UNION_META_KEYS = ("title", "description", "default", "examples") # copied onto replacements
def _empty_object() -> dict:
@@ -47,56 +40,45 @@ def sanitize_property_key(key: str) -> str:
def _rename_property_keys(props: dict, path: str) -> dict[str, str]:
"""{original_key: conforming_key} for one properties dict (identity entries omitted).
Deterministic (insertion order, numeric suffixes on collision) so the model-visible
schema and the dispatch-time reverse map from the registry's original schema agree."""
Deterministic (insertion order, numeric suffixes on collision) so the model-visible schema
and the dispatch-time reverse map from the registry's original schema agree."""
renames: dict[str, str] = {}
taken = {k for k in props if _PROP_KEY_RE.match(k)}
for key in props:
if _PROP_KEY_RE.match(key):
continue
for key in (k for k in props if not _PROP_KEY_RE.match(k)):
base = sanitize_property_key(key)
candidate, i = base, 2
while candidate in taken:
suffix = f"_{i}"
candidate = base[: 64 - len(suffix)] + suffix
i += 1
candidate, i = base[: 64 - len(f"_{i}")] + f"_{i}", i + 1
taken.add(candidate)
renames[key] = candidate
logger.debug(
"schema_sanitizer[%s]: renamed property key %r -> %r "
"(provider key-pattern compat)", path, key, candidate)
logger.debug("schema_sanitizer[%s]: renamed property key %r -> %r "
"(provider key-pattern compat)", path, key, candidate)
return renames
def unrename_tool_args(params_schema: Any, args: Any) -> Any:
"""Map sanitized property keys in model-emitted args back to wire names. ``params_schema``
is the ORIGINAL registry schema; recurses into objects/array items; unknown keys pass."""
if not isinstance(params_schema, dict) or not isinstance(args, dict):
return args
props = params_schema.get("properties")
if not isinstance(props, dict):
"""Map sanitized keys in model-emitted args back to wire names. ``params_schema`` is the
ORIGINAL registry schema; recurses into objects/array items; unknown keys pass through."""
props = params_schema.get("properties") if isinstance(params_schema, dict) else None
if not isinstance(props, dict) or not isinstance(args, dict):
return args
reverse = {v: k for k, v in _rename_property_keys(props, "<unrename>").items()}
out = {}
for key, value in args.items():
orig = reverse.get(key, key)
subschema = props.get(orig)
if isinstance(subschema, dict):
if isinstance(value, dict):
value = unrename_tool_args(subschema, value)
elif isinstance(value, list) and isinstance(subschema.get("items"), dict):
value = [
unrename_tool_args(subschema["items"], item) if isinstance(item, dict) else item
for item in value]
sub = props.get(orig) if isinstance(props.get(orig), dict) else {}
if isinstance(value, dict) and sub:
value = unrename_tool_args(sub, value)
elif isinstance(value, list) and isinstance(sub.get("items"), dict):
value = [unrename_tool_args(sub["items"], item) if isinstance(item, dict) else item
for item in value]
out[orig] = value
return out
def sanitize_tool_schemas(tools: list[dict]) -> list[dict]:
"""Deep-copied ``tools`` (OpenAI format) with sanitized parameter schemas; safe to mutate."""
if not tools:
return tools
return [_sanitize_single_tool(tool) for tool in tools]
return [_sanitize_single_tool(tool) for tool in tools] if tools else tools
def _sanitize_single_tool(tool: dict) -> dict:
@@ -110,9 +92,7 @@ def _sanitize_single_tool(tool: dict) -> dict:
return out
name = fn.get("name", "<tool>")
top = _sanitize_node(params, path=name)
# Guarantee the top level is an object with properties.
if not isinstance(top, dict):
top = {}
top = top if isinstance(top, dict) else {} # guarantee an object with properties on top
top["type"] = "object"
if not isinstance(top.get("properties"), dict):
top["properties"] = {}
@@ -124,16 +104,14 @@ def _sanitize_single_tool(tool: dict) -> dict:
return out
# Sibling keywords strict JSON Schema validators reject alongside ``$ref``.
_REF_FORBIDDEN_SIBLINGS = frozenset({"default"})
_REF_FORBIDDEN_SIBLINGS = frozenset({"default"}) # strict validators reject these beside ``$ref``
def _strip_ref_siblings(node: Any) -> Any:
"""Recursively drop forbidden siblings of ``$ref`` (Fireworks rejects ``default`` there)."""
def strip(out: dict) -> dict:
if "$ref" in out:
for key in _REF_FORBIDDEN_SIBLINGS:
out.pop(key, None)
for key in _REF_FORBIDDEN_SIBLINGS if "$ref" in out else ():
out.pop(key, None)
return out
return _rewrite(node, strip)
@@ -143,18 +121,15 @@ _TOP_LEVEL_FORBIDDEN_KEYS = ("allOf", "anyOf", "oneOf", "enum", "not")
def _strip_top_level_combinators(params: dict, *, path: str = "<tool>") -> dict:
"""Drop combinators from the TOP level only (Codex rejects them there). They are usually
conditional-required hints, so validity is unchanged (handlers re-validate); nested
combinators are preserved."""
conditional-required hints, so validity is unchanged (handlers re-validate); nested ones
stay."""
if not isinstance(params, dict):
return params
out = dict(params)
for key in _TOP_LEVEL_FORBIDDEN_KEYS:
if key in out:
logger.debug(
"schema_sanitizer[%s]: stripped top-level %r combinator "
"from tool parameters (strict-backend compat)",
path, key)
out.pop(key, None)
for key in [k for k in _TOP_LEVEL_FORBIDDEN_KEYS if k in out]:
logger.debug("schema_sanitizer[%s]: stripped top-level %r combinator "
"from tool parameters (strict-backend compat)", path, key)
del out[key]
return out
@@ -163,20 +138,18 @@ def _is_null_branch(item: Any) -> bool:
def _carry_union_meta(outer: dict, replacement: dict, *, skip_default_on_ref: bool) -> None:
"""Copy outer-union metadata onto *replacement* where absent."""
"""Copy outer-union metadata onto *replacement* where absent (``default`` is illegal beside
``$ref`` on strict backends, hence ``skip_default_on_ref``)."""
for meta_key in _UNION_META_KEYS:
if meta_key in outer and meta_key not in replacement:
# ``default`` is illegal alongside ``$ref`` on strict backends.
if skip_default_on_ref and meta_key == "default" and "$ref" in replacement:
continue
if meta_key in outer and meta_key not in replacement and not (
skip_default_on_ref and meta_key == "default" and "$ref" in replacement):
replacement[meta_key] = outer[meta_key]
def strip_nullable_unions(schema: Any, *, keep_nullable_hint: bool = True) -> Any:
"""Collapse ``anyOf``/``oneOf`` nullable unions to the single non-null branch. MCP/Pydantic
optional fields arrive as ``{"anyOf": [{"type": "string"}, {"type": "null"}]}``; Anthropic
rejects the null branch and optionality is already in the parent's ``required``. Only
collapses when a null branch was dropped AND exactly one non-null branch survives.
"""Collapse ``anyOf``/``oneOf`` nullable unions (MCP/Pydantic optional fields) to the single
non-null branch: Anthropic rejects the null branch and optionality already lives in the parent's
``required``. Only when a null branch was dropped AND exactly one non-null branch survives.
``keep_nullable_hint`` sets ``nullable: true`` for runtime ``"null"`` → ``None`` coercion."""
def collapse(stripped: dict) -> Any:
for key in _UNION_KEYS:
@@ -189,7 +162,7 @@ def strip_nullable_unions(schema: Any, *, keep_nullable_hint: bool = True) -> An
if keep_nullable_hint:
replacement.setdefault("nullable", True)
_carry_union_meta(stripped, replacement, skip_default_on_ref=True)
return _rewrite(replacement, collapse)
return _rewrite(replacement, collapse) # the survivor may itself be a union
return stripped
return _rewrite(schema, collapse)
@@ -199,25 +172,23 @@ _CONST_PRIMITIVE_TYPES: dict[type, str] = {
def _const_branch_type(branch: Any) -> str | None:
"""JSON-Schema primitive type of a pure ``const`` branch, else None: a primitive ``const``
whose declared ``type`` (if any) matches; only ``title``/``description`` may accompany it."""
"""Primitive JSON-Schema type of a pure ``const`` branch (declared ``type``, if any, must match;
only ``title``/``description`` may accompany it), else None."""
if not isinstance(branch, dict) or "const" not in branch \
or set(branch) - {"const", "type", "title", "description"}:
return None
# ``type(value)`` lookup (not isinstance): bool is a subclass of int.
json_type = _CONST_PRIMITIVE_TYPES.get(type(branch["const"]))
return json_type if json_type is not None and branch.get("type") in (None, json_type) else None
return json_type if branch.get("type") in (None, json_type) else None
def collapse_const_unions(schema: Any) -> Any:
"""Collapse ``anyOf``/``oneOf`` unions of same-typed consts to ``enum`` (ported from
block/goose ``tool_schema_normalize.rs``, Apache-2.0). Rust/TS MCP servers emit
``{"anyOf": [{"const": "red"}, {"const": "green"}]}``, which strict backends mishandle.
Applies only when EVERY non-null branch is a pure ``const`` of one primitive type
(``bool`` never merges with ``integer``); one ``{"type": "null"}`` branch is tolerated
and recorded as ``nullable: true`` (null+multi-const unions land here, not in
``strip_nullable_unions``). Branch order is kept; outer metadata carried over; input
never mutated."""
"""Collapse ``anyOf``/``oneOf`` unions of same-typed consts (Rust/TS MCP servers emit
``{"anyOf": [{"const": "red"}, {"const": "green"}]}``) to ``enum``; ported from block/goose
``tool_schema_normalize.rs`` (Apache-2.0). Only when EVERY non-null branch is a pure ``const``
of one primitive type (``bool`` never merges with ``integer``); one ``{"type": "null"}`` branch
is tolerated as ``nullable: true``. Branch order kept; outer metadata carried; input never
mutated."""
def collapse(out: dict) -> Any:
for key in _UNION_KEYS:
variants = out.get(key)
@@ -241,53 +212,48 @@ def collapse_const_unions(schema: Any) -> Any:
_BARE_TYPE_NAMES = frozenset({"object", "string", "number", "integer", "boolean", "array", "null"})
# Keys whose values are NOT schemas (recursing would treat "path" as a bare-string schema);
# passed through (``required`` follows property renames).
# Values that are NOT schemas (recursing would treat a required name like "path" as a bare schema).
_NON_SCHEMA_LIST_KEYS = frozenset({"required", "enum", "examples", "dependentRequired"})
def _normalize_type_array(value: list, out: dict) -> None:
"""Normalize a ``type: [...]`` array into *out* (llama.cpp and Gemini-via-OpenAI reject
arrays). Per AI-SDK: one non-null type → ``type: X`` (+ ``nullable`` if ``null`` present);
several → ``anyOf`` of single-type schemas so EVERY branch survives; none → ``null`` or
the object fallback. Ported from anomalyco/opencode#31877."""
"""Normalize a ``type: [...]`` array into *out* (llama.cpp and Gemini-via-OpenAI reject arrays).
Per AI-SDK: one non-null type → ``type: X`` (+ ``nullable`` if ``null`` present); several →
``anyOf`` of single-type schemas so EVERY branch survives; none → ``null``/object fallback."""
has_null = "null" in value
non_null = [t for t in value if isinstance(t, str) and t != "null"]
if len(non_null) == 1:
out["type"] = non_null[0]
elif len(non_null) >= 2:
out["anyOf"] = [{"type": t} for t in non_null]
else:
if not non_null:
out["type"] = "null" if has_null else "object"
return
if len(non_null) == 1:
out["type"] = non_null[0]
else:
out["anyOf"] = [{"type": t} for t in non_null]
if has_null:
out.setdefault("nullable", True)
def _sanitize_node(node: Any, path: str) -> Any:
"""Recursively sanitize a JSON-Schema fragment: bare-string schemas become ``{"type":
<value>}`` (unknown strings → permissive object); object nodes gain ``properties: {}``;
``type`` arrays are normalized; property keys are renamed to the provider-safe pattern
and ``required`` follows, with entries missing from ``properties`` pruned."""
"""Recursively sanitize a JSON-Schema fragment: bare-string schemas → ``{"type": <value>}``
(unknown strings → permissive object); object nodes gain ``properties: {}``; ``type`` arrays
are normalized; property keys are renamed to the provider-safe pattern and ``required``
follows, with entries missing from ``properties`` pruned."""
if isinstance(node, str):
if node in _BARE_TYPE_NAMES:
logger.debug(
"schema_sanitizer[%s]: replacing bare-string schema %r with {'type': %r}",
path, node, node)
logger.debug("schema_sanitizer[%s]: replacing bare-string schema %r with {'type': %r}",
path, node, node)
return _empty_object() if node == "object" else {"type": node}
logger.debug(
"schema_sanitizer[%s]: replacing non-schema string %r "
"with empty object schema", path, node)
logger.debug("schema_sanitizer[%s]: replacing non-schema string %r "
"with empty object schema", path, node)
return _empty_object()
if isinstance(node, list):
return [_sanitize_node(item, f"{path}[{i}]") for i, item in enumerate(node)]
if not isinstance(node, dict):
return node
# Renames computed up front so ``required`` remaps even when it precedes ``properties``.
prop_renames: dict[str, str] = {}
if isinstance(node.get("properties"), dict):
prop_renames = _rename_property_keys(node["properties"], f"{path}.properties")
props_in = node.get("properties")
prop_renames = (_rename_property_keys(props_in, f"{path}.properties")
if isinstance(props_in, dict) else {})
out: dict = {}
for key, value in node.items():
if key == "type" and isinstance(value, list):
@@ -295,64 +261,55 @@ def _sanitize_node(node: Any, path: str) -> Any:
elif key in {"properties", "$defs", "definitions"} and isinstance(value, dict):
renames = prop_renames if key == "properties" else {}
out[key] = {
renames.get(sub_k, sub_k): _sanitize_node(sub_v, f"{path}.{key}.{renames.get(sub_k, sub_k)}")
for sub_k, sub_v in value.items()}
renames.get(k, k): _sanitize_node(v, f"{path}.{key}.{renames.get(k, k)}")
for k, v in value.items()}
elif key in {"items", "additionalProperties"}:
# Bool ``additionalProperties`` is valid; bool ``items`` is non-standard but preserved.
out[key] = value if isinstance(value, bool) else _sanitize_node(value, f"{path}.{key}")
elif key in {"anyOf", "oneOf", "allOf"} and isinstance(value, list):
out[key] = [_sanitize_node(item, f"{path}.{key}[{i}]") for i, item in enumerate(value)]
elif key in _NON_SCHEMA_LIST_KEYS:
if key == "required" and prop_renames and isinstance(value, list):
out[key] = [prop_renames.get(r, r) if isinstance(r, str) else r for r in value]
else:
out[key] = copy.deepcopy(value) if isinstance(value, (list, dict)) else value
else:
else: # anyOf/oneOf/allOf and any other nested schema recurse (lists index the path)
out[key] = _sanitize_node(value, f"{path}.{key}") if isinstance(value, (dict, list)) else value
if out.get("type") == "object":
if not isinstance(out.get("properties"), dict):
out["properties"] = {}
if isinstance(out.get("required"), list):
props = out.get("properties") or {}
valid = [r for r in out["required"] if isinstance(r, str) and r in props]
if not valid:
out.pop("required", None)
elif len(valid) != len(out["required"]):
valid = [r for r in out["required"] if isinstance(r, str) and r in out["properties"]]
if valid:
out["required"] = valid
else:
del out["required"]
return out
# ---- Reactive strips — only invoked after a backend rejects a schema -----------------------
# ---- Reactive strips — only invoked after a backend rejects a schema ----
_STRIP_ON_RECOVERY_KEYS = frozenset({"pattern", "format"})
_SCHEMA_MARKERS = frozenset({"type", "anyOf", "oneOf", "allOf"}) # a node with one IS a schema
def _dict_nodes(node: Any):
"""Pre-order walk yielding every dict node; each is yielded before its values are visited,
so a consumer may mutate it in place."""
"""Pre-order walk over every dict node (yielded before its values, so it may be mutated)."""
if isinstance(node, dict):
yield node
for v in node.values():
yield from _dict_nodes(v)
elif isinstance(node, list):
for item in node:
yield from _dict_nodes(item)
children = node.values() if isinstance(node, dict) else node if isinstance(node, list) else ()
for child in children:
yield from _dict_nodes(child)
def _reactive_strip(
tools: list[dict], strip_node: Callable[[dict], int], log_msg: str) -> tuple[list[dict], int]:
"""Apply *strip_node* (returns keywords removed) to every dict node of each tool's
parameters, in place. Handles OpenAI (``{"function": {"parameters": ..}}``) and Responses
(``{"name": .., "parameters": ..}``) formats. Returns ``(tools, stripped_count)``."""
if not tools:
return tools, 0
"""Apply *strip_node* (-> keywords removed) to every dict node of each tool's parameters, in
place; OpenAI (``{"function": {"parameters"}}``) and Responses (``{"parameters"}``) formats."""
stripped = 0
for tool in tools:
for tool in tools or ():
if not isinstance(tool, dict):
continue
fn = tool.get("function")
params = fn.get("parameters") if isinstance(fn, dict) else None
if not isinstance(params, dict):
params = tool.get("parameters")
params = params if isinstance(params, dict) else tool.get("parameters")
if isinstance(params, dict):
stripped += sum(strip_node(node) for node in _dict_nodes(params))
if stripped:
@@ -361,16 +318,14 @@ def _reactive_strip(
def strip_pattern_and_format(tools: list[dict]) -> tuple[list[dict], int]:
"""Strip ``pattern``/``format`` keywords from tool schemas, in place. Reactive: only after
llama.cpp's grammar converter rejected a schema (its regex engine is a small ECMAScript
subset), since cloud providers use these as prompting hints. Only strips beside
``type``/combinators, so a property literally *named* ``pattern`` is untouched."""
"""Strip ``pattern``/``format`` in place — reactive, only after llama.cpp's grammar converter
rejected a schema (its regex engine is a small ECMAScript subset); cloud providers use these as
prompting hints. Only beside ``type``/combinators, so a property *named* ``pattern`` stays."""
def _strip(node: dict) -> int:
if not ("type" in node or "anyOf" in node or "oneOf" in node or "allOf" in node):
return 0
hits = [k for k in node if k in _STRIP_ON_RECOVERY_KEYS]
is_schema = bool(node.keys() & _SCHEMA_MARKERS)
hits = [k for k in node if k in _STRIP_ON_RECOVERY_KEYS] if is_schema else []
for k in hits:
node.pop(k, None)
del node[k]
return len(hits)
return _reactive_strip(
tools, _strip,
@@ -379,13 +334,12 @@ def strip_pattern_and_format(tools: list[dict]) -> tuple[list[dict], int]:
def strip_slash_enum(tools: list[dict]) -> tuple[list[dict], int]:
"""Strip ``enum`` keywords whose string values contain ``/``, in place. xAI compiles
schemas to a grammar that rejects ``/`` in enum values (HTTP 400 before any token) —
typically MCP enums of HuggingFace model IDs. The constraint is a prompting hint only."""
"""Strip ``enum`` keywords whose string values contain ``/``, in place: xAI's grammar compiler
rejects them (HTTP 400 before any token) — typically MCP enums of HuggingFace model IDs."""
def _strip(node: dict) -> int:
enum_val = node.get("enum")
if isinstance(enum_val, list) and any(isinstance(v, str) and "/" in v for v in enum_val):
node.pop("enum", None)
del node["enum"]
return 1
return 0
return _reactive_strip(
+112 -164
View File
@@ -2,6 +2,7 @@
from __future__ import annotations
import contextlib
import os
import re
import shlex
@@ -11,9 +12,7 @@ from pathlib import Path
from typing import Callable
from tools.approval import (
_bash_exec_payload,
_deobfuscate_shell_word_for_detection,
_iter_shell_command_starts,
_bash_exec_payload, _deobfuscate_shell_word_for_detection, _iter_shell_command_starts,
_read_shell_word)
# bisect drives repeated checkouts of the running root — the exact skew hazard guarded here.
@@ -24,8 +23,7 @@ _WORKTREE_TARGET_ACTIONS = frozenset({"move", "remove"})
_STASH_SAFE_ACTIONS = frozenset({"list", "show", "create", "store", "drop", "clear"})
_RESET_WORKTREE_MODES = frozenset({"--hard", "--merge", "--keep"})
# `reset`/`stash`/`clean`/`restore` reach this set only in their SAFE forms (_mutates_worktree
# classifies the dangerous forms first); listing them just skips a pointless
# `git config --get alias.<sub>` subprocess for `stash list`, `reset --soft`, `clean -n`.
# runs first); listing them skips a pointless `git config --get alias.<sub>` subprocess.
_KNOWN_GIT_BUILTINS = frozenset({
"add", "am", "apply", "blame", "branch", "bundle", "cat-file", "clean", "clone", "commit",
"config", "describe", "diff", "fetch", "format-patch", "grep", "help", "init", "log",
@@ -64,36 +62,31 @@ class _Heredoc:
@dataclass
class _ShellContext:
class _ShellContext: # one `(` / `$(` / backtick nesting level and its live quote state
kind: str
opener: int
quote: str | None = None
def get_running_source_root() -> Path | None:
"""Return the source checkout backing this process, if there is one."""
try:
"""The source checkout backing this process, if there is one."""
with contextlib.suppress(OSError, RuntimeError):
root = Path(__file__).resolve().parent.parent
except (OSError, RuntimeError):
return None
return root if (root / ".git").exists() else None
return root if (root / ".git").exists() else None
return None
def _resolve(path_str: str, base: Path) -> Path:
path = Path(os.path.expanduser(path_str))
if not path.is_absolute():
path = base / path
try:
path = base / Path(os.path.expanduser(path_str)) # ``/`` keeps an absolute right operand
with contextlib.suppress(OSError, RuntimeError, ValueError):
return path.resolve()
except (OSError, RuntimeError, ValueError):
return path
return path
def _is_within(path: Path, root: Path) -> bool:
try:
with contextlib.suppress(OSError, RuntimeError, ValueError):
return path == root or path.is_relative_to(root)
except (OSError, RuntimeError, ValueError):
return False
return False
def _executable_name(value: str) -> str:
@@ -101,6 +94,7 @@ def _executable_name(value: str) -> str:
def _shell_words_at(command: str, start: int) -> list[str]:
"""Deobfuscated words of the simple command at ``start`` (stops at a newline; max 64)."""
words: list[str] = []
cursor = start
for _ in range(64):
@@ -116,13 +110,10 @@ def _consume_options(
words: list[str], start: int, options_with_arg: frozenset[str] = _NO_OPTIONS) -> int:
"""Index of the first positional at/after ``start`` (``--`` ends options)."""
index = start
while index < len(words):
option = words[index]
if option == "--":
while index < len(words) and words[index].startswith("-") and words[index] != "-":
if words[index] == "--":
return index + 1
if not option.startswith("-") or option == "-":
break
index += 2 if "=" not in option and option in options_with_arg else 1
index += 2 if "=" not in words[index] and words[index] in options_with_arg else 1
return index
@@ -140,9 +131,8 @@ def _command_parts(words: list[str]) -> tuple[dict[str, str], str | None, list[s
wrapper_options = _WRAPPER_OPTIONS_WITH_ARG.get(executable)
if wrapper_options is None:
return env, words[index], words[index + 1 :]
# `command -v/-V` only reports; nothing runs.
if executable == "command" and words[index + 1 : index + 2] in (["-v"], ["-V"]):
return env, None, []
break # `command -v/-V` only reports; nothing runs
index = _consume_options(words, index + 1, wrapper_options)
return env, None, []
@@ -157,15 +147,14 @@ def _scope_keys(command: str, starts: list[int]) -> dict[int, tuple[int, ...]]:
context = contexts[-1]
quote = context.quote
char = command[cursor]
nested = len(contexts) > 1
if quote == "'":
if char == "'":
context.quote = None
closes = quote is None and len(contexts) > 1 # an unquoted closer may pop a scope
if quote is not None and char == quote:
context.quote = None
elif quote == "'":
pass # single quotes: no escapes, no substitutions
elif char == "\\" and cursor + 1 < start:
cursor += 1
elif quote == '"' and char == '"':
context.quote = None
elif quote is None and char in {"'", '"'}:
elif quote is None and char in "'\"":
context.quote = char
# Unquoted or inside double quotes: substitutions still open scopes.
elif command.startswith("$(", cursor):
@@ -173,28 +162,27 @@ def _scope_keys(command: str, starts: list[int]) -> dict[int, tuple[int, ...]]:
cursor += 1
elif quote is None and char == "(":
contexts.append(_ShellContext("(", cursor))
elif quote is None and char == ")" and nested and contexts[-1].kind in {"(", "$("}:
elif (char == ")" and closes and context.kind in {"(", "$("}) or (
char == "`" and closes and context.kind == "`"):
contexts.pop()
elif char == "`":
if quote is None and nested and contexts[-1].kind == "`":
contexts.pop()
else:
contexts.append(_ShellContext("`", cursor))
contexts.append(_ShellContext("`", cursor))
cursor += 1
scopes[start] = tuple(item.opener for item in contexts[1:])
return scopes
def _operator_before(command: str, start: int) -> str | None:
"""The list/grouping operator (or newline) immediately preceding a command start."""
head = command[:start].rstrip()
if head[-2:] in {"&&", "||"}:
return head[-2:]
if head[-1:] in {";", "|", "&", "(", "{"}:
return head[-1:]
for tail in (head[-2:], head[-1:]):
if tail in {"&&", "||", ";", "|", "&", "(", "{"}:
return tail
return "\n" if "\n" in command[len(head):start] else None
def _cd_target(executable: str, args: list[str], cwd: Path) -> Path | None:
"""Directory a ``cd``/``pushd`` would land in (existing dirs only), else None."""
if _executable_name(executable) not in {"cd", "pushd"}:
return None
index = _consume_options(args, 0)
@@ -205,17 +193,15 @@ def _cd_target(executable: str, args: list[str], cwd: Path) -> Path | None:
def _shell_script_arg(args: list[str]) -> str | None:
"""Return the script string owned by a shell's ``-c``, if present. approval.py's
``_bash_exec_payload`` parses bash's real option grammar (``-o pipefail -c '<script>'``
hides ``-c`` behind an operand); when it finds no ``-c``, fall back to a permissive
positional scan, since zsh/dash/ksh option letters (``zsh -yc``) fall outside bash's
alphabet and would otherwise make this block-guard fail open."""
"""Script string owned by a shell's ``-c``, if present. ``_bash_exec_payload`` parses bash's
real option grammar (``-o pipefail -c '<script>'``); when it finds no ``-c``, fall back to a
permissive scan: zsh/dash/ksh letters (``zsh -yc``) would otherwise fail this guard open."""
has_c, payload = _bash_exec_payload(args)
if has_c:
return payload
for index, arg in enumerate(args):
if arg == "--" or not arg.startswith("-"):
break
return None
if "c" in arg[1:]:
return args[index + 1] if index + 1 < len(args) else None
return None
@@ -233,7 +219,7 @@ def _heredoc_specs(line: str) -> list[_Heredoc]:
index += 1 # skip the escaped character too
elif char == quote:
quote = None
elif char in {"'", '"'}:
elif char in "'\"":
quote = char
if quote or not line.startswith("<<", index) or line.startswith("<<<", index):
index += 1
@@ -241,52 +227,47 @@ def _heredoc_specs(line: str) -> list[_Heredoc]:
opener = _HEREDOC_OPENER_RE.match(line, index)
if opener is None: # unterminated quoted delimiter: give up on this line
break
operator_at, index = index, opener.end()
strip_tabs = bool(opener.group("dash"))
header, index = line[:index], opener.end()
delimiter = opener.group("quoted") if opener.group("q") else opener.group("bare")
if not delimiter:
continue
header = line[:operator_at]
starts = list(_iter_shell_command_starts(header))
words = _shell_words_at(header, starts[-1]) if starts else []
_, executable, args = _command_parts(words)
_, executable, args = _command_parts(_shell_words_at(header, starts[-1]) if starts else [])
# A bare shell (no -c script, no script operand) executes the body itself.
execute_as_shell = bool(
executable
and _executable_name(executable) in _SHELL_EXECUTABLES
executable and _executable_name(executable) in _SHELL_EXECUTABLES
and _shell_script_arg(args) is None
and not any(arg and not arg.startswith("-") for arg in args))
specs.append(_Heredoc(delimiter, strip_tabs, execute_as_shell))
specs.append(_Heredoc(delimiter, bool(opener.group("dash")), execute_as_shell))
return specs
def _mask_heredocs(command: str) -> tuple[str, list[str]]:
"""Blank heredoc bodies; return (masked command, bodies a bare shell would execute).
Unterminated heredocs run to end of input and are still reported."""
"""Blank heredoc bodies -> (masked command, bodies a bare shell would execute). Unterminated
heredocs run to end of input and are still reported."""
output: list[str] = []
pending: list[_Heredoc] = []
finished: list[_Heredoc] = []
for line in command.splitlines(keepends=True):
if pending:
current = pending[0]
candidate = line.rstrip("\r\n")
if current.strip_tabs:
candidate = candidate.lstrip("\t")
if candidate == current.delimiter:
finished.append(pending.pop(0))
else:
current.body.append(line)
output.append("".join(char if char in {"\r", "\n"} else " " for char in line))
if not pending:
output.append(line)
pending.extend(_heredoc_specs(line))
continue
output.append(line)
pending.extend(_heredoc_specs(line))
current = pending[0]
candidate = line.rstrip("\r\n")
if (candidate.lstrip("\t") if current.strip_tabs else candidate) == current.delimiter:
finished.append(pending.pop(0))
else:
current.body.append(line)
output.append(re.sub(r"[^\r\n]", " ", line))
shell_scripts = ["".join(spec.body) for spec in finished + pending if spec.execute_as_shell]
return "".join(output), shell_scripts
def _record_alias(config: str, aliases: dict[str, str]) -> None:
"""Record an inline ``-c alias.<name>=<value>`` git config override."""
if config.lower().startswith("alias.") and "=" in config:
key, value = config.split("=", 1)
key, sep, value = config.partition("=")
if sep and key.lower().startswith("alias."):
aliases[key[6:].lower()] = value
@@ -305,88 +286,70 @@ def _git_target_and_subcommand(
break
if not arg.startswith("-"):
break
# Separate-argument form (`-C dir`) vs attached form (`-Cdir`, `--work-tree=dir`, `-cK=V`).
if arg in _GIT_GLOBAL_OPTIONS_WITH_ARG:
if index + 1 < len(args):
value = args[index + 1]
if arg == "-C":
target = _resolve(value, target)
elif arg == "--work-tree":
work_tree = value
elif arg == "-c":
_record_alias(value, aliases)
option, value = arg, args[index + 1] if index + 1 < len(args) else None
index += 2
continue
if arg.startswith("-C") and len(arg) > 2:
target = _resolve(arg[2:], target)
elif arg.startswith("--work-tree="):
work_tree = arg.split("=", 1)[1]
elif arg.startswith("-calias."):
_record_alias(arg[2:], aliases)
index += 1
else:
option, value = next(((o, arg[len(o):]) for o in ("-C", "--work-tree=", "-c")
if arg.startswith(o) and len(arg) > len(o)), (None, None))
index += 1
if option == "-C" and value is not None:
target = _resolve(value, target)
elif option and option.startswith("--work-tree") and value is not None:
work_tree = value
elif option == "-c" and value is not None:
_record_alias(value, aliases) # only alias.* overrides matter; others are ignored
explicit_work_tree = work_tree or env.get("GIT_WORK_TREE")
if explicit_work_tree:
target = _resolve(explicit_work_tree, target)
return target, args[index].lower() if index < len(args) else None, args[index + 1 :], aliases
def _has_short_flag(arg: str, letter: str) -> bool:
return arg.startswith("-") and letter in arg[1:]
def _reset_mutates(args: list[str]) -> bool:
return any(arg in _RESET_WORKTREE_MODES or _RESET_HARD_RE.fullmatch(arg) for arg in args)
def _stash_mutates(args: list[str]) -> bool:
action = next((arg for arg in args if not arg.startswith("-")), "push")
return action not in _STASH_SAFE_ACTIONS
def _clean_mutates(args: list[str]) -> bool:
return not any(
arg == "--dry-run" or (not arg.startswith("--") and _has_short_flag(arg, "n"))
def _has_flag(args: list[str], long: str, letter: str, short_only: bool = False) -> bool:
"""``--long`` present, or ``letter`` inside a dash-prefixed arg (``-fdn``; ``short_only``
additionally excludes ``--`` args from the letter scan)."""
return any(
arg == long or (arg.startswith("-") and letter in arg[1:]
and not (short_only and arg.startswith("--")))
for arg in args)
def _restore_mutates(args: list[str]) -> bool:
staged = any(arg == "--staged" or _has_short_flag(arg, "S") for arg in args)
worktree = any(arg == "--worktree" or _has_short_flag(arg, "W") for arg in args)
return worktree or not staged
# Subcommands whose worktree impact depends on their arguments.
# Subcommands whose worktree impact depends on their arguments -> predicate(args).
_CONDITIONAL_MUTATIONS: dict[str, Callable[[list[str]], bool]] = {
"reset": _reset_mutates,
"stash": _stash_mutates,
"clean": _clean_mutates,
"restore": _restore_mutates}
"reset": lambda args: any(
arg in _RESET_WORKTREE_MODES or _RESET_HARD_RE.fullmatch(arg) for arg in args),
"stash": lambda args: next(
(arg for arg in args if not arg.startswith("-")), "push") not in _STASH_SAFE_ACTIONS,
"clean": lambda args: not _has_flag(args, "--dry-run", "n", short_only=True),
# `restore` touches the worktree unless ONLY --staged was requested.
"restore": lambda args: (_has_flag(args, "--worktree", "W")
or not _has_flag(args, "--staged", "S"))}
def _mutates_worktree(subcommand: str, args: list[str]) -> bool:
check = _CONDITIONAL_MUTATIONS.get(subcommand)
return check(args) if check is not None else subcommand in _WORKTREE_MUTATIONS
check = _CONDITIONAL_MUTATIONS.get(subcommand, lambda _args: subcommand in _WORKTREE_MUTATIONS)
return check(args)
def _inspect_git_worktree(args: list[str], cwd: Path, root: Path) -> str | None:
"""Block `worktree remove|move` aimed at the running root, from any directory."""
action_index = _consume_options(args, 0)
action = args[action_index].lower() if action_index < len(args) else None
if action not in _WORKTREE_TARGET_ACTIONS:
return None
target_index = _consume_options(args, action_index + 1)
if target_index < len(args) and _resolve(args[target_index], cwd) == root:
if (action in _WORKTREE_TARGET_ACTIONS and target_index < len(args)
and _resolve(args[target_index], cwd) == root):
return f"git worktree {action}"
return None
def _read_git_alias(executable: str, target: Path, alias: str) -> str | None:
try:
with contextlib.suppress(OSError, subprocess.SubprocessError):
result = subprocess.run(
[executable, "-C", str(target), "config", "--get", f"alias.{alias}"],
capture_output=True, text=True, timeout=1, check=False)
except (OSError, subprocess.SubprocessError):
return None
return (result.stdout.strip() or None) if result.returncode == 0 else None
return (result.stdout.strip() or None) if result.returncode == 0 else None
return None
def _inspect_git(
@@ -396,8 +359,7 @@ def _inspect_git(
args, current_dir, env)
if subcommand is None:
return None
# `worktree` names its victim as an argument, so the cwd check does not apply.
if subcommand == "worktree":
if subcommand == "worktree": # names its victim as an argument: the cwd check does not apply
return _inspect_git_worktree(sub_args, target, root)
if not _is_within(target, root):
return None
@@ -405,18 +367,16 @@ def _inspect_git(
return f"git {subcommand}"
if subcommand in _KNOWN_GIT_BUILTINS or depth >= _MAX_RECURSION:
return None
alias = inline_aliases.get(subcommand)
if alias is None:
alias = _read_git_alias(executable, target, subcommand)
alias = (inline_aliases[subcommand] if subcommand in inline_aliases
else _read_git_alias(executable, target, subcommand))
if not alias:
return None
if alias.startswith("!"):
if alias.startswith("!"): # shell alias: scan it as a command
return _find_mutation(alias[1:], target, root, depth + 1)
try:
with contextlib.suppress(ValueError):
alias_args = shlex.split(alias, posix=True)
except ValueError:
return None
return _inspect_git(executable, [*alias_args, *sub_args], target, {}, root, depth + 1)
return _inspect_git(executable, [*alias_args, *sub_args], target, {}, root, depth + 1)
return None
def _inspect_github_cli(
@@ -425,9 +385,8 @@ def _inspect_github_cli(
if not _is_within(current_dir, root):
return None
index = _consume_options(args, 0, frozenset({"-R", "--repo", "--hostname"}))
if args[index : index + 2] == ["pr", "checkout"]:
return f"{_executable_name(executable)} pr checkout"
return None
is_checkout = args[index : index + 2] == ["pr", "checkout"]
return f"{_executable_name(executable)} pr checkout" if is_checkout else None
def _inspect_shell(
@@ -439,9 +398,7 @@ def _inspect_shell(
# executable name -> inspector(executable, args, current_dir, env, root, depth)
_INSPECTORS: dict[str, Callable[..., str | None]] = {
"git": _inspect_git,
"gh": _inspect_github_cli,
"hub": _inspect_github_cli,
"git": _inspect_git, "gh": _inspect_github_cli, "hub": _inspect_github_cli,
**{shell: _inspect_shell for shell in _SHELL_EXECUTABLES}}
@@ -451,8 +408,7 @@ def _find_mutation(command: str, cwd: Path, root: Path, depth: int = 0) -> str |
return None
masked_command, heredoc_scripts = _mask_heredocs(command)
for script in heredoc_scripts:
operation = _find_mutation(script, cwd, root, depth + 1)
if operation:
if operation := _find_mutation(script, cwd, root, depth + 1):
return operation
starts = sorted(set(_iter_shell_command_starts(masked_command)))
scopes = _scope_keys(masked_command, starts)
@@ -462,50 +418,42 @@ def _find_mutation(command: str, cwd: Path, root: Path, depth: int = 0) -> str |
for start in starts:
scope = scopes[start]
cwd_by_scope.setdefault(scope, cwd_by_scope.get(scope[:-1], cwd))
operator = _operator_before(masked_command, start)
pending = pending_cd.pop(scope, None)
if pending is not None and operator in {"&&", ";", "\n"}:
if pending is not None and _operator_before(masked_command, start) in {"&&", ";", "\n"}:
cwd_by_scope[scope] = pending
env, executable, args = _command_parts(_shell_words_at(masked_command, start))
if executable is None:
continue
current_dir = cwd_by_scope[scope]
cd_target = _cd_target(executable, args, current_dir)
if cd_target is not None:
if (cd_target := _cd_target(executable, args, current_dir)) is not None:
pending_cd[scope] = cd_target
continue
inspect = _INSPECTORS.get(_executable_name(executable))
if inspect is not None:
operation = inspect(executable, args, current_dir, env, root, depth)
if operation:
return operation
elif (inspect := _INSPECTORS.get(_executable_name(executable))) and (
operation := inspect(executable, args, current_dir, env, root, depth)):
return operation
return None
def guard_active() -> bool:
"""Whether the self-repo git guard applies on this platform. Windows-only: NTFS locks
loaded .py/.pyd files, so overwriting the live checkout can corrupt the running
process. On POSIX open handles keep the old inode alive; the mixed-module hazard is
limited to later lazy imports — not worth blocking every git workflow for."""
"""Windows-only: NTFS locks loaded .py/.pyd files, so overwriting the live checkout can
corrupt the running process. On POSIX open handles keep the old inode alive; the mixed-module
hazard is limited to later lazy imports — not worth blocking every git workflow for."""
return os.name == "nt"
def detect_self_repo_git_mutation(
command: str, cwd: str | None, source_root: Path | None = None) -> tuple[bool, str | None]:
"""Return whether a command would rewrite the live source checkout."""
"""-> (blocked, message): whether a command would rewrite the live source checkout."""
root = source_root if source_root is not None else get_running_source_root()
if root is None or not command:
return False, None
root = _resolve(str(root), Path("/"))
operation = _find_mutation(command, _resolve(cwd, Path("/")) if cwd else Path("/"), root)
operation = _find_mutation(command, _resolve(cwd or "/", Path("/")), root)
return (True, _block_message(operation, root)) if operation is not None else (False, None)
def _block_message(operation: str, root: Path) -> str:
# Suggest a disk-backed scratch dir: /tmp is usually tmpfs (see message).
hermes_home = os.environ.get("HERMES_HOME", "").strip()
home = Path(hermes_home).expanduser() if hermes_home else Path.home() / ".hermes"
scratch = home / "scratch"
scratch = (Path(hermes_home).expanduser() if hermes_home else Path.home() / ".hermes") / "scratch"
return (
f"Blocked: `{operation}` would rewrite Hermes's live source checkout "
f"({root}) and can mix module versions in this running process. "
+147 -234
View File
@@ -1,6 +1,7 @@
"""Standalone per-platform senders and error helpers for send_message."""
import asyncio
import contextlib
import logging
import os
import re
@@ -14,33 +15,24 @@ _IMAGE_EXTS = {".jpg", ".jpeg", ".png", ".webp", ".gif"}
_VIDEO_EXTS = {".mp4", ".mov", ".avi", ".mkv", ".webm", ".3gp"}
_AUDIO_EXTS = {".ogg", ".opus", ".mp3", ".m2a", ".wav", ".m4a", ".flac"}
_VOICE_EXTS = {".ogg", ".opus"}
# Telegram's sendAudio only accepts MP3 / M4A; other audio goes via sendVoice (Opus/OGG) or document.
_TELEGRAM_SEND_AUDIO_EXTS = {".mp3", ".m4a"}
# Extensions carrying a native caption on the media bubble. Voice/audio notes are excluded:
# a caption on a voice note reads as a separate label, so the text stays its own message.
_TELEGRAM_SEND_AUDIO_EXTS = {".mp3", ".m4a"} # sendAudio accepts only these; other audio -> sendVoice / document
# Captionable on the media bubble; voice/audio notes excluded (a caption there reads as a separate label).
_CAPTIONABLE_EXTS = _IMAGE_EXTS | _VIDEO_EXTS | {".pdf", ".doc", ".docx", ".txt", ".md", ".csv", ".xlsx", ".zip"}
# Native caption limits (chars): Telegram caps photo/video at 1024; one conservative shared ceiling elsewhere.
_TELEGRAM_CAPTION_LIMIT = 1024
_DEFAULT_CAPTION_LIMIT = 4096
def _media_caption_split(text, media_files, *, max_caption_len):
"""Single chokepoint deciding whether text rides on the media bubble as its caption.
``(caption, "")`` only for exactly one captionable file (not a voice/audio note)
whose text fits ``max_caption_len``; otherwise ``(None, text)`` — multi-file
caption→file association is ambiguous. Length is codepoints, which never
under-counts Telegram's UTF-16 units for BMP text (over-counting fails safe);
the Telegram sender re-checks the *formatted* caption since escaping inflates it.
"""
"""Single chokepoint deciding whether text rides on the media bubble as its caption:
``(caption, "")`` only for exactly one captionable file (not a voice/audio note) whose
text fits ``max_caption_len``, else ``(None, text)`` — multi-file caption→file association
is ambiguous. Length is codepoints (never under-counts Telegram's UTF-16 units for BMP
text); the Telegram sender re-checks the *formatted* caption since escaping inflates it."""
stripped = (text or "").strip()
media = media_files or []
if not stripped or len(media) != 1 or len(stripped) > max_caption_len:
return None, text
media_path, is_voice = media[0]
if is_voice or os.path.splitext(media_path)[1].lower() not in _CAPTIONABLE_EXTS:
if (not stripped or len(media) != 1 or len(stripped) > max_caption_len or media[0][1]
or os.path.splitext(media[0][0])[1].lower() not in _CAPTIONABLE_EXTS):
return None, text
return stripped, ""
@@ -68,20 +60,15 @@ def _success(platform: str, chat_id, warnings=None, **fields) -> dict:
**({"warnings": warnings} if warnings else {})}
def _display_chat_id(platform_name: str, chat_id: str) -> str:
"""Return a result-safe chat identifier for tool transcripts/log consumers."""
return "group:***" if platform_name == "signal" and str(chat_id).startswith("group:") else chat_id
_NO_DELIVERABLE = "No deliverable text or media remained after processing MEDIA tags"
_TELEGRAM_TRANSIENT_MARKERS = ("bad gateway", "502", "too many requests", "429",
"service unavailable", "503", "gateway timeout", "504")
_TELEGRAM_TRANSIENT_MARKERS = ("bad gateway", "502", "too many requests", "429", "service unavailable", "503",
"gateway timeout", "504")
def _telegram_retry_delay(exc: Exception, attempt: int) -> float | None:
"""Seconds to wait before retrying, or None when final. Honours ``retry_after``;
timeouts are never retried (the send may have gone through); 5xx/429 back off."""
"""Retry delay in seconds, or None when final: honours ``retry_after``; timeouts are
never retried (the send may have gone through); 5xx/429 back off exponentially."""
retry_after = getattr(exc, "retry_after", None)
if retry_after is not None:
try:
@@ -114,22 +101,19 @@ def _is_telegram_thread_not_found(error: Exception) -> bool:
def _telegram_bot(token):
"""Bot honouring TELEGRAM_PROXY (``telegram.proxy_url``) — without it the standalone
path times out where api.telegram.org is blocked. Falls back to a direct connection."""
"""Bot honouring TELEGRAM_PROXY (standalone sends time out where api.telegram.org is
blocked); falls back to a direct connection."""
from telegram import Bot
try:
from gateway.platforms.base import resolve_proxy_url
proxy = resolve_proxy_url("TELEGRAM_PROXY", target_hosts=["api.telegram.org"])
except Exception:
proxy = None
if proxy:
try:
from telegram.request import HTTPXRequest
logger.info("send_message: standalone Telegram send routed through proxy %s", proxy)
return Bot(token=token, request=HTTPXRequest(proxy=proxy),
get_updates_request=HTTPXRequest(proxy=proxy))
except Exception as proxy_err:
logger.warning("send_message: failed to attach Telegram proxy (%s), falling back to direct connection", proxy_err)
if not proxy:
return Bot(token=token)
from telegram.request import HTTPXRequest
logger.info("send_message: standalone Telegram send routed through proxy %s", proxy)
return Bot(token=token, request=HTTPXRequest(proxy=proxy), get_updates_request=HTTPXRequest(proxy=proxy))
except Exception as proxy_err:
logger.warning("send_message: failed to attach Telegram proxy (%s), falling back to direct connection", proxy_err)
return Bot(token=token)
@@ -141,8 +125,7 @@ def _telegram_thread_kwargs(thread_id):
try:
from plugins.platforms.telegram.adapter import TelegramAdapter
effective = TelegramAdapter._message_thread_id_for_send(str(thread_id))
except Exception:
# Explicit mapping if the adapter import fails (python-telegram-bot missing).
except Exception: # adapter import failed (python-telegram-bot missing): explicit mapping
effective = None if str(thread_id) == "1" else int(thread_id)
return {} if effective is None else {"message_thread_id": effective}
@@ -157,8 +140,8 @@ def _strip_mdv2_safe(text):
def _adapter_media_method(ext, voice, force_document=False):
"""Adapter media method name + kind for one file: document when forced, else image /
video / voice by extension (``voice`` already folds in the caller's audio rule)."""
"""``(adapter method, kind)``: document when forced, else image / video / voice by
extension (``voice`` already folds in the caller's audio rule)."""
if force_document:
return "send_document", "document"
if ext in _IMAGE_EXTS:
@@ -169,33 +152,26 @@ def _adapter_media_method(ext, voice, force_document=False):
async def _telegram_send_media(bot, chat_id, f, ext, is_voice, force_document, **kwargs):
"""Bot API media method by extension: photo (unless forced document), video, voice
note, sendAudio (MP3/M4A only), else document."""
if ext in _IMAGE_EXTS and not force_document:
return await bot.send_photo(chat_id=chat_id, photo=f, **kwargs)
if ext in _VIDEO_EXTS:
return await bot.send_video(chat_id=chat_id, video=f, **kwargs)
if ext in _VOICE_EXTS and is_voice:
return await bot.send_voice(chat_id=chat_id, voice=f, **kwargs)
if ext in _TELEGRAM_SEND_AUDIO_EXTS:
return await bot.send_audio(chat_id=chat_id, audio=f, **kwargs)
return await bot.send_document(chat_id=chat_id, document=f, **kwargs)
"""Bot API media method by extension: photo (unless forced document), video, voice note,
sendAudio (MP3/M4A only), else document."""
kind = next((k for exts, k in ((() if force_document else _IMAGE_EXTS, "photo"), (_VIDEO_EXTS, "video"),
(_VOICE_EXTS if is_voice else (), "voice"), (_TELEGRAM_SEND_AUDIO_EXTS, "audio"))
if ext in exts), "document")
return await getattr(bot, f"send_{kind}")(chat_id=chat_id, **{kind: f}, **kwargs)
async def _telegram_send_text_chunk(bot, chat_id, chunk, parse_mode, has_html, text_kwargs):
"""One formatted text chunk with adapter-matching fallbacks: thread-not-found -> retry
without ``message_thread_id`` (dropped from ``text_kwargs`` for later chunks too);
parse failure -> plain text."""
"""One text chunk with adapter-matching fallbacks: thread-not-found -> retry without
``message_thread_id`` (dropped from ``text_kwargs`` for later chunks too); parse failure
-> plain text."""
async def send(text, mode):
return await _send_telegram_message_with_retry(bot, chat_id=chat_id, text=text, parse_mode=mode, **text_kwargs)
try:
return await send(chunk, parse_mode)
except Exception as md_error:
if _is_telegram_thread_not_found(md_error) and text_kwargs.get("message_thread_id") is not None:
logger.warning("Thread %s not found in _send_telegram, retrying without message_thread_id",
text_kwargs.get("message_thread_id"))
text_kwargs.pop("message_thread_id", None)
text_kwargs.pop("message_thread_id"))
return await send(chunk, parse_mode)
err_text = str(md_error).lower()
if "parse" in err_text or "markdown" in err_text or "html" in err_text:
@@ -205,28 +181,22 @@ async def _telegram_send_text_chunk(bot, chat_id, chunk, parse_mode, has_html, t
raise
async def _telegram_send_one_media(
bot, chat_id, media_path, is_voice, *, caption, parse_mode, has_html, thread_kwargs, force_document
):
async def _telegram_send_one_media(bot, chat_id, media_path, is_voice, *, caption, parse_mode, has_html,
thread_kwargs, force_document):
"""Upload one file with adapter-matching fallbacks (thread-not-found -> no
``message_thread_id``; caption parse failure -> plain caption). Retries re-seek
the file because the first attempt consumed it."""
``message_thread_id``; caption parse failure -> plain caption); retries re-seek the file."""
ext = os.path.splitext(media_path)[1].lower()
voice_note = ext in _VOICE_EXTS and is_voice
media_kwargs = dict(thread_kwargs)
# ``caption`` is only set for a single captionable file, so this never
# double-captions a multi-file send or a voice note.
if caption is not None and not voice_note:
media_kwargs.update(caption=caption, parse_mode=parse_mode)
# ``caption`` is only set for a single captionable file, so this never double-captions
# a multi-file send or a voice note.
media_kwargs = {**thread_kwargs, **({"caption": caption, "parse_mode": parse_mode}
if caption is not None and not voice_note else {})}
if voice_note or ext in _TELEGRAM_SEND_AUDIO_EXTS:
try:
with contextlib.suppress(Exception):
from plugins.platforms.telegram.adapter import _probe_voice_duration_seconds
duration = await asyncio.to_thread(_probe_voice_duration_seconds, media_path)
if duration is not None:
media_kwargs["duration"] = duration
except Exception:
pass
with open(media_path, "rb") as f:
try:
return await _telegram_send_media(bot, chat_id, f, ext, is_voice, force_document, **media_kwargs)
@@ -255,54 +225,42 @@ def _telegram_format(message):
return message, ParseMode.HTML, True
try:
from plugins.platforms.telegram.adapter import TelegramAdapter
formatted = TelegramAdapter.__new__(TelegramAdapter).format_message(message)
return TelegramAdapter.__new__(TelegramAdapter).format_message(message), ParseMode.MARKDOWN_V2, False
except Exception:
formatted = message # formatting unavailable: send as-is
return formatted, ParseMode.MARKDOWN_V2, False
return message, ParseMode.MARKDOWN_V2, False # formatting unavailable: send as-is
async def _send_telegram(token, chat_id, message, media_files=None, thread_id=None, disable_link_previews=False, force_document=False):
"""One-shot Telegram Bot API send; parse failures fall back to plain text so the
message still delivers."""
"""One-shot Telegram Bot API send; parse failures fall back to plain text."""
try:
formatted, send_parse_mode, _has_html = _telegram_format(message)
bot = _telegram_bot(token)
from plugins.platforms.telegram.telegram_ids import normalize_telegram_chat_id
from gateway.platforms.base import BasePlatformAdapter, utf16_len
# Telegram accepts a numeric chat_id OR an @username string; never force-int.
int_chat_id = normalize_telegram_chat_id(chat_id)
media_files = media_files or []
thread_kwargs = _telegram_thread_kwargs(thread_id)
# disable_web_page_preview is only valid for send_message, not media sends.
text_kwargs = {**thread_kwargs, **({"disable_web_page_preview": True} if disable_link_previews else {})}
last_msg, warnings = None, []
# MEDIA caption: a single captionable file + short text rides on the bubble as its
# *formatted* caption. Formatting can inflate a raw <1024 string past Telegram's
# cap, so re-check in UTF-16 units and fall back to a separate body.
_tg_caption = None
from gateway.platforms.base import BasePlatformAdapter, utf16_len
last_msg, warnings, _tg_caption = None, [], None
# MEDIA caption rides on the bubble as its *formatted* caption; formatting can inflate a
# raw <1024 string past Telegram's cap, so re-check in UTF-16 units.
_cap, _ = _media_caption_split(message, media_files, max_caption_len=_TELEGRAM_CAPTION_LIMIT)
if _cap is not None and utf16_len(formatted) <= _TELEGRAM_CAPTION_LIMIT:
_tg_caption, formatted = formatted, "" # suppress the separate text send below
if formatted.strip():
# Chunk *after* formatting, in UTF-16 units: MarkdownV2/HTML escaping inflates
# text, so a raw-<4096 message can exceed the limit once formatted.
for chunk in BasePlatformAdapter.truncate_message(formatted, 4096, len_fn=utf16_len):
last_msg = await _telegram_send_text_chunk(
bot, int_chat_id, chunk, send_parse_mode, _has_html, text_kwargs)
# Chunk *after* formatting, in UTF-16 units: escaping can push a raw-<4096 message over.
for chunk in BasePlatformAdapter.truncate_message(formatted, 4096, len_fn=utf16_len) if formatted.strip() else ():
last_msg = await _telegram_send_text_chunk(bot, int_chat_id, chunk, send_parse_mode, _has_html, text_kwargs)
for media_path, is_voice in media_files:
if not os.path.exists(media_path):
warnings.append(f"Media file not found, skipping: {media_path}")
logger.warning(warnings[-1])
# Caption mode suppressed the text send; if the file it was meant to
# caption is gone, deliver the words on their own.
# Caption mode suppressed the text send; the file is gone, so deliver the words alone.
if _tg_caption is not None and last_msg is None:
try:
last_msg = await _send_telegram_message_with_retry(
bot, chat_id=int_chat_id, text=_tg_caption,
parse_mode=send_parse_mode, **text_kwargs)
bot, chat_id=int_chat_id, text=_tg_caption, parse_mode=send_parse_mode, **text_kwargs)
_tg_caption = None # delivered — don't re-caption a later file
except Exception as _cap_err:
logger.warning("Telegram caption-fallback send failed for missing media: %s",
@@ -310,13 +268,11 @@ async def _send_telegram(token, chat_id, message, media_files=None, thread_id=No
continue
try:
last_msg = await _telegram_send_one_media(
bot, int_chat_id, media_path, is_voice,
caption=_tg_caption, parse_mode=send_parse_mode, has_html=_has_html,
thread_kwargs=thread_kwargs, force_document=force_document)
bot, int_chat_id, media_path, is_voice, caption=_tg_caption, parse_mode=send_parse_mode,
has_html=_has_html, thread_kwargs=thread_kwargs, force_document=force_document)
except Exception as e:
warnings.append(_sanitize_error_text(f"Failed to send media {media_path}: {e}"))
logger.error(warnings[-1])
if last_msg is None:
return {"error": _NO_DELIVERABLE, **({"warnings": warnings} if warnings else {})}
return _success("telegram", chat_id, warnings, message_id=str(last_msg.message_id))
@@ -326,20 +282,15 @@ async def _send_telegram(token, chat_id, message, media_files=None, thread_id=No
return _error(f"Telegram send failed: {e}")
def _live_runner():
"""Return the in-process gateway runner, or None (standalone/cron)."""
def _live_adapter(platform, *, lookup_failed_warning=None):
"""``(runner, adapter)`` for the in-process gateway; ``(None, None)`` standalone (cron);
``(runner, None)`` when the lookup fails — logged when a warning is given, never silently
swallowed (a silent fall-through could recreate a reconnect storm)."""
try:
from gateway.run import _gateway_runner_ref
return _gateway_runner_ref()
runner = _gateway_runner_ref()
except Exception:
return None
def _live_adapter(platform, *, lookup_failed_warning=None):
"""Return ``(runner, adapter)`` for the running gateway, or ``(runner, None)``.
A runner whose adapter lookup raises is logged when a warning is given, never
silently swallowed (a silent fall-through could recreate a reconnect storm)."""
runner = _live_runner()
runner = None
if runner is None:
return None, None
try:
@@ -366,16 +317,13 @@ def _plugin_standalone_sender(platform_name, *, label=None, discover=True):
async def _registry_standalone_send(platform_name, pconfig, chat_id, message, thread_id=None):
"""One-shot text send through a plugin's ``standalone_sender_fn``."""
sender, err = _plugin_standalone_sender(platform_name)
if err:
return err
return await sender(pconfig, chat_id, message, thread_id=thread_id)
return err or await sender(pconfig, chat_id, message, thread_id=thread_id)
async def _resolve_slack_user_target(token, chat_id):
"""Resolve ``user:U...`` / ``user_name:<handle>`` to a D... DM conversation
(chat.postMessage needs a conversation ID). ``user_name:`` maps to a user id via
users.list first (stable handle match only); other ids pass through unchanged.
Returns ``(chat_id, None)`` or ``(None, error_dict)``."""
"""Resolve ``user:U...`` / ``user_name:<handle>`` to a D... DM conversation (chat.postMessage
needs a conversation ID); ``user_name:`` goes through users.list first (stable handle match
only); other ids pass through. ``(chat_id, None)`` or ``(None, error_dict)``."""
if not (chat_id.startswith("user:") or chat_id.startswith("user_name:")):
return chat_id, None
try:
@@ -386,65 +334,52 @@ async def _resolve_slack_user_target(token, chat_id):
from gateway.platforms.base import resolve_proxy_url, proxy_kwargs_for_aiohttp
_sess_kw, _req_kw = proxy_kwargs_for_aiohttp(resolve_proxy_url())
headers = {"Authorization": f"Bearer {token}", "Content-Type": "application/json"}
async def post_api(session, method, payload):
async with session.post(
f"https://slack.com/api/{method}", headers=headers, json=payload, **_req_kw) as resp:
return await resp.json()
async def resolve_user_name(session, name):
query = name.strip().lstrip("@").lower()
matches, cursor = [], None
for _page in range(20):
payload = {"limit": 200, **({"cursor": cursor} if cursor else {})}
data = await post_api(session, "users.list", payload)
if not data.get("ok"):
return None, f"Slack users.list error: {data.get('error', 'unknown')}"
# Stable handle only: display/real names are mutable and non-unique.
matches += [m for m in data.get("members", [])
if not (m.get("deleted") or m.get("is_bot"))
and str(m.get("name", "")).strip().lower() == query]
cursor = (data.get("response_metadata") or {}).get("next_cursor")
if not cursor:
break
if not matches:
return None, f"Could not resolve Slack user '@{name}'."
if len(matches) > 1:
return None, f"Slack user '@{name}' matched multiple Slack users. Use a Slack user ID instead."
return matches[0].get("id"), None
async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=30), **_sess_kw) as session:
if chat_id.startswith("user_name:"):
user_id, error = await resolve_user_name(session, chat_id[len("user_name:"):])
if error:
return None, _error(error)
chat_id = f"user:{user_id}"
async def post_api(method, payload):
async with session.post(f"https://slack.com/api/{method}", headers=headers, json=payload,
**_req_kw) as resp:
return await resp.json()
user_id = chat_id[len("user:"):]
opened = await post_api(session, "conversations.open", {"users": user_id})
if chat_id.startswith("user_name:"):
name = chat_id[len("user_name:"):]
query = name.strip().lstrip("@").lower()
matches, cursor = [], None
for _page in range(20):
data = await post_api("users.list", {"limit": 200, **({"cursor": cursor} if cursor else {})})
if not data.get("ok"):
return None, _error(f"Slack users.list error: {data.get('error', 'unknown')}")
# Stable handle only: display/real names are mutable and non-unique.
matches += [m for m in data.get("members", []) if not (m.get("deleted") or m.get("is_bot"))
and str(m.get("name", "")).strip().lower() == query]
cursor = (data.get("response_metadata") or {}).get("next_cursor")
if not cursor:
break
if not matches:
return None, _error(f"Could not resolve Slack user '@{name}'.")
if len(matches) > 1:
return None, _error(f"Slack user '@{name}' matched multiple Slack users. Use a Slack user ID instead.")
chat_id = f"user:{matches[0].get('id')}"
opened = await post_api("conversations.open", {"users": chat_id[len("user:"):]})
if not opened.get("ok"):
return None, _error(f"Slack conversations.open error: {opened.get('error', 'unknown')}. "
"Check bot permissions (im:write).")
dm_id = (opened.get("channel") or {}).get("id")
if not dm_id:
return None, _error("Slack conversations.open did not return a DM channel ID")
return dm_id, None
return (dm_id, None) if dm_id else (None, _error("Slack conversations.open did not return a DM channel ID"))
except Exception as e:
return None, _error(f"Slack DM resolution failed: {e}")
async def _signal_send_batch(post, scheduler, rl, idx, n_batches, att_batch, batch_message):
"""One Signal batch under the scheduler with rate-limit retries. None on success,
False when retries were exhausted (batch lost), error dict for a non-rate-limit RPC error."""
"""One Signal batch under the scheduler with rate-limit retries: None on success, False when
retries were exhausted (batch lost), error dict for a non-rate-limit RPC error."""
n, max_attempts = len(att_batch), rl.SIGNAL_RATE_LIMIT_MAX_ATTEMPTS
for attempt in range(1, max_attempts + 1):
try:
await scheduler.acquire(n)
_rpc_t0 = time.monotonic()
data = await post(att_batch, batch_message)
_rpc_duration = time.monotonic() - _rpc_t0
if "error" not in data:
await scheduler.report_rpc_duration(_rpc_duration, n)
await scheduler.report_rpc_duration(time.monotonic() - _rpc_t0, n)
return None
err = data["error"]
if not rl._is_signal_rate_limit_error(err):
@@ -453,12 +388,10 @@ async def _signal_send_batch(post, scheduler, rl, idx, n_batches, att_batch, bat
scheduler.feedback(server_retry_after, n)
retry_after_label = f"{server_retry_after:.0f}s" if server_retry_after else "unknown"
if attempt >= max_attempts:
logger.error("Signal: rate-limit retries exhausted on batch %d/%d "
"(%d attachments lost, server retry_after=%s)",
idx + 1, n_batches, n, retry_after_label)
logger.error("Signal: rate-limit retries exhausted on batch %d/%d (%d attachments lost, "
"server retry_after=%s)", idx + 1, n_batches, n, retry_after_label)
return False
logger.warning("Signal: rate-limited on batch %d/%d "
"(attempt %d/%d, server retry_after=%s); "
logger.warning("Signal: rate-limited on batch %d/%d (attempt %d/%d, server retry_after=%s); "
"scheduler will pace the retry",
idx + 1, n_batches, attempt, max_attempts, retry_after_label)
except Exception as e:
@@ -471,34 +404,29 @@ async def _signal_send_batch(post, scheduler, rl, idx, n_batches, att_batch, bat
async def _send_signal(extra, chat_id, message, media_files=None):
"""signal-cli JSON-RPC send. Attachments go in SIGNAL_MAX_ATTACHMENTS_PER_MSG batches
metered by the process-wide SignalAttachmentScheduler — the same bucket the gateway
adapter uses, so tool sends and inbound replies share rate-limit state."""
"""signal-cli JSON-RPC send; attachments go in SIGNAL_MAX_ATTACHMENTS_PER_MSG batches metered
by the process-wide SignalAttachmentScheduler (shared with the gateway adapter's rate-limit state)."""
try:
import httpx
except ImportError:
return {"error": "httpx not installed"}
from gateway.platforms import signal_rate_limit as rl
from gateway.platforms.signal_format import markdown_to_signal
try:
http_url = extra.get("http_url", "http://127.0.0.1:8080").rstrip("/")
account = extra.get("account", "")
http_url, account = extra.get("http_url", "http://127.0.0.1:8080").rstrip("/"), extra.get("account", "")
if not account:
return {"error": "Signal account not configured"}
valid_media = media_files or []
attachment_paths = [path for path, _is_voice in valid_media if os.path.exists(path)]
attachment_paths = []
for media_path, _is_voice in valid_media:
if not os.path.exists(media_path):
if os.path.exists(media_path):
attachment_paths.append(media_path)
else:
logger.warning("Signal media file not found, skipping: %s", media_path)
# No attachments still means one (text-only) batch; with attachments
# the text rides on batch #0 so it isn't repeated per batch.
# No attachments still means one (text-only) batch; text rides on batch #0 only.
per_batch = rl.SIGNAL_MAX_ATTACHMENTS_PER_MSG
att_batches = [attachment_paths[i:i + per_batch]
for i in range(0, len(attachment_paths), per_batch)] or [[]]
n_batches = len(att_batches)
plain_text, text_styles = markdown_to_signal(message)
att_batches = [attachment_paths[i:i + per_batch] for i in range(0, len(attachment_paths), per_batch)] or [[]]
n_batches, (plain_text, text_styles) = len(att_batches), markdown_to_signal(message)
recipient = {"groupId": chat_id[6:]} if chat_id.startswith("group:") else {"recipient": [chat_id]}
async def _rpc_send(text, *, id_prefix, timeout, attachments=None, styled=False):
@@ -518,20 +446,17 @@ async def _send_signal(extra, chat_id, message, media_files=None):
timeout=rl._signal_send_timeout(len(batch_attachments)))
resp.raise_for_status()
return resp.json()
scheduler = rl.get_scheduler()
logger.info("send_message Signal: scheduler state=%s, %d attachment(s) in %d batch(es)",
scheduler.state(), len(attachment_paths), n_batches)
failed_batches: list[int] = []
for idx, att_batch in enumerate(att_batches):
n = len(att_batch)
estimated = scheduler.estimate_wait(n) if n > 0 else 0.0
if n > 0 and estimated >= rl.SIGNAL_BATCH_PACING_NOTICE_THRESHOLD:
if n > 0 and (estimated := scheduler.estimate_wait(n)) >= rl.SIGNAL_BATCH_PACING_NOTICE_THRESHOLD:
# Best-effort one-shot RPC for a user-facing pacing notice.
notice = (f"(More images coming — pausing ~{rl._format_wait(estimated)} "
f"for Signal rate limit, batch {idx + 1}/{n_batches}.)")
try:
await _rpc_send(notice, id_prefix="notice", timeout=30.0)
await _rpc_send(f"(More images coming — pausing ~{rl._format_wait(estimated)} "
f"for Signal rate limit, batch {idx + 1}/{n_batches}.)", id_prefix="notice", timeout=30.0)
except Exception as _e:
logger.warning("Signal: inline notice failed: %s", _e)
outcome = await _signal_send_batch(_post, scheduler, rl, idx, n_batches, att_batch,
@@ -540,7 +465,6 @@ async def _send_signal(extra, chat_id, message, media_files=None):
failed_batches.append(idx + 1)
elif outcome is not None:
return outcome
warnings = []
if len(attachment_paths) < len(valid_media):
warnings.append("Some media files were skipped (not found on disk)")
@@ -549,16 +473,16 @@ async def _send_signal(extra, chat_id, message, media_files=None):
f"(#{', #'.join(str(b) for b in failed_batches)})")
if failed_batches and len(failed_batches) == n_batches:
return _error(f"Signal: every batch ({n_batches}) hit rate limit; no attachments delivered")
return _success("signal", _display_chat_id("signal", chat_id), warnings)
# Result-safe chat identifier for tool transcripts/log consumers.
return _success("signal", "group:***" if str(chat_id).startswith("group:") else chat_id, warnings)
except Exception as e:
return _error(f"Signal send failed: {e}")
async def _send_matrix_via_adapter(pconfig, chat_id, message, media_files=None, thread_id=None):
"""Matrix adapter send (native media preserved). Prefer the live gateway adapter's
persistent olm/megolm session: ephemeral per-send connects re-init E2EE and claim
one-time keys, which under bursts exhausts recipient OTKs and silently drops
messages — so the ephemeral connect/disconnect path is only for standalone/cron."""
"""Matrix adapter send (native media preserved). Prefer the live gateway adapter's persistent
olm/megolm session: ephemeral per-send connects re-init E2EE and claim one-time keys, which
under bursts exhausts recipient OTKs and silently drops messages — ephemeral is cron-only."""
media_files = media_files or []
metadata = {"thread_id": thread_id} if thread_id else None
from gateway.config import Platform
@@ -566,16 +490,12 @@ async def _send_matrix_via_adapter(pconfig, chat_id, message, media_files=None,
"Matrix: live gateway adapter lookup failed; falling back to an "
"ephemeral connect (may re-init E2EE per send)"))
if live_adapter is not None:
# Owned by the gateway — must NOT be disconnected; return before the
# ephemeral adapter (and its ``finally`` disconnect) exists.
# Owned by the gateway — must NOT be disconnected (return before the ephemeral ``finally``).
return await _matrix_send_core(live_adapter, chat_id, message, media_files, metadata)
# --- Fallback: ephemeral adapter (standalone / cron context) ---
try:
from plugins.platforms.matrix.adapter import MatrixAdapter
except ImportError:
return {"error": "Matrix dependencies not installed. Run: pip install 'mautrix[encryption]'"}
adapter = MatrixAdapter(pconfig)
try:
if not await adapter.connect():
@@ -584,10 +504,8 @@ async def _send_matrix_via_adapter(pconfig, chat_id, message, media_files=None,
except Exception as e:
return _error(f"Matrix send failed: {e}")
finally:
try:
with contextlib.suppress(Exception):
await adapter.disconnect()
except Exception:
pass
async def _matrix_send_core(adapter, chat_id, message, media_files, metadata):
@@ -597,58 +515,58 @@ async def _matrix_send_core(adapter, chat_id, message, media_files, metadata):
last_result = await adapter.send(chat_id, message, metadata=metadata)
if not last_result.success:
return _error(f"Matrix send failed: {last_result.error}")
for media_path, is_voice in media_files:
if not os.path.exists(media_path):
return _error(f"Media file not found: {media_path}")
ext = os.path.splitext(media_path)[1].lower()
method, _ = _adapter_media_method(ext, (ext in _VOICE_EXTS and is_voice) or ext in _AUDIO_EXTS)
last_result = await getattr(adapter, method)(chat_id, media_path, metadata=metadata)
if not last_result.success:
return _error(f"Matrix media send failed: {last_result.error}")
return {"error": _NO_DELIVERABLE} if last_result is None else _success("matrix", chat_id, message_id=last_result.message_id)
return {"error": _NO_DELIVERABLE} if last_result is None else _success(
"matrix", chat_id, message_id=last_result.message_id)
def _gateway_platform_module(name, *, unavailable, unmet):
"""``(gateway.platforms.<name>, None)`` once its ``check_<name>_requirements`` passes, else ``(None, error)``."""
import importlib
try:
module = importlib.import_module(f"gateway.platforms.{name}")
except ImportError:
return None, {"error": unavailable}
return (module, None) if getattr(module, f"check_{name}_requirements")() else (None, {"error": unmet})
async def _send_weixin(pconfig, chat_id, message, media_files=None):
"""Send via Weixin iLink using the native adapter helper."""
wx, err = _gateway_platform_module("weixin", unavailable="Weixin adapter not available.",
unmet="Weixin requirements not met. Need aiohttp + cryptography.")
if err:
return err
try:
from gateway.platforms.weixin import check_weixin_requirements, send_weixin_direct
if not check_weixin_requirements():
return {"error": "Weixin requirements not met. Need aiohttp + cryptography."}
except ImportError:
return {"error": "Weixin adapter not available."}
try:
return await send_weixin_direct(extra=pconfig.extra, token=pconfig.token, chat_id=chat_id,
message=message, media_files=media_files)
return await wx.send_weixin_direct(extra=pconfig.extra, token=pconfig.token, chat_id=chat_id,
message=message, media_files=media_files)
except Exception as e:
return _error(f"Weixin send failed: {e}")
async def _send_bluebubbles(extra, chat_id, message):
"""Send via BlueBubbles iMessage server using the adapter's REST API."""
try:
from gateway.platforms.bluebubbles import BlueBubblesAdapter, check_bluebubbles_requirements
if not check_bluebubbles_requirements():
return {"error": "BlueBubbles requirements not met (need aiohttp + httpx)."}
except ImportError:
return {"error": "BlueBubbles adapter not available."}
bb, err = _gateway_platform_module("bluebubbles", unavailable="BlueBubbles adapter not available.",
unmet="BlueBubbles requirements not met (need aiohttp + httpx).")
if err:
return err
try:
from gateway.config import PlatformConfig
adapter = BlueBubblesAdapter(PlatformConfig(extra=extra))
adapter = bb.BlueBubblesAdapter(PlatformConfig(extra=extra))
if not await adapter.connect():
return _error("BlueBubbles: failed to connect to server")
try:
result = await adapter.send(chat_id, message)
if not result.success:
return _error(f"BlueBubbles send failed: {result.error}")
return _success("bluebubbles", chat_id, message_id=result.message_id)
finally:
await adapter.disconnect()
if not result.success:
return _error(f"BlueBubbles send failed: {result.error}")
return _success("bluebubbles", chat_id, message_id=result.message_id)
except Exception as e:
return _error(f"BlueBubbles send failed: {e}")
@@ -667,7 +585,6 @@ async def _send_qqbot(pconfig, chat_id, message):
secret = pconfig.token or extra.get("client_secret") or _getenv("QQ_CLIENT_SECRET", "")
if not appid or not secret:
return _error("QQBot: QQ_APP_ID / QQ_CLIENT_SECRET not configured.")
try:
async with httpx.AsyncClient(timeout=15) as client:
token_resp = await client.post("https://bots.qq.com/app/getAppAccessToken",
@@ -681,10 +598,9 @@ async def _send_qqbot(pconfig, chat_id, message):
# Separate endpoints for guild channels, C2C (private) and groups; first 2xx wins.
headers = {"Authorization": f"QQBot {access_token}", "Content-Type": "application/json"}
payload = {"content": message[:4000], "msg_type": 0}
endpoints = (
("channel", f"https://api.sgroup.qq.com/channels/{chat_id}/messages"),
("c2c", f"https://api.sgroup.qq.com/v2/users/{chat_id}/messages"),
("group", f"https://api.sgroup.qq.com/v2/groups/{chat_id}/messages"))
endpoints = (("channel", f"https://api.sgroup.qq.com/channels/{chat_id}/messages"),
("c2c", f"https://api.sgroup.qq.com/v2/users/{chat_id}/messages"),
("group", f"https://api.sgroup.qq.com/v2/groups/{chat_id}/messages"))
statuses = []
for kind, url in endpoints:
resp = await client.post(url, json=payload, headers=headers)
@@ -697,17 +613,14 @@ async def _send_qqbot(pconfig, chat_id, message):
async def _send_yuanbao(chat_id, message, media_files=None):
"""Send via the running Yuanbao adapter's persistent WebSocket (no throwaway client
possible). chat_id: ``group:<code>``, ``direct:<id>`` or ``<id>``."""
"""Send via the running Yuanbao adapter's persistent WebSocket (no throwaway client possible)."""
try:
from gateway.platforms.yuanbao import get_active_adapter, send_yuanbao_direct
except ImportError:
return _error("Yuanbao adapter module not available.")
adapter = get_active_adapter()
if adapter is None:
return _error("Yuanbao adapter is not running. Start the gateway with yuanbao platform enabled first.")
try:
return await send_yuanbao_direct(adapter, chat_id, message, media_files=media_files)
except Exception as e:
+58 -94
View File
@@ -6,10 +6,11 @@ import re
logger = logging.getLogger("tools.send_message_tool")
_TELEGRAM_TOPIC_TARGET_RE = re.compile(r"^\s*(-?\d+)(?::(\d+))?\s*$")
_NUMERIC_TOPIC_RE = _TELEGRAM_TOPIC_TARGET_RE # Discord snowflakes: numeric, same "<id>[:<thread>]" shape
_FEISHU_TARGET_RE = re.compile(r"^\s*((?:oc|ou|on|chat|open)_[-A-Za-z0-9]+)(?::([-A-Za-z0-9_]+))?\s*$")
# Slack conversation IDs: C (public), G (private/group), D (DM); uppercase alnum, 9+ chars.
# User IDs (U...) become ``user:U...`` and are opened as D... conversations first (posting
# straight to a U/W id fails); ``@handle`` -> ``user_name:...`` resolves via users.list.
# Slack conversation IDs: C (public), G (private/group), D (DM); uppercase alnum, 9+ chars. User IDs
# (U...) become ``user:U...`` and are opened as D... conversations first (posting straight to a U/W
# id fails); ``@handle`` -> ``user_name:...`` resolves via users.list.
_SLACK_TARGET_RE = re.compile(r"^\s*([CGD][A-Z0-9]{8,})\s*$")
_SLACK_USER_ID_RE = re.compile(r"^\s*(U[A-Z0-9]{8,})\s*$")
_SLACK_USER_NAME_RE = re.compile(r"^\s*@([A-Za-z0-9._-]{1,80})\s*$")
@@ -18,24 +19,17 @@ _SLACK_MENTION_RE = re.compile(r"^\s*<@(U[A-Z0-9]{8,})(?:\|[^>]+)?>\s*$")
_SLACK_THREAD_TARGET_RE = re.compile(r"^\s*([CGD][A-Z0-9]{8,}):([^\s:]+)\s*$")
_WEIXIN_TARGET_RE = re.compile(r"^\s*((?:wxid|gh|v\d+|wm|wb)_[A-Za-z0-9_-]+|[A-Za-z0-9._-]+@chatroom|filehelper)\s*$")
_YUANBAO_TARGET_RE = re.compile(r"^\s*((?:group|direct):[^:]+)\s*$")
# Discord snowflake IDs are numeric, same regex pattern as Telegram topic targets.
_NUMERIC_TOPIC_RE = _TELEGRAM_TOPIC_TARGET_RE
# Platforms addressing recipients by E.164 phone number ("+1555..."): the '+' fails the
# isdigit() rule and channel-name resolution cannot resolve a raw number; keep the '+'.
# E.164 phone recipients ("+1555..."): the '+' fails the isdigit() rule and the channel directory
# cannot resolve a raw number, so keep the '+' and treat it as explicit.
_PHONE_PLATFORMS = frozenset({"photon", "signal", "sms", "whatsapp"})
_E164_TARGET_RE = re.compile(r"^\s*\+(\d{7,15})\s*$")
# Photon DM chat GUID (mirrors _DM_CHAT_GUID_RE in the photon adapter).
_PHOTON_DM_GUID_RE = re.compile(r"^any;-;\+\d{6,}$")
# WhatsApp JIDs (@g.us groups, @s.whatsapp.net users, @lid, broadcast/newsletter): native
# targets the bridge accepts verbatim — never home-channel.
_WHATSAPP_JID_RE = re.compile(
r"^\s*[\w-]+@(?:g\.us|s\.whatsapp\.net|lid|broadcast|newsletter)\s*$", re.IGNORECASE)
# Buzz channels/DMs are native UUIDs: explicit targets, never the home channel.
_BUZZ_UUID_RE = re.compile(
r"^\s*[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}\s*$", re.IGNORECASE)
# A valid address is an explicit email target, not a channel name to resolve.
_PHOTON_DM_GUID_RE = re.compile(r"^any;-;\+\d{6,}$") # mirrors _DM_CHAT_GUID_RE in the photon adapter
# WhatsApp JIDs (@g.us, @s.whatsapp.net, @lid, broadcast/newsletter) and Buzz UUIDs are native targets
# the adapter accepts verbatim — explicit, never home-channel. A valid email address likewise.
_WHATSAPP_JID_RE = re.compile(r"^\s*[\w-]+@(?:g\.us|s\.whatsapp\.net|lid|broadcast|newsletter)\s*$", re.IGNORECASE)
_BUZZ_UUID_RE = re.compile(r"^\s*[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}\s*$", re.IGNORECASE)
_EMAIL_TARGET_RE = re.compile(r"^\s*[A-Za-z0-9._%+-]+@[A-Za-z0-9.-]+\.[A-Za-z]{2,}\s*$")
# Exceptions to "<PLATFORM>_HOME_CHANNEL" (email reads EMAIL_HOME_ADDRESS) for error hints.
# Exceptions to "<PLATFORM>_HOME_CHANNEL" for error hints (email reads EMAIL_HOME_ADDRESS).
_HOME_CHANNEL_ENV_OVERRIDES = {"email": "EMAIL_HOME_ADDRESS"}
_UNRESOLVED = object() # sentinel: stop parsing, target is NOT explicit (skip generic rules)
@@ -45,17 +39,19 @@ _UNRESOLVED = object() # sentinel: stop parsing, target is NOT explicit (skip g
# through to the generic rules in _parse_target_ref, or _UNRESOLVED.
def _parse_regex_groups(regex, *, thread_group=True):
"""Explicit when ``regex`` fully matches: chat_id = group 1, thread = group 2 (or None)."""
def parse(ref):
match = regex.fullmatch(ref)
return (match.group(1), match.group(2) if thread_group else None) if match else None
return parse
return lambda ref: ((m.group(1), m.group(2) if thread_group else None)
if (m := regex.fullmatch(ref)) else None)
def _parse_regex_stripped(regex):
"""Explicit when ``regex`` fully matches; returns the stripped ref verbatim."""
def parse(ref):
return (ref.strip(), None) if regex.fullmatch(ref) else None
return parse
return lambda ref: (ref.strip(), None) if regex.fullmatch(ref) else None
def _parse_nonempty(ref):
# ntfy topics and WeCom ids (the adapter picks the send command) are explicit when non-empty.
stripped = ref.strip()
return (stripped, None) if stripped else None
def _parse_telegram(ref):
@@ -68,20 +64,14 @@ def _parse_telegram(ref):
# (regex, chat_id template, thread comes from group 2) — thread form before bare id.
_SLACK_FORMS = (
(_SLACK_THREAD_TARGET_RE, "{}", True),
(_SLACK_TARGET_RE, "{}", False),
(_SLACK_USER_ID_RE, "user:{}", False),
(_SLACK_MENTION_RE, "user:{}", False),
(_SLACK_USER_NAME_RE, "user_name:{}", False))
_SLACK_FORMS = ((_SLACK_THREAD_TARGET_RE, "{}", True), (_SLACK_TARGET_RE, "{}", False),
(_SLACK_USER_ID_RE, "user:{}", False), (_SLACK_MENTION_RE, "user:{}", False),
(_SLACK_USER_NAME_RE, "user_name:{}", False))
def _parse_slack(ref):
for regex, template, has_thread in _SLACK_FORMS:
match = regex.fullmatch(ref)
if match:
return template.format(match.group(1)), (match.group(2) if has_thread else None)
return None
return next(((template.format(m.group(1)), m.group(2) if has_thread else None)
for regex, template, has_thread in _SLACK_FORMS if (m := regex.fullmatch(ref))), None)
def _parse_matrix(ref):
@@ -89,9 +79,7 @@ def _parse_matrix(ref):
# "@user" go via the generic rule so the numeric check keeps precedence.
trimmed = ref.strip()
split_idx = trimmed.rfind(":$")
if split_idx > 0:
return trimmed[:split_idx], trimmed[split_idx + 1 :]
return None
return (trimmed[:split_idx], trimmed[split_idx + 1:]) if split_idx > 0 else None
def _parse_yuanbao(ref):
@@ -99,24 +87,16 @@ def _parse_yuanbao(ref):
match = _YUANBAO_TARGET_RE.fullmatch(ref)
if match:
return match.group(1), None
if ref.strip().isdigit():
return f"group:{ref.strip()}", None
return _UNRESOLVED
def _parse_nonempty(ref):
# ntfy topics and WeCom ids (the adapter picks the send command) are explicit when non-empty.
stripped = ref.strip()
return (stripped, None) if stripped else None
return (f"group:{ref.strip()}", None) if ref.strip().isdigit() else _UNRESOLVED
def _parse_signal(ref):
# "group:<id>" is a native group target; an empty id is not explicit.
stripped = ref.strip()
if stripped.startswith("group:"):
group_id = stripped[len("group:"):].strip()
return (f"group:{group_id}", None) if group_id else _UNRESOLVED
return None
if not stripped.startswith("group:"):
return None
group_id = stripped[len("group:"):].strip()
return (f"group:{group_id}", None) if group_id else _UNRESOLVED
_PLATFORM_PARSERS = {
@@ -161,36 +141,28 @@ def resolve_send_target(
platform_name: str, target_ref: str, *, pass_unresolved_references: bool = False
) -> tuple[str | None, str | None, str | None]:
"""Resolve one send target the same way for every caller (model tool, CLI, cron).
Channel-directory IDs are trusted; plugin parsers are the authority on native syntax.
By default an unresolvable target is an error the model can read and pick a listed
target instead. ``pass_unresolved_references=True`` (no model in the loop: cron,
react/unreact on native message ids) hands an unresolvable target on a built-in
platform, or a plugin platform declaring no parser, to the adapter exactly as written;
a plugin platform WITH a parser stays strict for every caller. The optional validator
has the final say over parser-normalized, directory-resolved and passed-through IDs.
"""
Channel-directory IDs are trusted; plugin parsers are the authority on native syntax. By
default an unresolvable target is an error the model can act on. ``pass_unresolved_references``
(no model in the loop: cron, react/unreact on native ids) hands an unresolvable target on a
built-in platform, or a plugin platform without a parser, to the adapter as written; a plugin
WITH a parser stays strict. The optional validator has the final say over every returned id."""
from gateway.config import Platform
from gateway.platform_registry import platform_registry
entry = platform_registry.get(platform_name)
def _validate(candidate: str) -> str | None:
def _validated(chat_id, thread_id):
"""``(chat_id, thread_id, None)`` when the plugin validator (if any) accepts, else an error."""
if entry is None or entry.validate_target_ref_fn is None:
return None
return chat_id, thread_id, None
try:
verdict = entry.validate_target_ref_fn(candidate)
verdict = entry.validate_target_ref_fn(chat_id)
except Exception:
logger.debug("Plugin target validator failed for %s", platform_name, exc_info=True)
return f"Target validator failed for platform '{platform_name}'"
return None, None, f"Target validator failed for platform '{platform_name}'"
if verdict is True:
return None
return chat_id, thread_id, None
detail = f": {verdict}" if isinstance(verdict, str) and verdict else ""
return f"Invalid target '{target_ref}' on {platform_name}{detail}"
def _validated(chat_id, thread_id):
error = _validate(chat_id)
return (None, None, error) if error else (chat_id, thread_id, None)
return None, None, f"Invalid target '{target_ref}' on {platform_name}{detail}"
if entry is not None and entry.parse_target_ref_fn is not None:
try:
parsed = entry.parse_target_ref_fn(target_ref)
@@ -202,11 +174,9 @@ def resolve_send_target(
or not parsed[0] or (parsed[1] is not None and not isinstance(parsed[1], str))):
return None, None, f"Target parser for platform '{platform_name}' returned an invalid result"
return _validated(*parsed)
parsed_chat_id, parsed_thread_id, explicit = _parse_target_ref(platform_name, target_ref)
if explicit and parsed_chat_id is not None:
return _validated(parsed_chat_id, parsed_thread_id)
resolution_failed = False
try:
from gateway.channel_directory import resolve_channel_name
@@ -217,27 +187,21 @@ def resolve_send_target(
if resolved:
parsed_chat_id, parsed_thread_id, _ = _parse_target_ref(platform_name, resolved)
return _validated(parsed_chat_id or resolved, parsed_thread_id)
is_builtin = platform_name in {member.value for member in Platform}
if entry is None and not is_builtin:
return None, None, f"Unknown or unregistered plugin platform: {platform_name}"
def _pass_through_unresolved():
"""Hand the raw target to the adapter unchanged (it validates)."""
error = _validate(target_ref)
if error:
return None, None, error
logger.debug("Handing unresolved target '%s' to the %s adapter unchanged "
"(the adapter validates it)", target_ref, platform_name)
return target_ref, None, None
if entry is not None and entry.source == "plugin" and not is_builtin:
if pass_unresolved_references and entry.parse_target_ref_fn is None:
return _pass_through_unresolved()
return (None, None, f"Could not resolve '{target_ref}' on {platform_name}. "
"The plugin parser did not recognize it and no channel-directory entry matched.")
if pass_unresolved_references:
return _pass_through_unresolved()
hint = ("Try using a numeric channel ID instead." if resolution_failed
else "Use send_message(action='list') to see available targets.")
is_plugin = entry is not None and entry.source == "plugin" and not is_builtin
if pass_unresolved_references and (not is_plugin or entry.parse_target_ref_fn is None):
# Hand the raw target to the adapter unchanged (it validates).
chat_id, thread_id, error = _validated(target_ref, None)
if not error:
logger.debug("Handing unresolved target '%s' to the %s adapter unchanged "
"(the adapter validates it)", target_ref, platform_name)
return chat_id, thread_id, error
if is_plugin:
hint = "The plugin parser did not recognize it and no channel-directory entry matched."
elif resolution_failed:
hint = "Try using a numeric channel ID instead."
else:
hint = "Use send_message(action='list') to see available targets."
return None, None, f"Could not resolve '{target_ref}' on {platform_name}. {hint}"
+93 -167
View File
@@ -44,20 +44,18 @@ def send_message_tool(args, **kw):
def _resolve_tool_target(target: str, *, pass_unresolved_references: bool = False):
"""``(platform_name, chat_id, thread_id, error)`` for a ``platform[:ref]`` target;
``chat_id`` is None when no ref was given (caller falls back to the home channel)."""
"""``(platform_name, chat_id, thread_id, error)``; ``chat_id`` is None when no ref was given
(caller falls back to the home channel)."""
platform_name, _, target_ref = target.partition(":")
platform_name, target_ref = platform_name.strip().lower(), target_ref.strip() or None
prepare_send_message_platforms()
if not target_ref:
return platform_name, None, None, None
chat_id, thread_id, resolution_error = resolve_send_target(
platform_name, target_ref, pass_unresolved_references=pass_unresolved_references)
return platform_name, chat_id, thread_id, resolution_error
return platform_name, *resolve_send_target(platform_name, target_ref,
pass_unresolved_references=pass_unresolved_references)
def _handle_list():
"""Return formatted list of available messaging targets."""
try:
from gateway.channel_directory import format_directory_for_display
return json.dumps({"targets": format_directory_for_display()})
@@ -66,37 +64,28 @@ def _handle_list():
def _handle_react(args, remove=False):
"""Attach (``remove=True``: retract) an emoji reaction via a live gateway adapter's
``add_reaction`` / ``remove_reaction``. No standalone fallback: reacting needs the
adapter's live message-id state."""
target = args.get("target", "")
emoji = (args.get("emoji") or "").strip()
"""Attach (``remove=True``: retract) an emoji reaction via the live gateway adapter; no
standalone fallback because reacting needs the adapter's live message-id state."""
target, emoji = args.get("target", ""), (args.get("emoji") or "").strip()
message_id = (args.get("message_id") or "").strip() or None
if not target or (not remove and not emoji):
return tool_error("'target' is required when action='unreact'" if remove
else "Both 'target' and 'emoji' are required when action='react'")
# Platform-native ids (e.g. photon GUIDs) match no parser/directory entry; the
# adapter validates them.
platform_name, chat_id, _thread_id, resolution_error = _resolve_tool_target(
target, pass_unresolved_references=True)
# Platform-native ids (e.g. photon GUIDs) match no parser/directory entry; the adapter validates.
platform_name, chat_id, _thread_id, resolution_error = _resolve_tool_target(target, pass_unresolved_references=True)
if resolution_error:
return tool_error(resolution_error)
platform, err = _platform_enum(platform_name)
if err:
return tool_error(err)
if not chat_id:
try:
from gateway.config import load_gateway_config
home = load_gateway_config().get_home_channel(platform)
chat_id = load_gateway_config().get_home_channel(platform).chat_id
except Exception:
home = None
if not home:
return tool_error(f"No chat specified and no home channel set for {platform_name}. "
f"Use '{platform_name}:chat_id'.")
chat_id = home.chat_id
_, adapter = _live_adapter(platform)
if adapter is None:
return tool_error(f"Reactions require a live {platform_name} adapter in the running "
@@ -104,45 +93,34 @@ def _handle_react(args, remove=False):
react_fn = getattr(adapter, "remove_reaction" if remove else "add_reaction", None)
if not callable(react_fn):
return tool_error(f"Platform '{platform_name}' does not support message reactions.")
kwargs = {"chat_id": chat_id, "message_id": message_id, **({} if remove else {"emoji": emoji})}
try:
from model_tools import _run_async
result = _run_async(react_fn(**kwargs))
result = _run_async(react_fn(chat_id=chat_id, message_id=message_id, **({} if remove else {"emoji": emoji})))
except Exception as e:
return json.dumps(_error(f"Reaction failed: {e}"))
return json.dumps(result if isinstance(result, dict) else {"success": bool(result)})
def _handle_send(args):
"""Send a message to a platform target."""
target = args.get("target", "")
message = args.get("message", "")
target, message = args.get("target", ""), args.get("message", "")
if not target or not message:
return tool_error("Both 'target' and 'message' are required when action='send'")
platform_name, chat_id, thread_id, resolution_error = _resolve_tool_target(target)
if resolution_error:
return tool_error(resolution_error)
from tools.interrupt import is_interrupted
if is_interrupted():
return tool_error("Interrupted")
try:
from gateway.config import load_gateway_config
config = load_gateway_config()
except Exception as e:
return json.dumps(_error(f"Failed to load gateway config: {e}"))
platform, pconfig, entry, err = _resolve_platform_config(platform_name, config)
if err:
return tool_error(err)
from gateway.platforms.base import BasePlatformAdapter
# Capture [[as_document]] before extract_media strips it: images then go through
# send_document so the original bytes survive (Telegram's sendPhoto recompresses).
# Capture [[as_document]] before extract_media strips it (images keep original bytes via send_document).
force_document_attachments = "[[as_document]]" in message
media_files, cleaned_message = BasePlatformAdapter.extract_media(message)
media_files = BasePlatformAdapter.filter_media_delivery_paths(media_files)
@@ -152,30 +130,24 @@ def _handle_send(args):
chat_id, err = _home_chat_id(config, platform, platform_name)
if err:
return tool_error(err)
duplicate_skip = _maybe_skip_cron_duplicate_send(platform_name, chat_id, thread_id)
if duplicate_skip:
if duplicate_skip := _maybe_skip_cron_duplicate_send(platform_name, chat_id, thread_id):
return json.dumps(duplicate_skip)
if platform_name == "slack" and chat_id:
chat_id, resolve_err = _slack_dm_chat_id(pconfig, chat_id)
if resolve_err:
return json.dumps(resolve_err)
try:
from model_tools import _run_async
send_kwargs = {"thread_id": thread_id, "media_files": media_files,
"force_document": force_document_attachments}
# Only custom plugin handlers receive the complete typed request.
if entry is not None and entry.send_message_handler is not None:
send_kwargs["args"] = args
result = _run_async(_send_to_platform(platform, pconfig, chat_id, cleaned_message, **send_kwargs))
handler_args = {"args": args} if entry is not None and entry.send_message_handler is not None else {}
result = _run_async(_send_to_platform(platform, pconfig, chat_id, cleaned_message, thread_id=thread_id,
media_files=media_files, force_document=force_document_attachments,
**handler_args))
if isinstance(result, dict) and result.get("success"):
if used_home_channel:
result["note"] = f"Sent to {platform_name} home channel (chat_id: {chat_id})"
if mirror_text and _mirror_sent_message(platform_name, chat_id, mirror_text, thread_id):
result["mirrored"] = True
if isinstance(result, dict) and "error" in result:
result["error"] = _sanitize_error_text(result["error"])
return json.dumps(result)
@@ -193,9 +165,8 @@ def _platform_enum(platform_name):
def _resolve_platform_config(platform_name, config):
"""``(platform, pconfig, registry_entry, error)`` for a send. Plugin platforms must be
registered; disabled/missing platforms error, except Weixin, which may be configured
purely via .env (synthesized pconfig so cron delivery works without gateway.yaml)."""
"""``(platform, pconfig, registry_entry, error)``. Plugin platforms must be registered;
disabled/missing platforms error, except Weixin, which may be configured purely via .env."""
from gateway.config import Platform
from gateway.platform_registry import platform_registry
entry = platform_registry.get(platform_name)
@@ -204,7 +175,6 @@ def _resolve_platform_config(platform_name, config):
platform, err = _platform_enum(platform_name)
if err:
return None, None, None, err
pconfig = config.platforms.get(platform)
if not pconfig or not pconfig.enabled:
pconfig = _weixin_env_pconfig() if platform_name == "weixin" else None
@@ -215,8 +185,7 @@ def _resolve_platform_config(platform_name, config):
def _home_chat_id(config, platform, platform_name):
"""Return ``(home chat_id, None)`` or ``(None, actionable error)``.
Weixin additionally honours the WEIXIN_HOME_CHANNEL env var."""
"""``(home chat_id, None)`` or ``(None, actionable error)``; Weixin also honours WEIXIN_HOME_CHANNEL."""
home = config.get_home_channel(platform)
if home:
return home.chat_id, None
@@ -230,9 +199,8 @@ def _home_chat_id(config, platform, platform_name):
def _slack_dm_chat_id(pconfig, chat_id):
"""Open Slack user targets (``user:U...`` / ``user_name:@handle`` from the parser, or a
bare U... id from session metadata / home-channel config) as DM conversations —
chat.postMessage needs a conversation ID. ``(chat_id, None)`` or ``(None, error_dict)``."""
"""Open Slack user targets (``user:``/``user_name:`` from the parser, or a bare U... id from
session metadata / home-channel config) as DM conversations. ``(chat_id, None)`` or ``(None, error_dict)``."""
dm_target = f"user:{chat_id}" if chat_id.startswith("U") and _SLACK_USER_ID_RE.fullmatch(chat_id) else chat_id
if not dm_target.startswith(("user:", "user_name:")):
return chat_id, None
@@ -261,8 +229,7 @@ def _weixin_env_pconfig():
return None
from gateway.config import PlatformConfig
return PlatformConfig(enabled=True, token=wx_token, extra={
"account_id": wx_account,
"base_url": get_secret("WEIXIN_BASE_URL", "").strip(),
"account_id": wx_account, "base_url": get_secret("WEIXIN_BASE_URL", "").strip(),
"cdn_base_url": get_secret("WEIXIN_CDN_BASE_URL", "").strip()})
@@ -286,14 +253,11 @@ def _maybe_skip_cron_duplicate_send(platform_name: str, chat_id: str, thread_id:
from gateway.session_context import get_session_env
auto_platform = get_session_env("HERMES_CRON_AUTO_DELIVER_PLATFORM", "").strip().lower()
auto_chat_id = get_session_env("HERMES_CRON_AUTO_DELIVER_CHAT_ID", "").strip()
if not auto_platform or not auto_chat_id:
return None
auto_thread_id = get_session_env("HERMES_CRON_AUTO_DELIVER_THREAD_ID", "").strip() or None
if not (auto_platform == platform_name and auto_chat_id == str(chat_id) and auto_thread_id == thread_id):
if not (auto_platform and auto_chat_id and auto_platform == platform_name and auto_chat_id == str(chat_id)
and (get_session_env("HERMES_CRON_AUTO_DELIVER_THREAD_ID", "").strip() or None) == thread_id):
return None
target_label = f"{platform_name}:{chat_id}" + (f":{thread_id}" if thread_id is not None else "")
return {
"success": True, "skipped": True, "reason": "cron_auto_delivery_duplicate_target", "target": target_label,
return {"success": True, "skipped": True, "reason": "cron_auto_delivery_duplicate_target", "target": target_label,
"note": (f"Skipped send_message to {target_label}. This cron job will already auto-deliver "
"its final response to that same target. Put the intended user-facing content in "
"your final response instead, or use a different target if you want an additional message.")}
@@ -305,28 +269,25 @@ def _bounded_send_error(detail, max_chars=900):
return text if len(text) <= max_chars else f"{text[: max_chars - 3]}..."
async def _send_live_adapter_media(
adapter, chat_id, message, media_files, *, thread_id=None, metadata=None, force_document=False):
"""Deliver text and every media descriptor through adapter media APIs. Adapters that
only inherit the BasePlatformAdapter stub for a kind are unsupported, not no-op'd."""
caption, separate_text = _media_caption_split(
message, media_files, max_caption_len=_DEFAULT_CAPTION_LIMIT)
async def _send_live_adapter_media(adapter, chat_id, message, media_files, *, thread_id=None, metadata=None,
force_document=False):
"""Deliver text and every media descriptor through adapter media APIs; adapters that only
inherit the BasePlatformAdapter stub for a kind are unsupported, not no-op'd."""
caption, separate_text = _media_caption_split(message, media_files, max_caption_len=_DEFAULT_CAPTION_LIMIT)
last_result = None
if separate_text and separate_text.strip():
last_result = await adapter.send(chat_id=chat_id, content=separate_text, metadata=metadata)
if not last_result.success:
return {"error": f"Adapter send failed: {_bounded_send_error(last_result.error)}"}
from gateway.platforms.base import BasePlatformAdapter
total = len(media_files)
for index, descriptor in enumerate(media_files):
media_path = descriptor[0] if isinstance(descriptor, (list, tuple)) and descriptor else None
if not isinstance(media_path, str) or not media_path:
return {"error": f"Adapter media send failed: invalid media descriptor {index + 1}/{total}"}
is_voice = bool(descriptor[1]) if len(descriptor) > 1 else False
is_voice = len(descriptor) > 1 and bool(descriptor[1])
if not os.path.exists(media_path):
return {"error": f"Adapter media send failed: media file {index + 1}/{total} was not found"}
ext = os.path.splitext(media_path)[1].lower()
method_name, media_kind = _adapter_media_method(ext, is_voice or ext in _AUDIO_EXTS, force_document)
adapter_method = getattr(type(adapter), method_name, None)
@@ -345,22 +306,16 @@ async def _send_live_adapter_media(
continue
detail = _bounded_send_error(last_result.error or "media send failed")
return {"error": f"Adapter media send failed after {index}/{total} files: {detail}"}
if last_result is None:
return {"error": _NO_DELIVERABLE}
return {"success": True, "message_id": last_result.message_id, "media_delivered": True}
async def _dispatch_on_gateway_loop(runner, make_coro, log_message):
"""Await ``make_coro()`` on the gateway's loop. adapter.send() uses queues/tasks bound
to that loop; awaiting it from another loop (the tool worker thread) deadlocks, so
cross-loop calls are scheduled threadsafe onto it."""
"""Await ``make_coro()`` on the gateway's loop: adapter.send() uses queues/tasks bound to it,
so awaiting from another loop (the tool worker thread) deadlocks."""
gateway_loop = getattr(runner, "_gateway_loop", None)
try:
current_loop = asyncio.get_running_loop()
except RuntimeError:
current_loop = None
if gateway_loop is None or current_loop is gateway_loop:
if gateway_loop is None or asyncio.get_running_loop() is gateway_loop:
return await make_coro() # same loop / no gateway loop (CLI, tests)
if not gateway_loop.is_running():
return {"error": "Gateway loop is not running; cannot dispatch adapter send"}
@@ -368,31 +323,29 @@ async def _dispatch_on_gateway_loop(runner, make_coro, log_message):
fut = safe_schedule_threadsafe(make_coro(), gateway_loop, logger=logger, log_message=log_message)
if fut is None:
return {"error": "Gateway loop unavailable for send dispatch"}
# shield: a cancelled caller (agent interrupt) must not cancel the enqueued send or a
# retry would duplicate it. No timeout: the adapter and outer _run_async bound the wait.
# shield: a cancelled caller must not cancel the enqueued send (a retry would duplicate it).
# No timeout: the adapter and outer _run_async bound the wait.
return await asyncio.shield(asyncio.wrap_future(fut))
async def _send_via_adapter(
platform, pconfig, chat_id, chunk, *, thread_id=None, media_files=None, force_document=False):
"""Live in-process gateway adapter first, else the plugin's ``standalone_sender_fn``
(gateway not in this process, e.g. cron), else an error naming both options. Media
goes through the adapter's native media APIs under the same cross-loop rules."""
async def _send_via_adapter(platform, pconfig, chat_id, chunk, *, thread_id=None, media_files=None,
force_document=False):
"""Live in-process gateway adapter first, else the plugin's ``standalone_sender_fn`` (cron),
else an error naming both; media uses the adapter's native media APIs under the same rules."""
platform_name = platform.value if hasattr(platform, "value") else str(platform)
runner, adapter = _live_adapter(platform)
if adapter is not None:
try:
metadata = {**({"thread_id": thread_id} if thread_id else {}),
**({"publish_topic": chat_id} if platform_name == "ntfy" and chat_id else {})} or None
if media_files:
return await _dispatch_on_gateway_loop(
runner, lambda: _send_live_adapter_media(
adapter, chat_id, chunk, media_files,
thread_id=thread_id, metadata=metadata, force_document=force_document),
"send_message: failed to schedule media send on gateway loop")
if media_files: # always a dict result, returned as-is below
make_coro = lambda: _send_live_adapter_media( # noqa: E731
adapter, chat_id, chunk, media_files, thread_id=thread_id, metadata=metadata,
force_document=force_document)
else:
make_coro = lambda: adapter.send(chat_id=chat_id, content=chunk, metadata=metadata) # noqa: E731
result = await _dispatch_on_gateway_loop(
runner, lambda: adapter.send(chat_id=chat_id, content=chunk, metadata=metadata),
"send_message: failed to schedule on gateway loop")
runner, make_coro, f"send_message: failed to schedule{' media send' if media_files else ''} on gateway loop")
except asyncio.CancelledError:
raise
except Exception as e:
@@ -402,48 +355,42 @@ async def _send_via_adapter(
if result.success:
return {"success": True, "message_id": result.message_id}
return {"error": f"Adapter send failed: {_bounded_send_error(result.error)}"}
try:
from gateway.platform_registry import platform_registry
entry = platform_registry.get(platform_name)
sender = platform_registry.get(platform_name).standalone_sender_fn
except Exception:
entry = None
if entry is None or entry.standalone_sender_fn is None:
return {"error": (
f"No live adapter for platform '{platform_name}'. Is the gateway running with this platform "
f"connected? For out-of-process delivery (e.g. cron in a separate process), the platform "
f"plugin must register a standalone_sender_fn on its PlatformEntry.")}
sender = None
if sender is None:
return {"error": (f"No live adapter for platform '{platform_name}'. Is the gateway running with this platform "
f"connected? For out-of-process delivery (e.g. cron in a separate process), the platform "
f"plugin must register a standalone_sender_fn on its PlatformEntry.")}
try:
result = await entry.standalone_sender_fn(pconfig, chat_id, chunk, thread_id=thread_id,
media_files=media_files, force_document=force_document)
result = await sender(pconfig, chat_id, chunk, thread_id=thread_id, media_files=media_files,
force_document=force_document)
except asyncio.CancelledError:
raise
except Exception as e:
logger.debug("Plugin standalone send for %s raised", platform_name, exc_info=True)
return {"error": f"Plugin standalone send failed: {_bounded_send_error(e)}"}
if isinstance(result, dict) and (result.get("success") or result.get("error")):
if result.get("error"):
return {**result, "error": _bounded_send_error(result["error"])}
return result
return {**result, "error": _bounded_send_error(result["error"])} if result.get("error") else result
return {"error": (f"Plugin standalone send for '{platform_name}' returned an invalid result: "
f"expected a dict with 'success' or 'error' keys, got {type(result).__name__}")}
async def _send_chunks(chunks, send_one):
"""``send_one(chunk, is_last)`` in order; stop at the first error dict, else last result."""
last_result = None
result = None
for i, chunk in enumerate(chunks):
result = await send_one(chunk, i == len(chunks) - 1)
if isinstance(result, dict) and result.get("error"):
return result
last_result = result
return last_result
break
return result
def _platform_max_length(platform):
"""Chunking limit: the adapter constant for Signal (its raw JSON-RPC path bypasses
SignalAdapter's own chunking), the registry's ``max_message_length`` for plugins,
else None (no chunking)."""
"""Chunking limit: Signal's adapter constant (its raw JSON-RPC path bypasses the adapter's
chunking), the registry's ``max_message_length`` for plugins, else None (no chunking)."""
from gateway.config import Platform
if platform == Platform.SIGNAL:
try:
@@ -459,24 +406,18 @@ def _platform_max_length(platform):
return None
# Plugin platforms whose media (Discord: all) sends bypass the live adapter on purpose for
# the registry ``standalone_sender_fn``: Discord's handles forum channels/threads/multipart
# uploads; Slack uploads via files_upload_v2; WhatsApp posts to the Baileys bridge
# /send-media so media arrive as native bubbles.
# platform -> (error label, run discover_plugins first, caption-capable,
# media_files sentinel for non-final chunks, forward force_document)
_PLUGIN_STANDALONE_MEDIA = {
"discord": ("Discord", False, True, [], False),
"feishu": ("Feishu", True, False, None, False),
"slack": ("Slack", True, True, [], False),
"whatsapp": ("WhatsApp", True, True, None, True)}
# Plugin platforms whose media (Discord: all) sends deliberately bypass the live adapter for the
# registry ``standalone_sender_fn`` (Discord: forums/threads/multipart; Slack: files_upload_v2;
# WhatsApp: Baileys /send-media). platform -> (error label, run discover_plugins first,
# caption-capable, media_files sentinel for non-final chunks, forward force_document)
_PLUGIN_STANDALONE_MEDIA = {"discord": ("Discord", False, True, [], False), "feishu": ("Feishu", True, False, None, False),
"slack": ("Slack", True, True, [], False), "whatsapp": ("WhatsApp", True, True, None, True)}
async def _send_plugin_standalone(
platform_name, pconfig, chat_id, message, chunks, media_files, *, thread_id, max_len, force_document
):
"""Chunked send through a plugin's standalone_sender_fn; a single captionable file +
short text rides as the media caption."""
async def _send_plugin_standalone(platform_name, pconfig, chat_id, message, chunks, media_files, *, thread_id,
max_len, force_document):
"""Chunked send through a plugin's standalone_sender_fn; one captionable file + short text
rides as the media caption."""
label, discover, captionable, empty_media, pass_force = _PLUGIN_STANDALONE_MEDIA[platform_name]
sender, err = _plugin_standalone_sender(platform_name, label=label, discover=discover)
if err:
@@ -492,30 +433,25 @@ async def _send_plugin_standalone(
pconfig, chat_id, chunk, thread_id=thread_id, media_files=media_files if is_last else empty_media, **extra))
# Native-media chunked routes for built-in platforms; media rides on the final chunk,
# non-final chunks get the empty-media sentinel. platform -> (media required, sentinel,
# sender(platform, pconfig, chat_id, chunk, media, thread_id, force_document)).
# Matrix: ALL sends use the native adapter so text is encrypted in E2EE rooms too.
# Signal: attachments ride the JSON-RPC ``attachments`` param.
# Yuanbao / WeCom: media needs the running gateway adapter.
# Slack (text; media intercepted above): prefer the live adapter — multi-workspace aware,
# honors gates like ignored_channels — else the plugin's standalone sender.
# Names resolve at call time so tests can monkeypatch e.g. ``_send_signal``.
def _via_adapter_route(p, pc, cid, chunk, media, tid, fd):
return _send_via_adapter(p, pc, cid, chunk, thread_id=tid, media_files=media, force_document=fd)
# Native-media chunked routes for built-in platforms; media rides on the final chunk, non-final
# chunks get the sentinel. platform -> (media required, sentinel, sender(platform, pconfig,
# chat_id, chunk, media, thread_id, force_document)). Matrix: ALL sends use the native adapter
# (E2EE text). Signal: attachments ride the JSON-RPC param. Yuanbao / WeCom: media needs the
# running gateway. Slack text: live adapter (multi-workspace, ignored_channels gates) else the
# plugin's standalone sender. Names resolve at call time so tests can monkeypatch ``_send_signal``.
_CHUNKED_ROUTES = {
"matrix": (False, [], lambda p, pc, cid, chunk, media, tid, fd: _send_matrix_via_adapter(
pc, cid, chunk, media_files=media, thread_id=tid)),
"signal": (True, [], lambda p, pc, cid, chunk, media, tid, fd: _send_signal(
pc.extra, cid, chunk, media_files=media)),
"yuanbao": (True, None, lambda p, pc, cid, chunk, media, tid, fd: _send_yuanbao(
cid, chunk, media_files=media)),
"yuanbao": (True, None, lambda p, pc, cid, chunk, media, tid, fd: _send_yuanbao(cid, chunk, media_files=media)),
"slack": (False, [], _via_adapter_route),
"wecom": (True, None, _via_adapter_route)}
# Text-only senders for built-in platforms (generic path; media is dropped with a
# warning). Signature: (pconfig, chat_id, chunk, thread_id) -> result.
_TEXT_SENDERS = {
@@ -530,66 +466,56 @@ _MEDIA_PLATFORMS_NOTE = "telegram, discord, matrix, weixin, signal, yuanbao, fei
async def _send_to_platform(platform, pconfig, chat_id, message, thread_id=None, media_files=None, force_document=False, args=None):
"""Route a message to the platform sender, chunking long text with the adapters' smart
splitter. Branch order matters: Weixin first (its native helper must not be blocked by
unrelated optional imports such as lark-oapi), Telegram (chunks itself), plugin
standalone media routes, native chunked routes, then the generic text path."""
"""Route to the platform sender, chunking long text with the adapters' splitter. Order matters:
Weixin first (its native helper must not be blocked by unrelated optional imports such as
lark-oapi), Telegram (chunks itself), plugin standalone media, native chunked, generic text."""
from gateway.config import Platform
platform_name = platform.value if hasattr(platform, "value") else str(platform)
media_files = media_files or []
if platform == Platform.WEIXIN:
return await _send_weixin(pconfig, chat_id, message, media_files=media_files)
# Telegram chunks internally on the *formatted* text (escaping inflates length).
if platform == Platform.TELEGRAM:
disable_link_previews = bool(getattr(pconfig, "extra", {}) and pconfig.extra.get("disable_link_previews"))
return await _send_telegram(pconfig.token, chat_id, message, media_files=media_files, thread_id=thread_id,
disable_link_previews=disable_link_previews, force_document=force_document)
return await _send_telegram(
pconfig.token, chat_id, message, media_files=media_files, thread_id=thread_id, force_document=force_document,
disable_link_previews=bool(getattr(pconfig, "extra", {}) and pconfig.extra.get("disable_link_previews")))
from gateway.platforms.base import BasePlatformAdapter
max_len = _platform_max_length(platform)
chunks = BasePlatformAdapter.truncate_message(message, max_len) if max_len else [message]
if platform_name == "discord" or (media_files and platform_name in _PLUGIN_STANDALONE_MEDIA):
return await _send_plugin_standalone(platform_name, pconfig, chat_id, message, chunks, media_files,
thread_id=thread_id, max_len=max_len, force_document=force_document)
route = _CHUNKED_ROUTES.get(platform_name)
if route is not None and (media_files or not route[0]):
_, empty_media, sender = route
return await _send_chunks(chunks, lambda chunk, is_last: sender(
platform, pconfig, chat_id, chunk, media_files if is_last else empty_media, thread_id, force_document))
# Generic path: text only. Buzz has verified native media delivery through
# _send_via_adapter (media-only sends included), so it is exempt from the warning.
# Generic path: text only. Buzz delivers media natively via _send_via_adapter, so no warning.
warning = None
if media_files and platform_name != "buzz":
if not message.strip():
return {"error": (
f"send_message MEDIA delivery is currently only supported for {_MEDIA_PLATFORMS_NOTE}; "
f"target {platform_name} had only media attachments")}
warning = (
f"MEDIA attachments were omitted for {platform_name}; "
f"native send_message media delivery is currently only supported for {_MEDIA_PLATFORMS_NOTE}")
return {"error": (f"send_message MEDIA delivery is currently only supported for {_MEDIA_PLATFORMS_NOTE}; "
f"target {platform_name} had only media attachments")}
warning = (f"MEDIA attachments were omitted for {platform_name}; "
f"native send_message media delivery is currently only supported for {_MEDIA_PLATFORMS_NOTE}")
text_sender = _TEXT_SENDERS.get(platform_name)
if text_sender is not None:
send_one = lambda chunk, is_last: text_sender(pconfig, chat_id, chunk, thread_id) # noqa: E731
else:
from gateway.platform_registry import platform_registry
entry = platform_registry.get(platform_name)
handler = entry.send_message_handler if entry is not None else None
if handler is not None:
if entry is not None and entry.send_message_handler is not None:
# Custom handler receives the full typed request once (not per chunk).
try:
import inspect
result = handler(args or {}, chat_id, platform_name, pconfig)
result = entry.send_message_handler(args or {}, chat_id, platform_name, pconfig)
return await result if inspect.isawaitable(result) else result
except Exception as e:
return {"error": f"Plugin send_message handler failed: {e}"}
# Plugin platform: live gateway adapter if available, else standalone_sender_fn.
send_one = lambda chunk, is_last: _via_adapter_route( # noqa: E731
platform, pconfig, chat_id, chunk, media_files if is_last else [], thread_id, force_document)
last_result = await _send_chunks(chunks, send_one)
if (warning and isinstance(last_result, dict) and last_result.get("success")
and not last_result.get("media_delivered")):
+160 -244
View File
@@ -15,33 +15,27 @@ 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 from browsing/searching — integrations (HERMES_SESSION_SOURCE=tool), delegate
# subagent runs, kanban workers are not the user's history.
_HIDDEN_SESSION_SOURCES = ("kanban", "subagent", "tool")
# Searchable but DEMOTED below interactive sessions: cron vocabulary dominates bare
# BM25 and starves out the user's own sessions ("recall blindness").
_DEMOTED_SESSION_SOURCES = ("cron",)
# FTS rows scanned before dedup-by-lineage — well above the distinct sessions a query
# returns, so interactive matches buried under cron hits survive the demotion pass.
_DISCOVER_SCAN_LIMIT = 300
# Raw FTS rows are only a plan input; the response hydrates its own window/bookends.
_DISCOVER_SEARCH_FIELDS = ("id", "session_id", "role", "snippet", "source", "model", "session_started")
# Compaction handoff summaries (agent/context_compressor.py); excluded from bookends.
_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 the two cannot drift.
# /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 the two cannot drift.
_FRESH_RESET_END_REASONS = frozenset(_RESET_END_REASONS) | {"new_session"}
def _quiet(fn, default, msg, *log_args, with_exc: bool = False):
"""``fn()``, or *default* after debug-logging *msg* (exception appended as a
final ``%s`` arg when *with_exc*) on any exception."""
"""``fn()``, or *default* after debug-logging *msg* (+ the exception when *with_exc*)."""
try:
return fn()
except Exception as e:
@@ -50,8 +44,8 @@ def _quiet(fn, default, msg, *log_args, with_exc: bool = False):
def _loud(fn, log_msg, error_prefix, *log_args):
"""``(value, None)`` from ``fn()``, or ``(None, tool_error_json)`` after an
error-level log — for DB calls whose failure the model must see."""
"""``(fn(), None)``, or ``(None, tool_error_json)`` after an error-level log — for DB
calls whose failure the model must see."""
try:
return fn(), None
except Exception as e:
@@ -60,46 +54,43 @@ def _loud(fn, log_msg, error_prefix, *log_args):
def _format_timestamp(ts: Union[int, float, str, None]) -> str:
"""Unix timestamp (number / numeric string) -> readable date; ISO strings pass
through; "unknown" for None; str(ts) if conversion fails."""
"""Unix timestamp -> readable date; ISO strings pass through; "unknown" for None."""
if ts is None:
return "unknown"
if isinstance(ts, str) and not ts.replace(".", "").replace("-", "").isdigit():
return ts
try:
return datetime.fromtimestamp(float(ts)).strftime("%B %d, %Y at %I:%M %p")
except Exception as e:
logging.debug("Failed to format timestamp %s: %s", ts, e, exc_info=True)
return str(ts)
return _quiet(lambda: datetime.fromtimestamp(float(ts)).strftime("%B %d, %Y at %I:%M %p"), str(ts),
"Failed to format timestamp %s: %s", ts, with_exc=True)
def _get_session_meta(db, session_id: str) -> dict:
"""``db.get_session`` that degrades to ``{}`` on error."""
return _quiet(lambda: db.get_session(session_id), None,
"get_session failed for %s: %s", session_id, with_exc=True) or {}
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 root -> ``(root_id, has_compression_hop)``. The
flag separates a compression-split lineage (parent summarised away) from a
delegation lineage (child still visible to the parent). Errors -> ``(session_id, False)``."""
"""Walk parent_session_id to the root -> ``(root_id, has_compression_hop)``; the flag
separates a compression-split lineage (parent summarised away) from a delegation
lineage (child still visible to the parent)."""
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
s = _get_session_meta(db, cur)
has_compression = has_compression or s.get("end_reason") == "compression"
if not s.get("parent_session_id"):
break
@@ -108,31 +99,31 @@ def _resolve_to_parent(db, session_id: str) -> tuple[str, bool]:
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 _same_lineage(db, a: str, b: str) -> bool:
a_root = _resolve_lineage(db, a)
return bool(a_root and a_root == _resolve_lineage(db, b))
def _session_left_live_context(db, session_id: str) -> bool:
"""True when the transcript left everyone's live context: ``compression``
(summarised into the child) or a fresh reset (child starts empty). Live delegation
children (``end_reason is None``) and ``branched`` parents (copied verbatim into
the branch) ARE the current context, so they stay excluded from recall."""
s = session_id and _quiet(lambda: db.get_session(session_id), None, "get_session failed for %s", session_id)
end_reason = (s.get("end_reason") or None) if s else None
end_reason = (session_id and _get_session_meta(db, session_id).get("end_reason")) or None
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
"""Owning session and visibility flags for *message_id* (None if missing/error)."""
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
row = message_id and _quiet(_lookup, None, "message storage-state lookup failed for %s", message_id)
return dict(row) if row else None
def _is_compacted_state(state: Optional[Dict[str, Any]]) -> bool:
@@ -142,33 +133,14 @@ def _is_compacted_state(state: Optional[Dict[str, Any]]) -> bool:
def _is_compacted_message(db, message_id) -> bool:
"""True for a compaction-archived row — content no longer in live context, so
it stays discoverable even on the current session. False on any error."""
"""True for a compaction-archived row: no longer in live context, so discoverable
even on the current session. False on any error."""
return _is_compacted_state(_get_message_storage_state(db, message_id))
def _annotate_rebuild_status(db, payload: Dict[str, Any]) -> None:
"""Note rebuild progress while the deferred FTS backfill runs so the agent can
explain thin results instead of treating them as ground truth. Never raises."""
status = _quiet(db.fts_rebuild_status, None, "fts_rebuild_status failed")
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 so interactive sessions rank above automation; BM25 order is
kept within each class, so a cron hit never displaces an interactive one."""
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. Keeps ``content`` even when empty (tool-call-only
assistant turns); with *max_content_len* truncates and flags it."""
"""Slim a message row; keeps ``content`` even when empty (tool-call-only turns)."""
content = m.get("content")
if isinstance(content, str) and "\x1b" in content: # archived terminal output carries ANSI
from tools.ansi_strip import strip_ansi
@@ -178,7 +150,8 @@ def _shape_message(m: Dict[str, Any], anchor_id: Optional[int] = None,
if anchor_id is not None and m.get("id") == anchor_id:
entry["anchor"] = True
if max_content_len and content and len(content) > max_content_len:
entry.update(content=content[:max_content_len] + "…", content_truncated=True, original_content_chars=len(content))
entry.update(content=content[:max_content_len] + "…", content_truncated=True,
original_content_chars=len(content))
return {k: v for k, v in entry.items() if v is not None or k == "content"}
@@ -186,23 +159,29 @@ def _session_link(session_id: str, profile: str = None) -> str:
"""The reference the agent writes for 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, just not 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")
def _active():
from hermes_cli.profiles import get_active_profile_name
resolved = get_active_profile_name()
return "" if resolved == "custom" else resolved
name = (profile or "").strip() or _quiet(_active, "", "get_active_profile_name failed for session link")
return f"@session:{name}/{session_id}" if name else f"@session:{session_id}"
def _discovery_entry(lineage_root: Optional[str], **fields) -> Dict[str, Any]:
"""Canonical key order; ``parent_session_id`` set when the hit lives in a child."""
entry = {k: fields[k] for k in (
"session_id", "when", "source", "model", "title", "matched_role", "match_message_id", "snippet",
"bookend_start", "messages", "bookend_end", "messages_before", "messages_after", "detail")}
if lineage_root and lineage_root != entry["session_id"]:
entry["parent_session_id"] = lineage_root
return entry
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."""
"""Discovery-shaped result when the query matches a session title, else None."""
title_query = query.strip().strip("`'\"") # models often quote a remembered title
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)
session_id = title_query and _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)
@@ -220,56 +199,28 @@ def _title_match_result(db, query: str, current_lineage_root: Optional[str]) ->
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)
title = session_meta.get("title") or title_query
entry = _discovery_entry(
def shape(key, fallback, anchor=None):
return [_shape_message(m, anchor_id=anchor) for m in (view.get(key) or fallback)]
return {**_discovery_entry(
lineage_root, 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=title, matched_role="session_title", match_message_id=anchor_id,
snippet=f"Session title matched: {title}",
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")
entry["_lineage_root"] = lineage_root
return entry
def _discovery_entry(lineage_root: Optional[str], **fields) -> Dict[str, Any]:
"""One discovery result in canonical key order; ``parent_session_id`` is set
when the hit lives in a child of its lineage root."""
entry = {k: fields[k] for k in (
"session_id", "when", "source", "model", "title", "matched_role", "match_message_id", "snippet",
"bookend_start", "messages", "bookend_end", "messages_before", "messages_after", "detail")}
if lineage_root and lineage_root != entry["session_id"]:
entry["parent_session_id"] = lineage_root
return entry
bookend_start=shape("bookend_start", messages[:3]), messages=shape("window", messages[:5], anchor_id),
bookend_end=shape("bookend_end", 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}
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. Current-lineage hits are skipped UNLESS the transcript left live
context (compression-ended, /new-reset predecessor, or an in-place compacted row
on the SAME session); a live delegation child (end_reason=None) stays excluded."""
for r in raw_results:
if len(seen_sessions) >= limit:
break
raw_sid = r["session_id"]
resolved_sid = _resolve_lineage(db, raw_sid)
is_compacted_hit = _is_compacted_message(db, r.get("id"))
in_current_lineage = bool(current_lineage_root) and resolved_sid == current_lineage_root
if in_current_lineage and not (_session_left_live_context(db, raw_sid) 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})
"""Discovery response; notes FTS backfill progress so the agent can explain thin
results instead of treating them as ground truth."""
status = _quiet(db.fts_rebuild_status, None, "fts_rebuild_status failed")
rebuild = {} if status is None else {"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.")}}
return _ok(mode="discover", query=query, detail=detail, results=results, count=len(results), **extra, **rebuild)
def _bookend(view: Dict[str, Any], key: str) -> List[Dict[str, Any]]:
@@ -278,18 +229,14 @@ def _bookend(view: Dict[str, Any], key: str) -> List[Dict[str, Any]]:
def _hydrate_hit(db, lineage_root: str, match_info: Dict[str, Any], result_detail: str) -> Optional[Dict[str, Any]]:
"""One discovery result from a surviving FTS row; None (hit dropped) if the
anchored view can't be loaded."""
hit_sid = match_info.get("session_id") or lineage_root
msg_id = match_info.get("id")
"""Discovery result from a surviving FTS row; None (dropped) if the view can't load."""
hit_sid, msg_id = match_info.get("session_id") or lineage_root, 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 = [m for m in (view.get("window") or []) if full or m.get("id") == msg_id]
session_meta, full = _get_session_meta(db, lineage_root), result_detail == "full"
return _discovery_entry(
lineage_root, session_id=hit_sid,
when=_format_timestamp(session_meta.get("started_at") or match_info.get("session_started")),
@@ -298,7 +245,8 @@ def _hydrate_hit(db, lineage_root: str, match_info: Dict[str, Any], result_detai
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],
messages=[_shape_message(m, anchor_id=msg_id, max_content_len=4000)
for m in (view.get("window") or []) if full or m.get("id") == msg_id],
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)
@@ -315,23 +263,35 @@ def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort
fields=_DISCOVER_SEARCH_FIELDS), "FTS5 search failed: %s", "Search failed")
if err:
return err
# 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)
# 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`; stable sort keeps BM25
# order within each class.
raw_results = sorted(raw_results, key=lambda r: (r.get("source") or "") in _DEMOTED_SESSION_SOURCES)
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)
results = [title_result] if title_result else []
if title_result and (title_lineage := title_result.pop("_lineage_root", None)):
seen_sessions[title_lineage] = {"_title_only": True}
# Dedupe by lineage (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.
# Current-lineage hits are skipped UNLESS the transcript left live context
# (compression-ended, /new-reset predecessor, or an in-place compacted row on the
# SAME session); a live delegation child (end_reason=None) stays excluded.
for r in raw_results:
if len(seen_sessions) >= limit:
break
raw_sid, resolved_sid = r["session_id"], _resolve_lineage(db, r["session_id"])
is_compacted_hit = _is_compacted_message(db, r.get("id"))
if current_lineage_root and resolved_sid == current_lineage_root and not (
_session_left_live_context(db, raw_sid) 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})
for lineage_root, match_info in seen_sessions.items():
if match_info.get("_title_only"):
continue
@@ -350,8 +310,7 @@ def _discover(db, query: str, role_filter: Optional[List[str]], limit: int, sort
def _resolve_profile_db(profile: str):
"""Another profile's ``state.db`` opened read-only (safe on a live DB), or None
for the current profile."""
"""Another profile's ``state.db`` opened read-only (safe on a live DB); None = current."""
if profile is None or not str(profile).strip():
return None
from hermes_cli import profiles as profiles_mod
@@ -364,17 +323,17 @@ def _resolve_profile_db(profile: str):
def _locate_session_db(session_id: str):
"""Scan every profile's ``state.db`` for a session id -> ``(db, profile_name)`` or
``(None, None)``. Ids are globally unique, so the first hit is authoritative."""
"""Scan every profile's ``state.db`` -> ``(db, profile_name)`` or ``(None, None)``.
Ids are globally unique, so the first hit is authoritative."""
from pathlib import Path
try:
from hermes_cli import profiles as profiles_mod
from hermes_state import SessionDB
except Exception:
return None, None
targets = [("default", profiles_mod.get_profile_dir("default"))]
targets += _quiet(lambda: [(info.name, info.path) for info in profiles_mod.list_profiles()],
[], "list_profiles failed during session locate")
targets = [("default", profiles_mod.get_profile_dir("default"))] + _quiet(
lambda: [(info.name, info.path) for info in profiles_mod.list_profiles()], [],
"list_profiles failed during session locate")
seen: set = set()
for name, home in targets:
db_path = Path(home) / "state.db"
@@ -382,20 +341,13 @@ def _locate_session_db(session_id: str):
continue
seen.add(str(db_path))
pdb = _quiet(lambda: SessionDB(db_path=db_path, read_only=True), None, "open %s failed", db_path)
if pdb and _quiet(lambda: pdb.get_session(session_id), None,
"get_session probe failed for %s in %s", session_id, name):
if pdb and _get_session_meta(pdb, session_id):
return pdb, name
if pdb:
pdb.close()
return None, None
def _get_session_meta(db, session_id: str) -> dict:
"""``db.get_session`` that degrades to ``{}`` on error."""
return _quiet(lambda: db.get_session(session_id), None,
"get_session failed for %s: %s", session_id, with_exc=True) or {}
def _read_session(db, session_id: str, head: int = 20, tail: int = 10, link_profile: str = None) -> str:
"""Read shape: whole session, or ``head`` + ``tail`` messages with a scroll pointer."""
meta = _get_session_meta(db, session_id)
@@ -406,13 +358,26 @@ def _read_session(db, session_id: str, head: int = 20, tail: int = 10, link_prof
if err:
return err
shaped = [_shape_message(m) for m in rows]
total = len(shaped)
truncated = total > head + tail
extra = {"message": (f"Session has {total} messages; showing first {head} + last {tail}. "
"Pass around_message_id (any id above) to scroll the middle.")} if truncated else {}
total, truncated = len(shaped), len(shaped) > head + tail
return _ok(mode="read", session_id=session_id, link=_session_link(session_id, link_profile),
session_meta=_session_meta_block(meta), message_count=total, truncated=truncated,
messages=shaped[:head] + shaped[-tail:] if truncated else shaped, **extra)
messages=shaped[:head] + shaped[-tail:] if truncated else shaped,
**({"message": (f"Session has {total} messages; showing first {head} + last {tail}. "
"Pass around_message_id (any id above) to scroll the middle.")} if truncated else {}))
def _read_with_profile_fallback(db, sid: str, profile: Optional[str]) -> str:
"""Read shape; on a miss scan every profile (the model may have dropped the owning
profile from the link) and tag the result with where it was found."""
result = _read_session(db, sid, link_profile=profile)
located, owner = (None, None) if json.loads(result).get("success") else _locate_session_db(sid)
if located is None:
return result
try:
found = json.loads(_read_session(located, sid, link_profile=owner))
finally:
located.close()
return json.dumps({**found, "profile": owner}, ensure_ascii=False) if found.get("success") else result
def _list_recent_sessions(db, limit: int, current_session_id: str = None, link_profile: str = None) -> str:
@@ -432,19 +397,14 @@ def _list_recent_sessions(db, limit: int, current_session_id: str = None, link_p
exclude_sources=list(_HIDDEN_SESSION_SOURCES), timeout_seconds=3.0)
current_root, has_compression_hop = (
_resolve_to_parent(db, current_session_id) if current_session_id else (None, False))
results = []
for s in sessions:
sid = s.get("id", "")
# Compression continuation: the root was summarised into the live child, so
# hide it. /new-reset children carry no transcript — keep that root browsable.
if sid == current_session_id or (has_compression_hop and current_root and sid == current_root):
continue
results.append({
"session_id": sid, "link": _session_link(sid, link_profile), "title": s.get("title") or None,
**{k: s.get(k, "") for k in ("source", "started_at", "last_active")},
"message_count": s.get("message_count", 0), "preview": s.get("preview", "")})
if len(results) >= limit:
break
# Compression continuation: the root was summarised into the live child, so hide
# it. /new-reset children carry no transcript — keep that root browsable.
hidden = {current_session_id, current_root if has_compression_hop and current_root else None}
results = [{
"session_id": s.get("id", ""), "link": _session_link(s.get("id", ""), link_profile),
"title": s.get("title") or None, **{k: s.get(k, "") for k in ("source", "started_at", "last_active")},
"message_count": s.get("message_count", 0), "preview": s.get("preview", "")}
for s in [x for x in sessions if x.get("id", "") not in hidden][:limit]]
return _ok(mode="browse", results=results, count=len(results), message=(
f"Showing {len(results)} most recent sessions. Pass a query= to search, "
"or session_id+around_message_id to scroll."))
@@ -461,41 +421,19 @@ def _clamp_int(value, default: int, lo: int, hi: int) -> int:
return max(lo, min(value, hi))
def _anchor_in_live_context(db, anchor_state, anchor_session_id: str, current_session_id: str) -> bool:
def _anchor_in_live_context(db, anchor_state, anchor_sid: str, current_session_id: str) -> bool:
"""True when the scroll anchor is still in the caller's active context (reject).
Same-lineage history that LEFT live context (compacted rows, compression-ended
parents, /new-reset predecessors) is allowed, so scroll never rejects a result
discovery just returned."""
if not _same_lineage(db, anchor_session_id, current_session_id) or _is_compacted_state(anchor_state):
parents, /new-reset predecessors) passes, so scroll never rejects a discovery result.
Rewind/undo rows (active=0, compacted!=1) never count as out-of-context history."""
if not _same_lineage(db, anchor_sid, current_session_id) or _is_compacted_state(anchor_state):
return False
# Rewind/undo rows (active=0, compacted!=1) never count as out-of-context history.
inactive_non_compacted = anchor_state is not None and anchor_state["active"] == 0 and anchor_state["compacted"] != 1
return inactive_non_compacted or not _session_left_live_context(db, anchor_session_id)
def _same_lineage(db, a: str, b: str) -> bool:
a_root, b_root = _resolve_lineage(db, a), _resolve_lineage(db, b)
return bool(a_root and b_root and a_root == b_root)
def _rebind_to_owner(db, session_id: str, owning: str, around_message_id: int, window: int):
"""Lineage rebind when the caller paired a parent session_id with a message id
living in a descendant. ``(view, warning)`` from the owner, or ``(None, None)``."""
rebind_view = _same_lineage(db, session_id, owning) and _quiet(
lambda: db.get_messages_around(owning, around_message_id, window=window),
None, "rebind get_messages_around failed: %s", with_exc=True)
if not (rebind_view and rebind_view.get("window")):
return None, None
return rebind_view, f"around_message_id {around_message_id} lives in {owning} (child of {session_id}); rebound transparently"
return (anchor_state is not None and anchor_state["active"] == 0) or not _session_left_live_context(db, anchor_sid)
def _scroll(db, session_id: str, around_message_id: int, window: int = 5,
current_session_id: str = None) -> str:
"""Scroll shape: a window centered on an anchor (no FTS5, no bookends);
rebinds silently if the anchor lives in a same-lineage child."""
if not isinstance(session_id, str) or not session_id.strip():
return tool_error("scroll requires session_id", success=False)
session_id = session_id.strip()
"""Scroll shape: a window centered on an anchor (no FTS5, no bookends)."""
try:
around_message_id = int(around_message_id)
except (TypeError, ValueError):
@@ -503,9 +441,8 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5,
window = _clamp_int(window, 5, 1, 20)
# Locate the anchor BEFORE the current-lineage guard (see _anchor_in_live_context).
anchor_state = _get_message_storage_state(db, around_message_id)
owning_session_id = anchor_state.get("session_id") if anchor_state is not None else None
if current_session_id and _anchor_in_live_context(
db, anchor_state, owning_session_id or session_id, current_session_id):
owning = (anchor_state or {}).get("session_id")
if current_session_id and _anchor_in_live_context(db, anchor_state, owning or session_id, current_session_id):
return tool_error("scroll rejected: anchor lives in the current session lineage (already in your active context)", success=False)
session_meta = _get_session_meta(db, session_id)
if not session_meta:
@@ -515,13 +452,18 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5,
if err:
return err
messages = view.get("window") or []
rebind_warning = None
if not messages and owning_session_id and owning_session_id != session_id:
rebind_view, rebind_warning = _rebind_to_owner(db, session_id, owning_session_id, around_message_id, window)
if rebind_view is not None:
view, messages = rebind_view, rebind_view["window"]
session_meta = _get_session_meta(db, owning_session_id) or session_meta
session_id = owning_session_id
extra = {}
if not messages and owning and owning != session_id:
# Lineage rebind: the caller paired a parent session_id with a message id
# living in a descendant — serve the owner's window transparently.
rebind_view = _same_lineage(db, session_id, owning) and _quiet(
lambda: db.get_messages_around(owning, around_message_id, window=window),
None, "rebind get_messages_around failed: %s", with_exc=True)
if rebind_view and rebind_view.get("window"):
extra["warning"] = (f"around_message_id {around_message_id} lives in {owning} "
f"(child of {session_id}); rebound transparently")
view, messages, session_id = rebind_view, rebind_view["window"], owning
session_meta = _get_session_meta(db, owning) or session_meta
if not messages:
return tool_error(f"around_message_id {around_message_id} not in session_id {session_id}", success=False)
return _ok(
@@ -532,27 +474,7 @@ def _scroll(db, session_id: str, around_message_id: int, window: int = 5,
hint=("Scroll forward: re-call with around_message_id = the LAST message's "
"id; backward: the FIRST message's id (the boundary message repeats "
"as an orientation marker). messages_before/messages_after < window "
"means you've hit that end of the session."),
**({"warning": rebind_warning} if rebind_warning else {}))
def _read_with_profile_fallback(db, sid: str, profile: Optional[str]) -> str:
"""Read shape; on a miss scan every profile (the model may have dropped the
owning profile from the link) and tag the result with where it was found."""
result = _read_session(db, sid, link_profile=profile)
if json.loads(result).get("success"):
return result
located, owner = _locate_session_db(sid)
if located is None:
return result
try:
found = json.loads(_read_session(located, sid, link_profile=owner))
finally:
located.close()
if not found.get("success"):
return result
found["profile"] = owner
return json.dumps(found, ensure_ascii=False)
"means you've hit that end of the session."), **extra)
def _dispatch(query, role_filter, limit, db, current_session_id, session_id,
@@ -578,39 +500,34 @@ def _dispatch(query, role_filter, limit, db, current_session_id, session_id,
owned_dbs.append(profile_db)
if isinstance(session_id, str) and session_id.strip():
if around_message_id is not None:
return _scroll(db, session_id, around_message_id, window, current_session_id)
return _scroll(db, session_id.strip(), around_message_id, window, current_session_id)
return _read_with_profile_fallback(db, session_id.strip(), profile)
limit = _clamp_int(limit, 3, 1, 10)
if not query or not isinstance(query, str) or not query.strip():
return _list_recent_sessions(db, limit, current_session_id, link_profile=profile)
role_list = ([r.strip() for r in role_filter.split(",") if r.strip()] or None) if isinstance(role_filter, str) else None
sort_norm = sort.strip().lower() if isinstance(sort, str) else None
sort_norm = sort_norm if sort_norm in ("newest", "oldest") else None
detail_norm = "full" if isinstance(detail, str) and detail.strip().lower() == "full" else "adaptive"
return _discover(
db=db, query=query.strip(), role_filter=role_list, limit=limit, sort=sort_norm,
detail=detail_norm, current_session_id=current_session_id, link_profile=profile)
db=db, query=query.strip(), limit=limit, sort=sort_norm if sort_norm in ("newest", "oldest") else None,
role_filter=([r.strip() for r in role_filter.split(",") if r.strip()] or None) if isinstance(role_filter, str) else None,
detail="full" if isinstance(detail, str) and detail.strip().lower() == "full" else "adaptive",
current_session_id=current_session_id, link_profile=profile)
def session_search(query: str = "", role_filter: str = None, limit: int = 3, db=None,
current_session_id: str = None, session_id: str = None, around_message_id: int = None,
window: int = 5, sort: str = None, profile: str = None, detail: str = "adaptive") -> str:
"""Run session search, closing DBs opened here. Positional order is frozen for old callers."""
from hermes_state import format_session_db_unavailable, get_shared_session_db, release_or_close
owned_dbs: List[Any] = []
if db is None:
try:
from hermes_state import get_shared_session_db
db = get_shared_session_db()
owned_dbs.append(db)
except Exception:
logging.debug("SessionDB unavailable for session_search", exc_info=True)
from hermes_state import format_session_db_unavailable
db = _quiet(get_shared_session_db, None, "SessionDB unavailable for session_search")
if db is None:
return tool_error(format_session_db_unavailable(), success=False)
owned_dbs.append(db)
try:
return _dispatch(query, role_filter, limit, db, current_session_id, session_id,
around_message_id, window, sort, profile, detail, owned_dbs)
finally:
from hermes_state import release_or_close
for owned_db in reversed(owned_dbs):
_quiet(lambda: release_or_close(owned_db), None, "Failed to close session_search SessionDB")
@@ -730,8 +647,7 @@ SESSION_SEARCH_SCHEMA = {
}
# --- Registry ---
from tools.registry import registry, tool_error
from tools.registry import registry, tool_error # noqa: E402 (registration at import time)
registry.register(
name="session_search",
+62 -111
View File
@@ -1,19 +1,18 @@
"""Conservative heredoc masking for shell-command scanners (terminal '&' guard, blocked-command
checks, cron lifecycle_guard) that false-positive on heredoc *bodies*. Stripping every body is
unsafe the other way (a fake ``<<`` in quotes can swallow a real operator; unquoted bodies
expand; ``bash <<'EOF'`` executes), so a body is masked ONLY when every delimiter is quoted,
every heredoc has an exact terminator line, the opener is a single command (no ``;|&``,
``$(...)``, backticks or process substitution) and the consumer is an allowlisted non-shell
interpreter. Otherwise the command is returned untouched: a false positive is acceptable,
hiding shell syntax from a guard is not. Masked bodies keep their newline count (re.MULTILINE)."""
"""Conservative heredoc masking for shell-command scanners ('&' guard, blocked-command checks,
cron lifecycle_guard) that false-positive on heredoc *bodies*. Stripping every body is unsafe the
other way (a fake ``<<`` in quotes can swallow an operator; unquoted bodies expand; ``bash <<'EOF'``
executes), so a body is masked ONLY when every delimiter is quoted, every heredoc has an exact
terminator line, the opener is a single command (no ``;|&``, ``$(...)``, backticks, process
substitution) and the consumer is an allowlisted non-shell interpreter. Otherwise the command is
returned untouched: a false positive is acceptable, hiding shell syntax from a guard is not.
Masked bodies keep their newline count (re.MULTILINE)."""
from __future__ import annotations
import re
# Non-shell interpreters whose quoted heredoc bodies are program text/data for
# THAT interpreter. Optional VAR=... assignments, ``env`` and a path prefix are
# allowed. Deliberately narrow: anything unmatched keeps its body visible.
# Non-shell interpreters whose quoted heredoc bodies are data for THAT interpreter; optional
# VAR=... assignments, ``env`` and a path prefix allowed. Narrow on purpose: unmatched = visible.
_INERT_HEREDOC_CONSUMER_RE = re.compile(
r"^\s*(?:[A-Z_][A-Z0-9_]*=\S+\s+)*(?:env\s+)?(?:[A-Za-z0-9_./-]+/)?"
r"(?:python(?:3(?:\.\d+)*)?|osascript|cat)(?=\s|$)",
@@ -24,12 +23,9 @@ def _span_end(command: str, cursor: int, closer: str) -> int:
"""Index just past the backslash-aware span opened at ``cursor``."""
end = cursor + 1
while end < len(command):
if command[end] == "\\" and end + 1 < len(command):
end += 2
continue
if command[end] == closer:
return end + 1
end += 1
end += 2 if command[end] == "\\" and end + 1 < len(command) else 1
return end
@@ -39,20 +35,15 @@ def _mask_simple_quotes(command: str) -> str:
cursor = 0
while cursor < len(command):
char = command[cursor]
if char == "'":
closing = command.find("'", cursor + 1)
if closing == -1:
result.append(command[cursor:])
break
result.append("''")
cursor = closing + 1
elif char == '"':
end = _span_end(command, cursor, '"')
if not command[cursor:end].endswith('"'):
result.append(command[cursor:])
break
if char in "'\"": # single quotes have no escapes; double quotes are backslash-aware
end = (command.find("'", cursor + 1) + 1 if char == "'"
else _span_end(command, cursor, '"'))
segment = command[cursor:end]
result.append(segment if "$(" in segment or "`" in segment else '""')
if not segment.endswith(char):
result.append(command[cursor:])
break
keep = char == '"' and ("$(" in segment or "`" in segment)
result.append(segment if keep else char * 2)
cursor = end
elif char == "`":
end = _span_end(command, cursor, "`")
@@ -68,118 +59,88 @@ def _parse_heredoc_operator(command: str, index: int):
"""Parse one ``<<`` opener -> ``(end_index, delimiter, strip_tabs, quoted)`` or None."""
if not command.startswith("<<", index) or command.startswith("<<<", index):
return None
cursor = index + 2
strip_tabs = cursor < len(command) and command[cursor] == "-"
if strip_tabs:
cursor += 1
strip_tabs = command.startswith("-", index + 2)
cursor = index + 3 if strip_tabs else index + 2
while cursor < len(command) and command[cursor] in " \t":
cursor += 1
if cursor >= len(command) or command[cursor] in "\r\n":
return None
delimiter: list[str] = []
quoted = False
while cursor < len(command):
while cursor < len(command) and not (command[cursor].isspace() or command[cursor] in ";&|<>()"):
char = command[cursor]
if char.isspace() or char in ";&|<>()":
break
if char == "\\":
if char == "\\": # backslash-escaped char: quoted, literal
if cursor + 1 >= len(command) or command[cursor + 1] in "\r\n":
return None
quoted = True
delimiter.append(command[cursor + 1])
cursor += 2
continue
if char in "'\"":
elif char in "'\"":
quoted = True
quote = char
cursor += 1
while cursor < len(command) and command[cursor] != quote:
while cursor < len(command) and command[cursor] != char:
current = command[cursor]
if current in "\r\n":
return None
if quote == '"' and current == "\\":
if char == '"' and current == "\\":
if cursor + 1 >= len(command):
return None
following = command[cursor + 1]
if following in {"$", "`", '"', "\\", "\n"}:
delimiter.append(following)
cursor += 2
continue
# In double quotes, backslash is literal before other chars.
if command[cursor + 1] in '$`"\\\n': # else backslash is literal in dquotes
cursor += 1
current = command[cursor]
delimiter.append(current)
cursor += 1
if cursor >= len(command):
if cursor >= len(command): # unterminated quote
return None
cursor += 1
continue
delimiter.append(char)
cursor += 1
else:
delimiter.append(char)
cursor += 1
if not delimiter and not quoted:
return None
return cursor, "".join(delimiter), strip_tabs, quoted
def _scan_heredoc_command_unit(command: str, start: int):
"""Scan one logical command -> ``(end, specs, unknown_operator, has_list_operator)``:
an unparseable ``<<`` (caller must fail closed) / unquoted ``;|&`` on the opener."""
"""Scan one logical command -> ``(end, specs, unknown_operator, has_list_operator)``: an
unparseable ``<<`` (caller must fail closed) / an unquoted ``;|&`` on the opener line."""
cursor = start
quote = None
comment = False
specs = []
unknown_operator = False
has_list_operator = False
while cursor < len(command):
char = command[cursor]
if comment:
if char == "\n":
return cursor, specs, unknown_operator, has_list_operator
cursor += 1
continue
if quote is not None:
if quote in {'"', "`"} and char == "\\" and cursor + 1 < len(command):
cursor += 2
continue
if char == "\n" and (comment or quote is None):
break
# Backslash escapes (incl. line continuations) outside single quotes skip the next char.
escaped = char == "\\" and quote != "'" and not comment and cursor + 1 < len(command)
if comment or quote is not None or escaped:
if char == quote:
quote = None
cursor += 1
continue
if char == "\\" and cursor + 1 < len(command):
# Includes line continuations: the logical command keeps going.
cursor += 2
continue
if char in "'\"`":
cursor += 2 if escaped else 1
elif char in "'\"`":
quote = char
cursor += 1
continue
if char == "#":
previous = command[cursor - 1] if cursor > start else ""
if cursor == start or previous.isspace() or previous in ";&|()":
comment = True
cursor += 1
continue
if char == "\n":
return cursor, specs, unknown_operator, has_list_operator
if command.startswith("<<<", cursor):
elif char == "#" and (cursor == start or command[cursor - 1].isspace()
or command[cursor - 1] in ";&|()"):
comment = True
cursor += 1
elif command.startswith("<<<", cursor):
cursor += 3
continue
if command.startswith("<<", cursor):
elif command.startswith("<<", cursor):
parsed = _parse_heredoc_operator(command, cursor)
if parsed is None:
unknown_operator = True
cursor += 2
continue
cursor, delimiter, strip_tabs, quoted = parsed
specs.append((delimiter, strip_tabs, quoted))
continue
if char in ";|&":
has_list_operator = True
cursor += 1
return len(command), specs, unknown_operator, has_list_operator
else:
cursor, delimiter, strip_tabs, quoted = parsed
specs.append((delimiter, strip_tabs, quoted))
else:
has_list_operator = has_list_operator or char in ";|&"
cursor += 1
return cursor, specs, unknown_operator, has_list_operator
def _find_heredoc_close(
@@ -200,14 +161,12 @@ def _find_heredoc_close(
def strip_inert_heredoc_bodies(command: str) -> str:
"""Mask heredoc bodies that are provably inert data (see module docstring)."""
# Runs on every terminal call: skip the state machine when no '<<' exists,
# and stop scanning once past the last '<<'.
# Runs on every terminal call: skip the state machine when no '<<' exists; stop past the last.
if "<<" not in command:
return command
last_opener_index = command.rfind("<<")
ranges: list[tuple[int, int]] = []
command_start = 0
while command_start <= last_opener_index:
command_end, specs, unknown_operator, has_list_operator = (
_scan_heredoc_command_unit(command, command_start))
@@ -219,9 +178,7 @@ def strip_inert_heredoc_bodies(command: str) -> str:
command_start = command_end + 1
continue
if command_end >= len(command):
# Opener with no body line: unterminated — leave visible.
return command
return command # opener with no body line: unterminated — leave visible
body_cursor = command_end + 1
body_ranges: list[tuple[int, int]] = []
for delimiter, strip_tabs, _quoted in specs:
@@ -230,22 +187,16 @@ def strip_inert_heredoc_bodies(command: str) -> str:
return command # unterminated
body_ranges.append((body_cursor, close_end))
body_cursor = close_end
if all(quoted for _delimiter, _strip_tabs, quoted in specs) and not has_list_operator:
masked_opener = _mask_simple_quotes(command[command_start:command_end])
nested_scope = any(m in masked_opener for m in ("$(", "`", "<(", ">("))
if not nested_scope and _INERT_HEREDOC_CONSUMER_RE.search(masked_opener):
if (not any(m in masked_opener for m in ("$(", "`", "<(", ">("))
and _INERT_HEREDOC_CONSUMER_RE.search(masked_opener)):
ranges.extend(body_ranges)
command_start = body_cursor
if not ranges:
return command
# Single-pass rebuild: ranges are sorted and non-overlapping.
# Single-pass rebuild (ranges are sorted and non-overlapping), bodies -> their newlines only.
parts: list[str] = []
previous = 0
for start, end in ranges:
parts.append(command[previous:start])
parts.append("\n" * command.count("\n", start, end))
parts += [command[previous:start], "\n" * command.count("\n", start, end)]
previous = end
parts.append(command[previous:])
return "".join(parts)
return "".join(parts) + command[previous:]
+76 -150
View File
@@ -3,10 +3,9 @@
Every skill mutation (any actor) appends one JSONL entry to
``~/.hermes/skills/.curator_ledger.jsonl`` with before/after file manifests whose
contents are stored content-addressed (sha256-deduped) under
``~/.hermes/.curator_backups/blobs/``. JSONL, not the state DB: durable,
human-greppable, survives DB resets. The ledger is TELEMETRY, NOT A GATE: every
public write path swallows and logs. The one exception is ``rollback_entry``,
which FAILS CLOSED when its own pre-rollback safety capture fails.
``~/.hermes/.curator_backups/blobs/``. JSONL, not the state DB: durable, greppable,
survives DB resets. TELEMETRY, NOT A GATE: every public write path swallows and
logs — except ``rollback_entry``, which FAILS CLOSED when its safety capture fails.
"""
from __future__ import annotations
@@ -28,14 +27,12 @@ from hermes_constants import get_hermes_home
logger = logging.getLogger(__name__)
# Snapshot-id shape used by agent.curator_backup (duplicated so the ledger can
# read the newest skills.tar.gz without importing the backup stack).
# Snapshot-id shape of agent.curator_backup (duplicated to avoid importing the backup stack).
_BACKUP_ID_RE = re.compile(r"^\d{4}-\d{2}-\d{2}T\d{2}-\d{2}-\d{2}Z(-\d{2})?$")
# ".archive/<name>-YYYYMMDDHHMMSS" collision suffix added by archive_skill.
_ARCHIVE_TS_SUFFIX_RE = re.compile(r"^(.+)-\d{14}$")
# Actions whose rollback must restore a COMPLETE package: consolidation may
# have re-homed support files out of the tree first, so a disk-only capture
# would make rollback restore a hollow skill.
# Rollback of these must restore a COMPLETE package: consolidation may have re-homed
# support files first, so a disk-only capture would restore a hollow skill.
_PACKAGE_RESTORE_ACTIONS = frozenset({"delete", "archive", "purge"})
_VALID_ACTORS = {"curator", "agent", "user"}
_NON_PACKAGE_TOPS = {".curator_backups", ".hub", ".archive"}
@@ -66,18 +63,18 @@ def derive_actor() -> str:
return "agent"
def _skills_dir() -> Path:
return get_hermes_home() / "skills"
def ledger_path() -> Path:
return get_hermes_home() / "skills" / ".curator_ledger.jsonl"
return _skills_dir() / ".curator_ledger.jsonl"
def blobs_dir() -> Path:
return get_hermes_home() / ".curator_backups" / "blobs"
def _skills_dir() -> Path:
return get_hermes_home() / "skills"
def ledger_enabled() -> bool:
"""Config gate ``skills.ledger`` (default True); lazy import keeps this importable without the CLI."""
try:
@@ -88,25 +85,17 @@ def ledger_enabled() -> bool:
return True
def _norm(path: Path | str) -> Path:
return Path(os.path.normpath(str(path)))
def _rel_posix(path: Path | str, root: Path) -> Optional[str]:
"""POSIX path of ``path`` relative to ``root`` (both normalized), or None when outside."""
try:
return _norm(path).relative_to(_norm(root)).as_posix()
except ValueError:
return Path(os.path.normpath(str(path))).relative_to(os.path.normpath(str(root))).as_posix()
except (ValueError, TypeError):
return None
def _is_within(root: Path, path: Path) -> bool:
"""True when *path* (normalized, no symlink resolution) sits under *root*."""
try:
root_r, path_r = _norm(root), _norm(path)
return path_r == root_r or root_r in path_r.parents
except Exception:
return False
return _rel_posix(path, root) is not None
def _store_blob(data: bytes) -> str:
@@ -125,34 +114,23 @@ def read_blob(sha256: str) -> Optional[bytes]:
"""Return blob content or None when missing/invalid."""
if not sha256 or not all(c in "0123456789abcdef" for c in sha256):
return None
try:
p = blobs_dir() / sha256
return p.read_bytes() if p.exists() else None
except OSError:
return None
with suppress(OSError):
return (blobs_dir() / sha256).read_bytes() if (blobs_dir() / sha256).exists() else None
return None
def snapshot_paths(root: Optional[Path], *, complete_package: bool = False) -> List[Dict[str, str]]:
"""{path, sha256} for every file under *root*, each stored as a blob.
Empty when root is None/missing. Raises on I/O failure — callers decide whether
that is fatal (rollback safety capture) or swallowed (telemetry).
``complete_package`` unions in the newest curator tarball's files (disk hashes win)."""
"""{path, sha256} for every file under *root*, each stored as a blob; [] when root is
None/missing. Raises on I/O failure — callers decide whether that is fatal (rollback safety
capture) or swallowed (telemetry). ``complete_package`` unions in the newest curator
tarball's files (disk hashes win)."""
if root is None:
return []
root = Path(root)
if root.is_file():
files = [root]
elif root.is_dir():
files = sorted(p for p in root.rglob("*") if p.is_file())
elif complete_package:
files = [] # gone from disk; the backup fill below may still recover it
else:
return []
root = Path(root) # gone from disk -> []; the complete_package fill may still recover it
files = ([root] if root.is_file()
else sorted(p for p in root.rglob("*") if p.is_file()) if root.is_dir() else [])
out = [{"path": str(f), "sha256": _store_blob(f.read_bytes())} for f in files]
if complete_package:
out = fill_snapshot_from_curator_backup(root, out)
return out
return fill_snapshot_from_curator_backup(root, out) if complete_package else out
def _package_rel(root: Path) -> Optional[str]:
@@ -170,8 +148,7 @@ def _strip_archive_timestamp(name: str) -> str:
def _skill_md_parents(items: Optional[List[Dict[str, str]]]) -> List[Path]:
paths = [Path(str(item.get("path", ""))) for item in items or []]
return [p.parent for p in paths if p.name == "SKILL.md"]
return [p.parent for p in (Path(str(i.get("path", ""))) for i in items or []) if p.name == "SKILL.md"]
def package_prefixes(
@@ -183,48 +160,34 @@ def package_prefixes(
candidates = [_package_rel(Path(root)) if root is not None else None]
candidates += [_package_rel(p) for p in _skill_md_parents(before)]
candidates += [skill, _strip_archive_timestamp(skill) if skill else None]
found: List[str] = []
for prefix in candidates:
prefix = (prefix or "").strip("/")
if prefix and prefix not in found:
found.append(prefix)
return found
def _latest_skills_tarball() -> Optional[Path]:
"""Newest ``skills.tar.gz`` under ``skills/.curator_backups/``."""
backups = _skills_dir() / ".curator_backups"
try:
children = list(backups.iterdir()) if backups.is_dir() else []
except OSError:
return None
candidates = [
child / "skills.tar.gz" for child in children
if child.is_dir() and _BACKUP_ID_RE.match(child.name) and (child / "skills.tar.gz").is_file()]
# Parent dirs sort lexicographically == chronologically for the id shape.
return max(candidates, key=lambda p: p.parent.name) if candidates else None
return list(dict.fromkeys(p for p in ((c or "").strip("/") for c in candidates) if p))
def _read_package_files_from_latest_backup(prefixes: List[str]) -> Dict[str, bytes]:
"""``{posix-relpath: bytes}`` under *prefixes* in the newest snapshot; malicious
member names (absolute, ``..`` traversal) are rejected."""
if not prefixes:
"""``{posix-relpath: bytes}`` under *prefixes* in the newest ``skills/.curator_backups/*/
skills.tar.gz``; malicious member names (absolute, ``..`` traversal) are rejected."""
backups = _skills_dir() / ".curator_backups"
try:
children = list(backups.iterdir()) if prefixes and backups.is_dir() else []
except OSError:
return {}
archive = _latest_skills_tarball()
if archive is None:
candidates = [
child / "skills.tar.gz" for child in children
if child.is_dir() and _BACKUP_ID_RE.match(child.name) and (child / "skills.tar.gz").is_file()]
if not candidates:
return {}
# Parent dirs sort lexicographically == chronologically for the id shape.
archive = max(candidates, key=lambda p: p.parent.name)
prefixed = tuple(p if p.endswith("/") else p + "/" for p in prefixes)
exact = set(prefixes)
out: Dict[str, bytes] = {}
try:
with tarfile.open(archive, "r:gz") as tf:
for member in tf.getmembers():
if not member.isfile():
continue
name = member.name.replace("\\", "/").lstrip("./")
if not name or name.startswith("/") or ".." in Path(name).parts:
continue
if name not in exact and not name.startswith(prefixed):
if (not member.isfile() or not name or name.startswith("/")
or ".." in Path(name).parts
or (name not in exact and not name.startswith(prefixed))):
continue
extracted = tf.extractfile(member)
if extracted is not None:
@@ -238,14 +201,12 @@ def _read_package_files_from_latest_backup(prefixes: List[str]) -> Dict[str, byt
def fill_snapshot_from_curator_backup(
root: Optional[Path], existing: Optional[List[Dict[str, str]]] = None, *,
skill: Optional[str] = None) -> List[Dict[str, str]]:
"""Union missing skill-package files from the newest curator snapshot.
Completeness fill, not a gate: failures return *existing* unchanged, and only
ABSENT paths are filled. Fill targets go where rollback must restore them:
under *root* when known (for purge that is ``.archive/<name>/``, NOT the live
tree), else the live skills dir; the tar's leading package-dir segment is
stripped when *root* already names the package. Every target must stay under
``skills/`` and HERMES_HOME."""
"""Union missing skill-package files from the newest curator snapshot. Completeness fill, not
a gate: failures return *existing* unchanged, and only ABSENT paths are filled. Fill targets go
where rollback must restore them: under *root* when known (for purge that is
``.archive/<name>/``, NOT the live tree), else the live skills dir; the tar's leading
package-dir segment is stripped when *root* already names the package. Every target must stay
under ``skills/`` and HERMES_HOME."""
out = list(existing or [])
prefixes = package_prefixes(root, skill, out)
if not prefixes:
@@ -260,9 +221,7 @@ def fill_snapshot_from_curator_backup(
skills = _skills_dir()
dest_root = Path(root) if root is not None else None
pkg_names = {dest_root.name, _strip_archive_timestamp(dest_root.name)} if dest_root else set()
have = {
rel for rel in (_rel_posix(str(item.get("path", "")), skills) for item in out) if rel is not None
}
have = {rel for rel in (_rel_posix(str(i.get("path", "")), skills) for i in out) if rel is not None}
for rel, data in extra.items():
parts = rel.split("/")
if dest_root is not None and parts and parts[0] in pkg_names:
@@ -293,14 +252,10 @@ def append_entry(
return None
try:
entry = {
"id": uuid.uuid4().hex[:12],
"ts": datetime.now(timezone.utc).isoformat(),
"id": uuid.uuid4().hex[:12], "ts": datetime.now(timezone.utc).isoformat(),
"actor": actor if actor in _VALID_ACTORS else derive_actor(),
"action": action,
"skill": skill,
"evidence": evidence or {},
"before": before or [],
"after": after or []}
"action": action, "skill": skill, "evidence": evidence or {},
"before": before or [], "after": after or []}
path = ledger_path()
path.parent.mkdir(parents=True, exist_ok=True)
with open(path, "a", encoding="utf-8") as fh:
@@ -315,10 +270,9 @@ def record_mutation(
action: str, skill: str, before_root: Optional[Path] = None,
before: Optional[List[Dict[str, str]]] = None, after_root: Optional[Path] = None,
actor: Optional[str] = None, evidence: Optional[Dict[str, Any]] = None) -> Optional[str]:
"""Mutation hook: after-state from *after_root* (before = pre-captured list or
captured from *before_root*), then append. NEVER raises. delete/archive/purge
capture a COMPLETE package (filled from the newest curator backup) so
rollback never restores a shell."""
"""Mutation hook: after-state from *after_root* (before = pre-captured list or captured from
*before_root*), then append. NEVER raises. delete/archive/purge capture a COMPLETE package
(filled from the newest curator backup) so rollback never restores a shell."""
if not ledger_enabled():
return None
try:
@@ -327,9 +281,8 @@ def record_mutation(
before = snapshot_paths(before_root, complete_package=_complete)
elif _complete:
before = fill_snapshot_from_curator_backup(before_root, before, skill=skill)
after = snapshot_paths(after_root)
return append_entry(
action, skill, before=before, after=after, actor=actor, evidence=evidence)
return append_entry(action, skill, before=before, after=snapshot_paths(after_root),
actor=actor, evidence=evidence)
except Exception as e:
logger.warning("skill_ledger: record_mutation failed (%s) — mutation unaffected", e)
return None
@@ -344,9 +297,7 @@ def capture_before(
return None
try:
captured = snapshot_paths(root)
if complete_package:
captured = fill_snapshot_from_curator_backup(root, captured, skill=skill)
return captured
return fill_snapshot_from_curator_backup(root, captured, skill=skill) if complete_package else captured
except Exception as e:
logger.warning("skill_ledger: before-capture failed (%s) — mutation unaffected", e)
return None
@@ -354,33 +305,22 @@ def capture_before(
def list_entries(skill: Optional[str] = None, limit: Optional[int] = None) -> List[Dict[str, Any]]:
"""Read the ledger, newest first. Malformed lines are skipped."""
path = ledger_path()
if not path.exists():
try:
lines = ledger_path().read_text(encoding="utf-8").splitlines()
except OSError: # missing or unreadable ledger == empty
return []
rows: List[Dict[str, Any]] = []
try:
with open(path, "r", encoding="utf-8") as fh:
for line in fh:
try:
row = json.loads(line) if line.strip() else None
except json.JSONDecodeError:
continue
if isinstance(row, dict):
rows.append(row)
except OSError:
return []
if skill:
rows = [r for r in rows if r.get("skill") == skill]
for line in lines:
with suppress(json.JSONDecodeError):
row = json.loads(line) if line.strip() else None
if isinstance(row, dict) and (not skill or row.get("skill") == skill):
rows.append(row)
rows.reverse()
if limit is not None and limit >= 0:
rows = rows[:limit]
return rows
return rows[:limit] if limit is not None and limit >= 0 else rows
def get_entry(entry_id: str) -> Optional[Dict[str, Any]]:
if not entry_id:
return None
return next((row for row in list_entries() if row.get("id") == entry_id), None)
return next((r for r in list_entries() if r.get("id") == entry_id), None) if entry_id else None
def _validate_entry_paths(entry: Dict[str, Any]) -> Optional[str]:
@@ -397,20 +337,15 @@ def _validate_entry_paths(entry: Dict[str, Any]) -> Optional[str]:
def rollback_entry(entry_id: str) -> Tuple[bool, str]:
"""Restore the before-state of mutation *entry_id*. Fail-closed (mirrors
agent/curator_backup.rollback): every before-blob must exist BEFORE any
change, and a pre-rollback safety entry of every touched path's CURRENT
state is appended first — if that fails, nothing is changed."""
agent/curator_backup.rollback): every before-blob must exist BEFORE any change, and a
pre-rollback safety entry of every touched path's CURRENT state is appended first."""
entry = get_entry(entry_id)
if entry is None:
return False, f"no ledger entry with id '{entry_id}'"
path_err = _validate_entry_paths(entry)
if path_err:
if path_err := _validate_entry_paths(entry):
return False, f"refusing rollback: {path_err}"
before = list(entry.get("before") or [])
after = list(entry.get("after") or [])
# Historical hollow delete/archive/purge entries (SKILL.md only): fill from the
# newest curator backup so the complete package is restored. Entry hashes win;
# only missing paths are added, and the filled set is re-validated.
@@ -418,24 +353,18 @@ def rollback_entry(entry_id: str) -> Tuple[bool, str]:
before = fill_snapshot_from_curator_backup(
next(iter(_skill_md_parents(before)), None), before,
skill=str(entry.get("skill") or "") or None)
path_err = _validate_entry_paths({**entry, "before": before, "after": after})
if path_err:
if path_err := _validate_entry_paths({**entry, "before": before, "after": after}):
return False, f"refusing rollback: {path_err}"
# Pre-check every blob we need so we never fail mid-restore.
for item in before:
if read_blob(str(item.get("sha256", ""))) is None:
return False, (f"missing blob {item.get('sha256')} for {item.get('path')}; "
"rollback aborted, nothing was changed")
# Safety entry: CURRENT state of every touched path, so the rollback itself is undoable.
touched = {str(i["path"]) for i in before + after if i.get("path")}
try:
safety_before: List[Dict[str, str]] = []
for p in sorted(touched):
fp = Path(p)
if fp.is_file():
safety_before.append({"path": p, "sha256": _store_blob(fp.read_bytes())})
safety_before = [{"path": p, "sha256": _store_blob(Path(p).read_bytes())}
for p in sorted(touched) if Path(p).is_file()]
safety_id = append_entry(
"pre-rollback", entry.get("skill", "?"), before=safety_before, after=safety_before,
evidence={"rollback_target": entry_id})
@@ -445,7 +374,6 @@ def rollback_entry(entry_id: str) -> Tuple[bool, str]:
if safety_id is None:
return False, ("pre-rollback safety capture failed (ledger disabled or "
"unwritable); rollback aborted and current skills were not changed")
# Restore: write every before-file, remove files the mutation created.
before_paths = {str(i["path"]) for i in before}
for item in before:
@@ -456,14 +384,12 @@ def rollback_entry(entry_id: str) -> Tuple[bool, str]:
for item in after:
p = str(item.get("path", ""))
if p and p not in before_paths:
fp = Path(p)
try:
if fp.is_file():
fp.unlink()
if Path(p).is_file():
Path(p).unlink()
removed += 1
except OSError as e:
logger.warning("skill_ledger: could not remove %s during rollback: %s", p, e)
append_entry(
"rollback", entry.get("skill", "?"), before=safety_before, after=before,
evidence={"rollback_target": entry_id, "restored": restored, "removed": removed})
+50 -73
View File
@@ -18,50 +18,41 @@ _BATCH_MAX_OPS = 20
def _validate_batch_ops(operations, default_name, tool_error):
"""Shape checks with no side effects. Returns (names, None) or (None, error_json)."""
from tools.skill_manager_guards import _background_review_preflight
def fail(i, msg):
return None, tool_error(f"operations[{i}]{msg}", success=False)
names = []
for i, op in enumerate(operations):
if not isinstance(op, dict) or not op.get("action"):
return None, tool_error(f"operations[{i}] needs an 'action'.", success=False)
return fail(i, " needs an 'action'.")
act = op["action"]
if act not in _BATCH_OP_ACTIONS:
return None, tool_error(
f"operations[{i}]: unknown action '{act}'. Batchable: "
f"{', '.join(sorted(_BATCH_OP_ACTIONS))}; delete must be sole.",
success=False)
return fail(i, f": unknown action '{act}'. Batchable: "
f"{', '.join(sorted(_BATCH_OP_ACTIONS))}; delete must be sole.")
nm = op.get("name") or default_name
if not nm:
return None, tool_error(f"operations[{i}] needs a 'name' (the skill it targets).", success=False)
return fail(i, " needs a 'name' (the skill it targets).")
names.append(nm)
if act == "create" and nm in names[:-1]:
return None, tool_error(
f"operations[{i}]: create for '{nm}' must precede that skill's other ops.",
success=False)
preflight = _background_review_preflight(act, nm)
if preflight is not None:
return fail(i, f": create for '{nm}' must precede that skill's other ops.")
if (preflight := _background_review_preflight(act, nm)) is not None:
return None, json.dumps(preflight, ensure_ascii=False)
# Clobber guard: a DESTRUCTIVE op (create/write_file/remove_file/full rewrite) on
# a file an earlier op touched would SILENTLY discard its work — reject it.
# Additive patches are always legal. Paths are normalized against spelling variants.
touched_files = set()
for i, op in enumerate(operations):
act = op["action"]
nm = names[i]
act, nm = op["action"], names[i]
# create and full-rewrite patch (content) always hit SKILL.md.
full_rewrite = act == "patch" and bool(op.get("content"))
fp = (op.get("file_path") or "").strip()
target = ("SKILL.md" if (act == "create" or full_rewrite or not fp)
else posixpath.normpath(fp.lstrip("/")))
key = (nm, target)
destructive = act in ("create", "write_file", "remove_file") or full_rewrite
if destructive and key in touched_files:
return None, tool_error(
f"operations[{i}]: {act} on '{target}' of skill '{nm}' — an earlier op in this "
f"batch already touched that file, and this op would silently discard its work. "
f"One destructive op (write_file/remove_file/full rewrite) per file per batch; put "
f"it first, or fold the change in. Patch chains are fine.",
success=False)
if (act in ("create", "write_file", "remove_file") or full_rewrite) and key in touched_files:
return fail(i, f": {act} on '{target}' of skill '{nm}' — an earlier op in this "
f"batch already touched that file, and this op would silently discard its work. "
f"One destructive op (write_file/remove_file/full rewrite) per file per batch; put "
f"it first, or fold the change in. Patch chains are fine.")
touched_files.add(key)
return names, None
@@ -72,9 +63,8 @@ def _snapshot_skills(names, snap_root, find_skill):
for nm in dict.fromkeys(names): # ordered unique
pre = find_skill(nm)
pre_dir = Path(pre["path"]) if pre else None
snap = None
if pre_dir is not None and pre_dir.is_dir():
snap = snap_root / nm
snap = snap_root / nm if pre_dir is not None and pre_dir.is_dir() else None
if snap is not None:
try:
shutil.copytree(pre_dir, snap)
except Exception as exc: # noqa: BLE001 — no snapshot, no atomicity
@@ -84,26 +74,27 @@ def _snapshot_skills(names, snap_root, find_skill):
def _restore_snapshot(pre_dir, snap, post_dir) -> None:
if snap is not None:
if post_dir is not None and post_dir.is_dir():
# Move the broken state aside and delete it only after the snapshot is
# back, so a failed copytree (disk full, locked file) can't mean total loss.
aside = post_dir.with_name(post_dir.name + ".rollback-broken")
shutil.rmtree(aside, ignore_errors=True)
post_dir.rename(aside)
try:
shutil.copytree(snap, pre_dir)
except Exception:
# Restore failed: put the half-applied state back rather than nothing.
shutil.rmtree(pre_dir, ignore_errors=True)
aside.rename(pre_dir)
raise
shutil.rmtree(aside, ignore_errors=True)
else:
shutil.copytree(snap, pre_dir)
elif post_dir is not None and post_dir.is_dir():
# Batch created this skill: remove the partial result.
shutil.rmtree(post_dir)
post_exists = post_dir is not None and post_dir.is_dir()
if snap is None:
if post_exists: # Batch created this skill: remove the partial result.
shutil.rmtree(post_dir)
return
if not post_exists:
shutil.copytree(snap, pre_dir)
return
# Move the broken state aside and delete it only after the snapshot is
# back, so a failed copytree (disk full, locked file) can't mean total loss.
aside = post_dir.with_name(post_dir.name + ".rollback-broken")
shutil.rmtree(aside, ignore_errors=True)
post_dir.rename(aside)
try:
shutil.copytree(snap, pre_dir)
except Exception:
# Restore failed: put the half-applied state back rather than nothing.
shutil.rmtree(pre_dir, ignore_errors=True)
aside.rename(pre_dir)
raise
shutil.rmtree(aside, ignore_errors=True)
def _rollback(snapshots, find_skill):
@@ -114,9 +105,8 @@ def _rollback(snapshots, find_skill):
post = find_skill(nm)
_restore_snapshot(pre_dir, snap, Path(post["path"]) if post else None)
except Exception as exc: # noqa: BLE001
notes.append(
f"ROLLBACK FAILED for '{nm}' ({exc}); snapshot preserved at '{snap}'"
if snap is not None else f"ROLLBACK FAILED for '{nm}' ({exc})")
notes.append(f"ROLLBACK FAILED for '{nm}' ({exc})"
+ (f"; snapshot preserved at '{snap}'" if snap is not None else ""))
return ("; ".join(notes) if notes else "all touched skills rolled back"), bool(notes)
@@ -129,67 +119,55 @@ def _skill_manage_batch(operations, default_name: str = None, task_id: str = Non
top-level ``name`` fallback (staged replay)."""
from tools import skill_manager_tool as _smt
from tools.registry import tool_error
if not isinstance(operations, list) or not operations:
return tool_error("operations must be a non-empty array.", success=False)
if len(operations) > _BATCH_MAX_OPS:
return tool_error(f"operations is capped at {_BATCH_MAX_OPS} ops per call.", success=False)
if any(isinstance(op, dict) and op.get("action") == "delete" for op in operations):
if len(operations) != 1:
return tool_error(
"delete must be the SOLE op in its call — it doesn't "
"compose with other ops' rollback.",
success=False)
op = operations[0]
nm = op.get("name") or default_name
return tool_error("delete must be the SOLE op in its call — it doesn't "
"compose with other ops' rollback.", success=False)
nm = operations[0].get("name") or default_name
if not nm:
return tool_error("operations[0] (delete) needs a 'name'.", success=False)
return _smt.skill_manage(
action="delete", name=nm, absorbed_into=op.get("absorbed_into"),
task_id=task_id, session_id=session_id)
return _smt.skill_manage(action="delete", name=nm, task_id=task_id, session_id=session_id,
absorbed_into=operations[0].get("absorbed_into"))
names, err = _validate_batch_ops(operations, default_name, tool_error)
if err is not None:
return err
if not _smt._skill_gate_bypass.get():
# Approval gate for the WHOLE batch as one pending write.
def _staging(wa):
acts = ", ".join(op["action"] for op in operations)
gist = f"batch({len(operations)} ops: {acts}) on {', '.join(sorted(set(names)))}"
return {"action": "batch", "operations": operations}, gist
staged = _smt._run_write_gate(_staging)
if staged is not None:
return staged
snap_root = Path(tempfile.mkdtemp(prefix="skill_batch_"))
snapshots, snap_err = _snapshot_skills(names, snap_root, _smt._find_skill)
if snap_err is not None:
shutil.rmtree(snap_root, ignore_errors=True)
return tool_error(snap_err, success=False)
# Single-op path with the gate bypassed (the batch already cleared/staged it).
results = []
rollback_failed = False
token = _smt._skill_gate_bypass.set(True)
try:
for i, op in enumerate(operations):
raw = _smt._skill_manage_from(
{**op, "name": names[i]}, task_id=task_id, session_id=session_id)
raw = _smt._skill_manage_from({**op, "name": names[i], "operations": None},
task_id=task_id, session_id=session_id)
try:
parsed = json.loads(raw)
except Exception: # noqa: BLE001
parsed = {"success": False, "error": "unparseable op result"}
if not parsed.get("success"):
note, rollback_failed = _rollback(snapshots, _smt._find_skill)
fail = {
fail = { # key order is wire-visible
"success": False,
"error": (
f"operations[{i}] ({op['action']} on '{names[i]}') failed: "
f"{parsed.get('error', 'unknown error')} — batch aborted, {note}."),
"failed_index": i,
"completed_before_failure": i}
"error": (f"operations[{i}] ({op['action']} on '{names[i]}') failed: "
f"{parsed.get('error', 'unknown error')} — batch aborted, {note}."),
"failed_index": i, "completed_before_failure": i}
# Carry the failing op's teaching payload (patch's file_preview /
# fuzzy-match hints) through — without it the model recovers blind.
for k, v in parsed.items():
@@ -205,7 +183,6 @@ def _skill_manage_batch(operations, default_name: str = None, task_id: str = Non
logger.warning("skill_manage batch rollback failed, snapshots kept at %s", snap_root)
else:
shutil.rmtree(snap_root, ignore_errors=True)
return json.dumps(
{"success": True, "operations_applied": len(results), "results": results},
ensure_ascii=False)
+85 -114
View File
@@ -1,13 +1,12 @@
"""Write/delete guards for ``skill_manage``.
Every guard returns ``None`` when the operation may proceed, otherwise a refusal
(error dict or message). Origin-owned state (``_find_skill``, ``_skills_dir``) is
reached lazily through ``tools.skill_manager_tool`` so test patches keep working.
"""
"""Write/delete guards for ``skill_manage``. Every guard returns ``None`` when the
operation may proceed, else a refusal (error dict or message). Origin-owned state
(``_find_skill``, ``_skills_dir``) is reached lazily via ``tools.skill_manager_tool``
so test patches keep working."""
import contextvars as _ctxvars
import logging
import threading
from contextlib import suppress
from pathlib import Path
from typing import Any, Dict, Optional
@@ -28,10 +27,9 @@ def _is_background_review() -> bool:
def _resolved_str(path: Path) -> str:
try:
with suppress(Exception):
return str(path.resolve())
except Exception:
return str(path)
return str(path)
class _BackgroundReviewReadMarks:
@@ -55,14 +53,12 @@ _background_review_read_paths: "_ctxvars.ContextVar[Optional[_BackgroundReviewRe
def mark_background_review_skill_read(path: Path) -> None:
"""Record that the active background-review fork has read a skill file.
The fork must not patch content it only inferred from the transcript:
skill_view/read_file call this, and the write guards require the mark."""
"""Record that the active background-review fork has read a skill file. The fork must not
patch content it only inferred from the transcript: skill_view/read_file call this, and
the write guards require the mark."""
if not _is_background_review():
return
marks = _background_review_read_paths.get()
if marks is None:
if (marks := _background_review_read_paths.get()) is None:
_background_review_read_paths.set(marks := _BackgroundReviewReadMarks())
marks.add(_resolved_str(path))
@@ -77,27 +73,29 @@ def _reset_background_review_read_marks() -> None:
_background_review_read_paths.set(_BackgroundReviewReadMarks())
def _containing_skills_root(skill_path: Path) -> Path:
"""Skills root (local or external_dirs) containing ``skill_path``; local dir if none match."""
def _resolved_roots(skill_path: Path):
"""``(resolved skill_path, [(root, resolved_root), ...])`` over every resolvable skills root."""
from agent.skill_utils import get_all_skills_dirs
from tools import skill_manager_tool as _smt
try:
resolved = skill_path.resolve()
except OSError:
resolved = skill_path
roots = []
for root in get_all_skills_dirs():
try:
if resolved.is_relative_to(root.resolve()):
return root
except OSError:
continue
return _smt._skills_dir()
with suppress(OSError):
roots.append((root, root.resolve()))
return resolved, roots
def _containing_skills_root(skill_path: Path) -> Path:
"""Skills root (local or external_dirs) containing ``skill_path``; local dir if none match."""
from tools import skill_manager_tool as _smt
resolved, roots = _resolved_roots(skill_path)
return next((root for root, r in roots if resolved.is_relative_to(r)), _smt._skills_dir())
def _is_path_redirect(path: Path) -> bool:
"""Symlink or (Windows 3.12+) junction — either lets a poisoned tree redirect
``shutil.rmtree`` outside the skills root."""
"""Symlink or (Windows 3.12+) junction — either lets a poisoned tree redirect rmtree outside."""
try:
return path.is_symlink() or (hasattr(path, "is_junction") and path.is_junction())
except OSError:
@@ -105,42 +103,39 @@ def _is_path_redirect(path: Path) -> bool:
def _validate_delete_target(skill_dir: Path) -> Optional[str]:
"""Last-line guard before ``shutil.rmtree(skill_dir)``: even a poisoned tree
must never delete (1) a path outside every known skills root, (2) a skills
root itself, or (3) a symlink/junction (rmtree would follow it)."""
from agent.skill_utils import get_all_skills_dirs
"""Last-line guard before rmtree: even a poisoned tree must never delete (1) a path outside
every known skills root, (2) a skills root itself, (3) a symlink/junction (rmtree follows it)."""
if _is_path_redirect(skill_dir):
return (
f"Refusing to delete '{skill_dir}': the skill directory is a "
f"symlink/junction. Remove the link target manually if intended.")
return (f"Refusing to delete '{skill_dir}': the skill directory is a "
f"symlink/junction. Remove the link target manually if intended.")
try:
resolved = skill_dir.resolve()
skill_dir.resolve()
except OSError as exc:
return f"Refusing to delete '{skill_dir}': could not resolve path ({exc})."
for root in get_all_skills_dirs():
try:
root = root.resolve()
except OSError:
continue
resolved, roots = _resolved_roots(skill_dir)
for _root, root in roots:
if resolved == root:
return (
f"Refusing to delete '{skill_dir}': resolves to the skills root "
f"itself, which would remove every installed skill.")
return (f"Refusing to delete '{skill_dir}': resolves to the skills root "
f"itself, which would remove every installed skill.")
if resolved.is_relative_to(root):
return None
return (
f"Refusing to delete '{skill_dir}': path does not resolve inside any "
f"known skills root.")
return f"Refusing to delete '{skill_dir}': path does not resolve inside any known skills root."
def _is_pinned(name: str, what: str) -> Optional[bool]:
"""skill_usage pinned flag; None (logged at debug) when the record is unreadable."""
try:
from tools import skill_usage
return bool(skill_usage.get_record(name).get("pinned"))
except Exception:
logger.debug("%s lookup failed for %s", what, name, exc_info=True)
return None
def _pinned_guard(name: str) -> Optional[str]:
"""Refusal message if *name* is pinned or essential, else None.
Pin only guards **deletion**; patches/edits stay allowed. ESSENTIAL_SKILLS are
permanently pinned (the system prompt references them). Best-effort: an
unreadable sidecar lets the delete through."""
"""Refusal message if *name* is pinned or essential, else None. Pin only guards DELETION;
patches/edits stay allowed. ESSENTIAL_SKILLS are permanently pinned (the system prompt
references them). Best-effort: an unreadable sidecar lets the delete through."""
try:
from agent.skill_utils import ESSENTIAL_SKILLS
if name in ESSENTIAL_SKILLS:
@@ -150,47 +145,34 @@ def _pinned_guard(name: str) -> Optional[str]:
f"cannot be deleted. Patches and edits are still allowed.")
except Exception:
logger.debug("essential-guard lookup failed for %s", name, exc_info=True)
try:
from tools import skill_usage
if skill_usage.get_record(name).get("pinned"):
return (
f"Skill '{name}' is pinned and cannot be deleted by skill_manage. Ask the user to "
f"run `hermes curator unpin {name}` if they want to delete it. Patches and edits "
f"are allowed on pinned skills; only deletion is blocked.")
except Exception:
logger.debug("pinned-guard lookup failed for %s", name, exc_info=True)
if _is_pinned(name, "pinned-guard"):
return (
f"Skill '{name}' is pinned and cannot be deleted by skill_manage. Ask the user to "
f"run `hermes curator unpin {name}` if they want to delete it. Patches and edits "
f"are allowed on pinned skills; only deletion is blocked.")
return None
def _background_review_write_guard(
name: str, skill_dir: Path, action: str) -> Optional[Dict[str, Any]]:
"""Refuse autonomous curator writes to anything but curator-owned sediment.
The background review fork has no user in the loop, so unlike foreground
agents it is also blocked on pinned/external/bundled/hub skills."""
"""Refuse autonomous curator writes to anything but curator-owned sediment. The review fork
has no user in the loop, so it is also blocked on pinned/external/bundled/hub skills."""
if not _is_background_review():
return None
try:
from tools import skill_usage
if skill_usage.get_record(name).get("pinned"):
return _refusal(
f"Refusing background curator {action} for pinned skill '{name}': pinned skills "
f"are off-limits to autonomous maintenance. Ask the user to run `hermes curator "
f"unpin {name}` if they want it changed.")
except Exception:
logger.debug("pinned skill guard lookup failed for %s", name, exc_info=True)
refuse = f"Refusing background curator {action} for"
if _is_pinned(name, "pinned skill guard"):
return _refusal(
f"{refuse} pinned skill '{name}': pinned skills "
f"are off-limits to autonomous maintenance. Ask the user to run `hermes curator "
f"unpin {name}` if they want it changed.")
try:
from agent.skill_utils import is_external_skill_path
if is_external_skill_path(skill_dir):
return _refusal(
f"Refusing background curator {action} for skill '{name}': "
"the skill lives in skills.external_dirs, which are "
"externally owned and read-only to autonomous curation.")
f"{refuse} skill '{name}': the skill lives in skills.external_dirs, which are "
f"externally owned and read-only to autonomous curation.")
except Exception:
logger.debug("external skill guard lookup failed for %s", name, exc_info=True)
try:
from tools import skill_usage
for predicate, label in (
@@ -198,8 +180,7 @@ def _background_review_write_guard(
(skill_usage.is_hub_installed, "hub-installed"),
(skill_usage.is_bundled, "bundled")):
if predicate(name):
return _refusal(
f"Refusing background curator {action} for {label} skill '{name}'.")
return _refusal(f"{refuse} {label} skill '{name}'.")
# Not curator-managed (no `created_by: "agent"`) => user-owned. A MISSING
# record and an explicit `created_by: null` must resolve IDENTICALLY (keying
# on presence made the policy depend on the guard's own side effect: the
@@ -209,13 +190,13 @@ def _background_review_write_guard(
_detail = (f"created_by={usage_rec.get('created_by')!r}" if isinstance(usage_rec, dict)
else "no usage record")
return _refusal(
f"Refusing background curator {action} for skill '{name}': the skill is not "
f"{refuse} skill '{name}': the skill is not "
f"curator-managed ({_detail}). User-owned skills are off-limits to autonomous "
f"curation. Run `hermes curator adopt {name}` to opt it in.")
except Exception:
logger.warning("owned skill guard lookup failed for %s", name, exc_info=True)
return _refusal(
f"Refusing background curator {action} for skill '{name}': agent ownership could not "
f"{refuse} skill '{name}': agent ownership could not "
f"be verified because the provenance record is unavailable or unreadable.")
return None
@@ -243,37 +224,33 @@ def _background_review_preflight(action: str, name: str) -> Optional[Dict[str, A
def _curator_consolidation_delete_guard(
name: str, absorbed_into: Optional[str]) -> Optional[Dict[str, Any]]:
"""Fail closed on unverified deletes during the curator consolidation pass.
The review fork's only legitimate delete is a verified consolidation declared
via ``absorbed_into=<umbrella>`` (existence validated in ``_delete_skill``).
The deterministic inactivity prune never calls ``skill_manage``, so a bare
delete here can only be the LLM pass pruning without evidence: refuse it."""
if not _is_background_review():
return None
if isinstance(absorbed_into, str) and absorbed_into.strip():
"""Fail closed on unverified deletes during the curator consolidation pass. The fork's only
legitimate delete is a consolidation declared via ``absorbed_into=<umbrella>`` (existence
validated in ``_delete_skill``); the deterministic inactivity prune never calls skill_manage,
so a bare delete here can only be the LLM pass pruning without evidence."""
if not _is_background_review() or (isinstance(absorbed_into, str) and absorbed_into.strip()):
return None
return _refusal(
f"Refusing background curator delete of skill '{name}': the consolidation pass may only "
f"archive a skill it has absorbed into an umbrella. Pass absorbed_into=<umbrella> (the "
f"umbrella must already exist) to record a verified consolidation. Pruning a skill with no "
f"forwarding target is not permitted here — the deterministic inactivity prune handles "
f"staleness archival "
"separately. Keeping '{name}' active.".format(name=name),
f"staleness archival separately. Keeping '{name}' active.",
_fail_closed=True)
def _maybe_auto_propose_org_edit(name: str, skill_path: Path) -> Optional[str]:
"""Submit an org-skill edit upstream when `sync.org_auto_propose` is on.
Returns a note for the tool result or None; never raises (the edit is
already saved locally and can be proposed later)."""
def _is_org_mirror(skill_path: Path) -> bool:
from agent.skill_utils import is_org_mirror_path
from tools import skill_manager_tool as _smt
return is_org_mirror_path(skill_path, _smt._skills_dir())
def _maybe_auto_propose_org_edit(name: str, skill_path: Path) -> Optional[str]:
"""Submit an org-skill edit upstream when `sync.org_auto_propose` is on. Returns a note for
the tool result or None; never raises (the edit is saved locally and can be proposed later)."""
try:
from agent.skill_utils import is_org_mirror_path
from tools import skills_sync_client as ssc
if not is_org_mirror_path(skill_path, _smt._skills_dir()):
if not _is_org_mirror(skill_path):
return None
if not ssc.sync_org_auto_propose():
return (
@@ -294,20 +271,14 @@ def _maybe_auto_propose_org_edit(name: str, skill_path: Path) -> Optional[str]:
def _org_mirror_write_guard(name: str, skill_path: Path, action: str) -> Optional[Dict[str, Any]]:
"""Org-shared skills are EDITABLE IN PLACE — this only blocks deletion.
Edits land in the mirror, survive the next org pull (baseline sidecar in
skills_sync_client) and reach the org via `hermes sync propose`. Deletion
stays refused: the mirror is a view of org HEAD, so a local delete just
comes back, and removing for everyone is an admin action."""
"""Org-shared skills are EDITABLE IN PLACE — this only blocks deletion. Edits land in the
mirror, survive the next org pull (baseline sidecar in skills_sync_client) and reach the org
via `hermes sync propose`. Deletion stays refused: the mirror is a view of org HEAD, so a
local delete just comes back, and removing for everyone is an admin action."""
if action not in {"delete", "remove_file"}:
return None
from tools import skill_manager_tool as _smt
try:
from agent.skill_utils import is_org_mirror_path
if is_org_mirror_path(skill_path, _smt._skills_dir()):
if _is_org_mirror(skill_path):
return _refusal(
f"Cannot {action} '{name}' locally: it is shared by your organisation, so a local "
f"delete would just come back on the next sync. Ask an org admin to remove it for "
+157 -269
View File
@@ -1,11 +1,10 @@
#!/usr/bin/env python3
"""Skill Manager Tool — agent-managed skill creation & editing.
Skills are the agent's procedural memory (narrow "how to do X"), as opposed to
MEMORY.md/USER.md (broad, declarative). New skills land in ~/.hermes/skills/
(or ``skills.create_dir``); existing skills (bundled, hub, user) are modified in
place. Layout: ``<skills>/[category/]<skill>/SKILL.md`` + optional
``references/ templates/ scripts/ assets/``.
Skills are the agent's procedural memory (narrow "how to do X"; MEMORY.md/USER.md are
broad, declarative). New skills land in ~/.hermes/skills/ (or ``skills.create_dir``);
existing skills (bundled, hub, user) are modified in place. Layout:
``<skills>/[category/]<skill>/SKILL.md`` + optional ``references/ templates/ scripts/ assets/``.
"""
import contextvars as _ctxvars
@@ -34,7 +33,7 @@ from tools.skill_manager_guards import ( # noqa: F401 — re-exported for calle
_background_review_write_guard, _containing_skills_root, _curator_consolidation_delete_guard,
_is_path_redirect, _maybe_auto_propose_org_edit, _org_mirror_write_guard, _pinned_guard,
_reset_background_review_read_marks, _validate_delete_target, _is_background_review,
mark_background_review_skill_read)
mark_background_review_skill_read, _refusal as _err)
from tools.skill_manager_batch import ( # noqa: F401
_BATCH_MAX_OPS, _BATCH_OP_ACTIONS, _skill_manage_batch)
from tools.skills_guard import scan_skill, should_allow_install, format_scan_report
@@ -43,20 +42,17 @@ logger = logging.getLogger(__name__)
def _guard_agent_created_enabled() -> bool:
"""skills.guard_agent_created (default False): the agent can already run the same
code via terminal() ungated, so the scan is opt-in belt-and-suspenders."""
"""skills.guard_agent_created (default False): opt-in — terminal() runs the same code ungated."""
try:
from hermes_cli.config import load_config
return is_truthy_value(
cfg_get(load_config(), "skills", "guard_agent_created"), default=False)
return is_truthy_value(cfg_get(load_config(), "skills", "guard_agent_created"), default=False)
except Exception:
return False
def _security_scan_skill(skill_dir: Path) -> Optional[str]:
"""Post-write scan; error string if blocked, else None. No-op unless
skills.guard_agent_created. An "ask" verdict (dangerous findings) is surfaced
as an error so the agent can retry with the flagged content removed."""
"""Post-write scan (opt-in); error string if blocked, else None. An "ask" verdict
(dangerous findings) is surfaced as an error so the agent can retry without them."""
if not _guard_agent_created_enabled():
return None
try:
@@ -78,9 +74,8 @@ _SKILLS_DIR_AT_IMPORT = SKILLS_DIR
def _skills_dir() -> Path:
"""Active profile's skills dir at call time: multi-profile runtimes import once
and bind a different profile per session. An explicitly patched module-level
``SKILLS_DIR`` (tests) wins, otherwise resolve from the live HERMES_HOME."""
"""Active profile's skills dir at call time (multi-profile runtimes rebind per session).
An explicitly patched module-level ``SKILLS_DIR`` (tests) wins over the live HERMES_HOME."""
configured = Path(SKILLS_DIR)
return configured if configured != _SKILLS_DIR_AT_IMPORT else get_hermes_home() / "skills"
@@ -104,45 +99,38 @@ def _display_create_dir() -> str:
return f"{display_hermes_home()}/skills/"
def _err(message: str) -> Dict[str, Any]:
return {"success": False, "error": message}
# --- Validation helpers -------------------------------------------------------
def _check_identifier(value: str, label: str, invalid: str) -> Optional[str]:
if len(value) > MAX_NAME_LENGTH:
return f"{label} exceeds {MAX_NAME_LENGTH} characters."
return None if VALID_NAME_RE.match(value) else invalid
def _validate_name(name: str) -> Optional[str]:
if not name:
return "Skill name is required."
if len(name) > MAX_NAME_LENGTH:
return f"Skill name exceeds {MAX_NAME_LENGTH} characters."
if not VALID_NAME_RE.match(name):
return f"Invalid skill name '{name}'. {_NAME_RULE} Must start with a letter or digit."
return None
return _check_identifier(
name, "Skill name", f"Invalid skill name '{name}'. {_NAME_RULE} Must start with a letter or digit.")
def _validate_category(category: Optional[str]) -> Optional[str]:
if category is None:
if category is None or (isinstance(category, str) and not category.strip()):
return None
if not isinstance(category, str):
return "Category must be a string."
category = category.strip()
if not category:
return None
invalid = (f"Invalid category '{category}'. {_NAME_RULE} "
"Categories must be a single directory name.")
if "/" in category or "\\" in category:
return invalid
if len(category) > MAX_NAME_LENGTH:
return f"Category exceeds {MAX_NAME_LENGTH} characters."
return None if VALID_NAME_RE.match(category) else invalid
return _check_identifier(category, "Category", invalid)
def _validate_frontmatter(content: str, *, new_skill: bool = False) -> Optional[str]:
"""Validate frontmatter (name + description) and a non-empty body.
``new_skill`` (create only) also enforces SKILL_PROMPT_DESC_LIMIT so new skills
never lose routing signal to index truncation; edit/patch skip it so existing
over-limit skills remain maintainable."""
"""Validate frontmatter (name + description) and a non-empty body. ``new_skill`` (create
only) also enforces SKILL_PROMPT_DESC_LIMIT so new skills never lose routing signal to
index truncation; edit/patch skip it so existing over-limit skills stay maintainable."""
if not content.strip():
return "Content cannot be empty."
content = content.lstrip("\ufeff") # tolerate a Windows UTF-8 BOM
@@ -157,10 +145,9 @@ def _validate_frontmatter(content: str, *, new_skill: bool = False) -> Optional[
return f"YAML frontmatter parse error: {e}"
if not isinstance(parsed, dict):
return "Frontmatter must be a YAML mapping (key: value pairs)."
if "name" not in parsed:
return "Frontmatter must include 'name' field."
if "description" not in parsed:
return "Frontmatter must include 'description' field."
for field in ("name", "description"):
if field not in parsed:
return f"Frontmatter must include '{field}' field."
desc = str(parsed["description"])
if len(desc) > MAX_DESCRIPTION_LENGTH:
return f"Description exceeds {MAX_DESCRIPTION_LENGTH} characters."
@@ -199,12 +186,10 @@ def _resolve_skill_dir(name: str, category: str = None) -> Path:
base = _skills_dir()
try:
from agent.skill_utils import get_skill_create_dir
create_dir = get_skill_create_dir()
if create_dir is not None:
base = create_dir
base = get_skill_create_dir() or base
except Exception:
logger.debug("skills.create_dir lookup failed", exc_info=True)
return base / category / name if category else base / name
return base / (category or "") / name
def _iter_skill_dirs(root: Path):
@@ -215,16 +200,12 @@ def _iter_skill_dirs(root: Path):
def _find_skill(name: str) -> Optional[Dict[str, Any]]:
"""Find a skill across the local skills dir then skills.external_dirs.
"""Find a skill (local skills dir, then skills.external_dirs) -> ``{"path": Path}`` | None.
Accepts the bare dir name (``axolotl``) and the categorized relative path
(``mlops/axolotl``) — the two forms skill_view resolves. Bare lookups compare
the skill's own dir name so category-nested skills still match.
Returns ``{"path": Path}`` or None."""
Accepts the bare dir name (``axolotl``; matches category-nested skills too) and the
categorized relative path (``mlops/axolotl``) — the two forms skill_view resolves. The
categorized form matches RELATIVE to the local root only (relative_to raises for external dirs)."""
from agent.skill_utils import get_all_skills_dirs
# The categorized form matches RELATIVE to the local root only (relative_to
# raises for external dirs).
local_root = None
if "/" in name or "\\" in name:
try:
@@ -234,7 +215,6 @@ def _find_skill(name: str) -> Optional[Dict[str, Any]]:
"skills dir resolve failed; categorized lookups fall back to the unresolved path",
exc_info=True)
local_root = _skills_dir()
for skills_dir in get_all_skills_dirs():
if not skills_dir.exists():
continue
@@ -250,8 +230,8 @@ def _find_skill(name: str) -> Optional[Dict[str, Any]]:
def _find_skill_in_other_profiles(name: str) -> List[Tuple[str, Path]]:
"""``(profile, skill_dir)`` pairs for OTHER profiles holding ``name`` (so the
not-found error can explain a wrong-profile mistake). Fail-quiet."""
"""``(profile, skill_dir)`` pairs for OTHER profiles holding ``name`` (so the not-found
error can explain a wrong-profile mistake). Fail-quiet."""
matches: List[Tuple[str, Path]] = []
try:
from hermes_constants import get_default_hermes_root
@@ -262,19 +242,16 @@ def _find_skill_in_other_profiles(name: str) -> List[Tuple[str, Path]]:
active_dir = _active.resolve() if _active.exists() else _active
# Every profile's skills dir EXCEPT the active one (already searched).
candidates: List[Tuple[str, Path]] = [("default", root / "skills")]
profiles_root = root / "profiles"
with suppress(OSError):
if profiles_root.is_dir():
candidates += [(e.name, e / "skills") for e in profiles_root.iterdir() if e.is_dir()]
if (root / "profiles").is_dir():
candidates += [(e.name, e / "skills") for e in (root / "profiles").iterdir() if e.is_dir()]
for profile_name, skills_dir in candidates:
try:
with suppress(OSError, RuntimeError):
if skills_dir.resolve() == active_dir or not skills_dir.is_dir():
continue
hit = next((d for d in _iter_skill_dirs(skills_dir) if d.name == name), None)
if hit is not None:
matches.append((profile_name, hit)) # one match per profile is enough
except (OSError, RuntimeError):
continue
return matches
@@ -302,7 +279,6 @@ def _skill_not_found_error(name: str, suffix: str = "") -> str:
def _validate_file_path(file_path: str) -> Optional[str]:
"""Validate a write_file/remove_file path: under an allowed subdir, no escape."""
from tools.path_security import has_traversal_component
if not file_path:
return "file_path is required."
parts = Path(file_path).parts
@@ -324,40 +300,31 @@ def _resolve_supporting_file(skill_dir: Path, file_path: str):
"""Validate ``file_path`` and resolve it inside ``skill_dir``
-> ``(target, None)`` | ``(None, error_dict)``."""
from tools.path_security import validate_within_dir
err = _validate_file_path(file_path)
if err:
return None, _err(err)
target = skill_dir / file_path
err = validate_within_dir(target, skill_dir)
if err:
return None, _err(err)
return target, None
target = skill_dir / (file_path or "")
err = _validate_file_path(file_path) or validate_within_dir(target, skill_dir)
return (None, _err(err)) if err else (target, None)
def _locate_for_write(name: str, action: str, not_found_suffix: str = ""):
"""Find the skill and run the org-mirror + background-review write guards
-> ``(skill_dir, None)`` | ``(None, error_dict)``."""
def _locate_for_write(name: str, action: str, not_found_suffix: str = "", *,
org_guard: bool = True):
"""Find the skill; run the org-mirror (unless ``org_guard=False``) and background-review
write guards -> ``(skill_dir, None)`` | ``(None, error_dict)``."""
existing = _find_skill(name)
if not existing:
return None, _err(_skill_not_found_error(name, not_found_suffix))
skill_dir = existing["path"]
guard = (_org_mirror_write_guard(name, skill_dir, action)
guard = ((org_guard and _org_mirror_write_guard(name, skill_dir, action))
or _background_review_write_guard(name, skill_dir, action))
if guard:
return None, guard
return skill_dir, None
return (None, guard) if guard else (skill_dir, None)
def _guarded_write(name: str, skill_dir: Path, target: Path, action: str, label: str,
content: str) -> Optional[Dict[str, Any]]:
"""Read-before-write guard (existing targets only), atomic write, then the
security scan; a blocked scan restores the original (or unlinks a new file).
Returns an error dict or None."""
"""Read-before-write guard (existing targets only), atomic write, then the security scan;
a blocked scan restores the original (or unlinks a new file). Error dict or None."""
original = None
if target.exists():
read_guard = _background_review_read_before_write_guard(name, target, action, label)
if read_guard:
if read_guard := _background_review_read_before_write_guard(name, target, action, label):
return read_guard
original = target.read_text(encoding="utf-8")
target.parent.mkdir(parents=True, exist_ok=True)
@@ -372,19 +339,20 @@ def _guarded_write(name: str, skill_dir: Path, target: Path, action: str, label:
return _err(scan_error)
def _attach_org_note(result: Dict[str, Any], name: str, skill_dir: Path) -> None:
org_note = _maybe_auto_propose_org_edit(name, skill_dir)
if org_note:
def _attach_org_note(result: Dict[str, Any], name: str, skill_dir: Path) -> Dict[str, Any]:
if org_note := _maybe_auto_propose_org_edit(name, skill_dir):
result["org_sharing"] = org_note
result["message"] = f"{result['message']} {org_note}"
return result
def _add_description_prompt_preview(result: Dict[str, Any], content: str) -> None:
def _add_description_prompt_preview(result: Dict[str, Any], content: str) -> Dict[str, Any]:
fm, _ = _parse_frontmatter(content)
if is_skill_description_truncated_for_prompt(fm):
result["system_prompt_preview"] = (
f"System prompt will show: \"{extract_skill_description(fm)}\" — keep the trigger "
f"self-contained in the first {SKILL_PROMPT_DESC_LIMIT - 3} chars.")
return result
def _attach_lint_findings(result: Dict[str, Any], skill_md: Path) -> None:
@@ -393,7 +361,7 @@ def _attach_lint_findings(result: Dict[str, Any], skill_md: Path) -> None:
from tools.skill_linter import lint_skill # local import: optional path
findings = lint_skill(skill_md)
except Exception:
return
findings = None
if not findings:
return
result["lint_warnings"] = [
@@ -403,17 +371,18 @@ def _attach_lint_findings(result: Dict[str, Any], skill_md: Path) -> None:
"— fix them with skill_manage(action='patch') to match Hermes skill standards.")
def _clip(text: str, n: int, ellipsis: str) -> str:
return text[:n] + (ellipsis if len(text) > n else "")
# --- Core actions -------------------------------------------------------------
def _create_skill(name: str, content: str, category: str = None) -> Dict[str, Any]:
err = (_validate_name(name) or _validate_category(category)
or _validate_frontmatter(content, new_skill=True) or _validate_content_size(content))
if err:
if err := (_validate_name(name) or _validate_category(category)
or _validate_frontmatter(content, new_skill=True) or _validate_content_size(content)):
return _err(err)
existing = _find_skill(name)
if existing:
if existing := _find_skill(name):
return _err(f"A skill named '{name}' already exists at {existing['path']}.")
skill_dir = _resolve_skill_dir(name, category)
skill_dir.mkdir(parents=True, exist_ok=True)
skill_md = skill_dir / "SKILL.md"
@@ -421,52 +390,39 @@ def _create_skill(name: str, content: str, category: str = None) -> Dict[str, An
if scan_error := _security_scan_skill(skill_dir):
shutil.rmtree(skill_dir, ignore_errors=True)
return _err(scan_error)
root = _skills_dir()
# Relative when under the profile dir; absolute when created under skills.create_dir.
root = _skills_dir() # display relative under the profile dir; absolute under skills.create_dir
display = skill_dir.relative_to(root) if skill_dir.is_relative_to(root) else skill_dir
result = {
"success": True, "message": f"Skill '{name}' created.", "path": str(display),
"skill_md": str(skill_md), "_change": {"description": _description_preview(content)},
}
if category:
result["category"] = category
result["hint"] = (
"To add reference files, templates, or scripts, use "
"skill_manage(action='write_file', name='{}', file_path='references/example.md', file_content='...')".format(name)
)
_add_description_prompt_preview(result, content)
_attach_lint_findings(result, skill_md)
**({"category": category} if category else {}),
"hint": "To add reference files, templates, or scripts, use "
f"skill_manage(action='write_file', name='{name}', file_path='references/example.md', "
"file_content='...')"}
_attach_lint_findings(_add_description_prompt_preview(result, content), skill_md)
return result
def _edit_skill(name: str, content: str) -> Dict[str, Any]:
"""Replace the SKILL.md of any existing skill (full rewrite)."""
err = _validate_frontmatter(content) or _validate_content_size(content)
if err:
if err := _validate_frontmatter(content) or _validate_content_size(content):
return _err(err)
skill_dir, guard = _locate_for_write(name, "edit")
if guard:
return guard
# SKILL.md always exists here (_find_skill requires it), so a blocked scan restores it.
guard = _guarded_write(name, skill_dir, skill_dir / "SKILL.md", "edit", "SKILL.md", content)
if guard:
if guard := guard or _guarded_write(name, skill_dir, skill_dir / "SKILL.md", "edit", "SKILL.md", content):
return guard
result = {
"success": True, "message": f"Skill '{name}' updated (full rewrite).",
"path": str(skill_dir), "_change": {"description": _description_preview(content)}}
_attach_org_note(result, name, skill_dir)
_add_description_prompt_preview(result, content)
return result
return _add_description_prompt_preview(_attach_org_note(result, name, skill_dir), content)
def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = None,
replace_all: bool = False) -> Dict[str, Any]:
"""Targeted find-and-replace within SKILL.md (default) or a supporting file;
requires a unique match unless replace_all."""
"""Targeted find-and-replace in SKILL.md (default) or a supporting file; unique match unless replace_all."""
if not old_string:
# A bare "required" error is a dead end: the model retries blindly and
# often escapes to action='write_file', clobbering the whole file.
# A bare "required" error is a dead end: the model retries blindly and often
# escapes to action='write_file', clobbering the whole file.
return _err(
"old_string is required for 'patch' and must be the EXACT text currently in the file. "
"Read the target file first (read_file on the skill's SKILL.md, or the file named by "
@@ -474,13 +430,12 @@ def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = N
"action='write_file' — that rewrites the entire file and destroys unrelated content.")
if new_string is None:
return _err("new_string is required for 'patch'. Use an empty string to delete matched text.")
# No old_string == new_string guard here: fuzzy_find_and_replace rejects
# that with a richer error (file_preview) this layer cannot produce.
# No old_string == new_string guard here: fuzzy_find_and_replace rejects that with a
# richer error (file_preview) this layer cannot produce.
skill_dir, guard = _locate_for_write(name, "patch")
if guard:
return guard
target_label = file_path or "SKILL.md"
if file_path:
target, err = _resolve_supporting_file(skill_dir, file_path)
if err:
@@ -489,14 +444,11 @@ def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = N
target = skill_dir / "SKILL.md"
if not target.exists():
return _err(f"File not found: {target.relative_to(skill_dir)}")
target_label = file_path or "SKILL.md"
read_guard = _background_review_read_before_write_guard(name, target, "patch", target_label)
if read_guard:
if read_guard := _background_review_read_before_write_guard(name, target, "patch", target_label):
return read_guard
content = target.read_text(encoding="utf-8")
# Same fuzzy engine as the file patch tool (whitespace/indent/escape
# normalization, block anchors) so minor formatting mismatches don't fail.
# Same fuzzy engine as the file patch tool (whitespace/indent/escape normalization,
# block anchors) so minor formatting mismatches don't fail.
from tools.fuzzy_match import fuzzy_find_and_replace
new_content, match_count, _strategy, match_error = fuzzy_find_and_replace(
content, old_string, new_string, replace_all)
@@ -504,42 +456,28 @@ def _patch_skill(name: str, old_string: str, new_string: str, file_path: str = N
with suppress(Exception):
from tools.fuzzy_match import format_no_match_hint
match_error += format_no_match_hint(match_error, match_count, old_string, content)
return _err(match_error) | {"file_preview": content[:500] + ("..." if len(content) > 500 else "")}
err = _validate_content_size(new_content, label=target_label)
if err:
return _err(match_error) | {"file_preview": _clip(content, 500, "...")}
if err := _validate_content_size(new_content, label=target_label):
return _err(err)
if not file_path:
err = _validate_frontmatter(new_content)
if err:
return _err(f"Patch would break SKILL.md structure: {err}")
guard = _guarded_write(name, skill_dir, target, "patch", target_label, new_content)
if guard:
if not file_path and (err := _validate_frontmatter(new_content)):
return _err(f"Patch would break SKILL.md structure: {err}")
if guard := _guarded_write(name, skill_dir, target, "patch", target_label, new_content):
return guard
result = {
"success": True,
"message": f"Patched {target_label} in skill '{name}' ({match_count} replacement{'s' if match_count > 1 else ''}).",
"_change": {"old": old_string[:200] + ("…" if len(old_string) > 200 else ""),
"new": new_string[:200] + ("…" if len(new_string) > 200 else "")}}
_attach_org_note(result, name, skill_dir)
return result
"_change": {"old": _clip(old_string, 200, "…"), "new": _clip(new_string, 200, "…")}}
return _attach_org_note(result, name, skill_dir)
def _delete_skill(name: str, absorbed_into: Optional[str] = None) -> Dict[str, Any]:
"""Delete a skill. ``absorbed_into``: None = undeclared (legacy, accepted); "" =
explicit prune; "<skill>" = absorbed into that umbrella, which must exist on
disk (validated here so the model can't claim a nonexistent umbrella)."""
"""Delete a skill. ``absorbed_into``: None = undeclared (legacy, accepted); "" = explicit prune;
"<skill>" = absorbed into that umbrella, which must exist (so the model can't claim one)."""
skill_dir, guard = _locate_for_write(name, "delete")
if guard:
if guard := guard or _curator_consolidation_delete_guard(name, absorbed_into):
return guard
fail_closed = _curator_consolidation_delete_guard(name, absorbed_into)
if fail_closed:
return fail_closed
pinned_err = _pinned_guard(name)
if pinned_err:
if pinned_err := _pinned_guard(name):
return _err(pinned_err)
absorbed_target = absorbed_into.strip() if isinstance(absorbed_into, str) else ""
if absorbed_target:
if absorbed_target == name:
@@ -547,14 +485,11 @@ def _delete_skill(name: str, absorbed_into: Optional[str] = None) -> Dict[str, A
if not _find_skill(absorbed_target):
return _err(f"absorbed_into='{absorbed_target}' does not exist. "
f"Create or patch the umbrella skill first, then retry the delete.")
skills_root = _containing_skills_root(skill_dir)
unsafe = _validate_delete_target(skill_dir) # defense-in-depth before rmtree
if unsafe:
if unsafe := _validate_delete_target(skill_dir): # defense-in-depth before rmtree
return _err(unsafe)
# Curator consolidations must be RECOVERABLE (`hermes curator restore`): archive
# instead of rmtree. Foreground deletes keep hard-delete semantics.
# Curator consolidations must be RECOVERABLE (`hermes curator restore`): archive instead
# of rmtree. Foreground deletes keep hard-delete semantics.
absorbed_note = f" Content absorbed into '{absorbed_target}'." if absorbed_target else ""
if _is_background_review():
try:
@@ -567,7 +502,6 @@ def _delete_skill(name: str, absorbed_into: Optional[str] = None) -> Dict[str, A
return {"success": True,
"message": f"Skill '{name}' archived ({archive_msg}).{absorbed_note}",
"_archived": True}
shutil.rmtree(skill_dir)
_rmdir_if_empty(skill_dir.parent, skills_root) # empty category dir, never the root
return {"success": True, "message": f"Skill '{name}' deleted.{absorbed_note}"}
@@ -580,60 +514,41 @@ def _rmdir_if_empty(parent: Path, stop: Path) -> None:
def _write_file(name: str, file_path: str, file_content: str) -> Dict[str, Any]:
"""Add or overwrite a supporting file within any skill directory."""
err = _validate_file_path(file_path)
if err:
if err := _validate_file_path(file_path):
return _err(err)
if not file_content and file_content != "":
return _err("file_content is required.")
content_bytes = len(file_content.encode("utf-8"))
if content_bytes > MAX_SKILL_FILE_BYTES:
if (content_bytes := len(file_content.encode("utf-8"))) > MAX_SKILL_FILE_BYTES:
return _err(f"File content is {content_bytes:,} bytes (limit: {MAX_SKILL_FILE_BYTES:,} "
f"bytes / 1 MiB). Consider splitting into smaller files.")
err = _validate_content_size(file_content, label=file_path)
if err:
if err := _validate_content_size(file_content, label=file_path):
return _err(err)
skill_dir, guard = _locate_for_write(name, "write_file", " Create it first with action='create'.")
if guard:
return guard
target, err = _resolve_supporting_file(skill_dir, file_path)
if err:
return err
guard = _guarded_write(name, skill_dir, target, "write_file", file_path, file_content)
if guard:
if guard := err or _guarded_write(name, skill_dir, target, "write_file", file_path, file_content):
return guard
result = {"success": True, "message": f"File '{file_path}' written to skill '{name}'.",
"path": str(target)}
_attach_org_note(result, name, skill_dir)
return result
return _attach_org_note({"success": True, "message": f"File '{file_path}' written to skill '{name}'.",
"path": str(target)}, name, skill_dir)
def _remove_file(name: str, file_path: str) -> Dict[str, Any]:
"""Remove a supporting file from any skill directory."""
err = _validate_file_path(file_path)
if err:
if err := _validate_file_path(file_path):
return _err(err)
existing = _find_skill(name)
if not existing:
return _err(_skill_not_found_error(name))
skill_dir = existing["path"]
guard = _background_review_write_guard(name, skill_dir, "remove_file")
skill_dir, guard = _locate_for_write(name, "remove_file", org_guard=False)
if guard:
return guard
target, err = _resolve_supporting_file(skill_dir, file_path)
if err:
return err
if not target.exists():
available = [
str(f.relative_to(skill_dir)) for subdir in ALLOWED_SUBDIRS
if (skill_dir / subdir).exists()
for f in (skill_dir / subdir).rglob("*") if f.is_file()]
return {"success": False, "error": f"File '{file_path}' not found in skill '{name}'.",
"available_files": available if available else None}
read_guard = _background_review_read_before_write_guard(name, target, "remove_file", file_path)
if read_guard:
if not target.exists(): # list what IS there so the model can pick the right path
available = [str(f.relative_to(skill_dir)) for subdir in ALLOWED_SUBDIRS
if (skill_dir / subdir).exists() for f in (skill_dir / subdir).rglob("*") if f.is_file()]
return _err(f"File '{file_path}' not found in skill '{name}'.", available_files=available or None)
if read_guard := _background_review_read_before_write_guard(name, target, "remove_file", file_path):
return read_guard
target.unlink()
_rmdir_if_empty(target.parent, skill_dir)
return {"success": True, "message": f"File '{file_path}' removed from skill '{name}'."}
@@ -641,18 +556,15 @@ def _remove_file(name: str, file_path: str) -> Dict[str, Any]:
# --- Main entry point ---------------------------------------------------------
# Set while replaying an already-approved staged skill write so skill_manage()
# does not re-gate (and re-stage) it.
# Set while replaying an approved staged skill write so skill_manage() does not re-gate it.
_skill_gate_bypass: "_ctxvars.ContextVar[bool]" = _ctxvars.ContextVar(
"skill_gate_bypass", default=False)
_GATED_ACTIONS = {"create", "edit", "patch", "delete", "write_file", "remove_file"}
def _run_write_gate(build_staging):
"""Shared write gate: None to proceed, else a JSON tool result (blocked/staged).
``build_staging(wa) -> (payload, gist)`` runs only when staging. Fails open
if write_approval cannot be imported."""
``build_staging(wa) -> (payload, gist)`` runs only when staging. Fails open if
write_approval cannot be imported."""
try:
from tools import write_approval as wa
except Exception:
@@ -669,26 +581,24 @@ def _run_write_gate(build_staging):
def _apply_skill_write_gate(action, name, **payload_kwargs):
"""Flat-shape gate: stage the full kwargs so approval can replay them;
bypassed during approved-pending replay."""
if action not in _GATED_ACTIONS or _skill_gate_bypass.get():
"""Flat-shape gate: stage the full kwargs so approval can replay them; bypassed during replay."""
if action not in _ACTION_HANDLERS or _skill_gate_bypass.get():
return None
def _staging(wa):
payload = {"action": action, "name": name,
**{k: v for k, v in payload_kwargs.items() if v is not None}}
gist_kw = {k: payload_kwargs.get(k) or ""
for k in ("content", "file_path", "old_string", "new_string")}
return payload, wa.skill_gist(action, name, **gist_kw)
return _run_write_gate(_staging)
_FLAT_OP_KEYS = ("content", "category", "file_path", "file_content", "old_string", "new_string")
_FLAT_OP_KEYS = ("content", "category", "file_path", "file_content", "old_string", "new_string",
"absorbed_into", "operations")
def _skill_manage_from(payload: Dict[str, Any], **extra) -> str:
"""Call ``skill_manage`` with the flat-shape fields taken from ``payload``."""
"""Call ``skill_manage`` with the flat-shape fields (and absorbed_into/operations) of ``payload``."""
return skill_manage(
action=payload.get("action", ""), name=payload.get("name", ""),
replace_all=payload.get("replace_all", False),
@@ -699,24 +609,20 @@ def apply_skill_pending(payload: Dict[str, Any]) -> str:
"""Replay a staged skill write, bypassing the gate (the /skills approve handler)."""
token = _skill_gate_bypass.set(True)
try:
return _skill_manage_from(
payload, absorbed_into=payload.get("absorbed_into"),
operations=payload.get("operations"))
return _skill_manage_from(payload)
finally:
_skill_gate_bypass.reset(token)
# Debounce state for the sync push hook: a burst of skill_manage writes
# collapses into one push after a quiet window, on a daemon timer.
# Sync push debounce: a burst of skill_manage writes collapses into one push on a daemon timer.
_sync_push_timer = None
_sync_push_lock = threading.Lock()
_SYNC_PUSH_DEBOUNCE_S = 5.0
def _maybe_debounced_sync_push(skill_name: str) -> None:
"""Debounced best-effort sync push after a skill write; never blocks the caller.
Skills not opted into sync do nothing (no auth, no network); the push itself
(``skills_sync_client.maybe_push_skills``) enforces the access gate."""
"""Debounced best-effort sync push after a skill write; never blocks the caller. Skills not
opted into sync do nothing (no auth/network); ``maybe_push_skills`` enforces the access gate."""
global _sync_push_timer
try:
from tools.skill_usage import is_sync_enabled
@@ -724,12 +630,10 @@ def _maybe_debounced_sync_push(skill_name: str) -> None:
return
except Exception:
return
def _fire():
with suppress(Exception):
from tools.skills_sync_client import maybe_push_skills
maybe_push_skills(message=f"sync: {skill_name}")
with _sync_push_lock:
if _sync_push_timer is not None:
_sync_push_timer.cancel() # only sets an Event; never raises
@@ -739,21 +643,18 @@ def _maybe_debounced_sync_push(skill_name: str) -> None:
def _act_patch(a):
# Two shapes: old_string/new_string = targeted replacement;
# content (alone) = full SKILL.md rewrite (absorbs the old 'edit').
"""Two shapes: old_string/new_string = targeted replacement (validated in _patch_skill so the
tool and the helper give the same guidance); content alone = full rewrite (the old 'edit')."""
if a["content"] and (a["old_string"] or a["new_string"] is not None):
return tool_error("Pass EITHER content (full SKILL.md rewrite) OR "
"old_string/new_string (targeted replacement), not both.", success=False)
if a["content"]:
return _edit_skill(a["name"], a["content"])
# Targeted-replacement validation lives in _patch_skill so the public
# tool and the helper return the same actionable guidance.
return _patch_skill(a["name"], a["old_string"], a["new_string"], a["file_path"], a["replace_all"])
# action -> handler(args dict). Handlers return a result dict, or a JSON string
# (tool_error) for argument-shape errors. "edit" is a legacy alias for a full
# rewrite (old transcripts/callers; not in the schema).
# action -> handler(args dict) returning a result dict, or a tool_error JSON string for
# argument-shape errors. "edit" is a legacy alias for a full rewrite (not in the schema).
_ACTION_HANDLERS = {
"create": lambda a: _create_skill(a["name"], a["content"], a["category"]),
"edit": lambda a: _edit_skill(a["name"], a["content"]),
@@ -762,33 +663,29 @@ _ACTION_HANDLERS = {
"write_file": lambda a: _write_file(a["name"], a["file_path"], a["file_content"]),
"remove_file": lambda a: _remove_file(a["name"], a["file_path"])}
# action -> (arg, is_missing, error) argument-shape checks run before the handler.
_MISSING, _IS_NONE = (lambda v: not v), (lambda v: v is None)
_REQUIRED_ARGS = {
"create": [("content", lambda v: not v,
"create": [("content", _MISSING,
"content is required for 'create'. Provide the full SKILL.md text (frontmatter + body).")],
"edit": [("content", lambda v: not v,
"edit": [("content", _MISSING,
"content is required for a full rewrite. Provide the full updated SKILL.md text.")],
"write_file": [("file_path", lambda v: not v,
"file_path is required for 'write_file'. Example: 'references/api-guide.md'"),
("file_content", lambda v: v is None, "file_content is required for 'write_file'.")],
"remove_file": [("file_path", lambda v: not v, "file_path is required for 'remove_file'.")]}
"write_file": [
("file_path", _MISSING, "file_path is required for 'write_file'. Example: 'references/api-guide.md'"),
("file_content", _IS_NONE, "file_content is required for 'write_file'.")],
"remove_file": [("file_path", _MISSING, "file_path is required for 'remove_file'.")]}
def _record_success(action, name, result, *, file_path, absorbed_into, task_id,
session_id, ledger_before) -> None:
"""Best-effort post-mutation side effects (never break the tool): audit ledger,
prompt-cache clear, curator telemetry, debounced sync push."""
"""Best-effort post-mutation side effects (never break the tool): ledger, prompt-cache
clear, curator telemetry, debounced sync push."""
with suppress(Exception):
from tools import skill_ledger as _ledger
_post = _find_skill(name)
_evidence = {}
if action == "delete":
# consolidation vs prune, and whether the recoverable archive handled it
_evidence["absorbed_into"] = absorbed_into
_evidence["archived"] = bool(result.get("_archived"))
if session_id:
_evidence["session_id"] = session_id
if file_path:
_evidence["file_path"] = file_path
# delete: consolidation vs prune, and whether the recoverable archive handled it
_evidence = ({"absorbed_into": absorbed_into, "archived": bool(result.get("_archived"))}
if action == "delete" else {})
_evidence.update({k: v for k, v in (("session_id", session_id), ("file_path", file_path)) if v})
_ledger.record_mutation(
action, name, before=ledger_before if ledger_before is not None else [],
after_root=_post["path"] if _post else None, evidence=_evidence)
@@ -808,8 +705,7 @@ def _record_success(action, name, result, *, file_path, absorbed_into, task_id,
bump_patch(name, action=action, task_id=task_id, session_id=session_id)
elif action == "delete" and not result.get("_archived"):
forget(name)
# Runs only AFTER the write gate passed (staged writes returned early), so
# un-reviewed content is never pushed.
# Only AFTER the write gate passed (staged writes returned early): never push un-reviewed content.
with suppress(Exception):
_maybe_debounced_sync_push(name)
@@ -819,44 +715,37 @@ def skill_manage(
file_content: str = None, old_string: str = None, new_string: str = None,
replace_all: bool = False, absorbed_into: str = None, task_id: str = None,
session_id: str = None, operations=None) -> str:
"""Dispatch to the action handler; returns a JSON string. ``operations`` (batch
shape, applied atomically by _skill_manage_batch) overrides the flat fields."""
"""Dispatch to the action handler -> JSON string. ``operations`` (atomic batch shape,
see _skill_manage_batch) overrides the flat fields."""
if operations is not None:
return _skill_manage_batch(
operations, default_name=name or None, task_id=task_id, session_id=session_id)
if (preflight := _background_review_preflight(action, name)) is not None:
return json.dumps(preflight, ensure_ascii=False)
# Approval gate: skills are too large to review inline, so they always stage
# regardless of origin; bypassed when replaying an approved staged write.
args = dict(
content=content, category=category, file_path=file_path,
file_content=file_content, old_string=old_string, new_string=new_string,
replace_all=replace_all, absorbed_into=absorbed_into)
# Approval gate: skills are too large to review inline, so they always stage regardless
# of origin; bypassed when replaying an approved staged write.
args = dict(content=content, category=category, file_path=file_path, file_content=file_content,
old_string=old_string, new_string=new_string, replace_all=replace_all,
absorbed_into=absorbed_into)
if (gate_result := _apply_skill_write_gate(action, name, **args)) is not None:
return gate_result
# Audit ledger pre-capture: telemetry, not a gate — failures must NEVER block the
# mutation. delete destroys the whole package (and consolidation may have re-homed
# support files first), so complete it from the newest curator backup or a restore is hollow.
# Ledger pre-capture: telemetry, not a gate — failures must NEVER block the mutation. delete
# destroys the whole package (consolidation may have re-homed support files first), so
# complete it from the newest curator backup or a restore is hollow.
_ledger_before = None
with suppress(Exception):
from tools import skill_ledger as _ledger
_pre = _find_skill(name)
_ledger_before = _ledger.capture_before(
_pre["path"] if _pre else None, complete_package=(action == "delete"), skill=name)
handler = _ACTION_HANDLERS.get(action)
if handler is None:
result = _err(f"Unknown action '{action}'. Use: create, edit, patch, delete, write_file, remove_file")
else:
for arg, missing, message in _REQUIRED_ARGS.get(action, ()):
if missing(args[arg]):
return tool_error(message, success=False)
result = handler({"name": name, **args})
if isinstance(result, str):
return result # tool_error JSON for argument-shape problems (patch)
for arg, missing, message in _REQUIRED_ARGS.get(action, ()):
if missing(args[arg]):
return tool_error(message, success=False)
handler = _ACTION_HANDLERS.get(action, lambda a: _err(
f"Unknown action '{action}'. Use: create, edit, patch, delete, write_file, remove_file"))
result = handler({"name": name, **args})
if isinstance(result, str):
return result # tool_error JSON for argument-shape problems (patch)
if result.get("success"):
_record_success(
action, name, result, file_path=file_path, absorbed_into=absorbed_into,
@@ -966,5 +855,4 @@ from tools.registry import registry, tool_error
registry.register(
name="skill_manage", toolset="skills", schema=SKILL_MANAGE_SCHEMA, emoji="📝",
handler=lambda args, **kw: _skill_manage_from(
args, absorbed_into=args.get("absorbed_into"), operations=args.get("operations"),
task_id=kw.get("task_id"), session_id=kw.get("session_id")))
args, task_id=kw.get("task_id"), session_id=kw.get("session_id")))