mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Fix stateful MCP initialize auth capture
This commit is contained in:
parent
f71c9f716c
commit
ac6986bd94
2 changed files with 54 additions and 2 deletions
|
|
@ -3399,8 +3399,12 @@ if MCP_AVAILABLE:
|
|||
async def wrapped_send(message: Message) -> None:
|
||||
if message.get("type") == "http.response.start":
|
||||
for key, value in message.get("headers", []):
|
||||
if key.lower() == b"mcp-session-id":
|
||||
session_id = value.decode()
|
||||
header_name = key if isinstance(key, bytes) else str(key).encode()
|
||||
if header_name.lower() == b"mcp-session-id":
|
||||
session_id = (
|
||||
value.decode() if isinstance(value, bytes) else str(value)
|
||||
)
|
||||
auth_context_var.set(auth_user)
|
||||
_stateful_session_auth_contexts[session_id] = auth_user
|
||||
_stateful_session_auth_context_last_seen[session_id] = (
|
||||
time.monotonic()
|
||||
|
|
|
|||
|
|
@ -1408,6 +1408,54 @@ async def test_stateful_mcp_requests_refresh_session_auth_context():
|
|||
mcp_server._stateful_session_auth_contexts.pop(session_id, None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_initialize_response_capture_accepts_str_headers_and_sets_auth_context():
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server import server as mcp_server
|
||||
except ImportError:
|
||||
pytest.skip("MCP server not available")
|
||||
|
||||
session_id = "initialize-session-1"
|
||||
auth_user = mcp_server.MCPAuthenticatedUser(
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="initialize-key", user_id="user-a")
|
||||
)
|
||||
previous_auth_user = mcp_server.MCPAuthenticatedUser(
|
||||
user_api_key_auth=UserAPIKeyAuth(api_key="previous-key", user_id="user-b")
|
||||
)
|
||||
sent_messages = []
|
||||
|
||||
async def send(message):
|
||||
sent_messages.append(message)
|
||||
|
||||
wrapped_send = mcp_server._wrap_send_with_stateful_session_auth_context(
|
||||
send,
|
||||
auth_user,
|
||||
"owner-fingerprint",
|
||||
)
|
||||
token = mcp_server.auth_context_var.set(previous_auth_user)
|
||||
try:
|
||||
await wrapped_send(
|
||||
{
|
||||
"type": "http.response.start",
|
||||
"headers": [("mcp-session-id", session_id)],
|
||||
}
|
||||
)
|
||||
|
||||
assert mcp_server.auth_context_var.get() is auth_user
|
||||
assert mcp_server._stateful_session_auth_contexts[session_id] is auth_user
|
||||
assert mcp_server._stateful_session_owners[session_id] == "owner-fingerprint"
|
||||
assert session_id in mcp_server._stateful_session_auth_context_last_seen
|
||||
assert sent_messages == [
|
||||
{
|
||||
"type": "http.response.start",
|
||||
"headers": [("mcp-session-id", session_id)],
|
||||
}
|
||||
]
|
||||
finally:
|
||||
mcp_server.auth_context_var.reset(token)
|
||||
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