|
| 1 | +from collections.abc import AsyncIterator |
| 2 | +from contextlib import asynccontextmanager |
| 3 | + |
| 4 | +import anyio |
| 5 | +import pytest |
| 6 | +from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream |
| 7 | +from mcp.client.auth import OAuthClientProvider |
| 8 | +from mcp.shared.message import SessionMessage |
| 9 | + |
| 10 | +from mcp_simple_auth_client import main as client_module |
| 11 | +from mcp_simple_auth_client.main import SimpleAuthClient |
| 12 | + |
| 13 | + |
| 14 | +@pytest.mark.anyio |
| 15 | +async def test_oauth_client_preserves_the_complete_connection_url(monkeypatch: pytest.MonkeyPatch) -> None: |
| 16 | + """The example passes the opaque MCP endpoint unchanged to its OAuth provider.""" |
| 17 | + resource_url = "https://mcp.example.com/prefix/mcp?tenant=mcp" |
| 18 | + providers: list[OAuthClientProvider] = [] |
| 19 | + sessions = 0 |
| 20 | + |
| 21 | + class FakeCallbackServer: |
| 22 | + def __init__(self, port: int) -> None: |
| 23 | + assert port == 3030 |
| 24 | + |
| 25 | + def start(self) -> None: |
| 26 | + pass |
| 27 | + |
| 28 | + @asynccontextmanager |
| 29 | + async def fake_sse_client( |
| 30 | + *, url: str, auth: OAuthClientProvider, timeout: float |
| 31 | + ) -> AsyncIterator[ |
| 32 | + tuple[MemoryObjectReceiveStream[SessionMessage | Exception], MemoryObjectSendStream[SessionMessage]] |
| 33 | + ]: |
| 34 | + assert url == resource_url |
| 35 | + assert timeout == 60.0 |
| 36 | + providers.append(auth) |
| 37 | + read_send, read_receive = anyio.create_memory_object_stream[SessionMessage | Exception](1) |
| 38 | + write_send, write_receive = anyio.create_memory_object_stream[SessionMessage](1) |
| 39 | + async with read_send, read_receive, write_send, write_receive: |
| 40 | + yield read_receive, write_send |
| 41 | + |
| 42 | + async def record_session( |
| 43 | + self: SimpleAuthClient, |
| 44 | + read_stream: MemoryObjectReceiveStream[SessionMessage | Exception], |
| 45 | + write_stream: MemoryObjectSendStream[SessionMessage], |
| 46 | + ) -> None: |
| 47 | + nonlocal sessions |
| 48 | + sessions += 1 |
| 49 | + |
| 50 | + monkeypatch.setattr(client_module, "CallbackServer", FakeCallbackServer) |
| 51 | + monkeypatch.setattr(client_module, "sse_client", fake_sse_client) |
| 52 | + monkeypatch.setattr(SimpleAuthClient, "_run_session", record_session) |
| 53 | + |
| 54 | + await SimpleAuthClient(resource_url, transport_type="sse").connect() |
| 55 | + |
| 56 | + assert sessions == 1 |
| 57 | + assert [str(provider.context.server_url) for provider in providers] == [resource_url] |
0 commit comments