refactor(tools): MCP registration/common/handlers comprehension and predicate collapses

This commit is contained in:
Teknium
2026-09-02 23:56:55 -07:00
parent d846133563
commit c74c187f42
4 changed files with 34 additions and 59 deletions

View File

@@ -67,18 +67,11 @@ def _jittered(seconds: float) -> float:
return max(0.0, seconds * random.uniform(1.0 - _BACKOFF_JITTER, 1.0 + _BACKOFF_JITTER))
# Credential patterns to strip from error messages.
# Credential patterns to strip from error messages: GitHub PAT, OpenAI-style key, Bearer token,
# and ``token= / key= / API_KEY= / password= / secret=`` assignments.
_CREDENTIAL_PATTERN = re.compile(
r"(?:ghp_[A-Za-z0-9_]{1,255}" # GitHub PAT
r"|sk-[A-Za-z0-9_]{1,255}" # OpenAI-style key
r"|Bearer\s+\S+" # Bearer token
r"|token=[^\s&,;\"']{1,255}" # token=...
r"|key=[^\s&,;\"']{1,255}" # key=...
r"|API_KEY=[^\s&,;\"']{1,255}" # API_KEY=...
r"|password=[^\s&,;\"']{1,255}" # password=...
r"|secret=[^\s&,;\"']{1,255}" # secret=...
r")",
re.IGNORECASE)
r"(?:ghp_[A-Za-z0-9_]{1,255}|sk-[A-Za-z0-9_]{1,255}|Bearer\s+\S+"
r"|(?:token|key|API_KEY|password|secret)=[^\s&,;\"']{1,255})", re.IGNORECASE)
def _env_ref_name(ref: str) -> str:
@@ -104,12 +97,11 @@ def _exc_str(exc: BaseException) -> str:
def _prepend_path(env: dict, directory: str) -> dict:
"""Prepend *directory* to env PATH if it is not already present."""
updated = dict(env or {})
if not directory:
return updated
parts = [part for part in updated.get("PATH", "").split(os.pathsep) if part]
if directory not in parts:
parts = [directory, *parts]
updated["PATH"] = os.pathsep.join(parts) if parts else directory
if directory:
parts = [part for part in updated.get("PATH", "").split(os.pathsep) if part]
if directory not in parts:
parts = [directory, *parts]
updated["PATH"] = os.pathsep.join(parts) if parts else directory
return updated

View File

@@ -43,11 +43,10 @@ async def _connect_server(name: str, config: dict) -> _core.MCPServerTask:
on the same loop). Raises on bad config, missing HTTP support or connect failure."""
server = _core.MCPServerTask(name)
claim = _core._connect_server_claim.get()
claim_token = None
if claim is not None:
claim(server)
# The run task copies this context: don't retain the discovery closure for its life.
claim_token = _core._connect_server_claim.set(None)
# The run task copies this context: don't retain the discovery closure for its life.
claim_token = _core._connect_server_claim.set(None) if claim is not None else None
try:
await server.start(config)
except asyncio.CancelledError:

View File

@@ -34,8 +34,8 @@ def _trust_gate_check(server_name: str, tool_name: str) -> Optional[str]:
try:
from tools.approval import request_elicitation_consent
answer = request_elicitation_consent(
f"MCP tool '{tool_name}' on UNTRUSTED server '{server_name}' wants to run. This "
f"tool is write-capable (no readOnlyHint=true annotation) and may modify external state.",
f"MCP tool '{tool_name}' on UNTRUSTED server '{server_name}' wants to run. This tool is write-capable "
f"(no readOnlyHint=true annotation) and may modify external state.",
f"Server '{server_name}' is configured 'trust: untrusted'. "
f"Approve to run '{tool_name}' once, or deny to block it.",
surface=f"mcp-trust/{server_name}")
@@ -109,10 +109,6 @@ def _strike(server_name: str, message: str, **extra) -> str:
return tool_error(message, **extra)
def _mcp_loop_running() -> bool:
return _core._mcp_loop is not None and _core._mcp_loop.is_running()
def _lookup_reconnectable_server(server_name: str, require_loop: bool = False):
"""The registered server object when it can be signalled to reconnect, else None.
With *require_loop*, also None unless the MCP loop is running (nothing to wait on)."""
@@ -123,6 +119,10 @@ def _lookup_reconnectable_server(server_name: str, require_loop: bool = False):
return srv
def _mcp_loop_running() -> bool:
return _core._mcp_loop is not None and _core._mcp_loop.is_running()
def _retry_once(server_name: str, retry_call, op_description: str, what: str):
"""Re-run ``retry_call`` after a recovery step. Returns the result (closing the breaker)
when it is not an error payload; None when the retry raised or errored (caller falls through)."""
@@ -267,8 +267,7 @@ async def _track_inflight_rpc(server: Any, server_name: str, op: str):
"""Register the running RPC so teardown can fail it fast. A deliberate teardown
(``_reconnecting`` set first) turns the cancel into a retryable RuntimeError; external
cancels propagate unchanged. Doubles without ``_inflight_tasks`` skip tracking."""
inflight = getattr(server, "_inflight_tasks", None)
task = asyncio.current_task()
inflight, task = getattr(server, "_inflight_tasks", None), asyncio.current_task()
tracked = task is not None and inflight is not None
if tracked:
inflight.add(task)
@@ -348,10 +347,8 @@ def _capped_structured_content(result):
"""``structuredContent`` (or None); over the hard cap it degrades to the head+tail
truncated JSON string (multi-MB JSON flood guard)."""
structured = mcp_field(result, "structured_content", "structuredContent")
if structured is None:
return None
try:
as_json = json.dumps(structured, ensure_ascii=False, default=str)
as_json = json.dumps(structured, ensure_ascii=False, default=str) if structured is not None else ""
except (TypeError, ValueError):
return structured
return _truncate_mcp_text_result(as_json) if len(as_json) > _MCP_HARD_RESULT_CAP_CHARS else structured

View File

@@ -99,8 +99,8 @@ def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: d
reason = _skip_reason(entry["handler_key"])
if reason:
logger.debug("MCP server '%s': skipping utility '%s' (%s)", server_name, entry["handler_key"], reason)
continue
selected.append(entry)
else:
selected.append(entry)
return selected
@@ -108,11 +108,9 @@ def _existing_tool_names() -> List[str]:
"""Tool names for all currently connected servers plus lazy (cache-registered) servers,
whose tools live only in the registry."""
names: List[str] = []
for _sname, server in _core._servers.items():
if hasattr(server, "_registered_tool_names"):
names.extend(server._registered_tool_names)
else:
names.extend(_core._convert_mcp_schema(server.name, t)["name"] for t in server._tools)
for server in _core._servers.values():
names.extend(server._registered_tool_names if hasattr(server, "_registered_tool_names")
else (_core._convert_mcp_schema(server.name, t)["name"] for t in server._tools))
with _core._lock:
names.extend(n for sname, tool_names in _core._lazy_server_tool_names.items()
if sname not in _core._servers for n in tool_names)
@@ -126,15 +124,10 @@ def _make_tool_filter(name: str, config: dict) -> Callable[[str], bool]:
tools_filter = config.get("tools") or {}
include_raw = tools_filter.get("include")
include_set = _normalize_name_filter(include_raw, f"mcp_servers.{name}.tools.include")
include_active = isinstance(include_raw, (str, list, tuple, set))
exclude_set = _normalize_name_filter(tools_filter.get("exclude"), f"mcp_servers.{name}.tools.exclude")
def _should_register(tool_name: str) -> bool:
if include_active:
return matches_name_filter(tool_name, include_set)
return not (exclude_set and matches_name_filter(tool_name, exclude_set))
return _should_register
if isinstance(include_raw, (str, list, tuple, set)):
return lambda tool_name: matches_name_filter(tool_name, include_set)
return lambda tool_name: not (exclude_set and matches_name_filter(tool_name, exclude_set))
class _CachedMCPTool:
@@ -192,9 +185,7 @@ def _utility_candidates(name: str, entries: Iterable[Any], tool_timeout) -> List
"""``{schema, handler_key}`` rows (live selection or cache) -> candidates; malformed rows dropped."""
out: List[_Candidate] = []
for raw in entries:
if not isinstance(raw, dict):
continue
schema, key = raw.get("schema"), raw.get("handler_key")
schema, key = (raw.get("schema"), raw.get("handler_key")) if isinstance(raw, dict) else (None, None)
if isinstance(schema, dict) and key in _UTILITY_HANDLER_FACTORIES and schema.get("name"):
handler = _UTILITY_HANDLER_FACTORIES[key](name, tool_timeout)
out.append(_Candidate(schema["name"], f"{_UTILITY_ORIGIN_PREFIX}{key!r}", schema, handler))
@@ -285,16 +276,12 @@ def _write_schema_cache(name: str, server: "MCPServerTask", config: dict, should
lazily without spawning it. Never raises."""
try:
from tools.mcp_schema_cache import config_fingerprint, write_cache_entry
tools_payload = []
for t in server._tools:
if should_register(t.name):
schema_obj = getattr(t, "inputSchema", None)
tools_payload.append({
"name": t.name,
"description": t.description or "",
"inputSchema": schema_obj if isinstance(schema_obj, dict) else {},
# Persisted so the lazy path trust-gates identically next startup.
"annotations": {"readOnlyHint": _annotation_read_only_hint(t)}})
tools_payload = [{
"name": t.name, "description": t.description or "",
"inputSchema": t.inputSchema if isinstance(getattr(t, "inputSchema", None), dict) else {},
# Persisted so the lazy path trust-gates identically next startup.
"annotations": {"readOnlyHint": _annotation_read_only_hint(t)},
} for t in server._tools if should_register(t.name)]
utility_payload = [{"schema": e["schema"], "handler_key": e["handler_key"]}
for e in _select_utility_schemas(name, server, config)]
cache_meta = getattr(server, "_list_cache_meta", None) or {}