diff --git a/litellm/llms/groq/chat/transformation.py b/litellm/llms/groq/chat/transformation.py index 089c0cac62c..80417d93090 100644 --- a/litellm/llms/groq/chat/transformation.py +++ b/litellm/llms/groq/chat/transformation.py @@ -39,6 +39,8 @@ from litellm.types.utils import ModelResponse, ModelResponseStream from ...openai_like.chat.transformation import OpenAILikeChatConfig +GROQ_COMPOUND_MODELS = frozenset({"compound", "compound-mini"}) + class GroqChatConfig(OpenAILikeChatConfig): frequency_penalty: Optional[int] = None @@ -84,6 +86,26 @@ class GroqChatConfig(OpenAILikeChatConfig): def get_config(cls): return super().get_config() + @staticmethod + def _get_groq_model_name(model: str) -> str: + return f"groq/{model}" if model in GROQ_COMPOUND_MODELS else model + + def transform_request( + self, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + headers: dict, + ) -> dict: + return super().transform_request( + model=self._get_groq_model_name(model), + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=headers, + ) + def get_model_response_iterator( self, streaming_response: Union[Iterator[str], AsyncIterator[str], ModelResponse], diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 70b6b05e6ec..6818a4baba3 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -24104,6 +24104,28 @@ "supports_response_schema": false, "supports_tool_choice": true }, + "groq/compound": { + "input_cost_per_token": 0.0, + "litellm_provider": "groq", + "max_input_tokens": 131072, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 0.0, + "supports_tool_choice": true, + "supports_web_search": true + }, + "groq/compound-mini": { + "input_cost_per_token": 0.0, + "litellm_provider": "groq", + "max_input_tokens": 131072, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 0.0, + "supports_tool_choice": true, + "supports_web_search": true + }, "groq/whisper-large-v3": { "input_cost_per_second": 3.083e-05, "litellm_provider": "groq", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b961c326625..dc28a54f734 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -24262,6 +24262,28 @@ "supports_response_schema": false, "supports_tool_choice": true }, + "groq/compound": { + "input_cost_per_token": 0.0, + "litellm_provider": "groq", + "max_input_tokens": 131072, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 0.0, + "supports_tool_choice": true, + "supports_web_search": true + }, + "groq/compound-mini": { + "input_cost_per_token": 0.0, + "litellm_provider": "groq", + "max_input_tokens": 131072, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 0.0, + "supports_tool_choice": true, + "supports_web_search": true + }, "groq/whisper-large-v3": { "input_cost_per_second": 3.083e-05, "litellm_provider": "groq", diff --git a/tests/test_litellm/llms/groq/chat/test_transformation.py b/tests/test_litellm/llms/groq/chat/test_transformation.py new file mode 100644 index 00000000000..fa0335488fa --- /dev/null +++ b/tests/test_litellm/llms/groq/chat/test_transformation.py @@ -0,0 +1,92 @@ +import json +from pathlib import Path +from unittest.mock import patch + +import httpx +import pytest + +import litellm +from litellm.llms.custom_httpx.http_handler import HTTPHandler +from litellm.llms.groq.chat.transformation import GroqChatConfig + + +@pytest.mark.parametrize( + "model, expected", + [ + ("compound", "groq/compound"), + ("compound-mini", "groq/compound-mini"), + ("llama-3.3-70b-versatile", "llama-3.3-70b-versatile"), + ("openai/gpt-oss-120b", "openai/gpt-oss-120b"), + ], +) +def test_get_groq_model_name(model, expected): + assert GroqChatConfig._get_groq_model_name(model) == expected + + +@pytest.mark.parametrize("model", ["groq/compound", "groq/compound-mini"]) +def test_compound_models_in_backup_cost_map(model): + """ + Regression test for https://github.com/BerriAI/litellm/issues/32467 + + The Groq compound systems must be registered so they route through the groq + provider and expose metadata. + """ + json_path = Path(__file__).parents[5] / "litellm" / "model_prices_and_context_window_backup.json" + with open(json_path) as f: + model_cost = json.load(f) + + info = model_cost.get(model) + assert info is not None, f"{model} missing from backup JSON" + assert info["litellm_provider"] == "groq" + assert info["mode"] == "chat" + assert info["max_input_tokens"] == 131072 + assert info["max_output_tokens"] == 8192 + + +@pytest.mark.parametrize( + "requested_model, expected_api_model", + [ + ("groq/compound", "groq/compound"), + ("groq/compound-mini", "groq/compound-mini"), + ("groq/llama-3.3-70b-versatile", "llama-3.3-70b-versatile"), + ], +) +def test_compound_model_name_sent_to_groq(requested_model, expected_api_model): + """ + Regression test for https://github.com/BerriAI/litellm/issues/32467 + + Groq exposes the compound systems with the "groq/" prefix as part of the + actual model id, so LiteLLM must send "groq/compound(-mini)" instead of the + prefix-stripped "compound(-mini)", which Groq rejects with model_not_found. + """ + client = HTTPHandler() + fake_response = httpx.Response( + status_code=200, + json={ + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 1, + "model": expected_api_model, + "service_tier": "auto", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "hi"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + }, + request=httpx.Request("POST", "https://api.groq.com/openai/v1/chat/completions"), + ) + + with patch.object(client, "post", return_value=fake_response) as mock_post: + litellm.completion( + model=requested_model, + messages=[{"role": "user", "content": "hi"}], + api_key="fake-key", + client=client, + ) + + sent_body = json.loads(mock_post.call_args.kwargs["data"]) + assert sent_body["model"] == expected_api_model