This commit is contained in:
devin-ai-integration[bot] 2026-09-30 22:28:19 +00:00 • committed by GitHub
commit ec0071831d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
22 changed files with 632 additions and 29 deletions

View file

@ -148,6 +148,7 @@ jobs:
tests/test_litellm/proxy/response_api_endpoints
tests/test_litellm/proxy/image_endpoints
tests/test_litellm/proxy/ocr_endpoints
tests/test_litellm/proxy/search_endpoints
tests/test_litellm/proxy/vector_store_endpoints
tests/test_litellm/proxy/agent_endpoints
tests/test_litellm/proxy/a2a

View file

@ -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,
@ -5962,7 +6000,7 @@ class BaseLLMHTTPHandler:
if error_headers is None and error_response:
error_headers = getattr(error_response, "headers", None)
if error_response and hasattr(error_response, "text"):
error_text = getattr(error_response, "text", error_text)
error_text = getattr(error_response, "text", None) or error_text
if error_headers:
error_headers = dict(error_headers)
else:
@ -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,

View file

@ -1,6 +1,8 @@
import asyncio
import io
from collections.abc import Sequence
from itertools import chain
from types import MappingProxyType
from typing import Final, get_type_hints
import orjson
@ -36,6 +38,7 @@ from litellm.types.llms.openai import ChatCompletionUserMessage
router: Final = APIRouter()
IMAGE_EDIT_NUMERIC_FORM_FIELDS: Final = numeric_form_fields(get_type_hints(ImageEditRequestParams))
IMAGE_EDIT_OPTIONAL_FIELD_DEFAULTS: Final = MappingProxyType({"prompt": None, "image": None})
IMAGE_ARRAY_FIELD: Final = "image[]"
MASK_ARRAY_FIELD: Final = "mask[]"
@ -294,12 +297,13 @@ async def image_edit_api(
#########################################################
# Read request body and convert UploadFiles to BytesIO
#########################################################
form_fields: Final = coerce_numeric_form_fields(
parsed_body=await _read_request_body(request=request),
numeric_fields=IMAGE_EDIT_NUMERIC_FORM_FIELDS,
)
data: Final = {
key: value
for key, value in coerce_numeric_form_fields(
parsed_body=await _read_request_body(request=request),
numeric_fields=IMAGE_EDIT_NUMERIC_FORM_FIELDS,
).items()
for key, value in chain(IMAGE_EDIT_OPTIONAL_FIELD_DEFAULTS.items(), form_fields.items())
if key not in BRACKETED_FILE_FIELDS
}
image_files: Final = await batch_to_bytesio(image)
@ -316,10 +320,6 @@ async def image_edit_api(
detail=f"'{_field}' must be provided as a multipart file upload, not a string.",
)
# Ensure prompt exists in data (default to None for models that don't require it)
if "prompt" not in data:
data["prompt"] = None
#########################################################
# Process request
#########################################################

View file

@ -1220,6 +1220,12 @@ async def get_daily_activity(
detail={"error": "Please provide start_date and end_date"},
)
if page < 1 or page_size < 1:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"page and page_size must be >= 1, got page={page}, page_size={page_size}",
)
try:
where_conditions: Final = _build_where_conditions(
entity_id_field=entity_id_field,

View file

@ -12216,6 +12216,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),

View file

@ -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
@ -162,8 +163,38 @@ REQUIRED_BODY_PARAMS_BY_ROUTE: Final[Mapping[str, tuple[str, ...]]] = {
"aembedding": ("input",),
"aresponses": ("input",),
"acreate_batch": ("input_file_id", "endpoint", "completion_window"),
"aspeech": ("input",),
"amoderation": ("input",),
"aimage_generation": ("prompt",),
"asearch": ("query",),
"atext_completion": ("prompt",),
"atranscription": ("file",),
"arerank": ("query", "documents"),
"acompact_responses": ("input",),
"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):
@ -175,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(
@ -631,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

View file

@ -10,6 +10,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
router: Final = APIRouter()
@ -134,6 +135,11 @@ async def search(
if search_tool_name is not None:
data["search_tool_name"] = search_tool_name
if not (
data.get("search_tool_name") or data.get("model") or general_settings.get("completion_model") or user_model
):
raise ProxyMissingRequiredParamError(route="/search", param="search_tool_name")
if "search_tool_name" in data and data["search_tool_name"]:
data["model"] = data["search_tool_name"]
search_tool_name_value: Final = data["search_tool_name"]

View file

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

View file

@ -1205,7 +1205,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):

View file

@ -1034,6 +1034,7 @@ def _raw_batches_request(body: Dict[str, Any]) -> MagicMock:
request.url.__str__.return_value = "http://localhost/v1/batches"
request.url.path = "/v1/batches"
request.method = "POST"
request.scope = {"type": "http", "path": "/v1/batches", "method": "POST"}
request.query_params = {}
request.headers = {"Content-Type": "application/json"}
request.client = MagicMock()

View file

@ -222,6 +222,28 @@ def test_image_edit_multipart_n_that_is_not_a_number_is_left_alone(monkeypatch):
assert captured["n"] == "two"
@pytest.mark.parametrize(
"files, form, missing",
[
({}, {"model": "stability.stable-style-transfer-v1:0", "prompt": "oil painting"}, "image"),
(
{"image": ("tree.png", b"\x89PNG\r\n\x1a\n", "image/png")},
{"model": "stability.stable-image-remove-background-v1:0"},
"prompt",
),
],
)
def test_image_edit_without_an_optional_field_reaches_the_provider_with_it_set_to_none(
monkeypatch, files, form, missing
):
captured: Dict[str, Any] = {}
response = _image_edit_client(monkeypatch, captured).post("/v1/images/edits", files=files or None, data=form)
assert response.status_code == 200, response.text
assert missing in captured and captured[missing] is None, captured
@pytest.mark.asyncio
async def test_a_model_the_router_cannot_serve_answers_an_openai_typed_error(monkeypatch: pytest.MonkeyPatch):
"""A bare HTTPException carries no type or param, so the tail used to ship the
@ -290,7 +312,9 @@ async def test_failure_log_carries_the_callers_litellm_call_id(
async def fake_add_litellm_data_to_request(**kwargs: object) -> object:
return kwargs["data"]
async def fake_pre_call_hook(*, user_api_key_dict: UserAPIKeyAuth, data: dict[str, object], call_type: str) -> dict[str, object]:
async def fake_pre_call_hook(
*, user_api_key_dict: UserAPIKeyAuth, data: dict[str, object], call_type: str
) -> dict[str, object]:
return data
async def fake_post_call_failure_hook(**_: object) -> None:
@ -327,7 +351,9 @@ async def test_failure_log_carries_the_callers_litellm_call_id(
)
with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"), pytest.raises(ProxyException) as raised:
await endpoints.image_generation(request=request, fastapi_response=Response(), user_api_key_dict=UserAPIKeyAuth())
await endpoints.image_generation(
request=request, fastapi_response=Response(), user_api_key_dict=UserAPIKeyAuth()
)
assert raised.value.headers["x-litellm-call-id"] == call_id
record = next(r for r in caplog.records if "Exception occured" in r.getMessage())
@ -378,7 +404,9 @@ async def test_failure_before_the_provider_call_bills_the_callers_litellm_call_i
)
with pytest.raises(ProxyException) as raised:
await endpoints.image_generation(request=request, fastapi_response=Response(), user_api_key_dict=UserAPIKeyAuth())
await endpoints.image_generation(
request=request, fastapi_response=Response(), user_api_key_dict=UserAPIKeyAuth()
)
assert raised.value.headers["x-litellm-call-id"] == call_id
assert [data["litellm_call_id"] for data in hook_request_data] == [call_id]

View file

@ -115,6 +115,36 @@ async def test_get_daily_activity_order_has_id_tiebreaker():
)
@pytest.mark.asyncio
@pytest.mark.parametrize("page, page_size", [(0, 10), (-1, 10), (1, 0), (1, -5)])
async def test_get_daily_activity_rejects_non_positive_pagination_with_400(page, page_size):
from fastapi import HTTPException
mock_prisma = MagicMock()
mock_table = MagicMock()
mock_table.count = AsyncMock(return_value=0)
mock_table.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_dailyteamspend = mock_table
with pytest.raises(HTTPException) as exc_info:
await get_daily_activity(
prisma_client=mock_prisma,
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id=None,
entity_metadata_field=None,
start_date="2026-09-18",
end_date="2026-09-25",
model=None,
api_key=None,
page=page,
page_size=page_size,
)
assert exc_info.value.status_code == 400, exc_info.value.detail
mock_table.find_many.assert_not_called()
def test_is_user_agent_tag():
"""Test _is_user_agent_tag function."""
# Test None and empty string

View file

@ -0,0 +1,54 @@
from unittest.mock import AsyncMock, MagicMock
import orjson
import pytest
from litellm.proxy import proxy_server
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
from litellm.proxy.search_endpoints.endpoints import search
def _json_request(body: dict[str, object]) -> MagicMock:
request = MagicMock()
request.body = AsyncMock(return_value=orjson.dumps(body))
return request
@pytest.mark.asyncio
@pytest.mark.parametrize("body", [{"query": "litellm"}, {"query": "litellm", "search_tool_name": ""}])
async def test_search_without_search_tool_name_or_model_is_a_400(body):
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
await search(
request=_json_request(body),
fastapi_response=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
)
assert exc_info.value.code == "400"
assert exc_info.value.param == "search_tool_name"
assert exc_info.value.message == "/search: Missing required parameter: 'search_tool_name'."
@pytest.mark.asyncio
@pytest.mark.parametrize("default_source", ["cli_model", "completion_model"])
async def test_search_with_only_a_query_falls_back_to_the_proxy_default_model(monkeypatch, default_source):
if default_source == "cli_model":
monkeypatch.setattr(proxy_server, "user_model", "perplexity-search")
else:
monkeypatch.setitem(proxy_server.general_settings, "completion_model", "perplexity-search")
search_result = {"object": "search", "results": []}
router = MagicMock()
router.asearch = AsyncMock(return_value=search_result)
monkeypatch.setattr(proxy_server, "llm_router", router)
response = await search(
request=_json_request({"query": "litellm"}),
fastapi_response=MagicMock(),
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
)
assert response == search_result, response
router.asearch.assert_awaited_once()
assert router.asearch.await_args.kwargs["query"] == "litellm"
assert router.asearch.await_args.kwargs["model"] == "perplexity-search"

View file

@ -601,6 +601,13 @@ def test_fallback_login_has_no_deprecation_banner(client_no_auth):
assert "<form" in html
def test_text_completion_without_a_prompt_returns_400_naming_prompt(client_no_auth):
response = client_no_auth.post("/v1/completions", json={"model": "vllm_embed_model"})
assert response.status_code == 400, response.text
assert response.json()["error"]["param"] == "prompt", response.text
@pytest.mark.parametrize(
"ui_logo_path",
[

View file

@ -1,8 +1,6 @@
import pytest
from typing import Final
from unittest.mock import MagicMock
@ -14,14 +12,14 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError, route_reque
@pytest.mark.parametrize(
"route_type, required_body_params",
[
("atext_completion", {}),
("atext_completion", {"prompt": "Hello"}),
("acompletion", {"messages": [{"role": "user", "content": "Hello"}]}),
("aembedding", {"input": "Hello"}),
("aimage_generation", {}),
("aspeech", {}),
("atranscription", {}),
("amoderation", {}),
("arerank", {}),
("aimage_generation", {"prompt": "a cat"}),
("aspeech", {"input": "Hello"}),
("atranscription", {"file": b"audio"}),
("amoderation", {"input": "Hello"}),
("arerank", {"query": "Hello", "documents": ["hi"]}),
],
)
@pytest.mark.asyncio
@ -253,7 +251,7 @@ async def test_route_request_no_model_required():
for route_type in test_cases:
# Test data without model parameter
data = {"input": "test input", "api_key": "test-key"}
data = {"input": "test input", "query": "test query", "api_key": "test-key"}
llm_router = MagicMock()
getattr(llm_router, route_type).return_value = "fake_response"
@ -284,6 +282,7 @@ async def test_route_request_no_model_required_with_router_settings():
# Test data with model parameter (it will be ignored for these route types)
data = {
"input": "test input",
"query": "test query",
"model": "test-model", # Include dummy model to avoid KeyError
}
@ -1045,9 +1044,51 @@ async def test_route_request_override_enable_tag_filtering_beats_body_value():
("aembedding", "input", "/embeddings"),
("aresponses", "input", "/responses"),
("acreate_batch", "input_file_id", "/batches"),
("aspeech", "input", "/audio/speech"),
("amoderation", "input", "/moderations"),
("aimage_generation", "prompt", "/image/generations"),
("asearch", "query", "/search"),
("atext_completion", "prompt", "/completions"),
("atranscription", "file", "/audio/transcriptions"),
("arerank", "query", "/rerank"),
("acompact_responses", "input", "/responses/compact"),
("anthropic_messages", "messages", "anthropic_messages"),
("agenerate_content", "contents", "agenerate_content"),
("aocr", "document", "/ocr"),
("acreate_fine_tuning_job", "training_file", "acreate_fine_tuning_job"),
("avector_store_search", "query", "avector_store_search"),
("avector_store_file_create", "file_id", "avector_store_file_create"),
("avector_store_file_update", "attributes", "avector_store_file_update"),
("avideo_generation", "prompt", "/videos"),
("avideo_remix", "prompt", "/videos/{video_id}/remix"),
("avideo_edit", "prompt", "/videos/edits"),
("avideo_extension", "prompt", "/videos/extensions"),
("avideo_create_character", "name", "/videos/characters"),
("acreate_container", "name", "/containers"),
("aupload_container_file", "file", "/containers/{container_id}/files"),
("acreate_agent", "name", "/v1beta/agents"),
("acreate_eval", "data_source_config", "/evals"),
("acreate_run", "data_source", "/evals/{eval_id}/runs"),
],
)
@pytest.mark.parametrize(
"data_extra",
[
{},
{
"messages": None,
"input": None,
"input_file_id": None,
"prompt": None,
"query": None,
"file": None,
"image": None,
"contents": None,
"document": None,
"name": None,
},
],
)
@pytest.mark.parametrize("data_extra", [{}, {"messages": None, "input": None, "input_file_id": None}])
def test_raise_if_required_body_param_missing_rejects_missing_param(route_type, param, route, data_extra):
from litellm.proxy.route_llm_request import (
ProxyMissingRequiredParamError,
@ -1084,17 +1125,54 @@ def test_raise_if_required_body_param_missing_names_first_missing_batch_param(da
assert exc_info.value.param == param
@pytest.mark.parametrize(
"route_type, data, param",
[
("arerank", {"model": "rerank-model", "query": "hi"}, "documents"),
("anthropic_messages", {"model": "claude", "messages": []}, "max_tokens"),
("avideo_extension", {"model": "sora-2", "prompt": "longer"}, "seconds"),
("avideo_create_character", {"name": "hero"}, "video"),
("acreate_eval", {"data_source_config": {"type": "custom"}}, "testing_criteria"),
("acreate_interaction", {"input": "hi"}, "model"),
("acreate_interaction", {"model": None, "agent": None, "input": "hi"}, "model"),
("acreate_interaction", {"model": "gemini-3-pro-preview"}, "input"),
],
)
def test_raise_if_required_body_param_missing_names_each_missing_param(route_type, data, param):
from litellm.proxy.route_llm_request import (
ProxyMissingRequiredParamError,
raise_if_required_body_param_missing,
)
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
raise_if_required_body_param_missing(route_type=route_type, data=data)
assert exc_info.value.code == "400"
assert exc_info.value.param == param
@pytest.mark.parametrize(
"route_type, data",
[
("acompletion", {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}),
("acompletion", {"model": "gpt-4o", "messages": []}),
("atext_completion", {"model": "gpt-4o"}),
("atext_completion", {"model": "gpt-4o", "prompt": "hi"}),
("aembedding", {"model": "text-embedding-3-small", "input": "hi"}),
("aresponses", {"model": "gpt-4o", "input": "hi"}),
("aresponses", {"model": "gpt-4o", "input": []}),
("arerank", {"model": "rerank-model"}),
("aimage_generation", {"model": "dall-e-3"}),
("arerank", {"model": "rerank-model", "query": "hi", "documents": ["hello"]}),
("aimage_edit", {"model": "gpt-image-1", "image": b"png", "prompt": "a hat"}),
("aimage_edit", {"model": "stability.stable-image-remove-background-v1:0", "image": b"png"}),
("aimage_edit", {"model": "stability.stable-style-transfer-v1:0", "init_image": b"png"}),
("anthropic_messages", {"model": "claude", "messages": [], "max_tokens": 16}),
("avideo_extension", {"model": "sora-2", "prompt": "longer", "seconds": "4"}),
("acreate_eval", {"data_source_config": {"type": "custom"}, "testing_criteria": []}),
("acreate_interaction", {"model": "gemini-3-pro-preview", "input": "hi"}),
("acreate_interaction", {"agent": "deep-research", "input": "hi"}),
("aimage_generation", {"model": "gpt-image-1", "prompt": "a cat"}),
("aspeech", {"model": "gpt-4o-mini-tts", "input": "hi", "voice": "alloy"}),
("amoderation", {"model": "omni-moderation-latest", "input": ""}),
("asearch", {"model": "perplexity-search", "query": "litellm"}),
(
"acreate_batch",
{"input_file_id": "file-abc", "endpoint": "/v1/chat/completions", "completion_window": "24h"},
@ -1257,6 +1335,7 @@ async def test_route_request_read_through_disabled_without_store_model_in_db(mon
assert table.find_many_wheres == []
@pytest.mark.asyncio
async def test_route_request_routing_group_name_passes_model_gate():
from unittest.mock import AsyncMock, patch
@ -1325,3 +1404,23 @@ def test_proxy_model_not_found_error_keeps_the_raw_model_only_in_the_client_resp
assert raw_model in error.detail["error"]
assert raw_model not in error.spend_log_error_message
assert error.spend_log_error_message.startswith("/chat/completions: Invalid model name passed in")
@pytest.mark.asyncio
async def test_route_request_without_model_on_model_routed_endpoint_is_a_400():
import litellm
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
router = litellm.Router(
model_list=[
{"model_name": "rerank-model", "litellm_params": {"model": "cohere/rerank-v3.5", "api_key": "fake"}}
]
)
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
await route_request(
data={"query": "hi", "documents": ["hello"]}, llm_router=router, user_model=None, route_type="arerank"
)
assert exc_info.value.code == "400"
assert exc_info.value.param == "model"

View file

@ -135,3 +135,17 @@ async def test_list_vector_stores_admin_not_blocked():
):
# Must not raise any HTTPException — admin is always allowed.
await list_vector_stores(user_api_key_dict=admin)
@pytest.mark.asyncio
@pytest.mark.parametrize("page, page_size", [(0, 10), (-1, 10), (1, 0), (1, -5)])
async def test_list_vector_stores_rejects_non_positive_pagination_with_400(page, page_size):
from litellm.proxy.vector_store_endpoints.management_endpoints import (
list_vector_stores,
)
with pytest.raises(HTTPException) as exc_info:
await list_vector_stores(user_api_key_dict=_make_internal_user(), page=page, page_size=page_size)
assert exc_info.value.status_code == 400, exc_info.value.detail

View file

@ -1,5 +1,6 @@
import asyncio
import base64
import inspect
import json
import logging
import threading
@ -40,6 +41,11 @@ from litellm.llms.azure.videos.transformation import AzureVideoConfig
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeMessagesConfig,
)
from litellm.llms.anthropic.skills.transformation import AnthropicSkillsConfig
from litellm.llms.openai.evals.transformation import OpenAIEvalsConfig
from litellm.llms.mistral.files.transformation import MistralFilesConfig
from litellm.llms.openai.vector_store_files.transformation import OpenAIVectorStoreFilesConfig
from litellm.llms.openai.vector_stores.transformation import OpenAIVectorStoreConfig
from litellm.llms.openai.videos.transformation import OpenAIVideoConfig
from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig
from litellm.types.llms.openai import HttpxBinaryResponseContent, ResponsesAPIResponse
@ -4302,3 +4308,231 @@ async def test_async_text_to_speech_handler_records_upstream_response_headers():
assert response.content == b"audio-bytes"
_assert_upstream_headers_recorded(response)
async def _get_by_id_with_upstream(handler_name: str, upstream_response: httpx.Response) -> 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
def _clients_answering_with(upstream_response: httpx.Response) -> tuple[HTTPHandler, AsyncHTTPHandler]:
sync_client: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(lambda _: upstream_response)))
async_client: Final = AsyncHTTPHandler()
async_client.client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _: upstream_response))
return sync_client, async_client
def _call_lookup_handler(name: str, is_async: bool, client: HTTPHandler | AsyncHTTPHandler) -> object:
handler: Final = BaseLLMHTTPHandler()
vector_store_params: Final = GenericLiteLLMParams(api_base="https://api.example.test/v1", api_key="sk-test")
files_params: Final = {"api_base": "https://api.example.test", "api_key": "sk-test"}
match name:
case "vector_store_retrieve":
return handler.vector_store_retrieve_handler(
vector_store_id="vs_missing",
vector_store_provider_config=OpenAIVectorStoreConfig(),
custom_llm_provider="openai",
litellm_params=vector_store_params,
logging_obj=Mock(),
client=client,
_is_async=is_async,
)
case "vector_store_list":
return handler.vector_store_list_handler(
after=None,
before=None,
limit=None,
order=None,
vector_store_provider_config=OpenAIVectorStoreConfig(),
custom_llm_provider="openai",
litellm_params=vector_store_params,
logging_obj=Mock(),
client=client,
_is_async=is_async,
)
case "vector_store_file_list":
return handler.vector_store_file_list_handler(
vector_store_id="vs_missing",
query_params={},
vector_store_files_provider_config=OpenAIVectorStoreFilesConfig(),
custom_llm_provider="openai",
litellm_params=vector_store_params,
logging_obj=Mock(),
client=client,
_is_async=is_async,
)
case "vector_store_file_retrieve":
return handler.vector_store_file_retrieve_handler(
vector_store_id="vs_missing",
file_id="file_missing",
vector_store_files_provider_config=OpenAIVectorStoreFilesConfig(),
custom_llm_provider="openai",
litellm_params=vector_store_params,
logging_obj=Mock(),
client=client,
_is_async=is_async,
)
case "file_retrieve":
return handler.retrieve_file(
file_id="file_missing",
provider_config=MistralFilesConfig(),
litellm_params=files_params,
headers={},
logging_obj=Mock(),
_is_async=is_async,
client=client,
)
case "vector_store_file_content":
return handler.vector_store_file_content_handler(
vector_store_id="vs_missing",
file_id="file_missing",
vector_store_files_provider_config=OpenAIVectorStoreFilesConfig(),
custom_llm_provider="openai",
litellm_params=vector_store_params,
logging_obj=Mock(),
client=client,
_is_async=is_async,
)
case "eval_list":
return handler.list_evals_handler(
url="https://api.example.test/v1/evals",
query_params={},
evals_api_provider_config=OpenAIEvalsConfig(),
custom_llm_provider="openai",
litellm_params=vector_store_params,
logging_obj=Mock(),
client=client,
_is_async=is_async,
)
case "eval_get":
return handler.get_eval_handler(
url="https://api.example.test/v1/evals/eval_missing",
evals_api_provider_config=OpenAIEvalsConfig(),
custom_llm_provider="openai",
litellm_params=vector_store_params,
logging_obj=Mock(),
client=client,
_is_async=is_async,
)
case "eval_run_list":
return handler.list_runs_handler(
url="https://api.example.test/v1/evals/eval_missing/runs",
query_params={},
evals_api_provider_config=OpenAIEvalsConfig(),
custom_llm_provider="openai",
litellm_params=vector_store_params,
logging_obj=Mock(),
client=client,
_is_async=is_async,
)
case "eval_run_get":
return handler.get_run_handler(
url="https://api.example.test/v1/evals/eval_missing/runs/run_missing",
evals_api_provider_config=OpenAIEvalsConfig(),
custom_llm_provider="openai",
litellm_params=vector_store_params,
logging_obj=Mock(),
client=client,
_is_async=is_async,
)
case "skill_list":
return handler.list_skills_handler(
url="https://api.example.test/v1/skills",
query_params={},
skills_api_provider_config=AnthropicSkillsConfig(),
custom_llm_provider="anthropic",
litellm_params=vector_store_params,
logging_obj=Mock(),
client=client,
_is_async=is_async,
)
case "skill_get":
return handler.get_skill_handler(
url="https://api.example.test/v1/skills/skill_missing",
skills_api_provider_config=AnthropicSkillsConfig(),
custom_llm_provider="anthropic",
litellm_params=vector_store_params,
logging_obj=Mock(),
client=client,
_is_async=is_async,
)
case _:
return handler.list_files(
purpose=None,
provider_config=MistralFilesConfig(),
litellm_params=files_params,
headers={},
logging_obj=Mock(),
_is_async=is_async,
client=client,
)
async def _run_lookup_handler(name: str, is_async: bool, client: HTTPHandler | AsyncHTTPHandler) -> object:
result: Final = _call_lookup_handler(name, is_async, client)
return await result if inspect.isawaitable(result) else result
@pytest.mark.asyncio
@pytest.mark.parametrize(
"name",
(
"vector_store_retrieve",
"vector_store_list",
"vector_store_file_list",
"vector_store_file_retrieve",
"vector_store_file_content",
"file_retrieve",
"file_list",
"eval_list",
"eval_get",
"eval_run_list",
"eval_run_get",
"skill_list",
"skill_get",
),
)
@pytest.mark.parametrize("is_async", (False, True))
@pytest.mark.parametrize("status_code", (404, 503))
async def test_lookup_handlers_raise_the_provider_error_status(name: str, is_async: bool, status_code: int) -> None:
sync_client, async_client = _clients_answering_with(
httpx.Response(status_code, json={"error": {"message": "No such object"}})
)
with pytest.raises(BaseLLMException) as error:
await _run_lookup_handler(name, is_async, async_client if is_async else sync_client)
assert error.value.status_code == status_code
assert "No such object" in error.value.message

View file

@ -39,6 +39,7 @@ def mock_request(request):
mock_req.headers = Headers({"content-type": "application/json"})
mock_req.method = "POST"
mock_req.url.path = request.param.get("path")
mock_req.scope = {"type": "http", "path": request.param.get("path"), "method": "POST"}
async def mock_body():
return json.dumps(request.param.get("payload", {})).encode("utf-8")

View file

@ -195,7 +195,7 @@ async def test_route_request_returns_403_when_model_is_fully_blocked(monkeypatch
with pytest.raises(litellm.PermissionDeniedError) as exc_info:
await route_request(
data={"model": "gpt-4o"},
data={"model": "gpt-4o", "data_source_config": {"type": "custom"}, "testing_criteria": []},
llm_router=router,
user_model=None,
route_type="acreate_eval",

View file

@ -6479,6 +6479,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

View file

@ -1109,6 +1109,7 @@ def test_video_content_handler_passes_variant_to_url():
mock_client = MagicMock(spec=HTTPHandler)
mock_response = MagicMock()
mock_response.content = b"thumbnail-bytes"
mock_response.status_code = 200
mock_client.get.return_value = mock_response
with patch(
@ -1154,6 +1155,7 @@ def test_video_content_handler_uses_get_for_openai():
mock_client = MagicMock(spec=HTTPHandler)
mock_response = MagicMock()
mock_response.content = b"mp4-bytes"
mock_response.status_code = 200
mock_client.get.return_value = mock_response
# Patch _get_httpx_client to ensure no real HTTP client is created