diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 4cc7446fd32..41673c7283d 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -3178,6 +3178,7 @@ if MCP_AVAILABLE: client_ip=_client_ip, session_id=session_id if use_stateful else None, touch_last_seen=(scope.get("method") or "").upper() != "DELETE", + copy_existing_session_auth_context=is_initialize, ) local_send = send if use_stateful and is_initialize: @@ -3409,6 +3410,7 @@ if MCP_AVAILABLE: client_ip: Optional[str] = None, session_id: Optional[str] = None, touch_last_seen: bool = True, + copy_existing_session_auth_context: bool = False, ) -> MCPAuthenticatedUser: auth_user = ( _stateful_session_auth_contexts.get(session_id) if session_id else None @@ -3416,6 +3418,16 @@ if MCP_AVAILABLE: if auth_user is not None and session_id is not None: if touch_last_seen: _stateful_session_auth_context_last_seen[session_id] = time.monotonic() + if copy_existing_session_auth_context: + return set_auth_context( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) _update_auth_context( auth_user=auth_user, user_api_key_auth=user_api_key_auth, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index cf6fdb04a23..f8fab3aff67 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -1655,6 +1655,15 @@ async def test_initialize_request_with_existing_session_tracks_new_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) + existing_auth_user = mcp_server.MCPAuthenticatedUser( + user_api_key_auth=owner_auth, + mcp_auth_header="old-mcp-auth", + mcp_servers=["old-server"], + mcp_server_auth_headers={"old-server": {"Authorization": "Bearer old-key"}}, + oauth2_headers={"Authorization": "Bearer old-oauth"}, + raw_headers={"x-old-header": "old"}, + client_ip="old-client-ip", + ) initialize_body = b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}' scope = { "type": "http", @@ -1694,6 +1703,14 @@ async def test_initialize_request_with_existing_session_tracks_new_session(): ) await mcp_server._purge_expired_stateful_session_auth_contexts(now=now) assert new_session_id in mcp_server._stateful_session_auth_contexts + assert ( + mcp_server._stateful_session_auth_contexts[new_session_id] + is not existing_auth_user + ) + assert ( + mcp_server._stateful_session_auth_contexts[new_session_id].mcp_auth_header + == "new-mcp-auth" + ) async def stateless_handle(s, r, se): raise AssertionError( @@ -1702,7 +1719,7 @@ async def test_initialize_request_with_existing_session_tracks_new_session(): try: mcp_server._stateful_session_auth_contexts[existing_session_id] = ( - mcp_server.MCPAuthenticatedUser(user_api_key_auth=owner_auth) + existing_auth_user ) mcp_server._stateful_session_auth_context_last_seen[existing_session_id] = 1.0 mcp_server._stateful_session_owners[existing_session_id] = owner_fingerprint @@ -1711,7 +1728,14 @@ async def test_initialize_request_with_existing_session_tracks_new_session(): patch( "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", new_callable=AsyncMock, - return_value=(owner_auth, None, None, None, None, None), + return_value=( + owner_auth, + "new-mcp-auth", + ["new-server"], + {"new-server": {"Authorization": "Bearer new-key"}}, + {"Authorization": "Bearer new-oauth"}, + {"x-new-header": "new"}, + ), ), patch( "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", @@ -1738,6 +1762,12 @@ async def test_initialize_request_with_existing_session_tracks_new_session(): 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 + assert ( + mcp_server._stateful_session_auth_contexts[existing_session_id] + is existing_auth_user + ) + assert existing_auth_user.mcp_auth_header == "old-mcp-auth" + assert existing_auth_user.mcp_servers == ["old-server"] finally: mcp_server._remove_stateful_session_tracking(existing_session_id) mcp_server._remove_stateful_session_tracking(new_session_id)