This commit is contained in:
Yug 2026-04-30 11:10:22 +05:30
parent 75c0ba6cfe
commit c76676373c
2 changed files with 28 additions and 16 deletions

View file

@ -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",

View file

@ -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"