fix(tools): narrow MCP OAuth lock scope
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user