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 a50ce1d3c48..bd3d12b9789 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py @@ -5,7 +5,19 @@ usage/spend data by querying the aggregated daily activity endpoints. import json from datetime import date -from typing import Any, AsyncIterator, Callable, Dict, List, Literal, Optional, cast +from typing import ( + TYPE_CHECKING, + Any, + AsyncIterator, + Awaitable, + Callable, + Dict, + List, + Literal, + Optional, + Union, + cast, +) from typing_extensions import TypedDict @@ -16,6 +28,10 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import ( SpendAnalyticsPaginatedResponse, ) +if TYPE_CHECKING: + from litellm.router import Router + from litellm.utils import CustomStreamWrapper, ModelResponse + # --------------------------------------------------------------------------- # Constants # --------------------------------------------------------------------------- @@ -518,13 +534,37 @@ async def _process_tool_call( ) +def _is_router_model(llm_router: "Router", model: str) -> bool: + """True when ``model`` is a name the proxy router can resolve: a configured + model group, a model_group_alias, or a concrete deployment id.""" + return ( + model in llm_router.model_names + or model in llm_router.model_group_alias + or llm_router.has_model_id(model) + ) + + +def _completion_fn( + llm_router: Optional["Router"], model: str +) -> Callable[..., Awaitable[Union["ModelResponse", "CustomStreamWrapper"]]]: + """Pick the completion entrypoint for ``model``: the proxy router when it + can resolve the name (a configured model group, a model_group_alias, or a + deployment id) so Usage AI chat accepts the same names as the UI dropdown, + otherwise a direct litellm call.""" + if llm_router is not None and _is_router_model(llm_router, model): + return llm_router.acompletion + return litellm.acompletion + + async def _stream_final_response( - model: str, chat_messages: List[Dict[str, Any]] + model: str, + chat_messages: List[Dict[str, Any]], + llm_router: Optional["Router"], ) -> 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 _completion_fn(llm_router, model)( model=model, messages=chat_messages, stream=True, @@ -541,6 +581,7 @@ async def stream_usage_ai_chat( model: Optional[str] = None, user_id: Optional[str] = None, is_admin: bool = False, + llm_router: Optional["Router"] = None, ) -> AsyncIterator[str]: """Stream SSE events: status → tool_call → chunk → done.""" resolved_model = (model or "").strip() or DEFAULT_COMPETITOR_DISCOVERY_MODEL @@ -555,7 +596,7 @@ 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 _completion_fn(llm_router, resolved_model)( model=resolved_model, messages=chat_messages, tools=tools, @@ -573,7 +614,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, llm_router + ): 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 d0df80fed0d..8d706c0310f 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py @@ -49,6 +49,7 @@ async def usage_ai_chat( from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import ( stream_usage_ai_chat, ) + from litellm.proxy.proxy_server import llm_router is_admin = _user_has_admin_view(user_api_key_dict) if is_admin: @@ -63,6 +64,7 @@ async def usage_ai_chat( model=data.model, user_id=user_id, is_admin=is_admin, + llm_router=llm_router, ), 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 e8a74e41dae..3896fa90893 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 @@ -12,12 +12,24 @@ from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import ( TOOLS_ADMIN, TOOLS_BASE, _build_system_prompt, + _is_router_model, _summarise_entity_data, _summarise_usage_data, stream_usage_ai_chat, ) +def _make_fake_router(model_names, *, alias=None, model_ids=()): + """A stand-in proxy router exposing only the attributes the model-resolution + logic touches, so tests can inject it without a real Router.""" + router = MagicMock() + router.model_names = set(model_names) + router.model_group_alias = dict(alias or {}) + router.has_model_id = lambda candidate: candidate in set(model_ids) + router.acompletion = AsyncMock() + return router + + SAMPLE_AGGREGATED_RESPONSE = { "results": [ { @@ -466,3 +478,136 @@ class TestUsageAiChatServiceAccountGuard: is_admin=False, ) assert "Endpoint-level guard missing" in str(exc_info.value) + + +class TestIsRouterModel: + """The predicate that decides whether a selected model can be resolved by + the proxy router rather than treated as a raw provider/model string.""" + + def test_configured_model_group(self): + router = _make_fake_router({"glm-5"}) + assert _is_router_model(router, "glm-5") is True + + def test_model_group_alias(self): + router = _make_fake_router(set(), alias={"my-alias": "openai/gpt-4o-mini"}) + assert _is_router_model(router, "my-alias") is True + + def test_concrete_deployment_id(self): + router = _make_fake_router(set(), model_ids={"deployment-123"}) + assert _is_router_model(router, "deployment-123") is True + + def test_unknown_model_is_not_router_resolvable(self): + router = _make_fake_router({"glm-5"}) + assert _is_router_model(router, "openai/gpt-4o-mini") is False + + +def _tool_call_then_stream(): + """Build (first_response_with_tool_call, streaming_response) mock pair that + drives stream_usage_ai_chat through both of its completion calls.""" + tool_call = MagicMock() + tool_call.id = "call_1" + tool_call.function.name = "get_usage_data" + tool_call.function.arguments = json.dumps( + {"start_date": "2025-01-01", "end_date": "2025-01-31"} + ) + + first_response = MagicMock() + first_response.choices = [MagicMock()] + first_response.choices[0].message.tool_calls = [tool_call] + first_response.choices[0].message.model_dump.return_value = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": { + "name": "get_usage_data", + "arguments": '{"start_date":"2025-01-01","end_date":"2025-01-31"}', + }, + } + ], + } + + async def stream(): + chunk = MagicMock() + chunk.choices = [MagicMock()] + chunk.choices[0].delta.content = "Total spend is $50.25" + yield chunk + + return first_response, stream() + + +class TestModelResolutionRouting: + """Regression for GitHub #24513: a selected proxy alias / model group must + resolve through the router instead of being passed to litellm.acompletion + as a raw provider/model string.""" + + @pytest.mark.asyncio + async def test_proxy_model_group_routes_through_router(self): + first_response, stream = _tool_call_then_stream() + fake_router = _make_fake_router({"mylitellmmodel"}) + fake_router.acompletion = AsyncMock(side_effect=[first_response, stream]) + + with ( + 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, + return_value=SAMPLE_AGGREGATED_RESPONSE, + ), + ): + mock_litellm.acompletion = AsyncMock() + + events = [] + async for event in stream_usage_ai_chat( + messages=[{"role": "user", "content": "What is my total spend?"}], + model="mylitellmmodel", + user_id="user-1", + is_admin=True, + llm_router=fake_router, + ): + events.append(json.loads(event.replace("data: ", "").strip())) + + mock_litellm.acompletion.assert_not_called() + assert fake_router.acompletion.await_count == 2 + assert ( + fake_router.acompletion.await_args_list[0].kwargs["model"] + == "mylitellmmodel" + ) + assert ( + fake_router.acompletion.await_args_list[1].kwargs["model"] + == "mylitellmmodel" + ) + assert any(e["type"] == "done" for e in events) + + @pytest.mark.asyncio + async def test_unconfigured_model_falls_back_to_litellm(self): + first_response, stream = _tool_call_then_stream() + fake_router = _make_fake_router({"some-other-group"}) + + with ( + 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, + return_value=SAMPLE_AGGREGATED_RESPONSE, + ), + ): + mock_litellm.acompletion = AsyncMock(side_effect=[first_response, stream]) + + async for _ in stream_usage_ai_chat( + messages=[{"role": "user", "content": "spend?"}], + model="openai/gpt-4o-mini", + user_id="user-1", + is_admin=True, + llm_router=fake_router, + ): + pass + + fake_router.acompletion.assert_not_called() + assert mock_litellm.acompletion.await_count == 2