mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
resolve
This commit is contained in:
parent
75c0ba6cfe
commit
c76676373c
2 changed files with 28 additions and 16 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue