mirror of
https://github.com/langflow-ai/langflow.git
synced 2026-07-25 16:09:31 +08:00
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:
@ -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."""
|
||||
|
||||
@ -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:
|
||||
|
||||
Reference in New Issue
Block a user