diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 20096a0e373..06c2990a4a4 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -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 diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index dd97db45a88..1fc9ffaacb2 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -3910,6 +3910,7 @@ class BaseLLMHTTPHandler: provider_config=provider_config, ) + self._raise_for_provider_error_status(response=batch_response, provider_config=provider_config) return provider_config.transform_retrieve_batch_response( model=model, raw_response=batch_response, @@ -4067,6 +4068,7 @@ class BaseLLMHTTPHandler: provider_config=provider_config, ) + self._raise_for_provider_error_status(response=batch_response, provider_config=provider_config) return provider_config.transform_retrieve_batch_response( model=model, raw_response=batch_response, @@ -4484,6 +4486,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=provider_config) + self._raise_for_provider_error_status(response=response, provider_config=provider_config) return provider_config.transform_retrieve_file_response( raw_response=response, logging_obj=logging_obj, @@ -4540,6 +4543,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=provider_config) + self._raise_for_provider_error_status(response=response, provider_config=provider_config) return provider_config.transform_retrieve_file_response( raw_response=response, logging_obj=logging_obj, @@ -4732,6 +4736,7 @@ class BaseLLMHTTPHandler: except Exception as e: raise self._handle_error(e=e, provider_config=provider_config) + self._raise_for_provider_error_status(response=response, provider_config=provider_config) files_per_page: Final = self._files_per_listing_page( response, provider_config, logging_obj, litellm_params, headers, sync_httpx_client, timeout ) @@ -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, diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index f6c86d75169..1a5386b20d0 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -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 diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 660b7a261b8..a6651114cd0 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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 diff --git a/litellm/proxy/image_endpoints/endpoints.py b/litellm/proxy/image_endpoints/endpoints.py index 16dc38575da..4dc147b6687 100644 --- a/litellm/proxy/image_endpoints/endpoints.py +++ b/litellm/proxy/image_endpoints/endpoints.py @@ -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 ######################################################### diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 27b0960c823..28dfbb09eab 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -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, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index caec0ed524a..2425fc3887a 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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), diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 7badf0e79bb..882aa240fad 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -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, diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 323299f98fb..7da09ddcb68 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -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 diff --git a/litellm/proxy/search_endpoints/endpoints.py b/litellm/proxy/search_endpoints/endpoints.py index 2676682c59d..9cc76024770 100644 --- a/litellm/proxy/search_endpoints/endpoints.py +++ b/litellm/proxy/search_endpoints/endpoints.py @@ -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"] diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index cae144bb266..fca591a69ea 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -369,6 +369,11 @@ async def list_vector_stores( - page: int - Page number for pagination (default: 1) - page_size: int - Number of items per page (default: 100) """ + if page_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 diff --git a/litellm/utils.py b/litellm/utils.py index 26f412d3b75..35fea4f5f4a 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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): diff --git a/tests/integration/_support/upstream.py b/tests/integration/_support/upstream.py index 40423ae4d3a..4a624c923fd 100644 --- a/tests/integration/_support/upstream.py +++ b/tests/integration/_support/upstream.py @@ -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( diff --git a/tests/integration/compatibility/test_missing_body_param_status.py b/tests/integration/compatibility/test_missing_body_param_status.py new file mode 100644 index 00000000000..3e76ebaeda5 --- /dev/null +++ b/tests/integration/compatibility/test_missing_body_param_status.py @@ -0,0 +1,1570 @@ +from __future__ import annotations + +import asyncio +import json +import os +import signal +import socket +import subprocess +import sys +import uuid +from collections.abc import Iterator +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager +from functools import partial +from pathlib import Path +from typing import Final + +import httpx +import psutil +import pytest +from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value +from integration._support.database import read_rows +from integration._support.process import owned_proxy_process +from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario +from openai import AsyncOpenAI, BadRequestError, OpenAI +from pydantic import JsonValue + +from litellm.responses.utils import ResponsesAPIRequestUtils +from litellm.types.videos.utils import decode_video_id_with_provider, encode_video_id_with_provider +from tests.integration.cost_calculation.cost_tracking_case import ( + BinaryResponse, + JsonResponse, + RoutedResponse, + SseResponse, +) + +_Route = tuple[str, str, tuple[str, ...], dict[str, JsonValue], str] +_ROUTES: Final[dict[str, _Route]] = { + "acompletion": ( + "/v1/chat/completions", + "/chat/completions", + ("messages",), + {"messages": [{"role": "user", "content": "chat"}]}, + "openai/gpt-4o-mini", + ), + "aembedding": ( + "/v1/embeddings", + "/embeddings", + ("input",), + {"input": ["embedding"]}, + "openai/text-embedding-3-small", + ), + "aresponses": ("/v1/responses", "/responses", ("input",), {"input": "response"}, "openai/gpt-4o-mini"), + "acreate_batch": ( + "/v1/batches", + "/batches", + ("input_file_id", "endpoint", "completion_window"), + {"input_file_id": "file-audit", "endpoint": "/v1/chat/completions", "completion_window": "24h"}, + "openai/gpt-4o-mini", + ), + "aspeech": ( + "/v1/audio/speech", + "/audio/speech", + ("input",), + {"input": "speech", "voice": "alloy"}, + "openai/gpt-4o-mini-tts", + ), + "amoderation": ("/v1/moderations", "/moderations", ("input",), {"input": "moderate"}, "openai/gpt-4o-mini"), + "aimage_generation": ( + "/v1/images/generations", + "/image/generations", + ("prompt",), + {"prompt": "image"}, + "openai/gpt-image-1", + ), + "asearch": ("/v1/search/{tool}", "/search", ("query",), {"query": "search"}, "openai/gpt-4o-mini"), + "atext_completion": ( + "/v1/completions", + "/completions", + ("prompt",), + {"prompt": "complete"}, + "openai/gpt-3.5-turbo-instruct", + ), + "atranscription": ( + "/v1/audio/transcriptions", + "/audio/transcriptions", + ("file",), + {}, + "openai/gpt-4o-mini-transcribe", + ), + "arerank": ( + "/v1/rerank", + "/rerank", + ("query", "documents"), + {"query": "rank", "documents": ["first"]}, + "cohere/rerank-v4.0", + ), + "acompact_responses": ( + "/v1/responses/compact", + "/responses/compact", + ("input",), + {"input": "response"}, + "openai/gpt-4o-mini", + ), + "anthropic_messages": ( + "/v1/messages", + "anthropic_messages", + ("messages", "max_tokens"), + {"messages": [{"role": "user", "content": "message"}], "max_tokens": 8}, + "anthropic/claude-haiku-4-5", + ), + "agenerate_content": ( + "/v1beta/models/{model}:generateContent", + "agenerate_content", + ("contents",), + {"contents": [{"parts": [{"text": "Gemini"}]}]}, + "gemini/gemini-2.5-flash", + ), + "aocr": ("/v1/ocr", "/ocr", ("document",), {}, "mistral/mistral-ocr-latest"), + "acreate_fine_tuning_job": ( + "/v1/fine_tuning/jobs", + "/fine_tuning/jobs", + ("training_file",), + {"training_file": "file-audit", "model": "gpt-4o-mini"}, + "openai/gpt-4o-mini", + ), + "avector_store_search": ( + "/v1/vector_stores/{vector_store_id}/search", + "avector_store_search", + ("query",), + {"query": "vector query"}, + "openai/text-embedding-3-small", + ), + "avector_store_file_create": ( + "/v1/vector_stores/{vector_store_id}/files", + "avector_store_file_create", + ("file_id",), + {"file_id": "file-audit"}, + "openai/text-embedding-3-small", + ), + "avector_store_file_update": ( + "/v1/vector_stores/{vector_store_id}/files/{file_id}", + "avector_store_file_update", + ("attributes",), + {"attributes": {"source": "audit"}}, + "openai/text-embedding-3-small", + ), + "avideo_generation": ("/v1/videos", "/videos", ("prompt",), {"prompt": "video"}, "openai/sora-2"), + "avideo_remix": ( + "/v1/videos/{video_id}/remix", + "/videos/{video_id}/remix", + ("prompt",), + {"prompt": "remix"}, + "openai/sora-2", + ), + "avideo_edit": ( + "/v1/videos/edits", + "/videos/edits", + ("prompt",), + {"prompt": "edit", "video": {"id": "video-audit"}}, + "openai/sora-2", + ), + "avideo_extension": ( + "/v1/videos/extensions", + "/videos/extensions", + ("prompt", "seconds"), + {"prompt": "extend", "seconds": 5, "video_id": "video-audit"}, + "openai/sora-2", + ), + "avideo_create_character": ( + "/v1/videos/characters", + "/videos/characters", + ("name", "video"), + {"name": "character"}, + "openai/sora-2", + ), + "acreate_container": ("/v1/containers", "/containers", ("name",), {"name": "container"}, "openai/gpt-4o-mini"), + "aupload_container_file": ( + "/v1/containers/container-audit/files", + "/containers/{container_id}/files", + ("file",), + {}, + "openai/gpt-4o-mini", + ), + "acreate_agent": ( + "/v1beta/agents", + "/v1beta/agents", + ("name",), + {"name": "agent", "base_agent": "waverunner", "instructions": "You are a helpful assistant."}, + "gemini/gemini-2.5-flash", + ), + "acreate_interaction": ( + "/interactions", + "/interactions", + ("input", "model"), + {"input": "interaction"}, + "gemini/gemini-2.5-flash", + ), + "acreate_eval": ( + "/v1/evals", + "/evals", + ("data_source_config", "testing_criteria"), + {"data_source_config": {"type": "custom"}, "testing_criteria": [{"type": "string_check"}]}, + "openai/gpt-4o-mini", + ), + "acreate_run": ( + "/v1/evals/eval-audit/runs", + "/evals/{eval_id}/runs", + ("data_source",), + {"data_source": {"type": "custom"}}, + "openai/gpt-4o-mini", + ), +} +_MISSING: Final = tuple( + route + for route in _ROUTES + if route not in {"acreate_fine_tuning_job", "atranscription", "avideo_create_character", "aupload_container_file"} +) +_SKIP_VALID: Final = frozenset( + { + "aspeech", + "asearch", + "atranscription", + "aocr", + "avideo_create_character", + "aupload_container_file", + "acreate_fine_tuning_job", + "acreate_agent", + } +) +_NO_MODEL_BODY: Final = frozenset( + { + "asearch", + "agenerate_content", + "avector_store_search", + "avector_store_file_create", + "avector_store_file_update", + "acreate_agent", + } +) +_DOCUMENTED_GAPS: Final = ( + pytest.param("/v1/audio/transcriptions", {}, 422, ("body", "file"), None, False, id="atranscription-gap"), + pytest.param( + "/v1/videos/characters", + {"name": "character"}, + 422, + ("body", "video"), + None, + True, + id="avideo_create_character-gap", + ), + pytest.param( + "/v1/containers/container-audit/files", + {}, + 400, + None, + {"detail": "Missing required 'file' field"}, + False, + id="aupload_container_file-gap", + ), +) +_BODIES: Final[dict[str, dict[str, JsonValue]]] = { + "acompletion": { + "id": "chatcmpl-$UNIQUE_ID", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "scripted"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + }, + "aembedding": { + "object": "list", + "data": [{"object": "embedding", "embedding": [0.1, 0.2], "index": 0}], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 1, "total_tokens": 1}, + }, + "aresponses": { + "id": "resp_$UNIQUE_ID", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [], + "usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}, + }, + "acreate_batch": { + "id": "batch_$UNIQUE_ID", + "object": "batch", + "endpoint": "/v1/chat/completions", + "input_file_id": "file-audit", + "completion_window": "24h", + "created_at": 1, + "status": "validating", + }, + "amoderation": { + "id": "modr-$UNIQUE_ID", + "model": "omni-moderation-latest", + "results": [{"flagged": False, "categories": {}, "category_scores": {}}], + }, + "aimage_generation": {"created": 1, "data": [{"url": "https://images.invalid/audit.png"}]}, + "arerank": {"id": "rerank-$UNIQUE_ID", "results": [{"index": 0, "relevance_score": 0.5}], "meta": {}}, + "asearch": {"object": "search", "results": []}, + "atext_completion": { + "id": "cmpl-$UNIQUE_ID", + "object": "text_completion", + "created": 1, + "model": "gpt-3.5-turbo-instruct", + "choices": [{"text": "scripted", "index": 0, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + }, + "anthropic_messages": { + "id": "msg-$UNIQUE_ID", + "type": "message", + "role": "assistant", + "model": "claude-haiku-4-5", + "content": [{"type": "text", "text": "scripted"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 1, "output_tokens": 1}, + }, + "agenerate_content": { + "candidates": [{"content": {"parts": [{"text": "scripted"}], "role": "model"}, "finishReason": "STOP"}], + "usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}, + }, + "acompact_responses": { + "id": "resp_$UNIQUE_ID", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [], + "usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2}, + }, + "acreate_fine_tuning_job": { + "id": "ftjob-$UNIQUE_ID", + "object": "fine_tuning.job", + "created_at": 1, + "error": None, + "fine_tuned_model": None, + "finished_at": None, + "hyperparameters": {"n_epochs": "auto"}, + "model": "gpt-4o-mini", + "organization_id": "org-audit", + "result_files": [], + "seed": 1, + "status": "validating_files", + "trained_tokens": None, + "training_file": "file-audit", + "validation_file": None, + }, + "avector_store_search": {"object": "vector_store.search_results.page", "search_query": "vector query", "data": []}, + "avector_store_file_create": { + "id": "file-audit", + "object": "vector_store.file", + "created_at": 1, + "usage_bytes": 0, + "vector_store_id": "vs-audit", + "status": "completed", + "last_error": None, + "attributes": {}, + }, + "avector_store_file_update": { + "id": "file-audit", + "object": "vector_store.file", + "created_at": 1, + "usage_bytes": 0, + "vector_store_id": "vs-audit", + "status": "completed", + "last_error": None, + "attributes": {"source": "audit"}, + }, + "avideo_generation": { + "id": "video-audit-generation", + "object": "video", + "created_at": 1, + "status": "queued", + "model": "sora-2", + }, + "avideo_remix": { + "id": "video-audit-remix", + "object": "video", + "created_at": 1, + "status": "queued", + "model": "sora-2", + "remixed_from_video_id": "video-audit", + }, + "avideo_extension": { + "id": "video-audit-extension", + "object": "video", + "created_at": 1, + "status": "queued", + "model": "sora-2", + "seconds": "5", + }, + "acreate_container": { + "id": "container-audit", + "object": "container", + "created_at": 1, + "status": "running", + "name": "container", + }, + "acreate_agent": {"id": "agent-$UNIQUE_ID", "name": "agent"}, + "acreate_interaction": { + "id": "interaction-$UNIQUE_ID", + "object": "interaction", + "status": "completed", + "model": "gemini-2.5-flash", + }, + "acreate_eval": { + "id": "eval-$UNIQUE_ID", + "object": "eval", + "created_at": 1, + "data_source_config": {"type": "custom"}, + "testing_criteria": [{"type": "string_check"}], + }, + "acreate_run": { + "id": "evalrun-$UNIQUE_ID", + "object": "eval.run", + "created_at": 1, + "status": "queued", + "data_source": {"type": "custom"}, + "eval_id": "eval-audit", + }, + "avideo_edit": { + "id": "video-audit-edit", + "object": "video", + "created_at": 1, + "status": "queued", + "model": "sora-2", + }, + "avideo_create_character": { + "id": "character-audit", + "object": "character", + "created_at": 1, + "name": "character", + }, + "aupload_container_file": { + "id": "container-file-audit", + "object": "container.file", + "container_id": "container-audit", + "created_at": 1, + "path": "notes.txt", + "source": "user", + }, +} +_STREAM_RESPONSES: Final[dict[str, SseResponse]] = { + "acompletion": SseResponse( + content_type="text/event-stream", + frames=( + ( + 'data: {"id":"chatcmpl-$UNIQUE_ID","object":"chat.completion.chunk","created":1,' + '"model":"gpt-4o-mini","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}' + ), + ( + 'data: {"id":"chatcmpl-$UNIQUE_ID","object":"chat.completion.chunk","created":1,' + '"model":"gpt-4o-mini","choices":[{"index":0,"delta":{"content":"streamed "},"finish_reason":null}]}' + ), + ( + 'data: {"id":"chatcmpl-$UNIQUE_ID","object":"chat.completion.chunk","created":1,' + '"model":"gpt-4o-mini","choices":[{"index":0,"delta":{"content":"response"},"finish_reason":null}]}' + ), + ( + 'data: {"id":"chatcmpl-$UNIQUE_ID","object":"chat.completion.chunk","created":1,' + '"model":"gpt-4o-mini","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}' + ), + "data: [DONE]", + ), + ), + "aresponses": SseResponse( + content_type="text/event-stream", + frames=( + ( + "event: response.created\n" + 'data: {"type":"response.created","response":{"id":"resp_$REQUEST_ID","object":"response",' + '"created_at":1,"status":"in_progress","model":"gpt-4o-mini","output":[],"usage":null}}' + ), + ( + "event: response.output_item.added\n" + 'data: {"type":"response.output_item.added","output_index":0,' + '"item":{"type":"message","id":"msg_$REQUEST_ID","status":"in_progress",' + '"role":"assistant","content":[]}}' + ), + ( + "event: response.content_part.added\n" + 'data: {"type":"response.content_part.added","item_id":"msg_$REQUEST_ID",' + '"output_index":0,"content_index":0,' + '"part":{"type":"output_text","text":"","annotations":[]}}' + ), + ( + "event: response.output_text.delta\n" + 'data: {"type":"response.output_text.delta","item_id":"msg_$REQUEST_ID",' + '"output_index":0,"content_index":0,"delta":"streamed "}' + ), + ( + "event: response.output_text.delta\n" + 'data: {"type":"response.output_text.delta","item_id":"msg_$REQUEST_ID",' + '"output_index":0,"content_index":0,"delta":"response"}' + ), + ( + "event: response.output_text.done\n" + 'data: {"type":"response.output_text.done","item_id":"msg_$REQUEST_ID",' + '"output_index":0,"content_index":0,"text":"streamed response"}' + ), + ( + "event: response.content_part.done\n" + 'data: {"type":"response.content_part.done","item_id":"msg_$REQUEST_ID",' + '"output_index":0,"content_index":0,' + '"part":{"type":"output_text","text":"streamed response","annotations":[]}}' + ), + ( + "event: response.output_item.done\n" + 'data: {"type":"response.output_item.done","output_index":0,' + '"item":{"type":"message","id":"msg_$REQUEST_ID","status":"completed",' + '"role":"assistant","content":[{"type":"output_text","text":"streamed response",' + '"annotations":[]}]}}' + ), + ( + "event: response.completed\n" + 'data: {"type":"response.completed","response":{"id":"resp_$REQUEST_ID",' + '"object":"response","created_at":1,"status":"completed","model":"gpt-4o-mini",' + '"output":[{"id":"msg_$REQUEST_ID","type":"message","status":"completed",' + '"role":"assistant","content":[{"type":"output_text","text":"streamed response",' + '"annotations":[]}]}],"usage":{"input_tokens":1,"output_tokens":2}}}' + ), + ), + ), + "anthropic_messages": SseResponse( + content_type="text/event-stream", + frames=( + ( + "event: message_start\n" + 'data: {"type":"message_start","message":{"id":"msg_$REQUEST_ID","type":"message",' + '"role":"assistant","model":"claude-haiku-4-5","content":[],"stop_reason":null,' + '"stop_sequence":null,"usage":{"input_tokens":1,"output_tokens":0}}}' + ), + ( + "event: content_block_start\n" + 'data: {"type":"content_block_start","index":0,' + '"content_block":{"type":"text","text":""}}' + ), + ( + "event: content_block_delta\n" + 'data: {"type":"content_block_delta","index":0,' + '"delta":{"type":"text_delta","text":"streamed "}}' + ), + ( + "event: content_block_delta\n" + 'data: {"type":"content_block_delta","index":0,' + '"delta":{"type":"text_delta","text":"response"}}' + ), + ('event: content_block_stop\ndata: {"type":"content_block_stop","index":0}'), + ( + "event: message_delta\n" + 'data: {"type":"message_delta","delta":{"stop_reason":"end_turn",' + '"stop_sequence":null},"usage":{"output_tokens":2}}' + ), + 'event: message_stop\ndata: {"type":"message_stop"}', + ), + ), +} + + +class _Observations: + def __init__(self, url: str) -> None: + self.url = 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?include_method=true").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 for_scenario(self, identity: str) -> tuple[dict[str, JsonValue], ...]: + return tuple(item for item in self.items if f"/{identity}/" in str(item.get("path"))) + + def provider_calls(self, identity: str) -> tuple[dict[str, JsonValue], ...]: + calls: Final = tuple( + item + for item in self.for_scenario(identity) + if not (item.get("method") == "GET" and str(item.get("path", "")).endswith(("/v1/models", "/models"))) + ) + return calls + + +def _response( + route: str, + *, + streaming: bool = False, +) -> BinaryResponse | JsonResponse | RoutedResponse | SseResponse: + if route == "aspeech": + return BinaryResponse(content_type="audio/mpeg", length=16) + if streaming: + return _STREAM_RESPONSES[route] + if route == "acreate_fine_tuning_job": + return RoutedResponse( + content_type="application/x-routed", + routes={ + "POST /files": JsonResponse( + content_type="application/json", + body={ + "id": "file-training", + "object": "file", + "purpose": "fine-tune", + "filename": "training.jsonl", + "bytes": 90, + "created_at": 1, + "status": "processed", + }, + ), + "POST /fine_tuning/jobs": JsonResponse( + content_type="application/json", + body=_BODIES[route], + ), + }, + ) + if route == "aupload_container_file": + return RoutedResponse( + content_type="application/x-routed", + routes={ + "POST /containers": JsonResponse( + content_type="application/json", + body=_BODIES["acreate_container"], + ), + "POST /containers/container-audit/files": JsonResponse( + content_type="application/json", + body=_BODIES[route], + ), + }, + ) + return JsonResponse( + content_type="application/json", body=_BODIES.get(route, {"id": "audit-$UNIQUE_ID", "object": "audit_response"}) + ) + + +def _assert_scripted_response(route: str, caller: dict[str, JsonValue]) -> None: + scripted: Final = _BODIES[route] + if route == "acompletion": + expected_choices: Final = scripted["choices"] + actual_choices: Final = caller["choices"] + assert isinstance(expected_choices, list) and isinstance(actual_choices, list) + expected_choice: Final = object_value(expected_choices[0]) + actual_choice: Final = object_value(actual_choices[0]) + assert object_value(expected_choice["message"])["content"] == object_value(actual_choice["message"])["content"] + elif route == "atext_completion": + expected_choices = scripted["choices"] + actual_choices = caller["choices"] + assert isinstance(expected_choices, list) and isinstance(actual_choices, list) + expected_choice = object_value(expected_choices[0]) + actual_choice = object_value(actual_choices[0]) + assert actual_choice["text"] == expected_choice["text"] + elif route == "aembedding": + expected_data: Final = scripted["data"] + actual_data: Final = caller["data"] + assert isinstance(expected_data, list) and isinstance(actual_data, list) + expected_item: Final = object_value(expected_data[0]) + actual_item: Final = object_value(actual_data[0]) + assert actual_item["embedding"] == expected_item["embedding"] + elif route == "amoderation": + expected_results: Final = scripted["results"] + actual_results: Final = caller["results"] + assert isinstance(expected_results, list) and isinstance(actual_results, list) + expected_result: Final = object_value(expected_results[0]) + actual_result: Final = object_value(actual_results[0]) + assert actual_result["flagged"] is expected_result["flagged"] + elif route == "aimage_generation": + expected_data = scripted["data"] + actual_data = caller["data"] + assert isinstance(expected_data, list) and isinstance(actual_data, list) + expected_item = object_value(expected_data[0]) + actual_item = object_value(actual_data[0]) + assert actual_item["url"] == expected_item["url"] + elif route == "arerank": + expected_results = scripted["results"] + actual_results = caller["results"] + assert isinstance(expected_results, list) and isinstance(actual_results, list) + expected_result = object_value(expected_results[0]) + actual_result = object_value(actual_results[0]) + assert actual_result["index"] == expected_result["index"] + assert actual_result["relevance_score"] == expected_result["relevance_score"] + elif route == "anthropic_messages": + expected_content_list: Final = scripted["content"] + actual_content_list: Final = caller["content"] + assert isinstance(expected_content_list, list) and isinstance(actual_content_list, list) + expected_content: Final = object_value(expected_content_list[0]) + actual_content: Final = object_value(actual_content_list[0]) + assert actual_content["text"] == expected_content["text"] + elif route == "agenerate_content": + expected_candidates: Final = scripted["candidates"] + actual_candidates: Final = caller["candidates"] + assert isinstance(expected_candidates, list) and isinstance(actual_candidates, list) + expected_candidate: Final = object_value(expected_candidates[0]) + actual_candidate: Final = object_value(actual_candidates[0]) + expected_parts: Final = object_value(expected_candidate["content"])["parts"] + actual_parts: Final = object_value(actual_candidate["content"])["parts"] + assert isinstance(expected_parts, list) and isinstance(actual_parts, list) + expected_part: Final = object_value(expected_parts[0]) + actual_part: Final = object_value(actual_parts[0]) + assert actual_part["text"] == expected_part["text"] + elif route in { + "aresponses", + "acreate_batch", + "acompact_responses", + "avector_store_search", + "avector_store_file_create", + "avector_store_file_update", + "avideo_generation", + "avideo_remix", + "avideo_extension", + "avideo_edit", + "avideo_create_character", + "aupload_container_file", + "acreate_container", + "acreate_agent", + "acreate_interaction", + "acreate_eval", + "acreate_run", + }: + for field in ("status", "object"): + if field in scripted: + assert caller.get(field) == scripted[field] + if "id" in scripted: + actual_id: Final = caller.get("id") + expected_id: Final = str(scripted["id"]) + assert isinstance(actual_id, str) and actual_id + if route in {"avideo_generation", "avideo_remix", "avideo_extension", "avideo_edit"}: + assert decode_video_id_with_provider(actual_id)["video_id"] == expected_id + elif route == "acreate_container": + assert ResponsesAPIRequestUtils.decode_container_id_to_original(actual_id) == expected_id + else: + expected_prefix: Final = expected_id.split("$UNIQUE_ID", maxsplit=1)[0] + assert actual_id.startswith(expected_prefix) + if route in {"avideo_create_character", "acreate_agent"}: + assert caller.get("name") == scripted["name"] + if route == "aupload_container_file": + assert caller.get("container_id") == scripted["container_id"] + + +def _register( + scenario: Scenario, + route: str, + *, + streaming: bool = False, + deployment_params: dict[str, JsonValue] | None = None, + provider_model: str | None = None, + response_route: str | None = None, +) -> tuple[str, str, ScenarioHandle]: + identity: Final = f"audit-{route}-{uuid.uuid4().hex}" + handle: Final = register_scenario(identity, _response(response_route or route, streaming=streaming)) + scenario.cleanups.callback(delete_scenario, handle) + model: Final = ( + "" + if route == "asearch" + else scenario.model( + model=provider_model or _ROUTES[route][4], + api_base=handle.api_base(), + api_key=identity, + **(deployment_params or {}), + ) + ) + return model, identity, handle + + +def _path(template: str, model: str, tool: str, store: str) -> str: + video: Final = ( + encode_video_id_with_provider("video-audit", "openai", model_id=model) if "{video_id}" in template else "" + ) + return ( + template.replace("{model}", model) + .replace("{tool}", tool) + .replace("{vector_store_id}", store) + .replace("{file_id}", "file-audit") + .replace("{video_id}", video) + ) + + +def _store(gateway: Gateway, scenario: Scenario, model: str, identity: str, handle: ScenarioHandle, store: str) -> None: + response: Final = gateway.request( + "POST", + "/vector_store/new", + { + "vector_store_id": store, + "custom_llm_provider": "openai", + "litellm_params": {"model": model, "api_base": handle.api_base(), "api_key": identity}, + }, + ) + assert response.status_code == 200, response.text + scenario.cleanups.callback(gateway.post, "/vector_store/delete", {"vector_store_id": store}) + + +def _search_tool(gateway: Gateway, scenario: Scenario, identity: str, handle: ScenarioHandle) -> str: + name: Final = f"audit-search-{uuid.uuid4().hex}" + created: Final = gateway.post( + "/search_tools", + { + "search_tool": { + "search_tool_name": name, + "litellm_params": {"search_provider": "exa_ai", "api_key": identity, "api_base": handle.api_base()}, + } + }, + ) + scenario.cleanups.callback( + lambda tool_id: gateway.request("DELETE", f"/search_tools/{tool_id}"), str(created["search_tool_id"]) + ) + return name + + +def _error(route: str, parameter: str) -> dict[str, JsonValue]: + message: Final = f"{route}: Missing required parameter: '{parameter}'." + return ( + {"type": "error", "error": {"type": "invalid_request_error", "message": message}} + if route == "anthropic_messages" + else {"error": {"message": message, "type": "invalid_request_error", "param": parameter, "code": "400"}} + ) + + +def _missing_response( + gateway: Gateway, + path: str, + body: dict[str, JsonValue], + expected: dict[str, JsonValue], +) -> httpx.Response: + response: Final = _post(gateway, path, body) + assert response.status_code == 400, response.text + assert response.json() == expected, response.text + return response + + +def _post(gateway: Gateway, path: str, body: dict[str, JsonValue]) -> httpx.Response: + return gateway.client.post(path, json=body, headers={"Authorization": f"Bearer {gateway.key}"}) + + +def _observed( + buffer: _Observations, + identity: str, + expected_count: int = 1, +) -> tuple[dict[str, JsonValue], ...]: + eventually(buffer.read, lambda _items: len(buffer.for_scenario(identity)) == expected_count, seconds=20) + return buffer.for_scenario(identity) + + +def _stream_event_payloads(lines: tuple[str, ...]) -> tuple[dict[str, JsonValue], ...]: + return tuple(JSON_OBJECT.validate_json(line.removeprefix("data: ")) for line in lines if line.startswith("data: {")) + + +def _stream_event_names(lines: tuple[str, ...]) -> tuple[str, ...]: + return tuple(line.removeprefix("event: ") for line in lines if line.startswith("event: ")) + + +def _stream_event_text(route: str, event: dict[str, JsonValue]) -> str: + if route == "acompletion": + choices: Final = event.get("choices") + if not isinstance(choices, list) or not choices: + return "" + delta: Final = object_value(object_value(choices[0]).get("delta")) + content: Final = delta.get("content") + return content if isinstance(content, str) else "" + if route == "aresponses" and event.get("type") == "response.output_text.delta": + delta: Final = event.get("delta") + return delta if isinstance(delta, str) else "" + if route == "anthropic_messages" and event.get("type") == "content_block_delta": + delta: Final = object_value(event.get("delta")) + if delta.get("type") != "text_delta": + return "" + text: Final = delta.get("text") + return text if isinstance(text, str) else "" + return "" + + +def _assembled_stream_text(route: str, lines: tuple[str, ...]) -> str: + return "".join(_stream_event_text(route, event) for event in _stream_event_payloads(lines)) + + +@pytest.mark.parametrize("route", _MISSING, ids=_MISSING) +def test_added_required_fields_return_exact_400(gateway: Gateway, route: str) -> None: + template, error_route, fields, valid_body, _provider_model = _ROUTES[route] + with gateway.scenario() as scenario: + model, identity, handle = _register(scenario, route) + tool: Final = _search_tool(gateway, scenario, identity, handle) if route == "asearch" else "" + store: Final = f"vs-{uuid.uuid4().hex}" + if route.startswith("avector_store_"): + _store(gateway, scenario, model, identity, handle, store) + path: Final = _path(template, model, tool, store) + observations: Final = _Observations(gateway.upstream_url) + for field in fields: + missing_body: Final = { + key: value + for key, value in {**valid_body, **({"model": model} if route not in _NO_MODEL_BODY else {})}.items() + if key != field + } + body: Final = { + **missing_body, + **( + { + "video": { + "id": encode_video_id_with_provider("video-audit", "openai", model_id=model), + } + } + if route == "avideo_edit" + else {} + ), + } + expected: Final = _error(error_route, field) + response: Final = _missing_response(gateway, path, body, expected) + assert response.status_code == 400, f"{route}.{field}: {response.text}" + assert response.json() == expected, response.text + observations.read() + assert observations.provider_calls(identity) == () + + +@pytest.mark.parametrize( + ("path", "body", "expected_status", "expected_loc", "expected_body", "as_form"), _DOCUMENTED_GAPS +) +def test_unchanged_from_base_documented_gaps( + gateway: Gateway, + path: str, + body: dict[str, JsonValue], + expected_status: int, + expected_loc: tuple[str, str] | None, + expected_body: dict[str, JsonValue] | None, + as_form: bool, +) -> None: + response: Final = ( + gateway.client.post(path, data=body, headers={"Authorization": f"Bearer {gateway.key}"}) + if as_form + else _post(gateway, path, body) + ) + assert response.status_code == expected_status, response.text + if expected_body is not None: + assert response.json() == expected_body, response.text + else: + assert expected_loc is not None + detail: Final = JSON_OBJECT.validate_python(response.json()).get("detail") + assert isinstance(detail, list) and detail, response.text + assert object_value(detail[0]).get("loc") == list(expected_loc), response.text + + +def test_fine_tuning_missing_training_file_returns_422_without_upstream_call(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model, identity, _handle = _register(scenario, "acreate_fine_tuning_job") + response: Final = _post(gateway, "/v1/fine_tuning/jobs", {"model": model}) + assert response.status_code == 422, response.text + detail: Final = JSON_OBJECT.validate_python(response.json()).get("detail") + assert isinstance(detail, list) and detail, response.text + assert object_value(detail[0]).get("loc") == ["body", "training_file"], response.text + observations: Final = _Observations(gateway.upstream_url) + observations.read() + assert observations.provider_calls(identity) == () + + +@pytest.mark.parametrize( + "route", + tuple(name for name in _ROUTES if name not in _SKIP_VALID), + ids=tuple(name for name in _ROUTES if name not in _SKIP_VALID), +) +def test_valid_required_fields_reach_upstream(gateway: Gateway, route: str) -> None: + template, _error_route, fields, body, provider_model = _ROUTES[route] + with gateway.scenario() as scenario: + model, identity, handle = _register(scenario, route) + store: Final = f"vs-{uuid.uuid4().hex}" + if route.startswith("avector_store_"): + _store(gateway, scenario, model, identity, handle, store) + path: Final = _path(template, model, "", store) + request_body: Final = { + **body, + **( + { + "video": { + "id": encode_video_id_with_provider("video-audit", "openai", model_id=model), + } + } + if route == "avideo_edit" + else {} + ), + **({"model": model} if route not in _NO_MODEL_BODY else {}), + **({"input": f"response-{identity}"} if route == "aresponses" else {}), + } + response: Final = eventually( + lambda: _post(gateway, path, request_body), + lambda result: not (result.status_code == 400 and "Invalid model name" in result.text), + seconds=30, + ) + assert response.status_code == 200, response.text + caller: Final = JSON_OBJECT.validate_python(response.json()) + _assert_scripted_response(route, caller) + matches: Final = _observed(_Observations(gateway.upstream_url), identity) + outbound: Final = object_value(matches[0]["body"]) + expected_model: Final = provider_model.removeprefix("gemini/") if route == "acreate_interaction" else model + assert all( + outbound.get(field) == (expected_model if field == "model" else request_body[field]) for field in fields + ), matches + if route == "avideo_edit": + assert outbound.get("video") == {"id": "video-audit"}, matches + + +def test_anthropic_messages_uses_deployment_max_tokens_default(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model, identity, _handle = _register( + scenario, + "anthropic_messages", + deployment_params={"max_tokens": 32}, + ) + response: Final = _post( + gateway, + "/v1/messages", + {"model": model, "messages": [{"role": "user", "content": "default max tokens"}]}, + ) + assert response.status_code == 200, response.text + _assert_scripted_response("anthropic_messages", JSON_OBJECT.validate_python(response.json())) + observations: Final = _Observations(gateway.upstream_url) + captured: Final = eventually( + observations.read, + lambda _items: any(item.get("method") == "POST" for item in observations.for_scenario(identity)), + seconds=20, + ) + provider_requests: Final = tuple( + item for item in observations.for_scenario(identity) if item.get("method") == "POST" + ) + assert len(provider_requests) == 1, captured + outbound: Final = object_value(provider_requests[0]["body"]) + assert outbound.get("max_tokens") == 32, provider_requests + + +def test_anthropic_messages_explicit_null_reaches_upstream(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model, identity, _handle = _register( + scenario, + "anthropic_messages", + deployment_params={"max_tokens": 32}, + provider_model="openai/gpt-4o-mini", + response_route="acompletion", + ) + response: Final = _post( + gateway, + "/v1/messages", + {"model": model, "messages": [{"role": "user", "content": "null max tokens"}], "max_tokens": None}, + ) + assert response.status_code == 200, response.text + observations: Final = _Observations(gateway.upstream_url) + captured: Final = eventually( + observations.read, + lambda _items: any(item.get("method") == "POST" for item in observations.for_scenario(identity)), + seconds=20, + ) + provider_requests: Final = tuple( + item for item in observations.for_scenario(identity) if item.get("method") == "POST" + ) + assert len(provider_requests) == 1, captured + + +def test_image_generation_null_prompt_reaches_upstream(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model, identity, _handle = _register(scenario, "aimage_generation") + response: Final = _post(gateway, "/v1/images/generations", {"model": model, "prompt": None}) + assert response.status_code == 200, response.text + _assert_scripted_response("aimage_generation", JSON_OBJECT.validate_python(response.json())) + observations: Final = _Observations(gateway.upstream_url) + captured: Final = eventually( + observations.read, + lambda _items: any(item.get("method") == "POST" for item in observations.for_scenario(identity)), + seconds=20, + ) + provider_requests: Final = tuple( + item for item in observations.for_scenario(identity) if item.get("method") == "POST" + ) + assert len(provider_requests) == 1, captured + outbound: Final = object_value(provider_requests[0]["body"]) + assert "prompt" in outbound and outbound["prompt"] is None, provider_requests + + +def test_valid_agent_creation_reaches_upstream(gateway: Gateway) -> None: + body: Final = _ROUTES["acreate_agent"][3] + with gateway.scenario() as scenario: + identity: Final = f"audit-acreate_agent-{uuid.uuid4().hex}" + handle: Final = register_scenario(identity, _response("acreate_agent")) + scenario.cleanups.callback(delete_scenario, handle) + request_body: Final = { + **body, + "litellm_params_template": {"api_base": handle.api_base(), "api_key": identity}, + } + response: Final = gateway.request("POST", "/v1beta/agents", request_body) + assert response.status_code == 200, response.text + caller: Final = JSON_OBJECT.validate_python(response.json()) + _assert_scripted_response("acreate_agent", caller) + matches: Final = _observed(_Observations(gateway.upstream_url), identity) + outbound: Final = object_value(matches[0]["body"]) + assert outbound == { + "name": "agent", + "base_agent": "waverunner", + "instructions": "You are a helpful assistant.", + }, matches + + +def test_valid_speech_returns_binary_audio_and_reaches_upstream(gateway: Gateway) -> None: + template, _error_route, fields, body, _provider_model = _ROUTES["aspeech"] + with gateway.scenario() as scenario: + model, identity, _handle = _register(scenario, "aspeech") + request_body: Final = {**body, "model": model} + response: Final = _post(gateway, template, request_body) + assert response.status_code == 200, response.text + assert response.headers.get("content-type") == "audio/mpeg" + assert response.content == b"\x00" * 16 + matches: Final = _observed(_Observations(gateway.upstream_url), identity) + outbound: Final = object_value(matches[0]["body"]) + assert all(outbound.get(field) == request_body[field] for field in fields), matches + assert outbound.get("voice") == request_body["voice"], matches + + +def test_valid_video_character_request_reaches_upstream(gateway: Gateway) -> None: + template, _error_route, _fields, _body, _provider_model = _ROUTES["avideo_create_character"] + with gateway.scenario() as scenario: + model, identity, _handle = _register(scenario, "avideo_create_character") + path: Final = _path(template, model, "", "") + response: Final = gateway.request_multipart( + path, + {"name": "character", "model": model}, + {"video": ("character.mp4", b"scripted-video", "video/mp4")}, + ) + assert response.status_code == 200, response.text + _assert_scripted_response("avideo_create_character", JSON_OBJECT.validate_python(response.json())) + matches: Final = _observed(_Observations(gateway.upstream_url), identity) + outbound: Final = object_value(matches[0]["body"]) + assert outbound.get("video") == { + "filename": "character.mp4", + "content_type": "video/mp4", + }, matches + assert outbound.get("name") == "character", matches + + +def test_valid_container_file_upload_reaches_upstream(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model, identity, _handle = _register(scenario, "aupload_container_file") + container_response: Final = gateway.request("POST", "/v1/containers", {"model": model, "name": "container"}) + assert container_response.status_code == 200, container_response.text + container: Final = object_value(JSON_OBJECT.validate_python(container_response.json())) + assert container.get("object") == "container", container + container_id: Final = string_value(container["id"]) + response: Final = gateway.request_multipart( + f"/v1/containers/{container_id}/files", + {}, + {"file": ("notes.txt", b"container file contents", "text/plain")}, + ) + assert response.status_code == 200, response.text + _assert_scripted_response("aupload_container_file", JSON_OBJECT.validate_python(response.json())) + matches: Final = _observed(_Observations(gateway.upstream_url), identity, expected_count=2) + create_request: Final = object_value(matches[0]["body"]) + upload_request: Final = object_value(matches[1]["body"]) + assert str(matches[0]["path"]).endswith("/containers"), matches + assert create_request == {"name": "container"}, matches + assert str(matches[1]["path"]).endswith("/containers/container-audit/files"), matches + assert upload_request == { + "file": { + "filename": "notes.txt", + "content_type": "text/plain", + } + }, matches + + +@pytest.mark.parametrize( + ("route", "expected_text"), + ( + pytest.param("acompletion", "streamed response", id="chat-completions"), + pytest.param("anthropic_messages", "streamed response", id="anthropic-messages"), + pytest.param("aresponses", "streamed response", id="responses"), + ), +) +def test_valid_streaming_required_fields_reach_upstream( + gateway: Gateway, + route: str, + expected_text: str, +) -> None: + template, _error_route, fields, body, _provider_model = _ROUTES[route] + with gateway.scenario() as scenario: + model, identity, _handle = _register(scenario, route, streaming=True) + request_body: Final = { + **body, + **( + {"messages": [{"role": "user", "content": f"stream-{identity}"}]} + if route in {"acompletion", "anthropic_messages"} + else {} + ), + **({"input": f"response-{identity}"} if route == "aresponses" else {}), + "model": model, + "stream": True, + } + headers: Final = {"Authorization": f"Bearer {gateway.key}"} + with gateway.client.stream("POST", template, json=request_body, headers=headers) as response: + assert response.status_code == 200, response.read().decode() + lines: Final = tuple(response.iter_lines()) + assert _assembled_stream_text(route, lines) == expected_text, lines + if route == "anthropic_messages": + events: Final = _stream_event_names(lines) + payloads: Final = _stream_event_payloads(lines) + assert events[-1:] == ("message_stop",), lines + assert payloads and payloads[-1].get("type") == "message_stop", lines + matches: Final = _observed(_Observations(gateway.upstream_url), identity) + outbound: Final = object_value(matches[0]["body"]) + assert outbound.get("stream") is True, matches + assert all(outbound.get(field) == request_body[field] for field in fields), matches + + +@pytest.mark.parametrize("client_kind", ("sync", "async"), ids=("sync", "async")) +def test_openai_sdk_missing_moderations_input_returns_bad_request(gateway: Gateway, client_kind: str) -> None: + with gateway.scenario() as scenario: + model, identity, _handle = _register(scenario, "amoderation") + _missing_response(gateway, "/v1/moderations", {"model": model}, _error("/moderations", "input")) + if client_kind == "sync": + with OpenAI(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0) as client: + with pytest.raises(BadRequestError) as raised: + client.post("/v1/moderations", body={"model": model}, cast_to=httpx.Response) + else: + + async def request() -> None: + async with AsyncOpenAI( + base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0 + ) as client: + await client.post("/v1/moderations", body={"model": model}, cast_to=httpx.Response) + + with pytest.raises(BadRequestError) as raised: + asyncio.run(request()) + assert raised.value.status_code == 400 + assert raised.value.response.json() == _error("/moderations", "input") + observations: Final = _Observations(gateway.upstream_url) + observations.read() + assert observations.provider_calls(identity) == () + + +def test_default_search_model_uses_query_without_model(gateway: Gateway, tmp_path: Path) -> None: + with gateway.scenario() as scenario: + identity: Final = f"default-search-{uuid.uuid4().hex}" + handle: Final = register_scenario(identity, _response("asearch")) + scenario.cleanups.callback(delete_scenario, handle) + config: Final = tmp_path / "search.yaml" + config.write_text( + json.dumps( + { + "model_list": [], + "general_settings": {"completion_model": "exa-search"}, + "search_tools": [ + { + "search_tool_name": "exa-search", + "litellm_params": { + "search_provider": "exa_ai", + "api_key": identity, + "api_base": handle.api_base(), + }, + } + ], + } + ), + 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) + query: Final = f"query-{uuid.uuid4().hex}" + response: Final = eventually( + lambda: _post(candidate, "/v1/search", {"query": query}), + lambda result: not (result.status_code == 400 and "Invalid model name" in result.text), + seconds=30, + ) + assert ( + response.status_code == 200 and JSON_OBJECT.validate_python(response.json()).get("object") == "search" + ), response.text + outbound: Final = object_value(_observed(_Observations(gateway.upstream_url), identity)[0]["body"]) + assert outbound.get("query") == query, outbound + + +def test_interaction_without_model_uses_completion_model(gateway: Gateway, tmp_path: Path) -> None: + with gateway.scenario() as scenario: + identity: Final = f"default-interaction-{uuid.uuid4().hex}" + handle: Final = register_scenario(identity, _response("acreate_interaction")) + scenario.cleanups.callback(delete_scenario, handle) + config: Final = tmp_path / "interaction.yaml" + config.write_text( + json.dumps( + { + "model_list": [ + { + "model_name": "interaction-default", + "litellm_params": { + "model": "gemini/gemini-2.5-flash", + "api_base": handle.api_base(), + "api_key": identity, + }, + } + ], + "general_settings": {"completion_model": "interaction-default"}, + } + ), + encoding="utf-8", + ) + with owned_proxy_process(gateway, tmp_path, {}, config=config) as owned: + candidate: Final = Gateway(owned.gateway.client, owned.gateway.key, gateway.upstream_url) + request_input: Final = f"interaction-{identity}" + response: Final = _post(candidate, "/interactions", {"input": request_input}) + assert response.status_code == 200, response.text + _assert_scripted_response("acreate_interaction", JSON_OBJECT.validate_python(response.json())) + observations: Final = _Observations(gateway.upstream_url) + eventually( + observations.read, + lambda _items: any(item.get("method") == "POST" for item in observations.for_scenario(identity)), + seconds=20, + ) + upstream_calls: Final = tuple( + item for item in observations.for_scenario(identity) if item.get("method") == "POST" + ) + assert len(upstream_calls) == 1, upstream_calls + outbound: Final = object_value(upstream_calls[0]["body"]) + assert outbound.get("input") == request_input, outbound + assert outbound.get("model") == "gemini-2.5-flash", outbound + + +def test_promptless_image_edit_reaches_upstream(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + identity: Final = f"image-edit-{uuid.uuid4().hex}" + handle: Final = register_scenario(identity, _response("aimage_generation")) + scenario.cleanups.callback(delete_scenario, handle) + model: Final = scenario.model(model="openai/gpt-image-1", api_base=handle.api_base(), api_key=identity) + response: Final = eventually( + lambda: gateway.client.post( + "/v1/images/edits", + data={"model": model}, + files={"image": ("audit.png", b"png", "image/png")}, + headers={"Authorization": f"Bearer {gateway.key}"}, + ), + lambda result: not (result.status_code == 400 and "Invalid model name" in result.text), + seconds=30, + ) + assert response.status_code == 200, response.text + outbound: Final = object_value(_observed(_Observations(gateway.upstream_url), identity)[0]["body"]) + assert "prompt" not in outbound and any(field in outbound for field in ("image", "image[]")), outbound + + +def _healthy(url: str) -> int: + try: + return httpx.get(f"{url}/health", timeout=2, trust_env=False).status_code + except httpx.TransportError: + return 0 + + +@contextmanager +def _upstream(directory: Path) -> Iterator[tuple[subprocess.Popen[bytes], str]]: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + port: Final = int(reserve.getsockname()[1]) + root: Final = Path(__file__).resolve().parents[2] + url: Final = f"http://127.0.0.1:{port}" + with (directory / "upstream.log").open("w") as log: + process: Final = subprocess.Popen( + [sys.executable, "-P", "-m", "integration._support.upstream", "--port", str(port)], + cwd=root, + env={**os.environ, "PYTHONPATH": str(root)}, + stdout=log, + stderr=subprocess.STDOUT, + start_new_session=True, + ) + try: + eventually(lambda: _healthy(url), lambda status: status == 200, seconds=30) + yield process, url + finally: + if process.poll() is None: + process.send_signal(signal.SIGCONT) + process.terminate() + process.wait(timeout=10) + + +def _register_owned(url: str, identity: str, scripted: JsonResponse) -> ScenarioHandle: + result: Final = httpx.post( + f"{url}/__scenarios", + json={"scenario_id": identity, "response": scripted.model_dump(mode="json")}, + timeout=10, + trust_env=False, + ) + result.raise_for_status() + return ScenarioHandle(identity, url) + + +def _delete_owned(handle: ScenarioHandle) -> None: + httpx.delete( + f"{handle.control_url}/__scenarios/{handle.scenario_id}", timeout=10, trust_env=False + ).raise_for_status() + + +def _workers(process: subprocess.Popen[bytes]) -> tuple[psutil.Process, ...]: + return tuple(psutil.Process(process.pid).children(recursive=True)) + + +def _process_tree_line(process: psutil.Process) -> str: + try: + return f"{process.pid} {' '.join(process.cmdline())}" + except psutil.Error: + return f"{process.pid} " + + +def _worker_alive(worker: psutil.Process) -> bool: + try: + return worker.is_running() and worker.status() != psutil.STATUS_ZOMBIE + except psutil.NoSuchProcess: + return False + + +def _chat_body(model: str, marker: str, valid: bool) -> dict[str, JsonValue]: + return {"model": model, "user": marker, **({"messages": [{"role": "user", "content": marker}]} if valid else {})} + + +def _spend_rows(request_id: str) -> list[dict[str, JsonValue]]: + return read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (request_id,)) + + +def _owned_model( + scenario: Scenario, + url: str, + route: str, + script: JsonResponse, +) -> tuple[str, str, ScenarioHandle]: + identity: Final = f"chaos-{route}-{uuid.uuid4().hex}" + handle: Final = _register_owned(url, identity, script) + scenario.cleanups.callback(_delete_owned, handle) + model: Final = scenario.model(model=_ROUTES[route][4], api_base=handle.api_base(), api_key=identity) + return model, identity, handle + + +def _missing_call( + route: str, + parameter: str, + model: str, + identity: str, +) -> tuple[str, dict[str, JsonValue], dict[str, JsonValue], str]: + template, error_route, _fields, valid_body, _provider_model = _ROUTES[route] + body: Final = {key: value for key, value in {**valid_body, "model": model}.items() if key != parameter} + return _path(template, model, "", ""), body, _error(error_route, parameter), identity + + +def test_upstream_pause_and_worker_kill_preserve_required_body_status( + gateway: Gateway, + tmp_path: Path, + record_property: pytest.RecordProperty, +) -> None: + with ( + _upstream(tmp_path) as (upstream, url), + owned_proxy_process( + gateway, + tmp_path, + {"INTEGRATION_UPSTREAM_URL": url}, + workers=2, + ) as owned, + ): + candidate: Final = Gateway(owned.gateway.client, owned.gateway.key, url) + with candidate.scenario() as scenario: + script: Final = JsonResponse(content_type="application/json", body=_BODIES["acompletion"]) + chat_model, chat_identity, _chat_handle = _owned_model(scenario, url, "acompletion", script) + probe: Final = eventually( + lambda: _post(candidate, "/v1/chat/completions", _chat_body(chat_model, "probe", True)), + lambda response: response.status_code == 200, + seconds=30, + ) + assert probe.status_code == 200, probe.text + process_root: Final = psutil.Process(owned.process.pid) + processes: Final = (process_root, *process_root.children(recursive=True)) + process_tree: Final = "\n".join(_process_tree_line(process) for process in processes) + record_property("owned_proxy_process_tree", process_tree) + workers: Final = eventually( + lambda: _workers(owned.process), + lambda children: len(children) >= 2, + seconds=30, + ) + observations: Final = _Observations(url) + missing_routes: Final = ( + ("aspeech", "input"), + ("aspeech", "input"), + ("amoderation", "input"), + ("amoderation", "input"), + ("aimage_generation", "prompt"), + ("aimage_generation", "prompt"), + ("atext_completion", "prompt"), + ("atext_completion", "prompt"), + ("arerank", "query"), + ("arerank", "documents"), + ) + missing_models: Final = tuple( + _owned_model(scenario, url, route, script) for route, _parameter in missing_routes + ) + missing_calls: Final = tuple( + _missing_call(route, parameter, model, identity) + for (route, parameter), (model, identity, _handle) in zip(missing_routes, missing_models) + ) + markers: Final = tuple(f"burst-{uuid.uuid4().hex}" for _ in range(20)) + valid_calls: Final = tuple( + ("/v1/chat/completions", _chat_body(chat_model, marker, True), marker) for marker in markers + ) + upstream.send_signal(signal.SIGSTOP) + try: + with ThreadPoolExecutor(max_workers=10) as pool: + missing_futures: Final = tuple( + pool.submit(_post, candidate, path, body) for path, body, _expected, _identity in missing_calls + ) + missing: Final = tuple( + (call, future.result(timeout=15)) for call, future in zip(missing_calls, missing_futures) + ) + assert all( + response.status_code == 400 and response.json() == expected + for (_path, _body, expected, _identity), response in missing + ), [response.text for _call, response in missing] + paused_statuses: Final = tuple(response.status_code for _call, response in missing) + record_property( + "chaos_paused_missing_status_counts", + str({status: paused_statuses.count(status) for status in sorted(set(paused_statuses))}), + ) + finally: + upstream.send_signal(signal.SIGCONT) + with ThreadPoolExecutor(max_workers=20) as pool: + valid_futures: Final = tuple( + pool.submit(_post, candidate, path, body) for path, body, _marker in valid_calls + ) + valid: Final = tuple(future.result(timeout=30) for future in valid_futures) + assert all(response.status_code == 200 for response in valid), [response.text for response in valid] + record_property("chaos_burst_size", len(missing_calls) + len(valid_calls)) + resumed_statuses: Final = tuple(response.status_code for response in valid) + record_property( + "chaos_resumed_valid_status_counts", + str({status: resumed_statuses.count(status) for status in sorted(set(resumed_statuses))}), + ) + eventually( + observations.read, + lambda _items: all( + sum( + object_value(item["body"]).get("user") == marker + for item in observations.for_scenario(chat_identity) + ) + == 1 + for _path, _body, marker in valid_calls + ), + seconds=30, + ) + missing_observations: Final = { + identity: len(observations.provider_calls(identity)) + for _path, _body, _expected, identity in missing_calls + } + assert all(count == 0 for count in missing_observations.values()), missing_observations + record_property( + "chaos_missing_split", + str(tuple(f"{route}:{parameter}" for route, parameter in missing_routes)), + ) + record_property("chaos_missing_upstream_provider_call_counts", str(missing_observations)) + request_ids: Final = tuple(str(JSON_OBJECT.validate_python(response.json())["id"]) for response in valid) + spend_rows: Final = tuple( + eventually(partial(_spend_rows, request_id), lambda values: len(values) == 1, seconds=60) + for request_id in request_ids + ) + assert all(rows[0]["request_id"] == request_id for rows, request_id in zip(spend_rows, request_ids)), ( + spend_rows + ) + record_property("chaos_spend_query", 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = %s') + record_property("chaos_spend_response_id_count", len(request_ids)) + record_property("chaos_spend_row_counts", str(tuple(len(rows) for rows in spend_rows))) + workers[0].kill() + eventually(lambda: _worker_alive(workers[0]), lambda alive: not alive, seconds=10) + post_kill: Final = tuple( + _post(candidate, path, body) for path, body, _expected, _identity in missing_calls[:5] + ) + assert all( + response.status_code == 400 and response.json() == expected + for response, (_path, _body, expected, _identity) in zip(post_kill, missing_calls[:5]) + ), [response.text for response in post_kill] + recovered: Final = _post(candidate, "/v1/chat/completions", _chat_body(chat_model, "recovered", True)) + assert recovered.status_code == 200, recovered.text diff --git a/tests/integration/management/test_vector_store_config_ownership.py b/tests/integration/management/test_vector_store_config_ownership.py index e4ea4e324ff..24bea397ced 100644 --- a/tests/integration/management/test_vector_store_config_ownership.py +++ b/tests/integration/management/test_vector_store_config_ownership.py @@ -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: diff --git a/tests/integration/providers/test_provider_lookup_status.py b/tests/integration/providers/test_provider_lookup_status.py new file mode 100644 index 00000000000..2bc28d2b648 --- /dev/null +++ b/tests/integration/providers/test_provider_lookup_status.py @@ -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") diff --git a/tests/unit/llms/custom_httpx/test_llm_http_handler.py b/tests/unit/llms/custom_httpx/test_llm_http_handler.py index f3332cb513c..d283cc6c64c 100644 --- a/tests/unit/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_llm_http_handler.py @@ -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 diff --git a/tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py b/tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py index 3dcfede92ea..b19678ffb59 100644 --- a/tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py +++ b/tests/unit/proxy/google_endpoints/test_google_endpoint_routing.py @@ -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") diff --git a/tests/unit/proxy/image_endpoints/test_endpoints.py b/tests/unit/proxy/image_endpoints/test_endpoints.py index ad0901e9eee..f4aebecc11e 100644 --- a/tests/unit/proxy/image_endpoints/test_endpoints.py +++ b/tests/unit/proxy/image_endpoints/test_endpoints.py @@ -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] diff --git a/tests/unit/proxy/management_endpoints/test_common_daily_activity.py b/tests/unit/proxy/management_endpoints/test_common_daily_activity.py index 8808b73f89d..e6a6680d3e4 100644 --- a/tests/unit/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/unit/proxy/management_endpoints/test_common_daily_activity.py @@ -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 diff --git a/tests/unit/proxy/search_endpoints/__init__.py b/tests/unit/proxy/search_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/search_endpoints/test_endpoints.py b/tests/unit/proxy/search_endpoints/test_endpoints.py new file mode 100644 index 00000000000..bd6460e3dfb --- /dev/null +++ b/tests/unit/proxy/search_endpoints/test_endpoints.py @@ -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" diff --git a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py index 5dfd2f57ca6..292b2904546 100644 --- a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py +++ b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py @@ -600,6 +600,13 @@ def test_fallback_login_has_no_deprecation_banner(client_no_auth): assert " 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" diff --git a/tests/unit/proxy/vector_store_endpoints/test_vector_store_rbac.py b/tests/unit/proxy/vector_store_endpoints/test_vector_store_rbac.py index b5164ca61df..87cfddd1ae3 100644 --- a/tests/unit/proxy/vector_store_endpoints/test_vector_store_rbac.py +++ b/tests/unit/proxy/vector_store_endpoints/test_vector_store_rbac.py @@ -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"] == [] diff --git a/tests/unit/test_model_block_unblock.py b/tests/unit/test_model_block_unblock.py index da63ed4a95a..7045cd77439 100644 --- a/tests/unit/test_model_block_unblock.py +++ b/tests/unit/test_model_block_unblock.py @@ -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", diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index eaa38532f24..404fe4fa6f6 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -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 diff --git a/tests/unit/test_video_generation.py b/tests/unit/test_video_generation.py index 5c1d0bfa884..a1e5a335fd5 100644 --- a/tests/unit/test_video_generation.py +++ b/tests/unit/test_video_generation.py @@ -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