diff --git a/litellm/utils.py b/litellm/utils.py index 8d55783bf9c..4bb85c46249 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1033,17 +1033,47 @@ def function_setup( # noqa: PLR0915 call_type == CallTypes.image_generation.value or call_type == CallTypes.aimage_generation.value ): - messages = args[0] if len(args) > 0 else kwargs["prompt"] + if len(args) > 0: + messages = args[0] + else: + prompt_value = kwargs.get("prompt") + if prompt_value is None: + raise BadRequestError( + message="Missing required parameter: 'prompt' for image generation.", + model=kwargs.get("model") or "", + llm_provider="", + ) + messages = prompt_value elif ( call_type == CallTypes.moderation.value or call_type == CallTypes.amoderation.value ): - messages = args[1] if len(args) > 1 else kwargs["input"] + if len(args) > 1: + messages = args[1] + else: + input_value = kwargs.get("input") + if input_value is None: + raise BadRequestError( + message="Missing required parameter: 'input' for moderation.", + model=kwargs.get("model") or "", + llm_provider="", + ) + messages = input_value elif ( call_type == CallTypes.atext_completion.value or call_type == CallTypes.text_completion.value ): - messages = args[0] if len(args) > 0 else kwargs["prompt"] + if len(args) > 0: + messages = args[0] + else: + prompt_value = kwargs.get("prompt") + if prompt_value is None: + raise BadRequestError( + message="Missing required parameter: 'prompt' for text completion.", + model=kwargs.get("model") or "", + llm_provider="", + ) + messages = prompt_value elif ( call_type == CallTypes.rerank.value or call_type == CallTypes.arerank.value ): @@ -1052,7 +1082,17 @@ def function_setup( # noqa: PLR0915 call_type == CallTypes.atranscription.value or call_type == CallTypes.transcription.value ): - _file_obj: FileTypes = args[1] if len(args) > 1 else kwargs["file"] + if len(args) > 1: + _file_obj: FileTypes = args[1] + else: + file_value = kwargs.get("file") + if file_value is None: + raise BadRequestError( + message="Missing required parameter: 'file' for transcription.", + model=kwargs.get("model") or "", + llm_provider="", + ) + _file_obj = file_value # Lazy import audio_utils.utils only when needed for transcription calls audio_utils = _get_cached_audio_utils() file_checksum = audio_utils.get_audio_file_content_hash(file_obj=_file_obj) diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 50dc3c6c6ec..c190e76d3b9 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -3865,3 +3865,60 @@ class TestValidateAndFixThinkingParam: validate_and_fix_thinking_param(thinking=thinking) assert "budgetTokens" in thinking assert "budget_tokens" not in thinking + + +class TestFunctionSetupMissingRequiredField: + """ + Regression tests: function_setup() must raise BadRequestError (status 400) + instead of a raw KeyError when the caller omits the required input field + for text-completion, image-generation, or moderation. + """ + + def _call_function_setup(self, original_function: str, **kwargs): + import uuid + from datetime import datetime + + from litellm.utils import Rules, function_setup + + return function_setup( + original_function=original_function, + rules_obj=Rules(), + start_time=datetime.now(), + litellm_call_id=str(uuid.uuid4()), + **kwargs, + ) + + def test_text_completion_missing_prompt_raises_bad_request(self): + with pytest.raises(litellm.BadRequestError) as exc_info: + self._call_function_setup( + original_function="atext_completion", + model="text-davinci-003", + ) + assert exc_info.value.status_code == 400 + assert "prompt" in str(exc_info.value) + + def test_text_completion_with_prompt_succeeds(self): + logging_obj, kwargs = self._call_function_setup( + original_function="atext_completion", + model="text-davinci-003", + prompt="hello", + ) + assert kwargs["prompt"] == "hello" + + def test_image_generation_missing_prompt_raises_bad_request(self): + with pytest.raises(litellm.BadRequestError) as exc_info: + self._call_function_setup( + original_function="aimage_generation", + model="dall-e-3", + ) + assert exc_info.value.status_code == 400 + assert "prompt" in str(exc_info.value) + + def test_moderation_missing_input_raises_bad_request(self): + with pytest.raises(litellm.BadRequestError) as exc_info: + self._call_function_setup( + original_function="amoderation", + model="text-moderation-latest", + ) + assert exc_info.value.status_code == 400 + assert "input" in str(exc_info.value)