From 0217f704df4dace81f47d8cbfe041adb5f650cbb Mon Sep 17 00:00:00 2001 From: Eric Hare Date: Thu, 23 Apr 2026 11:25:12 -0700 Subject: [PATCH] fix: serialize concurrent MCP session access to prevent race conditions (#12761) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: serialize concurrent MCP session access to prevent race conditions Two MCPTools components pointing at the same SSE URL share a single MCPSessionManager via the component cache. Under concurrent flow execution (10+ runs against the same server), ~40-50% of runs failed with one of: Error updating tool list: 'streamable_http__0' Timeout updating tool list: ... Error updating tool list: Three races caused this: 1. get_session() was not serialized per-server. Concurrent callers iterating self.sessions_by_server[server_key]["sessions"] could both pass the health check, then both invoke _cleanup_session_by_id(). 2. _cleanup_session_by_id() used del sessions[session_id] in a finally block. Two callers that both passed the `if session_id not in sessions` guard would race on the delete — the loser raised KeyError: 'streamable_http__0', matching the reported symptom. 3. Session ids were generated from len(sessions), so removing and re-adding sessions could silently produce colliding ids. Fixes: - Per-server asyncio.Lock (guarded by a module-level lock for creation) serializes session reuse/creation/cleanup. - _cleanup_session_by_id() now pops the session entry up front; only the winning caller runs teardown, the rest no-op. - Session ids come from a monotonic per-server counter. Regression tests cover all three races: 10 concurrent get_session calls must share a single created session; 10 concurrent cleanups must not raise; session ids must not recycle "_0" after cleanup. Fixes langflow-ai/langflow#9860 * fix: extend per-server lock to cleanup paths and reclaim per-key maps Addresses review findings on the previous commit: HIGH: cleanup paths bypassed the per-server lock. Both `_cleanup_session(context_id)` (invoked from client disconnect) and `_cleanup_idle_sessions()` (invoked from the periodic background task) still mutated `sessions_by_server` without holding the lock. A concurrent `get_session()` could finish validating a session while idle-cleanup popped and cancelled its task; `get_session()` then returned a dead session and registered a dangling refcount entry. MEDIUM: per-server maps (`_server_locks`, `_session_id_counters`) grew unboundedly. The previous patch only reclaimed server entries from `sessions_by_server`; rotating auth/session headers (which change `server_key` via `_get_server_key`) leaked entries in the other two maps forever in long-lived processes. Changes: - Replace `_get_server_lock()` with `_server_lock()`, an async context manager that pin-counts the entry so reclamation can't race a task that holds or is about to acquire the lock. - Wrap the mutating regions of `_cleanup_session()` and `_cleanup_idle_sessions()` in `_server_lock(server_key)`. - On lock release, reclaim both `_server_locks[server_key]` and `_session_id_counters[server_key]` once pins drop to 0, the lock is unheld, and the server has no remaining sessions. - `cleanup_all()` also clears the two maps under `_locks_guard`. New regression tests: - `test_cleanup_idle_vs_get_session_are_serialized` — an in-flight `get_session` blocks a concurrent idle-cleanup pass; the session is returned live, not torn down underneath the caller. - `test_concurrent_cleanup_session_and_get_session_safe` — refcount transitions stay consistent across concurrent connect/disconnect. - `test_server_lock_and_counter_reclaimed_when_unused` — after all sessions for a `server_key` are removed, both `_server_locks` and `_session_id_counters` entries go away. * fix: CAS context mapping pop to preserve cross-server handoffs Addresses the review finding on cross-server reuse of `context_id`: `_cleanup_session(context_id)` runs under the per-server lock of the server the mapping pointed at when cleanup started. That lock does NOT serialize a concurrent `get_session(context_id, different_server)` which runs under a *different* per-server lock and atomically re-points `_context_to_session[context_id]` at the new session. The old cleanup path then ran `self._context_to_session.pop(context_id, None)` unconditionally, wiping out the fresh mapping. The new session on the other server was left with refcount 1 and no context mapping, so the next disconnect was a no-op and the session leaked indefinitely. Fix: CAS the pop. After the refcount decrement / teardown, only drop the mapping if `_context_to_session[context_id]` still points at `(server_key, session_id)` — the pair we just cleaned up. The `get()` and `pop()` run synchronously (no `await` between them) so asyncio cannot interleave another coroutine between the check and the mutation. New regression test `test_cleanup_does_not_wipe_cross_server_handoff` slows down `_cleanup_session_by_id` for server A while a concurrent `get_session(ctx, serverB)` re-points the context at server B. Without the CAS, the assertion that `_context_to_session[ctx] == (server_B, ...)` fails. With the CAS, the fresh B mapping survives and its refcount is intact. * refactor(mcp): address review findings on concurrent-access fix IMPORTANT - Replace `dict[str, Any]` with `_ServerLockEntry` TypedDict so the lock/pins shape is visible to static analysis (finding #1). - Add TODO marker above `MCPSessionManager` flagging the module's file-size debt as follow-up work (finding #2). RECOMMENDED - Replace timing-based `asyncio.sleep(0.02)` / `asyncio.sleep(0.05)` in concurrency regression tests with deterministic Event-based rendezvous and pin-count polling (finding #3). - Introduce `_sessions_for(server_key)` helper and use it in `_cleanup_idle_sessions`, `_cleanup_session_by_id`, and `cleanup_all` to stop reaching through `sessions_by_server[server_key]["sessions"]` at every call site (finding #4). - Document why `_cleanup_session_by_id` keeps a broad `except Exception` (transport-layer teardown raises many different exception hierarchies — leak-on-cleanup is worse than a swallowed error) (finding #5). - Log a warning when pin count goes negative in `_release_server_lock_if_idle` so a missing acquire / double release surfaces in telemetry instead of being silently swept (finding #6). - Document the "caller must hold `_server_lock(server_key)`" invariant on `_next_session_id` (finding #7). NICE TO HAVE - Annotate the `_server_lock` async context manager yield type as `AsyncIterator[None]` for IDE support (finding #8). All 175 tests in `test_mcp_util.py` still pass; ruff clean. --- .../tests/unit/base/mcp/test_mcp_util.py | 348 ++++++++++++++++ src/lfx/src/lfx/base/mcp/util.py | 391 ++++++++++++------ 2 files changed, 617 insertions(+), 122 deletions(-) diff --git a/src/backend/tests/unit/base/mcp/test_mcp_util.py b/src/backend/tests/unit/base/mcp/test_mcp_util.py index 80bdcc922d..954b7ba4b8 100644 --- a/src/backend/tests/unit/base/mcp/test_mcp_util.py +++ b/src/backend/tests/unit/base/mcp/test_mcp_util.py @@ -127,6 +127,354 @@ class TestMCPSessionManager: assert session1 != session2 assert mock_create.call_count == 2 + async def test_concurrent_get_session_same_server_reuses_one_session(self, session_manager): + """Concurrent get_session calls for the same server must share one session. + + Regression test for the race condition reported in + https://github.com/langflow-ai/langflow/issues/9860 where two MCPTools + components pointing at the same SSE URL would race on session + creation/cleanup under concurrent flow execution and intermittently + fail with errors such as: + Error updating tool list: 'streamable_http__0' + Timeout updating tool list: ... + """ + connection_params = {"url": "http://example.test/sse", "headers": {}} + + mock_session = AsyncMock() + mock_task = AsyncMock() + mock_task.done = MagicMock(return_value=False) + + create_calls = 0 + # Deterministic rendezvous: the first caller enters ``fake_create`` and + # parks on ``release``, giving every other caller time to queue on the + # per-server lock. Without the lock, they would all enter the creation + # path and ``create_calls`` would exceed 1. This replaces a timing- + # based ``asyncio.sleep(0.05)`` so the test does not rely on CI speed. + create_started = asyncio.Event() + release = asyncio.Event() + + async def fake_create(_session_id, _params, _preferred_transport=None): + nonlocal create_calls + create_calls += 1 + create_started.set() + await release.wait() + return mock_session, mock_task, "streamable_http", False + + with ( + patch.object(session_manager, "_create_streamable_http_session", side_effect=fake_create), + patch.object(session_manager, "_validate_session_connectivity", return_value=True), + ): + getters = [ + asyncio.create_task(session_manager.get_session(f"ctx_{i}", connection_params, "streamable_http")) + for i in range(10) + ] + # Wait until the first caller is in ``fake_create``; every other + # caller is now queued on the per-server lock. + await create_started.wait() + release.set() + results = await asyncio.gather(*getters) + + assert all(s is mock_session for s in results) + # All 10 concurrent callers must share the single created session. + assert create_calls == 1 + server_key = session_manager._get_server_key(connection_params, "streamable_http") + assert len(session_manager.sessions_by_server[server_key]["sessions"]) == 1 + + async def test_concurrent_cleanup_same_session_is_idempotent(self, session_manager): + """_cleanup_session_by_id must be safe under concurrent invocation. + + Previously the finally-block `del sessions[session_id]` raised + `KeyError` when two callers both passed the existence guard — the + visible symptom from issue #9860. + """ + server_key = "streamable_http_test" + session_id = f"{server_key}_0" + + mock_task = AsyncMock() + mock_task.done = MagicMock(return_value=False) + mock_task.cancel = MagicMock() + + session_manager.sessions_by_server[server_key] = { + "sessions": { + session_id: { + "session": AsyncMock(), + "task": mock_task, + "type": "streamable_http", + "last_used": 0, + } + }, + "last_cleanup": 0, + } + + # Fire many concurrent cleanups; only one should actually tear the + # session down, the rest should be no-ops (not KeyError). + await asyncio.gather(*[session_manager._cleanup_session_by_id(server_key, session_id) for _ in range(10)]) + + assert session_id not in session_manager.sessions_by_server[server_key]["sessions"] + mock_task.cancel.assert_called_once() + + async def test_session_ids_are_monotonic_not_len_based(self, session_manager): + """Session ids must come from a monotonic counter, not len(sessions). + + With `len(sessions)` as the id source, removing and re-adding sessions + can produce colliding ids and silently overwrite a live session entry. + """ + connection_params = {"url": "http://example.test/sse", "headers": {}} + + async def fake_create(_session_id, _params, _preferred_transport=None): + # Return a fresh session/task for each creation. + s = AsyncMock() + t = AsyncMock() + t.done = MagicMock(return_value=False) + t.cancel = MagicMock() + return s, t, "streamable_http", False + + server_key = session_manager._get_server_key(connection_params, "streamable_http") + + with ( + patch.object(session_manager, "_create_streamable_http_session", side_effect=fake_create), + # Force health check to fail so each call creates a fresh session. + patch.object(session_manager, "_validate_session_connectivity", return_value=False), + ): + ids_seen: list[str] = [] + for i in range(3): + await session_manager.get_session(f"ctx_{i}", connection_params, "streamable_http") + ids_seen.extend(session_manager.sessions_by_server[server_key]["sessions"].keys()) + + # Counter keeps advancing even as old sessions are cleaned up. + assert session_manager._session_id_counters[server_key] == 3 + # And the new id is unique (never recycles "_0"). + assert f"{server_key}_0" not in session_manager.sessions_by_server[server_key]["sessions"] + + async def test_cleanup_idle_vs_get_session_are_serialized(self, session_manager): + """Idle cleanup must not race a concurrent get_session(). + + Without holding the per-server lock, the idle-cleanup task could pop + and cancel a session mid-validation; `get_session()` would then hand + the caller a dead session plus a dangling refcount entry. + """ + connection_params = {"url": "http://example.test/sse", "headers": {}} + server_key = session_manager._get_server_key(connection_params, "streamable_http") + + mock_session = AsyncMock() + mock_task = AsyncMock() + mock_task.done = MagicMock(return_value=False) + mock_task.cancel = MagicMock() + + # Seed one idle session (last_used in the distant past). + session_manager.sessions_by_server[server_key] = { + "sessions": { + f"{server_key}_0": { + "session": mock_session, + "task": mock_task, + "type": "streamable_http", + "last_used": 0, # definitely past the idle timeout + } + }, + "last_cleanup": 0, + } + + started = asyncio.Event() + release = asyncio.Event() + + async def slow_validate(_session): + started.set() + await release.wait() + return True + + with patch.object(session_manager, "_validate_session_connectivity", side_effect=slow_validate): + # Start a get_session() that will block inside the health check + # while holding the per-server lock. + getter = asyncio.create_task( + session_manager.get_session("ctx_reader", connection_params, "streamable_http") + ) + await started.wait() + + # Fire the idle cleanup concurrently. It must wait for the lock. + cleaner = asyncio.create_task(session_manager._cleanup_idle_sessions()) + # Deterministic rendezvous: wait until the cleaner has pinned the + # per-server lock (but is still blocked on the getter releasing + # it). Polling ``pins`` with scheduler-only yields replaces a + # timing-based ``asyncio.sleep(0.02)`` so the test does not rely + # on CI speed to expose a race. ``asyncio.Lock`` does not expose + # a waiters count, so internal-state polling is the deterministic + # primitive we have. + while session_manager._server_locks.get(server_key, {}).get("pins", 0) < 2: # noqa: ASYNC110 + await asyncio.sleep(0) + + # While the getter still holds the lock, the session must not have + # been cleaned up. + assert f"{server_key}_0" in session_manager.sessions_by_server[server_key]["sessions"] + assert not mock_task.cancel.called + + # Let the getter finish. + release.set() + result = await getter + await cleaner + + # Getter returned the healthy cached session. Because `get_session` + # bumps `last_used` on reuse, the idle-cleanup pass — which ran only + # after the getter released the lock — correctly left the session + # alone. Without the lock, the cleaner could have torn it down + # mid-validation and the getter would have returned a dead session. + assert result is mock_session + assert not mock_task.cancel.called + assert f"{server_key}_0" in session_manager.sessions_by_server[server_key]["sessions"] + + async def test_concurrent_cleanup_session_and_get_session_safe(self, session_manager): + """Disconnect path must honor the per-server lock. + + `_cleanup_session(context_id)` previously mutated refcounts without + the per-server lock, so a concurrent `get_session()` could observe + state mid-transition and return a session that's about to be torn + down by the other caller's disconnect path. + """ + connection_params = {"url": "http://example.test/sse", "headers": {}} + + mock_session = AsyncMock() + mock_task = AsyncMock() + mock_task.done = MagicMock(return_value=False) + mock_task.cancel = MagicMock() + + async def fake_create(_session_id, _params, _preferred_transport=None): + return mock_session, mock_task, "streamable_http", False + + with ( + patch.object(session_manager, "_create_streamable_http_session", side_effect=fake_create), + patch.object(session_manager, "_validate_session_connectivity", return_value=True), + ): + # Establish the session with one context. + s1 = await session_manager.get_session("ctx_a", connection_params, "streamable_http") + assert s1 is mock_session + + # Concurrently: a second caller grabs the session (bumping refcount) + # while the first caller disconnects (decrementing refcount). + results = await asyncio.gather( + session_manager.get_session("ctx_b", connection_params, "streamable_http"), + session_manager._cleanup_session("ctx_a"), + ) + + s2 = results[0] + # ctx_b still has a live session because ctx_a's cleanup only decremented + # refcount to 1, not zero. + assert s2 is mock_session + assert not mock_task.cancel.called + server_key = session_manager._get_server_key(connection_params, "streamable_http") + assert f"{server_key}_0" in session_manager.sessions_by_server[server_key]["sessions"] + + async def test_cleanup_does_not_wipe_cross_server_handoff(self, session_manager): + """Concurrent reconnect to a different server must not be wiped out. + + Regression test for a cross-server race on `_context_to_session`: + when `_cleanup_session(ctx)` is running for server A and a concurrent + `get_session(ctx, serverB)` re-points the same context at server B, + the cleanup previously ran `self._context_to_session.pop(context_id)` + unconditionally — destroying the fresh mapping. The new B session + then had refcount 1 forever, leaking on subsequent disconnect. + """ + server_a_params = {"url": "http://a.example.test/sse", "headers": {}} + server_b_params = {"url": "http://b.example.test/sse", "headers": {}} + server_key_a = session_manager._get_server_key(server_a_params, "streamable_http") + server_key_b = session_manager._get_server_key(server_b_params, "streamable_http") + + mock_session_a = AsyncMock() + mock_task_a = AsyncMock() + mock_task_a.done = MagicMock(return_value=False) + mock_task_a.cancel = MagicMock() + mock_session_b = AsyncMock() + mock_task_b = AsyncMock() + mock_task_b.done = MagicMock(return_value=False) + mock_task_b.cancel = MagicMock() + + async def fake_create(_session_id, params, _preferred_transport=None): + if params["url"] == server_a_params["url"]: + return mock_session_a, mock_task_a, "streamable_http", False + return mock_session_b, mock_task_b, "streamable_http", False + + with ( + patch.object(session_manager, "_create_streamable_http_session", side_effect=fake_create), + patch.object(session_manager, "_validate_session_connectivity", return_value=True), + ): + # Establish context on server A. + s_a = await session_manager.get_session("ctx_move", server_a_params, "streamable_http") + assert s_a is mock_session_a + + # Pause cleanup-of-A inside its lock so the concurrent + # get-for-B can race the context_to_session.pop(). + release = asyncio.Event() + original_cleanup = session_manager._cleanup_session_by_id + + async def slow_cleanup(server_key, session_id): + if server_key == server_key_a: + # Let the B-reconnect proceed while we hold server_A's lock. + await release.wait() + await original_cleanup(server_key, session_id) + + with patch.object(session_manager, "_cleanup_session_by_id", side_effect=slow_cleanup): + cleanup_task = asyncio.create_task(session_manager._cleanup_session("ctx_move")) + # Give cleanup a chance to enter server_A's lock and start awaiting. + await asyncio.sleep(0) + + # Concurrent reconnect to server B for the same context. This + # runs under server_B's lock, so it is not blocked. + s_b = await session_manager.get_session("ctx_move", server_b_params, "streamable_http") + assert s_b is mock_session_b + + # Now let the A-cleanup finish. With the CAS check, it must + # NOT pop the fresh (server_B, _) mapping. + release.set() + await cleanup_task + + # The fresh B mapping survives. + assert session_manager._context_to_session.get("ctx_move") == (server_key_b, f"{server_key_b}_0") + # A's session is gone; B's session is live with refcount 1. + assert f"{server_key_a}_0" not in session_manager.sessions_by_server.get(server_key_a, {}).get("sessions", {}) + assert f"{server_key_b}_0" in session_manager.sessions_by_server[server_key_b]["sessions"] + assert session_manager._session_refcount.get((server_key_b, f"{server_key_b}_0")) == 1 + + async def test_server_lock_and_counter_reclaimed_when_unused(self, session_manager): + """Per-server locks and id counters must be reclaimed with the server. + + Without this, rotating auth/session headers (which change `server_key` + via `_get_server_key`) causes `_server_locks` and + `_session_id_counters` to grow without bound in long-lived processes. + """ + connection_params = {"url": "http://example.test/sse", "headers": {}} + server_key = session_manager._get_server_key(connection_params, "streamable_http") + + mock_session = AsyncMock() + mock_task = AsyncMock() + mock_task.done = MagicMock(return_value=False) + mock_task.cancel = MagicMock() + + async def fake_create(_session_id, _params, _preferred_transport=None): + return mock_session, mock_task, "streamable_http", False + + with ( + patch.object(session_manager, "_create_streamable_http_session", side_effect=fake_create), + patch.object(session_manager, "_validate_session_connectivity", return_value=True), + ): + await session_manager.get_session("ctx_gc", connection_params, "streamable_http") + + # While the session is live, the lock entry exists (pin count back to 0 + # but the sessions_by_server entry holds the server_key alive). + assert server_key in session_manager._server_locks + assert server_key in session_manager._session_id_counters + + # Disconnect the only context. + await session_manager._cleanup_session("ctx_gc") + + # Trigger the periodic cleanup to remove the now-empty server entry. + # Mark the session as already expired so it gets swept on this pass, + # then run cleanup. + # (After _cleanup_session above, the sessions dict is already empty + # for this server_key, so _cleanup_idle_sessions will drop the entry.) + await session_manager._cleanup_idle_sessions() + + assert server_key not in session_manager.sessions_by_server + assert server_key not in session_manager._session_id_counters + assert server_key not in session_manager._server_locks + class TestHeaderValidation: """Test the header validation functionality.""" diff --git a/src/lfx/src/lfx/base/mcp/util.py b/src/lfx/src/lfx/base/mcp/util.py index e84d360711..3346a83ec1 100644 --- a/src/lfx/src/lfx/base/mcp/util.py +++ b/src/lfx/src/lfx/base/mcp/util.py @@ -9,9 +9,9 @@ import shlex import shutil import subprocess import unicodedata -from collections.abc import Awaitable, Callable +from collections.abc import AsyncIterator, Awaitable, Callable from types import UnionType -from typing import Any, Union, get_args, get_origin +from typing import Any, TypedDict, Union, get_args, get_origin from urllib.parse import urlparse from uuid import UUID @@ -816,6 +816,22 @@ def _is_mcp_session_bust_error(exc: BaseException) -> bool: return False +class _ServerLockEntry(TypedDict): + """Shape of each value in ``MCPSessionManager._server_locks``. + + ``pins`` is the number of callers that have obtained (but not yet + released) the lock via ``_server_lock``; it gates reclamation of the + entry so a new caller can't grab a fresh lock while an older caller is + about to enter the old one. + """ + + lock: asyncio.Lock + pins: int + + +# TODO(langflow-ai/langflow#12541-followup): MCPSessionManager lives in this +# 2k+ line module; extract it (and the concurrency primitives below) into a +# dedicated ``mcp/session_manager.py`` so future edits stay small. class MCPSessionManager: """Manages persistent MCP sessions with proper context manager lifecycle. @@ -838,9 +854,103 @@ class MCPSessionManager: # Cache which transport works for each server to avoid retrying failed transports # server_key -> "streamable_http" | "sse" self._transport_preference: dict[str, str] = {} + # Per-server asyncio locks to serialize session create/reuse/cleanup under + # concurrent access. Without this, two concurrent flow executions sharing + # the same MCP server URL can race on the sessions dict and raise a + # KeyError from `del sessions[session_id]` in `_cleanup_session_by_id`, or + # create colliding session_ids from `len(sessions)`. + # + # Each entry is a `_ServerLockEntry` {"lock": asyncio.Lock(), "pins": int}. + # The pin count is the number of callers that have obtained (but not yet + # released) the lock via `_server_lock`. We reclaim the entry only when + # pins == 0 and the lock is not held, to avoid a new caller grabbing a + # fresh lock while an older caller is about to enter the old one. + self._server_locks: dict[str, _ServerLockEntry] = {} + self._locks_guard = asyncio.Lock() + # Monotonic counter per server_key to generate unique session_ids even + # when sessions are removed between allocations. + self._session_id_counters: dict[str, int] = {} self._cleanup_task = None self._start_cleanup_task() + @contextlib.asynccontextmanager + async def _server_lock(self, server_key: str) -> AsyncIterator[None]: + """Acquire the per-server lock with pin counting for safe reclamation. + + The pin count prevents reclaiming a lock that another task is about to + enter (e.g. between obtaining a reference and calling ``async with``). + Reclamation in `_cleanup_idle_sessions` / `_release_server_lock_if_idle` + only runs when pins drop to zero *and* the lock is not held. + """ + async with self._locks_guard: + entry = self._server_locks.get(server_key) + if entry is None: + entry = _ServerLockEntry(lock=asyncio.Lock(), pins=0) + self._server_locks[server_key] = entry + entry["pins"] += 1 + lock = entry["lock"] + try: + async with lock: + yield + finally: + await self._release_server_lock_if_idle(server_key) + + async def _release_server_lock_if_idle(self, server_key: str): + """Drop the pin and, once the server is fully idle, reclaim the maps. + + Reclamation is deliberately conservative: we only drop the lock entry + (and the matching session-id counter) when *both* conditions hold — + pin count is zero and the server has no remaining sessions. This + prevents two problems: + - Churning the lock on every `get_session` call while a server is + actively in use (pin count oscillates 0↔1 between callers). + - Rotating auth/session headers (which change `server_key` via + `_get_server_key`) leaking per-key entries forever in long-lived + processes. + """ + async with self._locks_guard: + entry = self._server_locks.get(server_key) + if entry is None: + return + entry["pins"] -= 1 + if entry["pins"] < 0: + # A negative pin count means a missing acquire or a double release. + # Log loudly so it surfaces in telemetry instead of being swept. + await logger.awarning( + f"Negative pin count ({entry['pins']}) for server_key {server_key}; " + "this indicates a missing _server_lock acquire or a double release.", + ) + if entry["pins"] <= 0 and not entry["lock"].locked() and server_key not in self.sessions_by_server: + self._server_locks.pop(server_key, None) + self._session_id_counters.pop(server_key, None) + + def _next_session_id(self, server_key: str) -> str: + """Generate a monotonically unique session_id for *server_key*. + + Caller must hold ``self._server_lock(server_key)`` while invoking this. + The increment is otherwise unsynchronised — two concurrent callers + without the lock would race on ``_session_id_counters[server_key]`` and + produce colliding ids. + """ + current = self._session_id_counters.get(server_key, 0) + self._session_id_counters[server_key] = current + 1 + return f"{server_key}_{current}" + + def _sessions_for(self, server_key: str) -> dict[str, dict[str, Any]]: + """Return the sessions dict for *server_key* (empty dict if absent). + + Encapsulates the ``sessions_by_server[server_key]["sessions"]`` shape + so callers don't have to reach through the outer envelope. Handles the + legacy structure (sessions stored directly under the server_key) + uniformly as well. + """ + server_data = self.sessions_by_server.get(server_key) + if server_data is None: + return {} + if isinstance(server_data, dict) and "sessions" in server_data: + return server_data["sessions"] + return server_data # legacy flat structure + def _start_cleanup_task(self): """Start the periodic cleanup task.""" if self._cleanup_task is None or self._cleanup_task.done(): @@ -861,30 +971,37 @@ class MCPSessionManager: await logger.awarning(f"Error in periodic cleanup: {e}") async def _cleanup_idle_sessions(self): - """Clean up sessions that have been idle for too long.""" + """Clean up sessions that have been idle for too long. + + Acquires the per-server lock before mutating the sessions dict so we + don't race with `get_session()` — otherwise a concurrent `get_session` + could finish validating a session while this task pops and cancels it, + handing the caller a dead session plus a dangling refcount entry. + """ current_time = asyncio.get_event_loop().time() - servers_to_remove = [] - for server_key, server_data in self.sessions_by_server.items(): - sessions = server_data.get("sessions", {}) - sessions_to_remove = [] + # Snapshot keys to avoid mutating-while-iterating. + for server_key in list(self.sessions_by_server.keys()): + async with self._server_lock(server_key): + sessions = self._sessions_for(server_key) + if not sessions and server_key not in self.sessions_by_server: + continue - for session_id, session_info in list(sessions.items()): - if current_time - session_info["last_used"] > get_session_idle_timeout(): - sessions_to_remove.append(session_id) + sessions_to_remove = [ + session_id + for session_id, session_info in list(sessions.items()) + if current_time - session_info["last_used"] > get_session_idle_timeout() + ] - # Clean up idle sessions - for session_id in sessions_to_remove: - await logger.ainfo(f"Cleaning up idle session {session_id} for server {server_key}") - await self._cleanup_session_by_id(server_key, session_id) + for session_id in sessions_to_remove: + await logger.ainfo(f"Cleaning up idle session {session_id} for server {server_key}") + await self._cleanup_session_by_id(server_key, session_id) - # Remove server entry if no sessions left - if not sessions: - servers_to_remove.append(server_key) - - # Clean up empty server entries - for server_key in servers_to_remove: - del self.sessions_by_server[server_key] + # Remove server entry if no sessions left. The counter for + # this server_key is reclaimed by `_release_server_lock_if_idle` + # once this lock's pin count hits zero. + if not sessions: + self.sessions_by_server.pop(server_key, None) def _get_server_key(self, connection_params, transport_type: str) -> str: """Generate a consistent server key based on connection parameters.""" @@ -957,85 +1074,95 @@ class MCPSessionManager: The key insight is that we should reuse sessions based on the server identity (command + args for stdio, URL for Streamable HTTP) rather than the context_id. This prevents creating a new subprocess for each unique context. + + Concurrent callers for the same server are serialized via a per-server + lock. This is required to keep the `sessions` dict consistent across + concurrent flow executions that share a single `MCPSessionManager` + (e.g. two `MCPTools` components pointing at the same SSE URL). """ server_key = self._get_server_key(connection_params, transport_type) - # Ensure server entry exists - if server_key not in self.sessions_by_server: - self.sessions_by_server[server_key] = {"sessions": {}, "last_cleanup": asyncio.get_event_loop().time()} + async with self._server_lock(server_key): + # Ensure server entry exists + if server_key not in self.sessions_by_server: + self.sessions_by_server[server_key] = { + "sessions": {}, + "last_cleanup": asyncio.get_event_loop().time(), + } - server_data = self.sessions_by_server[server_key] - sessions = server_data["sessions"] + server_data = self.sessions_by_server[server_key] + sessions = server_data["sessions"] - # Try to find a healthy existing session - for session_id, session_info in list(sessions.items()): - session = session_info["session"] - task = session_info["task"] + # Try to find a healthy existing session + for session_id, session_info in list(sessions.items()): + session = session_info["session"] + task = session_info["task"] - # Check if session is still alive - if not task.done(): - # Update last used time - session_info["last_used"] = asyncio.get_event_loop().time() + # Check if session is still alive + if not task.done(): + # Update last used time + session_info["last_used"] = asyncio.get_event_loop().time() - # Quick health check - if await self._validate_session_connectivity(session): - await logger.adebug(f"Reusing existing session {session_id} for server {server_key}") - # record mapping & bump ref-count for backwards compatibility - self._context_to_session[context_id] = (server_key, session_id) - self._session_refcount[(server_key, session_id)] = ( - self._session_refcount.get((server_key, session_id), 0) + 1 - ) - return session - await logger.ainfo(f"Session {session_id} for server {server_key} failed health check, cleaning up") - await self._cleanup_session_by_id(server_key, session_id) + # Quick health check + if await self._validate_session_connectivity(session): + await logger.adebug(f"Reusing existing session {session_id} for server {server_key}") + # record mapping & bump ref-count for backwards compatibility + self._context_to_session[context_id] = (server_key, session_id) + self._session_refcount[(server_key, session_id)] = ( + self._session_refcount.get((server_key, session_id), 0) + 1 + ) + return session + await logger.ainfo(f"Session {session_id} for server {server_key} failed health check, cleaning up") + await self._cleanup_session_by_id(server_key, session_id) + else: + # Task is done, clean up + await logger.ainfo(f"Session {session_id} for server {server_key} task is done, cleaning up") + await self._cleanup_session_by_id(server_key, session_id) + + # Check if we've reached the maximum number of sessions for this server + if len(sessions) >= get_max_sessions_per_server(): + # Remove the oldest session + oldest_session_id = min(sessions.keys(), key=lambda x: sessions[x]["last_used"]) + await logger.ainfo( + f"Maximum sessions reached for server {server_key}, removing oldest session {oldest_session_id}" + ) + await self._cleanup_session_by_id(server_key, oldest_session_id) + + # Create new session. Use a monotonic counter so removed sessions + # don't cause id collisions with newly-created sessions. + session_id = self._next_session_id(server_key) + await logger.ainfo(f"Creating new session {session_id} for server {server_key}") + + if transport_type == "stdio": + session, task = await self._create_stdio_session(session_id, connection_params) + actual_transport = "stdio" + elif transport_type == "streamable_http": + # Pass the cached transport preference if available (SSE only when last success required it) + preferred_transport = self._transport_preference.get(server_key) + session, task, actual_transport, sse_pref_lock = await self._create_streamable_http_session( + session_id, connection_params, preferred_transport + ) + if actual_transport == "streamable_http": + self._transport_preference[server_key] = "streamable_http" + elif sse_pref_lock: + self._transport_preference[server_key] = "sse" else: - # Task is done, clean up - await logger.ainfo(f"Session {session_id} for server {server_key} task is done, cleaning up") - await self._cleanup_session_by_id(server_key, session_id) + msg = f"Unknown transport type: {transport_type}" + raise ValueError(msg) - # Check if we've reached the maximum number of sessions for this server - if len(sessions) >= get_max_sessions_per_server(): - # Remove the oldest session - oldest_session_id = min(sessions.keys(), key=lambda x: sessions[x]["last_used"]) - await logger.ainfo( - f"Maximum sessions reached for server {server_key}, removing oldest session {oldest_session_id}" - ) - await self._cleanup_session_by_id(server_key, oldest_session_id) + # Store session info with the actual transport used + sessions[session_id] = { + "session": session, + "task": task, + "type": actual_transport, + "last_used": asyncio.get_event_loop().time(), + } - # Create new session - session_id = f"{server_key}_{len(sessions)}" - await logger.ainfo(f"Creating new session {session_id} for server {server_key}") + # register mapping & initial ref-count for the new session + self._context_to_session[context_id] = (server_key, session_id) + self._session_refcount[(server_key, session_id)] = 1 - if transport_type == "stdio": - session, task = await self._create_stdio_session(session_id, connection_params) - actual_transport = "stdio" - elif transport_type == "streamable_http": - # Pass the cached transport preference if available (SSE only when last success required it) - preferred_transport = self._transport_preference.get(server_key) - session, task, actual_transport, sse_pref_lock = await self._create_streamable_http_session( - session_id, connection_params, preferred_transport - ) - if actual_transport == "streamable_http": - self._transport_preference[server_key] = "streamable_http" - elif sse_pref_lock: - self._transport_preference[server_key] = "sse" - else: - msg = f"Unknown transport type: {transport_type}" - raise ValueError(msg) - - # Store session info with the actual transport used - sessions[session_id] = { - "session": session, - "task": task, - "type": actual_transport, - "last_used": asyncio.get_event_loop().time(), - } - - # register mapping & initial ref-count for the new session - self._context_to_session[context_id] = (server_key, session_id) - self._session_refcount[(server_key, session_id)] = 1 - - return session + return session async def _create_stdio_session(self, session_id: str, connection_params): """Create a new stdio session as a background task to avoid context issues.""" @@ -1257,22 +1384,24 @@ class MCPSessionManager: raise ValueError(msg) from timeout_err async def _cleanup_session_by_id(self, server_key: str, session_id: str): - """Clean up a specific session by server key and session ID.""" - if server_key not in self.sessions_by_server: + """Clean up a specific session by server key and session ID. + + Safe against concurrent cleanup of the same session: we `pop` the entry + up front so two concurrent callers don't both try to cancel the same + task or `del` the same key (which raised `KeyError: 'streamable_http_..._0'` + previously under concurrent flow execution). + """ + sessions = self._sessions_for(server_key) + if not sessions and server_key not in self.sessions_by_server: return - server_data = self.sessions_by_server[server_key] - # Handle both old and new session structure - if isinstance(server_data, dict) and "sessions" in server_data: - sessions = server_data["sessions"] - else: - # Handle old structure where sessions were stored directly - sessions = server_data - - if session_id not in sessions: + # Atomically remove the session entry; only the caller that wins this + # pop performs the actual teardown. Concurrent callers get None and + # return early instead of racing on del/task.cancel(). + session_info = sessions.pop(session_id, None) + if session_info is None: return - session_info = sessions[session_id] try: # First try to properly close the session if it exists if "session" in session_info: @@ -1318,10 +1447,11 @@ class MCPSessionManager: except asyncio.CancelledError: await logger.ainfo(f"Cancelled task for session {session_id}") except Exception as e: # noqa: BLE001 + # Teardown is load-bearing: MCP transports (stdio subprocess, SSE, + # streamable HTTP) all raise their own exception hierarchies on + # shutdown, and a leak on cleanup is far worse than a swallowed + # error. Log and continue rather than propagating. await logger.awarning(f"Error cleaning up session {session_id}: {e}") - finally: - # Remove from sessions dict - del sessions[session_id] async def cleanup_all(self): """Clean up all sessions.""" @@ -1333,15 +1463,7 @@ class MCPSessionManager: # Clean up all sessions for server_key in list(self.sessions_by_server.keys()): - server_data = self.sessions_by_server[server_key] - # Handle both old and new session structure - if isinstance(server_data, dict) and "sessions" in server_data: - sessions = server_data["sessions"] - else: - # Handle old structure where sessions were stored directly - sessions = server_data - - for session_id in list(sessions.keys()): + for session_id in list(self._sessions_for(server_key).keys()): await self._cleanup_session_by_id(server_key, session_id) # Clear the sessions_by_server structure completely @@ -1351,6 +1473,13 @@ class MCPSessionManager: self._context_to_session.clear() self._session_refcount.clear() + # Reclaim per-server lock and counter maps. Safe here because + # cleanup_all is a shutdown/reset operation; no other manager state + # should be in use past this point. + async with self._locks_guard: + self._server_locks.clear() + self._session_id_counters.clear() + # Clear all background tasks for task in list(self._background_tasks): if not task.done(): @@ -1368,6 +1497,19 @@ class MCPSessionManager: Decrements the ref-count for the session used by *context_id* and only tears the session down when the last context that references it goes away. + + Acquires the per-server lock so concurrent `get_session()` calls don't + observe a half-torn-down session (e.g. returning a ClientSession whose + background task was just cancelled out from under them). + + Uses a compare-and-swap on `_context_to_session[context_id]` before + popping it: if a concurrent `get_session()` has re-pointed the same + context at a *different* server (e.g. a component reconnecting to a + new MCP URL while the old disconnect is in flight), we must not wipe + out the fresh mapping — otherwise the new session leaks. The per- + server lock doesn't cover this case because the new and old sessions + live under different server_keys, so the two operations run in + parallel. """ mapping = self._context_to_session.get(context_id) if not mapping: @@ -1375,17 +1517,22 @@ class MCPSessionManager: return server_key, session_id = mapping - ref_key = (server_key, session_id) - remaining = self._session_refcount.get(ref_key, 1) - 1 + async with self._server_lock(server_key): + ref_key = (server_key, session_id) + remaining = self._session_refcount.get(ref_key, 1) - 1 - if remaining <= 0: - await self._cleanup_session_by_id(server_key, session_id) - self._session_refcount.pop(ref_key, None) - else: - self._session_refcount[ref_key] = remaining + if remaining <= 0: + await self._cleanup_session_by_id(server_key, session_id) + self._session_refcount.pop(ref_key, None) + else: + self._session_refcount[ref_key] = remaining - # Remove the mapping for this context - self._context_to_session.pop(context_id, None) + # CAS: only drop the context->session mapping if it still points + # at the session we just cleaned up. The get() and pop() below run + # synchronously with no `await` between them, so no other coroutine + # can interleave and re-point the mapping after our check. + if self._context_to_session.get(context_id) == (server_key, session_id): + self._context_to_session.pop(context_id, None) class MCPStdioClient: