refactor(tools/mcp_oauth): extract mcp_oauth_provider; compact oauth manager/dashboard bridge and schema_sanitizer; repoint tests

This commit is contained in:
Teknium
2026-09-02 14:41:58 -07:00
parent 8b3dde5dcc
commit 20258aad9e
7 changed files with 1081 additions and 1867 deletions
+16 -5
View File
@@ -15,16 +15,27 @@ from tools.mcp_oauth import (
OAuthNonInteractiveError,
build_oauth_auth,
remove_oauth_tokens,
_find_free_port,
_can_open_browser,
_is_interactive,
_wait_for_callback,
_make_callback_handler,
_make_redirect_handler,
_paste_callback_reader,
)
def _find_free_port() -> int:
import socket
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind(("127.0.0.1", 0))
return s.getsockname()[1]
async def _wait_for_callback():
"""Await the per-flow waiter on the legacy module-level port (the removed shim)."""
import tools.mcp_oauth as mod
return await mod._make_callback_waiter(mod._oauth_port)()
def _set_interactive_stdin(monkeypatch, *, is_tty: bool = True) -> None:
mock_stdin = MagicMock()
mock_stdin.isatty.return_value = is_tty
@@ -503,7 +514,7 @@ class TestCallbackPortReservation:
monkeypatch.setattr(mod, "_raise_if_non_interactive", lambda lead: None)
async def drive():
task = asyncio.create_task(mod._wait_for_callback())
task = asyncio.create_task(_wait_for_callback())
threading.Thread(
target=_hit_callback_when_ready,
args=(f"http://127.0.0.1:{port}/callback?code=abc123&state=xyz",),
@@ -814,7 +825,7 @@ class TestNonInteractiveFailFastAtCallbackBoundary:
monkeypatch.setattr(mod.asyncio, "sleep", no_sleep)
with pytest.raises(OAuthNonInteractiveError, match="interactive session"):
asyncio.run(mod._wait_for_callback())
asyncio.run(_wait_for_callback())
fake_server.assert_not_called()
def test_redirect_handler_rejects_and_does_not_open_browser(self, monkeypatch, capsys):
@@ -1072,7 +1083,7 @@ def test_wait_for_callback_port_in_use_reports_clear_error(monkeypatch):
mo, "HTTPServer", side_effect=OSError("address already in use")
):
with pytest.raises(mo.OAuthNonInteractiveError) as excinfo:
asyncio.run(mo._wait_for_callback())
asyncio.run(_wait_for_callback())
msg = str(excinfo.value)
assert "54321" in msg
+2 -2
View File
@@ -108,7 +108,7 @@ class TestManagerOAuthProviderMetadata:
provider = _manager_provider_with_context(storage, oauth_metadata=None)
with patch.object(
_HERMES_PROVIDER_CLS.__bases__[0], "_initialize", new=AsyncMock()
_HERMES_PROVIDER_CLS.__bases__[-1], "_initialize", new=AsyncMock()
):
asyncio.run(provider._initialize())
@@ -136,7 +136,7 @@ class TestManagerOAuthProviderMetadata:
manager.invalidate_if_disk_changed = AsyncMock(return_value=False)
with patch.object(
_HERMES_PROVIDER_CLS.__bases__[0],
_HERMES_PROVIDER_CLS.__bases__[-1],
"async_auth_flow",
new=fake_parent_flow,
), patch("tools.mcp_oauth_manager.get_manager", return_value=manager):
+13 -6
View File
@@ -40,6 +40,7 @@ class DashboardOAuthFlow:
_lock: threading.Lock = field(default_factory=threading.Lock, init=False, repr=False)
async def publish_authorization_url(self, url: str) -> None:
"""Record the SDK's authorization URL (with its ``state``) for the dashboard to show."""
state = parse_qs(urlparse(url).query).get("state", [None])[0]
if not state:
raise ValueError("OAuth authorization URL did not include state")
@@ -51,10 +52,15 @@ class DashboardOAuthFlow:
self.status = "authorization_required"
self._authorization_ready.set()
@staticmethod
async def _await_event(event: threading.Event, timeout: float, message: str) -> None:
if not await asyncio.to_thread(event.wait, timeout):
raise TimeoutError(message)
async def wait_for_authorization_url(self, timeout: float = 30.0) -> str:
ready = await asyncio.to_thread(self._authorization_ready.wait, timeout)
if not ready:
raise TimeoutError("Timed out waiting for MCP authorization URL")
await self._await_event(
self._authorization_ready, timeout, "Timed out waiting for MCP authorization URL"
)
if not self.authorization_url:
raise RuntimeError(self.error or "MCP OAuth flow ended before authorization")
return self.authorization_url
@@ -66,6 +72,7 @@ class DashboardOAuthFlow:
state: str | None,
error: str | None,
) -> None:
"""Hand the browser redirect to the waiting flow; ``state`` must match exactly."""
with self._lock:
if self._callback_ready.is_set():
raise ValueError("OAuth callback already received")
@@ -84,9 +91,9 @@ class DashboardOAuthFlow:
self._callback_ready.set()
async def wait_for_callback(self, timeout: float = 300.0) -> tuple[str, str | None]:
ready = await asyncio.to_thread(self._callback_ready.wait, timeout)
if not ready:
raise TimeoutError("Timed out waiting for MCP OAuth callback")
await self._await_event(
self._callback_ready, timeout, "Timed out waiting for MCP OAuth callback"
)
if self._callback_error:
raise RuntimeError(f"OAuth authorization failed: {self._callback_error}")
if self._callback is None:
+507 -1031
View File
File diff suppressed because it is too large Load Diff
+152 -411
View File
@@ -1,35 +1,20 @@
#!/usr/bin/env python3
"""Central manager for per-server MCP OAuth state.
One instance shared across the process. Holds per-server OAuth provider
instances and coordinates:
One instance per process. Holds per-server provider instances and coordinates:
- **Cross-process token reload** via mtime-based disk watch. When an external
process (e.g. a user cron job) refreshes tokens on disk, the next auth flow
picks them up without requiring a process restart.
- **401 deduplication** via in-flight futures. When N concurrent tool calls
all hit 401 with the same access_token, only one recovery attempt fires;
the rest await the same result.
- **Reconnect signalling** for long-lived MCP sessions. The manager itself
does not drive reconnection — the `MCPServerTask` in `mcp_tool.py` does —
but the manager is the single source of truth that decides when reconnect
is warranted.
- **Cross-process token reload** via mtime-based disk watch, so tokens
refreshed by another process (cron, another CLI) are picked up without a
restart (Claude Code's ``invalidateOAuthCacheIfDiskChanged`` bug class).
- **401 deduplication** via in-flight futures: N concurrent tool calls hitting
401 with the same access_token trigger one recovery attempt.
- **Reconnect signalling** — ``MCPServerTask`` in ``mcp_tool.py`` drives the
reconnect; the manager decides when it is warranted.
Replaces what used to be scattered across eight call sites in `mcp_oauth.py`,
`mcp_tool.py`, and `hermes_cli/mcp_config.py`. This module is the ONLY place
that instantiates the MCP SDK's `OAuthClientProvider` — all other code paths
go through `get_manager()`.
Design reference:
- Claude Code's ``invalidateOAuthCacheIfDiskChanged``
(``claude-code/src/utils/auth.ts:1320``, CC-1096 / GH#24317). Identical
external-refresh staleness bug class.
- Codex's ``refresh_oauth_if_needed`` / ``persist_if_needed``
(``codex-rs/rmcp-client/src/rmcp_client.rs:805``). We lean on the MCP SDK's
lazy refresh rather than calling refresh before every op, because one
``stat()`` per tool call is cheaper than an ``await`` + potential refresh
round-trip, and the SDK's in-memory expiry path is already correct.
This module is the ONLY place that instantiates the SDK's ``OAuthClientProvider``
for runtime use; other code paths go through ``get_manager()``. We lean on the
SDK's lazy refresh rather than refreshing before every op: one ``stat()`` per
tool call is cheaper than an await + refresh round-trip.
"""
from __future__ import annotations
@@ -72,23 +57,10 @@ def _same_endpoint(a: str, b: str) -> bool:
@dataclass
class _ProviderEntry:
"""Per-server OAuth state tracked by the manager.
Fields:
server_url: The MCP server URL used to build the provider. Tracked
so we can discard a cached provider if the URL changes.
oauth_config: Optional dict from ``mcp_servers.<name>.oauth``.
provider: The ``httpx.Auth``-compatible provider wrapping the MCP
SDK. None until first use.
last_mtime_ns: Last-seen ``st_mtime_ns`` of the on-disk tokens file.
Zero if never read. Used by :meth:`MCPOAuthManager.invalidate_if_disk_changed`
to detect external refreshes.
lock: Serialises concurrent access to this entry's state. Bound to
whichever asyncio loop first awaits it (the MCP event loop).
pending_401: In-flight 401-handler futures keyed by the failed
access_token, for deduplicating thundering-herd 401s. Mirrors
Claude Code's ``pending401Handlers`` map.
"""
"""Per-server OAuth state. ``last_mtime_ns`` is the last-seen tokens-file
mtime (0 = never read) for external-refresh detection; ``lock`` binds to
whichever asyncio loop first awaits it (the MCP event loop);
``pending_401`` dedupes thundering-herd 401s by failed access_token."""
server_url: str
oauth_config: Optional[dict]
@@ -104,191 +76,79 @@ class _ProviderEntry:
def _make_hermes_provider_class() -> Optional[type]:
"""Lazy-import the SDK base class and return our subclass.
Wrapped in a function so this module imports cleanly even when the
MCP SDK's OAuth module is unavailable (e.g. older mcp versions).
"""
"""Lazy-import the SDK base class and return our subclass (None if the
SDK's OAuth module is unavailable, so this module still imports)."""
try:
from mcp.client.auth.oauth2 import OAuthClientProvider
except ImportError: # pragma: no cover — SDK required in CI
return None
from tools.mcp_oauth_provider import HermesProviderMixin
class HermesMCPOAuthProvider(OAuthClientProvider):
class HermesMCPOAuthProvider(HermesProviderMixin, OAuthClientProvider):
"""OAuthClientProvider with pre-flow disk-mtime reload.
Before every ``async_auth_flow`` invocation, asks the manager to
check whether the tokens file on disk has been modified externally.
If so, the manager resets ``_initialized`` so the next flow
re-reads from storage.
This makes external-process refreshes (cron, another CLI instance)
visible to the running MCP session without requiring a restart.
Reference: Claude Code's ``invalidateOAuthCacheIfDiskChanged``
(``src/utils/auth.ts:1320``, CC-1096 / GH#24317).
Before every ``async_auth_flow`` the manager checks whether the tokens
file changed on disk and, if so, resets ``_initialized`` so the next
flow re-reads storage — making external refreshes visible to a running
session. Token-endpoint fixes come from ``HermesProviderMixin``.
"""
_hermes_logger = logger
def __init__(
self,
*args: Any,
server_name: str = "",
preregistered: bool = False,
token_user_agent: "str | None" = None,
**kwargs: Any,
):
super().__init__(*args, **kwargs)
# mcp 2.0.0 uses a task-owned anyio.Lock and holds it across the
# yielded resource request. A session-long GET therefore blocks
# every concurrent POST, and HTTPX may later close the auth-flow
# generator from a different task than the lock owner. A binary
# semaphore preserves mutual exclusion without task ownership;
# async_auth_flow below narrows its scope around resource I/O.
# mcp 2.0 uses a task-owned anyio.Lock held across the yielded
# resource request: a session-long GET blocks every concurrent
# POST, and HTTPX may close the auth-flow generator from another
# task. A binary semaphore keeps mutual exclusion without task
# ownership; async_auth_flow narrows its scope around resource I/O.
import anyio
self.context.lock = anyio.Semaphore(1, max_value=1)
self._hermes_server_name = server_name
self._hermes_home = ""
# When the client_id comes from config.yaml (pre-registered), an
# invalid_client rejection means the *config* is wrong — deleting
# client.json would just be re-seeded from config and re-running
# registration can't help. Only auto-heal dynamically-registered
# clients. See _maybe_flag_poisoned_client.
# A config-supplied (pre-registered) client_id rejected as
# invalid_client means the *config* is wrong — re-registration
# can't help, so only dynamically-registered clients auto-heal.
self._hermes_preregistered = preregistered
# oauth.user_agent — stamped onto token-endpoint requests only;
# some authorization servers/WAFs reject httpx's default (#75576).
self._hermes_token_user_agent = token_user_agent
def _stamp_token_user_agent(self, request):
ua = getattr(self, "_hermes_token_user_agent", None)
if ua:
request.headers["User-Agent"] = ua
return request
def _hermes_storage(self):
"""The context storage when it is a ``HermesTokenStorage``, else None."""
from tools.mcp_oauth import HermesTokenStorage
def _coerce_client_secret_post(self) -> None:
"""Use client_secret_post when dynamic registration returned a secret.
Some MCP OAuth providers, notably Supabase, return a
``client_secret`` from dynamic client registration but omit
``token_endpoint_auth_method``. The MCP SDK treats the missing
value as public-client auth (``none``), so token exchange omits the
secret and Supabase rejects it with ``Required parameter:
client_secret``. Coerce the in-memory client info before token and
refresh requests.
"""
info = getattr(self.context, "client_info", None)
if not info or not getattr(info, "client_secret", None):
return
method = getattr(info, "token_endpoint_auth_method", None)
if method not in (None, "none", ""):
return
from mcp.shared.auth import OAuthClientInformationFull
data = info.model_dump(mode="json", exclude_none=True)
data["token_endpoint_auth_method"] = "client_secret_post"
self.context.client_info = OAuthClientInformationFull.model_validate(data)
async def _exchange_token_authorization_code(self, *args: Any, **kwargs: Any):
self._coerce_client_secret_post()
request = await super()._exchange_token_authorization_code(*args, **kwargs)
return self._stamp_token_user_agent(request)
async def _refresh_token(self):
self._coerce_client_secret_post()
request = await super()._refresh_token()
return self._stamp_token_user_agent(request)
async def _handle_token_response(self, response):
"""Accept any 2xx token response and avoid leaking token bodies in errors."""
if 200 <= response.status_code < 300:
from mcp.client.auth.utils import handle_token_response_scopes
from mcp.client.auth.oauth2 import OAuthTokenError
from httpx import HTTPError
try:
token_response = await handle_token_response_scopes(response)
except (HTTPError, OAuthTokenError):
raise OAuthTokenError("Invalid token response") from None
self.context.current_tokens = token_response
self.context.update_token_expiry(token_response)
await self.context.storage.set_tokens(token_response)
return
from mcp.client.auth.oauth2 import OAuthTokenError
raise OAuthTokenError(f"Token exchange failed ({response.status_code})")
async def _handle_refresh_response(self, response) -> bool:
"""Accept any 2xx refresh response and avoid logging token bodies."""
if not (200 <= response.status_code < 300):
logger.warning("Token refresh failed: %s", response.status_code)
self.context.clear_tokens()
return False
from mcp.shared.auth import OAuthToken
from httpx import HTTPError
from pydantic import ValidationError
try:
content = await response.aread()
token_response = OAuthToken.model_validate_json(content)
self.context.current_tokens = token_response
self.context.update_token_expiry(token_response)
await self.context.storage.set_tokens(token_response)
return True
except (HTTPError, ValidationError):
logger.warning("Invalid refresh response: %s", response.status_code)
self.context.clear_tokens()
return False
storage = self.context.storage
return storage if isinstance(storage, HermesTokenStorage) else None
async def _initialize(self) -> None:
"""Load stored tokens + client info AND seed token_expiry_time.
"""Load stored state, seed ``token_expiry_time``, restore/prefetch metadata.
Also eagerly fetches OAuth authorization-server metadata (PRM +
ASM) when we have stored tokens but no cached metadata, so the
SDK's ``_refresh_token`` can build the correct token_endpoint
URL on the preemptive-refresh path. Without this, the SDK
falls back to ``{mcp_server_url}/token`` (wrong for providers
whose AS is a different origin — BetterStack's MCP lives at
``https://mcp.betterstack.com`` but its token endpoint is at
``https://betterstack.com/oauth/token``), the refresh 404s, and
we drop through to full browser reauth.
The SDK's ``_initialize`` populates ``current_tokens`` but never
calls ``update_token_expiry``, so ``is_token_valid()`` is True for
any loaded token regardless of age and a restarted process ships
stale Bearer tokens (some providers answer 200 with an app-level
auth error the transport can't see). Seeding the expiry makes the
SDK take ``can_refresh_token()`` and refresh before the first
request; ``HermesTokenStorage`` persists absolute ``expires_at`` so
the TTL reflects wall-clock age.
The SDK's base ``_initialize`` populates ``current_tokens`` but
does NOT call ``update_token_expiry``, so ``token_expiry_time``
stays ``None`` and ``is_token_valid()`` returns True for any
loaded token regardless of actual age. After a process restart
this ships stale Bearer tokens to the server; some providers
return HTTP 401 (caught by the 401 handler), others return 200
with an app-level auth error (invisible to the transport layer,
e.g. BetterStack returning "No teams found. Please check your
authentication.").
Seeding ``token_expiry_time`` from the reloaded token fixes that:
``is_token_valid()`` correctly reports False for expired tokens,
``async_auth_flow`` takes the ``can_refresh_token()`` branch,
and the SDK quietly refreshes before the first real request.
Paired with :class:`HermesTokenStorage` persisting an absolute
``expires_at`` timestamp (``mcp_oauth.py:set_tokens``) so the
remaining TTL we compute here reflects real wall-clock age.
Metadata is restored from disk, else discovered pre-flight when we
hold tokens but no metadata: otherwise ``_refresh_token`` guesses
``{server_url}/token`` (wrong for split-origin providers such as
BetterStack), the refresh 404s and we fall through to browser reauth.
"""
await super()._initialize()
tokens = self.context.current_tokens
if tokens is not None and tokens.expires_in is not None:
self.context.update_token_expiry(tokens)
# Cold-load: restore OAuth server metadata from disk before any
# refresh attempt. Without this, a restarted process with cached
# tokens but no in-memory metadata would fall back to the SDK's
# guessed ``{server_url}/token`` path (returns 404 on most real
# providers) and require a full browser re-authorization.
storage = self.context.storage
from tools.mcp_oauth import HermesTokenStorage
if (
isinstance(storage, HermesTokenStorage)
and self.context.oauth_metadata is None
):
storage = self._hermes_storage()
if storage is not None and self.context.oauth_metadata is None:
meta = storage.load_oauth_metadata()
if meta is not None:
self.context.oauth_metadata = meta
@@ -299,20 +159,11 @@ def _make_hermes_provider_class() -> Optional[type]:
meta.token_endpoint,
)
# Pre-flight OAuth AS discovery so ``_refresh_token`` has a
# correct ``token_endpoint`` before the first refresh attempt.
# Only runs when we have tokens on cold-load but no cached
# metadata — i.e. the exact scenario where the SDK's built-in
# 401-branch discovery hasn't had a chance to run yet.
if (
tokens is not None
and self.context.oauth_metadata is None
):
if tokens is not None and self.context.oauth_metadata is None:
try:
await self._prefetch_oauth_metadata()
except Exception as exc: # pragma: no cover — defensive
# Non-fatal: if discovery fails, the SDK's normal 401-
# branch discovery will run on the next request.
# Non-fatal: the SDK's 401-branch discovery runs next request.
logger.debug(
"MCP OAuth '%s': pre-flight metadata discovery "
"failed (non-fatal): %s",
@@ -320,18 +171,15 @@ def _make_hermes_provider_class() -> Optional[type]:
)
async def _prefetch_oauth_metadata(self) -> None:
"""Fetch PRM + ASM from the well-known endpoints, cache on context.
"""Fetch PRM + ASM from the well-known endpoints and cache on context.
Mirrors the SDK's 401-branch discovery (oauth2.py ~line 511-551)
but runs synchronously before the first request instead of
inside the httpx auth_flow generator. Uses the SDK's own URL
builders and response handlers so we track whatever the SDK
version we're pinned to expects.
Mirrors the SDK's 401-branch discovery but runs before the first
request. Uses the SDK's own URL builders/response handlers so we
track whatever the pinned SDK version expects.
"""
# The SDK's httpx flavour, not Hermes' — mcp 2.0 builds on httpx2,
# and `create_oauth_metadata_request` below returns one of *its*
# Request objects, which only its own AsyncClient can send. See
# tools.mcp_tool.sdk_httpx.
# The SDK's httpx flavour, not Hermes' — mcp 2.0 builds on httpx2 and
# `create_oauth_metadata_request` returns *its* Request objects,
# which only its own AsyncClient can send (tools.mcp_tool.sdk_httpx).
from tools.mcp_tool import sdk_httpx
httpx = sdk_httpx()
if httpx is None: # pragma: no cover — SDK import would have failed
@@ -345,53 +193,46 @@ def _make_hermes_provider_class() -> Optional[type]:
)
server_url = self.context.server_url
async def _send(client, url: str, label: str):
try:
return await client.send(create_oauth_metadata_request(url))
except httpx.HTTPError as exc:
logger.debug(
"MCP OAuth '%s': %s discovery to %s failed: %s",
self._hermes_server_name, label, url, exc,
)
return None
async with httpx.AsyncClient(timeout=10.0) as client:
# Step 1: PRM discovery to learn the authorization_server URL.
for url in build_protected_resource_metadata_discovery_urls(
None, server_url
):
req = create_oauth_metadata_request(url)
try:
resp = await client.send(req)
except httpx.HTTPError as exc:
logger.debug(
"MCP OAuth '%s': PRM discovery to %s failed: %s",
self._hermes_server_name, url, exc,
)
for url in build_protected_resource_metadata_discovery_urls(None, server_url):
resp = await _send(client, url, "PRM")
if resp is None:
continue
prm = await handle_protected_resource_response(resp)
if prm:
self.context.protected_resource_metadata = prm
if prm.authorization_servers:
self.context.auth_server_url = str(
prm.authorization_servers[0]
)
self.context.auth_server_url = str(prm.authorization_servers[0])
break
# Step 2: ASM discovery against the auth_server_url (or
# server_url fallback for legacy providers).
# Step 2: ASM discovery against auth_server_url (server_url
# fallback for legacy providers).
for url in build_oauth_authorization_server_metadata_discovery_urls(
self.context.auth_server_url, server_url
):
req = create_oauth_metadata_request(url)
try:
resp = await client.send(req)
except httpx.HTTPError as exc:
logger.debug(
"MCP OAuth '%s': ASM discovery to %s failed: %s",
self._hermes_server_name, url, exc,
)
resp = await _send(client, url, "ASM")
if resp is None:
continue
ok, asm = await handle_auth_metadata_response(resp)
if not ok:
break
if asm:
self.context.oauth_metadata = asm
# Persist immediately so a subsequent cold-load can
# skip discovery entirely.
storage = self.context.storage
from tools.mcp_oauth import HermesTokenStorage
if isinstance(storage, HermesTokenStorage):
# Persist now so a later cold-load skips discovery.
storage = self._hermes_storage()
if storage is not None:
storage.save_oauth_metadata(asm)
logger.debug(
"MCP OAuth '%s': pre-flight ASM discovered "
@@ -401,61 +242,37 @@ def _make_hermes_provider_class() -> Optional[type]:
break
def _persist_oauth_metadata_if_changed(self) -> None:
"""Persist discovered OAuth metadata for future process restarts.
Called after the SDK's normal 401-branch auth flow completes so
metadata discovered via the lazy path (not pre-flight) is also
saved. No-op when nothing to persist or metadata hasn't changed.
"""
"""Save metadata the SDK discovered lazily (401 branch) for future
restarts; no-op when absent, not our storage, or unchanged."""
meta = self.context.oauth_metadata
if meta is None:
return
storage = self.context.storage
from tools.mcp_oauth import HermesTokenStorage
if not isinstance(storage, HermesTokenStorage):
storage = self._hermes_storage()
if meta is None or storage is None:
return
existing = storage.load_oauth_metadata()
if (
existing is None
or str(existing.token_endpoint) != str(meta.token_endpoint)
):
if existing is None or str(existing.token_endpoint) != str(meta.token_endpoint):
storage.save_oauth_metadata(meta)
async def _maybe_flag_poisoned_client(self, response: Any) -> None:
"""Detect a dead client registration and force re-registration.
When the IdP rejects our ``client_id`` with ``invalid_client`` on
the token endpoint (token exchange or refresh), the cached client
registration is provably dead server-side. We delete ``client.json``
(+ stale metadata) so the SDK's next ``async_auth_flow`` takes the
``if not client_info`` branch and re-runs RFC 7591 dynamic client
registration. This addresses the recurring manual-reset ritual in
GH#36767 for the auto-detectable subset (token-endpoint rejection);
the browser-side "Redirect URI Mismatch" case has no HTTP signal
and is handled by ``hermes mcp reauth``.
An ``invalid_client`` rejection of our ``client_id`` at the token
endpoint (exchange or refresh) proves the cached registration is
dead server-side; delete ``client.json`` (+ stale metadata) so the
SDK re-runs DCR next flow. The browser-side "Redirect URI Mismatch"
case has no HTTP signal and is left to ``hermes mcp reauth``.
Conservative by construction — acts ONLY when all hold:
* status is 400/401,
* the request hit the discovered ``token_endpoint`` (the only
request carrying our ``client_id``), and
* the body carries the ``invalid_client`` error code
(word-boundary match, so RFC 7591's ``invalid_client_metadata``
registration error does not trip it).
Pre-registered (config-supplied) clients are never poisoned.
Fully best-effort: any failure here is swallowed so a detection
miss never breaks the live auth flow.
Covers both the authorization-code token exchange and the
preemptive refresh — but only when ``token_endpoint`` was
discovered (``_initialize`` prefetches it on cold-load). If that
discovery was skipped, the guard returns early and the user falls
back to ``hermes mcp reauth``.
Conservative by construction — acts ONLY when status is 400/401,
the request hit the discovered ``token_endpoint`` (the only request
carrying our ``client_id``), and the body carries ``invalid_client``
as a whole word (so RFC 7591's ``invalid_client_metadata`` does not
trip it). Pre-registered clients are never poisoned. Best-effort:
any failure is swallowed so a miss never breaks the live flow. If
``token_endpoint`` was never discovered the guard returns early.
"""
try:
if self._hermes_preregistered:
return
status = getattr(response, "status_code", None)
if status not in (400, 401):
if getattr(response, "status_code", None) not in (400, 401):
return
meta = getattr(self.context, "oauth_metadata", None)
token_endpoint = (
@@ -470,22 +287,15 @@ def _make_hermes_provider_class() -> Optional[type]:
if not _same_endpoint(req_url, token_endpoint):
return
body = await response.aread()
# Word-boundary match: matches `"error":"invalid_client"` but
# not the RFC 7591 registration error `invalid_client_metadata`
# (the trailing `_metadata` removes the right-hand boundary).
if not re.search(rb"\binvalid_client\b", body.lower()):
return
storage = self.context.storage
from tools.mcp_oauth import HermesTokenStorage
# When the rejected client_id was our Client ID Metadata
# Document URL, re-presenting it next flow would loop: the
# server has already fetched that document and refused it.
# Dropping the URL sends the retry down the DCR branch
# instead, and the marker on disk keeps the next process from
# walking back into the same refusal. `hermes mcp login`
# clears the marker, so a fixed document gets another chance.
storage = self._hermes_storage()
# If the rejected client_id was our CIMD URL, re-presenting it
# would loop (the server already fetched and refused it). Drop
# the URL so the retry takes DCR, and mark it on disk so the
# next process doesn't walk back into the same refusal
# (`hermes mcp login` clears the marker).
cimd_url = getattr(self.context, "client_metadata_url", None)
rejected_id = getattr(self.context.client_info, "client_id", None)
if cimd_url and rejected_id == cimd_url:
@@ -496,10 +306,10 @@ def _make_hermes_provider_class() -> Optional[type]:
self._hermes_server_name, cimd_url,
)
self.context.client_metadata_url = None
if isinstance(storage, HermesTokenStorage):
if storage is not None:
storage.mark_cimd_rejected()
if isinstance(storage, HermesTokenStorage):
if storage is not None:
storage.poison_client_registration()
# Drop the in-memory client so the SDK re-registers next flow.
self.context.client_info = None
@@ -511,9 +321,7 @@ def _make_hermes_provider_class() -> Optional[type]:
)
async def async_auth_flow(self, request): # type: ignore[override]
# Pre-flow hook: ask the manager to refresh from disk if needed.
# Any failure here is non-fatal — we just log and proceed with
# whatever state the SDK already has.
# Pre-flow hook: reload from disk if it changed (non-fatal on error).
try:
await get_manager().invalidate_if_disk_changed(
self._hermes_server_name,
@@ -525,20 +333,11 @@ def _make_hermes_provider_class() -> Optional[type]:
self._hermes_server_name, exc,
)
# Manually bridge the bidirectional generator protocol. httpx's
# auth_flow driver (httpx._client._send_handling_auth) calls
# ``auth_flow.asend(response)`` to feed HTTP responses back into
# the generator. A naive wrapper using ``async for item in inner:
# yield item`` DISCARDS those .asend(response) values and resumes
# the inner generator with None, so the SDK's
# ``response = yield request`` branch in
# mcp/client/auth/oauth2.py sees response=None and crashes at
# ``if response.status_code == 401`` with AttributeError.
#
# The bridge below forwards each .asend() value into the inner
# generator via inner.asend(incoming), preserving the bidirectional
# contract. Regression from PR #11383 caught by
# tests/tools/test_mcp_oauth_bidirectional.py.
# Bridge the bidirectional generator protocol by hand: httpx feeds
# responses back via ``auth_flow.asend(response)``. A naive
# ``async for item in inner: yield item`` DISCARDS those values, so
# the SDK's ``response = yield request`` sees None and crashes on
# ``response.status_code`` (tests/tools/test_mcp_oauth_bidirectional.py).
inner = super().async_auth_flow(request)
resource_lock_released = False
sent_access_token = None
@@ -546,10 +345,9 @@ def _make_hermes_provider_class() -> Optional[type]:
try:
outgoing = await inner.__anext__()
while True:
# The SDK holds context.lock for its entire generator,
# including while HTTPX waits on the actual MCP request.
# Release it only for that request. OAuth discovery,
# refresh, registration, and token exchange remain
# The SDK holds context.lock for its whole generator, even
# while HTTPX waits on the MCP request. Release it for that
# request only; discovery/refresh/registration/exchange stay
# serialized exactly as the SDK implements them.
if outgoing is request:
tokens = self.context.current_tokens
@@ -562,9 +360,8 @@ def _make_hermes_provider_class() -> Optional[type]:
if resource_lock_released:
await self.context.lock.acquire()
resource_lock_released = False
# A different request may have completed refresh or full
# authorization while this resource request was in
# flight. Retry with that token instead of starting a
# Another request may have refreshed/authorized while this
# one was in flight: retry with that token instead of a
# duplicate OAuth transition from the stale 401/403.
tokens = self.context.current_tokens
if (
@@ -577,22 +374,18 @@ def _make_hermes_provider_class() -> Optional[type]:
await inner.aclose()
retry_after_concurrent_auth = True
break
# Sniff the response for a dead-client-registration signal
# before handing it back to the SDK (best-effort, GH#36767).
# Sniff for a dead-client-registration signal (best-effort).
await self._maybe_flag_poisoned_client(incoming)
outgoing = await inner.asend(incoming)
except StopAsyncIteration:
# Persist any metadata the SDK discovered lazily during the
# 401 branch so a subsequent cold-load skips discovery.
# Persist metadata discovered lazily in the 401 branch.
self._persist_oauth_metadata_if_changed()
return
finally:
if resource_lock_released:
# Balance the SDK's surrounding ``async with`` even when
# HTTPX cancels or closes the flow while the resource
# request is still in flight. Shield only this local
# bookkeeping; general inner-generator teardown remains
# the separate concern tracked by the cleanup PR.
# HTTPX cancels/closes the flow mid-request. Shield only
# this local bookkeeping.
import anyio
with anyio.CancelScope(shield=True):
@@ -626,9 +419,8 @@ class MCPOAuthManager:
def __init__(self) -> None:
self._entries: dict[tuple[str, str], _ProviderEntry] = {}
self._entries_lock = threading.Lock()
# Holds strong references to in-flight 401 handler tasks so the
# event loop's weak-reference bookkeeping cannot GC them mid-run
# and leave `await pending` waiters hanging forever.
# Strong refs to in-flight 401 tasks so the loop's weak bookkeeping
# cannot GC them mid-run and leave `await pending` hanging forever.
self._inflight_tasks: set[asyncio.Task] = set()
# -- Provider construction / caching -------------------------------------
@@ -686,15 +478,8 @@ class MCPOAuthManager:
server_name: str,
entry: _ProviderEntry,
) -> Optional[Any]:
"""Build the underlying OAuth provider.
Constructs :class:`HermesMCPOAuthProvider` directly using the helpers
extracted from ``tools.mcp_oauth``. The subclass injects a pre-flow
disk-watch hook so external token refreshes (cron, other CLI
instances) are visible to running MCP sessions.
Returns None if the MCP SDK's OAuth support is unavailable.
"""
"""Build a :class:`HermesMCPOAuthProvider` from the shared
``tools.mcp_oauth`` helpers; None if the SDK's OAuth support is unavailable."""
if _HERMES_PROVIDER_CLS is None:
logger.warning(
"MCP OAuth '%s': SDK auth module unavailable", server_name,
@@ -702,30 +487,13 @@ class MCPOAuthManager:
return None
# Local imports avoid circular deps at module import time.
from tools.mcp_oauth import (
HermesTokenStorage,
OAuthNonInteractiveError,
_OAUTH_AVAILABLE,
_build_client_metadata,
_configure_callback_port,
_is_interactive,
_maybe_preregister_client,
_make_callback_waiter,
_make_redirect_handler,
cimd_provider_kwargs,
token_request_user_agent,
)
from tools.mcp_oauth import _OAUTH_AVAILABLE, OAuthNonInteractiveError, _is_interactive
from tools.mcp_oauth_provider import build_provider_kwargs, prepare_oauth_config
if not _OAUTH_AVAILABLE:
return None
cfg = dict(entry.oauth_config or {})
from tools.mcp_oauth import apply_oauth_provider_defaults
apply_oauth_provider_defaults(
cfg, server_name=server_name, server_url=entry.server_url
)
storage = HermesTokenStorage(server_name)
cfg, storage = prepare_oauth_config(server_name, entry.server_url, entry.oauth_config)
from tools.mcp_dashboard_oauth import get_dashboard_oauth_flow
@@ -742,29 +510,11 @@ class MCPOAuthManager:
"authorization."
)
_configure_callback_port(cfg, storage)
client_metadata = _build_client_metadata(cfg)
_maybe_preregister_client(storage, cfg, client_metadata)
resolved_port = cfg.get("_resolved_port", 0)
redirect_handler = _make_redirect_handler(resolved_port)
# mcp 2.0 removed OAuthClientProvider's `timeout` argument, so the
# configured `oauth.timeout` now bounds the callback waiter's own poll
# loop instead — that is where the browser round-trip is awaited.
callback_handler = _make_callback_waiter(
resolved_port, cfg.get("_cimd_url"), timeout=float(cfg.get("timeout", 300))
)
return _HERMES_PROVIDER_CLS(
server_name=server_name,
preregistered=bool(cfg.get("client_id")),
server_url=entry.server_url,
client_metadata=client_metadata,
storage=storage,
redirect_handler=redirect_handler,
callback_handler=callback_handler,
token_user_agent=token_request_user_agent(cfg),
**cimd_provider_kwargs(cfg),
**build_provider_kwargs(cfg, storage, ssh_proxy_hint=False),
)
def remove(
@@ -778,9 +528,7 @@ class MCPOAuthManager:
Called by ``hermes mcp remove <name>`` and (indirectly) by
``hermes mcp login <name>`` during forced re-auth.
"""
with self._entries_lock:
entry = self._entries.pop(self._key(server_name, hermes_home), None)
entry = self.evict(server_name, hermes_home=hermes_home)
from tools.mcp_oauth import remove_oauth_tokens
remove_oauth_tokens(server_name, hermes_home=hermes_home)
logger.info(
@@ -807,10 +555,10 @@ class MCPOAuthManager:
server_name: str,
*,
hermes_home: str | Path | None = None,
) -> None:
) -> _ProviderEntry | None:
"""Drop only the in-process provider, preserving persisted OAuth state."""
with self._entries_lock:
self._entries.pop(self._key(server_name, hermes_home), None)
return self._entries.pop(self._key(server_name, hermes_home), None)
# -- Disk watch ----------------------------------------------------------
@@ -820,13 +568,10 @@ class MCPOAuthManager:
*,
hermes_home: str | Path | None = None,
) -> bool:
"""If the tokens file on disk has a newer mtime than last-seen, force
the MCP SDK provider to reload its in-memory state.
"""Force the SDK provider to reload when the tokens file mtime changed.
Returns True if the cache was invalidated (mtime differed). This is
the core fix for the external-refresh workflow: a cron job writes
fresh tokens to disk, and on the next tool call the running MCP
session picks them up without a restart.
Returns True if invalidated. This is the external-refresh fix: a cron
job writes fresh tokens and the next tool call picks them up.
"""
from tools.mcp_oauth import _get_token_dir, _safe_filename
@@ -844,9 +589,8 @@ class MCPOAuthManager:
if mtime_ns != entry.last_mtime_ns:
old = entry.last_mtime_ns
entry.last_mtime_ns = mtime_ns
# Force the SDK's OAuthClientProvider to reload from storage
# on its next auth flow. `_initialized` is private API but
# stable across the MCP SDK versions we pin (>=1.26.0).
# `_initialized` is private SDK API but stable across the
# versions we pin (>=1.26.0); resetting it forces a reload.
if hasattr(entry.provider, "_initialized"):
entry.provider._initialized = False # noqa: SLF001
logger.info(
@@ -901,19 +645,16 @@ class MCPOAuthManager:
pending.set_result(True)
return
# Step 2: No disk change — if the SDK can refresh
# in-place, let the caller retry. The SDK's httpx.Auth
# flow will issue the refresh on the next request.
provider = entry.provider
ctx = getattr(provider, "context", None)
can_refresh = False
if ctx is not None:
can_refresh_fn = getattr(ctx, "can_refresh_token", None)
if callable(can_refresh_fn):
try:
can_refresh = bool(can_refresh_fn())
except Exception:
can_refresh = False
# Step 2: No disk change — if the SDK can refresh in
# place, let the caller retry (the httpx.Auth flow
# refreshes on the next request).
can_refresh_fn = getattr(
getattr(entry.provider, "context", None), "can_refresh_token", None
)
try:
can_refresh = bool(can_refresh_fn()) if callable(can_refresh_fn) else False
except Exception:
can_refresh = False
if not pending.done():
pending.set_result(can_refresh)
except Exception as exc: # pragma: no cover — defensive
+155
View File
@@ -0,0 +1,155 @@
"""Shared ``OAuthClientProvider`` customizations for Hermes MCP OAuth.
Two code paths build an SDK provider — ``tools.mcp_oauth.build_oauth_auth``
(legacy public API) and ``tools.mcp_oauth_manager.MCPOAuthManager`` — and both
need the same real-world fixes and the same config → constructor-kwargs
plumbing. This module holds that shared core once; the origin modules keep
their own subclass (logger name, disk-watch hooks) on top of it.
"""
from __future__ import annotations
import logging
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from tools.mcp_oauth import HermesTokenStorage
logger = logging.getLogger(__name__)
class HermesProviderMixin:
"""Token-endpoint fixes layered over the SDK's ``OAuthClientProvider``.
- Supabase-style dynamic registration returns a ``client_secret`` but omits
``token_endpoint_auth_method``; the SDK then treats the client as public,
omits the secret, and the token endpoint rejects the exchange (looping the
browser authorization page). Coerce the in-memory client info to
``client_secret_post`` right before token and refresh requests.
- ``token_user_agent`` (``oauth.user_agent``) is stamped onto token-endpoint
requests only — some authorization servers and WAFs reject httpx's default.
- Any 2xx token/refresh response is accepted, and token bodies never leak
into exception text or log output.
Must precede the SDK class in the MRO. Subclasses set ``_hermes_logger`` so
warnings keep their origin module's logger name.
"""
_hermes_logger: logging.Logger = logger
def __init__(self, *args: Any, token_user_agent: str | None = None, **kwargs: Any):
super().__init__(*args, **kwargs)
self._hermes_token_user_agent = token_user_agent
def _stamp_token_user_agent(self, request):
ua = getattr(self, "_hermes_token_user_agent", None) # tests build via __new__
if ua:
request.headers["User-Agent"] = ua
return request
def _coerce_client_secret_post(self) -> None:
info = self.context.client_info
if not info or not getattr(info, "client_secret", None):
return
if getattr(info, "token_endpoint_auth_method", None) not in (None, "none", ""):
return
from mcp.shared.auth import OAuthClientInformationFull
data = info.model_dump(mode="json", exclude_none=True)
data["token_endpoint_auth_method"] = "client_secret_post"
self.context.client_info = OAuthClientInformationFull.model_validate(data)
async def _exchange_token_authorization_code(self, *args: Any, **kwargs: Any):
self._coerce_client_secret_post()
request = await super()._exchange_token_authorization_code(*args, **kwargs)
return self._stamp_token_user_agent(request)
async def _refresh_token(self):
self._coerce_client_secret_post()
request = await super()._refresh_token()
return self._stamp_token_user_agent(request)
async def _store_tokens(self, token_response) -> None:
self.context.current_tokens = token_response
self.context.update_token_expiry(token_response)
await self.context.storage.set_tokens(token_response)
async def _handle_token_response(self, response):
"""Accept any 2xx token response; never echo the body into errors."""
from mcp.client.auth.oauth2 import OAuthTokenError
if not (200 <= response.status_code < 300):
raise OAuthTokenError(f"Token exchange failed ({response.status_code})")
from httpx import HTTPError
from mcp.client.auth.utils import handle_token_response_scopes
try:
token_response = await handle_token_response_scopes(response)
except (HTTPError, OAuthTokenError):
raise OAuthTokenError("Invalid token response") from None
await self._store_tokens(token_response)
async def _handle_refresh_response(self, response) -> bool:
"""Accept any 2xx refresh response; never log the body."""
if not (200 <= response.status_code < 300):
self._hermes_logger.warning("Token refresh failed: %s", response.status_code)
self.context.clear_tokens()
return False
from httpx import HTTPError
from mcp.shared.auth import OAuthToken
from pydantic import ValidationError
try:
token_response = OAuthToken.model_validate_json(await response.aread())
except (HTTPError, ValidationError):
self._hermes_logger.warning("Invalid refresh response: %s", response.status_code)
self.context.clear_tokens()
return False
await self._store_tokens(token_response)
return True
def prepare_oauth_config(
server_name: str, server_url: str, oauth_config: dict | None
) -> tuple[dict, "HermesTokenStorage"]:
"""Copy the ``oauth:`` block, apply provider defaults, open its token storage.
The copy matters: later steps record ``_resolved_port`` / ``_cimd_url`` in
the dict, which must never leak back into the caller's config.
"""
from tools import mcp_oauth as mo
cfg = dict(oauth_config or {})
mo.apply_oauth_provider_defaults(cfg, server_name=server_name, server_url=server_url)
return cfg, mo.HermesTokenStorage(server_name)
def build_provider_kwargs(
cfg: dict, storage: "HermesTokenStorage", *, ssh_proxy_hint: bool
) -> dict[str, Any]:
"""Resolve the callback port and return the shared provider constructor kwargs.
Runs the port → client-metadata → pre-registration sequence (order matters:
metadata needs the resolved port, pre-registration needs the metadata).
``ssh_proxy_hint`` lets the redirect handler tailor its remote-session hint
to a configured proxy ``redirect_uri``. Helpers are looked up on
``tools.mcp_oauth`` at call time so tests can patch them there.
"""
from tools import mcp_oauth as mo
port = mo._configure_callback_port(cfg, storage)
client_metadata = mo._build_client_metadata(cfg)
mo._maybe_preregister_client(storage, cfg, client_metadata)
redirect_uri = (cfg.get("redirect_uri") or None) if ssh_proxy_hint else None
return {
"client_metadata": client_metadata,
"storage": storage,
"redirect_handler": mo._make_redirect_handler(port, redirect_uri=redirect_uri),
# mcp 2.0 dropped OAuthClientProvider's own `timeout`; the configured
# `oauth.timeout` bounds the callback waiter's poll loop instead.
"callback_handler": mo._make_callback_waiter(
port, cfg.get("_cimd_url"), timeout=float(cfg.get("timeout", 300))
),
"token_user_agent": mo.token_request_user_agent(cfg),
**mo.cimd_provider_kwargs(cfg),
}
+236 -412
View File
@@ -1,37 +1,21 @@
"""Sanitize tool JSON schemas for broad LLM-backend compatibility.
Some local inference backends (notably llama.cpp's ``json-schema-to-grammar``
converter used to build GBNF tool-call parsers) are strict about what JSON
Schema shapes they accept. Schemas that OpenAI / Anthropic / most cloud
providers silently accept can make llama.cpp fail the entire request with:
Some backends are strict about JSON Schema shapes that OpenAI/Anthropic/most
cloud providers silently accept — llama.cpp's ``json-schema-to-grammar`` fails
the whole request (``Unrecognized schema: "object"``), Anthropic rejects
nullable ``anyOf`` at the top of ``input_schema``, Fireworks rejects ``default``
beside ``$ref``, OpenAI's Codex backend rejects top-level combinators. Known
hostile constructs:
HTTP 400: Unable to generate parser for this template.
Automatic parser generation failed: JSON schema conversion failed:
Unrecognized schema: "object"
* ``{"type": "object"}`` with no ``properties``.
* A bare string (``"object"``) where a schema dict belongs (malformed MCP output).
* ``"type": ["string", "null"]`` array types.
* ``anyOf``/``oneOf`` unions whose only purpose is to permit ``null``.
* ``default`` (etc.) alongside ``$ref`` — e.g. ``{"$ref": "#/$defs/Foo", "default": null}``.
The failure modes we've seen in the wild:
* ``{"type": "object"}`` with no ``properties`` — rejected as a node the
grammar generator can't constrain.
* A schema value that is the bare string ``"object"`` instead of a dict
(malformed MCP server output, e.g. ``additionalProperties: "object"``).
* ``"type": ["string", "null"]`` array types — many converters only accept
single-string ``type``.
* ``anyOf`` / ``oneOf`` unions whose only purpose is to permit ``null`` for
optional fields (common Pydantic/MCP shape). Anthropic rejects these at
the top of ``input_schema``; collapse them to the non-null branch.
* Unconstrained ``additionalProperties`` on objects with empty properties.
* ``default`` (and other annotation keywords) alongside ``$ref`` — strict
backends (Fireworks-hosted Kimi, JSON Schema draft-07 validators) reject
sibling keywords at the same level as ``$ref``. Common MCP/Pydantic shape
after nullable-union collapse::
{"$ref": "#/$defs/Foo", "default": null}
This module walks the final tool schema tree (after MCP-level normalization
and any per-tool dynamic rebuilds) and fixes the known-hostile constructs
in-place on a deep copy. It is intentionally conservative: it only modifies
shapes the LLM backend couldn't use anyway.
This module walks the final tool schema tree (after MCP normalization and any
per-tool dynamic rebuilds) and fixes those in place on a deep copy. It is
deliberately conservative: it only modifies shapes the backend couldn't use.
"""
from __future__ import annotations
@@ -39,33 +23,37 @@ from __future__ import annotations
import copy
import logging
import re
from typing import Any
from typing import Any, Callable
logger = logging.getLogger(__name__)
# Anthropic (and Bedrock/Vertex/Azure fronting it) reject tool input schemas
# whose property keys don't match this pattern. Cloudflare's flat API MCP
# ships 61 such keys (query-filter params like ``issue_class~neq`` and
# ``meta.<field>[<operator>]``) — one bad key anywhere in the tools array
# 400s the entire request.
# whose property keys don't match this pattern; one bad key anywhere in the
# tools array 400s the entire request (Cloudflare's MCP ships 61 such keys).
_PROP_KEY_RE = re.compile(r"^[a-zA-Z0-9_.-]{1,64}$")
_PROP_KEY_BAD_CHARS = re.compile(r"[^a-zA-Z0-9_.-]")
_UNION_KEYS = ("anyOf", "oneOf")
# Outer-node metadata carried onto a union's replacement node.
_UNION_META_KEYS = ("title", "description", "default", "examples")
def _empty_object() -> dict:
return {"type": "object", "properties": {}}
def sanitize_property_key(key: str) -> str:
"""Deterministically map an arbitrary property key to a conforming one."""
new = _PROP_KEY_BAD_CHARS.sub("_", key)[:64]
return new or "param"
return _PROP_KEY_BAD_CHARS.sub("_", key)[:64] or "param"
def _rename_property_keys(props: dict, path: str) -> dict[str, str]:
"""Return {original_key: conforming_key} for one properties dict.
Identity entries are omitted. Deterministic: keys are processed in
insertion order and collisions deduped with numeric suffixes, so the
model-visible schema AND the dispatch-time reverse map (computed
independently from the registry's original schema) always agree.
Identity entries are omitted. Deterministic (insertion order, numeric
suffixes on collision) so the model-visible schema and the dispatch-time
reverse map computed from the registry's original schema always agree.
"""
renames: dict[str, str] = {}
taken = {k for k in props if _PROP_KEY_RE.match(k)}
@@ -90,9 +78,8 @@ def _rename_property_keys(props: dict, path: str) -> dict[str, str]:
def unrename_tool_args(params_schema: Any, args: Any) -> Any:
"""Map sanitized property keys in model-emitted args back to wire names.
``params_schema`` is the ORIGINAL (unsanitized) parameters schema from the
registry. Recurses into object-typed values and array items so nested
renamed keys are restored too. Unknown keys pass through untouched.
``params_schema`` is the ORIGINAL (unsanitized) registry schema. Recurses
into object values and array items; unknown keys pass through untouched.
"""
if not isinstance(params_schema, dict) or not isinstance(args, dict):
return args
@@ -118,21 +105,11 @@ def unrename_tool_args(params_schema: Any, args: Any) -> Any:
def sanitize_tool_schemas(tools: list[dict]) -> list[dict]:
"""Return a copy of ``tools`` with each tool's parameter schema sanitized.
Input is an OpenAI-format tool list:
``[{"type": "function", "function": {"name": ..., "parameters": {...}}}]``
The returned list is a deep copy — callers can safely mutate it without
affecting the original registry entries.
"""
"""Return a deep-copied ``tools`` list (OpenAI format) with each tool's
parameter schema sanitized; callers may mutate the result freely."""
if not tools:
return tools
sanitized: list[dict] = []
for tool in tools:
sanitized.append(_sanitize_single_tool(tool))
return sanitized
return [_sanitize_single_tool(tool) for tool in tools]
def _sanitize_single_tool(tool: dict) -> dict:
@@ -143,34 +120,27 @@ def _sanitize_single_tool(tool: dict) -> dict:
return out
params = fn.get("parameters")
# Missing / non-dict parameters → substitute the minimal valid shape.
if not isinstance(params, dict):
fn["parameters"] = {"type": "object", "properties": {}}
if not isinstance(params, dict): # missing / non-dict → minimal valid shape
fn["parameters"] = _empty_object()
return out
fn["parameters"] = _sanitize_node(params, path=fn.get("name", "<tool>"))
# After recursion, guarantee the top-level is an object with properties.
top = fn["parameters"]
name = fn.get("name", "<tool>")
top = _sanitize_node(params, path=name)
# Guarantee the top level is an object with properties.
if not isinstance(top, dict):
fn["parameters"] = {"type": "object", "properties": {}}
top = _empty_object()
else:
if top.get("type") != "object":
top["type"] = "object"
if "properties" not in top or not isinstance(top.get("properties"), dict):
if not isinstance(top.get("properties"), dict):
top["properties"] = {}
# Final pass: collapse nullable anyOf/oneOf unions that the recursive
# sanitizer above leaves intact (it only handles the array-form
# ``type: [X, "null"]``). Keep the ``nullable: true`` hint so runtime
# argument coercion (``model_tools._schema_allows_null``) can still
# map a model-emitted ``"null"`` string to Python ``None``.
fn["parameters"] = strip_nullable_unions(fn["parameters"], keep_nullable_hint=True)
# Strip top-level combinators that strict backends (OpenAI's Codex
# endpoint at chatgpt.com/backend-api/codex) reject outright. Nested
# combinators inside properties are preserved.
fn["parameters"] = _strip_top_level_combinators(
fn["parameters"], path=fn.get("name", "<tool>")
)
fn["parameters"] = _strip_ref_siblings(fn["parameters"])
# Collapse nullable unions the recursive pass leaves intact (it only
# handles the array-form ``type: [X, "null"]``); keep ``nullable: true`` so
# runtime coercion (``model_tools._schema_allows_null``) still maps a
# model-emitted ``"null"`` string to Python ``None``.
top = strip_nullable_unions(top, keep_nullable_hint=True)
top = _strip_top_level_combinators(top, path=name)
fn["parameters"] = _strip_ref_siblings(top)
return out
@@ -179,26 +149,16 @@ _REF_FORBIDDEN_SIBLINGS = frozenset({"default"})
def _strip_ref_siblings(node: Any) -> Any:
"""Drop forbidden sibling keywords from nodes that carry ``$ref``.
Fireworks (and other draft-07-strict backends) fail tool requests with::
JSON Schema not supported: keyword(s) ['default'] not allowed at
the same level as $ref.
Nullable-union collapse and MCP ingestion can leave ``default`` on a
``$ref`` node; strip it recursively.
"""
"""Recursively drop forbidden sibling keywords from nodes carrying ``$ref``
(Fireworks: ``keyword(s) ['default'] not allowed at the same level as $ref``)."""
if isinstance(node, list):
return [_strip_ref_siblings(item) for item in node]
if not isinstance(node, dict):
return node
out = {key: _strip_ref_siblings(value) for key, value in node.items()}
if "$ref" in out:
for key in _REF_FORBIDDEN_SIBLINGS:
if key in out:
out.pop(key, None)
out.pop(key, None)
return out
@@ -206,22 +166,12 @@ _TOP_LEVEL_FORBIDDEN_KEYS = ("allOf", "anyOf", "oneOf", "enum", "not")
def _strip_top_level_combinators(params: dict, *, path: str = "<tool>") -> dict:
"""Drop combinator keywords from the top-level of a function parameters schema.
"""Drop combinator keywords from the TOP level of a parameters schema only.
OpenAI's Codex backend (``chatgpt.com/backend-api/codex``) is stricter
than the public Functions API and rejects requests with::
Invalid schema for function 'X': schema must have type 'object' and
not have 'oneOf'/'anyOf'/'allOf'/'enum'/'not' at the top level.
These keywords are typically used for conditional required-fields hints
(``allOf: [{if: ..., then: {required: [...]}}]``). Removing them at the
top level discards the hint but does not change which argument *values*
are valid — the tool handler always re-validates required fields.
Only the *top* level is stripped; combinators nested inside a property's
schema are preserved (the strict rule only applies to the outermost
parameters object).
OpenAI's Codex backend rejects ``oneOf/anyOf/allOf/enum/not`` at the top
level. They are usually conditional-required hints; dropping them does not
change which argument values are valid (handlers re-validate). Nested
combinators are preserved.
"""
if not isinstance(params, dict):
return params
@@ -237,36 +187,34 @@ def _strip_top_level_combinators(params: dict, *, path: str = "<tool>") -> dict:
return out
def _is_null_branch(item: Any) -> bool:
return isinstance(item, dict) and item.get("type") == "null"
def _carry_union_meta(outer: dict, replacement: dict, *, skip_default_on_ref: bool) -> None:
"""Copy outer-union metadata onto *replacement* where absent."""
for meta_key in _UNION_META_KEYS:
if meta_key in outer and meta_key not in replacement:
# ``default`` is illegal alongside ``$ref`` on strict backends.
if skip_default_on_ref and meta_key == "default" and "$ref" in replacement:
continue
replacement[meta_key] = outer[meta_key]
def strip_nullable_unions(
schema: Any,
*,
keep_nullable_hint: bool = True,
) -> Any:
"""Collapse ``anyOf`` / ``oneOf`` nullable unions to the non-null branch.
"""Collapse ``anyOf``/``oneOf`` nullable unions to the single non-null branch.
MCP / Pydantic optional fields commonly arrive as::
{"anyOf": [{"type": "string"}, {"type": "null"}], "default": null}
Anthropic's tool input-schema validator rejects the null branch. Tool
optionality is already represented by the parent object's ``required``
array, so we collapse the union to the single non-null variant.
Metadata (``title``, ``description``, ``default``, ``examples``) on the
outer union node is carried over to the replacement variant.
Args:
schema: JSON-Schema fragment (dict, list, or scalar).
keep_nullable_hint: If True, set ``nullable: true`` on the replacement
to preserve the "this field may be None" signal for downstream
consumers that care (e.g. runtime argument coercion that maps the
literal string ``"null"`` to Python ``None``). Anthropic's
validator accepts ``nullable: true`` but strict producers may
prefer False.
Returns:
The schema with nullable unions collapsed. Non-union nodes are
returned unchanged.
MCP/Pydantic optional fields arrive as
``{"anyOf": [{"type": "string"}, {"type": "null"}], "default": null}``;
Anthropic rejects the null branch, and optionality is already expressed by
the parent's ``required``. Only collapses when a null branch was dropped
AND exactly one non-null branch survives. Outer metadata is carried over.
``keep_nullable_hint`` sets ``nullable: true`` on the replacement for
downstream consumers (runtime ``"null"`` → ``None`` coercion).
"""
if isinstance(schema, list):
return [strip_nullable_unions(item, keep_nullable_hint=keep_nullable_hint) for item in schema]
@@ -277,27 +225,16 @@ def strip_nullable_unions(
k: strip_nullable_unions(v, keep_nullable_hint=keep_nullable_hint)
for k, v in schema.items()
}
for key in ("anyOf", "oneOf"):
for key in _UNION_KEYS:
variants = stripped.get(key)
if not isinstance(variants, list):
continue
non_null = [
item for item in variants
if not (isinstance(item, dict) and item.get("type") == "null")
]
# Only collapse when we actually dropped a null branch AND exactly
# one non-null branch survives (otherwise the union is meaningful
# and we leave it alone).
non_null = [item for item in variants if not _is_null_branch(item)]
if len(non_null) == 1 and len(non_null) != len(variants):
replacement = dict(non_null[0]) if isinstance(non_null[0], dict) else {}
if keep_nullable_hint:
replacement.setdefault("nullable", True)
for meta_key in ("title", "description", "default", "examples"):
if meta_key in stripped and meta_key not in replacement:
# ``default`` is illegal alongside ``$ref`` on strict backends.
if meta_key == "default" and "$ref" in replacement:
continue
replacement[meta_key] = stripped[meta_key]
_carry_union_meta(stripped, replacement, skip_default_on_ref=True)
return strip_nullable_unions(replacement, keep_nullable_hint=keep_nullable_hint)
return stripped
@@ -311,59 +248,41 @@ _CONST_PRIMITIVE_TYPES: dict[type, str] = {
def _const_branch_type(branch: Any) -> str | None:
"""Return the JSON-Schema primitive type of a pure ``const`` branch.
"""JSON-Schema primitive type of a pure ``const`` branch, else None.
A branch qualifies when it is a dict carrying ``const`` with a primitive
value, and any declared ``type`` matches the const value's type. Branch
metadata (``title``, ``description``) does not disqualify it, but any
other constraining keyword does. Returns ``None`` for non-qualifying
branches.
Qualifies when the dict carries a primitive ``const`` and any declared
``type`` matches it; ``title``/``description`` are allowed, any other
constraining keyword disqualifies.
"""
if not isinstance(branch, dict) or "const" not in branch:
return None
extra = set(branch) - {"const", "type", "title", "description"}
if extra:
if set(branch) - {"const", "type", "title", "description"}:
return None
value = branch["const"]
# bool is a subclass of int in Python; check it first so True/False never
# classify as integers.
for py_type, json_type in _CONST_PRIMITIVE_TYPES.items():
if type(value) is py_type:
declared = branch.get("type")
if declared is not None and declared != json_type:
return None
return json_type
return None
# ``type(value) is`` (not isinstance): bool is a subclass of int.
json_type = _CONST_PRIMITIVE_TYPES.get(type(value))
if json_type is None:
return None
declared = branch.get("type")
if declared is not None and declared != json_type:
return None
return json_type
def collapse_const_unions(schema: Any) -> Any:
"""Collapse ``anyOf`` / ``oneOf`` unions of same-typed consts to ``enum``.
"""Collapse ``anyOf``/``oneOf`` unions of same-typed consts to ``enum``.
Ported from block/goose ``tool_schema_normalize.rs`` (Apache-2.0).
Ported from block/goose ``tool_schema_normalize.rs`` (Apache-2.0). MCP
servers generated from Rust/TS union types emit
``{"anyOf": [{"const": "red"}, {"const": "green"}]}``; strict backends
mishandle these while ``{"type": "string", "enum": [...]}`` is universal.
MCP servers (particularly ones generated from Rust/TypeScript union types)
commonly emit closed value sets as const unions::
{"anyOf": [{"const": "red"}, {"const": "green"}, {"const": "blue"}]}
Strict tool-calling backends reject or mishandle these, while the
equivalent property-level ``enum`` form is universally supported::
{"type": "string", "enum": ["red", "green", "blue"]}
The collapse applies only when EVERY non-null branch is a pure ``const``
of the same primitive type (bool/int/float/str — ``bool`` never merges
with ``integer``). Mixed unions and non-uniform const types pass through
untouched. A single ``{"type": "null"}`` branch is tolerated: it is
dropped and recorded as ``nullable: true`` (matching the
``strip_nullable_unions`` convention), since strip_nullable_unions only
collapses unions with exactly one non-null branch and therefore leaves
null+multi-const unions for us.
Outer-node metadata (``title``, ``description``, ``default``,
``examples``) is carried onto the replacement. Enum order preserves
branch order, so output is deterministic and byte-stable across
discoveries. Input is never mutated.
Applies only when EVERY non-null branch is a pure ``const`` of one
primitive type (``bool`` never merges with ``integer``). One
``{"type": "null"}`` branch is tolerated and recorded as ``nullable: true``
(``strip_nullable_unions`` only handles single-non-null unions, so
null+multi-const unions land here). Enum order preserves branch order;
outer metadata is carried over; input is never mutated.
"""
if isinstance(schema, list):
return [collapse_const_unions(item) for item in schema]
@@ -371,13 +290,12 @@ def collapse_const_unions(schema: Any) -> Any:
return schema
out = {k: collapse_const_unions(v) for k, v in schema.items()}
for key in ("anyOf", "oneOf"):
for key in _UNION_KEYS:
variants = out.get(key)
if not isinstance(variants, list) or not variants:
continue
null_branches = [
item for item in variants
if isinstance(item, dict) and item.get("type") == "null" and "const" not in item
item for item in variants if _is_null_branch(item) and "const" not in item
]
const_branches = [item for item in variants if item not in null_branches]
if len(null_branches) > 1 or not const_branches:
@@ -391,46 +309,67 @@ def collapse_const_unions(schema: Any) -> Any:
}
if null_branches:
replacement["nullable"] = True
for meta_key in ("title", "description", "default", "examples"):
if meta_key in out and meta_key not in replacement:
replacement[meta_key] = out[meta_key]
_carry_union_meta(out, replacement, skip_default_on_ref=False)
return replacement
return out
_BARE_TYPE_NAMES = frozenset({"object", "string", "number", "integer", "boolean", "array", "null"})
# Sibling keywords whose values are NOT schemas: recursing would mistake literal
# strings like "path" for bare-string schemas. Passed through unchanged
# (``required`` remapped through property renames).
_NON_SCHEMA_LIST_KEYS = frozenset({"required", "enum", "examples", "dependentRequired"})
def _normalize_type_array(value: list, out: dict) -> None:
"""Normalize a ``type: [...]`` array into *out*.
Several backends reject array types (llama.cpp's grammar generator; Gemini
via OpenAI-compatible transports 400s). Per the AI-SDK behavior: one
non-null type → ``type: X`` (+ ``nullable`` if ``null`` present); several →
``anyOf`` of single-type schemas so EVERY branch survives; none → ``null``
or the object fallback. Ported from anomalyco/opencode#31877.
"""
has_null = "null" in value
non_null = [t for t in value if isinstance(t, str) and t != "null"]
if len(non_null) == 1:
out["type"] = non_null[0]
elif len(non_null) >= 2:
out["anyOf"] = [{"type": t} for t in non_null]
else:
out["type"] = "null" if has_null else "object"
return
if has_null:
out.setdefault("nullable", True)
def _sanitize_node(node: Any, path: str) -> Any:
"""Recursively sanitize a JSON-Schema fragment.
- Replaces bare-string schema values ("object", "string", ...) with
``{"type": <value>}`` so downstream consumers see a dict.
- Injects ``properties: {}`` into object-typed nodes missing it.
- Normalizes ``type: [X, "null"]`` arrays to single ``type: X`` (keeping
``nullable: true`` as a hint), and multi-type arrays like
``["number", "string"]`` to an ``anyOf`` of single-type schemas so no
branch is dropped (ported from anomalyco/opencode#31877).
- Bare-string schema values become ``{"type": <value>}`` (unknown strings
become a permissive object schema rather than something backends reject).
- Object-typed nodes gain ``properties: {}`` (llama.cpp can't constrain a
free-form object).
- ``type`` arrays are normalized (see ``_normalize_type_array``).
- Recurses into ``properties``, ``items``, ``additionalProperties``,
``anyOf``, ``oneOf``, ``allOf``, and ``$defs`` / ``definitions``.
``anyOf``/``oneOf``/``allOf`` and ``$defs``/``definitions``; property
keys are renamed to the provider-safe pattern and ``required`` follows.
- ``required`` entries that don't exist in ``properties`` are pruned
(malformed MCP schemas; built-in/plugin tools skip the MCP-level check).
"""
# Malformed: the schema position holds a bare string like "object".
if isinstance(node, str):
if node in {"object", "string", "number", "integer", "boolean", "array", "null"}:
if node in _BARE_TYPE_NAMES:
logger.debug(
"schema_sanitizer[%s]: replacing bare-string schema %r "
"with {'type': %r}",
path, node, node,
)
return {"type": node} if node != "object" else {
"type": "object",
"properties": {},
}
# Any other stray string is not a schema — drop it by replacing with
# a permissive object schema rather than propagate something the
# backend will reject.
return _empty_object() if node == "object" else {"type": node}
logger.debug(
"schema_sanitizer[%s]: replacing non-schema string %r "
"with empty object schema", path, node,
)
return {"type": "object", "properties": {}}
return _empty_object()
if isinstance(node, list):
return [_sanitize_node(item, f"{path}[{i}]") for i, item in enumerate(node)]
@@ -438,227 +377,70 @@ def _sanitize_node(node: Any, path: str) -> Any:
if not isinstance(node, dict):
return node
# Compute property-key renames up front so the ``required`` branch below
# can remap regardless of dict iteration order (``required`` may precede
# ``properties`` in the source dict).
# Renames are computed up front so ``required`` can be remapped even when
# it precedes ``properties`` in the source dict.
prop_renames: dict[str, str] = {}
if isinstance(node.get("properties"), dict):
prop_renames = _rename_property_keys(node["properties"], f"{path}.properties")
out: dict = {}
for key, value in node.items():
# JSON Schema ``type`` arrays (e.g. ``["number", "string"]``, common
# in MCP tool schemas) are rejected by several tool-call backends:
# * llama.cpp's grammar generator only accepts a singular string type.
# * Gemini (including OpenAI-compatible transports such as GitHub
# Copilot proxying to Gemini) rejects the array form outright —
# plain @ai-sdk/google rewrites it, but the OpenAI-compatible path
# forwards it verbatim and the backend 400s.
#
# Normalize per the SDK's behavior:
# * single non-null type → ``type: X`` (+ ``nullable: true`` if the
# array also contained "null"). No data lost.
# * multiple non-null types → ``anyOf`` of single-type schemas, so
# EVERY branch survives instead of silently dropping all but the
# first. ``null`` is lifted into ``nullable: true``.
# * all-null / empty → ``type: "null"`` (or object fallback).
# Ported from anomalyco/opencode#31877.
if key == "type" and isinstance(value, list):
has_null = "null" in value
non_null = [t for t in value if isinstance(t, str) and t != "null"]
if len(non_null) == 1:
out["type"] = non_null[0]
if has_null:
out.setdefault("nullable", True)
continue
if len(non_null) >= 2:
# Preserve all branches as a union instead of dropping them.
out["anyOf"] = [{"type": t} for t in non_null]
if has_null:
out.setdefault("nullable", True)
continue
# No usable non-null type: all-null array → type: "null";
# otherwise an empty/garbage array → object fallback.
out["type"] = "null" if has_null else "object"
continue
if key in {"properties", "$defs", "definitions"} and isinstance(value, dict):
_normalize_type_array(value, out)
elif key in {"properties", "$defs", "definitions"} and isinstance(value, dict):
renames = prop_renames if key == "properties" else {}
new_props = {}
for sub_k, sub_v in value.items():
out_k = renames.get(sub_k, sub_k)
new_props[out_k] = _sanitize_node(sub_v, f"{path}.{key}.{out_k}")
out[key] = new_props
out[key] = {
renames.get(sub_k, sub_k): _sanitize_node(sub_v, f"{path}.{key}.{renames.get(sub_k, sub_k)}")
for sub_k, sub_v in value.items()
}
elif key in {"items", "additionalProperties"}:
if isinstance(value, bool):
# Keep bool ``additionalProperties`` as-is — it's a valid form
# and widely accepted. ``items: true/false`` is non-standard
# but we preserve rather than drop.
out[key] = value
else:
out[key] = _sanitize_node(value, f"{path}.{key}")
# Bool ``additionalProperties`` is valid and widely accepted;
# ``items: true/false`` is non-standard but preserved rather than dropped.
out[key] = value if isinstance(value, bool) else _sanitize_node(value, f"{path}.{key}")
elif key in {"anyOf", "oneOf", "allOf"} and isinstance(value, list):
out[key] = [
_sanitize_node(item, f"{path}.{key}[{i}]")
for i, item in enumerate(value)
]
elif key in {"required", "enum", "examples", "dependentRequired"}:
# Schema "sibling" keywords whose values are NOT schemas:
# - ``required``: list of property-name strings
# - ``enum``: list of literal values (any JSON type)
# - ``examples``: list of example values (any JSON type)
# - ``dependentRequired``: mapping of property names to lists of
# required property-name strings (JSON Schema 2020-12)
# Recursing into these with _sanitize_node() would mis-interpret
# literal strings like "path" as bare-string schemas and replace
# them with {"type": "object"} dicts. Pass through unchanged
# (remapping ``required`` entries through the property renames).
out[key] = [_sanitize_node(item, f"{path}.{key}[{i}]") for i, item in enumerate(value)]
elif key in _NON_SCHEMA_LIST_KEYS:
if key == "required" and prop_renames and isinstance(value, list):
out[key] = [prop_renames.get(r, r) if isinstance(r, str) else r
for r in value]
out[key] = [prop_renames.get(r, r) if isinstance(r, str) else r for r in value]
else:
out[key] = copy.deepcopy(value) if isinstance(value, (list, dict)) else value
else:
out[key] = _sanitize_node(value, f"{path}.{key}") if isinstance(value, (dict, list)) else value
# Object nodes without properties: inject empty properties dict.
# llama.cpp's grammar generator can't constrain a free-form object.
if out.get("type") == "object" and not isinstance(out.get("properties"), dict):
out["properties"] = {}
# Prune ``required`` entries that don't exist in properties (defense
# against malformed MCP schemas; also caught upstream for MCP tools, but
# built-in tools or plugin tools may not have been through that path).
if out.get("type") == "object" and isinstance(out.get("required"), list):
props = out.get("properties") or {}
valid = [r for r in out["required"] if isinstance(r, str) and r in props]
if not valid:
out.pop("required", None)
elif len(valid) != len(out["required"]):
out["required"] = valid
if out.get("type") == "object":
if not isinstance(out.get("properties"), dict):
out["properties"] = {}
if isinstance(out.get("required"), list):
props = out.get("properties") or {}
valid = [r for r in out["required"] if isinstance(r, str) and r in props]
if not valid:
out.pop("required", None)
elif len(valid) != len(out["required"]):
out["required"] = valid
return out
# =============================================================================
# Reactive strip — only invoked when llama.cpp rejects a schema
# Reactive strips — only invoked after a backend rejects a schema
# =============================================================================
_STRIP_ON_RECOVERY_KEYS = frozenset({"pattern", "format"})
def strip_pattern_and_format(tools: list[dict]) -> tuple[list[dict], int]:
"""Strip ``pattern`` and ``format`` JSON Schema keywords from tool schemas.
This is a *reactive* sanitizer invoked only when llama.cpp's
``json-schema-to-grammar`` converter has rejected a tool schema with an
HTTP 400 grammar-parse error. llama.cpp's regex engine supports only a
small subset of ECMAScript regex (literals, ``.``, ``[...]``, ``|``,
``*``, ``+``, ``?``, ``{n,m}``) — it rejects escape classes like ``\\d``,
``\\w``, ``\\s`` and most ``format`` values. Cloud providers (OpenAI,
Anthropic, OpenRouter, Gemini) accept these keywords fine and rely on
them as prompting hints, so we keep them in the default schema and only
strip on demand.
The strip operates on a sibling of ``type`` (so schema keywords are
removed) — a property literally *named* ``pattern`` (e.g. the first arg
of the built-in ``search_files`` tool) is not affected because property
names live in the ``properties`` dict, not as siblings of ``type``.
Args:
tools: OpenAI-format tool list, mutated in place for efficiency.
Callers that need to preserve the original should deep-copy first.
Returns:
``(tools, stripped_count)`` — the same list reference plus a count of
how many ``pattern``/``format`` keywords were removed across all tools.
"""
def _reactive_strip(tools: list[dict], strip_node: Callable[[dict], int], log_msg: str) -> tuple[list[dict], int]:
"""Walk every tool's parameters in place, applying *strip_node* to each dict
node (it returns how many keywords it removed). Handles OpenAI format
(``{"function": {"parameters": ...}}``) and Responses format
(``{"name": ..., "parameters": ...}`` — codex_responses mode, xAI, etc.).
Returns ``(tools, stripped_count)`` — the same list reference."""
if not tools:
return tools, 0
stripped = 0
def _walk(node: Any) -> None:
nonlocal stripped
if isinstance(node, dict):
# Only strip as a sibling of ``type`` — i.e. when this node is
# itself a schema. This avoids stripping literal property keys
# named "pattern" (search_files.pattern, etc.) because those live
# inside a ``properties`` dict, not as siblings of ``type``.
is_schema_node = "type" in node or "anyOf" in node or "oneOf" in node or "allOf" in node
for key in list(node.keys()):
if is_schema_node and key in _STRIP_ON_RECOVERY_KEYS:
node.pop(key, None)
stripped += 1
continue
_walk(node[key])
elif isinstance(node, list):
for item in node:
_walk(item)
for tool in tools:
if not isinstance(tool, dict):
continue
# OpenAI-format: {"function": {"parameters": {...}}}
fn = tool.get("function")
if isinstance(fn, dict):
params = fn.get("parameters")
if isinstance(params, dict):
_walk(params)
continue
# Responses-format: {"name": "...", "parameters": {...}}
# (used by codex_responses API mode — xAI, OpenAI Codex, etc.)
params = tool.get("parameters")
if isinstance(params, dict):
_walk(params)
continue
if stripped:
logger.info(
"schema_sanitizer: stripped %d pattern/format keyword(s) from "
"tool schemas (llama.cpp grammar-parse recovery)",
stripped,
)
return tools, stripped
def strip_slash_enum(tools: list[dict]) -> tuple[list[dict], int]:
"""Strip ``enum`` keywords whose string values contain a forward slash.
xAI's ``/v1/responses`` and ``/v1/chat/completions`` endpoints compile
tool schemas to a grammar that rejects ``enum`` values containing ``/``
(the request fails with HTTP 400 "Invalid arguments passed to the
model" before any token is emitted). Most commonly hit by MCP-derived
tools whose enum lists HuggingFace model IDs (``Qwen/Qwen3.5-0.8B``,
``openai/gpt-oss-20b``) or owner/name environment IDs. The constraint
is purely a prompting hint; dropping it lets the model still see the
field description and pick a value, without xAI tripping on the slash.
Args:
tools: OpenAI-format or Responses-format tool list, mutated in
place. Callers that need to preserve the original should
deep-copy first.
Returns:
``(tools, stripped_count)`` — same list reference plus a count of
how many ``enum`` keywords were removed.
"""
if not tools:
return tools, 0
stripped = 0
def _walk(node: Any) -> None:
nonlocal stripped
if isinstance(node, dict):
enum_val = node.get("enum")
if isinstance(enum_val, list) and any(
isinstance(v, str) and "/" in v for v in enum_val
):
node.pop("enum", None)
stripped += 1
stripped += strip_node(node)
for v in node.values():
_walk(v)
elif isinstance(node, list):
@@ -669,19 +451,61 @@ def strip_slash_enum(tools: list[dict]) -> tuple[list[dict], int]:
if not isinstance(tool, dict):
continue
fn = tool.get("function")
if isinstance(fn, dict):
params = fn.get("parameters")
if isinstance(params, dict):
_walk(params)
continue
params = tool.get("parameters")
if isinstance(params, dict):
_walk(params)
if isinstance(fn, dict) and isinstance(fn.get("parameters"), dict):
_walk(fn["parameters"])
continue
if isinstance(tool.get("parameters"), dict):
_walk(tool["parameters"])
if stripped:
logger.info(
"schema_sanitizer: stripped %d enum keyword(s) containing '/' "
"from tool schemas (xAI Responses grammar-compile recovery)",
stripped,
)
logger.info(log_msg, stripped)
return tools, stripped
def strip_pattern_and_format(tools: list[dict]) -> tuple[list[dict], int]:
"""Strip ``pattern``/``format`` keywords from tool schemas, in place.
Reactive: invoked only after llama.cpp's grammar converter rejected a
schema with HTTP 400. Its regex engine supports a small ECMAScript subset
(no ``\\d``/``\\w``/``\\s``) and most ``format`` values; cloud providers rely
on these as prompting hints, so they stay in the default schema.
Only strips as a sibling of ``type``/combinators (i.e. on schema nodes), so
a property literally *named* ``pattern`` (``search_files``) is untouched —
property names live inside ``properties``, not beside ``type``.
"""
def _strip(node: dict) -> int:
if not ("type" in node or "anyOf" in node or "oneOf" in node or "allOf" in node):
return 0
hits = [k for k in node if k in _STRIP_ON_RECOVERY_KEYS]
for k in hits:
node.pop(k, None)
return len(hits)
return _reactive_strip(
tools, _strip,
"schema_sanitizer: stripped %d pattern/format keyword(s) from "
"tool schemas (llama.cpp grammar-parse recovery)",
)
def strip_slash_enum(tools: list[dict]) -> tuple[list[dict], int]:
"""Strip ``enum`` keywords whose string values contain ``/``, in place.
xAI's ``/v1/responses`` and ``/v1/chat/completions`` compile schemas to a
grammar that rejects ``/`` in enum values (HTTP 400 before any token) —
typically MCP enums of HuggingFace model IDs or owner/name env IDs. The
constraint is a prompting hint only; the model still sees the description.
"""
def _strip(node: dict) -> int:
enum_val = node.get("enum")
if isinstance(enum_val, list) and any(isinstance(v, str) and "/" in v for v in enum_val):
node.pop("enum", None)
return 1
return 0
return _reactive_strip(
tools, _strip,
"schema_sanitizer: stripped %d enum keyword(s) containing '/' "
"from tool schemas (xAI Responses grammar-compile recovery)",
)