Jeremiah Lowin commited on
Commit
917b76a
·
1 Parent(s): 74262ae

Ensure close() cleans up clients

Browse files
Files changed (1) hide show
  1. src/fastmcp/client/client.py +39 -13
src/fastmcp/client/client.py CHANGED
@@ -150,13 +150,6 @@ class Client(Generic[ClientTransportT]):
150
  self.transport = cast(ClientTransportT, infer_transport(transport))
151
  if auth is not None:
152
  self.transport._set_auth(auth)
153
- self._session: ClientSession | None = None
154
- self._exit_stack: AsyncExitStack | None = None
155
- self._nesting_counter: int = 0
156
- self._context_lock = anyio.Lock()
157
- self._session_task: asyncio.Task | None = None
158
- self._ready_event = asyncio.Event()
159
- self._stop_event = asyncio.Event()
160
  self._initialize_result: mcp.types.InitializeResult | None = None
161
 
162
  if log_handler is None:
@@ -196,7 +189,15 @@ class Client(Generic[ClientTransportT]):
196
  self._session_kwargs["sampling_callback"] = create_sampling_callback(
197
  sampling_handler
198
  )
199
- # self._session_manager = self._context_manager()
 
 
 
 
 
 
 
 
200
 
201
  @property
202
  def session(self) -> ClientSession:
@@ -252,6 +253,14 @@ class Client(Generic[ClientTransportT]):
252
  self._initialize_result = None
253
 
254
  async def __aenter__(self):
 
 
 
 
 
 
 
 
255
  async with self._context_lock:
256
  need_to_start = self._session_task is None or self._session_task.done()
257
  if need_to_start:
@@ -262,19 +271,37 @@ class Client(Generic[ClientTransportT]):
262
  self._nesting_counter += 1
263
  return self
264
 
265
- async def __aexit__(self, exc_type, exc_val, exc_tb):
 
266
  async with self._context_lock:
267
- self._nesting_counter -= 1
268
- if self._nesting_counter != 0:
 
 
 
 
 
 
 
 
 
 
 
 
269
  return
270
  self._stop_event.set()
271
  runner_task = self._session_task
272
  self._session_task = None
 
 
273
  if runner_task:
274
  await runner_task
 
275
  # Reset for future reconnects
276
  self._stop_event = anyio.Event()
277
  self._ready_event = anyio.Event()
 
 
278
 
279
  async def _session_runner(self):
280
  async with AsyncExitStack() as stack:
@@ -289,9 +316,8 @@ class Client(Generic[ClientTransportT]):
289
  self._ready_event.set()
290
 
291
  async def close(self):
 
292
  await self.transport.close()
293
- self._session = None
294
- self._initialize_result = None
295
 
296
  # --- MCP Client Methods ---
297
 
 
150
  self.transport = cast(ClientTransportT, infer_transport(transport))
151
  if auth is not None:
152
  self.transport._set_auth(auth)
 
 
 
 
 
 
 
153
  self._initialize_result: mcp.types.InitializeResult | None = None
154
 
155
  if log_handler is None:
 
189
  self._session_kwargs["sampling_callback"] = create_sampling_callback(
190
  sampling_handler
191
  )
192
+
193
+ # session context management
194
+ self._session: ClientSession | None = None
195
+ self._exit_stack: AsyncExitStack | None = None
196
+ self._nesting_counter: int = 0
197
+ self._context_lock = anyio.Lock()
198
+ self._session_task: asyncio.Task | None = None
199
+ self._ready_event = asyncio.Event()
200
+ self._stop_event = asyncio.Event()
201
 
202
  @property
203
  def session(self) -> ClientSession:
 
253
  self._initialize_result = None
254
 
255
  async def __aenter__(self):
256
+ await self._connect()
257
+ return self
258
+
259
+ async def __aexit__(self, exc_type, exc_val, exc_tb):
260
+ await self._disconnect()
261
+
262
+ async def _connect(self):
263
+ # ensure only one session is running at a time to avoid race conditions
264
  async with self._context_lock:
265
  need_to_start = self._session_task is None or self._session_task.done()
266
  if need_to_start:
 
271
  self._nesting_counter += 1
272
  return self
273
 
274
+ async def _disconnect(self, force: bool = False):
275
+ # ensure only one session is running at a time to avoid race conditions
276
  async with self._context_lock:
277
+ # if we are forcing a disconnect, reset the nesting counter
278
+ if force:
279
+ self._nesting_counter = 0
280
+
281
+ # otherwise decrement to check if we are done nesting
282
+ else:
283
+ self._nesting_counter = max(0, self._nesting_counter - 1)
284
+
285
+ # if we are still nested, return
286
+ if self._nesting_counter > 0:
287
+ return
288
+
289
+ # stop the active seesion
290
+ if self._session_task is None:
291
  return
292
  self._stop_event.set()
293
  runner_task = self._session_task
294
  self._session_task = None
295
+
296
+ # wait for the session to finish
297
  if runner_task:
298
  await runner_task
299
+
300
  # Reset for future reconnects
301
  self._stop_event = anyio.Event()
302
  self._ready_event = anyio.Event()
303
+ self._session = None
304
+ self._initialize_result = None
305
 
306
  async def _session_runner(self):
307
  async with AsyncExitStack() as stack:
 
316
  self._ready_event.set()
317
 
318
  async def close(self):
319
+ await self._disconnect(force=True)
320
  await self.transport.close()
 
 
321
 
322
  # --- MCP Client Methods ---
323