yihuang2026 commited on
Commit
d4d4e69
·
unverified ·
1 Parent(s): f1fb9bb

simply nesting handling

Browse files
Files changed (1) hide show
  1. src/fastmcp/client/client.py +17 -19
src/fastmcp/client/client.py CHANGED
@@ -44,10 +44,9 @@ class Client:
44
  read_timeout_seconds: datetime.timedelta | None = None,
45
  ):
46
  self.transport = infer_transport(transport)
47
- # stack to record nested context manager, None is pushed if reuse existing session
48
- self._session_cms: list[
49
- tuple[AbstractAsyncContextManager[ClientSession] | None, ClientSession]
50
- ] = []
51
 
52
  self._session_kwargs: SessionKwargs = {
53
  "sampling_callback": None,
@@ -66,12 +65,11 @@ class Client:
66
  @property
67
  def session(self) -> ClientSession:
68
  """Get the current active session. Raises RuntimeError if not connected."""
69
- if not self._session_cms:
70
  raise RuntimeError(
71
  "Client is not connected. Use 'async with client:' context manager first."
72
  )
73
- _, session = self._session_cms[-1]
74
- return session
75
 
76
  def set_roots(self, roots: RootsList | RootsHandler) -> None:
77
  """Set the roots for the client. This does not automatically call `send_roots_list_changed`."""
@@ -85,24 +83,24 @@ class Client:
85
 
86
  def is_connected(self) -> bool:
87
  """Check if the client is currently connected."""
88
- return len(self._session_cms) > 0
89
 
90
  async def __aenter__(self):
91
- if self._session_cms:
92
- # share the current session, push a None as context manager to avoid close it in aexit
93
- _, session = self._session_cms[-1]
94
- self._session_cms.append((None, session))
95
- else:
96
  # create new session
97
- session_cm = self.transport.connect_session(**self._session_kwargs)
98
- session = await session_cm.__aenter__()
99
- self._session_cms.append((session_cm, session))
 
100
  return self
101
 
102
  async def __aexit__(self, exc_type, exc_val, exc_tb):
103
- cm, _ = self._session_cms.pop()
104
- if cm is not None:
105
- await cm.__aexit__(exc_type, exc_val, exc_tb)
 
 
 
106
 
107
  # --- MCP Client Methods ---
108
  async def ping(self) -> None:
 
44
  read_timeout_seconds: datetime.timedelta | None = None,
45
  ):
46
  self.transport = infer_transport(transport)
47
+ self._session: ClientSession | None = None
48
+ self._session_cms: AbstractAsyncContextManager[ClientSession] | None = None
49
+ self._nesting_counter: int = 0
 
50
 
51
  self._session_kwargs: SessionKwargs = {
52
  "sampling_callback": None,
 
65
  @property
66
  def session(self) -> ClientSession:
67
  """Get the current active session. Raises RuntimeError if not connected."""
68
+ if not self._session:
69
  raise RuntimeError(
70
  "Client is not connected. Use 'async with client:' context manager first."
71
  )
72
+ return self._session
 
73
 
74
  def set_roots(self, roots: RootsList | RootsHandler) -> None:
75
  """Set the roots for the client. This does not automatically call `send_roots_list_changed`."""
 
83
 
84
  def is_connected(self) -> bool:
85
  """Check if the client is currently connected."""
86
+ return self._session is not None
87
 
88
  async def __aenter__(self):
89
+ if self._nesting_counter == 0:
 
 
 
 
90
  # create new session
91
+ self._session_cm = self.transport.connect_session(**self._session_kwargs)
92
+ self._session = await self._session_cm.__aenter__()
93
+
94
+ self._nesting_counter += 1
95
  return self
96
 
97
  async def __aexit__(self, exc_type, exc_val, exc_tb):
98
+ self._nesting_counter -= 0
99
+
100
+ if self._nesting_counter == 0 and self._session_cms is not None:
101
+ await self._session_cms.__aexit__(exc_type, exc_val, exc_tb)
102
+ self._session_cms = None
103
+ self._session = None
104
 
105
  # --- MCP Client Methods ---
106
  async def ping(self) -> None: