Merge pull request #509 from jlowin/initialize-result

Store the initialize result on the client
This commit is contained in:
Jeremiah Lowin 2025-05-20 10:48:32 -04:00 committed by GitHub
commit 86cff91711
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 76 additions and 29 deletions

View file

@ -1,5 +1,5 @@
import datetime
from contextlib import AsyncExitStack
from contextlib import AsyncExitStack, asynccontextmanager
from pathlib import Path
from typing import Any, cast
@ -84,6 +84,7 @@ class Client:
self._session: ClientSession | None = None
self._exit_stack: AsyncExitStack | None = None
self._nesting_counter: int = 0
self._initialize_result: mcp.types.InitializeResult | None = None
if log_handler is None:
log_handler = default_log_handler
@ -117,10 +118,19 @@ class Client:
"""Get the current active session. Raises RuntimeError if not connected."""
if self._session is None:
raise RuntimeError(
"Client is not connected. Use 'async with client:' context manager first."
"Client is not connected. Use the 'async with client:' context manager first."
)
return self._session
@property
def initialize_result(self) -> mcp.types.InitializeResult:
"""Get the result of the initialization request."""
if self._initialize_result is None:
raise RuntimeError(
"Client is not connected. Use the 'async with client:' context manager first."
)
return self._initialize_result
def set_roots(self, roots: RootsList | RootsHandler) -> None:
"""Set the roots for the client. This does not automatically call `send_roots_list_changed`."""
self._session_kwargs["list_roots_callback"] = create_roots_callback(roots)
@ -135,27 +145,35 @@ class Client:
"""Check if the client is currently connected."""
return self._session is not None
@asynccontextmanager
async def _context_manager(self):
with catch(get_catch_handlers()):
async with self.transport.connect_session(
**self._session_kwargs
) as session:
self._session = session
# Initialize the session
self._initialize_result = await self._session.initialize()
try:
yield
finally:
self._exit_stack = None
self._session = None
self._initialize_result = None
async def __aenter__(self):
if self._nesting_counter == 0:
# Create exit stack to manage both context managers
stack = AsyncExitStack()
await stack.__aenter__()
# Add the exception handling context
stack.enter_context(catch(get_catch_handlers()))
await stack.enter_async_context(self._context_manager())
# the above catch will only apply once this __aenter__ finishes so
# we need to wrap the session creation in a new context in case it
# raises errors itself
with catch(get_catch_handlers()):
# Create and enter the transport session using the exit stack
session_cm = self.transport.connect_session(**self._session_kwargs)
self._session = await stack.enter_async_context(session_cm)
# Store the stack for cleanup in __aexit__
self._exit_stack = stack
self._nesting_counter += 1
return self
async def __aexit__(self, exc_type, exc_val, exc_tb):
@ -168,7 +186,6 @@ class Client:
await self._exit_stack.__aexit__(exc_type, exc_val, exc_tb)
finally:
self._exit_stack = None
self._session = None
# --- MCP Client Methods ---

View file

@ -46,6 +46,7 @@ class ClientTransport(abc.ABC):
A Transport is responsible for establishing and managing connections
to an MCP server, and providing a ClientSession within an async context.
"""
@abc.abstractmethod
@ -54,7 +55,9 @@ class ClientTransport(abc.ABC):
self, **session_kwargs: Unpack[SessionKwargs]
) -> AsyncIterator[ClientSession]:
"""
Establishes a connection and yields an active, initialized ClientSession.
Establishes a connection and yields an active ClientSession.
The ClientSession is *not* expected to be initialized in this context manager.
The session is guaranteed to be valid only within the scope of the
async context manager. Connection setup and teardown are handled
@ -65,7 +68,7 @@ class ClientTransport(abc.ABC):
constructor (e.g., callbacks, timeouts).
Yields:
An initialized mcp.ClientSession instance.
A mcp.ClientSession instance.
"""
raise NotImplementedError
yield None # type: ignore
@ -94,7 +97,6 @@ class WSTransport(ClientTransport):
async with ClientSession(
read_stream, write_stream, **session_kwargs
) as session:
await session.initialize() # Initialize after session creation
yield session
def __repr__(self) -> str:
@ -143,7 +145,6 @@ class SSETransport(ClientTransport):
async with ClientSession(
read_stream, write_stream, **session_kwargs
) as session:
await session.initialize()
yield session
def __repr__(self) -> str:
@ -189,7 +190,6 @@ class StreamableHttpTransport(ClientTransport):
async with ClientSession(
read_stream, write_stream, **session_kwargs
) as session:
await session.initialize()
yield session
def __repr__(self) -> str:
@ -237,7 +237,6 @@ class StdioTransport(ClientTransport):
async with ClientSession(
read_stream, write_stream, **session_kwargs
) as session:
await session.initialize()
yield session
def __repr__(self) -> str:

View file

@ -256,18 +256,51 @@ async def test_read_resource_mcp(fastmcp_server):
async def test_client_connection(fastmcp_server):
"""Test that the client connects and disconnects properly."""
"""Test that connect is idempotent."""
client = Client(transport=FastMCPTransport(fastmcp_server))
# Before connection
# Connect idempotently
async with client:
assert client.is_connected()
# Make a request to ensure connection is working
await client.ping()
assert not client.is_connected()
# During connection
async def test_initialize_result_connected(fastmcp_server):
"""Test that initialize_result returns the correct result when connected."""
client = Client(transport=FastMCPTransport(fastmcp_server))
# Initialize result should not be accessible before connection
with pytest.raises(RuntimeError, match="Client is not connected"):
_ = client.initialize_result
async with client:
# Once connected, initialize_result should be available
result = client.initialize_result
# Verify the initialize result has expected properties
assert hasattr(result, "serverInfo")
assert result.serverInfo.name == "TestServer"
assert result.serverInfo.version is not None
async def test_initialize_result_disconnected(fastmcp_server):
"""Test that initialize_result raises an error when not connected."""
client = Client(transport=FastMCPTransport(fastmcp_server))
# Initialize result should not be accessible before connection
with pytest.raises(RuntimeError, match="Client is not connected"):
_ = client.initialize_result
# Connect and then disconnect
async with client:
assert client.is_connected()
# After connection
# After disconnection, initialize_result should raise an error
assert not client.is_connected()
with pytest.raises(RuntimeError, match="Client is not connected"):
_ = client.initialize_result
async def test_client_nested_context_manager(fastmcp_server):

View file

@ -640,7 +640,6 @@ class TestToolContextInjection:
assert len(result) == 1
content = result[0]
assert isinstance(content, TextContent)
assert content.text == "1"
async def test_async_context(self):
"""Test that context works in async functions."""
@ -656,8 +655,7 @@ class TestToolContextInjection:
assert len(result) == 1
content = result[0]
assert isinstance(content, TextContent)
assert "Async request" in content.text
assert "42" in content.text
assert content.text == "Async request 2: 42"
async def test_optional_context(self):
"""Test that context is optional."""
@ -798,7 +796,7 @@ class TestResourceContext:
async with Client(mcp) as client:
result = await client.read_resource(AnyUrl("resource://test"))
assert isinstance(result[0], TextResourceContents)
assert result[0].text == "1"
assert result[0].text == "2"
class TestResourceTemplates:
@ -1096,7 +1094,7 @@ class TestResourceTemplateContext:
async with Client(mcp) as client:
result = await client.read_resource(AnyUrl("resource://test"))
assert isinstance(result[0], TextResourceContents)
assert result[0].text == "Resource template: test 1"
assert result[0].text.startswith("Resource template: test 2")
class TestPrompts: