From f37c5e36083f34350407c3b77de33fb8b3ed0020 Mon Sep 17 00:00:00 2001 From: shivam Date: Wed, 30 Sep 2026 19:43:39 +0000 Subject: [PATCH] fix(proxy): return 4xx for missing required params across all LLM routes and propagate provider status on lookups Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/llms/custom_httpx/llm_http_handler.py | 65 +++++++++++ litellm/proxy/proxy_server.py | 2 + litellm/proxy/route_llm_request.py | 44 +++++++- .../management_endpoints.py | 5 + litellm/utils.py | 2 +- tests/test_litellm/proxy/test_proxy_server.py | 7 ++ .../proxy/test_route_llm_request.py | 105 ++++++++++++++++-- .../test_vector_store_rbac.py | 14 +++ .../custom_httpx/test_llm_http_handler.py | 40 +++++++ tests/unit/test_utils.py | 5 + 10 files changed, 279 insertions(+), 10 deletions(-) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 8d65aa7b0ca..1d880616cdf 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -3910,6 +3910,7 @@ class BaseLLMHTTPHandler: provider_config=provider_config, ) + self._raise_for_provider_error_status(response=batch_response, provider_config=provider_config) return provider_config.transform_retrieve_batch_response( model=model, raw_response=batch_response, @@ -4067,6 +4068,7 @@ class BaseLLMHTTPHandler: provider_config=provider_config, ) + self._raise_for_provider_error_status(response=batch_response, provider_config=provider_config) return provider_config.transform_retrieve_batch_response( model=model, raw_response=batch_response, @@ -4484,6 +4486,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=provider_config) + self._raise_for_provider_error_status(response=response, provider_config=provider_config) return provider_config.transform_retrieve_file_response( raw_response=response, logging_obj=logging_obj, @@ -4540,6 +4543,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=provider_config) + self._raise_for_provider_error_status(response=response, provider_config=provider_config) return provider_config.transform_retrieve_file_response( raw_response=response, logging_obj=logging_obj, @@ -4732,6 +4736,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=provider_config) + self._raise_for_provider_error_status(response=response, provider_config=provider_config) files_per_page: Final = self._files_per_listing_page( response, provider_config, logging_obj, litellm_params, headers, sync_httpx_client, timeout ) @@ -4789,6 +4794,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=provider_config) + self._raise_for_provider_error_status(response=response, provider_config=provider_config) files_per_page: Final = self._files_per_async_listing_page( response, provider_config, logging_obj, litellm_params, headers, async_httpx_client, timeout ) @@ -5925,6 +5931,38 @@ class BaseLLMHTTPHandler: return None + def _raise_for_provider_error_status( + self, + response: httpx.Response, + provider_config: Union[ + BaseConfig, + BaseRerankConfig, + BaseResponsesAPIConfig, + BaseImageEditConfig, + BaseImageGenerationConfig, + BaseVectorStoreConfig, + BaseVectorStoreFilesConfig, + BaseGoogleGenAIGenerateContentConfig, + BaseAnthropicMessagesConfig, + BaseBatchesConfig, + BaseVideoConfig, + BaseSearchConfig, + BaseTextToSpeechConfig, + BaseSkillsAPIConfig, + "BasePassthroughConfig", + "BaseContainerConfig", + BaseEvalsAPIConfig, + BaseRealtimeHTTPConfig, + ], + ) -> None: + if not httpx.codes.is_error(response.status_code): + return + raise provider_config.get_error_class( + error_message=response.text, + status_code=response.status_code, + headers=response.headers, + ) + def _handle_error( self, e: Exception, @@ -7337,6 +7375,7 @@ class BaseLLMHTTPHandler: ) # Transform the response using the provider config + self._raise_for_provider_error_status(response=response, provider_config=video_content_provider_config) return video_content_provider_config.transform_video_content_response( raw_response=response, logging_obj=logging_obj, @@ -7415,6 +7454,7 @@ class BaseLLMHTTPHandler: ) # Transform the response using the provider config + self._raise_for_provider_error_status(response=response, provider_config=video_content_provider_config) return await video_content_provider_config.async_transform_video_content_response( raw_response=response, logging_obj=logging_obj, @@ -8388,6 +8428,7 @@ class BaseLLMHTTPHandler: params=params, ) + self._raise_for_provider_error_status(response=response, provider_config=video_list_provider_config) return video_list_provider_config.transform_video_list_response( raw_response=response, logging_obj=logging_obj, @@ -8569,6 +8610,7 @@ class BaseLLMHTTPHandler: headers=headers, ) + self._raise_for_provider_error_status(response=response, provider_config=video_status_provider_config) return video_status_provider_config.transform_video_status_retrieve_response( raw_response=response, logging_obj=logging_obj, @@ -8659,6 +8701,7 @@ class BaseLLMHTTPHandler: url=url, headers=headers, ) + self._raise_for_provider_error_status(response=response, provider_config=video_status_provider_config) return await video_status_provider_config.async_transform_video_status_retrieve_response( raw_response=response, logging_obj=logging_obj, @@ -10156,6 +10199,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=vector_store_provider_config) + self._raise_for_provider_error_status(response=response, provider_config=vector_store_provider_config) return vector_store_provider_config.transform_create_vector_store_response( response=response, ) @@ -10220,6 +10264,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=vector_store_provider_config) + self._raise_for_provider_error_status(response=response, provider_config=vector_store_provider_config) return vector_store_provider_config.transform_create_vector_store_response( response=response, ) @@ -10286,6 +10331,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=vector_store_provider_config) + self._raise_for_provider_error_status(response=response, provider_config=vector_store_provider_config) return response.json() def vector_store_list_handler( @@ -10364,6 +10410,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=vector_store_provider_config) + self._raise_for_provider_error_status(response=response, provider_config=vector_store_provider_config) return response.json() async def async_vector_store_update_handler( @@ -10832,6 +10879,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=vector_store_files_provider_config) + self._raise_for_provider_error_status(response=response, provider_config=vector_store_files_provider_config) return vector_store_files_provider_config.transform_list_vector_store_files_response(response=response) def vector_store_file_list_handler( @@ -10908,6 +10956,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=vector_store_files_provider_config) + self._raise_for_provider_error_status(response=response, provider_config=vector_store_files_provider_config) return vector_store_files_provider_config.transform_list_vector_store_files_response(response=response) async def async_vector_store_file_retrieve_handler( @@ -10967,6 +11016,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=vector_store_files_provider_config) + self._raise_for_provider_error_status(response=response, provider_config=vector_store_files_provider_config) return vector_store_files_provider_config.transform_retrieve_vector_store_file_response(response=response) def vector_store_file_retrieve_handler( @@ -11037,6 +11087,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=vector_store_files_provider_config) + self._raise_for_provider_error_status(response=response, provider_config=vector_store_files_provider_config) return vector_store_files_provider_config.transform_retrieve_vector_store_file_response(response=response) async def async_vector_store_file_content_handler( @@ -11096,6 +11147,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=vector_store_files_provider_config) + self._raise_for_provider_error_status(response=response, provider_config=vector_store_files_provider_config) return vector_store_files_provider_config.transform_retrieve_vector_store_file_content_response( response=response ) @@ -11168,6 +11220,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=vector_store_files_provider_config) + self._raise_for_provider_error_status(response=response, provider_config=vector_store_files_provider_config) return vector_store_files_provider_config.transform_retrieve_vector_store_file_content_response( response=response ) @@ -12128,6 +12181,7 @@ class BaseLLMHTTPHandler: provider_config=skills_api_provider_config, ) + self._raise_for_provider_error_status(response=response, provider_config=skills_api_provider_config) return skills_api_provider_config.transform_list_skills_response( raw_response=response, logging_obj=logging_obj, @@ -12175,6 +12229,7 @@ class BaseLLMHTTPHandler: provider_config=skills_api_provider_config, ) + self._raise_for_provider_error_status(response=response, provider_config=skills_api_provider_config) return skills_api_provider_config.transform_list_skills_response( raw_response=response, logging_obj=logging_obj, @@ -12231,6 +12286,7 @@ class BaseLLMHTTPHandler: provider_config=skills_api_provider_config, ) + self._raise_for_provider_error_status(response=response, provider_config=skills_api_provider_config) return skills_api_provider_config.transform_get_skill_response( raw_response=response, logging_obj=logging_obj, @@ -12276,6 +12332,7 @@ class BaseLLMHTTPHandler: provider_config=skills_api_provider_config, ) + self._raise_for_provider_error_status(response=response, provider_config=skills_api_provider_config) return skills_api_provider_config.transform_get_skill_response( raw_response=response, logging_obj=logging_obj, @@ -12546,6 +12603,7 @@ class BaseLLMHTTPHandler: provider_config=evals_api_provider_config, ) + self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config) return evals_api_provider_config.transform_list_evals_response( raw_response=response, logging_obj=logging_obj, @@ -12593,6 +12651,7 @@ class BaseLLMHTTPHandler: provider_config=evals_api_provider_config, ) + self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config) return evals_api_provider_config.transform_list_evals_response( raw_response=response, logging_obj=logging_obj, @@ -12649,6 +12708,7 @@ class BaseLLMHTTPHandler: provider_config=evals_api_provider_config, ) + self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config) return evals_api_provider_config.transform_get_eval_response( raw_response=response, logging_obj=logging_obj, @@ -12694,6 +12754,7 @@ class BaseLLMHTTPHandler: provider_config=evals_api_provider_config, ) + self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config) return evals_api_provider_config.transform_get_eval_response( raw_response=response, logging_obj=logging_obj, @@ -13171,6 +13232,7 @@ class BaseLLMHTTPHandler: provider_config=evals_api_provider_config, ) + self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config) return evals_api_provider_config.transform_list_runs_response( raw_response=response, logging_obj=logging_obj, @@ -13218,6 +13280,7 @@ class BaseLLMHTTPHandler: provider_config=evals_api_provider_config, ) + self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config) return evals_api_provider_config.transform_list_runs_response( raw_response=response, logging_obj=logging_obj, @@ -13274,6 +13337,7 @@ class BaseLLMHTTPHandler: provider_config=evals_api_provider_config, ) + self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config) return evals_api_provider_config.transform_get_run_response( raw_response=response, logging_obj=logging_obj, @@ -13319,6 +13383,7 @@ class BaseLLMHTTPHandler: provider_config=evals_api_provider_config, ) + self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config) return evals_api_provider_config.transform_get_run_response( raw_response=response, logging_obj=logging_obj, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index a994293a4d8..472a2ce3e5f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -12172,6 +12172,8 @@ async def completion( ) litellm_call_id: Final = request_litellm_call_id(data) log_llm_api_exception(e, litellm_call_id) + if isinstance(e, ProxyException): + raise with_litellm_call_id(e, litellm_call_id) error_msg: Final = f"{e}" raise ProxyException( message=getattr(e, "message", error_msg), diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 810ab7c018c..59deac0462f 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -1,5 +1,6 @@ import asyncio from collections.abc import Mapping +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal import httpx @@ -165,8 +166,35 @@ REQUIRED_BODY_PARAMS_BY_ROUTE: Final[Mapping[str, tuple[str, ...]]] = { "amoderation": ("input",), "aimage_generation": ("prompt",), "asearch": ("query",), + "atext_completion": ("prompt",), + "atranscription": ("file",), + "arerank": ("query", "documents"), + "acompact_responses": ("input",), + "aimage_edit": ("image", "prompt"), + "anthropic_messages": ("messages", "max_tokens"), + "agenerate_content": ("contents",), + "aocr": ("document",), + "acreate_fine_tuning_job": ("training_file",), + "avector_store_search": ("query",), + "avector_store_file_create": ("file_id",), + "avector_store_file_update": ("attributes",), + "avideo_generation": ("prompt",), + "avideo_remix": ("prompt",), + "avideo_edit": ("prompt",), + "avideo_extension": ("prompt", "seconds"), + "avideo_create_character": ("name", "video"), + "acreate_container": ("name",), + "aupload_container_file": ("file",), + "acreate_agent": ("name",), + "acreate_interaction": ("input",), + "acreate_eval": ("data_source_config", "testing_criteria"), + "acreate_run": ("data_source",), } +REQUIRED_ONE_OF_BODY_PARAMS_BY_ROUTE: Final[Mapping[str, tuple[str, str]]] = MappingProxyType( + {"acreate_interaction": ("model", "agent")} +) + class ProxyMissingRequiredParamError(ProxyException): def __init__(self, route: str, param: str): @@ -178,11 +206,18 @@ class ProxyMissingRequiredParamError(ProxyException): ) -def raise_if_required_body_param_missing(route_type: str, data: Mapping[str, object]) -> None: - missing_param: Final = next( +def _find_missing_required_body_param(route_type: str, data: Mapping[str, object]) -> str | None: + one_of_params: Final = REQUIRED_ONE_OF_BODY_PARAMS_BY_ROUTE.get(route_type) + if one_of_params is not None and all(data.get(param) is None for param in one_of_params): + return one_of_params[0] + return next( (param for param in REQUIRED_BODY_PARAMS_BY_ROUTE.get(route_type, ()) if data.get(param) is None), None, ) + + +def raise_if_required_body_param_missing(route_type: str, data: Mapping[str, object]) -> None: + missing_param: Final = _find_missing_required_body_param(route_type, data) if missing_param is None: return raise ProxyMissingRequiredParamError( @@ -634,6 +669,11 @@ async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited pr # These endpoints don't need a model, use custom_llm_provider directly return getattr(litellm, f"{route_type}")(**data) + if "model" not in data: + raise ProxyMissingRequiredParamError( + route=ROUTE_ENDPOINT_MAPPING.get(route_type, route_type), + param="model", + ) team_model_name: Final = llm_router.map_team_model(data["model"], team_id) if team_id is not None else None if team_model_name is not None: data["model"] = team_model_name diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index cae144bb266..e65b29dadf8 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -369,6 +369,11 @@ async def list_vector_stores( - page: int - Page number for pagination (default: 1) - page_size: int - Number of items per page (default: 100) """ + if page < 1 or page_size < 1: + raise HTTPException( + status_code=400, + detail=f"page and page_size must be >= 1, got page={page}, page_size={page_size}", + ) await check_feature_access_for_user(user_api_key_dict, "vector_stores") from litellm.proxy.proxy_server import prisma_client diff --git a/litellm/utils.py b/litellm/utils.py index eeccd27c1d8..d47bb643a60 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1188,7 +1188,7 @@ def function_setup( elif call_type == CallTypes.moderation.value or call_type == CallTypes.amoderation.value: messages = args[1] if len(args) > 1 else kwargs["input"] elif call_type == CallTypes.atext_completion.value or call_type == CallTypes.text_completion.value: - messages = args[0] if len(args) > 0 else kwargs["prompt"] + messages = args[0] if len(args) > 0 else kwargs.get("prompt") elif call_type == CallTypes.rerank.value or call_type == CallTypes.arerank.value: messages = kwargs.get("query") elif call_type in (CallTypes.search.value, CallTypes.asearch.value): diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 23d319159ca..924734eb3a4 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -601,6 +601,13 @@ def test_fallback_login_has_no_deprecation_banner(client_no_auth): assert " object: + async_client: Final = AsyncHTTPHandler() + await async_client.close() + async_client.client = httpx.AsyncClient(transport=httpx.MockTransport(lambda request: upstream_response)) + handler: Final = BaseLLMHTTPHandler() + if handler_name == "get_eval": + return await handler.async_get_eval_handler( + url="https://api.example.test/v1/evals/eval_missing", + evals_api_provider_config=OpenAIEvalsConfig(), + custom_llm_provider="openai", + litellm_params=GenericLiteLLMParams(), + logging_obj=Mock(), + client=async_client, + ) + return await handler.async_get_skill_handler( + url="https://api.example.test/v1/skills/skill_missing", + skills_api_provider_config=AnthropicSkillsConfig(), + custom_llm_provider="anthropic", + litellm_params=GenericLiteLLMParams(), + logging_obj=Mock(), + client=async_client, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("handler_name", ("get_eval", "get_skill")) +@pytest.mark.parametrize("status_code", (400, 401, 404, 429, 503)) +async def test_get_by_id_handlers_raise_the_provider_error_status(handler_name: str, status_code: int) -> None: + upstream_response: Final = httpx.Response(status_code, json={"error": {"message": "No such object"}}) + + with pytest.raises(BaseLLMException) as error: + await _get_by_id_with_upstream(handler_name, upstream_response) + + assert error.value.status_code == status_code + assert "No such object" in error.value.message diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index afc449d22c4..a8bb5147a88 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -6437,6 +6437,11 @@ def test_function_setup_logs_the_search_query_edit_prompt_and_ocr_document_summa assert _logged_request_messages(original_function, *args, **kwargs) == [{"role": "user", "content": expected}] +@pytest.mark.parametrize("original_function", ("atext_completion", "text_completion")) +def test_function_setup_without_a_prompt_leaves_the_missing_prompt_to_request_validation(original_function: str) -> None: + assert _logged_request_messages(original_function, model="gpt-4o") is None + + def test_search_with_a_mixed_type_query_list_still_reaches_its_own_validation_error() -> None: mixed_query: Final = cast(list[str], ["Eiffel Tower", 7]) # cast-ok: the invalid list is the point of the test