diff --git a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py index 4de29e04092..f4ba44fb7b0 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py @@ -509,15 +509,56 @@ async def _process_tool_call( ) +def _get_proxy_router(): + """Return the global proxy router when available.""" + try: + from litellm.proxy.proxy_server import llm_router + except Exception: + return None + return llm_router + + +async def _usage_ai_acompletion( + *, + model: str, + messages: List[Dict[str, Any]], + team_id: Optional[str] = None, + **kwargs: Any, +): + """ + Route Usage AI requests through the proxy router when available so proxy + aliases and model groups resolve consistently with normal proxy traffic. + """ + router = _get_proxy_router() + request_kwargs: Dict[str, Any] = { + "model": model, + "messages": messages, + **kwargs, + } + + if team_id: + metadata = dict(cast(Dict[str, Any], request_kwargs.get("metadata") or {})) + metadata["user_api_key_team_id"] = team_id + request_kwargs["metadata"] = metadata + + if router is not None: + return await router.acompletion(**request_kwargs) + + return await litellm.acompletion(**request_kwargs) + + async def _stream_final_response( - model: str, chat_messages: List[Dict[str, Any]] + model: str, + chat_messages: List[Dict[str, Any]], + team_id: Optional[str] = None, ) -> AsyncIterator[str]: """Stream the final LLM response after tool results are appended.""" yield _sse({"type": "status", "message": "Analyzing results..."}) - response = await litellm.acompletion( + response = await _usage_ai_acompletion( model=model, messages=chat_messages, + team_id=team_id, stream=True, temperature=USAGE_AI_TEMPERATURE, ) @@ -532,6 +573,7 @@ async def stream_usage_ai_chat( model: Optional[str] = None, user_id: Optional[str] = None, is_admin: bool = False, + team_id: Optional[str] = None, ) -> AsyncIterator[str]: """Stream SSE events: status → tool_call → chunk → done.""" resolved_model = (model or "").strip() or DEFAULT_COMPETITOR_DISCOVERY_MODEL @@ -546,9 +588,10 @@ async def stream_usage_ai_chat( try: yield _sse({"type": "status", "message": "Thinking..."}) tools = get_tools_for_role(is_admin) - response = await litellm.acompletion( + response = await _usage_ai_acompletion( model=resolved_model, messages=chat_messages, + team_id=team_id, tools=tools, temperature=USAGE_AI_TEMPERATURE, ) @@ -564,7 +607,9 @@ async def stream_usage_ai_chat( for tc in choice.message.tool_calls: async for event in _process_tool_call(tc, chat_messages, user_id, is_admin): yield event - async for event in _stream_final_response(resolved_model, chat_messages): + async for event in _stream_final_response( + resolved_model, chat_messages, team_id + ): yield event yield _sse({"type": "done"}) diff --git a/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py index 0dbe518afb7..2e91c52b9df 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py @@ -59,6 +59,7 @@ async def usage_ai_chat( model=data.model, user_id=user_id, is_admin=is_admin, + team_id=user_api_key_dict.team_id, ), media_type="text/event-stream", headers={"Cache-Control": "no-cache", "X-Accel-Buffering": "no"}, diff --git a/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py index f9303bd13a6..769e0209895 100644 --- a/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py +++ b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py @@ -177,11 +177,13 @@ class TestSummariseEntityData: class TestStreamUsageAiChat: - @pytest.mark.asyncio - async def test_stream_emits_status_events(self): + @staticmethod + def _make_tool_call_response( + tool_name: str = "get_usage_data", + ) -> tuple[MagicMock, MagicMock]: mock_tool_call = MagicMock() mock_tool_call.id = "call_123" - mock_tool_call.function.name = "get_usage_data" + mock_tool_call.function.name = tool_name mock_tool_call.function.arguments = json.dumps( { "start_date": "2025-01-01", @@ -200,18 +202,24 @@ class TestStreamUsageAiChat: "id": "call_123", "type": "function", "function": { - "name": "get_usage_data", + "name": tool_name, "arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31"}', }, } ], } + return mock_tool_call, mock_first_response - async def mock_stream(): - chunk = MagicMock() - chunk.choices = [MagicMock()] - chunk.choices[0].delta.content = "Total spend is $50.25" - yield chunk + @staticmethod + async def _mock_stream(content: str = "Total spend is $50.25"): + chunk = MagicMock() + chunk.choices = [MagicMock()] + chunk.choices[0].delta.content = content + yield chunk + + @pytest.mark.asyncio + async def test_stream_emits_status_events(self): + _, mock_first_response = self._make_tool_call_response() with patch( "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm" @@ -222,7 +230,7 @@ class TestStreamUsageAiChat: mock_litellm.acompletion = AsyncMock( side_effect=[ mock_first_response, - mock_stream(), + self._mock_stream(), ] ) mock_fetch.return_value = SAMPLE_AGGREGATED_RESPONSE @@ -251,39 +259,13 @@ class TestStreamUsageAiChat: @pytest.mark.asyncio async def test_stream_handles_team_tool(self): - mock_tool_call = MagicMock() - mock_tool_call.id = "call_team" - mock_tool_call.function.name = "get_team_usage_data" - mock_tool_call.function.arguments = json.dumps( - { - "start_date": "2025-01-01", - "end_date": "2025-01-31", - } + _, mock_first_response = self._make_tool_call_response( + tool_name="get_team_usage_data" ) - - mock_first_response = MagicMock() - mock_first_response.choices = [MagicMock()] - mock_first_response.choices[0].message.tool_calls = [mock_tool_call] - mock_first_response.choices[0].message.model_dump.return_value = { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_team", - "type": "function", - "function": { - "name": "get_team_usage_data", - "arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31"}', - }, - } - ], - } - - async def mock_stream(): - chunk = MagicMock() - chunk.choices = [MagicMock()] - chunk.choices[0].delta.content = "Engineering is the top team." - yield chunk + mock_first_response.choices[0].message.tool_calls[0].id = "call_team" + mock_first_response.choices[0].message.model_dump.return_value["tool_calls"][0][ + "id" + ] = "call_team" with patch( "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm" @@ -294,7 +276,7 @@ class TestStreamUsageAiChat: mock_litellm.acompletion = AsyncMock( side_effect=[ mock_first_response, - mock_stream(), + self._mock_stream(content="Engineering is the top team."), ] ) mock_fetch.return_value = SAMPLE_TEAM_RESPONSE @@ -330,41 +312,26 @@ class TestStreamUsageAiChat: @pytest.mark.asyncio async def test_non_admin_enforces_user_id(self): - mock_tool_call = MagicMock() - mock_tool_call.id = "call_456" - mock_tool_call.function.name = "get_usage_data" - mock_tool_call.function.arguments = json.dumps( + _, mock_first_response = self._make_tool_call_response() + mock_first_response.choices[0].message.tool_calls[0].id = "call_456" + mock_first_response.choices[0].message.tool_calls[ + 0 + ].function.arguments = json.dumps( { "start_date": "2025-01-01", "end_date": "2025-01-31", "user_id": "other-user", } ) - - mock_first_response = MagicMock() - mock_first_response.choices = [MagicMock()] - mock_first_response.choices[0].message.tool_calls = [mock_tool_call] - mock_first_response.choices[0].message.model_dump.return_value = { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call_456", - "type": "function", - "function": { - "name": "get_usage_data", - "arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31","user_id":"other-user"}', - }, - } - ], + mock_first_response.choices[0].message.model_dump.return_value["tool_calls"][0] = { + "id": "call_456", + "type": "function", + "function": { + "name": "get_usage_data", + "arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31","user_id":"other-user"}', + }, } - async def mock_stream(): - chunk = MagicMock() - chunk.choices = [MagicMock()] - chunk.choices[0].delta.content = "Data." - yield chunk - mock_fetch = AsyncMock(return_value=SAMPLE_AGGREGATED_RESPONSE) with patch( @@ -382,7 +349,7 @@ class TestStreamUsageAiChat: mock_litellm.acompletion = AsyncMock( side_effect=[ mock_first_response, - mock_stream(), + self._mock_stream(content="Data."), ] ) @@ -400,3 +367,79 @@ class TestStreamUsageAiChat: end_date="2025-01-31", user_id="my-user-id", ) + + @pytest.mark.asyncio + async def test_stream_uses_router_acompletion_for_proxy_alias_models(self): + _, mock_first_response = self._make_tool_call_response() + mock_router = MagicMock() + mock_router.acompletion = AsyncMock( + side_effect=[ + mock_first_response, + self._mock_stream(), + ] + ) + + with patch( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat._get_proxy_router", + return_value=mock_router, + ), patch( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm" + ) as mock_litellm, patch( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat._fetch_usage_data", + new_callable=AsyncMock, + ) as mock_fetch: + mock_fetch.return_value = SAMPLE_AGGREGATED_RESPONSE + + events = [] + async for event in stream_usage_ai_chat( + messages=[{"role": "user", "content": "What is my total spend?"}], + model="mylitellmmodel", + team_id="team-123", + is_admin=True, + ): + events.append(json.loads(event.replace("data: ", "").strip())) + + assert any(event["type"] == "done" for event in events) + assert mock_router.acompletion.await_count == 2 + mock_litellm.acompletion.assert_not_called() + + first_call = mock_router.acompletion.await_args_list[0].kwargs + second_call = mock_router.acompletion.await_args_list[1].kwargs + + assert first_call["model"] == "mylitellmmodel" + assert first_call["metadata"]["user_api_key_team_id"] == "team-123" + assert second_call["metadata"]["user_api_key_team_id"] == "team-123" + assert second_call["stream"] is True + + @pytest.mark.asyncio + async def test_stream_falls_back_to_litellm_when_router_missing(self): + _, mock_first_response = self._make_tool_call_response() + + with patch( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat._get_proxy_router", + return_value=None, + ), patch( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm" + ) as mock_litellm, patch( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat._fetch_usage_data", + new_callable=AsyncMock, + ) as mock_fetch: + mock_litellm.acompletion = AsyncMock( + side_effect=[ + mock_first_response, + self._mock_stream(), + ] + ) + mock_fetch.return_value = SAMPLE_AGGREGATED_RESPONSE + + events = [] + async for event in stream_usage_ai_chat( + messages=[{"role": "user", "content": "What is my total spend?"}], + model="gpt-4o-mini", + team_id="team-123", + is_admin=True, + ): + events.append(json.loads(event.replace("data: ", "").strip())) + + assert any(event["type"] == "done" for event in events) + assert mock_litellm.acompletion.await_count == 2 diff --git a/ui/litellm-dashboard/src/components/UsagePage/components/UsageAIChatPanel.test.tsx b/ui/litellm-dashboard/src/components/UsagePage/components/UsageAIChatPanel.test.tsx index 85ce11e605a..827d0facae7 100644 --- a/ui/litellm-dashboard/src/components/UsagePage/components/UsageAIChatPanel.test.tsx +++ b/ui/litellm-dashboard/src/components/UsagePage/components/UsageAIChatPanel.test.tsx @@ -1,6 +1,7 @@ -import { screen } from "@testing-library/react"; -import { beforeAll, describe, expect, it, vi } from "vitest"; +import { act, fireEvent, screen, waitFor } from "@testing-library/react"; +import { beforeAll, beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders } from "../../../../tests/test-utils"; +import { modelHubCall, usageAiChatStream } from "../../networking"; import UsageAIChatPanel from "./UsageAIChatPanel"; beforeAll(() => { @@ -20,7 +21,7 @@ vi.mock("../../networking", () => ({ { model_group: "claude-3-opus" }, ], }), - usageAiChatStream: vi.fn(), + usageAiChatStream: vi.fn().mockResolvedValue(undefined), })); const defaultProps = { @@ -30,6 +31,10 @@ const defaultProps = { }; describe("UsageAIChatPanel", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + it("should render the panel when open", () => { renderWithProviders(); @@ -82,4 +87,53 @@ describe("UsageAIChatPanel", () => { expect(screen.getByTestId("usage-ai-chat-panel")).not.toHaveClass("translate-x-full"); expect(screen.getByTestId("usage-ai-chat-panel")).toHaveClass("translate-x-0"); }); + + it("should submit the selected model value unchanged", async () => { + renderWithProviders(); + + await waitFor(() => { + expect(modelHubCall).toHaveBeenCalledWith("test-token"); + }); + + const modelSelect = screen.getByRole("combobox"); + await act(async () => { + fireEvent.mouseDown(modelSelect); + }); + + await waitFor(() => { + expect( + screen.getByRole("option", { name: "claude-3-opus" }) + ).toBeInTheDocument(); + }); + + await act(async () => { + fireEvent.click(screen.getByRole("option", { name: "claude-3-opus" })); + }); + + await act(async () => { + fireEvent.change(screen.getByPlaceholderText("Ask about your usage..."), { + target: { value: "hello" }, + }); + }); + + await act(async () => { + fireEvent.click(screen.getByRole("button", { name: "Send" })); + }); + + await waitFor(() => { + expect(usageAiChatStream).toHaveBeenCalled(); + }); + + expect(usageAiChatStream).toHaveBeenCalledWith( + "test-token", + [{ role: "user", content: "hello" }], + "claude-3-opus", + expect.any(Function), + expect.any(Function), + expect.any(Function), + expect.any(Function), + expect.any(Function), + expect.any(AbortSignal), + ); + }); });