refactor(mcp): unify live/cache registration into one candidate loop; phase-split sampling and elicitation handlers

This commit is contained in:
Teknium
2026-09-02 15:38:42 -07:00
parent 88b74d6ef0
commit d6ec8c3018
2 changed files with 466 additions and 648 deletions

View File

@@ -1,13 +1,16 @@
"""Registering a connected (or schema-cached) MCP server's tools into the tool
registry: include/exclude filtering, trust-tier metadata capture, utility-tool
selection, name-collision resolution and the schema-cache write-through."""
selection, name-collision resolution and the schema-cache write-through.
Both entry points (``_register_server_tools`` for a live server,
``_register_from_cache_sync`` for a lazy cached manifest) build a list of
``_Candidate`` records and hand them to the single ``_register_candidates`` loop."""
import logging
from types import SimpleNamespace
from typing import TYPE_CHECKING, Any, Callable, Dict, List
from tools.mcp_tool_common import _parse_boolish, _core
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Optional
from tools.mcp_tool_common import _parse_boolish, _core, _resolve_tool_timeout
from tools.mcp_tool_handlers import _make_check_fn, _make_get_prompt_handler, _make_list_prompts_handler, _make_list_resources_handler, _make_read_resource_handler
from tools.mcp_tool_common import _resolve_tool_timeout
from tools.mcp_tool_schema import _UTILITY_CAPABILITY_ATTRS, _UTILITY_CAPABILITY_METHODS, _build_utility_schemas, _normalize_name_filter, matches_name_filter
if TYPE_CHECKING: # pragma: no cover
@@ -15,6 +18,16 @@ if TYPE_CHECKING: # pragma: no cover
logger = logging.getLogger("tools.mcp_tool")
_UTILITY_ORIGIN_PREFIX = "generated utility "
# Utility tool key -> handler factory; each takes (server_name, tool_timeout).
_UTILITY_HANDLER_FACTORIES = {
"list_resources": _make_list_resources_handler,
"read_resource": _make_read_resource_handler,
"list_prompts": _make_list_prompts_handler,
"get_prompt": _make_get_prompt_handler,
}
def _normalize_server_trust(value: Any) -> str:
"""Normalize a config ``trust`` value: None -> ``full`` (backward-compatible
@@ -23,10 +36,8 @@ def _normalize_server_trust(value: Any) -> str:
if value is None:
return _core._TRUST_FULL
text = str(value).strip().lower()
if text == _core._TRUST_FULL:
return _core._TRUST_FULL
if text == _core._TRUST_UNTRUSTED:
return _core._TRUST_UNTRUSTED
if text in (_core._TRUST_FULL, _core._TRUST_UNTRUSTED):
return text
logger.warning(
"MCP trust: unrecognized trust value %r — treating as 'untrusted' "
"(valid values: full, untrusted)", value,
@@ -42,20 +53,16 @@ def _annotation_read_only_hint(mcp_tool: Any) -> bool:
if annotations is None:
return False
if isinstance(annotations, dict):
hint = annotations.get("readOnlyHint")
else:
hint = getattr(annotations, "readOnlyHint", None)
return hint is True
return annotations.get("readOnlyHint") is True
return getattr(annotations, "readOnlyHint", None) is True
def _record_tool_trust_metadata(
server_name: str, config: dict, tools: List[Any]
) -> None:
"""Capture per-server trust and per-tool readOnlyHint at discovery."""
def _record_tool_trust_metadata(server_name: str, config: dict, tools: List[Any]) -> None:
"""Capture per-server trust and per-tool readOnlyHint at discovery (the
security boundary: the call-time gate classifies from data we control,
never re-read server-supplied state)."""
with _core._lock:
_core._server_trust_levels[server_name] = _normalize_server_trust(
(config or {}).get("trust")
)
_core._server_trust_levels[server_name] = _normalize_server_trust((config or {}).get("trust"))
hints = _core._tool_read_only_hints.setdefault(server_name, {})
for tool in tools:
name = getattr(tool, "name", None)
@@ -78,51 +85,39 @@ def _forget_mcp_tool_server(tool_name: str) -> None:
def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: dict) -> List[dict]:
"""Select utility schemas based on config and server capabilities."""
tools_filter = config.get("tools") or {}
resources_enabled = _parse_boolish(tools_filter.get("resources"), default=True)
prompts_enabled = _parse_boolish(tools_filter.get("prompts"), default=True)
family_enabled = {
family: _parse_boolish(tools_filter.get(family), default=True)
for family in ("resources", "prompts")
}
# ``initialize_result.capabilities`` is the source of truth: its sub-objects
# are non-None iff the server advertises that request family. The old
# ``hasattr(server.session, ...)`` gate never filtered anything because
# ClientSession defines all four methods on the class.
advertised_caps = None
init_result = getattr(server, "initialize_result", None)
if init_result is not None:
advertised_caps = getattr(init_result, "capabilities", None)
advertised_caps = getattr(init_result, "capabilities", None) if init_result is not None else None
selected: List[dict] = []
for entry in _build_utility_schemas(server_name):
handler_key = entry["handler_key"]
if handler_key in {"list_resources", "read_resource"} and not resources_enabled:
logger.debug("MCP server '%s': skipping utility '%s' (resources disabled)", server_name, handler_key)
family = _UTILITY_CAPABILITY_ATTRS[handler_key]
if not family_enabled[family]:
logger.debug("MCP server '%s': skipping utility '%s' (%s disabled)", server_name, handler_key, family)
continue
if handler_key in {"list_prompts", "get_prompt"} and not prompts_enabled:
logger.debug("MCP server '%s': skipping utility '%s' (prompts disabled)", server_name, handler_key)
continue
if advertised_caps is not None:
cap_attr = _UTILITY_CAPABILITY_ATTRS[handler_key]
if getattr(advertised_caps, cap_attr, None) is None:
if getattr(advertised_caps, family, None) is None:
logger.debug(
"MCP server '%s': skipping utility '%s' "
"(server does not advertise '%s' capability)",
server_name,
handler_key,
cap_attr,
)
continue
else:
# Legacy fallback when initialize_result wasn't captured (test
# fixtures, older paths): register every stub, as before.
required_method = _UTILITY_CAPABILITY_METHODS[handler_key]
if not hasattr(server.session, required_method):
logger.debug(
"MCP server '%s': skipping utility '%s' (session lacks %s)",
server_name,
handler_key,
required_method,
"MCP server '%s': skipping utility '%s' (server does not advertise '%s' capability)",
server_name, handler_key, family,
)
continue
# Legacy fallback when initialize_result wasn't captured (test
# fixtures, older paths): register every stub the session can serve.
elif not hasattr(server.session, _UTILITY_CAPABILITY_METHODS[handler_key]):
logger.debug(
"MCP server '%s': skipping utility '%s' (session lacks %s)",
server_name, handler_key, _UTILITY_CAPABILITY_METHODS[handler_key],
)
continue
selected.append(entry)
return selected
@@ -135,30 +130,19 @@ def _existing_tool_names() -> List[str]:
names.extend(server._registered_tool_names)
continue
for mcp_tool in server._tools:
schema = _core._convert_mcp_schema(server.name, mcp_tool)
names.append(schema["name"])
names.append(_core._convert_mcp_schema(server.name, mcp_tool)["name"])
# Lazy servers registered from the schema cache have no MCPServerTask yet —
# their tools live only in the registry.
with _core._lock:
lazy_names = [
names.extend(
n
for sname, tool_names in _core._lazy_server_tool_names.items()
if sname not in _core._servers
for n in tool_names
]
names.extend(lazy_names)
)
return names
# Utility tool key -> handler factory; each takes (server_name, tool_timeout).
_UTILITY_HANDLER_FACTORIES = {
"list_resources": _make_list_resources_handler,
"read_resource": _make_read_resource_handler,
"list_prompts": _make_list_prompts_handler,
"get_prompt": _make_get_prompt_handler,
}
def _make_tool_filter(name: str, config: dict) -> Callable[[str], bool]:
"""Build the include/exclude predicate for a server's tool names.
@@ -171,9 +155,7 @@ def _make_tool_filter(name: str, config: dict) -> Callable[[str], bool]:
include_raw = tools_filter.get("include")
include_set = _normalize_name_filter(include_raw, f"mcp_servers.{name}.tools.include")
include_active = isinstance(include_raw, (str, list, tuple, set))
exclude_set = _normalize_name_filter(
tools_filter.get("exclude"), f"mcp_servers.{name}.tools.exclude"
)
exclude_set = _normalize_name_filter(tools_filter.get("exclude"), f"mcp_servers.{name}.tools.exclude")
def _should_register(tool_name: str) -> bool:
if include_active:
@@ -185,69 +167,206 @@ def _make_tool_filter(name: str, config: dict) -> Callable[[str], bool]:
return _should_register
def _resolve_name_collisions(name: str, candidates: List[dict]):
class _CachedMCPTool:
"""Minimal stand-in for MCP Tool objects loaded from the schema cache.
Missing/non-dict ``annotations`` (older cache files) fail closed to
write-capable via ``_annotation_read_only_hint``."""
__slots__ = ("name", "description", "inputSchema", "annotations")
def __init__(self, name: str, description: str, inputSchema: dict, annotations: Optional[dict] = None):
self.name = name
self.description = description
self.inputSchema = inputSchema or {}
self.annotations = annotations if isinstance(annotations, dict) else None
@classmethod
def from_cache_dicts(cls, raws: Iterable[Any]) -> List["_CachedMCPTool"]:
"""Cached tool rows -> stand-ins; rows that are not dicts or lack a name are dropped."""
tools = []
for raw in raws:
if not isinstance(raw, dict) or not raw.get("name"):
continue
schema = raw.get("inputSchema")
tools.append(cls(
raw["name"],
raw.get("description") or "",
schema if isinstance(schema, dict) else {},
raw.get("annotations"),
))
return tools
@dataclass
class _Candidate:
"""One registry registration attempt for a server: a native tool or a
generated utility. ``origin`` is the human-readable provenance used in
collision diagnostics."""
registry_name: str
origin: str
schema: dict
handler: Callable
@property
def is_utility(self) -> bool:
return self.origin.startswith(_UTILITY_ORIGIN_PREFIX)
@property
def description(self) -> str:
return self.schema.get("description") or ""
def _tool_candidates(name: str, tools: Iterable[Any], should_register: Callable[[str], bool], tool_timeout) -> List[_Candidate]:
"""Native tools (live SDK objects or ``_CachedMCPTool``) -> candidates.
The description scan runs on BOTH paths: the cache file is user-writable JSON."""
candidates: List[_Candidate] = []
for mcp_tool in tools:
if not should_register(mcp_tool.name):
logger.debug("MCP server '%s': skipping tool '%s' (filtered by config)", name, mcp_tool.name)
continue
_core._scan_mcp_description(name, mcp_tool.name, mcp_tool.description or "")
schema = _core._convert_mcp_schema(name, mcp_tool)
candidates.append(_Candidate(
schema["name"], f"tool {mcp_tool.name!r}", schema,
_core._make_tool_handler(name, mcp_tool.name, tool_timeout),
))
return candidates
def _utility_candidates(name: str, entries: Iterable[Any], tool_timeout) -> List[_Candidate]:
"""``{schema, handler_key}`` rows (live selection or cache) -> candidates;
malformed rows are dropped."""
candidates: List[_Candidate] = []
for raw in entries:
if not isinstance(raw, dict):
continue
schema, handler_key = raw.get("schema"), raw.get("handler_key")
if not isinstance(schema, dict) or handler_key not in _UTILITY_HANDLER_FACTORIES or not schema.get("name"):
continue
candidates.append(_Candidate(
schema["name"], f"{_UTILITY_ORIGIN_PREFIX}{handler_key!r}", schema,
_UTILITY_HANDLER_FACTORIES[handler_key](name, tool_timeout),
))
return candidates
def _resolve_name_collisions(name: str, candidates: List[_Candidate]) -> List[_Candidate]:
"""Preflight registry-name collisions among one server's candidates.
Returns ``(unique_candidates, ambiguous_names, shadowed_utilities)``. Exact
duplicate rows (same name + origin) are dropped silently; a generated
Exact duplicate rows (same name + origin) are dropped silently; a generated
utility that normalizes onto a server-native tool's name is shadowed (the
native tool wins); any other multi-origin collision is ambiguous and every
colliding entry is skipped (fail closed).
colliding entry is skipped (fail closed). Returns the survivors in order.
"""
unique_candidates: List[dict] = []
seen_candidates: set[tuple[str, str]] = set()
unique: List[_Candidate] = []
seen: set[tuple[str, str]] = set()
origins_by_name: Dict[str, set[str]] = {}
for candidate in candidates:
key = (candidate["registry_name"], candidate["origin"])
if key in seen_candidates:
for c in candidates:
if (c.registry_name, c.origin) in seen:
logger.debug(
"MCP server '%s': duplicate registration candidate %s for '%s'; "
"keeping one",
name,
candidate["origin"],
candidate["registry_name"],
"MCP server '%s': duplicate registration candidate %s for '%s'; keeping one",
name, c.origin, c.registry_name,
)
continue
seen_candidates.add(key)
unique_candidates.append(candidate)
origins_by_name.setdefault(candidate["registry_name"], set()).add(
candidate["origin"]
)
seen.add((c.registry_name, c.origin))
unique.append(c)
origins_by_name.setdefault(c.registry_name, set()).add(c.origin)
ambiguous_names: Dict[str, List[str]] = {}
shadowed_utilities: set[tuple[str, str]] = set()
ambiguous: Dict[str, List[str]] = {}
shadowed: set[tuple[str, str]] = set()
for registry_name, origins in origins_by_name.items():
if len(origins) <= 1:
continue
utility_origins = sorted(
o for o in origins if o.startswith("generated utility ")
)
utility_origins = sorted(o for o in origins if o.startswith(_UTILITY_ORIGIN_PREFIX))
native_origins = sorted(origins - set(utility_origins))
if len(native_origins) == 1 and utility_origins:
for util_origin in utility_origins:
shadowed_utilities.add((registry_name, util_origin))
shadowed.update((registry_name, o) for o in utility_origins)
logger.info(
"MCP server '%s': generated utility %s normalizes onto "
"server-native %s — keeping the native tool and dropping the "
"utility (the utility only applies when the server has no such "
"tool of its own)",
name,
", ".join(utility_origins),
native_origins[0],
name, ", ".join(utility_origins), native_origins[0],
)
continue
ambiguous_names[registry_name] = sorted(origins)
ambiguous[registry_name] = sorted(origins)
for registry_name, origins in sorted(ambiguous_names.items()):
for registry_name, origins in sorted(ambiguous.items()):
logger.error(
"MCP server '%s': name normalization collision for '%s' from %s; "
"skipping every colliding entry instead of choosing an arbitrary "
"handler",
name,
registry_name,
", ".join(origins),
"skipping every colliding entry instead of choosing an arbitrary handler",
name, registry_name, ", ".join(origins),
)
return unique_candidates, ambiguous_names, shadowed_utilities
return [
c for c in unique
if c.registry_name not in ambiguous and (c.registry_name, c.origin) not in shadowed
]
def _log_foreign_owner(name: str, c: _Candidate, existing_toolset: str, lazy: bool) -> None:
"""Diagnostics for a candidate whose registry name is already owned by
another toolset (skipped to preserve the existing owner)."""
if lazy:
if not c.is_utility:
logger.warning(
"MCP server '%s' (lazy): cached tool '%s' collides with toolset '%s' — skipping",
name, c.registry_name, existing_toolset,
)
elif existing_toolset.startswith("mcp-"):
logger.error(
"MCP server '%s': %s normalizes to '%s', already owned by MCP toolset '%s' "
"— skipping to preserve the existing owner",
name, c.origin, c.registry_name, existing_toolset,
)
else:
logger.warning(
"MCP server '%s': %s (→ '%s') collides with built-in tool in toolset '%s' "
"— skipping to preserve built-in",
name, c.origin, c.registry_name, existing_toolset,
)
def _register_candidates(
name: str, candidates: List[_Candidate], *, check_fn: Callable, scope: Callable[[], Optional[str]], lazy: bool,
) -> List[str]:
"""Register candidates under toolset ``mcp-{name}``; returns the names that
actually landed. The ownership pre-check is advisory only — multiple
servers connect in parallel, so ``ToolRegistry.register()`` is the atomic
ownership gate and its verdict is re-read after every call."""
from tools.registry import registry
toolset_name = f"mcp-{name}"
registered: List[str] = []
for c in candidates:
existing_toolset = registry.get_toolset_for_tool(c.registry_name)
if existing_toolset and existing_toolset != toolset_name:
_log_foreign_owner(name, c, existing_toolset, lazy)
continue
registry.register(
name=c.registry_name,
toolset=toolset_name,
schema=c.schema,
handler=c.handler,
check_fn=check_fn,
is_async=False,
description=c.description,
scope=scope(),
)
if registry.get_toolset_for_tool(c.registry_name) != toolset_name:
if not lazy:
logger.error(
"MCP server '%s': registration of %s as '%s' was rejected by "
"the registry; skipping provenance/count updates",
name, c.origin, c.registry_name,
)
continue
_core._track_mcp_tool_server(c.registry_name, name)
registered.append(c.registry_name)
if registered:
registry.register_toolset_alias(name, toolset_name)
return registered
def _write_schema_cache(name: str, server: "MCPServerTask", config: dict, should_register) -> None:
@@ -256,7 +375,7 @@ def _write_schema_cache(name: str, server: "MCPServerTask", config: dict, should
try:
from tools.mcp_schema_cache import config_fingerprint, write_cache_entry
tools_payload: List[dict] = []
tools_payload = []
for mcp_tool in server._tools:
if not should_register(mcp_tool.name):
continue
@@ -266,9 +385,7 @@ def _write_schema_cache(name: str, server: "MCPServerTask", config: dict, should
"description": mcp_tool.description or "",
"inputSchema": schema_obj if isinstance(schema_obj, dict) else {},
# Persisted so the lazy path trust-gates identically next startup.
"annotations": {
"readOnlyHint": _annotation_read_only_hint(mcp_tool),
},
"annotations": {"readOnlyHint": _annotation_read_only_hint(mcp_tool)},
})
utility_payload = [
{"schema": entry["schema"], "handler_key": entry["handler_key"]}
@@ -295,243 +412,44 @@ def _register_server_tools(name: str, server: "MCPServerTask", config: dict) ->
live registry rather than mutating ``toolsets.TOOLSETS``. Lossy name
normalization can map distinct raw names (``read-file``/``read_file``) to
one registry name; such collisions fail closed — every ambiguous entry is
skipped. Returns the registered prefixed names.
skipped. Generated utilities share the namespace and join the same
preflight. Returns the registered prefixed names.
"""
from tools.registry import registry
registered_names: List[str] = []
toolset_name = f"mcp-{name}"
_should_register = _make_tool_filter(name, config)
should_register = _make_tool_filter(name, config)
check_fn = _make_check_fn(name)
candidates: List[dict] = []
# Security boundary: capture trust tier and readOnlyHint NOW, at discovery,
# so the call-time gate classifies from data we control, not re-read
# server-supplied state.
_record_tool_trust_metadata(name, config, server._tools)
for mcp_tool in server._tools:
if not _should_register(mcp_tool.name):
logger.debug(
"MCP server '%s': skipping tool '%s' (filtered by config)",
name,
mcp_tool.name,
)
continue
_core._scan_mcp_description(name, mcp_tool.name, mcp_tool.description or "")
schema = _core._convert_mcp_schema(name, mcp_tool)
candidates.append(
{
"registry_name": schema["name"],
"origin": f"tool {mcp_tool.name!r}",
"schema": schema,
"handler": _core._make_tool_handler(
name, mcp_tool.name, server.tool_timeout
),
"check_fn": check_fn,
}
)
# Generated resource/prompt utility tools share the same namespace as raw
# MCP tools, so they must participate in the same collision preflight.
for entry in _select_utility_schemas(name, server, config):
schema = entry["schema"]
handler_key = entry["handler_key"]
candidates.append(
{
"registry_name": schema["name"],
"origin": f"generated utility {handler_key!r}",
"schema": schema,
"handler": _UTILITY_HANDLER_FACTORIES[handler_key](
name, server.tool_timeout
),
"check_fn": check_fn,
}
)
unique_candidates, ambiguous_names, shadowed_utilities = _resolve_name_collisions(name, candidates)
for candidate in unique_candidates:
registry_name = candidate["registry_name"]
if registry_name in ambiguous_names:
continue
if (registry_name, candidate["origin"]) in shadowed_utilities:
continue
existing_toolset = registry.get_toolset_for_tool(registry_name)
if existing_toolset and existing_toolset != toolset_name:
if existing_toolset.startswith("mcp-"):
logger.error(
"MCP server '%s': %s normalizes to '%s', already owned by "
"MCP toolset '%s' — skipping to preserve the existing owner",
name,
candidate["origin"],
registry_name,
existing_toolset,
)
else:
logger.warning(
"MCP server '%s': %s (→ '%s') collides with built-in tool "
"in toolset '%s' — skipping to preserve built-in",
name,
candidate["origin"],
registry_name,
existing_toolset,
)
continue
registry.register(
name=registry_name,
toolset=toolset_name,
schema=candidate["schema"],
handler=candidate["handler"],
check_fn=candidate["check_fn"],
is_async=False,
description=candidate["schema"]["description"],
scope=_core._server_registry_scope(name),
)
# The pre-check above is advisory only. Multiple servers connect in
# parallel, so ToolRegistry.register() is the atomic ownership gate.
if registry.get_toolset_for_tool(registry_name) != toolset_name:
logger.error(
"MCP server '%s': registration of %s as '%s' was rejected by "
"the registry; skipping provenance/count updates",
name,
candidate["origin"],
registry_name,
)
continue
_core._track_mcp_tool_server(registry_name, name)
registered_names.append(registry_name)
if registered_names:
registry.register_toolset_alias(name, toolset_name)
_write_schema_cache(name, server, config, _should_register)
return registered_names
class _CachedMCPTool:
"""Minimal stand-in for MCP Tool objects loaded from the schema cache."""
__slots__ = ("name", "description", "inputSchema")
def __init__(self, name: str, description: str, inputSchema: dict):
self.name = name
self.description = description
self.inputSchema = inputSchema or {}
candidates = _tool_candidates(name, server._tools, should_register, server.tool_timeout)
candidates += _utility_candidates(name, _select_utility_schemas(name, server, config), server.tool_timeout)
registered = _register_candidates(
name, _resolve_name_collisions(name, candidates),
check_fn=check_fn, scope=lambda: _core._server_registry_scope(name), lazy=False,
)
if registered:
_write_schema_cache(name, server, config, should_register)
return registered
def _register_from_cache_sync(name: str, config: dict, entry: dict) -> List[str]:
"""Lazy startup: register a server's tools from a cached manifest with no
child process. The first real call routes through
``_get_connected_server_for_call`` -> ``_ensure_lazy_server_connected``."""
from tools.registry import registry
from tools.mcp_schema_cache import (
config_fingerprint,
tools_from_cache_entry,
utility_tools_from_cache_entry,
)
``_get_connected_server_for_call`` -> ``_ensure_lazy_server_connected``.
Trust metadata is recorded first so the call-time gate is identical whether
the server was spawned live or registered from cache."""
from tools.mcp_schema_cache import config_fingerprint, tools_from_cache_entry, utility_tools_from_cache_entry
registered_names: List[str] = []
toolset_name = f"mcp-{name}"
fingerprint = config_fingerprint(config)
tool_timeout = _resolve_tool_timeout(config)
_should_register = _make_tool_filter(name, config)
check_fn = _make_check_fn(name)
# Record trust metadata before registration so the call-time gate is
# identical whether the server was spawned live or registered from cache.
# Missing "annotations" in older cache files fails closed to write-capable.
cached_tool_objs = [
SimpleNamespace(
name=raw.get("name"),
annotations=raw.get("annotations")
if isinstance(raw.get("annotations"), dict) else None,
)
for raw in tools_from_cache_entry(entry)
if isinstance(raw, dict) and raw.get("name")
]
_record_tool_trust_metadata(name, config, cached_tool_objs)
for raw in tools_from_cache_entry(entry):
if not isinstance(raw, dict):
continue
raw_name = raw.get("name")
if not raw_name or not _should_register(raw_name):
continue
raw_schema = raw.get("inputSchema")
mcp_tool = _CachedMCPTool(
raw_name,
raw.get("description") or "",
raw_schema if isinstance(raw_schema, dict) else {},
)
# Defense-in-depth: the cache file is user-writable JSON, so apply the
# same injection scan as eager discovery.
_core._scan_mcp_description(name, mcp_tool.name, mcp_tool.description or "")
schema = _core._convert_mcp_schema(name, mcp_tool)
registry_name = schema["name"]
existing_toolset = registry.get_toolset_for_tool(registry_name)
if existing_toolset and existing_toolset != toolset_name:
logger.warning(
"MCP server '%s' (lazy): cached tool '%s' collides with "
"toolset '%s' — skipping",
name, registry_name, existing_toolset,
)
continue
registry.register(
name=registry_name,
toolset=toolset_name,
schema=schema,
handler=_core._make_tool_handler(name, raw_name, tool_timeout),
check_fn=check_fn,
is_async=False,
description=schema["description"],
scope=_core._mcp_registry_scope(),
)
if registry.get_toolset_for_tool(registry_name) != toolset_name:
continue
_core._track_mcp_tool_server(registry_name, name)
registered_names.append(registry_name)
for raw in utility_tools_from_cache_entry(entry):
if not isinstance(raw, dict):
continue
schema = raw.get("schema")
handler_key = raw.get("handler_key")
if not isinstance(schema, dict) or handler_key not in _UTILITY_HANDLER_FACTORIES:
continue
util_name = schema.get("name") or ""
if not util_name:
continue
existing_toolset = registry.get_toolset_for_tool(util_name)
if existing_toolset and existing_toolset != toolset_name:
continue
registry.register(
name=util_name,
toolset=toolset_name,
schema=schema,
handler=_UTILITY_HANDLER_FACTORIES[handler_key](name, tool_timeout),
check_fn=check_fn,
is_async=False,
description=schema.get("description") or "",
scope=_core._mcp_registry_scope(),
)
if registry.get_toolset_for_tool(util_name) != toolset_name:
continue
_core._track_mcp_tool_server(util_name, name)
registered_names.append(util_name)
if registered_names:
registry.register_toolset_alias(name, toolset_name)
cached_tools = _CachedMCPTool.from_cache_dicts(tools_from_cache_entry(entry))
_record_tool_trust_metadata(name, config, cached_tools)
candidates = _tool_candidates(name, cached_tools, _make_tool_filter(name, config), tool_timeout)
candidates += _utility_candidates(name, utility_tools_from_cache_entry(entry), tool_timeout)
registered = _register_candidates(
name, candidates, check_fn=check_fn, scope=_core._mcp_registry_scope, lazy=True,
)
if registered:
with _core._lock:
_core._lazy_server_configs[name] = dict(config)
_core._lazy_server_fingerprints[name] = fingerprint
_core._lazy_server_tool_names[name] = list(registered_names)
logger.info(
"MCP server '%s' (lazy): registered %d tool(s) from schema cache",
name, len(registered_names),
)
return registered_names
_core._lazy_server_fingerprints[name] = config_fingerprint(config)
_core._lazy_server_tool_names[name] = list(registered)
logger.info("MCP server '%s' (lazy): registered %d tool(s) from schema cache", name, len(registered))
return registered

View File

@@ -1,16 +1,109 @@
"""MCP client-side handlers for server-initiated requests: sampling (sampling/createMessage, text and tool-use results) and elicitation. Split from tools/mcp_tool.py."""
"""MCP client-side handlers for server-initiated requests: sampling
(sampling/createMessage, text and tool-use results) and elicitation."""
import asyncio
import json
import logging
import time
from typing import List, Optional
from typing import Callable, List, Optional
from tools.mcp_tool_common import _MISSING, _exc_str, _safe_numeric, _sanitize_error, mcp_field, _core
from tools.mcp_tool_schema import _normalize_mcp_input_schema
logger = logging.getLogger("tools.mcp_tool")
def _tool_use_id(block):
"""Tool-use id (the discriminator for a tool *result* block), read under both
SDK spellings — on mcp 2.x a bare ``hasattr(b, "toolUseId")`` is False and
would silently drop tool results."""
return mcp_field(block, "tool_use_id", "toolUseId", _MISSING)
def _is_tool_use(block) -> bool:
return hasattr(block, "name") and hasattr(block, "input")
def _tool_result_text(block) -> str:
"""Text of a ToolResultContent block ("" when it carries no content)."""
content = getattr(block, "content", None)
if content is None:
return ""
items = content if isinstance(content, list) else [content]
return "\n".join(item.text for item in items if hasattr(item, "text"))
def _content_part(block) -> Optional[dict]:
"""One OpenAI content part for a text/image block; None when unsupported."""
if hasattr(block, "text"):
return {"type": "text", "text": block.text}
mime = mcp_field(block, "mime_type", "mimeType", _MISSING)
if hasattr(block, "data") and mime is not _MISSING:
return {"type": "image_url", "image_url": {"url": f"data:{mime};base64,{block.data}"}}
logger.warning("Unsupported sampling content block type: %s (skipped)", type(block).__name__)
return None
def _tool_call_dict(tu, index: int) -> dict:
args = tu.input
return {
"id": getattr(tu, "id", f"call_{index}"),
"type": "function",
"function": {
"name": tu.name,
"arguments": json.dumps(args, ensure_ascii=False) if isinstance(args, dict) else str(args),
},
}
def _convert_sampling_message(msg) -> List[dict]:
"""One MCP SamplingMessage -> OpenAI-format messages (tool results first,
then either an assistant tool_calls message or plain content)."""
blocks = msg.content_as_list if hasattr(msg, "content_as_list") else (
msg.content if isinstance(msg.content, list) else [msg.content]
)
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)]}
text_parts = [b.text for b in content_blocks if hasattr(b, "text")]
if text_parts:
msg_dict["content"] = "\n".join(text_parts)
out.append(msg_dict)
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})
return out
def _parse_tool_call_arguments(server_name: str, args) -> dict:
"""LLM tool_calls arguments -> dict; malformed JSON / non-dict values are
wrapped as ``{"_raw": ...}`` rather than dropped."""
if isinstance(args, str):
try:
return json.loads(args)
except (json.JSONDecodeError, ValueError):
logger.warning(
"MCP server '%s': malformed tool_calls arguments from LLM (wrapping as raw): %.100s",
server_name, args,
)
return {"_raw": args}
return args if isinstance(args, dict) else {"_raw": str(args)}
def _response_total_tokens(response, default):
return getattr(getattr(response, "usage", None), "total_tokens", default)
class SamplingHandler:
"""Handles sampling/createMessage requests for one MCP server.
@@ -24,23 +117,17 @@ class SamplingHandler:
"""
_STOP_REASON_MAP = {"stop": "endTurn", "length": "maxTokens", "tool_calls": "toolUse"}
_LOG_LEVELS = {"debug": logging.DEBUG, "info": logging.INFO, "warning": logging.WARNING}
def __init__(self, server_name: str, config: dict):
self.server_name = server_name
self.max_rpm = _safe_numeric(config.get("max_rpm", 10), 10, int)
self.timeout = _safe_numeric(config.get("timeout", 30), 30, float)
self.max_tokens_cap = _safe_numeric(config.get("max_tokens_cap", 4096), 4096, int)
self.max_tool_rounds = _safe_numeric(
config.get("max_tool_rounds", 5), 5, int, minimum=0,
)
self.max_tool_rounds = _safe_numeric(config.get("max_tool_rounds", 5), 5, int, minimum=0)
self.model_override = config.get("model")
self.allowed_models = config.get("allowed_models", [])
_log_levels = {"debug": logging.DEBUG, "info": logging.INFO, "warning": logging.WARNING}
self.audit_level = _log_levels.get(
str(config.get("log_level", "info")).lower(), logging.INFO,
)
self.audit_level = self._LOG_LEVELS.get(str(config.get("log_level", "info")).lower(), logging.INFO)
self._rate_timestamps: List[float] = []
self._tool_loop_count = 0
self.metrics = {"requests": 0, "errors": 0, "tokens_used": 0, "tool_use_count": 0}
@@ -59,100 +146,20 @@ class SamplingHandler:
"""Config override > server hint > None (use default)."""
if self.model_override:
return self.model_override
if preferences and hasattr(preferences, "hints") and preferences.hints:
for hint in preferences.hints:
if hasattr(hint, "name") and hint.name:
return hint.name
for hint in (getattr(preferences, "hints", None) or []):
if getattr(hint, "name", None):
return hint.name
return None
@staticmethod
def _extract_tool_result_text(block) -> str:
"""Extract text from a ToolResultContent block."""
if not hasattr(block, "content") or block.content is None:
return ""
items = block.content if isinstance(block.content, list) else [block.content]
return "\n".join(item.text for item in items if hasattr(item, "text"))
return _tool_result_text(block)
def _convert_messages(self, params) -> List[dict]:
"""Convert MCP SamplingMessages to OpenAI format.
Uses ``msg.content_as_list`` when the SDK provides it; dispatches per
block by duck-typing.
"""
# A tool-use id is the discriminator for a tool *result* block; it must be
# read under both spellings (mcp_field) — on mcp 2.x a bare
# ``hasattr(b, "toolUseId")`` is False, silently dropping tool results.
def _tool_use_id(block):
return mcp_field(block, "tool_use_id", "toolUseId", _MISSING)
def _is_tool_use(block):
return hasattr(block, "name") and hasattr(block, "input")
messages: List[dict] = []
for msg in params.messages:
blocks = msg.content_as_list if hasattr(msg, "content_as_list") else (
msg.content if isinstance(msg.content, list) else [msg.content]
)
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)
]
for tr in tool_results:
messages.append({
"role": "tool",
"tool_call_id": _tool_use_id(tr),
"content": self._extract_tool_result_text(tr),
})
if tool_uses:
tc_list = []
for tu in tool_uses:
tc_list.append({
"id": getattr(tu, "id", f"call_{len(tc_list)}"),
"type": "function",
"function": {
"name": tu.name,
"arguments": json.dumps(tu.input, ensure_ascii=False) if isinstance(tu.input, dict) else str(tu.input),
},
})
msg_dict: dict = {"role": msg.role, "tool_calls": tc_list}
text_parts = [b.text for b in content_blocks if hasattr(b, "text")]
if text_parts:
msg_dict["content"] = "\n".join(text_parts)
messages.append(msg_dict)
elif content_blocks:
# Pure text/image content.
if len(content_blocks) == 1 and hasattr(content_blocks[0], "text"):
messages.append({"role": msg.role, "content": content_blocks[0].text})
else:
parts = []
for block in content_blocks:
block_mime = mcp_field(
block, "mime_type", "mimeType", _MISSING
)
if hasattr(block, "text"):
parts.append({"type": "text", "text": block.text})
elif hasattr(block, "data") and block_mime is not _MISSING:
parts.append({
"type": "image_url",
"image_url": {"url": f"data:{block_mime};base64,{block.data}"},
})
else:
logger.warning(
"Unsupported sampling content block type: %s (skipped)",
type(block).__name__,
)
if parts:
messages.append({"role": msg.role, "content": parts})
return messages
"""Convert MCP SamplingMessages to OpenAI format (``content_as_list``
when the SDK provides it; per-block duck-typed dispatch)."""
return [m for msg in params.messages for m in _convert_sampling_message(msg)]
@staticmethod
def _error(message: str, code: int = -1):
@@ -161,6 +168,11 @@ class SamplingHandler:
return _core.ErrorData(code=code, message=message)
raise Exception(message)
def _fail(self, message: str):
"""Count an error and return the ErrorData for it."""
self.metrics["errors"] += 1
return self._error(message)
def _build_tool_use_result(self, choice, response):
"""Build a CreateMessageResultWithTools from an LLM tool_calls response."""
self.metrics["tool_use_count"] += 1
@@ -168,71 +180,41 @@ class SamplingHandler:
# Tool-loop governance.
if self.max_tool_rounds == 0:
self._tool_loop_count = 0
return self._error(
f"Tool loops disabled for server '{self.server_name}' (max_tool_rounds=0)"
)
return self._error(f"Tool loops disabled for server '{self.server_name}' (max_tool_rounds=0)")
self._tool_loop_count += 1
if self._tool_loop_count > self.max_tool_rounds:
self._tool_loop_count = 0
return self._error(
f"Tool loop limit exceeded for server '{self.server_name}' "
f"(max {self.max_tool_rounds} rounds)"
f"Tool loop limit exceeded for server '{self.server_name}' (max {self.max_tool_rounds} rounds)"
)
content_blocks = []
for tc in choice.message.tool_calls:
args = tc.function.arguments
if isinstance(args, str):
try:
parsed = json.loads(args)
except (json.JSONDecodeError, ValueError):
logger.warning(
"MCP server '%s': malformed tool_calls arguments "
"from LLM (wrapping as raw): %.100s",
self.server_name, args,
)
parsed = {"_raw": args}
else:
parsed = args if isinstance(args, dict) else {"_raw": str(args)}
content_blocks.append(_core.ToolUseContent(
type="tool_use",
id=tc.id,
name=tc.function.name,
input=parsed,
))
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
]
logger.log(
self.audit_level,
"MCP server '%s' sampling response: model=%s, tokens=%s, tool_calls=%d",
self.server_name, response.model,
getattr(getattr(response, "usage", None), "total_tokens", "?"),
len(content_blocks),
self.server_name, response.model, _response_total_tokens(response, "?"), len(content_blocks),
)
return _core.CreateMessageResultWithTools(
role="assistant",
content=content_blocks,
model=response.model,
stopReason="toolUse",
role="assistant", content=content_blocks, model=response.model, stopReason="toolUse",
)
def _build_text_result(self, choice, response):
"""Build a CreateMessageResult from a normal text response (resets the tool loop)."""
self._tool_loop_count = 0
response_text = choice.message.content or ""
logger.log(
self.audit_level,
"MCP server '%s' sampling response: model=%s, tokens=%s",
self.server_name, response.model,
getattr(getattr(response, "usage", None), "total_tokens", "?"),
self.server_name, response.model, _response_total_tokens(response, "?"),
)
return _core.CreateMessageResult(
role="assistant",
content=_core.TextContent(type="text", text=_sanitize_error(response_text)),
content=_core.TextContent(type="text", text=_sanitize_error(choice.message.content or "")),
model=response.model,
stopReason=self._STOP_REASON_MAP.get(choice.finish_reason, "endTurn"),
)
@@ -241,130 +223,89 @@ class SamplingHandler:
"""Kwargs to pass to ClientSession for sampling support."""
return {
"sampling_callback": self,
"sampling_capabilities": _core.SamplingCapability(
tools=_core.SamplingToolsCapability(),
),
"sampling_capabilities": _core.SamplingCapability(tools=_core.SamplingToolsCapability()),
}
async def __call__(self, context, params):
"""SDK sampling callback (``SamplingFnT``). Returns CreateMessageResult,
CreateMessageResultWithTools, or ErrorData."""
def _admit(self, params):
"""Rate-limit + allowed_models gate. Returns ``(resolved_model, None)``
or ``(None, ErrorData)``."""
if not self._check_rate_limit():
logger.warning(
"MCP server '%s' sampling rate limit exceeded (%d/min)",
self.server_name, self.max_rpm,
logger.warning("MCP server '%s' sampling rate limit exceeded (%d/min)", self.server_name, self.max_rpm)
return None, self._fail(
f"Sampling rate limit exceeded for server '{self.server_name}' ({self.max_rpm} requests/minute)"
)
self.metrics["errors"] += 1
return self._error(
f"Sampling rate limit exceeded for server '{self.server_name}' "
f"({self.max_rpm} requests/minute)"
)
model = self._resolve_model(
mcp_field(params, "model_preferences", "modelPreferences")
)
from agent.auxiliary_client import call_llm
model = self._resolve_model(mcp_field(params, "model_preferences", "modelPreferences"))
resolved_model = model or self.model_override or ""
if self.allowed_models and resolved_model and resolved_model not in self.allowed_models:
logger.warning(
"MCP server '%s' requested model '%s' not in allowed_models",
self.server_name, resolved_model,
"MCP server '%s' requested model '%s' not in allowed_models", self.server_name, resolved_model,
)
self.metrics["errors"] += 1
return self._error(
return None, self._fail(
f"Model '{resolved_model}' not allowed for server "
f"'{self.server_name}'. Allowed: {', '.join(self.allowed_models)}"
)
return resolved_model, None
def _build_llm_call(self, params, resolved_model: str) -> Callable[[], object]:
"""Translate the sampling params into a zero-arg sync ``call_llm`` thunk
(run off-loop so the MCP loop is not blocked)."""
from agent.auxiliary_client import call_llm
messages = self._convert_messages(params)
system_prompt = mcp_field(params, "system_prompt", "systemPrompt")
if system_prompt:
messages.insert(0, {"role": "system", "content": system_prompt})
max_tokens = min(
mcp_field(params, "max_tokens", "maxTokens", self.max_tokens_cap),
self.max_tokens_cap,
)
call_temperature = None
if hasattr(params, "temperature") and params.temperature is not None:
call_temperature = params.temperature
max_tokens = min(mcp_field(params, "max_tokens", "maxTokens", self.max_tokens_cap), self.max_tokens_cap)
temperature = getattr(params, "temperature", None)
# Forward server-provided tools.
call_tools = None
server_tools = getattr(params, "tools", None)
if server_tools:
call_tools = [
{
"type": "function",
"function": {
"name": getattr(t, "name", ""),
"description": getattr(t, "description", "") or "",
"parameters": _normalize_mcp_input_schema(
mcp_field(t, "input_schema", "inputSchema")
),
},
}
for t in server_tools
]
tools = [
{
"type": "function",
"function": {
"name": getattr(t, "name", ""),
"description": getattr(t, "description", "") or "",
"parameters": _normalize_mcp_input_schema(mcp_field(t, "input_schema", "inputSchema")),
},
}
for t in server_tools
] if server_tools else None
logger.log(
self.audit_level,
"MCP server '%s' sampling request: model=%s, max_tokens=%d, messages=%d",
self.server_name, resolved_model, max_tokens, len(messages),
)
return lambda: call_llm(
task="mcp", model=resolved_model or None, messages=messages, temperature=temperature,
max_tokens=max_tokens, tools=tools, timeout=self.timeout,
)
# Offload the sync LLM call so the MCP loop is not blocked.
def _sync_call():
return call_llm(
task="mcp",
model=resolved_model or None,
messages=messages,
temperature=call_temperature,
max_tokens=max_tokens,
tools=call_tools,
timeout=self.timeout,
)
async def __call__(self, context, params):
"""SDK sampling callback (``SamplingFnT``). Returns CreateMessageResult,
CreateMessageResultWithTools, or ErrorData."""
resolved_model, err = self._admit(params)
if err is not None:
return err
sync_call = self._build_llm_call(params, resolved_model)
try:
response = await asyncio.wait_for(
asyncio.to_thread(_sync_call), timeout=self.timeout,
)
response = await asyncio.wait_for(asyncio.to_thread(sync_call), timeout=self.timeout)
except asyncio.TimeoutError:
self.metrics["errors"] += 1
return self._error(
f"Sampling LLM call timed out after {self.timeout}s "
f"for server '{self.server_name}'"
)
return self._fail(f"Sampling LLM call timed out after {self.timeout}s for server '{self.server_name}'")
except Exception as exc:
self.metrics["errors"] += 1
return self._error(
f"Sampling LLM call failed: {_sanitize_error(_exc_str(exc))}"
)
return self._fail(f"Sampling LLM call failed: {_sanitize_error(_exc_str(exc))}")
# Empty choices happen on content filtering / provider errors.
if not getattr(response, "choices", None):
self.metrics["errors"] += 1
return self._error(
f"LLM returned empty response (no choices) for server "
f"'{self.server_name}'"
)
return self._fail(f"LLM returned empty response (no choices) for server '{self.server_name}'")
choice = response.choices[0]
self.metrics["requests"] += 1
total_tokens = getattr(getattr(response, "usage", None), "total_tokens", 0)
total_tokens = _response_total_tokens(response, 0)
if isinstance(total_tokens, int):
self.metrics["tokens_used"] += total_tokens
if (
choice.finish_reason == "tool_calls"
and hasattr(choice.message, "tool_calls")
and choice.message.tool_calls
):
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)
@@ -377,16 +318,11 @@ 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():
field_type = ""
field_desc = ""
if isinstance(field_spec, dict):
field_type = str(field_spec.get("type", "") or "")
field_desc = str(field_spec.get("description", "") or "")
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 ""
if field_desc:
lines.append(f" - {field_name}{suffix}: {field_desc}")
else:
lines.append(f" - {field_name}{suffix}")
lines.append(f" - {field_name}{suffix}: {field_desc}" if field_desc else f" - {field_name}{suffix}")
return "\n".join(lines)
@@ -411,111 +347,75 @@ class ElicitationHandler:
# Back-reference for the agent's contextvars snapshot; optional so the
# handler stays unit-testable in isolation.
self.owner = owner
self.metrics = {
"requests": 0,
"accepted": 0,
"declined": 0,
"errors": 0,
}
self.metrics = {"requests": 0, "accepted": 0, "declined": 0, "errors": 0}
def session_kwargs(self) -> dict:
"""Kwargs to pass to ClientSession for elicitation support."""
return {"elicitation_callback": self}
def _result(self, action: str, metric: str):
"""Count *metric* and return ``ElicitResult(action)`` (accept carries empty content)."""
self.metrics[metric] += 1
if action == "accept":
return _core.ElicitResult(action="accept", content={})
return _core.ElicitResult(action=action)
def _consent_thunk(self, message: str, description: str) -> Callable[[], str]:
"""Sync consent call, replaying the agent's contextvars snapshot when the
owner captured one: the recv-loop task does NOT inherit them, and
gateway-platform detection needs them. ``Context.run`` executes a
context once, so it is copied per elicitation."""
from tools.approval import request_elicitation_consent
kwargs = {"timeout_seconds": int(self.timeout), "surface": f"mcp-elicitation/{self.server_name}"}
captured = getattr(self.owner, "_pending_call_context", None) if self.owner else None
if captured is None:
return lambda: request_elicitation_consent(message, description, **kwargs)
return lambda: captured.copy().run(request_elicitation_consent, message, description, **kwargs)
async def __call__(self, context, params):
"""SDK elicitation callback (``ElicitationFnT``). Returns ElicitResult or ErrorData."""
self.metrics["requests"] += 1
# URL-mode (OAuth, payment) would need a browser + waiting for
# notifications/elicitation/complete — not implemented; decline cleanly.
mode = getattr(params, "mode", "form")
if mode == "url":
if getattr(params, "mode", "form") == "url":
logger.info(
"MCP server '%s' requested URL-mode elicitation; "
"declining (URL-mode elicitation not implemented)",
"MCP server '%s' requested URL-mode elicitation; declining (URL-mode elicitation not implemented)",
self.server_name,
)
self.metrics["declined"] += 1
return _core.ElicitResult(action="decline")
return self._result("decline", "declined")
message = getattr(params, "message", "") or (
f"MCP server '{self.server_name}' is requesting your approval"
)
message = getattr(params, "message", "") or f"MCP server '{self.server_name}' is requesting your approval"
# ``requestedSchema`` on mcp 1.x, ``requested_schema`` on 2.0 (pydantic
# aliases don't apply to attribute access) — read both or the user is
# asked to approve without seeing which fields the server wants.
schema = (
getattr(params, "requestedSchema", None)
or getattr(params, "requested_schema", None)
or {}
)
schema = getattr(params, "requestedSchema", None) or getattr(params, "requested_schema", None) or {}
description = _format_elicitation_schema_summary(schema, self.server_name)
logger.info(
"MCP server '%s' elicitation request: %s",
self.server_name, _sanitize_error(message)[:200],
)
logger.info("MCP server '%s' elicitation request: %s", self.server_name, _sanitize_error(message)[:200])
# Lazy import avoids import-order coupling with early-bootstrap tools.approval.
try:
from tools.approval import request_elicitation_consent
invoke_consent = self._consent_thunk(message, description)
except Exception as exc: # pragma: no cover -- defensive
logger.error(
"MCP server '%s' elicitation: approval system unavailable: %s",
self.server_name, exc,
)
self.metrics["errors"] += 1
return _core.ElicitResult(action="decline")
logger.error("MCP server '%s' elicitation: approval system unavailable: %s", self.server_name, exc)
return self._result("decline", "errors")
# Offload the sync consent flow to a thread — inline it would freeze the
# MCP loop and every other RPC on this session. The recv-loop task does
# NOT inherit the agent's contextvars, so replay the snapshot captured on
# owner._pending_call_context for gateway-platform detection.
captured = getattr(self.owner, "_pending_call_context", None) if self.owner else None
def _invoke_consent() -> str:
if captured is None:
return request_elicitation_consent(
message,
description,
timeout_seconds=int(self.timeout),
surface=f"mcp-elicitation/{self.server_name}",
)
# Context.run executes a context once — copy so multiple
# elicitations within one tool call work.
return captured.copy().run(
request_elicitation_consent,
message,
description,
timeout_seconds=int(self.timeout),
surface=f"mcp-elicitation/{self.server_name}",
)
# MCP loop and every other RPC on this session.
try:
answer = await asyncio.wait_for(
asyncio.to_thread(_invoke_consent),
timeout=self.timeout + self._OUTER_TIMEOUT_GRACE_SECONDS,
asyncio.to_thread(invoke_consent), timeout=self.timeout + self._OUTER_TIMEOUT_GRACE_SECONDS,
)
except asyncio.TimeoutError:
logger.warning(
"MCP server '%s' elicitation timed out after %ds",
self.server_name, int(self.timeout),
)
self.metrics["errors"] += 1
return _core.ElicitResult(action="cancel")
logger.warning("MCP server '%s' elicitation timed out after %ds", self.server_name, int(self.timeout))
return self._result("cancel", "errors")
except Exception as exc:
logger.error(
"MCP server '%s' elicitation failed: %s",
self.server_name, exc, exc_info=True,
)
self.metrics["errors"] += 1
return _core.ElicitResult(action="decline")
logger.error("MCP server '%s' elicitation failed: %s", self.server_name, exc, exc_info=True)
return self._result("decline", "errors")
if answer == "accept":
self.metrics["accepted"] += 1
return _core.ElicitResult(action="accept", content={})
return self._result("accept", "accepted")
if answer == "cancel":
self.metrics["errors"] += 1
return _core.ElicitResult(action="cancel")
self.metrics["declined"] += 1
return _core.ElicitResult(action="decline")
return self._result("cancel", "errors")
return self._result("decline", "declined")