hopeful0 commited on
Commit
0a8046f
·
unverified ·
1 Parent(s): 7e77681

Refactor Client context management to avoid concurrency issue (#1054)

Browse files
Files changed (1) hide show
  1. src/fastmcp/client/client.py +21 -40
src/fastmcp/client/client.py CHANGED
@@ -285,23 +285,7 @@ class Client(Generic[ClientTransportT]):
285
  self._initialize_result = None
286
 
287
  async def __aenter__(self):
288
- await self._connect()
289
-
290
- # Check if session task failed and raise error immediately
291
- if (
292
- self._session_task is not None
293
- and self._session_task.done()
294
- and not self._session_task.cancelled()
295
- ):
296
- exception = self._session_task.exception()
297
- if isinstance(exception, httpx.HTTPStatusError):
298
- raise exception
299
- elif exception is not None:
300
- raise RuntimeError(
301
- f"Client failed to connect: {exception}"
302
- ) from exception
303
-
304
- return self
305
 
306
  async def __aexit__(self, exc_type, exc_val, exc_tb):
307
  await self._disconnect()
@@ -311,10 +295,21 @@ class Client(Generic[ClientTransportT]):
311
  async with self._context_lock:
312
  need_to_start = self._session_task is None or self._session_task.done()
313
  if need_to_start:
 
314
  self._stop_event = anyio.Event()
315
  self._ready_event = anyio.Event()
316
  self._session_task = asyncio.create_task(self._session_runner())
317
- await self._ready_event.wait()
 
 
 
 
 
 
 
 
 
 
318
  self._nesting_counter += 1
319
  return self
320
 
@@ -337,35 +332,21 @@ class Client(Generic[ClientTransportT]):
337
  if self._session_task is None:
338
  return
339
  self._stop_event.set()
340
- runner_task = self._session_task
 
341
  self._session_task = None
342
 
343
- # wait for the session to finish
344
- if runner_task:
345
- await runner_task
346
-
347
- # Reset for future reconnects
348
- self._stop_event = anyio.Event()
349
- self._ready_event = anyio.Event()
350
- self._session = None
351
- self._initialize_result = None
352
-
353
  async def _session_runner(self):
354
  try:
355
  async with AsyncExitStack() as stack:
356
- try:
357
- await stack.enter_async_context(self._context_manager())
358
- # Session/context is now ready
359
- self._ready_event.set()
360
- # Wait until disconnect/stop is requested
361
- await self._stop_event.wait()
362
- finally:
363
- # On exit, ensure ready event is set (idempotent)
364
- self._ready_event.set()
365
- except Exception:
366
  # Ensure ready event is set even if context manager entry fails
367
  self._ready_event.set()
368
- raise
369
 
370
  async def close(self):
371
  await self._disconnect(force=True)
 
285
  self._initialize_result = None
286
 
287
  async def __aenter__(self):
288
+ return await self._connect()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
289
 
290
  async def __aexit__(self, exc_type, exc_val, exc_tb):
291
  await self._disconnect()
 
295
  async with self._context_lock:
296
  need_to_start = self._session_task is None or self._session_task.done()
297
  if need_to_start:
298
+ assert self._nesting_counter == 0
299
  self._stop_event = anyio.Event()
300
  self._ready_event = anyio.Event()
301
  self._session_task = asyncio.create_task(self._session_runner())
302
+ await self._ready_event.wait()
303
+
304
+ if self._session_task.done():
305
+ exception = self._session_task.exception()
306
+ assert exception is not None
307
+ if isinstance(exception, httpx.HTTPStatusError):
308
+ raise exception
309
+ raise RuntimeError(
310
+ f"Client failed to connect: {exception}"
311
+ ) from exception
312
+
313
  self._nesting_counter += 1
314
  return self
315
 
 
332
  if self._session_task is None:
333
  return
334
  self._stop_event.set()
335
+ # wait for session to finish to ensure state has been reset
336
+ await self._session_task
337
  self._session_task = None
338
 
 
 
 
 
 
 
 
 
 
 
339
  async def _session_runner(self):
340
  try:
341
  async with AsyncExitStack() as stack:
342
+ await stack.enter_async_context(self._context_manager())
343
+ # Session/context is now ready
344
+ self._ready_event.set()
345
+ # Wait until disconnect/stop is requested
346
+ await self._stop_event.wait()
347
+ finally:
 
 
 
 
348
  # Ensure ready event is set even if context manager entry fails
349
  self._ready_event.set()
 
350
 
351
  async def close(self):
352
  await self._disconnect(force=True)