refactor(tools): fold legacy HTTP transport into streamable path, tighten MCP classifier/sampling bodies

This commit is contained in:
Teknium
2026-09-02 23:35:14 -07:00
parent 6a7d04e89d
commit 672cd01c7c
4 changed files with 77 additions and 108 deletions

View File

@@ -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.<name>.command to an absolute path and include "
"that directory in mcp_servers.<name>.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.<name>.command to an absolute path and include "
"that directory in mcp_servers.<name>.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).

View File

@@ -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)

View File

@@ -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)

View File

@@ -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))