simply nesting handling

This commit is contained in:
yihuang 2025-04-21 13:04:07 +08:00
commit 0af58a85a6
No known key found for this signature in database
GPG key ID: 5D8B72A08439BE01

View file

@ -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: