refactor(tools): fold single-use MCP helpers, collapse defensive layers, compact docstrings
_stdio_children_dead_impl/_refresh_tools_task folded into their methods, _recover_401/_is_invalid_client_at_token_endpoint defensive getattr chains collapsed, lifecycle pid ledgers and drain loop tightened, WHY-preserving docstring compaction across the group. Schemas byte-identical.
This commit is contained in:
+32
-50
@@ -53,10 +53,9 @@ 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 concurrent POST, and HTTPX may close the auth-flow
|
||||
# generator from another task. A binary semaphore keeps mutual exclusion without task
|
||||
# ownership; async_auth_flow narrows its scope around resource I/O.
|
||||
# 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.
|
||||
import anyio
|
||||
self.context.lock = anyio.Semaphore(1, max_value=1)
|
||||
self._hermes_server_name = server_name
|
||||
@@ -75,14 +74,12 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
|
||||
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 (some providers answer 200 with an app-level auth error). Seeding the expiry makes
|
||||
the SDK refresh first; ``HermesTokenStorage`` persists absolute ``expires_at`` so the TTL
|
||||
reflects wall-clock age. Metadata is restored from disk, else discovered pre-flight when
|
||||
we hold tokens but no metadata: otherwise ``_refresh_token`` guesses ``{server_url}/token``
|
||||
"""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."""
|
||||
await super()._initialize()
|
||||
tokens = self.context.current_tokens
|
||||
@@ -113,8 +110,7 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
|
||||
from mcp.client.auth.utils import (
|
||||
build_oauth_authorization_server_metadata_discovery_urls,
|
||||
build_protected_resource_metadata_discovery_urls, create_oauth_metadata_request,
|
||||
handle_auth_metadata_response, handle_protected_resource_response,
|
||||
)
|
||||
handle_auth_metadata_response, handle_protected_resource_response)
|
||||
server_url = self.context.server_url
|
||||
|
||||
async def _send(client, url: str, label: str):
|
||||
@@ -167,14 +163,12 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
|
||||
``invalid_client_metadata`` does not trip it. The body is read only after the endpoint
|
||||
matches."""
|
||||
from urllib.parse import urlsplit
|
||||
meta = getattr(self.context, "oauth_metadata", None)
|
||||
token_endpoint = str(meta.token_endpoint) if meta is not None and getattr(meta, "token_endpoint", None) else None
|
||||
token_endpoint = getattr(getattr(self.context, "oauth_metadata", None), "token_endpoint", None)
|
||||
req = getattr(response, "request", None)
|
||||
req_url = str(req.url) if req is not None else None
|
||||
if not token_endpoint or not req_url:
|
||||
if not token_endpoint or req is None:
|
||||
return False
|
||||
try:
|
||||
pa, pb = urlsplit(req_url), urlsplit(token_endpoint)
|
||||
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()
|
||||
@@ -184,15 +178,12 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
|
||||
return re.search(rb"\binvalid_client\b", body.lower()) is not None
|
||||
|
||||
async def _maybe_flag_poisoned_client(self, response: Any) -> None:
|
||||
"""Detect a dead client registration and force re-registration.
|
||||
|
||||
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 status 400/401 at the discovered
|
||||
"""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 so a
|
||||
miss never breaks the live flow. The browser-side "Redirect URI Mismatch" case has no
|
||||
HTTP signal and is left to ``hermes mcp reauth``."""
|
||||
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
|
||||
@@ -226,9 +217,8 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
|
||||
except Exception as exc: # pragma: no cover — defensive
|
||||
self._log_nonfatal("pre-flow disk-watch", exc)
|
||||
|
||||
# Bridge the bidirectional generator protocol by hand: httpx feeds responses back via
|
||||
# ``auth_flow.asend(response)``. A naive ``async for item in inner: yield item`` DISCARDS
|
||||
# those values, so the SDK's ``response = yield request`` sees None and crashes.
|
||||
# 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
|
||||
sent_access_token = None
|
||||
@@ -237,8 +227,7 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
|
||||
outgoing = await inner.__anext__()
|
||||
while True:
|
||||
# The SDK holds context.lock for its whole generator, even while HTTPX waits on
|
||||
# the MCP request. Release it for that request only; discovery/refresh/
|
||||
# registration/exchange stay serialized exactly as the SDK implements them.
|
||||
# the MCP request. Release it for that request only; OAuth transitions stay serialized.
|
||||
if outgoing is request:
|
||||
tokens = self.context.current_tokens
|
||||
sent_access_token = tokens.access_token if tokens is not None else None
|
||||
@@ -249,8 +238,7 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
|
||||
await self.context.lock.acquire()
|
||||
resource_lock_released = False
|
||||
# Another request may have refreshed/authorized while this one was in flight:
|
||||
# retry with that token instead of a duplicate OAuth transition from the stale
|
||||
# 401/403.
|
||||
# retry with that token instead of a duplicate OAuth transition from a stale 401/403.
|
||||
tokens = self.context.current_tokens
|
||||
if (getattr(incoming, "status_code", None) in (401, 403) and self.context.is_token_valid()
|
||||
and tokens is not None and tokens.access_token != sent_access_token):
|
||||
@@ -262,7 +250,6 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
|
||||
outgoing = await inner.asend(incoming)
|
||||
except StopAsyncIteration:
|
||||
self._persist_oauth_metadata_if_changed() # metadata discovered lazily in the 401 branch
|
||||
return
|
||||
finally:
|
||||
if resource_lock_released:
|
||||
# Balance the SDK's surrounding ``async with`` even when HTTPX cancels/closes
|
||||
@@ -274,7 +261,6 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES):
|
||||
if retry_after_concurrent_auth:
|
||||
yield request
|
||||
self._persist_oauth_metadata_if_changed()
|
||||
return
|
||||
|
||||
|
||||
# Cached at import time; None when the SDK's OAuth module is unavailable.
|
||||
@@ -282,9 +268,8 @@ _HERMES_PROVIDER_CLS: Optional[type] = HermesMCPOAuthProvider if _SDK_BASES else
|
||||
|
||||
|
||||
class MCPOAuthManager:
|
||||
"""Single source of truth for per-server MCP OAuth state. Thread-safe: ``_entries`` is
|
||||
guarded by ``_entries_lock`` for get-or-create; per-entry state by the entry's own
|
||||
``asyncio.Lock`` (used from the MCP event loop thread)."""
|
||||
"""Single source of truth for per-server MCP OAuth state. ``_entries`` is guarded by
|
||||
``_entries_lock`` (get-or-create); per-entry state by the entry's ``asyncio.Lock``."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._entries: dict[tuple[str, str], _ProviderEntry] = {}
|
||||
@@ -294,9 +279,8 @@ class MCPOAuthManager:
|
||||
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. If ``server_url``
|
||||
changes for a name the cached entry is discarded and rebuilt. 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)
|
||||
@@ -388,10 +372,9 @@ class MCPOAuthManager:
|
||||
# caller retry (the httpx.Auth flow refreshes on the next request).
|
||||
can_refresh = True
|
||||
if not await self.invalidate_if_disk_changed(server_name):
|
||||
can_refresh_fn = getattr(getattr(entry.provider, "context", None), "can_refresh_token", None)
|
||||
try:
|
||||
can_refresh = bool(can_refresh_fn()) if callable(can_refresh_fn) else False
|
||||
except Exception:
|
||||
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)
|
||||
@@ -403,11 +386,10 @@ class MCPOAuthManager:
|
||||
entry.pending_401.pop(key, None)
|
||||
|
||||
async def handle_401(self, server_name: str, failed_access_token: Optional[str] = None) -> bool:
|
||||
"""Handle a 401 from a tool call, deduplicated across concurrent callers. True: a
|
||||
(possibly new) access token is available — reconnect and retry. False: no recovery
|
||||
path — surface a ``needs_reauth`` error so the model stops hallucinating manual
|
||||
refreshes. N 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
|
||||
|
||||
+11
-25
@@ -13,25 +13,19 @@ from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tools.mcp_oauth import HermesTokenStorage
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class HermesProviderMixin:
|
||||
"""Token-endpoint fixes layered over the SDK's ``OAuthClientProvider``.
|
||||
"""Token-endpoint fixes layered over the SDK's ``OAuthClientProvider`` (must precede it in
|
||||
the MRO; subclasses set ``_hermes_logger`` to keep their own logger name).
|
||||
|
||||
- Supabase-style dynamic registration returns a ``client_secret`` but omits
|
||||
``token_endpoint_auth_method``; the SDK then treats the client as public, omits the
|
||||
secret, and the token endpoint rejects the exchange (looping the browser authorization
|
||||
page). Coerce the in-memory client info to ``client_secret_post`` before token requests.
|
||||
- ``token_user_agent`` (``oauth.user_agent``) is stamped onto token-endpoint requests
|
||||
only — some authorization servers and WAFs reject httpx's default.
|
||||
- Any 2xx token/refresh response is accepted, and token bodies never leak into
|
||||
exception text or log output.
|
||||
|
||||
Must precede the SDK class in the MRO. Subclasses set ``_hermes_logger`` so warnings
|
||||
keep their origin module's logger name.
|
||||
"""
|
||||
``token_endpoint_auth_method``; the SDK then treats the client as public and the token
|
||||
endpoint rejects the exchange (looping the browser page) — coerce ``client_secret_post``.
|
||||
- ``token_user_agent`` (``oauth.user_agent``) is stamped onto token-endpoint requests only
|
||||
(some authorization servers/WAFs reject httpx's default).
|
||||
- Any 2xx token/refresh response is accepted; token bodies never leak into errors/logs."""
|
||||
|
||||
_hermes_logger: logging.Logger = logger
|
||||
|
||||
@@ -54,7 +48,6 @@ class HermesProviderMixin:
|
||||
return
|
||||
from mcp.shared.auth import OAuthClientInformationFull
|
||||
from tools.mcp_oauth import HermesTokenStorage
|
||||
|
||||
data = info.model_dump(mode="json", exclude_none=True)
|
||||
if HermesTokenStorage._coerce_secret_auth_method(data):
|
||||
self.context.client_info = OAuthClientInformationFull.model_validate(data)
|
||||
@@ -75,12 +68,10 @@ class HermesProviderMixin:
|
||||
async def _handle_token_response(self, response):
|
||||
"""Accept any 2xx token response; never echo the body into errors."""
|
||||
from mcp.client.auth.oauth2 import OAuthTokenError
|
||||
|
||||
if not (200 <= response.status_code < 300):
|
||||
raise OAuthTokenError(f"Token exchange failed ({response.status_code})")
|
||||
from httpx import HTTPError
|
||||
from mcp.client.auth.utils import handle_token_response_scopes
|
||||
|
||||
try:
|
||||
token_response = await handle_token_response_scopes(response)
|
||||
except (HTTPError, OAuthTokenError):
|
||||
@@ -96,7 +87,6 @@ class HermesProviderMixin:
|
||||
from httpx import HTTPError
|
||||
from mcp.shared.auth import OAuthToken
|
||||
from pydantic import ValidationError
|
||||
|
||||
try:
|
||||
token_response = OAuthToken.model_validate_json(await response.aread())
|
||||
except (HTTPError, ValidationError):
|
||||
@@ -112,21 +102,17 @@ def prepare_oauth_config(server_name: str, server_url: str, oauth_config: dict |
|
||||
matters: later steps record ``_resolved_port`` / ``_cimd_url`` in the dict, which must
|
||||
never leak back into the caller's config."""
|
||||
from tools import mcp_oauth as mo
|
||||
|
||||
cfg = dict(oauth_config or {})
|
||||
mo.apply_oauth_provider_defaults(cfg, server_name=server_name, server_url=server_url)
|
||||
return cfg, mo.HermesTokenStorage(server_name)
|
||||
|
||||
|
||||
def build_provider_kwargs(cfg: dict, storage: "HermesTokenStorage", *, ssh_proxy_hint: bool) -> dict[str, Any]:
|
||||
"""Resolve the callback port and return the shared provider constructor kwargs.
|
||||
|
||||
Runs port → client-metadata → pre-registration (order matters: metadata needs the
|
||||
resolved port, pre-registration needs the metadata). ``ssh_proxy_hint`` lets the redirect
|
||||
handler tailor its remote-session hint to a configured proxy ``redirect_uri``. Helpers are
|
||||
looked up on ``tools.mcp_oauth`` at call time so tests can patch them there."""
|
||||
"""Resolve the callback port and return the shared provider constructor kwargs. Order
|
||||
matters: metadata needs the resolved port, pre-registration needs the metadata.
|
||||
``ssh_proxy_hint`` lets the redirect handler tailor its remote-session hint to a configured
|
||||
proxy ``redirect_uri``. Helpers are looked up on ``tools.mcp_oauth`` so tests can patch them."""
|
||||
from tools import mcp_oauth as mo
|
||||
|
||||
port = mo._configure_callback_port(cfg, storage)
|
||||
client_metadata = mo._build_client_metadata(cfg)
|
||||
mo._maybe_preregister_client(storage, cfg, client_metadata)
|
||||
|
||||
+20
-35
@@ -33,8 +33,7 @@ def _resolve_refresh_toolsets(agent, enabled_override, disabled_override):
|
||||
if enabled_override is not None or disabled_override is not None:
|
||||
enabled = enabled_override if enabled_override is not None else enabled
|
||||
disabled = disabled_override if disabled_override is not None else disabled
|
||||
agent.enabled_toolsets = enabled
|
||||
agent.disabled_toolsets = disabled
|
||||
agent.enabled_toolsets, agent.disabled_toolsets = enabled, disabled
|
||||
return enabled, disabled
|
||||
|
||||
|
||||
@@ -50,8 +49,7 @@ def _tool_defs_content_changed(agent, new_defs: list) -> bool:
|
||||
|
||||
def _publish_tool_snapshot(
|
||||
agent, new_defs: list, new_names: set, *, snapshot_generation: int,
|
||||
staged_engine_names: set, content_aware: bool, prefix_registered: Optional[set],
|
||||
) -> Optional[set]:
|
||||
staged_engine_names: set, content_aware: bool, prefix_registered: Optional[set]) -> Optional[set]:
|
||||
"""Single atomic read-diff-publish under ``_agent_tools_lock`` so ``added`` matches what
|
||||
was published and a stale (older-generation) rebuild can't overwrite a newer one. Returns
|
||||
the added names, or None when nothing was published (unchanged, or a newer snapshot won)."""
|
||||
@@ -85,35 +83,27 @@ def refresh_agent_mcp_tools(
|
||||
agent, *, enabled_override=None, disabled_override=None, quiet_mode: bool = True,
|
||||
content_aware: bool = False, preserve_prefix: bool = False) -> set:
|
||||
"""Re-derive an already-built agent's tool snapshot from the live registry; returns the
|
||||
newly-added tool names (empty when unchanged).
|
||||
newly-added tool names (empty when unchanged). The agent snapshots ``agent.tools`` at build
|
||||
time, so servers that connect later (slow OAuth, ``/reload-mcp``) are invisible until
|
||||
rebuilt. Shared by the TUI RPC, gateway reload, late-binding thread and between-turns
|
||||
refresh: respects the toolset filter, diffs by tool NAME (a count compare misses an
|
||||
equal-size swap), re-injects the memory-provider / context-engine tools ``agent_init``
|
||||
appends after ``get_tool_definitions``, publishes ``(tools, valid_tool_names)`` together.
|
||||
|
||||
The agent snapshots ``agent.tools`` once at build time, so servers that connect later
|
||||
(slow OAuth server, ``/reload-mcp``) are invisible until rebuilt. Single shared rebuild for
|
||||
the TUI RPC, gateway reload, late-binding thread and between-turns refresh: respects the
|
||||
agent's toolset filter, diffs by tool NAME (a count compare misses an equal-size swap),
|
||||
re-injects the memory-provider / context-engine (``lcm_*``) tools ``agent_init`` appends
|
||||
after ``get_tool_definitions``, and publishes ``(tools, valid_tool_names)`` together
|
||||
under ``_agent_tools_lock``.
|
||||
|
||||
``preserve_prefix`` is for rebuilds inside a live conversation, where the tool array is a
|
||||
cached request prefix and any moved byte re-prefills the whole history: existing tools
|
||||
keep their slot (schemas still refresh), a still-registered tool whose ``check_fn`` merely
|
||||
flapped is carried forward (safe: ``check_fn`` gates exposure, never invocation), a
|
||||
deregistered tool is dropped, new tools append at the tail. The caller owns the
|
||||
prompt-cache contract (turn-boundary policy differs per caller)."""
|
||||
``preserve_prefix``: for rebuilds inside a live conversation the tool array is a cached
|
||||
request prefix and any moved byte re-prefills the whole history — existing tools keep their
|
||||
slot (schemas still refresh), a still-registered tool whose ``check_fn`` merely flapped is
|
||||
carried forward (``check_fn`` gates exposure, never invocation), a deregistered tool is
|
||||
dropped, new tools append at the tail. The caller owns the prompt-cache contract."""
|
||||
from model_tools import get_tool_definitions
|
||||
from tools.registry import registry
|
||||
|
||||
enabled, disabled = _resolve_refresh_toolsets(agent, enabled_override, disabled_override)
|
||||
# Capture the registry generation BEFORE the slow get_tool_definitions call; at publish
|
||||
# time a slower caller holding an OLDER set must not clobber a newer set already published.
|
||||
# Generation captured BEFORE the slow get_tool_definitions call (a slower caller holding an
|
||||
# OLDER set must not clobber a newer one); definitions computed OUTSIDE the lock.
|
||||
snapshot_generation = registry._generation
|
||||
# Computed OUTSIDE the lock (can be slow); diff + publish happen together in one critical
|
||||
# section so concurrent callers can't torn-publish.
|
||||
new_defs = list(get_tool_definitions(enabled_toolsets=enabled, disabled_toolsets=disabled, quiet_mode=quiet_mode) or [])
|
||||
new_names = {_def_name(t) for t in new_defs}
|
||||
# Re-append the post-build families on LOCALS only; live agent attributes are untouched
|
||||
# until the single atomic publish.
|
||||
# Post-build families re-appended on LOCALS only; live attributes untouched until publish.
|
||||
staged_engine_names = _core._reinject_post_build_tools(agent, new_defs, new_names)
|
||||
# Registry membership is read OUTSIDE ``_agent_tools_lock``: taking ``registry._lock``
|
||||
# under the tools lock would be the first nesting of the two.
|
||||
@@ -128,9 +118,7 @@ def refresh_agent_mcp_tools(
|
||||
staged_engine_names=staged_engine_names, content_aware=content_aware, prefix_registered=prefix_registered)
|
||||
if added is None:
|
||||
return set()
|
||||
# Re-pin the session's tool order so a rebuild-for-existing-session (gateway agent-cache
|
||||
# eviction) restores exactly these names.
|
||||
persist_agent_tool_names(agent)
|
||||
persist_agent_tool_names(agent) # re-pin so a rebuild after agent-cache eviction restores this order
|
||||
return added
|
||||
|
||||
|
||||
@@ -140,7 +128,6 @@ def reprobe_tool_availability() -> None:
|
||||
otherwise replay the stale verdicts)."""
|
||||
from model_tools import _clear_tool_defs_cache
|
||||
from tools.registry import invalidate_check_fn_cache
|
||||
|
||||
invalidate_check_fn_cache()
|
||||
_clear_tool_defs_cache()
|
||||
|
||||
@@ -159,14 +146,12 @@ def persist_agent_tool_names(agent) -> None:
|
||||
|
||||
def restore_agent_tool_prefix(agent, saved_names: list) -> bool:
|
||||
"""Fold a freshly built agent's ``tools`` onto the session's saved order; True if changed.
|
||||
The gateway rebuilds a NEW AIAgent for an existing session after agent-cache eviction, with
|
||||
no predecessor to preserve, so the saved name list stands in (same merge rule as
|
||||
``_merge_preserving_prefix``; a saved tool still registered but failing its probe is
|
||||
carried forward from the registry schema)."""
|
||||
After agent-cache eviction the gateway rebuilds a NEW AIAgent with no predecessor to
|
||||
preserve, so the saved name list stands in (``_merge_preserving_prefix`` rule; a saved
|
||||
tool still registered but failing its probe is carried forward from the registry schema)."""
|
||||
if not saved_names:
|
||||
return False
|
||||
from tools.registry import registry
|
||||
|
||||
fresh_defs = _agent_tool_defs(agent)
|
||||
fresh = {_def_name(t): t for t in fresh_defs}
|
||||
|
||||
|
||||
+19
-32
@@ -66,7 +66,8 @@ _SAFE_ENV_KEYS_CASE_INSENSITIVE = frozenset({
|
||||
"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"\$\{([^}]+)\}")
|
||||
@@ -77,7 +78,6 @@ def _workspace_folder() -> str:
|
||||
(terminal cwd / task override / $TERMINAL_CWD), else cwd."""
|
||||
try:
|
||||
from tools.file_tools import _authoritative_workspace_root
|
||||
|
||||
root = _authoritative_workspace_root()
|
||||
if root:
|
||||
return root
|
||||
@@ -97,7 +97,8 @@ _CONTEXT_VAR_RESOLVERS = {
|
||||
"workspaceFolder": lambda: _core._workspace_folder(),
|
||||
"workspaceFolderBasename": _workspace_basename,
|
||||
"pathSeparator": lambda: os.sep,
|
||||
"/": lambda: os.sep}
|
||||
"/": lambda: os.sep,
|
||||
}
|
||||
|
||||
|
||||
def _build_safe_env(user_env: Optional[dict]) -> dict:
|
||||
@@ -142,9 +143,8 @@ def _node_fallback(command: str) -> str:
|
||||
candidates = [
|
||||
os.path.join(hermes_home, "node", "bin", command),
|
||||
os.path.join(home, ".local", "bin", command),
|
||||
# Canonical Node location for from-source Linux builds, the Hermes Docker image and
|
||||
# Intel Homebrew. Needed when a hand-authored env.PATH omits it: npx's shebang re-execs
|
||||
# /usr/bin/env node, so a symlink workaround fails one layer deeper.
|
||||
# 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):
|
||||
@@ -189,18 +189,14 @@ def _wrap_command_with_watchdog(command: str, args: list) -> tuple[str, list]:
|
||||
|
||||
|
||||
def _interpolate_env_vars(value):
|
||||
"""Recursively resolve ``${VAR}`` / Cursor ``${env:VAR}`` placeholders plus the Cursor
|
||||
context vars. Env refs resolve from the active profile's secret scope when multiplexing
|
||||
(so ``${API_KEY}`` picks up the routed profile's value, not another profile's in
|
||||
``os.environ``). Unset vars keep the literal placeholder."""
|
||||
"""Recursively resolve ``${VAR}`` / Cursor ``${env:VAR}`` placeholders and context vars. Env
|
||||
refs resolve from the active profile's secret scope when multiplexing (the routed profile's
|
||||
value, not another profile's in ``os.environ``). Unset vars keep the literal placeholder."""
|
||||
from agent.secret_scope import get_secret as _get_secret
|
||||
|
||||
if isinstance(value, str):
|
||||
def _replace(m):
|
||||
resolver = _CONTEXT_VAR_RESOLVERS.get(m.group(1).strip())
|
||||
if resolver is not None:
|
||||
return resolver()
|
||||
return _get_secret(_env_ref_name(m.group(1)), m.group(0)) or m.group(0)
|
||||
return resolver() if resolver is not None else (_get_secret(_env_ref_name(m.group(1)), m.group(0)) or m.group(0))
|
||||
return _ENV_VAR_PATTERN.sub(_replace, value)
|
||||
if isinstance(value, dict):
|
||||
return {k: _interpolate_env_vars(v) for k, v in value.items()}
|
||||
@@ -215,16 +211,14 @@ _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/connect failures and is invisible in config.yaml.
|
||||
Advisory only: values are never mutated (whitespace could be intentional) and never
|
||||
logged (often secrets). Returns the flagged key paths."""
|
||||
"""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."""
|
||||
flagged: List[str] = []
|
||||
|
||||
def _walk(value: Any, path: str) -> None:
|
||||
if isinstance(value, str):
|
||||
if value != value.strip():
|
||||
flagged.append(path)
|
||||
if isinstance(value, str) and value != value.strip():
|
||||
flagged.append(path)
|
||||
elif isinstance(value, dict):
|
||||
for k, v in value.items():
|
||||
_walk(v, f"{path}.{k}" if path else str(k))
|
||||
@@ -234,10 +228,9 @@ def _warn_hidden_whitespace(server_name: str, config: dict) -> List[str]:
|
||||
|
||||
_walk(config, "")
|
||||
for key_path in flagged:
|
||||
dedupe_key = (server_name, key_path)
|
||||
if dedupe_key in _whitespace_warned:
|
||||
if (server_name, key_path) in _whitespace_warned:
|
||||
continue
|
||||
_whitespace_warned.add(dedupe_key)
|
||||
_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 "
|
||||
@@ -269,7 +262,6 @@ def _portable_mcp_servers(safe_servers: Dict[str, dict]) -> None:
|
||||
on a name clash. Never raises."""
|
||||
try:
|
||||
from hermes_cli.plugins import discover_plugins, get_plugin_manager
|
||||
|
||||
discover_plugins()
|
||||
portable = get_plugin_manager().get_portable_mcp_servers()
|
||||
for name, cfg in _core._filter_suspicious_mcp_servers(portable).items():
|
||||
@@ -283,25 +275,20 @@ def _portable_mcp_servers(safe_servers: Dict[str, dict]) -> None:
|
||||
|
||||
def _load_mcp_config() -> Dict[str, dict]:
|
||||
"""Read ``mcp_servers`` from config.yaml as ``{name: config}`` (empty on error or in safe
|
||||
mode). Entries carry ``command``/``args``/``env`` (stdio) or ``url``/``headers`` (HTTP)
|
||||
plus optional timeout/auth keys; ``${VAR}`` placeholders are interpolated after ``.env``
|
||||
is loaded."""
|
||||
mode); ``${VAR}`` placeholders are interpolated after ``.env`` is loaded."""
|
||||
try:
|
||||
from hermes_cli.config import load_config
|
||||
from utils import env_var_enabled as _env_enabled
|
||||
|
||||
if _env_enabled("HERMES_SAFE_MODE"):
|
||||
return {}
|
||||
servers = load_config().get("mcp_servers")
|
||||
if not isinstance(servers, dict):
|
||||
servers = {}
|
||||
try: # ensure .env vars are available for interpolation
|
||||
from hermes_cli.env_loader import load_hermes_dotenv
|
||||
load_hermes_dotenv()
|
||||
except Exception:
|
||||
pass
|
||||
safe_servers: Dict[str, dict] = {}
|
||||
for name, cfg in _core._filter_suspicious_mcp_servers(servers).items():
|
||||
for name, cfg in _core._filter_suspicious_mcp_servers(servers if isinstance(servers, dict) else {}).items():
|
||||
interpolated = _interpolate_env_vars(cfg)
|
||||
if isinstance(interpolated, dict):
|
||||
_warn_hidden_whitespace(name, interpolated)
|
||||
|
||||
+12
-32
@@ -13,14 +13,11 @@ from tools.mcp_tool_schema import mcp_prefixed_tool_name
|
||||
logger = logging.getLogger("tools.mcp_tool")
|
||||
|
||||
|
||||
# Hard allocation ceiling for one MCP text payload (chars): the first line of defense against
|
||||
# a multi-megabyte flood being JSON-encoded and handed downstream. Deliberately far ABOVE the
|
||||
# budget layer's 50K spillover threshold so ordinary large results reach spillover intact.
|
||||
# Hard ceiling for one MCP text payload (chars), deliberately far ABOVE the budget layer's 50K
|
||||
# spillover threshold so ordinary large results reach spillover intact; only floods are lossy.
|
||||
_MCP_HARD_RESULT_CAP_CHARS = 2_000_000
|
||||
|
||||
# Hard cap on decoded resource bytes from one block, so a misbehaving server can't fill the
|
||||
# cache disk. Base64 expands ~4/3; oversized payloads are rejected BEFORE decoding so a
|
||||
# multi-GB blob string is never transiently doubled in memory.
|
||||
# Cap on decoded resource bytes per block (a misbehaving server can't fill the cache disk).
|
||||
# Base64 expands ~4/3; oversized payloads are rejected BEFORE decoding (never doubled in memory).
|
||||
_MCP_RESOURCE_MAX_BYTES = 50 * 1024 * 1024
|
||||
_MCP_RESOURCE_MAX_B64_CHARS = _MCP_RESOURCE_MAX_BYTES * 4 // 3 + 4
|
||||
|
||||
@@ -33,11 +30,8 @@ def _truncate_mcp_text_result(text: str, max_chars: int = _MCP_HARD_RESULT_CAP_C
|
||||
head_chars = int(max_chars * 0.4)
|
||||
tail_chars = max_chars - head_chars
|
||||
omitted = len(text) - head_chars - tail_chars
|
||||
return (
|
||||
text[:head_chars]
|
||||
+ f"\n\n... [MCP RESULT TRUNCATED - {omitted:,} chars omitted "
|
||||
f"out of {len(text):,} total] ...\n\n"
|
||||
+ text[-tail_chars:])
|
||||
return (text[:head_chars] + f"\n\n... [MCP RESULT TRUNCATED - {omitted:,} chars omitted "
|
||||
f"out of {len(text):,} total] ...\n\n" + text[-tail_chars:])
|
||||
|
||||
|
||||
def _is_reserved_mcp_meta_key(key: str) -> bool:
|
||||
@@ -93,12 +87,10 @@ def _decode_block_b64(data, what: str, label: str, *, cap_what: Optional[str] =
|
||||
def _write_block_cache(writer: str, what: str, skip_label: str, *args,
|
||||
unavailable: str = "", failed: str = "", **kwargs) -> Tuple[Optional[str], str]:
|
||||
"""Call ``gateway.platforms.base.<writer>(*args, **kwargs)``: ``(path, "")`` or ``(None,
|
||||
marker)``. Fail-open: gateway deps missing (e.g. cron without gateway) → debug log +
|
||||
``unavailable``; any other cache error → warning + ``failed``. One bad block must never
|
||||
kill the tool result."""
|
||||
marker)``. Fail-open so one bad block never kills the tool result: gateway deps missing
|
||||
(cron without gateway) → ``unavailable``; any other cache error → warning + ``failed``."""
|
||||
try:
|
||||
import gateway.platforms.base as _base
|
||||
|
||||
return getattr(_base, writer)(*args, **kwargs), ""
|
||||
except ImportError:
|
||||
logger.debug("MCP %s caching skipped — gateway.platforms.base unavailable", skip_label)
|
||||
@@ -146,22 +138,18 @@ def _mcp_resource_filename(uri: str, mime_type: str) -> str:
|
||||
import re as _re
|
||||
from pathlib import Path
|
||||
from urllib.parse import urlparse, unquote
|
||||
|
||||
name = ""
|
||||
if uri:
|
||||
try:
|
||||
name = Path(unquote(urlparse(str(uri)).path or "")).name
|
||||
except (ValueError, TypeError):
|
||||
name = ""
|
||||
pass
|
||||
# Strip control chars (hostile URIs could inject newlines/ANSI into the filename and
|
||||
# transcript marker) and cap length, preserving the extension.
|
||||
name = _re.sub(r"[\x00-\x1f\x7f]", "", name).strip()
|
||||
if len(name) > 150:
|
||||
stem, dot, ext = name.rpartition(".")
|
||||
if dot and 0 < len(ext) <= 12:
|
||||
name = stem[: 150 - len(ext) - 1] + "." + ext
|
||||
else:
|
||||
name = name[:150]
|
||||
name = stem[: 150 - len(ext) - 1] + "." + ext if dot and 0 < len(ext) <= 12 else name[:150]
|
||||
if not name or name in {".", ".."}:
|
||||
ext = mimetypes.guess_extension(_base_mime(mime_type)) or ".bin"
|
||||
name = f"resource{ext}"
|
||||
@@ -175,22 +163,15 @@ def _render_mcp_resource_block(block, server_name: str = "") -> str:
|
||||
are only readable via the originating session). "" for non-resource blocks; failures are
|
||||
reported inline rather than silently dropped."""
|
||||
block_type = getattr(block, "type", "")
|
||||
|
||||
if block_type == "resource_link" or (
|
||||
hasattr(block, "uri") and not hasattr(block, "resource") and block_type != "text"):
|
||||
if block_type == "resource_link" or (hasattr(block, "uri") and not hasattr(block, "resource") and block_type != "text"):
|
||||
uri = getattr(block, "uri", None)
|
||||
if not uri:
|
||||
return ""
|
||||
name = getattr(block, "name", "") or ""
|
||||
mime = mcp_field(block, "mime_type", "mimeType", "") or ""
|
||||
details = f"uri={uri}"
|
||||
if name:
|
||||
details += f", name={name}"
|
||||
if mime:
|
||||
details += f", mimeType={mime}"
|
||||
details = f"uri={uri}" + (f", name={name}" if name else "") + (f", mimeType={mime}" if mime else "")
|
||||
reader = mcp_prefixed_tool_name(server_name, "read_resource") if server_name else "the MCP server's read_resource tool"
|
||||
return f"[MCP resource link: {details} — fetch it with {reader}]"
|
||||
|
||||
resource = getattr(block, "resource", None)
|
||||
if resource is None:
|
||||
return ""
|
||||
@@ -200,7 +181,6 @@ def _render_mcp_resource_block(block, server_name: str = "") -> str:
|
||||
blob = getattr(resource, "blob", None)
|
||||
if blob is None:
|
||||
return ""
|
||||
|
||||
uri = str(getattr(resource, "uri", "") or "")
|
||||
mime = str(mcp_field(resource, "mime_type", "mimeType", "") or "")
|
||||
raw_bytes, err = _decode_block_b64(
|
||||
|
||||
+35
-54
@@ -17,24 +17,6 @@ logger = logging.getLogger("tools.mcp_tool")
|
||||
_KEEPALIVE_RPC_TIMEOUT = 30.0
|
||||
|
||||
|
||||
def _stdio_children_dead_impl(pids, is_http: bool) -> bool:
|
||||
"""True when every pid has exited. Best-effort: False (unknown → don't fail fast) for HTTP,
|
||||
no captured PIDs, missing psutil, or a failed probe."""
|
||||
if not pids or is_http:
|
||||
return False
|
||||
try:
|
||||
import psutil
|
||||
except ImportError:
|
||||
return False
|
||||
for pid in pids:
|
||||
try:
|
||||
if psutil.pid_exists(pid): # handles Windows without signal-permission noise
|
||||
return False
|
||||
except Exception:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
class MCPServerHealthMixin:
|
||||
"""Methods of :class:`tools.mcp_tool.MCPServerTask` (mixed in; relies on its attributes)."""
|
||||
|
||||
@@ -85,15 +67,15 @@ class MCPServerHealthMixin:
|
||||
|
||||
# -- notifications / logs --
|
||||
|
||||
async def _refresh_tools_task(self):
|
||||
try:
|
||||
await self._refresh_tools()
|
||||
except Exception:
|
||||
logger.exception("MCP server '%s': dynamic tool refresh failed", self.name)
|
||||
|
||||
def _schedule_tools_refresh(self) -> asyncio.Task:
|
||||
"""Schedule a background tool refresh and keep it strongly referenced."""
|
||||
task = asyncio.create_task(self._refresh_tools_task())
|
||||
"""Schedule a background tool refresh (failures logged) and keep it strongly referenced."""
|
||||
async def _run():
|
||||
try:
|
||||
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
|
||||
@@ -129,19 +111,15 @@ class MCPServerHealthMixin:
|
||||
return
|
||||
if not (_core._MCP_NOTIFICATION_TYPES and isinstance(message, _core.ServerNotification)):
|
||||
return
|
||||
# mcp 2.0 made ServerNotification a plain union (payload IS the message) instead
|
||||
# of a RootModel (payload under ``.root``). ``isinstance`` accepts both; only the
|
||||
# unwrap differs — without it ``.root`` raises into the catch-all and refreshes stop.
|
||||
# mcp 2.0 made ServerNotification a plain union (payload IS the message) instead of
|
||||
# a RootModel (payload under ``.root``); without this unwrap refreshes silently stop.
|
||||
payload = getattr(message, "root", message)
|
||||
if isinstance(payload, _core.ToolListChangedNotification):
|
||||
logger.info("MCP server '%s': received tools/list_changed notification", self.name)
|
||||
# Refresh in a separate task: some servers emit list_changed right after
|
||||
# initialize while another request is in flight, and refreshing synchronously
|
||||
# inside the handler can wedge the stdio JSON-RPC stream.
|
||||
# Separate task: refreshing synchronously inside the handler can wedge the stdio
|
||||
# JSON-RPC stream when list_changed arrives while another request is in flight.
|
||||
self._schedule_tools_refresh()
|
||||
# Yield one tick so short-lived notification contexts (and tests) can observe
|
||||
# the scheduled refresh.
|
||||
await asyncio.sleep(0)
|
||||
await asyncio.sleep(0) # one tick so short-lived contexts (and tests) observe it
|
||||
elif isinstance(payload, _core.PromptListChangedNotification):
|
||||
logger.debug("MCP server '%s': prompts/list_changed (ignored)", self.name)
|
||||
elif isinstance(payload, _core.ResourceListChangedNotification):
|
||||
@@ -154,7 +132,6 @@ class MCPServerHealthMixin:
|
||||
"""Deregister *tool_names* this server's toolset still owns. Never removes a colliding
|
||||
name currently 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:
|
||||
@@ -172,13 +149,11 @@ class MCPServerHealthMixin:
|
||||
old_tool_names = set(self._registered_tool_names)
|
||||
async with self._rpc_lock:
|
||||
new_mcp_tools = await _core._paginate_full_list(self.session.list_tools, "tools", self.name)
|
||||
# Remove only stale names first — no nuke-and-repave: live agent turns may hold
|
||||
# tool-call IDs pointing at existing handlers, and in-place replacement avoids
|
||||
# transient "tool not connected" races.
|
||||
# Remove only stale names first — no nuke-and-repave: live turns may hold tool-call
|
||||
# IDs pointing at existing handlers; in-place replacement avoids "not connected" races.
|
||||
self._deregister_owned(old_tool_names - {mcp_prefixed_tool_name(self.name, tool.name) for tool in new_mcp_tools})
|
||||
# Re-register; the helper may skip names ambiguous after normalization. A raw name
|
||||
# can become ambiguous without changing its normalized name, so the pre-pass misses
|
||||
# it: drop any old entry the final collision-checked registration no longer owns.
|
||||
# Re-register; a raw name can become ambiguous after normalization without changing
|
||||
# its normalized name, so also drop old entries the final registration no longer owns.
|
||||
self._tools = new_mcp_tools
|
||||
registered_names = _core._register_server_tools(self.name, self, self._config)
|
||||
self._deregister_owned(old_tool_names - set(registered_names))
|
||||
@@ -197,9 +172,8 @@ class MCPServerHealthMixin:
|
||||
|
||||
async def _keepalive_probe(self) -> None:
|
||||
"""Exercise the session; raise on a genuine connection failure. ``ping`` first (cheap,
|
||||
OPTIONAL utility). On -32601 latch ``_ping_unsupported`` and fall back to ``list_tools``
|
||||
when the server advertises tools; otherwise the -32601 propagates (no liveness primitive
|
||||
left). The latch resets on each fresh transport connection."""
|
||||
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."""
|
||||
if not self._ping_unsupported:
|
||||
try:
|
||||
await asyncio.wait_for(self.session.send_ping(), timeout=_KEEPALIVE_RPC_TIMEOUT)
|
||||
@@ -212,9 +186,8 @@ class MCPServerHealthMixin:
|
||||
logger.info("MCP server '%s': does not implement the optional 'ping' utility (-32601); "
|
||||
"using 'list_tools' for keepalive on this connection.", self.name)
|
||||
elif isinstance(exc, (TimeoutError, asyncio.TimeoutError)) and self._advertises_tools():
|
||||
# A server that silently drops ping looks like a dead transport. Confirm with
|
||||
# list_tools before declaring it dead; if that also fails, propagate the
|
||||
# original failure.
|
||||
# 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)
|
||||
except Exception:
|
||||
@@ -282,11 +255,10 @@ class MCPServerHealthMixin:
|
||||
return True
|
||||
|
||||
def _fail_inflight_calls(self, reason: str) -> None:
|
||||
"""Cancel every in-flight RPC on this connection. Called from lifecycle exits BEFORE the
|
||||
transport unwinds: the SDK does not always fail pending requests when streams close, so a
|
||||
call would otherwise wait out the full tool timeout. Cancelling anything flags
|
||||
``_teardown_race`` so run() treats the next reconnect as recovery rather than charging
|
||||
the rapid-drop budget."""
|
||||
"""Cancel every in-flight RPC BEFORE the transport unwinds: the SDK does not always fail
|
||||
pending requests when streams close, so a call would otherwise wait out the full tool
|
||||
timeout. Cancelling anything flags ``_teardown_race`` so run() treats the next reconnect
|
||||
as recovery rather than charging the rapid-drop budget."""
|
||||
victims = [t for t in self._inflight_tasks if not t.done()]
|
||||
if not victims:
|
||||
return
|
||||
@@ -297,7 +269,16 @@ class MCPServerHealthMixin:
|
||||
task.cancel()
|
||||
|
||||
def _stdio_children_dead(self) -> bool:
|
||||
return _stdio_children_dead_impl(getattr(self, "_stdio_child_pids", None), self._is_http())
|
||||
"""True when every stdio child we spawned has exited. Best-effort: False (unknown → don't
|
||||
fail fast) for HTTP, no captured PIDs, missing psutil, or a failed probe."""
|
||||
pids = getattr(self, "_stdio_child_pids", None)
|
||||
if not pids or self._is_http():
|
||||
return False
|
||||
try:
|
||||
import psutil
|
||||
return not any(psutil.pid_exists(pid) for pid in pids) # Windows-safe, no signal noise
|
||||
except Exception: # missing psutil or failed probe → unknown → don't fail fast
|
||||
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
|
||||
|
||||
+46
-68
@@ -10,21 +10,16 @@ from tools.mcp_tool_common import _core
|
||||
|
||||
logger = logging.getLogger("tools.mcp_tool")
|
||||
|
||||
# Live stdio MCP children (pid -> server_name), added after connection and
|
||||
# removed on normal shutdown, so they can be force-killed if SDK teardown fails.
|
||||
# Live stdio MCP children (pid -> server_name), added after connection and removed on normal
|
||||
# shutdown, so they can be force-killed if SDK teardown fails.
|
||||
_stdio_pids: Dict[int, str] = {}
|
||||
|
||||
# PIDs that survived their session context exit (SDK teardown failed to kill
|
||||
# them); detected in _run_stdio's finally, reaped by _kill_orphaned_mcp_children().
|
||||
# Kept separate from _stdio_pids so cleanup sweeps never race active sessions.
|
||||
# PIDs that survived their session context exit (detected in _run_stdio's finally, reaped by
|
||||
# _kill_orphaned_mcp_children). Separate from _stdio_pids so sweeps never race active sessions.
|
||||
_orphan_stdio_pids: set = set()
|
||||
_orphan_stdio_pid_servers: Dict[int, str] = {}
|
||||
|
||||
# pid -> pgid captured at spawn. The SDK spawns children with
|
||||
# start_new_session=True (PGID == PID); grandchildren inherit that PGID and
|
||||
# keep it after the direct child exits, so killpg still reaches them. Tracked
|
||||
# separately from _stdio_pids so the PGID survives the child's removal.
|
||||
# Empty on Windows (os.getpgid is POSIX-only).
|
||||
# pid -> pgid captured at spawn. The SDK spawns with start_new_session=True (PGID == PID);
|
||||
# grandchildren keep that PGID after the direct child exits, so killpg still reaches them.
|
||||
# Separate from _stdio_pids so the PGID survives the child's removal. Empty on Windows.
|
||||
_stdio_pgids: Dict[int, int] = {}
|
||||
|
||||
|
||||
@@ -53,53 +48,49 @@ def _snapshot_child_pids() -> set:
|
||||
return set()
|
||||
|
||||
|
||||
# argv markers of non-MCP gateway children that can race into the snapshot
|
||||
# delta during an MCP spawn (defense-in-depth; LSP/slash_worker already use
|
||||
# start_new_session). Matched against argv[1:] because Python/Java children
|
||||
# start with the interpreter path.
|
||||
# argv markers of non-MCP gateway children that can race into the snapshot delta during an
|
||||
# MCP spawn (defense-in-depth; LSP/slash_worker already use start_new_session). Matched against
|
||||
# argv[1:] because Python/Java children start with the interpreter path.
|
||||
_NON_MCP_CHILD_CMDLINE_MARKERS: tuple[str, ...] = (
|
||||
"tui_gateway.slash_worker",
|
||||
"tui_gateway.entry",
|
||||
"-dorg.eclipse.equinox.launcher", # jdtls (legacy arg style)
|
||||
"eclipse.jdt.ls",
|
||||
"org.eclipse.equinox.launcher_")
|
||||
"tui_gateway.slash_worker", "tui_gateway.entry",
|
||||
"-dorg.eclipse.equinox.launcher", "eclipse.jdt.ls", "org.eclipse.equinox.launcher_", # jdtls
|
||||
)
|
||||
|
||||
|
||||
def _filter_mcp_children(pids: set) -> set:
|
||||
"""Drop non-MCP children from a PID snapshot delta. Tracking a stray child in
|
||||
_stdio_pgids is catastrophic if it lacks start_new_session: its pgid can be
|
||||
the TUI parent's, so the shutdown killpg() would kill the TUI itself."""
|
||||
"""Drop non-MCP children from a PID snapshot delta. Tracking a stray child in _stdio_pgids
|
||||
is catastrophic if it lacks start_new_session: its pgid can be the TUI parent's, so the
|
||||
shutdown killpg() would kill the TUI itself."""
|
||||
if not pids:
|
||||
return pids
|
||||
try:
|
||||
import psutil
|
||||
except ImportError:
|
||||
return pids # keep all PIDs (prior behavior)
|
||||
|
||||
def _is_mcp(pid: int) -> bool:
|
||||
kept = set()
|
||||
for pid in pids:
|
||||
try:
|
||||
argv = psutil.Process(pid).cmdline()
|
||||
except (psutil.NoSuchProcess, psutil.AccessDenied, OSError):
|
||||
return False # raced away or zombie — cannot be our fresh server, unsafe to track
|
||||
return not any(marker in arg for arg in argv[1:] for marker in _NON_MCP_CHILD_CMDLINE_MARKERS)
|
||||
|
||||
return {pid for pid in pids if _is_mcp(pid)}
|
||||
continue # raced away or zombie — cannot be our fresh server, unsafe to track
|
||||
if not any(marker in arg for arg in argv[1:] for marker in _NON_MCP_CHILD_CMDLINE_MARKERS):
|
||||
kept.add(pid)
|
||||
return kept
|
||||
|
||||
|
||||
def _clear_connect_cooldowns() -> None:
|
||||
"""Drop connect-retry cooldowns: a restart must re-attempt every server
|
||||
immediately, not honour a stale per-server backoff. Caller holds ``_core._lock``."""
|
||||
"""Drop connect-retry cooldowns: a restart must re-attempt every server immediately, not
|
||||
honour a stale per-server backoff. Caller holds ``_core._lock``."""
|
||||
_core._server_connect_retry_after.clear()
|
||||
_core._server_connect_failures.clear()
|
||||
|
||||
|
||||
def shutdown_mcp_servers(*, scope: Optional[str] = None):
|
||||
"""Close MCP server connections (in parallel) and stop the background loop.
|
||||
Each server Task is signalled to exit its own ``async with`` so the anyio
|
||||
cancel-scope cleanup runs in the Task that opened it. ``scope`` restricts
|
||||
teardown to one multiplexed profile's servers (its ``/reload-mcp`` must not
|
||||
kill other profiles') and leaves the shared loop running if anything else is
|
||||
still connected."""
|
||||
"""Close MCP server connections (in parallel) and stop the background loop. Each server
|
||||
Task is signalled to exit its own ``async with`` so the anyio cancel-scope cleanup runs in
|
||||
the Task that opened it. ``scope`` restricts teardown to one multiplexed profile's servers
|
||||
(its ``/reload-mcp`` must not kill other profiles') and leaves the shared loop running if
|
||||
anything else is still connected."""
|
||||
with _core._lock:
|
||||
selected = [name for name in _core._servers if scope is None or _core._server_scope_keys.get(name) == scope]
|
||||
servers_snapshot = [_core._servers[name] for name in selected]
|
||||
@@ -143,8 +134,8 @@ def _take_reapable_pids(include_active: bool, server_name: Optional[str]) -> tup
|
||||
|
||||
with _core._lock:
|
||||
pids = _owned({opid: _orphan_stdio_pid_servers.get(opid, "orphan") for opid in _orphan_stdio_pids})
|
||||
_orphan_stdio_pids.difference_update(pids)
|
||||
for opid in pids:
|
||||
_orphan_stdio_pids.discard(opid)
|
||||
_orphan_stdio_pid_servers.pop(opid, None)
|
||||
if include_active:
|
||||
active = _owned(_stdio_pids)
|
||||
@@ -156,18 +147,16 @@ def _take_reapable_pids(include_active: bool, server_name: Optional[str]) -> tup
|
||||
|
||||
|
||||
def _signal_mcp_process(pid: int, sig: int, server_name: str, pgid: Optional[int], my_pgid: Optional[int]) -> None:
|
||||
"""SIGTERM/SIGKILL via the spawn-time pgroup on POSIX (reaches reparented
|
||||
grandchildren), falling back to a per-pid signal."""
|
||||
"""SIGTERM/SIGKILL via the spawn-time pgroup on POSIX (reaches reparented grandchildren),
|
||||
falling back to a per-pid signal."""
|
||||
killpg = getattr(os, "killpg", None)
|
||||
if pgid is not None and killpg is not None:
|
||||
if my_pgid is not None and pgid == my_pgid:
|
||||
# Child shares the gateway's pgroup: killpg would kill the gateway too, so use
|
||||
# per-pid kill. Warn because per-pid kill can't reach grandchildren in this group.
|
||||
logger.warning(
|
||||
"MCP server '%s' pgid %d matches gateway pgid; skipping "
|
||||
"killpg to avoid self-kill and using per-pid kill — any "
|
||||
"grandchildren in this group may not be reaped",
|
||||
server_name, pgid)
|
||||
logger.warning("MCP server '%s' pgid %d matches gateway pgid; skipping "
|
||||
"killpg to avoid self-kill and using per-pid kill — any "
|
||||
"grandchildren in this group may not be reaped", server_name, pgid)
|
||||
else:
|
||||
try:
|
||||
killpg(pgid, sig)
|
||||
@@ -183,14 +172,12 @@ def _signal_mcp_process(pid: int, sig: int, server_name: str, pgid: Optional[int
|
||||
|
||||
|
||||
def _kill_orphaned_mcp_children(include_active: bool = False, server_name: Optional[str] = None) -> None:
|
||||
"""Best-effort reap of stdio MCP subprocesses: SIGTERM, wait 2s, SIGKILL survivors.
|
||||
By default only ``_orphan_stdio_pids`` (PIDs that outlived their session
|
||||
context) are reaped so concurrent cron jobs / live sessions are untouched;
|
||||
``include_active=True`` also kills every ``_stdio_pids`` entry and is only
|
||||
for final shutdown after the MCP loop has stopped. ``server_name`` limits the
|
||||
sweep to one server (stdio reconnects cleaning up their old transport)."""
|
||||
"""Best-effort reap of stdio MCP subprocesses: SIGTERM, wait 2s, SIGKILL survivors. By
|
||||
default only ``_orphan_stdio_pids`` are reaped so concurrent cron jobs / live sessions are
|
||||
untouched; ``include_active=True`` also kills every ``_stdio_pids`` entry and is only for
|
||||
final shutdown after the MCP loop has stopped. ``server_name`` limits the sweep to one
|
||||
server (stdio reconnects cleaning up their old transport)."""
|
||||
import signal as _signal
|
||||
|
||||
pids, pgids = _take_reapable_pids(include_active, server_name)
|
||||
if not pids: # skip the 2s sleep every MCP-free shutdown would otherwise pay
|
||||
return
|
||||
@@ -203,17 +190,13 @@ def _kill_orphaned_mcp_children(include_active: bool = False, server_name: Optio
|
||||
for pid, owner in pids.items():
|
||||
_signal_mcp_process(pid, _signal.SIGTERM, owner, pgids.get(pid), my_pgid)
|
||||
logger.debug("Sent SIGTERM to orphaned MCP process %d (%s)", pid, owner)
|
||||
|
||||
time.sleep(2)
|
||||
|
||||
sigkill = getattr(_signal, "SIGKILL", _signal.SIGTERM)
|
||||
# ``os.kill(pid, 0)`` is NOT a no-op on Windows; use the portable check.
|
||||
from gateway.status import _pid_exists
|
||||
from gateway.status import _pid_exists # ``os.kill(pid, 0)`` is NOT a no-op on Windows
|
||||
for pid, owner in pids.items():
|
||||
if not _pid_exists(pid):
|
||||
continue # exited after SIGTERM
|
||||
_signal_mcp_process(pid, sigkill, owner, pgids.get(pid), my_pgid)
|
||||
logger.warning("Force-killed MCP process %d (%s) after SIGTERM timeout", pid, owner)
|
||||
if _pid_exists(pid): # survived SIGTERM
|
||||
_signal_mcp_process(pid, sigkill, owner, pgids.get(pid), my_pgid)
|
||||
logger.warning("Force-killed MCP process %d (%s) after SIGTERM timeout", pid, owner)
|
||||
|
||||
|
||||
def _stop_mcp_loop_if_idle() -> bool:
|
||||
@@ -239,13 +222,8 @@ async def _drain_mcp_loop_tasks(*, timeout: Optional[float] = None) -> None:
|
||||
task.cancel()
|
||||
done, still_pending = await asyncio.wait(pending, timeout=timeout)
|
||||
for task in done:
|
||||
try:
|
||||
if not task.cancelled():
|
||||
task.exception()
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
except Exception as exc:
|
||||
logger.debug("Pending MCP loop task ended during shutdown: %s", exc)
|
||||
if not task.cancelled():
|
||||
task.exception() # mark retrieved so asyncio doesn't warn "exception was never retrieved"
|
||||
if still_pending:
|
||||
logger.warning("%d MCP loop task(s) still pending after %.1fs drain", len(still_pending), timeout)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user