diff --git a/tests/tools/test_mcp_oauth_bidirectional.py b/tests/tools/test_mcp_oauth_bidirectional.py index ea5eccbc2c..7759c465c5 100644 --- a/tests/tools/test_mcp_oauth_bidirectional.py +++ b/tests/tools/test_mcp_oauth_bidirectional.py @@ -28,6 +28,8 @@ the bridge forwards responses correctly into the inner SDK generator. """ from __future__ import annotations +import asyncio + import pytest @@ -206,6 +208,105 @@ async def test_hermes_provider_forwards_401_triggers_refresh(tmp_path, monkeypat await flow.aclose() +@pytest.mark.asyncio +async def test_long_lived_resource_request_does_not_block_concurrent_post( + tmp_path, monkeypatch +): + """A session-long MCP GET must not hold the provider state lock. + + MCP SDK 2.0.0 wraps its entire auth-flow generator in one lock. Leaving + the GET response pending then prevents a concurrent POST from even + acquiring its Bearer token. Hermes narrows that lock around resource I/O + while retaining it for OAuth state transitions. + """ + from tools.mcp_tool import sdk_httpx + httpx = sdk_httpx() + from mcp.shared.auth import OAuthClientInformationFull, OAuthClientMetadata, OAuthToken + from pydantic import AnyUrl + + from tools.mcp_oauth import HermesTokenStorage + from tools.mcp_oauth_manager import _HERMES_PROVIDER_CLS, reset_manager_for_tests + + assert _HERMES_PROVIDER_CLS is not None + + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + reset_manager_for_tests() + + storage = HermesTokenStorage("srv") + await storage.set_tokens( + OAuthToken( + access_token="access-token", + token_type="Bearer", + expires_in=3600, + refresh_token="refresh-token", + ) + ) + await storage.set_client_info( + OAuthClientInformationFull( + client_id="test-client", + redirect_uris=[AnyUrl("http://127.0.0.1:12345/callback")], + grant_types=["authorization_code", "refresh_token"], + response_types=["code"], + token_endpoint_auth_method="none", + ) + ) + + provider = _HERMES_PROVIDER_CLS( + server_name="srv", + server_url="https://example.com/mcp", + client_metadata=OAuthClientMetadata( + redirect_uris=[AnyUrl("http://127.0.0.1:12345/callback")], + client_name="Hermes Agent", + ), + storage=storage, + redirect_handler=_noop_redirect, + callback_handler=_noop_callback, + ) + + get_request = httpx.Request("GET", "https://example.com/mcp") + get_flow = provider.async_auth_flow(get_request) + get_outbound = await get_flow.__anext__() + + # Keep the GET open, as streamable HTTP does for the session lifetime. + # The POST must still authenticate and reach HTTPX without waiting for it. + post_request = httpx.Request("POST", "https://example.com/mcp") + post_flow = provider.async_auth_flow(post_request) + post_outbound = await asyncio.wait_for(post_flow.__anext__(), timeout=2.0) + + assert post_outbound is post_request + assert post_outbound.headers["authorization"] == "Bearer access-token" + + with pytest.raises(StopAsyncIteration): + await post_flow.asend(httpx.Response(200, request=post_outbound)) + + # A completed concurrent refresh must make the pending GET retry with the + # new token rather than start a second OAuth transition from its stale 401. + provider.context.current_tokens = OAuthToken( + access_token="new-access-token", + token_type="Bearer", + expires_in=3600, + refresh_token="refresh-token", + ) + provider.context.update_token_expiry(provider.context.current_tokens) + get_retry = await get_flow.asend( + httpx.Response( + 401, + request=get_outbound, + headers={ + "www-authenticate": ( + 'Bearer resource_metadata="https://example.com/' + '.well-known/oauth-protected-resource"' + ) + }, + ) + ) + assert get_retry is get_request + assert get_retry.headers["authorization"] == "Bearer new-access-token" + + with pytest.raises(StopAsyncIteration): + await get_flow.asend(httpx.Response(200, request=get_retry)) + + async def _noop_redirect(_url: str) -> None: """Redirect handler that does nothing (won't be invoked in these tests).""" return None diff --git a/tests/tools/test_mcp_oauth_manager.py b/tests/tools/test_mcp_oauth_manager.py index 80b6eebda4..457f4b8f80 100644 --- a/tests/tools/test_mcp_oauth_manager.py +++ b/tests/tools/test_mcp_oauth_manager.py @@ -334,9 +334,10 @@ def test_bridge_forwards_requests_and_poisons_on_token_endpoint_400( async def fake_base_flow(self, request): # Mimic the SDK: yield the request, receive the response, then finish. - forwarded.append(("out", request)) - response = yield request - forwarded.append(("in", response)) + async with self.context.lock: + forwarded.append(("out", request)) + response = yield request + forwarded.append(("in", response)) from mcp.client.auth.oauth2 import OAuthClientProvider monkeypatch.setattr(OAuthClientProvider, "async_auth_flow", fake_base_flow) diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index d0bca49967..ee681a6335 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -138,6 +138,15 @@ def _make_hermes_provider_class() -> Optional[type]: **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. + 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 @@ -531,10 +540,43 @@ def _make_hermes_provider_class() -> Optional[type]: # contract. Regression from PR #11383 caught by # tests/tools/test_mcp_oauth_bidirectional.py. inner = super().async_auth_flow(request) + resource_lock_released = False + sent_access_token = None + retry_after_concurrent_auth = False 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 + # serialized exactly as the SDK implements them. + if outgoing is request: + tokens = self.context.current_tokens + sent_access_token = ( + tokens.access_token if tokens is not None else None + ) + self.context.lock.release() + resource_lock_released = True incoming = yield outgoing + 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 + # duplicate OAuth transition from the stale 401/403. + tokens = self.context.current_tokens + if ( + getattr(incoming, "status_code", None) in (401, 403) + and self.context.is_token_valid() + and tokens is not None + and tokens.access_token != sent_access_token + ): + self._add_auth_header(request) + 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). await self._maybe_flag_poisoned_client(incoming) @@ -544,6 +586,22 @@ def _make_hermes_provider_class() -> Optional[type]: # 401 branch so a subsequent cold-load skips discovery. 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. + import anyio + + with anyio.CancelScope(shield=True): + await self.context.lock.acquire() + + if retry_after_concurrent_auth: + yield request + self._persist_oauth_metadata_if_changed() + return return HermesMCPOAuthProvider