|
| 1 | +"""Shared live-broker checks for the reference adapters.""" |
| 2 | + |
| 3 | +from collections.abc import AsyncIterator, Awaitable, Callable |
| 4 | +from contextlib import AsyncExitStack, asynccontextmanager |
| 5 | +from dataclasses import dataclass |
| 6 | +from functools import partial |
| 7 | +from typing import Any, TypeAlias |
| 8 | +from uuid import uuid4 |
| 9 | + |
| 10 | +import anyio |
| 11 | +from mcp import Client |
| 12 | +from mcp.server import Server, ServerRequestContext |
| 13 | +from mcp.server.mcpserver import Context, MCPServer |
| 14 | +from mcp.shared.transport import MessageMetadata, Transport, TransportContext |
| 15 | +from mcp.types import CallToolRequestParams, CallToolResult, ListToolsResult, PaginatedRequestParams, TextContent, Tool |
| 16 | + |
| 17 | +TransportFactory: TypeAlias = Callable[[AsyncExitStack, str, str, bool], Awaitable[Transport]] |
| 18 | + |
| 19 | + |
| 20 | +@dataclass(kw_only=True, frozen=True) |
| 21 | +class BrokerContext(TransportContext): |
| 22 | + principal: str |
| 23 | + |
| 24 | + |
| 25 | +def peer_context(metadata: MessageMetadata, *, principal: str, kind: str) -> BrokerContext: |
| 26 | + return BrokerContext(kind=kind, can_send_request=True, principal=principal) |
| 27 | + |
| 28 | + |
| 29 | +async def verify(factory: TransportFactory, *, kind: str, highlevel: bool, mode: str) -> None: |
| 30 | + """Check real concurrent calls, peer metadata, both server APIs, and shared lifespan.""" |
| 31 | + entered = {"alice": anyio.Event(), "bob": anyio.Event()} |
| 32 | + sessions = {principal: uuid4().hex for principal in entered} |
| 33 | + lifecycle: list[str] = [] |
| 34 | + |
| 35 | + async def identity(transport: TransportContext | None) -> str: |
| 36 | + assert isinstance(transport, BrokerContext) |
| 37 | + principal = transport.principal |
| 38 | + entered[principal].set() |
| 39 | + await entered["bob" if principal == "alice" else "alice"].wait() |
| 40 | + return principal |
| 41 | + |
| 42 | + @asynccontextmanager |
| 43 | + async def lifespan(server: Server[Any] | MCPServer[Any]) -> AsyncIterator[None]: |
| 44 | + lifecycle.append("start") |
| 45 | + try: |
| 46 | + yield None |
| 47 | + finally: |
| 48 | + lifecycle.append("stop") |
| 49 | + |
| 50 | + if highlevel: |
| 51 | + server = MCPServer("Broker", lifespan=lifespan) |
| 52 | + |
| 53 | + @server.tool() |
| 54 | + async def identify(ctx: Context) -> str: |
| 55 | + return await identity(ctx.transport) |
| 56 | + |
| 57 | + else: |
| 58 | + |
| 59 | + async def list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams | None) -> ListToolsResult: |
| 60 | + return ListToolsResult(tools=[Tool(name="identify", input_schema={"type": "object"})]) |
| 61 | + |
| 62 | + async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: |
| 63 | + assert params.name == "identify" |
| 64 | + return CallToolResult(content=[TextContent(text=await identity(ctx.transport))]) |
| 65 | + |
| 66 | + server = Server("Broker", lifespan=lifespan, on_list_tools=list_tools, on_call_tool=call_tool) |
| 67 | + |
| 68 | + results: dict[str, str] = {} |
| 69 | + |
| 70 | + async def call(client: Client, principal: str) -> None: |
| 71 | + result = await client.call_tool("identify") |
| 72 | + content = result.content[0] |
| 73 | + assert isinstance(content, TextContent) |
| 74 | + results[principal] = content.text |
| 75 | + |
| 76 | + async with AsyncExitStack() as stack: |
| 77 | + transports = { |
| 78 | + principal: await factory(stack, principal, session, True) for principal, session in sessions.items() |
| 79 | + } |
| 80 | + runtime = await stack.enter_async_context(server.serve()) |
| 81 | + clients: dict[str, Client] = {} |
| 82 | + for principal, transport in transports.items(): |
| 83 | + await runtime.connect( |
| 84 | + transport, |
| 85 | + session_id=sessions[principal], |
| 86 | + transport_builder=partial(peer_context, principal=principal, kind=kind), |
| 87 | + ) |
| 88 | + client_transport = await factory(stack, principal, sessions[principal], False) |
| 89 | + clients[principal] = await stack.enter_async_context( |
| 90 | + Client(client_transport, mode=mode, read_timeout_seconds=5) |
| 91 | + ) |
| 92 | + async with anyio.create_task_group() as tg: |
| 93 | + for principal, client in clients.items(): |
| 94 | + tg.start_soon(call, client, principal) |
| 95 | + assert results == {"alice": "alice", "bob": "bob"} |
| 96 | + assert lifecycle == ["start"] |
| 97 | + assert lifecycle == ["start", "stop"] |
0 commit comments