refactor(mcp): table-drive elicitation answers, unify sampling response log, tighten collision diagnostics
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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"})
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]:
|
||||
|
||||
@@ -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")))
|
||||
|
||||
@@ -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}
|
||||
|
||||
Reference in New Issue
Block a user