mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(utils): raise BadRequestError for missing required field in function_setup
This commit is contained in:
parent
e64d98f725
commit
212f24459a
2 changed files with 101 additions and 4 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue