fix: serialize concurrent MCP session access to prevent race conditions (#12761)

* 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_<hash>_0'
  Timeout updating tool list: ...
  Error updating tool list: <connection error>

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_<hash>_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.
This commit is contained in:
Eric Hare
2026-04-23 11:25:12 -07:00
committed by GitHub
parent ac654b2e71
commit 0217f704df
2 changed files with 617 additions and 122 deletions

View File

@ -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_<hash>_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."""

View File

@ -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: