fix(utils): raise BadRequestError for missing required field in function_setup

This commit is contained in:
michelligabriele 2026-04-14 06:16:27 +02:00
parent e64d98f725
commit 212f24459a
No known key found for this signature in database
2 changed files with 101 additions and 4 deletions

View file

@ -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)

View file

@ -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)