refactor(tools): fold legacy HTTP transport into streamable path, tighten MCP classifier/sampling bodies
This commit is contained in:
@@ -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).
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
|
||||
Reference in New Issue
Block a user