From 5bbad74044d0748ba92c8b2992766eb30e869ecc Mon Sep 17 00:00:00 2001 From: yihuang Date: Mon, 21 Apr 2025 13:14:46 +0800 Subject: [PATCH] add unit test --- src/fastmcp/client/client.py | 10 ++++----- tests/client/test_client.py | 39 +++++++++++++++++++++++++++++++++++- 2 files changed, 43 insertions(+), 6 deletions(-) diff --git a/src/fastmcp/client/client.py b/src/fastmcp/client/client.py index 82f77dd46..1bf34ba4c 100644 --- a/src/fastmcp/client/client.py +++ b/src/fastmcp/client/client.py @@ -45,7 +45,7 @@ class Client: ): self.transport = infer_transport(transport) self._session: ClientSession | None = None - self._session_cms: AbstractAsyncContextManager[ClientSession] | None = None + self._session_cm: AbstractAsyncContextManager[ClientSession] | None = None self._nesting_counter: int = 0 self._session_kwargs: SessionKwargs = { @@ -95,11 +95,11 @@ class Client: return self async def __aexit__(self, exc_type, exc_val, exc_tb): - self._nesting_counter -= 0 + self._nesting_counter -= 1 - if self._nesting_counter == 0 and self._session_cms is not None: - await self._session_cms.__aexit__(exc_type, exc_val, exc_tb) - self._session_cms = None + if self._nesting_counter == 0 and self._session_cm is not None: + await self._session_cm.__aexit__(exc_type, exc_val, exc_tb) + self._session_cm = None self._session = None # --- MCP Client Methods --- diff --git a/tests/client/test_client.py b/tests/client/test_client.py index 17c299e8c..10eee4c35 100644 --- a/tests/client/test_client.py +++ b/tests/client/test_client.py @@ -1,10 +1,15 @@ from typing import cast +from typing_extensions import Unpack +from collections.abc import AsyncIterator +from mcp import ClientSession +import contextlib +from mcp.shared.memory import create_client_server_memory_streams import pytest from pydantic import AnyUrl from fastmcp.client import Client -from fastmcp.client.transports import FastMCPTransport +from fastmcp.client.transports import FastMCPTransport, ClientTransport, SessionKwargs from fastmcp.server.server import FastMCP @@ -159,6 +164,38 @@ async def test_client_connection(fastmcp_server): # After connection assert not client.is_connected() +async def test_client_nested_context_manager(fastmcp_server): + """Test that the client connects and disconnects once in nested context manager.""" + class MockTransport(ClientTransport): + def __init__(self): + self._connected = False + + @contextlib.asynccontextmanager + async def connect_session( + self, **session_kwargs: Unpack[SessionKwargs], + ) -> AsyncIterator[ClientSession]: + assert not self._connected, "Transport is connected multiple times" + self._connected = True + async with create_client_server_memory_streams() as ( + _, + server_streams, + ): + yield ClientSession(*server_streams) + + client = Client(transport=MockTransport()) + + # Before connection + assert not client.is_connected() + + # During connection + async with client: + assert client.is_connected() + + async with client: + assert client.is_connected() + + # After connection + assert not client.is_connected() async def test_resource_template(fastmcp_server): """Test using a resource template with InMemoryClient."""