diff --git a/litellm/proxy/_experimental/mcp_server/sampling_handler.py b/litellm/proxy/_experimental/mcp_server/sampling_handler.py index 590268077ea..814cdad7e22 100644 --- a/litellm/proxy/_experimental/mcp_server/sampling_handler.py +++ b/litellm/proxy/_experimental/mcp_server/sampling_handler.py @@ -472,20 +472,26 @@ async def handle_sampling_create_message( # 6. Inject auth context for cost tracking if user_api_key_auth: - completion_kwargs["user"] = getattr(user_api_key_auth, "user_id", None) + from fastapi import Request - # Pass user_api_key_dict directly so proxy hooks can attribute the cost - # litellm_pre_call_utils usually checks for this in kwargs or metadata - if "metadata" not in completion_kwargs: - completion_kwargs["metadata"] = {} + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.proxy_server import proxy_config - api_key = getattr(user_api_key_auth, "api_key", None) - if api_key: - completion_kwargs["metadata"]["user_api_key"] = api_key - - team_id = getattr(user_api_key_auth, "team_id", None) - if team_id: - completion_kwargs["metadata"]["user_api_key_team_id"] = team_id + # We need a dummy FastAPI request object because add_litellm_data_to_request expects it + _dummy_request = Request( + scope={ + "type": "http", + "method": "POST", + "path": "/mcp/sampling/createMessage", + "headers": [(b"content-type", b"application/json")], + } + ) + completion_kwargs = await add_litellm_data_to_request( + data=completion_kwargs, + request=_dummy_request, + user_api_key_dict=user_api_key_auth, + proxy_config=proxy_config, + ) verbose_logger.debug( "MCP sampling: calling litellm.acompletion with model=%s, num_messages=%d, has_tools=%s", diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index c3f78bb7fe2..8f32bb86a8b 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -461,17 +461,21 @@ async def test_sse_mcp_handler_mock(): mock_sse = MagicMock() mock_sse.connect_sse = MagicMock() - + # Mock connect_sse to return an async context manager yielding dummy streams mock_context_manager = AsyncMock() mock_context_manager.__aenter__.return_value = (AsyncMock(), AsyncMock()) mock_sse.connect_sse.return_value = mock_context_manager - + mock_server = MagicMock() mock_server.run = AsyncMock() mock_server.create_initialization_options = MagicMock() with ( + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), patch( "litellm.proxy._experimental.mcp_server.server.sse", mock_sse, @@ -493,7 +497,9 @@ async def test_sse_mcp_handler_mock(): ) await handle_sse_mcp_endpoint(mock_scope, mock_receive, mock_send) - mock_sse.connect_sse.assert_called_once_with(mock_scope, mock_receive, mock_send) + mock_sse.connect_sse.assert_called_once_with( + mock_scope, mock_receive, mock_send + ) mock_server.run.assert_called_once() @@ -525,7 +531,7 @@ async def test_sse_post_messages_auth_failure(): with pytest.raises(HTTPException) as exc_info: await handle_sse_post_messages(mock_scope, mock_receive, mock_send) - + assert exc_info.value.status_code == 401 assert exc_info.value.detail == "Unauthorized"