Spaces:
Running
Running
hopeful0 commited on
Refactor Client context management to avoid concurrency issue (#1054)
Browse files- 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
| 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 |
-
|
| 357 |
-
|
| 358 |
-
|
| 359 |
-
|
| 360 |
-
|
| 361 |
-
|
| 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)
|