Merge simp/r3-33 (late tail) into hermes/simplify-codebase
This commit is contained in:
+54
-83
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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}")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
@@ -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)",
|
||||
|
||||
@@ -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}]")
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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})
|
||||
|
||||
@@ -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
@@ -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
@@ -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")))
|
||||
|
||||
Reference in New Issue
Block a user