refactor(mcp): table-drive elicitation answers, unify sampling response log, tighten collision diagnostics

This commit is contained in:
Teknium
2026-09-02 15:57:42 -07:00
parent ade852f1d5
commit 97fbd7e232
7 changed files with 22 additions and 50 deletions

View File

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

View File

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

View File

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

View File

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

View File

@@ -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]:

View File

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

View File

@@ -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__<server>__<tool>``: 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}