refactor(mcp): unify live/cache registration into one candidate loop; phase-split sampling and elicitation handlers
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user