mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(usage-ai-chat): resolve proxy model aliases before acompletion
This commit is contained in:
parent
934ecdca78
commit
1959d42546
2 changed files with 83 additions and 1 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue