refactor(tools): inline trivial sampling accessors, merge discovery branches

This commit is contained in:
Teknium
2026-09-02 23:54:08 -07:00
parent a8f0249a31
commit e0bdc2324e
2 changed files with 12 additions and 23 deletions

View File

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

View File

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