From 0af58a85a6bcbde9e120a6a7d1f60e73eea7e5a4 Mon Sep 17 00:00:00 2001 From: yihuang Date: Mon, 21 Apr 2025 13:04:07 +0800 Subject: [PATCH] simply nesting handling --- src/fastmcp/client/client.py | 36 +++++++++++++++++------------------- 1 file changed, 17 insertions(+), 19 deletions(-) diff --git a/src/fastmcp/client/client.py b/src/fastmcp/client/client.py index 4200d4d8d..82f77dd46 100644 --- a/src/fastmcp/client/client.py +++ b/src/fastmcp/client/client.py @@ -44,10 +44,9 @@ class Client: read_timeout_seconds: datetime.timedelta | None = None, ): self.transport = infer_transport(transport) - # stack to record nested context manager, None is pushed if reuse existing session - self._session_cms: list[ - tuple[AbstractAsyncContextManager[ClientSession] | None, ClientSession] - ] = [] + self._session: ClientSession | None = None + self._session_cms: AbstractAsyncContextManager[ClientSession] | None = None + self._nesting_counter: int = 0 self._session_kwargs: SessionKwargs = { "sampling_callback": None, @@ -66,12 +65,11 @@ class Client: @property def session(self) -> ClientSession: """Get the current active session. Raises RuntimeError if not connected.""" - if not self._session_cms: + if not self._session: raise RuntimeError( "Client is not connected. Use 'async with client:' context manager first." ) - _, session = self._session_cms[-1] - return session + return self._session def set_roots(self, roots: RootsList | RootsHandler) -> None: """Set the roots for the client. This does not automatically call `send_roots_list_changed`.""" @@ -85,24 +83,24 @@ class Client: def is_connected(self) -> bool: """Check if the client is currently connected.""" - return len(self._session_cms) > 0 + return self._session is not None async def __aenter__(self): - if self._session_cms: - # share the current session, push a None as context manager to avoid close it in aexit - _, session = self._session_cms[-1] - self._session_cms.append((None, session)) - else: + if self._nesting_counter == 0: # create new session - session_cm = self.transport.connect_session(**self._session_kwargs) - session = await session_cm.__aenter__() - self._session_cms.append((session_cm, session)) + self._session_cm = self.transport.connect_session(**self._session_kwargs) + self._session = await self._session_cm.__aenter__() + + self._nesting_counter += 1 return self async def __aexit__(self, exc_type, exc_val, exc_tb): - cm, _ = self._session_cms.pop() - if cm is not None: - await cm.__aexit__(exc_type, exc_val, exc_tb) + self._nesting_counter -= 0 + + 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 + self._session = None # --- MCP Client Methods --- async def ping(self) -> None: