From 97fbd7e232e601cdd8705b793c8aac983aa1fb31 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 15:57:42 -0700 Subject: [PATCH] refactor(mcp): table-drive elicitation answers, unify sampling response log, tighten collision diagnostics --- tools/mcp_tool_agent.py | 1 - tools/mcp_tool_common.py | 5 ----- tools/mcp_tool_config.py | 4 ---- tools/mcp_tool_lifecycle.py | 2 -- tools/mcp_tool_registration.py | 17 +++++++-------- tools/mcp_tool_sampling.py | 38 ++++++++++++++-------------------- tools/mcp_tool_schema.py | 5 ----- 7 files changed, 22 insertions(+), 50 deletions(-) diff --git a/tools/mcp_tool_agent.py b/tools/mcp_tool_agent.py index 7e921ca573..963af71bda 100644 --- a/tools/mcp_tool_agent.py +++ b/tools/mcp_tool_agent.py @@ -10,7 +10,6 @@ from tools.mcp_tool_common import _core logger = logging.getLogger("tools.mcp_tool") - # Serializes in-place swaps of ``agent.tools`` / ``agent.valid_tool_names`` by # the reload RPC, gateway reload and late-binding refresh thread; the run loop # reads them during tool iteration and must never see a half-updated pair. diff --git a/tools/mcp_tool_common.py b/tools/mcp_tool_common.py index 0e96d04200..bdd91de8aa 100644 --- a/tools/mcp_tool_common.py +++ b/tools/mcp_tool_common.py @@ -26,7 +26,6 @@ class _OriginProxy: return getattr(mcp_tool, name) - _core = _OriginProxy() _MISSING = object() @@ -42,7 +41,6 @@ def mcp_field(obj, snake: str, camel: str, default=None): value = getattr(obj, camel, _MISSING) return default if value is _MISSING else value - _DEFAULT_TOOL_TIMEOUT = 300 # seconds for tool calls @@ -63,7 +61,6 @@ def _resolve_tool_timeout(config: dict) -> float: logger.debug("mcp.tool_call timeout resolution failed", exc_info=True) return _DEFAULT_TOOL_TIMEOUT - # Jitter on reconnect backoff so servers that lost the same backend don't # retry in lockstep (thundering herd, synchronized log bursts). _BACKOFF_JITTER = 0.2 # +/-20% @@ -73,7 +70,6 @@ def _jittered(seconds: float) -> float: """``seconds`` with +/-20% uniform jitter, floored at 0.""" return max(0.0, seconds * random.uniform(1.0 - _BACKOFF_JITTER, 1.0 + _BACKOFF_JITTER)) - # Credential patterns to strip from error messages. _CREDENTIAL_PATTERN = re.compile( r"(?:" @@ -135,7 +131,6 @@ def _safe_numeric(value, default, coerce=int, minimum=1): except (TypeError, ValueError, OverflowError): return default - _TRUE_WORDS = frozenset({"true", "1", "yes", "on"}) _FALSE_WORDS = frozenset({"false", "0", "no", "off"}) diff --git a/tools/mcp_tool_config.py b/tools/mcp_tool_config.py index 4bf532814c..940fdc0e9e 100644 --- a/tools/mcp_tool_config.py +++ b/tools/mcp_tool_config.py @@ -14,7 +14,6 @@ from tools.mcp_tool_common import _env_ref_name, _prepend_path, _core logger = logging.getLogger("tools.mcp_tool") - _mcp_stderr_log_fh: Optional[Any] = None _mcp_stderr_log_lock = threading.Lock() @@ -56,7 +55,6 @@ def _write_stderr_log_header(server_name: str) -> None: except Exception: pass - # Env vars safe to pass to stdio subprocesses (no secrets). _SAFE_ENV_KEYS = frozenset({"PATH", "HOME", "USER", "LANG", "LC_ALL", "TERM", "SHELL", "TMPDIR"}) @@ -93,7 +91,6 @@ def _workspace_basename() -> str: root = _core._workspace_folder() return os.path.basename(root.rstrip("/\\")) or root - # Cursor's case-sensitive context vars -> resolver. _CONTEXT_VAR_RESOLVERS = { "userHome": lambda: os.path.expanduser("~"), @@ -228,7 +225,6 @@ def _interpolate_env_vars(value): return [_interpolate_env_vars(v) for v in value] return value - # (server_name, dotted key path) pairs already warned about; config loads # happen on every discovery pass, so warn once per process. _whitespace_warned: Set[Tuple[str, str]] = set() diff --git a/tools/mcp_tool_lifecycle.py b/tools/mcp_tool_lifecycle.py index b8e799697a..45890dad4f 100644 --- a/tools/mcp_tool_lifecycle.py +++ b/tools/mcp_tool_lifecycle.py @@ -10,7 +10,6 @@ from tools.mcp_tool_common import _core logger = logging.getLogger("tools.mcp_tool") - # Live stdio MCP children (pid -> server_name), added after connection and # removed on normal shutdown, so they can be force-killed if SDK teardown fails. _stdio_pids: Dict[int, str] = {} @@ -58,7 +57,6 @@ def _snapshot_child_pids() -> set: return set() - # argv markers of non-MCP gateway children that can race into the snapshot # delta during an MCP spawn (defense-in-depth; LSP/slash_worker already use # start_new_session). Matched against argv[1:] because Python/Java children diff --git a/tools/mcp_tool_registration.py b/tools/mcp_tool_registration.py index ecdfa24bb6..061c0caa5c 100644 --- a/tools/mcp_tool_registration.py +++ b/tools/mcp_tool_registration.py @@ -257,16 +257,13 @@ def _log_foreign_owner(name: str, c: _Candidate, existing_toolset: str, lazy: bo 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, - ) + return + log, fmt = ( + (logger.error, "MCP server '%s': %s normalizes to '%s', already owned by MCP toolset '%s' — skipping to preserve the existing owner") + if existing_toolset.startswith("mcp-") else + (logger.warning, "MCP server '%s': %s (→ '%s') collides with built-in tool in toolset '%s' — skipping to preserve built-in") + ) + log(fmt, 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]: diff --git a/tools/mcp_tool_sampling.py b/tools/mcp_tool_sampling.py index a1aa99c2c5..9fef4414db 100644 --- a/tools/mcp_tool_sampling.py +++ b/tools/mcp_tool_sampling.py @@ -163,11 +163,16 @@ class SamplingHandler: 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 + def _log_response(self, response, suffix: str = "", *args) -> None: + logger.log( + self.audit_level, "MCP server '%s' sampling response: model=%s, tokens=%s" + suffix, + self.server_name, response.model, _response_total_tokens(response, "?"), *args, + ) - # Tool-loop governance. + def _build_tool_use_result(self, choice, response): + """Build a CreateMessageResultWithTools from an LLM tool_calls response, + subject to tool-loop governance (``max_tool_rounds``; 0 disables).""" + self.metrics["tool_use_count"] += 1 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)") @@ -177,7 +182,6 @@ class SamplingHandler: return self._error( f"Tool loop limit exceeded for server '{self.server_name}' (max {self.max_tool_rounds} rounds)" ) - content_blocks = [ _core.ToolUseContent( type="tool_use", id=tc.id, name=tc.function.name, @@ -185,11 +189,7 @@ class SamplingHandler: ) 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, _response_total_tokens(response, "?"), len(content_blocks), - ) + self._log_response(response, ", tool_calls=%d", len(content_blocks)) return _core.CreateMessageResultWithTools( role="assistant", content=content_blocks, model=response.model, stopReason="toolUse", ) @@ -197,11 +197,7 @@ class SamplingHandler: def _build_text_result(self, choice, response): """Build a CreateMessageResult from a normal text response (resets the tool loop).""" self._tool_loop_count = 0 - logger.log( - self.audit_level, - "MCP server '%s' sampling response: model=%s, tokens=%s", - self.server_name, response.model, _response_total_tokens(response, "?"), - ) + self._log_response(response) return _core.CreateMessageResult( role="assistant", content=_core.TextContent(type="text", text=_sanitize_error(choice.message.content or "")), @@ -326,6 +322,8 @@ class ElicitationHandler: # asyncio-side safety net over the approval's own input() timeout so the # MCP loop never blocks indefinitely if the inner timeout is bypassed. _OUTER_TIMEOUT_GRACE_SECONDS = 5 + # consent answer -> (ElicitResult action, metric); anything else declines. + _ANSWER_RESULTS = {"accept": ("accept", "accepted"), "cancel": ("cancel", "errors")} def __init__(self, server_name: str, config: dict, owner: Optional["MCPServerTask"] = None): self.server_name = server_name @@ -344,9 +342,7 @@ class ElicitationHandler: 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) + return _core.ElicitResult(action=action, content={}) if action == "accept" else _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 @@ -402,8 +398,4 @@ class ElicitationHandler: logger.error("MCP server '%s' elicitation failed: %s", self.server_name, exc, exc_info=True) return self._result("decline", "errors") - if answer == "accept": - return self._result("accept", "accepted") - if answer == "cancel": - return self._result("cancel", "errors") - return self._result("decline", "declined") + return self._result(*self._ANSWER_RESULTS.get(answer, ("decline", "declined"))) diff --git a/tools/mcp_tool_schema.py b/tools/mcp_tool_schema.py index ce3c62c188..a28c3ead51 100644 --- a/tools/mcp_tool_schema.py +++ b/tools/mcp_tool_schema.py @@ -11,7 +11,6 @@ from tools.mcp_tool_common import mcp_field logger = logging.getLogger("tools.mcp_tool") - # Prompt-injection indicators in MCP tool descriptions. WARNING-level only: # log but never block, since false positives would break legitimate servers. _MCP_INJECTION_PATTERNS = [ @@ -44,7 +43,6 @@ def _scan_mcp_description(server_name: str, tool_name: str, description: str) -> ) return findings - _EMPTY_OBJECT_SCHEMA = {"type": "object", "properties": {}} @@ -125,7 +123,6 @@ def sanitize_mcp_name_component(value: str) -> str: the historical behavior) so generated names pass provider validation.""" return re.sub(r"[^A-Za-z0-9_]", "_", str(value or "")) - # ``mcp____``: the convention shared by Claude Code, Codex and # OpenCode. The double underscore disambiguates the server/tool boundary even # when either contains underscores, and matches the Anthropic-OAuth wire form. @@ -147,7 +144,6 @@ def _convert_mcp_schema(server_name: str, mcp_tool) -> dict: "parameters": _normalize_mcp_input_schema(mcp_field(mcp_tool, "input_schema", "inputSchema")), } - # Utility tools generated per server: handler_key -> (description template, # parameter properties, required names). Schemas are FROZEN wire bytes — the # key order emitted by ``_build_utility_schemas`` must not change. @@ -209,7 +205,6 @@ def matches_name_filter(tool_name: str, patterns: set[str]) -> bool: return True return any(fnmatch.fnmatchcase(tool_name, p) for p in patterns if "*" in p or "?" in p or "[" in p) - # Utility handler -> ClientSession method it needs (legacy gate when no # initialize_result was captured). _UTILITY_CAPABILITY_METHODS = {key: key for key, *_ in _UTILITY_TOOL_SPECS}