182 lines
6.0 KiB
Python
182 lines
6.0 KiB
Python
"""Small pure helpers shared by the tools.mcp_tool_* modules: SDK 1.x/2.x field
|
|
access, error-text sanitising, numeric/bool coercion, timeouts and jitter. No
|
|
origin state."""
|
|
|
|
import logging
|
|
import math
|
|
import os
|
|
import random
|
|
import re
|
|
from typing import Any, Optional
|
|
|
|
logger = logging.getLogger("tools.mcp_tool")
|
|
|
|
|
|
class _OriginProxy:
|
|
"""Attribute proxy for ``tools.mcp_tool`` resolved at access time.
|
|
|
|
The split modules read origin state (``_servers``, ``_lock``, SDK symbols,
|
|
patchable helpers) through this so ``mock.patch("tools.mcp_tool.X")`` and
|
|
origin-side rebinds stay effective, and so no split module needs the origin
|
|
imported first (the origin imports them while it is still initialising).
|
|
"""
|
|
|
|
__slots__ = ()
|
|
|
|
def __getattr__(self, name: str):
|
|
from tools import mcp_tool
|
|
|
|
return getattr(mcp_tool, name)
|
|
|
|
|
|
_core = _OriginProxy()
|
|
|
|
|
|
_MISSING = object()
|
|
|
|
|
|
def mcp_field(obj, snake: str, camel: str, default=None):
|
|
"""Read an MCP model field across the 1.x -> 2.x rename to snake_case.
|
|
|
|
Pydantic aliases don't apply to attribute access, so ``getattr(result,
|
|
"isError", False)`` silently returns the default on 2.x — failed calls read
|
|
as successful, schemas as empty. Trying both spellings stays correct on
|
|
either SDK generation (``mcp`` is an optional extra at the user's version).
|
|
"""
|
|
value = getattr(obj, snake, _MISSING)
|
|
if value is not _MISSING:
|
|
return value
|
|
value = getattr(obj, camel, _MISSING)
|
|
return default if value is _MISSING else value
|
|
|
|
|
|
_DEFAULT_TOOL_TIMEOUT = 300 # seconds for tool calls
|
|
|
|
|
|
def _resolve_tool_timeout(config: dict) -> float:
|
|
"""Per-server tool-call timeout. Precedence: ``mcp_servers.<name>.timeout``
|
|
> ``timeouts.mcp.tool_call`` > the 300s default; values are platform-clamped
|
|
by ``resolve_timeout``."""
|
|
per_server = config.get("timeout")
|
|
if per_server is not None:
|
|
return per_server
|
|
try:
|
|
from agent.deadline import resolve_timeout
|
|
|
|
resolved = resolve_timeout("mcp.tool_call", default=_DEFAULT_TOOL_TIMEOUT)
|
|
if resolved is not None:
|
|
return resolved
|
|
except Exception:
|
|
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%
|
|
|
|
|
|
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"(?:"
|
|
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,
|
|
)
|
|
|
|
|
|
def _env_ref_name(ref: str) -> str:
|
|
"""Bare env-var name from a ``${...}`` body; strips a Cursor-style ``env:`` prefix."""
|
|
ref = ref.strip()
|
|
if ref.startswith("env:"):
|
|
ref = ref[len("env:"):].strip()
|
|
return ref
|
|
|
|
|
|
def _sanitize_error(text: str) -> str:
|
|
"""Replace credential-like patterns with [REDACTED] before text reaches the LLM."""
|
|
return _CREDENTIAL_PATTERN.sub("[REDACTED]", text)
|
|
|
|
|
|
def _exc_str(exc: BaseException) -> str:
|
|
"""Non-empty string for *exc*: some exceptions (``anyio.ClosedResourceError``)
|
|
carry no message, so fall back to ``repr`` to keep diagnostics."""
|
|
text = str(exc).strip()
|
|
return text or repr(exc)
|
|
|
|
|
|
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
|
|
|
|
existing = updated.get("PATH", "")
|
|
parts = [part for part in existing.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
|
|
|
|
|
|
def _safe_numeric(value, default, coerce=int, minimum=1):
|
|
"""Coerce a config value (YAML strings included) to a number, clamped to
|
|
*minimum*; *default* on failure or non-finite floats."""
|
|
try:
|
|
result = coerce(value)
|
|
if isinstance(result, float) and not math.isfinite(result):
|
|
return default
|
|
return max(result, minimum)
|
|
except (TypeError, ValueError, OverflowError):
|
|
return default
|
|
|
|
|
|
def _parse_boolish(value: Any, default: bool = True) -> bool:
|
|
"""Parse a bool-like config value with safe fallback."""
|
|
if value is None:
|
|
return default
|
|
if isinstance(value, bool):
|
|
return value
|
|
if isinstance(value, str):
|
|
lowered = value.strip().lower()
|
|
if lowered in {"true", "1", "yes", "on"}:
|
|
return True
|
|
if lowered in {"false", "0", "no", "off"}:
|
|
return False
|
|
logger.warning("MCP config expected a boolean-ish value, got %r; using default=%s", value, default)
|
|
return default
|
|
|
|
|
|
def _get_lifecycle_seconds(config: dict, key: str) -> Optional[float]:
|
|
"""Return an optional positive lifecycle timeout from top-level/nested config."""
|
|
raw = config.get(key)
|
|
lifecycle = config.get("lifecycle")
|
|
if raw is None and isinstance(lifecycle, dict):
|
|
raw = lifecycle.get(key)
|
|
if raw is None:
|
|
return None
|
|
try:
|
|
seconds = float(raw)
|
|
except (TypeError, ValueError):
|
|
logger.warning("MCP config %s must be a number of seconds; ignoring %r", key, raw)
|
|
return None
|
|
if seconds == 0:
|
|
return None
|
|
if seconds < 0:
|
|
logger.warning("MCP config %s must be positive; ignoring %r", key, raw)
|
|
return None
|
|
return seconds
|