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..af79b29428d 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py @@ -426,6 +426,29 @@ def _sse(event: SSEEvent) -> str: return f"data: {json.dumps(event)}\n\n" +def _resolve_usage_chat_model_alias(model: str) -> str: + """Resolve model aliases configured on the proxy before calling acompletion.""" + resolved_model = model.strip() + if not resolved_model: + return resolved_model + + if litellm.model_alias_map and resolved_model in litellm.model_alias_map: + return litellm.model_alias_map[resolved_model] + + try: + from litellm.proxy.proxy_server import llm_router + + if ( + llm_router is not None + and resolved_model in getattr(llm_router, "model_group_alias", {}) + ): + return llm_router._get_model_from_alias(resolved_model) + except Exception: + pass + + return resolved_model + + def _resolve_fetch_kwargs( fn_name: str, fn_args: Dict[str, str], @@ -534,7 +557,8 @@ async def stream_usage_ai_chat( is_admin: bool = False, ) -> AsyncIterator[str]: """Stream SSE events: status → tool_call → chunk → done.""" - resolved_model = (model or "").strip() or DEFAULT_COMPETITOR_DISCOVERY_MODEL + selected_model = (model or "").strip() or DEFAULT_COMPETITOR_DISCOVERY_MODEL + resolved_model = _resolve_usage_chat_model_alias(selected_model) truncated = ( messages[-MAX_CHAT_MESSAGES:] if len(messages) > MAX_CHAT_MESSAGES else messages ) 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 66a18e2edb4..767ecf37839 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,6 +177,64 @@ class TestSummariseEntityData: class TestStreamUsageAiChat: + @pytest.mark.asyncio + async def test_stream_resolves_global_model_alias_before_acompletion(self): + mock_first_response = MagicMock() + mock_first_response.choices = [MagicMock()] + mock_first_response.choices[0].message.tool_calls = [] + mock_first_response.choices[0].message.content = "ok" + + with patch( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm" + ) as mock_litellm: + mock_litellm.model_alias_map = {"usage-alias": "gpt-4o-mini"} + mock_litellm.acompletion = AsyncMock(return_value=mock_first_response) + + events = [] + async for event in stream_usage_ai_chat( + messages=[{"role": "user", "content": "hi"}], + model="usage-alias", + is_admin=True, + ): + events.append(event) + + first_call_kwargs = mock_litellm.acompletion.call_args_list[0].kwargs + assert first_call_kwargs["model"] == "gpt-4o-mini" + + @pytest.mark.asyncio + async def test_stream_resolves_router_model_group_alias_before_acompletion(self): + mock_first_response = MagicMock() + mock_first_response.choices = [MagicMock()] + mock_first_response.choices[0].message.tool_calls = [] + mock_first_response.choices[0].message.content = "ok" + + mock_router = MagicMock() + mock_router.model_group_alias = {"friendly-alias": "hidden-route"} + mock_router._get_model_from_alias.return_value = "anthropic/claude-3-5-sonnet" + + with ( + patch( + "litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat.litellm" + ) as mock_litellm, + patch.dict( + "sys.modules", + {"litellm.proxy.proxy_server": MagicMock(llm_router=mock_router)}, + ), + ): + mock_litellm.model_alias_map = {} + mock_litellm.acompletion = AsyncMock(return_value=mock_first_response) + + events = [] + async for event in stream_usage_ai_chat( + messages=[{"role": "user", "content": "hi"}], + model="friendly-alias", + is_admin=True, + ): + events.append(event) + + first_call_kwargs = mock_litellm.acompletion.call_args_list[0].kwargs + assert first_call_kwargs["model"] == "anthropic/claude-3-5-sonnet" + @pytest.mark.asyncio async def test_stream_emits_status_events(self): mock_tool_call = MagicMock()