Jeremiah Lowin commited on
Commit
0c34dfa
·
unverified ·
2 Parent(s): 13c4937132eee4

Merge pull request #643 from jlowin/concurrency-2

Browse files

Ensure close() cleans up client context appropriately

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 = anyio.Event()
200
+ self._stop_event = anyio.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
 
src/fastmcp/client/transports.py CHANGED
@@ -18,6 +18,7 @@ from typing import (
18
  overload,
19
  )
20
 
 
21
  import httpx
22
  from mcp import ClientSession, StdioServerParameters
23
  from mcp.client.session import (
@@ -327,8 +328,8 @@ class StdioTransport(ClientTransport):
327
 
328
  self._session: ClientSession | None = None
329
  self._connect_task: asyncio.Task | None = None
330
- self._ready_event = asyncio.Event()
331
- self._stop_event = asyncio.Event()
332
 
333
  @contextlib.asynccontextmanager
334
  async def connect_session(
@@ -391,8 +392,8 @@ class StdioTransport(ClientTransport):
391
 
392
  # reset variables and events for potential future reconnects
393
  self._connect_task = None
394
- self._stop_event = asyncio.Event()
395
- self._ready_event = asyncio.Event()
396
 
397
  async def close(self):
398
  await self.disconnect()
 
18
  overload,
19
  )
20
 
21
+ import anyio
22
  import httpx
23
  from mcp import ClientSession, StdioServerParameters
24
  from mcp.client.session import (
 
328
 
329
  self._session: ClientSession | None = None
330
  self._connect_task: asyncio.Task | None = None
331
+ self._ready_event = anyio.Event()
332
+ self._stop_event = anyio.Event()
333
 
334
  @contextlib.asynccontextmanager
335
  async def connect_session(
 
392
 
393
  # reset variables and events for potential future reconnects
394
  self._connect_task = None
395
+ self._stop_event = anyio.Event()
396
+ self._ready_event = anyio.Event()
397
 
398
  async def close(self):
399
  await self.disconnect()
tests/client/test_client.py CHANGED
@@ -342,6 +342,52 @@ async def test_client_nested_context_manager(fastmcp_server):
342
  assert client._session is None
343
 
344
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
345
  async def test_resource_template(fastmcp_server):
346
  """Test using a resource template with InMemoryClient."""
347
  client = Client(transport=FastMCPTransport(fastmcp_server))
 
342
  assert client._session is None
343
 
344
 
345
+ async def test_concurrent_client_context_managers():
346
+ """
347
+ Test that concurrent client usage doesn't cause cross-task cancel scope issues.
348
+ https://github.com/jlowin/fastmcp/pull/643
349
+ """
350
+ # Create a simple server
351
+ server = FastMCP("Test Server")
352
+
353
+ @server.tool()
354
+ def echo(text: str) -> str:
355
+ """Echo tool"""
356
+ return text
357
+
358
+ # Create client
359
+ client = Client(server)
360
+
361
+ # Track results
362
+ results = {}
363
+ errors = []
364
+
365
+ async def use_client(task_id: str, delay: float = 0):
366
+ """Use the client with a small delay to ensure overlap"""
367
+ try:
368
+ async with client:
369
+ # Add a small delay to ensure contexts overlap
370
+ await asyncio.sleep(delay)
371
+ # Make an actual call to exercise the session
372
+ tools = await client.list_tools()
373
+ results[task_id] = len(tools)
374
+ except Exception as e:
375
+ errors.append((task_id, str(e)))
376
+
377
+ # Run multiple tasks concurrently
378
+ # The key is having them enter and exit the context at different times
379
+ await asyncio.gather(
380
+ use_client("task1", 0.0),
381
+ use_client("task2", 0.01), # Slight delay to ensure overlap
382
+ use_client("task3", 0.02),
383
+ return_exceptions=False,
384
+ )
385
+
386
+ assert len(errors) == 0, f"Errors occurred: {errors}"
387
+ assert len(results) == 3
388
+ assert all(count == 1 for count in results.values()) # All should see 1 tool
389
+
390
+
391
  async def test_resource_template(fastmcp_server):
392
  """Test using a resource template with InMemoryClient."""
393
  client = Client(transport=FastMCPTransport(fastmcp_server))
tests/server/test_logging.py CHANGED
@@ -2,6 +2,7 @@ import asyncio
2
  import logging
3
  from unittest.mock import AsyncMock, Mock, patch
4
 
 
5
  import pytest
6
 
7
  from fastmcp.server.server import FastMCP
@@ -27,7 +28,7 @@ async def test_uvicorn_logging_default_level(
27
  """Tests that FastMCP passes log_level to uvicorn.Config if no log_config is given."""
28
  mock_server_instance = AsyncMock()
29
  mock_uvicorn_server_constructor.return_value = mock_server_instance
30
- serve_finished_event = asyncio.Event()
31
  mock_server_instance.serve.side_effect = serve_finished_event.wait
32
 
33
  test_log_level = "warning"
@@ -63,7 +64,7 @@ async def test_uvicorn_logging_with_custom_log_config(
63
  """Tests that FastMCP passes log_config to uvicorn.Config and not log_level."""
64
  mock_server_instance = AsyncMock()
65
  mock_uvicorn_server_constructor.return_value = mock_server_instance
66
- serve_finished_event = asyncio.Event()
67
  mock_server_instance.serve.side_effect = serve_finished_event.wait
68
 
69
  sample_log_config = {
@@ -123,7 +124,7 @@ async def test_uvicorn_logging_custom_log_config_overrides_log_level_param(
123
  """Tests log_config precedence if log_level is also passed to run_http_async."""
124
  mock_server_instance = AsyncMock()
125
  mock_uvicorn_server_constructor.return_value = mock_server_instance
126
- serve_finished_event = asyncio.Event()
127
  mock_server_instance.serve.side_effect = serve_finished_event.wait
128
 
129
  sample_log_config = {
 
2
  import logging
3
  from unittest.mock import AsyncMock, Mock, patch
4
 
5
+ import anyio
6
  import pytest
7
 
8
  from fastmcp.server.server import FastMCP
 
28
  """Tests that FastMCP passes log_level to uvicorn.Config if no log_config is given."""
29
  mock_server_instance = AsyncMock()
30
  mock_uvicorn_server_constructor.return_value = mock_server_instance
31
+ serve_finished_event = anyio.Event()
32
  mock_server_instance.serve.side_effect = serve_finished_event.wait
33
 
34
  test_log_level = "warning"
 
64
  """Tests that FastMCP passes log_config to uvicorn.Config and not log_level."""
65
  mock_server_instance = AsyncMock()
66
  mock_uvicorn_server_constructor.return_value = mock_server_instance
67
+ serve_finished_event = anyio.Event()
68
  mock_server_instance.serve.side_effect = serve_finished_event.wait
69
 
70
  sample_log_config = {
 
124
  """Tests log_config precedence if log_level is also passed to run_http_async."""
125
  mock_server_instance = AsyncMock()
126
  mock_uvicorn_server_constructor.return_value = mock_server_instance
127
+ serve_finished_event = anyio.Event()
128
  mock_server_instance.serve.side_effect = serve_finished_event.wait
129
 
130
  sample_log_config = {