From b3273d0eb45c59602b1f76f4d97bb0c5da61e5d7 Mon Sep 17 00:00:00 2001 From: Carl Taylor Date: Tue, 4 Aug 2026 12:36:02 +1000 Subject: [PATCH 1/3] fix(oauth): avoid lock contention for long-running requests --- src/mcp/client/auth/oauth2.py | 107 +++++++-- tests/client/test_auth.py | 396 ++++++++++++++++++++++++++++++++++ 2 files changed, 483 insertions(+), 20 deletions(-) diff --git a/src/mcp/client/auth/oauth2.py b/src/mcp/client/auth/oauth2.py index 0ec0879688..2c9014e519 100644 --- a/src/mcp/client/auth/oauth2.py +++ b/src/mcp/client/auth/oauth2.py @@ -10,7 +10,8 @@ import secrets import string import time -from collections.abc import AsyncGenerator, Awaitable, Callable +from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable +from contextlib import asynccontextmanager from dataclasses import dataclass, field from typing import Any, Protocol from urllib.parse import quote, urlencode, urljoin, urlparse @@ -115,7 +116,13 @@ class OAuthContext: token_expiry_time: float | None = None # State - lock: anyio.Lock = field(default_factory=anyio.Lock) + # Semaphores are intentionally task-agnostic: HTTPX can resume or close an + # auth-flow generator from a different task than the one that yielded it. + # Normal resource requests are still yielded outside these critical sections. + lock: anyio.Semaphore = field(default_factory=lambda: anyio.Semaphore(1, max_value=1)) + # Refresh and authorization transitions share one single-flight lock while + # normal resource requests remain independent. + flow_lock: anyio.Semaphore = field(default_factory=lambda: anyio.Semaphore(1, max_value=1)) def get_authorization_base_url(self, server_url: str) -> str: """Extract base URL by removing path component.""" @@ -488,30 +495,86 @@ async def _handle_oauth_metadata_response(self, response: httpx.Response) -> Non metadata = OAuthMetadata.model_validate_json(content) self.context.oauth_metadata = metadata - async def async_auth_flow(self, request: httpx.Request) -> AsyncGenerator[httpx.Request, httpx.Response]: - """HTTPX auth flow integration.""" + async def _prepare_request(self, request: httpx.Request) -> tuple[bool, str | None]: + """Initialize state and capture request-specific protocol state.""" + protocol_version = request.headers.get(MCP_PROTOCOL_VERSION) async with self.context.lock: if not self._initialized: await self._initialize() # pragma: no cover - # Capture protocol version from request headers - self.context.protocol_version = request.headers.get(MCP_PROTOCOL_VERSION) - - if not self.context.is_token_valid() and self.context.can_refresh_token(): - # Try to refresh token - refresh_request = await self._refresh_token() # pragma: no cover - refresh_response = yield refresh_request # pragma: no cover + self.context.protocol_version = protocol_version + needs_refresh = not self.context.is_token_valid() and self.context.can_refresh_token() + return needs_refresh, protocol_version - if not await self._handle_refresh_response(refresh_response): # pragma: no cover - # Refresh failed, need full re-authentication - self._initialized = False - - if self.context.is_token_valid(): + async def _add_valid_auth_header(self, request: httpx.Request) -> str | None: + """Add the current valid token and return the token that was sent.""" + async with self.context.lock: + current_tokens = self.context.current_tokens + if self.context.is_token_valid() and current_tokens is not None: self._add_auth_header(request) + return current_tokens.access_token + return None - response = yield request + @asynccontextmanager + async def _serialized_transition(self) -> AsyncIterator[None]: + """Serialize token-changing OAuth transitions and their state writes.""" + async with self.context.flow_lock: + async with self.context.lock: + yield - if response.status_code == 401: + async def async_auth_flow(self, request: httpx.Request) -> AsyncGenerator[httpx.Request, httpx.Response]: + """HTTPX auth flow integration.""" + needs_refresh, protocol_version = await self._prepare_request(request) + + if needs_refresh: + async with self.context.flow_lock: + refresh_request: httpx.Request | None = None + async with self.context.lock: + self.context.protocol_version = protocol_version + # Another request may have refreshed the token while this + # request was waiting for the single-flight refresh lock. + if not self.context.is_token_valid() and self.context.can_refresh_token(): + refresh_request = await self._refresh_token() # pragma: no cover + + if refresh_request is not None: + # Do not hold the general provider-state lock across + # network I/O. ``flow_lock`` deliberately remains held to + # keep all token-changing transitions single-flight. + refresh_response = yield refresh_request # pragma: no cover + + async with self.context.lock: + if not await self._handle_refresh_response(refresh_response): # pragma: no cover + # Refresh failed, need full re-authentication + self._initialized = False + + sent_access_token = await self._add_valid_auth_header(request) + + # A GET SSE request can remain open for the session lifetime. Yield it + # outside the state lock so concurrent POST requests can authenticate. + response = yield request + + if response.status_code not in (401, 403): + return + + # Serialize the exceptional 401/403 state transitions. Their existing + # full authorization flow remains unchanged. Re-check the token only + # after acquiring the lock so concurrent 401 responses cannot start + # redundant authorization flows. + retry_after_authorization = False + async with self._serialized_transition(): + self.context.protocol_version = protocol_version + current_tokens = self.context.current_tokens + token_changed_since_request = ( + response.status_code in (401, 403) + and self.context.is_token_valid() + and current_tokens is not None + and current_tokens.access_token != sent_access_token + ) + + if token_changed_since_request: + self._add_auth_header(request) + retry_after_authorization = True + elif response.status_code == 401: # Perform full OAuth flow try: # OAuth flow must be inline due to generator constraints @@ -602,7 +665,7 @@ async def async_auth_flow(self, request: httpx.Request) -> AsyncGenerator[httpx. # Retry with new tokens self._add_auth_header(request) - yield request + retry_after_authorization = True elif response.status_code == 403: # Step 1: Extract error field from WWW-Authenticate header error = extract_field_from_www_auth(response, "error") @@ -624,4 +687,8 @@ async def async_auth_flow(self, request: httpx.Request) -> AsyncGenerator[httpx. # Retry with new tokens self._add_auth_header(request) - yield request + retry_after_authorization = True + + # The retried resource request can itself be a session-long GET. + if retry_after_authorization: + yield request diff --git a/tests/client/test_auth.py b/tests/client/test_auth.py index 5f8bc14107..e48b952f3e 100644 --- a/tests/client/test_auth.py +++ b/tests/client/test_auth.py @@ -7,6 +7,7 @@ from unittest import mock from urllib.parse import unquote +import anyio import httpx import pytest from inline_snapshot import Is, snapshot @@ -498,6 +499,7 @@ async def test_oauth_discovery_fallback_conditions(self, oauth_provider: OAuthCl assert final_request.headers["Authorization"] == "Bearer new_access_token" assert final_request.method == "GET" assert str(final_request.url) == "https://api.example.com/v1/mcp" + assert oauth_provider.context.lock.value == 1 # Send final success response to properly close the generator final_response = httpx.Response(200, request=final_request) @@ -1026,6 +1028,7 @@ async def test_auth_flow_with_no_tokens(self, oauth_provider: OAuthClientProvide assert final_request.headers["Authorization"] == "Bearer new_access_token" assert final_request.method == "GET" assert str(final_request.url) == "https://api.example.com/mcp" + assert oauth_provider.context.lock.value == 1 # Send final success response to properly close the generator final_response = httpx.Response(200, request=final_request) @@ -1165,6 +1168,7 @@ async def mock_callback() -> tuple[str, str | None]: # Should get final retry request final_request = await auth_flow.asend(token_response) + assert oauth_provider.context.lock.value == 1 # Send success response - flow should complete success_response = httpx.Response(200, request=final_request) @@ -2113,3 +2117,395 @@ async def test_get_resource_url_falls_back_when_prm_mismatches( # get_resource_url should return the canonical server URL, not the PRM resource assert provider.context.get_resource_url() == "https://api.example.com/v1/mcp" + + +@pytest.mark.anyio +async def test_concurrent_request_not_blocked_by_pending_long_running_request( + oauth_provider: OAuthClientProvider, valid_tokens: OAuthToken +) -> None: + """A long-running request must not block another request's auth flow.""" + oauth_provider.context.current_tokens = valid_tokens + oauth_provider.context.token_expiry_time = time.time() + 1800 + oauth_provider.context.client_info = OAuthClientInformationFull( + client_id="test_client_id", + client_secret="test_client_secret", + redirect_uris=[AnyUrl("http://localhost:3030/callback")], + ) + oauth_provider._initialized = True + + slow_request = httpx.Request("GET", "https://api.example.com/v1/mcp") + slow_flow = oauth_provider.async_auth_flow(slow_request) + yielded_slow = await slow_flow.__anext__() + assert yielded_slow.headers.get("Authorization") == "Bearer test_access_token" + + fast_request = httpx.Request("POST", "https://api.example.com/v1/mcp") + fast_flow = oauth_provider.async_auth_flow(fast_request) + with anyio.fail_after(5): + yielded_fast = await fast_flow.__anext__() + assert yielded_fast.headers.get("Authorization") == "Bearer test_access_token" + + await fast_flow.aclose() + await slow_flow.aclose() + + +@pytest.mark.anyio +async def test_concurrent_token_refresh_is_single_flight( + oauth_provider: OAuthClientProvider, valid_tokens: OAuthToken +) -> None: + """Concurrent requests must share one token refresh.""" + oauth_provider.context.current_tokens = valid_tokens + oauth_provider.context.token_expiry_time = time.time() - 100 + oauth_provider.context.client_info = OAuthClientInformationFull( + client_id="test_client_id", + client_secret="test_client_secret", + redirect_uris=[AnyUrl("http://localhost:3030/callback")], + ) + oauth_provider._initialized = True + + request_a = httpx.Request("GET", "https://api.example.com/v1/mcp") + flow_a = oauth_provider.async_auth_flow(request_a) + refresh_request = await flow_a.__anext__() + assert "grant_type=refresh_token" in refresh_request.read().decode() + + request_b = httpx.Request("POST", "https://api.example.com/v1/mcp") + flow_b = oauth_provider.async_auth_flow(request_b) + flow_b_done = anyio.Event() + flow_b_result: dict[str, httpx.Request] = {} + request_a_after_refresh: httpx.Request | None = None + + async def drive_flow_b() -> None: + flow_b_result["request"] = await flow_b.__anext__() + flow_b_done.set() + + async with anyio.create_task_group() as task_group: + task_group.start_soon(drive_flow_b) + + refresh_response = httpx.Response( + 200, + content=( + b'{"access_token": "new_access_token", "token_type": "Bearer", ' + b'"expires_in": 3600, "refresh_token": "new_refresh_token"}' + ), + request=refresh_request, + ) + request_a_after_refresh = await flow_a.asend(refresh_response) + with anyio.fail_after(5): + await flow_b_done.wait() + + request_b_after_refresh = flow_b_result["request"] + assert request_a_after_refresh is not None + assert request_a_after_refresh.url == request_a.url + assert request_b_after_refresh.url == request_b.url + assert request_a_after_refresh.headers["Authorization"] == "Bearer new_access_token" + assert request_b_after_refresh.headers["Authorization"] == "Bearer new_access_token" + + await flow_b.aclose() + await flow_a.aclose() + + +@pytest.mark.anyio +async def test_refresh_and_401_authorization_are_single_flight( + oauth_provider: OAuthClientProvider, valid_tokens: OAuthToken +) -> None: + """A 401 flow waits for an in-flight refresh, then uses its token.""" + oauth_provider.context.current_tokens = valid_tokens + oauth_provider.context.token_expiry_time = time.time() + 1800 + oauth_provider.context.client_info = OAuthClientInformationFull( + client_id="test_client_id", + client_secret="test_client_secret", + redirect_uris=[AnyUrl("http://localhost:3030/callback")], + ) + oauth_provider._initialized = True + + waiting_request = httpx.Request("POST", "https://api.example.com/v1/mcp") + waiting_flow = oauth_provider.async_auth_flow(waiting_request) + first_attempt = await waiting_flow.__anext__() + + oauth_provider.context.token_expiry_time = time.time() - 100 + refreshing_request = httpx.Request("GET", "https://api.example.com/v1/mcp") + refreshing_flow = oauth_provider.async_auth_flow(refreshing_request) + refresh_request = await refreshing_flow.__anext__() + + retry_result: dict[str, httpx.Request] = {} + refreshing_retry: httpx.Request | None = None + + async def drive_waiting_401() -> None: + retry_result["request"] = await waiting_flow.asend(httpx.Response(401, request=first_attempt)) + + async with anyio.create_task_group() as task_group: + task_group.start_soon(drive_waiting_401) + with anyio.fail_after(5): + while oauth_provider.context.flow_lock.statistics().tasks_waiting == 0: + await anyio.sleep(0) + + refresh_response = httpx.Response( + 200, + content=( + b'{"access_token": "new_access_token", "token_type": "Bearer", ' + b'"expires_in": 3600, "refresh_token": "new_refresh_token"}' + ), + request=refresh_request, + ) + refreshing_retry = await refreshing_flow.asend(refresh_response) + + waiting_retry = retry_result["request"] + assert refreshing_retry is not None + assert refreshing_retry.headers["Authorization"] == "Bearer new_access_token" + assert waiting_retry.headers["Authorization"] == "Bearer new_access_token" + assert oauth_provider.context.flow_lock.value == 1 + assert oauth_provider.context.lock.value == 1 + await waiting_flow.aclose() + await refreshing_flow.aclose() + + +@pytest.mark.anyio +async def test_401_restores_request_protocol_version_before_authorization( + oauth_provider: OAuthClientProvider, valid_tokens: OAuthToken +) -> None: + """Concurrent requests cannot change RFC 8707 behavior for a 401 flow.""" + oauth_provider.context.current_tokens = valid_tokens + oauth_provider.context.token_expiry_time = time.time() + 1800 + oauth_provider._initialized = True + + old_request = httpx.Request( + "GET", + "https://api.example.com/v1/mcp", + headers={"MCP-Protocol-Version": "2025-03-26"}, + ) + old_flow = oauth_provider.async_auth_flow(old_request) + old_attempt = await old_flow.__anext__() + + new_request = httpx.Request( + "POST", + "https://api.example.com/v1/mcp", + headers={"MCP-Protocol-Version": "2025-06-18"}, + ) + new_flow = oauth_provider.async_auth_flow(new_request) + await new_flow.__anext__() + assert oauth_provider.context.protocol_version == "2025-06-18" + + await old_flow.asend(httpx.Response(401, request=old_attempt)) + assert oauth_provider.context.protocol_version == "2025-03-26" + assert not oauth_provider.context.should_include_resource_param(oauth_provider.context.protocol_version) + + await old_flow.aclose() + await new_flow.aclose() + + +@pytest.mark.anyio +async def test_failed_refresh_clears_tokens(oauth_provider: OAuthClientProvider, valid_tokens: OAuthToken) -> None: + """A rejected refresh clears stale credentials before the request proceeds.""" + oauth_provider.context.current_tokens = valid_tokens + oauth_provider.context.token_expiry_time = time.time() - 100 + oauth_provider.context.client_info = OAuthClientInformationFull( + client_id="test_client_id", + client_secret="test_client_secret", + redirect_uris=[AnyUrl("http://localhost:3030/callback")], + ) + oauth_provider._initialized = True + + request = httpx.Request("POST", "https://api.example.com/v1/mcp") + flow = oauth_provider.async_auth_flow(request) + refresh_request = await flow.__anext__() + yielded_request = await flow.asend(httpx.Response(401, request=refresh_request)) + + assert yielded_request.url == request.url + assert "Authorization" not in yielded_request.headers + assert oauth_provider.context.current_tokens is None + assert oauth_provider._initialized is False + + await flow.aclose() + + +@pytest.mark.anyio +async def test_stale_401_retries_with_concurrently_refreshed_token_without_locking( + oauth_provider: OAuthClientProvider, valid_tokens: OAuthToken +) -> None: + """A stale 401 retries with the newer token without holding the state lock.""" + oauth_provider.context.current_tokens = valid_tokens + oauth_provider.context.token_expiry_time = time.time() + 1800 + oauth_provider._initialized = True + + stale_request = httpx.Request("GET", "https://api.example.com/v1/mcp") + stale_flow = oauth_provider.async_auth_flow(stale_request) + first_attempt = await stale_flow.__anext__() + assert first_attempt.headers["Authorization"] == "Bearer test_access_token" + + oauth_provider.context.current_tokens = OAuthToken( + access_token="new_access_token", + token_type="Bearer", + expires_in=3600, + refresh_token="new_refresh_token", + ) + oauth_provider.context.token_expiry_time = time.time() + 3600 + + retry = await stale_flow.asend(httpx.Response(401, request=first_attempt)) + assert retry.headers["Authorization"] == "Bearer new_access_token" + + # The retry can be a session-long GET, so it must not block another flow. + concurrent_flow = oauth_provider.async_auth_flow(httpx.Request("POST", "https://api.example.com/v1/mcp")) + with anyio.fail_after(5): + concurrent_request = await concurrent_flow.__anext__() + assert concurrent_request.headers["Authorization"] == "Bearer new_access_token" + + with pytest.raises(StopAsyncIteration): + await stale_flow.asend(httpx.Response(200, request=retry)) + await concurrent_flow.aclose() + + +@pytest.mark.anyio +async def test_401_rechecks_token_after_waiting_for_state_lock( + oauth_provider: OAuthClientProvider, valid_tokens: OAuthToken +) -> None: + """A waiting 401 flow uses a token installed while it waited.""" + oauth_provider.context.current_tokens = valid_tokens + oauth_provider.context.token_expiry_time = time.time() + 1800 + oauth_provider._initialized = True + + request = httpx.Request("POST", "https://api.example.com/v1/mcp") + flow = oauth_provider.async_auth_flow(request) + first_attempt = await flow.__anext__() + result: dict[str, httpx.Request] = {} + + async def drive_401() -> None: + result["retry"] = await flow.asend(httpx.Response(401, request=first_attempt)) + + async with anyio.create_task_group() as task_group: + async with oauth_provider.context.lock: + task_group.start_soon(drive_401) + with anyio.fail_after(5): + while oauth_provider.context.lock.statistics().tasks_waiting == 0: + await anyio.sleep(0) + + oauth_provider.context.current_tokens = OAuthToken( + access_token="new_access_token", + token_type="Bearer", + expires_in=3600, + refresh_token="new_refresh_token", + ) + oauth_provider.context.token_expiry_time = time.time() + 3600 + + assert result["retry"] is request + assert request.headers["Authorization"] == "Bearer new_access_token" + assert oauth_provider.context.lock.value == 1 + await flow.aclose() + + +@pytest.mark.anyio +async def test_success_response_does_not_wait_for_oauth_transition( + oauth_provider: OAuthClientProvider, valid_tokens: OAuthToken +) -> None: + """A successful resource response bypasses OAuth transition locks.""" + oauth_provider.context.current_tokens = valid_tokens + oauth_provider.context.token_expiry_time = time.time() + 1800 + oauth_provider._initialized = True + + request = httpx.Request("POST", "https://api.example.com/v1/mcp") + flow = oauth_provider.async_auth_flow(request) + first_attempt = await flow.__anext__() + completed = anyio.Event() + + async def complete_request() -> None: + with pytest.raises(StopAsyncIteration): + await flow.asend(httpx.Response(200, request=first_attempt)) + completed.set() + + async with anyio.create_task_group() as task_group: + async with oauth_provider.context.flow_lock: + task_group.start_soon(complete_request) + with anyio.fail_after(5): + await completed.wait() + + +@pytest.mark.anyio +async def test_403_rechecks_token_after_waiting_for_transition_lock( + oauth_provider: OAuthClientProvider, valid_tokens: OAuthToken +) -> None: + """A waiting 403 flow retries with a token installed while it waited.""" + oauth_provider.context.current_tokens = valid_tokens + oauth_provider.context.token_expiry_time = time.time() + 1800 + oauth_provider._initialized = True + + request = httpx.Request("POST", "https://api.example.com/v1/mcp") + flow = oauth_provider.async_auth_flow(request) + first_attempt = await flow.__anext__() + result: dict[str, httpx.Request] = {} + + async def drive_403() -> None: + result["retry"] = await flow.asend(httpx.Response(403, request=first_attempt)) + + async with anyio.create_task_group() as task_group: + async with oauth_provider.context.flow_lock: + task_group.start_soon(drive_403) + with anyio.fail_after(5): + while oauth_provider.context.flow_lock.statistics().tasks_waiting == 0: + await anyio.sleep(0) + + async with oauth_provider.context.lock: + oauth_provider.context.current_tokens = OAuthToken( + access_token="new_access_token", + token_type="Bearer", + expires_in=3600, + refresh_token="new_refresh_token", + ) + oauth_provider.context.token_expiry_time = time.time() + 3600 + + assert result["retry"] is request + assert request.headers["Authorization"] == "Bearer new_access_token" + assert oauth_provider.context.flow_lock.value == 1 + assert oauth_provider.context.lock.value == 1 + await flow.aclose() + + +@pytest.mark.anyio +async def test_401_does_not_retry_with_an_expired_unsent_token(oauth_provider: OAuthClientProvider) -> None: + """An expired token omitted from the request must not count as a newer token.""" + oauth_provider.context.current_tokens = OAuthToken( + access_token="expired_access_token", + token_type="Bearer", + expires_in=3600, + ) + oauth_provider.context.token_expiry_time = time.time() - 1 + oauth_provider._initialized = True + + request = httpx.Request("POST", "https://api.example.com/v1/mcp") + flow = oauth_provider.async_auth_flow(request) + first_attempt = await flow.__anext__() + assert "Authorization" not in first_attempt.headers + + discovery_request = await flow.asend(httpx.Response(401, request=first_attempt)) + assert discovery_request is not request + assert "Authorization" not in request.headers + await flow.aclose() + + +@pytest.mark.anyio +async def test_authorization_flow_can_close_from_a_different_task( + oauth_provider: OAuthClientProvider, valid_tokens: OAuthToken +) -> None: + """HTTPX may close a suspended auth generator from another task.""" + oauth_provider.context.current_tokens = valid_tokens + oauth_provider.context.token_expiry_time = time.time() + 1800 + oauth_provider._initialized = True + + request = httpx.Request("POST", "https://api.example.com/v1/mcp") + flow = oauth_provider.async_auth_flow(request) + first_attempt = await flow.__anext__() + + await flow.asend(httpx.Response(401, request=first_attempt)) + assert oauth_provider.context.flow_lock.value == 0 + assert oauth_provider.context.lock.value == 0 + + closed = anyio.Event() + + async def close_flow() -> None: + await flow.aclose() + closed.set() + + async with anyio.create_task_group() as task_group: + task_group.start_soon(close_flow) + with anyio.fail_after(5): + await closed.wait() + + assert oauth_provider.context.flow_lock.value == 1 + assert oauth_provider.context.lock.value == 1 From c180fac922e5b46987c188ea0bbfa88a95da661e Mon Sep 17 00:00:00 2001 From: Carl Taylor Date: Tue, 4 Aug 2026 12:44:34 +1000 Subject: [PATCH 2/3] test(oauth): mark invariant branches for coverage --- src/mcp/client/auth/oauth2.py | 4 ++-- tests/client/test_auth.py | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/src/mcp/client/auth/oauth2.py b/src/mcp/client/auth/oauth2.py index 2c9014e519..18c0543a71 100644 --- a/src/mcp/client/auth/oauth2.py +++ b/src/mcp/client/auth/oauth2.py @@ -666,7 +666,7 @@ async def async_auth_flow(self, request: httpx.Request) -> AsyncGenerator[httpx. # Retry with new tokens self._add_auth_header(request) retry_after_authorization = True - elif response.status_code == 403: + elif response.status_code == 403: # pragma: no branch # Step 1: Extract error field from WWW-Authenticate header error = extract_field_from_www_auth(response, "error") @@ -690,5 +690,5 @@ async def async_auth_flow(self, request: httpx.Request) -> AsyncGenerator[httpx. retry_after_authorization = True # The retried resource request can itself be a session-long GET. - if retry_after_authorization: + if retry_after_authorization: # pragma: no branch yield request diff --git a/tests/client/test_auth.py b/tests/client/test_auth.py index e48b952f3e..bbca5eaad4 100644 --- a/tests/client/test_auth.py +++ b/tests/client/test_auth.py @@ -2413,7 +2413,7 @@ async def complete_request() -> None: async with anyio.create_task_group() as task_group: async with oauth_provider.context.flow_lock: task_group.start_soon(complete_request) - with anyio.fail_after(5): + with anyio.fail_after(5): # pragma: no branch await completed.wait() @@ -2441,7 +2441,7 @@ async def drive_403() -> None: while oauth_provider.context.flow_lock.statistics().tasks_waiting == 0: await anyio.sleep(0) - async with oauth_provider.context.lock: + async with oauth_provider.context.lock: # pragma: no branch oauth_provider.context.current_tokens = OAuthToken( access_token="new_access_token", token_type="Bearer", From 047afc92a7cb59f9b1020c48471363d7fdc94d3d Mon Sep 17 00:00:00 2001 From: Carl Taylor Date: Tue, 4 Aug 2026 12:56:50 +1000 Subject: [PATCH 3/3] test(oauth): require concurrent refresh lock contention --- tests/client/test_auth.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/client/test_auth.py b/tests/client/test_auth.py index bbca5eaad4..3e752d8d80 100644 --- a/tests/client/test_auth.py +++ b/tests/client/test_auth.py @@ -2179,6 +2179,9 @@ async def drive_flow_b() -> None: async with anyio.create_task_group() as task_group: task_group.start_soon(drive_flow_b) + with anyio.fail_after(5): + while oauth_provider.context.flow_lock.statistics().tasks_waiting == 0: + await anyio.sleep(0) refresh_response = httpx.Response( 200,