mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
refactor: extract stale session handling into _strip_stale_mcp_session_header helper
This commit is contained in:
parent
51eb0e9d98
commit
a9a8bb71c4
2 changed files with 128 additions and 25 deletions
|
|
@ -1840,6 +1840,43 @@ if MCP_AVAILABLE:
|
|||
raw_headers,
|
||||
)
|
||||
|
||||
def _strip_stale_mcp_session_header(
|
||||
scope: Scope,
|
||||
mgr: "StreamableHTTPSessionManager",
|
||||
) -> None:
|
||||
"""
|
||||
Strip stale ``mcp-session-id`` headers so the session manager
|
||||
creates a fresh session instead of returning 404 "Session not found".
|
||||
|
||||
When clients like VSCode reconnect after a reload they may resend a
|
||||
session id that has already been cleaned up. Rather than letting the
|
||||
SDK return a 404 error loop, we detect the stale id and remove the
|
||||
header so a brand-new session is created transparently.
|
||||
|
||||
Fixes https://github.com/BerriAI/litellm/issues/20292
|
||||
"""
|
||||
_mcp_session_header = b"mcp-session-id"
|
||||
_session_id: Optional[str] = None
|
||||
for header_name, header_value in scope.get("headers", []):
|
||||
if header_name == _mcp_session_header:
|
||||
_session_id = header_value.decode("utf-8", errors="replace")
|
||||
break
|
||||
|
||||
if _session_id is None:
|
||||
return
|
||||
|
||||
known_sessions = getattr(mgr, "_server_instances", None)
|
||||
if known_sessions is not None and _session_id not in known_sessions:
|
||||
verbose_logger.warning(
|
||||
"MCP session ID '%s' not found in active sessions. "
|
||||
"Stripping stale header to force new session creation.",
|
||||
_session_id,
|
||||
)
|
||||
scope["headers"] = [
|
||||
(k, v) for k, v in scope["headers"]
|
||||
if k != _mcp_session_header
|
||||
]
|
||||
|
||||
async def handle_streamable_http_mcp(
|
||||
scope: Scope, receive: Receive, send: Send
|
||||
) -> None:
|
||||
|
|
@ -1896,31 +1933,7 @@ if MCP_AVAILABLE:
|
|||
# Give it a moment to start up
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
# Handle stale mcp-session-id headers (Fixes #20292)
|
||||
# When clients like VSCode reconnect after a reload, they may send a
|
||||
# stale mcp-session-id that no longer exists in the session manager.
|
||||
# This causes a 404 "Session not found" error loop. Strip the header
|
||||
# so the session manager creates a fresh session instead.
|
||||
_mcp_session_header = b"mcp-session-id"
|
||||
_stale_session_id: Optional[str] = None
|
||||
for header_name, header_value in scope.get("headers", []):
|
||||
if header_name == _mcp_session_header:
|
||||
_stale_session_id = header_value.decode("utf-8", errors="replace")
|
||||
break
|
||||
|
||||
if _stale_session_id is not None:
|
||||
# Check if this session ID exists in the session manager
|
||||
_known_sessions = getattr(session_manager, "_server_instances", None)
|
||||
if _known_sessions is not None and _stale_session_id not in _known_sessions:
|
||||
verbose_logger.warning(
|
||||
"MCP session ID '%s' not found in active sessions. "
|
||||
"Stripping stale mcp-session-id header to force new session creation.",
|
||||
_stale_session_id,
|
||||
)
|
||||
scope["headers"] = [
|
||||
(k, v) for k, v in scope["headers"]
|
||||
if k != _mcp_session_header
|
||||
]
|
||||
_strip_stale_mcp_session_header(scope, session_manager)
|
||||
|
||||
await session_manager.handle_request(scope, receive, send)
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -11,6 +11,96 @@ import pytest
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
||||
class TestStripStaleMcpSessionHeader:
|
||||
"""Unit tests for the _strip_stale_mcp_session_header helper."""
|
||||
|
||||
def test_strips_stale_session_id(self):
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_strip_stale_mcp_session_header,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
scope = {
|
||||
"headers": [
|
||||
(b"content-type", b"application/json"),
|
||||
(b"mcp-session-id", b"stale-id"),
|
||||
],
|
||||
}
|
||||
mgr = MagicMock()
|
||||
mgr._server_instances = {} # no active sessions
|
||||
|
||||
_strip_stale_mcp_session_header(scope, mgr)
|
||||
|
||||
header_names = [k for k, _ in scope["headers"]]
|
||||
assert b"mcp-session-id" not in header_names
|
||||
|
||||
def test_preserves_valid_session_id(self):
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_strip_stale_mcp_session_header,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
scope = {
|
||||
"headers": [
|
||||
(b"content-type", b"application/json"),
|
||||
(b"mcp-session-id", b"valid-id"),
|
||||
],
|
||||
}
|
||||
mgr = MagicMock()
|
||||
mgr._server_instances = {"valid-id": MagicMock()}
|
||||
|
||||
_strip_stale_mcp_session_header(scope, mgr)
|
||||
|
||||
header_names = [k for k, _ in scope["headers"]]
|
||||
assert b"mcp-session-id" in header_names
|
||||
|
||||
def test_no_op_when_no_session_header(self):
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_strip_stale_mcp_session_header,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
scope = {
|
||||
"headers": [
|
||||
(b"content-type", b"application/json"),
|
||||
],
|
||||
}
|
||||
mgr = MagicMock()
|
||||
mgr._server_instances = {}
|
||||
|
||||
_strip_stale_mcp_session_header(scope, mgr)
|
||||
|
||||
assert len(scope["headers"]) == 1
|
||||
|
||||
def test_no_op_when_server_instances_missing(self):
|
||||
"""If _server_instances attr doesn't exist, don't crash."""
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_strip_stale_mcp_session_header,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
scope = {
|
||||
"headers": [
|
||||
(b"mcp-session-id", b"some-id"),
|
||||
],
|
||||
}
|
||||
mgr = MagicMock(spec=[]) # no attributes
|
||||
|
||||
_strip_stale_mcp_session_header(scope, mgr)
|
||||
|
||||
# Should keep the header since we can't verify
|
||||
header_names = [k for k, _ in scope["headers"]]
|
||||
assert b"mcp-session-id" in header_names
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stale_mcp_session_id_is_stripped():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue