From e0bdc2324e469745d80ec377d3f53aa406c41780 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 23:54:08 -0700 Subject: [PATCH] refactor(tools): inline trivial sampling accessors, merge discovery branches --- tools/mcp_tool_sampling.py | 17 +++++------------ tools/mcp_tool_transport.py | 18 +++++++----------- 2 files changed, 12 insertions(+), 23 deletions(-) diff --git a/tools/mcp_tool_sampling.py b/tools/mcp_tool_sampling.py index aa8854597d..0d2797618d 100644 --- a/tools/mcp_tool_sampling.py +++ b/tools/mcp_tool_sampling.py @@ -18,10 +18,6 @@ def _tool_use_id(block): 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) @@ -54,8 +50,9 @@ def _convert_sampling_message(msg) -> List[dict]: 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)] + others = [b for b in blocks if _tool_use_id(b) is _MISSING] + tool_uses = [b for b in others if hasattr(b, "name") and hasattr(b, "input")] + content_blocks = [b for b in others if not (hasattr(b, "name") and hasattr(b, "input"))] 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)]} @@ -84,10 +81,6 @@ def _parse_tool_call_arguments(server_name: str, args) -> dict: 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: """``sampling_callback`` for one MCP server (per-instance rate-limit, metrics, tool-loop state). Runs on the MCP loop; the sync LLM call is offloaded via ``asyncio.to_thread``. Deprecated @@ -144,7 +137,7 @@ class SamplingHandler: 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) + self.server_name, response.model, getattr(getattr(response, "usage", None), "total_tokens", "?"), *args) def _build_tool_use_result(self, choice, response): """CreateMessageResultWithTools from a tool_calls response, under ``max_tool_rounds`` (0 disables).""" @@ -227,7 +220,7 @@ class SamplingHandler: 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 = _response_total_tokens(response, 0) + total_tokens = getattr(getattr(response, "usage", None), "total_tokens", 0) self.metrics["tokens_used"] += total_tokens if isinstance(total_tokens, int) else 0 if choice.finish_reason == "tool_calls" and getattr(choice.message, "tool_calls", None): return self._build_tool_use_result(choice, response) diff --git a/tools/mcp_tool_transport.py b/tools/mcp_tool_transport.py index cf77b6c6ef..b87e618d8f 100644 --- a/tools/mcp_tool_transport.py +++ b/tools/mcp_tool_transport.py @@ -384,7 +384,6 @@ class MCPServerTransportMixin: headers = dict(config.get("headers") or {}) # Agent Plugins v1 strict_redirect_headers: configured headers MUST NOT follow a cross-origin # redirect. Capture their names BEFORE client-generated headers are merged in. - strict_cfg_headers = bool(config.get("strict_redirect_headers")) configured_header_names = {key.lower() for key in headers} # Optional per-user identity header; explicit headers of the same name win. headers = _apply_identity_header(self.name, config, headers) @@ -394,10 +393,8 @@ class MCPServerTransportMixin: if not any(key.lower() == "mcp-protocol-version" for key in headers): headers["mcp-protocol-version"] = _core.LATEST_HANDSHAKE_VERSION connect_timeout = config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT) - ssl_verify = config.get("ssl_verify", True) - client_cert = _resolve_client_cert(self.name, config) - oauth_auth = self._build_oauth_auth(url, config) - common = (url, headers, connect_timeout, ssl_verify, client_cert, oauth_auth, strict_cfg_headers) + common = (url, headers, connect_timeout, config.get("ssl_verify", True), _resolve_client_cert(self.name, config), + self._build_oauth_auth(url, config), bool(config.get("strict_redirect_headers"))) if config.get("transport") == "sse": transport, label = self._sse_transport(*common), "SSE" else: @@ -418,12 +415,11 @@ class MCPServerTransportMixin: logger.info("MCP server '%s': does not advertise 'tools' capability — " "skipping tools/list (prompts/resources remain available)", self.name) self._tools = [] - self._register_discovered_tools_if_needed() - return - async with self._rpc_lock: - self._list_cache_meta = {} - self._tools = await _core._paginate_full_list( - self.session.list_tools, "tools", self.name, cache_meta_out=self._list_cache_meta) + else: + async with self._rpc_lock: + self._list_cache_meta = {} + self._tools = await _core._paginate_full_list( + self.session.list_tools, "tools", self.name, cache_meta_out=self._list_cache_meta) self._register_discovered_tools_if_needed() def _register_discovered_tools_if_needed(self) -> None: