Add tests for OAuth generator cleanup and use aclosing (#2759)

This commit is contained in:
Jeremiah Lowin 2025-12-26 16:01:27 -05:00 committed by GitHub
commit e454294ea8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 110 additions and 10 deletions

View file

@ -3,6 +3,7 @@ from __future__ import annotations
import time
import webbrowser
from collections.abc import AsyncGenerator
from contextlib import aclosing
from typing import Any
import anyio
@ -296,9 +297,8 @@ class OAuth(OAuthClientProvider):
"""
try:
# First attempt with potentially cached credentials
gen = super().async_auth_flow(request)
response = None
try:
async with aclosing(super().async_auth_flow(request)) as gen:
response = None
while True:
try:
# First iteration sends None, subsequent iterations send response
@ -306,8 +306,6 @@ class OAuth(OAuthClientProvider):
response = yield yielded_request
except StopAsyncIteration:
break
finally:
await gen.aclose()
except ClientNotFoundError:
logger.debug(
@ -318,14 +316,11 @@ class OAuth(OAuthClientProvider):
await self.token_storage_adapter.clear()
# Retry with fresh registration
gen = super().async_auth_flow(request)
response = None
try:
async with aclosing(super().async_auth_flow(request)) as gen:
response = None
while True:
try:
yielded_request = await gen.asend(response) # ty: ignore[invalid-argument-type]
response = yield yielded_request
except StopAsyncIteration:
break
finally:
await gen.aclose()

View file

@ -1,3 +1,4 @@
from unittest.mock import patch
from urllib.parse import urlparse
import httpx
@ -173,3 +174,107 @@ class TestOAuthClientUrlHandling:
# Token storage should key by the full URL, not just the host
assert oauth.token_storage_adapter._server_url == mcp_url
class TestOAuthGeneratorCleanup:
"""Tests for OAuth async generator cleanup (issue #2643).
The MCP SDK's OAuthClientProvider.async_auth_flow() holds a lock via
`async with self.context.lock`. If the generator is not explicitly closed,
GC may clean it up from a different task, causing:
RuntimeError: The current task is not holding this lock
"""
async def test_generator_closed_on_successful_flow(self):
"""Verify aclose() is called on the parent generator after successful flow."""
oauth = OAuth(mcp_url="https://example.com")
# Track generator lifecycle using a wrapper class
class TrackedGenerator:
def __init__(self):
self.aclose_called = False
self._exhausted = False
def __aiter__(self):
return self
async def __anext__(self):
if self._exhausted:
raise StopAsyncIteration
self._exhausted = True
return httpx.Request("GET", "https://example.com")
async def asend(self, value):
if self._exhausted:
raise StopAsyncIteration
self._exhausted = True
return httpx.Request("GET", "https://example.com")
async def athrow(self, exc_type, exc_val=None, exc_tb=None):
raise StopAsyncIteration
async def aclose(self):
self.aclose_called = True
tracked_gen = TrackedGenerator()
# Patch the parent class to return our tracked generator
with patch.object(
OAuth.__bases__[0], "async_auth_flow", return_value=tracked_gen
):
# Drive the OAuth flow
flow = oauth.async_auth_flow(httpx.Request("GET", "https://example.com"))
try:
# First asend(None) starts the generator per async generator protocol
await flow.asend(None) # ty: ignore[invalid-argument-type]
try:
await flow.asend(httpx.Response(200))
except StopAsyncIteration:
pass
except StopAsyncIteration:
pass
assert tracked_gen.aclose_called, (
"Generator aclose() was not called after flow completion"
)
async def test_generator_closed_on_exception(self):
"""Verify aclose() is called even when an exception occurs mid-flow."""
oauth = OAuth(mcp_url="https://example.com")
class FailingGenerator:
def __init__(self):
self.aclose_called = False
self._first_call = True
def __aiter__(self):
return self
async def __anext__(self):
return await self.asend(None)
async def asend(self, value):
if self._first_call:
self._first_call = False
return httpx.Request("GET", "https://example.com")
raise ValueError("Simulated failure")
async def athrow(self, exc_type, exc_val=None, exc_tb=None):
raise StopAsyncIteration
async def aclose(self):
self.aclose_called = True
tracked_gen = FailingGenerator()
with patch.object(
OAuth.__bases__[0], "async_auth_flow", return_value=tracked_gen
):
flow = oauth.async_auth_flow(httpx.Request("GET", "https://example.com"))
with pytest.raises(ValueError, match="Simulated failure"):
await flow.asend(None) # ty: ignore[invalid-argument-type]
await flow.asend(httpx.Response(200))
assert tracked_gen.aclose_called, (
"Generator aclose() was not called after exception"
)