mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Fix MCP initialize session active tracking
Co-authored-by: Yassin Kortam <yassin@berri.ai>
This commit is contained in:
parent
dd4af19c49
commit
eaf99edfe3
2 changed files with 125 additions and 16 deletions
|
|
@ -3148,10 +3148,29 @@ if MCP_AVAILABLE:
|
|||
session_id, asyncio.Lock()
|
||||
)
|
||||
|
||||
track_active_stateful_request = bool(use_stateful and session_id)
|
||||
if track_active_stateful_request and session_id:
|
||||
_stateful_session_active_request_counts[session_id] = (
|
||||
_stateful_session_active_request_counts.get(session_id, 0) + 1
|
||||
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
|
||||
)
|
||||
+ 1
|
||||
)
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
async def _dispatch() -> None:
|
||||
|
|
@ -3174,6 +3193,7 @@ if MCP_AVAILABLE:
|
|||
_owner_fingerprint_for(
|
||||
user_api_key_auth, oauth2_headers, _client_ip
|
||||
),
|
||||
_track_initialized_stateful_session,
|
||||
)
|
||||
|
||||
async with _gateway_initialize_instructions_request_scope(
|
||||
|
|
@ -3192,32 +3212,38 @@ if MCP_AVAILABLE:
|
|||
else:
|
||||
await _dispatch()
|
||||
finally:
|
||||
if track_active_stateful_request and session_id:
|
||||
if active_request_session_id:
|
||||
active_request_count = (
|
||||
_stateful_session_active_request_counts.get(session_id, 0) - 1
|
||||
_stateful_session_active_request_counts.get(
|
||||
active_request_session_id, 0
|
||||
)
|
||||
- 1
|
||||
)
|
||||
if active_request_count > 0:
|
||||
_stateful_session_active_request_counts[session_id] = (
|
||||
active_request_count
|
||||
)
|
||||
_stateful_session_active_request_counts[
|
||||
active_request_session_id
|
||||
] = active_request_count
|
||||
else:
|
||||
_stateful_session_active_request_counts.pop(session_id, None)
|
||||
_stateful_session_active_request_counts.pop(
|
||||
active_request_session_id, None
|
||||
)
|
||||
|
||||
if (
|
||||
scope.get("method") != "DELETE"
|
||||
and session_id in _stateful_session_auth_contexts
|
||||
and active_request_session_id in _stateful_session_auth_contexts
|
||||
):
|
||||
_stateful_session_auth_context_last_seen[session_id] = (
|
||||
time.monotonic()
|
||||
)
|
||||
_stateful_session_auth_context_last_seen[
|
||||
active_request_session_id
|
||||
] = time.monotonic()
|
||||
|
||||
# Periodic cleanup iterates _stateful_session_auth_context_last_seen,
|
||||
# so locks for untracked sessions must be dropped here.
|
||||
if (
|
||||
active_request_count <= 0
|
||||
and session_id not in _stateful_session_auth_contexts
|
||||
and active_request_session_id
|
||||
not in _stateful_session_auth_contexts
|
||||
):
|
||||
_stateful_session_locks.pop(session_id, None)
|
||||
_stateful_session_locks.pop(active_request_session_id, None)
|
||||
except HTTPException:
|
||||
# Re-raise HTTP exceptions to preserve status codes and details
|
||||
raise
|
||||
|
|
@ -3422,6 +3448,7 @@ if MCP_AVAILABLE:
|
|||
send: Send,
|
||||
auth_user: MCPAuthenticatedUser,
|
||||
owner_fingerprint: str,
|
||||
on_session_registered: Optional[Callable[[str], None]] = None,
|
||||
) -> Send:
|
||||
async def wrapped_send(message: Message) -> None:
|
||||
if message.get("type") == "http.response.start":
|
||||
|
|
@ -3431,6 +3458,8 @@ if MCP_AVAILABLE:
|
|||
session_id = (
|
||||
value.decode() if isinstance(value, bytes) else str(value)
|
||||
)
|
||||
if on_session_registered is not None:
|
||||
on_session_registered(session_id)
|
||||
auth_context_var.set(auth_user)
|
||||
_stateful_session_auth_contexts[session_id] = auth_user
|
||||
_stateful_session_auth_context_last_seen[session_id] = (
|
||||
|
|
|
|||
|
|
@ -1559,6 +1559,86 @@ async def test_initialize_response_capture_accepts_str_headers_and_sets_auth_con
|
|||
mcp_server._remove_stateful_session_tracking(session_id)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_request_tracks_active_session_after_response_header():
|
||||
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")
|
||||
|
||||
session_id = "initialize-active-session-1"
|
||||
owner_auth = UserAPIKeyAuth(api_key="initialize-key", user_id="user-a")
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/mcp",
|
||||
"headers": [
|
||||
(b"content-type", b"application/json"),
|
||||
(b"authorization", b"Bearer initialize-key"),
|
||||
],
|
||||
}
|
||||
receive = AsyncMock(
|
||||
return_value={
|
||||
"type": "http.request",
|
||||
"body": b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}',
|
||||
"more_body": False,
|
||||
}
|
||||
)
|
||||
|
||||
async def stateful_handle(s, r, se):
|
||||
await se(
|
||||
{
|
||||
"type": "http.response.start",
|
||||
"headers": [(b"mcp-session-id", session_id.encode())],
|
||||
}
|
||||
)
|
||||
assert mcp_server._stateful_session_active_request_counts[session_id] == 1
|
||||
now = (
|
||||
mcp_server._stateful_session_auth_context_last_seen[session_id]
|
||||
+ mcp_server._STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS
|
||||
)
|
||||
await mcp_server._purge_expired_stateful_session_auth_contexts(now=now)
|
||||
assert session_id in mcp_server._stateful_session_auth_contexts
|
||||
|
||||
async def stateless_handle(s, r, se):
|
||||
raise AssertionError("initialize request should use stateful manager")
|
||||
|
||||
try:
|
||||
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", {}),
|
||||
):
|
||||
await handle_streamable_http_mcp(scope, receive, AsyncMock())
|
||||
|
||||
assert session_id not in mcp_server._stateful_session_active_request_counts
|
||||
assert session_id in mcp_server._stateful_session_auth_contexts
|
||||
finally:
|
||||
mcp_server._remove_stateful_session_tracking(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