mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(mcp): bound client terminated session retention
This commit is contained in:
parent
e7b15b8bb6
commit
9906303fb4
2 changed files with 115 additions and 19 deletions
|
|
@ -106,6 +106,7 @@ _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS: Final = 30 * 60
|
|||
# the cap is still hit (every session in flight), the new `initialize` is
|
||||
# rejected with 429.
|
||||
_MAX_STATEFUL_SESSIONS_PER_OWNER: Final = 100
|
||||
_MAX_CLIENT_TERMINATED_SESSIONS: Final = 10000
|
||||
# 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
|
||||
|
|
@ -629,6 +630,7 @@ if MCP_AVAILABLE:
|
|||
_stateful_session_active_request_counts: Final[dict[str, int]] = {}
|
||||
_stateful_session_client_info: Final[dict[str, Implementation]] = {} # mutable-ok: cleared on session teardown
|
||||
_terminated_session_ids: Final[dict[str, float]] = {} # mutable-ok: explicitly closed id -> last replay
|
||||
_client_terminated_session_owners: Final[dict[str, str]] = {}
|
||||
|
||||
class _TerminableTransport(Protocol):
|
||||
async def terminate(self) -> None: ...
|
||||
|
|
@ -1322,8 +1324,10 @@ if MCP_AVAILABLE:
|
|||
session_id
|
||||
for session_id, last_replayed in _terminated_session_ids.items()
|
||||
if now - last_replayed >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS
|
||||
and not _stateful_session_active_request_counts.get(session_id)
|
||||
]:
|
||||
del _terminated_session_ids[session_id]
|
||||
_client_terminated_session_owners.pop(session_id, None)
|
||||
|
||||
def _is_terminated_session_id(session_id: str, now: float) -> bool:
|
||||
last_replayed: Final = _terminated_session_ids.get(session_id)
|
||||
|
|
@ -1331,6 +1335,7 @@ if MCP_AVAILABLE:
|
|||
return False
|
||||
if now - last_replayed >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS:
|
||||
del _terminated_session_ids[session_id]
|
||||
_client_terminated_session_owners.pop(session_id, None)
|
||||
return False
|
||||
_terminated_session_ids[session_id] = now
|
||||
return True
|
||||
|
|
@ -2231,6 +2236,10 @@ if MCP_AVAILABLE:
|
|||
_increment_active_request_session(initialized_session_id)
|
||||
|
||||
async def _dispatch() -> None:
|
||||
# Another DELETE may have completed while this request waited for the lock.
|
||||
if use_stateful and session_id and scope.get("method") == "DELETE":
|
||||
if await _handle_stale_mcp_session(scope, receive, send, target_manager):
|
||||
return
|
||||
_otel_publish_transport_span_on_scope(scope)
|
||||
_otel_publish_request_destinations_on_scope(scope)
|
||||
auth_user: Final = _set_or_update_auth_context(
|
||||
|
|
@ -2255,22 +2264,45 @@ if MCP_AVAILABLE:
|
|||
client_info=_extract_initialize_client_info(body),
|
||||
)
|
||||
|
||||
async with _gateway_initialize_instructions_request_scope(
|
||||
user_api_key_auth,
|
||||
mcp_servers,
|
||||
_client_ip,
|
||||
scoped_server_endpoint=scoped_server_endpoint,
|
||||
is_initialize=is_initialize,
|
||||
):
|
||||
await target_manager.handle_request(scope, receive, local_send)
|
||||
deleting_session: Final = session_id if use_stateful and scope.get("method") == "DELETE" else None
|
||||
if deleting_session:
|
||||
_forget_expired_terminated_session_ids(time.monotonic())
|
||||
delete_owner: Final = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip)
|
||||
if (
|
||||
use_stateful
|
||||
and session_id
|
||||
and scope.get("method") == "DELETE"
|
||||
and session_id not in _stateful_server_instances()
|
||||
len(_client_terminated_session_owners) >= _MAX_CLIENT_TERMINATED_SESSIONS
|
||||
or sum(owner == delete_owner for owner in _client_terminated_session_owners.values())
|
||||
>= _MAX_STATEFUL_SESSIONS_PER_OWNER
|
||||
):
|
||||
_terminated_session_ids[session_id] = time.monotonic()
|
||||
_remove_stateful_session_tracking(session_id)
|
||||
delete_capacity_response: Final = JSONResponse(
|
||||
status_code=429,
|
||||
content={
|
||||
"error": "Too Many Requests",
|
||||
"details": "Too many recently terminated MCP sessions. Retry after idle records expire.",
|
||||
},
|
||||
)
|
||||
await delete_capacity_response(scope, receive, local_send)
|
||||
return
|
||||
# Reserve before yielding so concurrent DELETEs cannot exceed the cap.
|
||||
# Never evict fresh records: that would allow terminated IDs to replay.
|
||||
_client_terminated_session_owners[deleting_session] = delete_owner
|
||||
_terminated_session_ids[deleting_session] = time.monotonic()
|
||||
try:
|
||||
async with _gateway_initialize_instructions_request_scope(
|
||||
user_api_key_auth,
|
||||
mcp_servers,
|
||||
_client_ip,
|
||||
scoped_server_endpoint=scoped_server_endpoint,
|
||||
is_initialize=is_initialize,
|
||||
):
|
||||
await target_manager.handle_request(scope, receive, local_send)
|
||||
finally:
|
||||
if deleting_session:
|
||||
if deleting_session in _stateful_server_instances():
|
||||
_client_terminated_session_owners.pop(deleting_session, None)
|
||||
_terminated_session_ids.pop(deleting_session, None)
|
||||
else:
|
||||
_terminated_session_ids[deleting_session] = time.monotonic()
|
||||
_remove_stateful_session_tracking(deleting_session)
|
||||
|
||||
try:
|
||||
if session_lock is not None:
|
||||
|
|
|
|||
|
|
@ -370,7 +370,7 @@ async def test_delete_stale_mcp_session_returns_success():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("outcome", ("deleted", "rejected", "exception"))
|
||||
@pytest.mark.parametrize("outcome", ("deleted", "rejected", "exception", "cancelled", "owner-cap", "worker-cap", "expired", "other-owner", "queued-delete"))
|
||||
async def test_delete_preserves_tracking_until_the_session_is_terminated(outcome):
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
|
|
@ -387,6 +387,8 @@ async def test_delete_preserves_tracking_until_the_session_is_terminated(outcome
|
|||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
|
||||
session_id = "delete-failure-session"
|
||||
from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import MCPAuthenticatedUser
|
||||
|
||||
|
|
@ -415,16 +417,28 @@ async def test_delete_preserves_tracking_until_the_session_is_terminated(outcome
|
|||
_stateful_session_owners[session_id] = _owner_fingerprint_for(user_auth)
|
||||
_stateful_session_locks[session_id] = session_lock
|
||||
|
||||
queued_deletes = []
|
||||
queued_send = AsyncMock()
|
||||
|
||||
async def dispatch(scope, receive, send):
|
||||
if outcome == "queued-delete":
|
||||
queued_deletes.append(asyncio.create_task(handle_streamable_http_mcp(dict(scope), receive, queued_send)))
|
||||
await asyncio.sleep(0)
|
||||
if outcome == "cancelled":
|
||||
raise asyncio.CancelledError
|
||||
if outcome == "exception":
|
||||
raise RuntimeError("delete failed")
|
||||
if outcome == "deleted":
|
||||
if outcome in ("deleted", "expired", "other-owner", "queued-delete"):
|
||||
mock_instances.pop(session_id)
|
||||
await send({"type": "http.response.start", "status": 200 if outcome == "deleted" else 403, "headers": []})
|
||||
await send({"type": "http.response.start", "status": 200 if outcome in ("deleted", "expired", "other-owner", "queued-delete") else 403, "headers": []})
|
||||
await send({"type": "http.response.body", "body": b""})
|
||||
|
||||
try:
|
||||
with (
|
||||
patch.object(server, "_terminated_session_ids", {"retained": 0.0 if outcome == "expired" else server.time.monotonic()}),
|
||||
patch.object(server, "_client_terminated_session_owners", {"retained": _owner_fingerprint_for(user_auth) if outcome == "owner-cap" else "other-owner"}, create=True),
|
||||
patch.object(server, "_MAX_STATEFUL_SESSIONS_PER_OWNER", 1 if outcome in ("owner-cap", "expired", "other-owner") else 100),
|
||||
patch.object(server, "_MAX_CLIENT_TERMINATED_SESSIONS", 1 if outcome == "worker-cap" else 10000, create=True),
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
|
||||
new_callable=AsyncMock,
|
||||
|
|
@ -446,8 +460,24 @@ async def test_delete_preserves_tracking_until_the_session_is_terminated(outcome
|
|||
mock_instances,
|
||||
),
|
||||
):
|
||||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
if outcome == "deleted":
|
||||
if outcome == "cancelled":
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
else:
|
||||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
if queued_deletes:
|
||||
await asyncio.gather(*queued_deletes)
|
||||
assert queued_send.await_args_list[0].args[0]["status"] == 200
|
||||
assert mock_handle_request.await_count == 1
|
||||
if outcome in ("owner-cap", "worker-cap"):
|
||||
assert send.await_args_list[0].args[0]["status"] == 429
|
||||
assert mock_handle_request.await_count == 0
|
||||
assert session_id in mock_instances
|
||||
assert session_id not in server._terminated_session_ids
|
||||
assert _stateful_session_owners[session_id] == _owner_fingerprint_for(user_auth)
|
||||
return
|
||||
if outcome in ("deleted", "expired", "other-owner", "queued-delete"):
|
||||
assert server._client_terminated_session_owners[session_id] == _owner_fingerprint_for(user_auth)
|
||||
assert session_id not in _stateful_session_auth_contexts
|
||||
assert session_id not in _stateful_session_owners
|
||||
assert session_id not in _stateful_session_locks
|
||||
|
|
@ -462,6 +492,9 @@ async def test_delete_preserves_tracking_until_the_session_is_terminated(outcome
|
|||
assert repeated_delete.await_args_list[0].args[0]["status"] == 200
|
||||
return
|
||||
|
||||
assert session_id not in server._client_terminated_session_owners
|
||||
assert session_id not in server._terminated_session_ids
|
||||
|
||||
assert mock_handle_request.await_count == 1
|
||||
assert _stateful_session_auth_contexts[session_id] is auth_context
|
||||
assert _stateful_session_auth_context_last_seen[session_id] == 1.0
|
||||
|
|
@ -476,6 +509,37 @@ async def test_delete_preserves_tracking_until_the_session_is_terminated(outcome
|
|||
_stateful_session_locks.pop(session_id, None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("expiry_path", ("cleanup", "replay"))
|
||||
def test_expired_client_termination_releases_capacity(expiry_path):
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
|
||||
with (
|
||||
patch.object(server, "_terminated_session_ids", {"expired": 0.0}),
|
||||
patch.object(server, "_client_terminated_session_owners", {"expired": "owner"}),
|
||||
patch.object(server, "_stateful_session_active_request_counts", {}),
|
||||
):
|
||||
now = server._STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS
|
||||
if expiry_path == "cleanup":
|
||||
server._forget_expired_terminated_session_ids(now)
|
||||
else:
|
||||
assert server._is_terminated_session_id("expired", now) is False
|
||||
assert server._terminated_session_ids == {}
|
||||
assert server._client_terminated_session_owners == {}
|
||||
|
||||
|
||||
def test_inflight_delete_keeps_its_reserved_capacity_until_dispatch_finishes():
|
||||
from litellm.proxy._experimental.mcp_server import server
|
||||
|
||||
with (
|
||||
patch.object(server, "_terminated_session_ids", {"deleting": 0.0}),
|
||||
patch.object(server, "_client_terminated_session_owners", {"deleting": "owner"}),
|
||||
patch.object(server, "_stateful_session_active_request_counts", {"deleting": 1}),
|
||||
):
|
||||
server._forget_expired_terminated_session_ids(server._STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS)
|
||||
assert server._terminated_session_ids == {"deleting": 0.0}
|
||||
assert server._client_terminated_session_owners == {"deleting": "owner"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_valid_mcp_session_id_is_preserved():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue