fix(tools): narrow MCP OAuth lock scope

This commit is contained in:
Gille
2026-08-28 15:04:56 -06:00
committed by Teknium
parent f7c79efbac
commit 9a1eef7a29
3 changed files with 163 additions and 3 deletions
+101
View File
@@ -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
+4 -3
View File
@@ -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)
+58
View File
@@ -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