fix(mcp): synchronize session pool singleton lifecycle (#3797)

* fix(mcp): synchronize session pool singleton lifecycle

Parity follow-up to #3778 (skill storage) and #3730 (sandbox provider).
get_session_pool() already serialised creation with _pool_lock, but its
fast-path check and final return read the global separately, and
reset_session_pool() cleared it with no lock. A reset_mcp_tools_cache()
(reachable via the /api/mcp/cache/reset admin endpoint) racing a
concurrent get could null the singleton between the None-check and the
return, handing the caller None despite the -> MCPSessionPool annotation.

Build and return the pool inside _pool_lock with a double check, and
clear it under the same lock in reset_session_pool(). The critical
section is tiny and never awaits, so holding the threading.Lock is safe
from both the async and sync/worker-thread paths. No behavior change for
single-threaded callers.

Signed-off-by: Yufeng He <40085740+he-yufeng@users.noreply.github.com>

* Apply suggestions from code review

Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com>

---------

Signed-off-by: Yufeng He <40085740+he-yufeng@users.noreply.github.com>
Co-authored-by: Willem Jiang <willem.jiang@gmail.com>
Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com>
This commit is contained in:
Yufeng He 2026-07-07 23:09:06 +08:00 committed by GitHub
parent 658c39ccf7
commit a4a88fda19
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 111 additions and 7 deletions

View File

@ -442,14 +442,19 @@ _pool_lock = threading.Lock()
def get_session_pool() -> MCPSessionPool:
"""Return the global session-pool singleton."""
global _pool
if _pool is None:
with _pool_lock:
if _pool is None:
_pool = MCPSessionPool()
return _pool
# Build and return under the lock so racing cold-start callers construct
# exactly one pool and reset_session_pool() can't null the global between
# reading it and returning it (which previously could hand back None). The
# critical section is tiny and never awaits, so a threading.Lock is safe to
# hold from both the async and sync/worker-thread paths.
with _pool_lock:
if _pool is None:
_pool = MCPSessionPool()
return _pool
def reset_session_pool() -> None:
"""Reset the singleton (for tests)."""
"""Reset the singleton (used in tests and the MCP cache reset path)."""
global _pool
_pool = None
with _pool_lock:
_pool = None

View File

@ -0,0 +1,99 @@
"""Concurrency regression tests for the MCP session-pool singleton lifecycle.
These guard the module-level ``get_session_pool`` / ``reset_session_pool``
singleton in ``deerflow.mcp.session_pool``. ``reset_session_pool`` is reachable
in production through the ``/api/mcp/cache/reset`` admin endpoint
(``reset_mcp_tools_cache`` closes the pool so it is rebuilt on the next tool
load), and the harness runs the main event loop alongside channel threads on
their own loops, so a reset can race a concurrent ``get_session_pool``. Before
the lock was extended to cover the return, ``get_session_pool`` re-read the
global after its fast-path ``None`` check, so a ``reset_session_pool`` landing in
that window handed the caller ``None`` despite the ``-> MCPSessionPool``
annotation.
This mirrors ``test_skill_storage_lifecycle.py`` the sibling singleton fixed
the same way in #3778 — adapted to the session pool, whose ``MCPSessionPool`` is
cheap to construct and already serialises creation, so the gap that mattered
here was the reset racing the get's return.
"""
import sys
import threading
from deerflow.mcp.session_pool import (
MCPSessionPool,
get_session_pool,
reset_session_pool,
)
def test_get_session_pool_returns_one_singleton_under_concurrent_cold_start():
"""Threads racing a cold start all observe the same single instance."""
reset_session_pool()
n_threads = 8
pools: list[MCPSessionPool] = []
pools_lock = threading.Lock()
# Barrier makes all threads enter get_session_pool() together, so the
# cold-start race is triggered rather than left to chance.
barrier = threading.Barrier(n_threads)
def get_pool() -> None:
barrier.wait()
pool = get_session_pool()
with pools_lock:
pools.append(pool)
threads = [threading.Thread(target=get_pool) for _ in range(n_threads)]
for thread in threads:
thread.start()
for thread in threads:
thread.join()
try:
assert len(pools) == n_threads
assert len({id(pool) for pool in pools}) == 1
finally:
reset_session_pool()
def test_reset_racing_get_never_returns_none():
"""A reset racing concurrent gets must never hand back ``None``.
Getters and a resetter run in tight loops while the interpreter is forced to
switch threads very often, so the reset repeatedly lands while a getter is
between its fast-path ``None`` check and its return the interleaving that
the unlocked check-then-return path turned into a ``None`` return. Without
the lock covering the return this reliably observes ``None``; with it, never.
"""
reset_session_pool()
none_seen: list[int] = []
none_seen_lock = threading.Lock()
stop = threading.Event()
def getter() -> None:
while not stop.is_set():
if get_session_pool() is None:
with none_seen_lock:
none_seen.append(1)
def resetter() -> None:
for _ in range(100000):
reset_session_pool()
stop.set()
previous_interval = sys.getswitchinterval()
sys.setswitchinterval(1e-6)
try:
getters = [threading.Thread(target=getter) for _ in range(4)]
reset_thread = threading.Thread(target=resetter)
for thread in getters:
thread.start()
reset_thread.start()
for thread in getters:
thread.join()
reset_thread.join()
finally:
sys.setswitchinterval(previous_interval)
reset_session_pool()
assert not none_seen, "get_session_pool() returned None while a reset raced it"