From 672cd01c7c0ec0811c8b63bbd90312a052a9d79a Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 23:35:14 -0700 Subject: [PATCH] refactor(tools): fold legacy HTTP transport into streamable path, tighten MCP classifier/sampling bodies --- tools/mcp_tool_errors.py | 76 ++++++++++++++++--------------------- tools/mcp_tool_loop.py | 24 +++++------- tools/mcp_tool_sampling.py | 28 ++++++-------- tools/mcp_tool_transport.py | 57 ++++++++++++---------------- 4 files changed, 77 insertions(+), 108 deletions(-) diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py index f2ac066da1..5f3d95853f 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -136,22 +136,20 @@ def _resolve_client_cert(server_name: str, config: dict): raise FileNotFoundError(f"{prefix}{label} not found at {expanded!r}") return expanded - if isinstance(raw_cert, (list, tuple)): - 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") - 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]")) - if len(raw_cert) == 2: - return pair - if not isinstance(raw_cert[2], str): - raise ValueError(f"{prefix}client_cert[2] (key passphrase) must be a string") - return (*pair, raw_cert[2]) - cert_path = _expand(raw_cert, "client_cert") + 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: - return (cert_path, _expand(raw_key, "client_key")) - return cert_path # single combined PEM (cert + key) + raise ValueError(f"{prefix}specify either client_cert as a list [cert, key] OR " + f"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]")) + if len(raw_cert) == 2: + return pair + if not isinstance(raw_cert[2], str): + raise ValueError(f"{prefix}client_cert[2] (key passphrase) must be a string") + return (*pair, raw_cert[2]) def _resolve_identity_header(server_name: str, config: dict): @@ -193,8 +191,8 @@ def _apply_identity_header(server_name: str, config: dict, headers: dict) -> dic if any(key.lower() == name.lower() for key in headers): logger.debug("MCP server '%s': identity_header '%s' already set via explicit " "headers config — keeping the explicit value", server_name, name) - return headers - headers[name] = value + else: + headers[name] = value return headers @@ -206,10 +204,8 @@ def _make_redirect_header_stripper(original_url, *, strict: bool = False, origin = (original_url.scheme, original_url.host, original_url.port) async def _strip_on_cross_origin_redirect(response): - if not (response.is_redirect and response.next_request): - return - target = response.next_request.url - if (target.scheme, target.host, target.port) == origin: + target = response.next_request.url if response.is_redirect and response.next_request else None + if target is None or (target.scheme, target.host, target.port) == origin: return headers = response.next_request.headers headers.pop("authorization", None) @@ -225,9 +221,7 @@ def _make_redirect_header_stripper(original_url, *, strict: bool = False, def _exc_children(exc: BaseException) -> List[BaseException]: """Sub-exceptions of a group, else ``__cause__``/``__context__`` when they are exceptions.""" nested = getattr(exc, "exceptions", None) - if nested: - return list(nested) - return [c for c in (exc.__cause__, exc.__context__) if isinstance(c, BaseException)] + return list(nested) if nested else [c for c in (exc.__cause__, exc.__context__) if isinstance(c, BaseException)] def _format_connect_error(exc: BaseException) -> str: @@ -244,21 +238,18 @@ def _format_connect_error(exc: BaseException) -> str: def _flatten_messages(current: BaseException) -> List[str]: # A group's own str() is opaque — only its children speak. text = "" if getattr(current, "exceptions", None) else str(current).strip() - messages = [text] if text else [] - for child in _exc_children(current): - messages.extend(_flatten_messages(child)) + 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 missing: - message = f"missing executable '{missing}'" - if os.path.basename(missing) in {"npx", "npm", "node"}: - message += (" (ensure Node.js is installed and PATH includes its bin directory, " - "or set mcp_servers..command to an absolute path and include " - "that directory in mcp_servers..env.PATH)") - return _sanitize_error(message) - deduped = list(dict.fromkeys(_flatten_messages(exc))) - return _sanitize_error("; ".join(deduped[:3])) + if not missing: + return _sanitize_error("; ".join(list(dict.fromkeys(_flatten_messages(exc)))[:3])) + message = f"missing executable '{missing}'" + if os.path.basename(missing) in {"npx", "npm", "node"}: + message += (" (ensure Node.js is installed and PATH includes its bin directory, " + "or set mcp_servers..command to an absolute path and include " + "that directory in mcp_servers..env.PATH)") + return _sanitize_error(message) # Lazily-built caches so this module imports without the SDK OAuth module. @@ -292,11 +283,10 @@ def _get_auth_error_types() -> tuple: (which still need the 401 check in :func:`_is_auth_error`).""" global _AUTH_ERROR_TYPES if not _AUTH_ERROR_TYPES: - _AUTH_ERROR_TYPES = tuple( - _optional_types("mcp.client.auth", "OAuthFlowError", "OAuthTokenError") - + _optional_types("mcp.client.auth", "UnauthorizedError") # older SDKs - + _optional_types("tools.mcp_oauth", "OAuthNonInteractiveError") - + list(_http_status_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()) return _AUTH_ERROR_TYPES @@ -304,9 +294,7 @@ 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()): return False - if isinstance(exc, _http_status_error_types()): - return getattr(exc.response, "status_code", None) == 401 - return True + return getattr(exc.response, "status_code", None) == 401 if isinstance(exc, _http_status_error_types()) else True # Lower-cased substrings meaning the transport session expired / was GC'd (OAuth token still valid). diff --git a/tools/mcp_tool_loop.py b/tools/mcp_tool_loop.py index a2c175fd88..4b86022c3b 100644 --- a/tools/mcp_tool_loop.py +++ b/tools/mcp_tool_loop.py @@ -101,7 +101,7 @@ def _wrap_with_home_override(coro: "Coroutine") -> "Coroutine": from hermes_constants import get_hermes_home_override, reset_hermes_home_override, set_hermes_home_override home_override = get_hermes_home_override() except Exception: - return coro + home_override = None if not home_override: return coro @@ -121,7 +121,7 @@ def _wrap_with_dashboard_oauth_flow(coro): from tools.mcp_dashboard_oauth import dashboard_oauth_flow, get_dashboard_oauth_flow flow = get_dashboard_oauth_flow() except Exception: - return coro + flow = None if flow is None: return coro @@ -163,16 +163,13 @@ def _run_on_mcp_loop(coro_or_factory, timeout: float = 30): if is_interrupted(): future.cancel() raise InterruptedError("User sent a new message") - wait_timeout = 0.1 - if deadline is not None: - remaining = deadline - time.monotonic() - if remaining <= 0: - future.cancel() - raise TimeoutError(f"MCP call timed out after {time.monotonic() - start_time:.1f}s " - f"(configured timeout: {float(timeout):.1f}s)") - wait_timeout = min(wait_timeout, remaining) + remaining = 0.1 if deadline is None else deadline - time.monotonic() + if remaining <= 0: + future.cancel() + raise TimeoutError(f"MCP call timed out after {time.monotonic() - start_time:.1f}s " + f"(configured timeout: {float(timeout):.1f}s)") try: - return future.result(timeout=wait_timeout) + return future.result(timeout=min(0.1, remaining)) except concurrent.futures.TimeoutError: # Aliases builtin TimeoutError, so it also fires for the coroutine's own timeout: a done # future must yield its outcome. @@ -270,7 +267,7 @@ def _stop_mcp_loop(*, only_if_idle: bool = False) -> bool: if loop is not None: # Drain before stopping: tasks still suspended when the loop closes get resumed by the GC # against a closed loop. shutdown_mcp_servers only reaps _servers; everything else ends here. - stop_owned_by_loop = False + future = None if loop.is_running(): from agent.async_utils import safe_schedule_threadsafe @@ -278,7 +275,6 @@ def _stop_mcp_loop(*, only_if_idle: bool = False) -> bool: _core._drain_and_stop_mcp_loop(), loop, logger=logger, log_message="MCP loop drain: failed to schedule", log_level=logging.WARNING) if future is not None: - stop_owned_by_loop = True try: future.result(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT + 1) except TimeoutError: @@ -291,7 +287,7 @@ def _stop_mcp_loop(*, only_if_idle: bool = False) -> bool: loop.run_until_complete(_core._drain_mcp_loop_tasks(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT)) except BaseException as exc: logger.warning("Error draining stopped MCP loop tasks: %s", exc) - if not stop_owned_by_loop and loop.is_running(): + if future is None and loop.is_running(): # drain-and-stop wasn't scheduled: stop it ourselves loop.call_soon_threadsafe(loop.stop) if thread is not None: thread.join(timeout=5) diff --git a/tools/mcp_tool_sampling.py b/tools/mcp_tool_sampling.py index 3c12586bbc..f58867c364 100644 --- a/tools/mcp_tool_sampling.py +++ b/tools/mcp_tool_sampling.py @@ -56,7 +56,6 @@ def _convert_sampling_message(msg) -> List[dict]: tool_results = [b for b in blocks if _tool_use_id(b) is not _MISSING] tool_uses = [b for b in blocks if _is_tool_use(b) and _tool_use_id(b) is _MISSING] content_blocks = [b for b in blocks if _tool_use_id(b) is _MISSING and not _is_tool_use(b)] - out = [{"role": "tool", "tool_call_id": _tool_use_id(tr), "content": _tool_result_text(tr)} for tr in tool_results] if tool_uses: msg_dict: dict = {"role": msg.role, "tool_calls": [_tool_call_dict(tu, i) for i, tu in enumerate(tool_uses)]} @@ -64,13 +63,12 @@ def _convert_sampling_message(msg) -> List[dict]: if text_parts: msg_dict["content"] = "\n".join(text_parts) out.append(msg_dict) + elif len(content_blocks) == 1 and hasattr(content_blocks[0], "text"): + out.append({"role": msg.role, "content": content_blocks[0].text}) elif content_blocks: - if len(content_blocks) == 1 and hasattr(content_blocks[0], "text"): - out.append({"role": msg.role, "content": content_blocks[0].text}) - else: - parts = [p for p in map(_content_part, content_blocks) if p is not None] - if parts: - out.append({"role": msg.role, "content": parts}) + parts = [p for p in map(_content_part, content_blocks) if p is not None] + if parts: + out.append({"role": msg.role, "content": parts}) return out @@ -159,10 +157,9 @@ class SamplingHandler: self._tool_loop_count = 0 return self._error( f"Tool loop limit exceeded for server '{self.server_name}' (max {self.max_tool_rounds} rounds)") - content_blocks = [ - _core.ToolUseContent(type="tool_use", id=tc.id, name=tc.function.name, - input=_parse_tool_call_arguments(self.server_name, tc.function.arguments)) - for tc in choice.message.tool_calls] + content_blocks = [_core.ToolUseContent(type="tool_use", id=tc.id, name=tc.function.name, + input=_parse_tool_call_arguments(self.server_name, tc.function.arguments)) + for tc in choice.message.tool_calls] self._log_response(response, ", tool_calls=%d", len(content_blocks)) return _core.CreateMessageResultWithTools( role="assistant", content=content_blocks, model=response.model, stopReason="toolUse") @@ -234,8 +231,7 @@ class SamplingHandler: choice = response.choices[0] self.metrics["requests"] += 1 total_tokens = _response_total_tokens(response, 0) - if isinstance(total_tokens, int): - self.metrics["tokens_used"] += total_tokens + self.metrics["tokens_used"] += total_tokens if isinstance(total_tokens, int) else 0 if choice.finish_reason == "tool_calls" and getattr(choice.message, "tool_calls", None): return self._build_tool_use_result(choice, response) return self._build_text_result(choice, response) @@ -249,10 +245,8 @@ def _format_elicitation_schema_summary(schema: dict, server_name: str) -> str: lines = [f"Fields requested by MCP server '{server_name}':"] for field_name, field_spec in props.items(): spec = field_spec if isinstance(field_spec, dict) else {} - field_type = str(spec.get("type", "") or "") - field_desc = str(spec.get("description", "") or "") - suffix = f" ({field_type})" if field_type else "" - lines.append(f" - {field_name}{suffix}: {field_desc}" if field_desc else f" - {field_name}{suffix}") + field_type, field_desc = str(spec.get("type", "") or ""), str(spec.get("description", "") or "") + lines.append(f" - {field_name}" + (f" ({field_type})" if field_type else "") + (f": {field_desc}" if field_desc else "")) return "\n".join(lines) diff --git a/tools/mcp_tool_transport.py b/tools/mcp_tool_transport.py index ce65b8b869..b3841236b4 100644 --- a/tools/mcp_tool_transport.py +++ b/tools/mcp_tool_transport.py @@ -98,34 +98,31 @@ class MCPServerTransportMixin: 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.""" - def initialize(): - return asyncio.wait_for(session.initialize(), timeout=connect_timeout) - - def discover(): - return asyncio.wait_for(session.discover(), timeout=connect_timeout) + def call(method: str): + return asyncio.wait_for(getattr(session, method)(), timeout=connect_timeout) async def attempt(primary, fallback, should_fallback, log_fmt, *log_extra): try: - return await primary() + return await call(primary) except Exception as exc: if isinstance(exc, asyncio.TimeoutError) or not should_fallback(exc): raise logger.info(log_fmt, self.name, exc, *log_extra) - return await fallback() + 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, + return await attempt("discover", "initialize", lambda exc: True, "MCP server '%s': server/discover rejected (%s) despite " "protocol=%s — falling back to the legacy handshake", mode) if mode in ("legacy", "handshake"): - return await initialize() + return await call("initialize") if mode != "auto": logger.warning("MCP server '%s': unknown protocol=%r — treating as 'auto' " "(valid: auto, stateless, legacy)", self.name, mode) # 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"), + "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)") @@ -173,8 +170,7 @@ class MCPServerTransportMixin: """Ledger the freshly spawned stdio children (pids, pgids, machine spawn ledger).""" new_pgids = _capture_pgids(new_pids) with _core._lock: - for _pid in new_pids: - _stdio_pids[_pid] = self.name + _stdio_pids.update(dict.fromkeys(new_pids, self.name)) _stdio_pgids.update(new_pgids) # Machine spawn ledger (startup sweeps reap orphans after an unclean exit); best-effort. for _pid in new_pids: @@ -349,10 +345,10 @@ class MCPServerTransportMixin: _httpx_mod = _core.sdk_httpx() def _mcp_http_client_factory(headers=None, timeout=None, auth=None): - kwargs: dict = {"follow_redirects": True, "verify": ssl_verify, - "timeout": timeout if timeout is not None else _httpx_mod.Timeout(30.0, read=300.0)} - kwargs.update({k: v for k, v in (("headers", headers), ("auth", auth), ("cert", client_cert)) if v is not None}) - return _httpx_mod.AsyncClient(**kwargs) + return _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}) sse_kwargs["httpx_client_factory"] = _mcp_http_client_factory return _core.sse_client(**sse_kwargs) @@ -360,9 +356,16 @@ class MCPServerTransportMixin: def _streamable_http_transport(self, url: str, headers: dict, connect_timeout: float, ssl_verify, client_cert, oauth_auth, strict_cfg_headers: bool, configured_header_names: set): - """Streamable HTTP context manager (mcp >= 1.24.0: caller-owned httpx client).""" + """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: - return self._legacy_http_transport(url, headers, connect_timeout, ssl_verify, oauth_auth, strict_cfg_headers) + 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. httpx = _core.sdk_httpx() @@ -382,17 +385,6 @@ class MCPServerTransportMixin: return _owned_client_streams() - def _legacy_http_transport(self, url: str, headers: dict, connect_timeout: float, - ssl_verify, oauth_auth, strict_cfg_headers: bool): - """Deprecated API (mcp < 1.24.0): the SDK owns the httpx client.""" - 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 {})) - async def _run_http(self, config: dict): """Run the server using HTTP/StreamableHTTP (or SSE) transport.""" _core._ensure_mcp_sdk() @@ -417,12 +409,11 @@ class MCPServerTransportMixin: ssl_verify = config.get("ssl_verify", True) client_cert = _resolve_client_cert(self.name, config) oauth_auth = self._build_oauth_auth(url, config) + common = (url, headers, connect_timeout, ssl_verify, client_cert, oauth_auth, strict_cfg_headers) if config.get("transport") == "sse": - transport = self._sse_transport(url, headers, connect_timeout, ssl_verify, client_cert, oauth_auth, strict_cfg_headers) - label = "SSE" + transport, label = self._sse_transport(*common), "SSE" else: - transport = self._streamable_http_transport(url, headers, connect_timeout, ssl_verify, client_cert, - oauth_auth, strict_cfg_headers, configured_header_names) + transport = self._streamable_http_transport(*common, configured_header_names) label = "HTTP" if _core._MCP_NEW_HTTP else "legacy HTTP" return await self._serve_transport(transport, label, float(connect_timeout))