refactor(tools/mcp_oauth): extract mcp_oauth_provider; compact oauth manager/dashboard bridge and schema_sanitizer; repoint tests
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
+152
-411
@@ -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
|
||||
|
||||
@@ -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
@@ -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)",
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user