mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Fix MCP reinitialize session tracking
Co-authored-by: Yassin Kortam <yassin@berri.ai>
This commit is contained in:
parent
225c11c2d9
commit
897aaa5dc1
2 changed files with 120 additions and 22 deletions
|
|
@ -72,8 +72,8 @@ _byok_cred_cache: Dict[Tuple[str, str], Tuple[Optional[str], float]] = {}
|
|||
_BYOK_CRED_CACHE_TTL = 60 # seconds
|
||||
_BYOK_CRED_CACHE_MAX_SIZE = 4096 # cap to prevent unbounded growth
|
||||
_STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS = 30 * 60
|
||||
# Maximum bytes to peek when sniffing the JSON-RPC method on a no-session-id
|
||||
# POST. An `initialize` envelope is a few hundred bytes; capping the peek
|
||||
# Maximum bytes to peek when sniffing the JSON-RPC method on a POST.
|
||||
# An `initialize` envelope is a few hundred bytes; capping the peek
|
||||
# prevents an authenticated client from forcing the proxy to buffer an
|
||||
# arbitrarily large body just to make a routing decision.
|
||||
_MCP_ROUTING_PEEK_MAX_BYTES = 4096
|
||||
|
|
@ -3083,7 +3083,7 @@ if MCP_AVAILABLE:
|
|||
return
|
||||
session_id = _get_session_id_from_scope(scope)
|
||||
|
||||
if scope.get("method") == "POST" and not session_id:
|
||||
if scope.get("method") == "POST":
|
||||
consumed_messages, body = await _read_request_body_for_routing(receive)
|
||||
is_initialize = _is_initialize_request(body)
|
||||
|
||||
|
|
@ -3148,30 +3148,24 @@ if MCP_AVAILABLE:
|
|||
session_id, asyncio.Lock()
|
||||
)
|
||||
|
||||
active_request_session_id = (
|
||||
session_id if use_stateful and session_id else None
|
||||
)
|
||||
if active_request_session_id:
|
||||
_stateful_session_active_request_counts[active_request_session_id] = (
|
||||
_stateful_session_active_request_counts.get(
|
||||
active_request_session_id, 0
|
||||
)
|
||||
active_request_session_ids: List[str] = []
|
||||
|
||||
def _increment_active_request_session(session_id_to_track: str) -> None:
|
||||
if session_id_to_track in active_request_session_ids:
|
||||
return
|
||||
active_request_session_ids.append(session_id_to_track)
|
||||
_stateful_session_active_request_counts[session_id_to_track] = (
|
||||
_stateful_session_active_request_counts.get(session_id_to_track, 0)
|
||||
+ 1
|
||||
)
|
||||
|
||||
if use_stateful and session_id:
|
||||
_increment_active_request_session(session_id)
|
||||
|
||||
def _track_initialized_stateful_session(
|
||||
initialized_session_id: str,
|
||||
) -> None:
|
||||
nonlocal active_request_session_id
|
||||
if active_request_session_id is not None:
|
||||
return
|
||||
active_request_session_id = initialized_session_id
|
||||
_stateful_session_active_request_counts[initialized_session_id] = (
|
||||
_stateful_session_active_request_counts.get(
|
||||
initialized_session_id, 0
|
||||
)
|
||||
+ 1
|
||||
)
|
||||
_increment_active_request_session(initialized_session_id)
|
||||
|
||||
async def _dispatch() -> None:
|
||||
auth_user = _set_or_update_auth_context(
|
||||
|
|
@ -3212,7 +3206,7 @@ if MCP_AVAILABLE:
|
|||
else:
|
||||
await _dispatch()
|
||||
finally:
|
||||
if active_request_session_id:
|
||||
for active_request_session_id in active_request_session_ids:
|
||||
active_request_count = (
|
||||
_stateful_session_active_request_counts.get(
|
||||
active_request_session_id, 0
|
||||
|
|
|
|||
|
|
@ -1639,6 +1639,110 @@ async def test_initialize_request_tracks_active_session_after_response_header():
|
|||
mcp_server._remove_stateful_session_tracking(session_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_request_with_existing_session_tracks_new_session():
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_server
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
handle_streamable_http_mcp,
|
||||
session_manager_stateful,
|
||||
session_manager_stateless,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
existing_session_id = "existing-initialize-session"
|
||||
new_session_id = "reinitialized-session"
|
||||
owner_auth = UserAPIKeyAuth(api_key="initialize-key", user_id="user-a")
|
||||
owner_fingerprint = mcp_server._owner_fingerprint_for(owner_auth)
|
||||
initialize_body = b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}'
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp",
|
||||
"headers": [
|
||||
(b"content-type", b"application/json"),
|
||||
(b"authorization", b"Bearer initialize-key"),
|
||||
(b"mcp-session-id", existing_session_id.encode()),
|
||||
],
|
||||
}
|
||||
receive = AsyncMock(
|
||||
return_value={
|
||||
"type": "http.request",
|
||||
"body": initialize_body,
|
||||
"more_body": False,
|
||||
}
|
||||
)
|
||||
stateful_called = []
|
||||
|
||||
async def stateful_handle(s, r, se):
|
||||
stateful_called.append(1)
|
||||
message = await r()
|
||||
assert message["body"] == initialize_body
|
||||
await se(
|
||||
{
|
||||
"type": "http.response.start",
|
||||
"headers": [(b"mcp-session-id", new_session_id.encode())],
|
||||
}
|
||||
)
|
||||
assert mcp_server._stateful_session_auth_contexts[new_session_id]
|
||||
assert mcp_server._stateful_session_owners[new_session_id] == owner_fingerprint
|
||||
assert mcp_server._stateful_session_active_request_counts[new_session_id] == 1
|
||||
now = (
|
||||
mcp_server._stateful_session_auth_context_last_seen[new_session_id]
|
||||
+ mcp_server._STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS
|
||||
)
|
||||
await mcp_server._purge_expired_stateful_session_auth_contexts(now=now)
|
||||
assert new_session_id in mcp_server._stateful_session_auth_contexts
|
||||
|
||||
async def stateless_handle(s, r, se):
|
||||
raise AssertionError(
|
||||
"initialize request with session should use stateful manager"
|
||||
)
|
||||
|
||||
try:
|
||||
mcp_server._stateful_session_auth_contexts[existing_session_id] = (
|
||||
mcp_server.MCPAuthenticatedUser(user_api_key_auth=owner_auth)
|
||||
)
|
||||
mcp_server._stateful_session_auth_context_last_seen[existing_session_id] = 1.0
|
||||
mcp_server._stateful_session_owners[existing_session_id] = owner_fingerprint
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
return_value=(owner_auth, None, None, None, None, None),
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
|
||||
True,
|
||||
),
|
||||
patch.object(
|
||||
session_manager_stateful,
|
||||
"handle_request",
|
||||
side_effect=stateful_handle,
|
||||
),
|
||||
patch.object(
|
||||
session_manager_stateless,
|
||||
"handle_request",
|
||||
side_effect=stateless_handle,
|
||||
),
|
||||
patch.object(
|
||||
session_manager_stateful,
|
||||
"_server_instances",
|
||||
{existing_session_id: MagicMock()},
|
||||
),
|
||||
):
|
||||
await handle_streamable_http_mcp(scope, receive, AsyncMock())
|
||||
|
||||
assert stateful_called
|
||||
assert new_session_id not in mcp_server._stateful_session_active_request_counts
|
||||
assert new_session_id in mcp_server._stateful_session_auth_contexts
|
||||
finally:
|
||||
mcp_server._remove_stateful_session_tracking(existing_session_id)
|
||||
mcp_server._remove_stateful_session_tracking(new_session_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stateful_mcp_auth_contexts_expire_with_idle_sessions():
|
||||
"""Expired session auth contexts should not remain in memory indefinitely."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue