fix(proxy): return 4xx instead of 500 for missing required params, invalid pagination and unknown ids (#43787)

* fix(proxy): return 400 instead of 500 for missing required params and invalid pagination

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* ci: run search_endpoints tests in proxy-endpoints shard

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

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

* fix(llm_http_handler): keep provider error text when re-raising mapped errors

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): allow promptless image edits and default search models

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): default missing image edit image to None

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* refactor(proxy): build image edit defaults without mutating request data

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): keep image edit defaults within type-discipline budget and give request mocks a scope

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): inject a fake router for the search default model test

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(llms): cover provider error status on vector store and file lookup handlers

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(llms): keep the lookup handler raise block to a single statement

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(llms): cover provider error status on eval, eval run, skill and vector store file content lookups

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(integration): cover missing required body params and provider lookup status codes

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* ci: run tests/unit/proxy/search_endpoints in the proxy-endpoints shard

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(integration): bind spend-row request id with partial to satisfy B023

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): only reject non-positive page_size on vector store list

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): remove unreachable fine-tuning body validation

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(integration): cover streaming anthropic messages reaching the upstream

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(integration): count only provider calls when asserting missing params never reach the upstream

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): preserve merge-base request compatibility

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): preserve interaction completion model defaults

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): retry model read-through before rejecting params a DB-only deployment may default

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: shivam <shivam@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: yucheng <yucheng@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-10-02 22:48:53 -07:00 • committed by GitHub
parent 8efb4a21f6
commit 6d8434f940
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
28 changed files with 2788 additions and 64 deletions

View file

@ -165,6 +165,7 @@ jobs:
tests/unit/proxy/response_api_endpoints
tests/unit/proxy/image_endpoints
tests/unit/proxy/ocr_endpoints
tests/unit/proxy/search_endpoints
tests/unit/proxy/vector_store_endpoints
tests/unit/proxy/agent_endpoints
tests/unit/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
)
@ -4787,6 +4792,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
)
@ -5921,6 +5927,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,
@ -5958,7 +5996,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:
@ -7333,6 +7371,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,
@ -7411,6 +7450,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,
@ -8384,6 +8424,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,
@ -8565,6 +8606,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,
@ -8655,6 +8697,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,
@ -10152,6 +10195,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,
)
@ -10216,6 +10260,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,
)
@ -10282,6 +10327,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(
@ -10360,6 +10406,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(
@ -10828,6 +10875,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(
@ -10904,6 +10952,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(
@ -10963,6 +11012,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(
@ -11033,6 +11083,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(
@ -11092,6 +11143,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
)
@ -11164,6 +11216,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
)
@ -12124,6 +12177,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,
@ -12171,6 +12225,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,
@ -12227,6 +12282,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,
@ -12272,6 +12328,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,
@ -12542,6 +12599,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,
@ -12589,6 +12647,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,
@ -12645,6 +12704,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,
@ -12690,6 +12750,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,
@ -13167,6 +13228,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,
@ -13214,6 +13276,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,
@ -13270,6 +13333,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,
@ -13315,6 +13379,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

@ -283,7 +283,7 @@ async def create_batch(
)
data["metadata"] = sanitize_openai_provider_metadata(data.get("metadata"))
raise_if_required_body_param_missing(route_type="acreate_batch", data=data)
raise_if_required_body_param_missing(route_type="acreate_batch", data=data, llm_router=llm_router)
## check if model is a loadbalanced model
router_model: str | None = None

View file

@ -114,7 +114,9 @@ from litellm.proxy.common_utils.sse_keepalive import (
from litellm.proxy.dd_span_tagger import DDSpanTagger
from litellm.proxy.guardrails.auto_router_compression import arm_pre_call as _arm_auto_router_compression
from litellm.proxy.native_compaction import with_proxy_compaction_executor
from litellm.proxy.route_llm_request import route_request
from litellm.proxy.route_llm_request import (
route_request,
)
from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails
from litellm.router import Router
from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict

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

@ -907,6 +907,13 @@ async def get_daily_activity(
date_range: Final = parse_canonical_date_range(start_date, end_date)
if isinstance(date_range, InvalidDateRange):
raise_public(date_range)
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:
scope: Final = daily_activity_scope(
table_name,

View file

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

@ -310,7 +310,7 @@ async def responses_api(
route_type="aresponses",
llm_router=llm_router,
)
raise_if_required_body_param_missing(route_type="aresponses", data=data)
raise_if_required_body_param_missing(route_type="aresponses", data=data, llm_router=llm_router)
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,

View file

@ -1,9 +1,12 @@
import asyncio
from collections.abc import Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal
import httpx
from fastapi import HTTPException, status
from pydantic import TypeAdapter, ValidationError
import litellm
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
@ -164,6 +167,42 @@ REQUIRED_BODY_PARAMS_BY_ROUTE: Final[Mapping[str, tuple[str, ...]]] = {
"acreate_batch": ("input_file_id", "endpoint", "completion_window"),
}
REQUIRED_PRESENT_BODY_PARAMS_BY_ROUTE: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType(
{
"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",),
"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")}
)
JSON_OBJECT_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object])
class ProxyMissingRequiredParamError(ProxyException):
def __init__(self, route: str, param: str):
@ -175,16 +214,91 @@ class ProxyMissingRequiredParamError(ProxyException):
)
def raise_if_required_body_param_missing(route_type: str, data: Mapping[str, object]) -> None:
missing_param: Final = next(
class ProxyMissingParamWithoutLoadedModelError(ProxyMissingRequiredParamError):
pass
@dataclass(frozen=True, slots=True)
class MissingBodyParam:
name: str
model_deployments_loaded: bool
def _find_missing_required_body_param(
route_type: str,
data: Mapping[str, object],
llm_router: LitellmRouter | None,
) -> MissingBodyParam | 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 MissingBodyParam(name=one_of_params[0], model_deployments_loaded=True)
missing_merge_base_param: Final = next(
(param for param in REQUIRED_BODY_PARAMS_BY_ROUTE.get(route_type, ()) if data.get(param) is None),
None,
)
if missing_merge_base_param is not None:
return MissingBodyParam(name=missing_merge_base_param, model_deployments_loaded=True)
missing_present_params: Final = tuple(
param for param in REQUIRED_PRESENT_BODY_PARAMS_BY_ROUTE.get(route_type, ()) if param not in data
)
if not missing_present_params:
return None
candidate_litellm_params: Final = _candidate_deployment_litellm_params(data, llm_router)
missing_param: Final = next(
(
param
for param in missing_present_params
if not any(deployment_params.get(param) is not None for deployment_params in candidate_litellm_params)
),
None,
)
if missing_param is None:
return None
return MissingBodyParam(name=missing_param, model_deployments_loaded=bool(candidate_litellm_params))
def _candidate_deployment_litellm_params(
data: Mapping[str, object],
llm_router: LitellmRouter | None,
) -> tuple[dict[str, object], ...]:
model_name: Final = data.get("model")
if llm_router is None or not isinstance(model_name, str):
return ()
deployments: Final = (
llm_router.get_model_list(
model_name=model_name,
team_id=get_team_id_from_data(dict(data)),
)
or ()
)
return tuple(
params for deployment in deployments if (params := _validated_deployment_litellm_params(deployment)) is not None
)
def _validated_deployment_litellm_params(deployment: Mapping[str, object]) -> dict[str, object] | None:
try:
return JSON_OBJECT_ADAPTER.validate_python(deployment.get("litellm_params"))
except ValidationError:
return None
def raise_if_required_body_param_missing(
route_type: str,
data: Mapping[str, object],
llm_router: LitellmRouter | None,
) -> None:
missing_param: Final = _find_missing_required_body_param(route_type, data, llm_router)
if missing_param is None:
return
raise ProxyMissingRequiredParamError(
error_class: Final = (
ProxyMissingRequiredParamError
if missing_param.model_deployments_loaded
else ProxyMissingParamWithoutLoadedModelError
)
raise error_class(
route=ROUTE_ENDPOINT_MAPPING.get(route_type, route_type),
param=missing_param,
param=missing_param.name,
)
@ -442,9 +556,13 @@ async def route_request(
route_type=route_type,
user_api_key_dict=user_api_key_dict,
)
except ProxyModelNotFoundError as e:
except (ProxyModelNotFoundError, ProxyMissingParamWithoutLoadedModelError) as e:
requested_model: Final = data.get("model", "")
if not e.retryable_with_model_read_through or not isinstance(requested_model, str) or not requested_model:
if (
(isinstance(e, ProxyModelNotFoundError) and not e.retryable_with_model_read_through)
or not isinstance(requested_model, str)
or not requested_model
):
raise
from litellm.proxy import proxy_server
from litellm.proxy.common_utils.registry_read_through import (
@ -469,7 +587,7 @@ async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited pr
route_type: RouteType,
user_api_key_dict: UserAPIKeyAuth | None = None,
):
raise_if_required_body_param_missing(route_type=route_type, data=data)
raise_if_required_body_param_missing(route_type=route_type, data=data, llm_router=llm_router)
await add_shared_session_to_data(data)
@ -631,6 +749,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_size < 1:
raise HTTPException(
status_code=400,
detail=f"page_size must be >= 1, got 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

@ -1207,7 +1207,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

@ -31,6 +31,7 @@ from integration.cost_calculation.cost_tracking_case import (
)
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError
from starlette.applications import Starlette
from starlette.datastructures import UploadFile as StarletteUploadFile
from starlette.requests import Request
from starlette.responses import JSONResponse, Response, StreamingResponse
from starlette.routing import Route, WebSocketRoute
@ -59,6 +60,12 @@ def error_type(status: int) -> str:
return "invalid_request_error" if status < 500 else "server_error"
def _form_observation_value(value: str | StarletteUploadFile) -> JsonValue:
if isinstance(value, StarletteUploadFile):
return {"filename": value.filename, "content_type": value.content_type}
return value
@dataclass(frozen=True, slots=True)
class Observation:
path: str
@ -154,7 +161,9 @@ class Provider:
async def chat(self, request: Request) -> Response:
body: Final = JSON_OBJECT.validate_json(await request.body())
self.observations.put(Observation(request.url.path, request.headers.get("authorization", ""), body))
self.observations.put(
Observation(request.url.path, request.headers.get("authorization", ""), body, request.method)
)
leaked: Final = tuple(sorted(INTERNAL_FIELDS.intersection(body)))
if leaked:
return JSONResponse({"error": {"message": f"Unexpected provider fields: {leaked}"}}, status_code=400)
@ -188,7 +197,9 @@ class Provider:
async def vector_store_search(self, request: Request) -> Response:
body: Final = JSON_OBJECT.validate_json(await request.body())
self.observations.put(Observation(request.url.path, request.headers.get("authorization", ""), body))
self.observations.put(
Observation(request.url.path, request.headers.get("authorization", ""), body, request.method)
)
query: Final = body.get("query")
if not isinstance(query, str) or not query:
return JSONResponse({"error": {"message": "query is required"}}, status_code=400)
@ -227,7 +238,7 @@ class Provider:
self.scripts[name] = deque(int(str(value)) for value in statuses)
return JSONResponse({"configured": len(statuses)})
async def observed(self, _request: Request) -> Response:
async def observed(self, request: Request) -> Response:
values: Final = tuple(self.observations.get() for _ in range(self.observations.qsize()))
return JSONResponse(
{
@ -280,7 +291,8 @@ class Provider:
response: Final = self.scenario_store.get(scenario_id)
if response is None:
return JSONResponse({"error": "Unknown scenario"}, status_code=404)
if request.method == "POST" and "json" in request.headers.get("content-type", ""):
content_type: Final = request.headers.get("content-type", "")
if request.method == "POST" and "json" in content_type:
raw_body: Final = await request.body()
if raw_body:
body: Final = JSON_OBJECT.validate_json(raw_body)
@ -290,9 +302,20 @@ class Provider:
request.url.path,
request.headers.get("authorization", ""),
body,
api_key=request.headers.get("x-goog-api-key", ""),
request.method,
request.headers.get("x-goog-api-key", ""),
)
)
elif request.method == "POST" and "multipart/form-data" in content_type:
fields: Final = await request.form()
body: Final = {name: _form_observation_value(value) for name, value in fields.items()}
self.observations.put(
Observation(request.url.path, request.headers.get("authorization", ""), body, request.method)
)
elif request.method == "GET":
self.observations.put(
Observation(request.url.path, request.headers.get("authorization", ""), {}, request.method)
)
if isinstance(response, RoutedResponse):
route_key: Final = f"{request.method} /{'/'.join(segments[1:])}"
route: Final = next(

File diff suppressed because it is too large Load diff

View file

@ -67,6 +67,24 @@ def assert_config_write_refused(gateway: Gateway) -> None:
assert "config file" in str(error["error"]), refused.text
def test_list_page_zero_returns_same_stores_as_page_one_and_page_size_zero_is_400(gateway: Gateway) -> None:
page_one: Final = gateway.request("GET", "/vector_store/list?page=1&page_size=100")
assert page_one.status_code == 200, page_one.text
page_zero: Final = gateway.request("GET", "/vector_store/list?page=0&page_size=100")
assert page_zero.status_code == 200, page_zero.text
page_one_ids: Final = {str(row["vector_store_id"]) for row in listed_rows(page_one)}
page_zero_ids: Final = {str(row["vector_store_id"]) for row in listed_rows(page_zero)}
assert page_one_ids == page_zero_ids
assert CONFIG_STORE_ID in page_one_ids
assert CONFIG_STORE_ID in page_zero_ids
assert object_value(page_zero.json())["current_page"] == 0, page_zero.text
zero_page_size: Final = gateway.request("GET", "/vector_store/list?page=1&page_size=0")
assert zero_page_size.status_code == 400, zero_page_size.text
assert "page_size must be >= 1" in zero_page_size.text, zero_page_size.text
def burst_list(gateway: Gateway) -> tuple[int, str]:
response: Final = gateway.request("GET", "/vector_store/list")
if response.status_code != 200:

View file

@ -0,0 +1,274 @@
from __future__ import annotations
import asyncio
import json
import uuid
from pathlib import Path
from typing import Final, Literal
import httpx
import pytest
from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value
from integration._support.process import owned_proxy_process
from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario
from openai import APIStatusError, AsyncOpenAI, OpenAI
from pydantic import JsonValue
from litellm.proxy.openai_files_endpoints.common_utils import encode_file_id_with_model
from litellm.types.videos.utils import encode_video_id_with_provider
from tests.integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse
def _provider_error(status: int) -> dict[str, JsonValue]:
return {
"error": {
"message": f"scripted provider status {status}",
"type": "rate_limit_error"
if status == 429
else "server_error"
if status >= 500
else "invalid_request_error",
"code": str(status),
}
}
class _ObservationBuffer:
def __init__(self, upstream_url: str) -> None:
self._url = upstream_url.rstrip("/")
self._items: tuple[dict[str, JsonValue], ...] = ()
def read(self) -> tuple[dict[str, JsonValue], ...]:
with httpx.Client(timeout=10, trust_env=False) as client:
payload: Final = JSON_OBJECT.validate_python(client.get(f"{self._url}/__observations").json())
requests: Final = payload.get("requests")
assert isinstance(requests, list)
self._items = (*self._items, *(object_value(item) for item in requests if isinstance(item, dict)))
return self._items
def route(self, scenario_id: str, suffix: str) -> tuple[dict[str, JsonValue], ...]:
return tuple(
item
for item in self._items
if f"/{scenario_id}/" in str(item.get("path")) and str(item.get("path")).endswith(suffix)
)
def _ready(gateway: Gateway, model: str) -> None:
eventually(
lambda: gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "model readiness"}]},
),
lambda response: not (response.status_code == 400 and "Invalid model name" in response.text),
seconds=30,
)
def _assert_observation(gateway: Gateway, scenario_id: str, suffix: str) -> tuple[dict[str, JsonValue], ...]:
buffer: Final = _ObservationBuffer(gateway.upstream_url)
result: Final = eventually(
buffer.read,
lambda _items: len(buffer.route(scenario_id, suffix)) >= 1,
seconds=20,
)
observations: Final = buffer.route(scenario_id, suffix)
assert observations, result
return observations
def _add_vector_store(
gateway: Gateway,
scenario: Scenario,
vector_store_id: str,
alias: str,
handle: ScenarioHandle,
) -> None:
created: Final = gateway.request(
"POST",
"/vector_store/new",
{
"vector_store_id": vector_store_id,
"custom_llm_provider": "openai",
"litellm_params": {"model": alias, "api_base": handle.api_base(), "api_key": handle.scenario_id},
},
)
assert created.status_code == 200, created.text
scenario.cleanups.callback(gateway.post, "/vector_store/delete", {"vector_store_id": vector_store_id})
def _assert_provider_error(response: httpx.Response, status: int) -> None:
assert response.status_code == status, response.text
body: Final = JSON_OBJECT.validate_python(response.json())
error: Final = object_value(body["error"])
assert str(error.get("code")) == str(status), response.text
assert "scripted provider status" in str(error.get("message")), response.text
@pytest.mark.parametrize(
"status", (400, 401, 404, 429, 500), ids=("bad-request", "unauthorized", "not-found", "rate-limit", "server-error")
)
def test_vector_store_lookup_preserves_provider_status(gateway: Gateway, status: int) -> None:
with gateway.scenario() as scenario:
vector_store_id: Final = f"vs-{uuid.uuid4().hex}"
handle: Final = register_scenario(
f"vector-{uuid.uuid4().hex}",
RoutedResponse(
content_type="application/x-routed",
routes={
f"GET /vector_stores/{vector_store_id}": JsonResponse(
content_type="application/json",
status=status,
body=_provider_error(status),
)
},
),
)
scenario.cleanups.callback(delete_scenario, handle)
alias: Final = scenario.model(api_base=handle.api_base(), api_key=handle.scenario_id)
_add_vector_store(gateway, scenario, vector_store_id, alias, handle)
_ready(gateway, alias)
response: Final = gateway.request("GET", f"/v1/vector_stores/{vector_store_id}")
_assert_provider_error(response, status)
_assert_observation(gateway, handle.scenario_id, f"/vector_stores/{vector_store_id}")
@pytest.mark.parametrize(
("provider_route", "path", "model_name", "status"),
(
("GET /videos/video-id", "/v1/videos/video-id", "openai/gpt-4o-mini", 404),
("GET /v1/evals/eval-id", "/v1/evals/eval-id", "openai/gpt-4o-mini", 404),
("GET /v1/skills/skill-id", "/v1/skills/skill-id?beta=true", "anthropic/claude-3-5-haiku-20241022", 404),
("GET /v1/batch/jobs/batch-id", "/v1/batches/batch-id", "mistral/mistral-large-latest", 404),
("GET /v1/messages/batches/batch-id", "/v1/batches/batch-id", "anthropic/claude-3-5-haiku-20241022", 500),
),
ids=("video", "eval", "skill", "mistral-batch", "anthropic-batch-gap"),
)
def test_model_scoped_lookup_returns_scripted_provider_404(
gateway: Gateway,
provider_route: str,
path: str,
model_name: str,
status: int,
) -> None:
with gateway.scenario() as scenario:
route_path: Final = provider_route.partition(" ")[2]
handle: Final = register_scenario(
f"scoped-{uuid.uuid4().hex}",
RoutedResponse(
content_type="application/x-routed",
routes={
provider_route: JsonResponse(content_type="application/json", status=404, body=_provider_error(404))
},
),
)
scenario.cleanups.callback(delete_scenario, handle)
alias: Final = scenario.model(model=model_name, api_base=handle.api_base(), api_key=handle.scenario_id)
_ready(gateway, alias)
request_path: Final = (
f"/v1/videos/{encode_video_id_with_provider('video-id', 'openai', model_id=alias)}"
if "/videos/" in route_path
else f"/v1/batches/{encode_file_id_with_model('batch-id', alias, id_type='batch')}"
if "/batches/" in route_path or "/batch/jobs/" in route_path
else path
)
headers: Final = {"x-litellm-model": alias} if "/skills/" in request_path else {}
params: Final = {"model": alias} if "/evals/" in request_path else None
response: Final = gateway.request("GET", request_path, params=params, headers=headers)
assert response.status_code == status, response.text
body: Final = JSON_OBJECT.validate_python(response.json())
message: Final = str(object_value(body["error"]).get("message"))
assert (
"Client error '404 Not Found'" in message and "/v1/messages/batches/batch-id" in message
if status == 500
else "scripted provider status 404" in message
), response.text
suffix: Final = "/videos/video-id" if "/videos/" in route_path else route_path
_assert_observation(gateway, handle.scenario_id, suffix)
def _sdk_file_lookup(gateway: Gateway, file_id: str) -> httpx.Response:
with OpenAI(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0) as client:
try:
client.files.retrieve(file_id)
except APIStatusError as error:
return error.response
pytest.fail("OpenAI SDK file lookup unexpectedly succeeded")
async def _async_sdk_file_lookup(gateway: Gateway, file_id: str) -> httpx.Response:
async with AsyncOpenAI(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0) as client:
try:
await client.files.retrieve(file_id)
except APIStatusError as error:
return error.response
pytest.fail("OpenAI async SDK file lookup unexpectedly succeeded")
@pytest.mark.parametrize("client_kind", ("sync", "async"), ids=("sync", "async"))
def test_openai_sdk_file_lookup_returns_head_provider_error(
gateway: Gateway,
client_kind: Literal["sync", "async"],
tmp_path: Path,
) -> None:
with gateway.scenario() as scenario:
scenario_id: Final = f"d12-file-lookup-{client_kind}"
handle: Final = register_scenario(
scenario_id,
RoutedResponse(
content_type="application/x-routed",
routes={
"GET /files/file-id": JsonResponse(
content_type="application/json",
status=404,
body=_provider_error(404),
)
},
),
)
scenario.cleanups.callback(delete_scenario, handle)
model: Final = f"audit-file-lookup-{uuid.uuid4().hex}"
config: Final = tmp_path / f"d12-{client_kind}.yaml"
config.write_text(
json.dumps(
{
"model_list": [
{
"model_name": model,
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_base": handle.api_base(),
"api_key": scenario_id,
},
}
]
}
),
encoding="utf-8",
)
with owned_proxy_process(
gateway,
tmp_path,
{},
config=config,
workers=2,
) as owned:
candidate: Final = Gateway(owned.gateway.client, owned.gateway.key, gateway.upstream_url)
file_id: Final = encode_file_id_with_model("file-id", model)
response: Final = (
_sdk_file_lookup(candidate, file_id)
if client_kind == "sync"
else asyncio.run(_async_sdk_file_lookup(candidate, file_id))
)
expected: Final = {
"error": {
"message": f"Error code: 404 - {_provider_error(404)}",
"type": "invalid_request_error",
"param": None,
"code": "404",
}
}
assert response.status_code == 404, response.text
assert response.json() == expected, response.text
_assert_observation(candidate, scenario_id, "/files/file-id")

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

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

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

@ -600,6 +600,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,17 +1044,44 @@ 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"),
("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}])
def test_raise_if_required_body_param_missing_rejects_missing_param(route_type, param, route, data_extra):
def test_raise_if_required_body_param_missing_rejects_missing_param(route_type, param, route):
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={"model": "gpt-4o", **data_extra})
raise_if_required_body_param_missing(
route_type=route_type,
data={"model": "gpt-4o"},
llm_router=None,
)
assert exc_info.value.code == "400"
assert exc_info.value.param == param
@ -1079,22 +1105,149 @@ def test_raise_if_required_body_param_missing_names_first_missing_batch_param(da
)
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
raise_if_required_body_param_missing(route_type="acreate_batch", data=data)
raise_if_required_body_param_missing(route_type="acreate_batch", data=data, llm_router=None)
assert exc_info.value.param == param
def test_raise_if_required_body_param_missing_rejects_null_for_merge_base_route() -> None:
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="aembedding",
data={"model": "text-embedding-3-small", "input": None},
llm_router=None,
)
assert exc_info.value.param == "input"
@pytest.mark.parametrize(
("route_type", "data"),
(
pytest.param(
"anthropic_messages",
{"model": "claude", "messages": [], "max_tokens": None},
id="anthropic-max-tokens",
),
pytest.param(
"aimage_generation",
{"model": "gpt-image-1", "prompt": None},
id="image-prompt",
),
),
)
def test_required_present_body_param_accepts_explicit_null(route_type: str, data: dict[str, object]) -> None:
from litellm.proxy.route_llm_request import raise_if_required_body_param_missing
raise_if_required_body_param_missing(route_type=route_type, data=data, llm_router=None)
def test_required_present_body_param_uses_router_deployment_default() -> None:
import litellm
from litellm.proxy.route_llm_request import raise_if_required_body_param_missing
router = litellm.Router(
model_list=[
{
"model_name": "claude-default",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "test-key",
"max_tokens": 32,
},
}
]
)
raise_if_required_body_param_missing(
route_type="anthropic_messages",
data={"model": "claude-default", "messages": []},
llm_router=router,
)
def test_required_present_body_param_without_router_default_still_raises() -> None:
import litellm
from litellm.proxy.route_llm_request import (
ProxyMissingRequiredParamError,
raise_if_required_body_param_missing,
)
router = litellm.Router(
model_list=[
{
"model_name": "claude-without-default",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "test-key",
},
}
]
)
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
raise_if_required_body_param_missing(
route_type="anthropic_messages",
data={"model": "claude-without-default", "messages": []},
llm_router=router,
)
assert exc_info.value.param == "max_tokens"
@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, llm_router=None)
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"},
@ -1104,7 +1257,7 @@ def test_raise_if_required_body_param_missing_names_first_missing_batch_param(da
def test_raise_if_required_body_param_missing_allows_valid_requests(route_type, data):
from litellm.proxy.route_llm_request import raise_if_required_body_param_missing
raise_if_required_body_param_missing(route_type=route_type, data=data)
raise_if_required_body_param_missing(route_type=route_type, data=data, llm_router=None)
@pytest.mark.asyncio
@ -1257,6 +1410,66 @@ 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_read_through_supplies_db_model_default_for_missing_param(monkeypatch):
import litellm
import litellm.proxy.proxy_server as proxy_server
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
model_name = "e2e-db-only-max-tokens-default"
router = litellm.Router(
model_list=[{"model_name": "some-other-model", "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake"}}]
)
db_row = SimpleNamespace(
model_id=f"{model_name}-id",
model_name=model_name,
litellm_params={"model": "anthropic/claude-sonnet-4-5", "api_key": "fake", "max_tokens": 64},
model_info={},
blocked=False,
)
fake_prisma, table = _fake_prisma_client_with_models([db_row])
monkeypatch.setattr(proxy_server, "prisma_client", fake_prisma)
monkeypatch.setattr(proxy_server, "store_model_in_db", True)
monkeypatch.setattr(proxy_server, "llm_router", router)
data = {"model": model_name, "messages": [{"role": "user", "content": "hi"}]}
with patch.object(router, "anthropic_messages", new=AsyncMock(return_value="db_default_used")) as spy:
response = await (await route_request(data, router, None, "anthropic_messages"))
assert response == "db_default_used"
spy.assert_called_once()
assert table.find_many_wheres[0] == {"model_name": model_name}
@pytest.mark.asyncio
async def test_route_request_missing_param_for_unknown_model_still_400s_after_read_through(monkeypatch):
import litellm
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
model_name = "e2e-unknown-model-missing-max-tokens"
router = litellm.Router(
model_list=[{"model_name": "some-other-model", "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake"}}]
)
fake_prisma, table = _fake_prisma_client_with_models([])
monkeypatch.setattr(proxy_server, "prisma_client", fake_prisma)
monkeypatch.setattr(proxy_server, "store_model_in_db", True)
monkeypatch.setattr(proxy_server, "llm_router", router)
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
await route_request(
{"model": model_name, "messages": [{"role": "user", "content": "hi"}]},
router,
None,
"anthropic_messages",
)
assert (exc_info.value.code, exc_info.value.param) == ("400", "max_tokens")
assert table.find_many_wheres[0] == {"model_name": model_name}
@pytest.mark.asyncio
async def test_route_request_routing_group_name_passes_model_gate():
from unittest.mock import AsyncMock, patch
@ -1325,3 +1538,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

@ -5,6 +5,7 @@ Verifies that check_feature_access_for_user is called and that a 403 is
raised when vector stores are disabled for internal users.
"""
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -40,9 +41,7 @@ async def test_list_vector_stores_blocked_when_disabled():
)
user = _make_internal_user()
with patch.dict(
"litellm.proxy.proxy_server.general_settings", _DISABLED_GS, clear=True
):
with patch.dict("litellm.proxy.proxy_server.general_settings", _DISABLED_GS, clear=True):
with pytest.raises(HTTPException) as exc_info:
await list_vector_stores(user_api_key_dict=user)
assert exc_info.value.status_code == 403
@ -59,13 +58,9 @@ async def test_list_vector_stores_allowed_when_not_disabled():
user = _make_internal_user()
mock_prisma = MagicMock()
mock_prisma.db.litellm_managedvectorstorestable.find_many = AsyncMock(
return_value=[]
)
mock_prisma.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=[])
with patch.dict(
"litellm.proxy.proxy_server.general_settings", _ENABLED_GS, clear=True
):
with patch.dict("litellm.proxy.proxy_server.general_settings", _ENABLED_GS, clear=True):
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
with patch.object(litellm, "vector_store_registry", None):
with patch(
@ -92,9 +87,7 @@ async def test_new_vector_store_blocked_when_disabled():
user = _make_internal_user()
vs = LiteLLM_ManagedVectorStore(vector_store_id="vs-1", custom_llm_provider="openai") # type: ignore[call-arg]
with patch.dict(
"litellm.proxy.proxy_server.general_settings", _DISABLED_GS, clear=True
):
with patch.dict("litellm.proxy.proxy_server.general_settings", _DISABLED_GS, clear=True):
with pytest.raises(HTTPException) as exc_info:
await new_vector_store(vector_store=vs, user_api_key_dict=user)
assert exc_info.value.status_code == 403
@ -120,13 +113,9 @@ async def test_list_vector_stores_admin_not_blocked():
)
mock_prisma = MagicMock()
mock_prisma.db.litellm_managedvectorstorestable.find_many = AsyncMock(
return_value=[]
)
mock_prisma.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=[])
with patch.dict(
"litellm.proxy.proxy_server.general_settings", _DISABLED_GS, clear=True
):
with patch.dict("litellm.proxy.proxy_server.general_settings", _DISABLED_GS, clear=True):
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
with patch.object(litellm, "vector_store_registry", None):
with patch(
@ -135,3 +124,48 @@ 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_size", [0, -5])
async def test_list_vector_stores_rejects_non_positive_page_size_with_400(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=1, page_size=page_size)
assert exc_info.value.status_code == 400, exc_info.value.detail
assert "page_size" in exc_info.value.detail
@pytest.mark.asyncio
@pytest.mark.parametrize("page", [0, -1])
async def test_list_vector_stores_accepts_non_positive_page_like_base(page):
from litellm.proxy.vector_store_endpoints.management_endpoints import (
list_vector_stores,
)
import litellm
admin: Final = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN.value,
user_id="admin-1",
)
mock_prisma: Final = MagicMock()
mock_prisma.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=[])
with patch.dict("litellm.proxy.proxy_server.general_settings", _DISABLED_GS, clear=True):
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma):
with patch.object(litellm, "vector_store_registry", None):
with patch(
"litellm.proxy.vector_store_endpoints.management_endpoints.VectorStoreRegistry._get_vector_stores_from_db",
new=AsyncMock(return_value=[]),
):
response: Final = await list_vector_stores(user_api_key_dict=admin, page=page, page_size=10)
assert response["current_page"] == page
assert response["total_count"] == 0
assert response["data"] == []

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

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