mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge e0eb130ddc into 632b69b5c8
This commit is contained in:
commit
ec0071831d
22 changed files with 632 additions and 29 deletions
1
.github/workflows/test-unit.yml
vendored
1
.github/workflows/test-unit.yml
vendored
|
|
@ -148,6 +148,7 @@ jobs:
|
|||
tests/test_litellm/proxy/response_api_endpoints
|
||||
tests/test_litellm/proxy/image_endpoints
|
||||
tests/test_litellm/proxy/ocr_endpoints
|
||||
tests/test_litellm/proxy/search_endpoints
|
||||
tests/test_litellm/proxy/vector_store_endpoints
|
||||
tests/test_litellm/proxy/agent_endpoints
|
||||
tests/test_litellm/proxy/a2a
|
||||
|
|
|
|||
|
|
@ -3910,6 +3910,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=provider_config,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=batch_response, provider_config=provider_config)
|
||||
return provider_config.transform_retrieve_batch_response(
|
||||
model=model,
|
||||
raw_response=batch_response,
|
||||
|
|
@ -4067,6 +4068,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=provider_config,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=batch_response, provider_config=provider_config)
|
||||
return provider_config.transform_retrieve_batch_response(
|
||||
model=model,
|
||||
raw_response=batch_response,
|
||||
|
|
@ -4484,6 +4486,7 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=provider_config)
|
||||
return provider_config.transform_retrieve_file_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -4540,6 +4543,7 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=provider_config)
|
||||
return provider_config.transform_retrieve_file_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -4732,6 +4736,7 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=provider_config)
|
||||
files_per_page: Final = self._files_per_listing_page(
|
||||
response, provider_config, logging_obj, litellm_params, headers, sync_httpx_client, timeout
|
||||
)
|
||||
|
|
@ -4789,6 +4794,7 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=provider_config)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=provider_config)
|
||||
files_per_page: Final = self._files_per_async_listing_page(
|
||||
response, provider_config, logging_obj, litellm_params, headers, async_httpx_client, timeout
|
||||
)
|
||||
|
|
@ -5925,6 +5931,38 @@ class BaseLLMHTTPHandler:
|
|||
|
||||
return None
|
||||
|
||||
def _raise_for_provider_error_status(
|
||||
self,
|
||||
response: httpx.Response,
|
||||
provider_config: Union[
|
||||
BaseConfig,
|
||||
BaseRerankConfig,
|
||||
BaseResponsesAPIConfig,
|
||||
BaseImageEditConfig,
|
||||
BaseImageGenerationConfig,
|
||||
BaseVectorStoreConfig,
|
||||
BaseVectorStoreFilesConfig,
|
||||
BaseGoogleGenAIGenerateContentConfig,
|
||||
BaseAnthropicMessagesConfig,
|
||||
BaseBatchesConfig,
|
||||
BaseVideoConfig,
|
||||
BaseSearchConfig,
|
||||
BaseTextToSpeechConfig,
|
||||
BaseSkillsAPIConfig,
|
||||
"BasePassthroughConfig",
|
||||
"BaseContainerConfig",
|
||||
BaseEvalsAPIConfig,
|
||||
BaseRealtimeHTTPConfig,
|
||||
],
|
||||
) -> None:
|
||||
if not httpx.codes.is_error(response.status_code):
|
||||
return
|
||||
raise provider_config.get_error_class(
|
||||
error_message=response.text,
|
||||
status_code=response.status_code,
|
||||
headers=response.headers,
|
||||
)
|
||||
|
||||
def _handle_error(
|
||||
self,
|
||||
e: Exception,
|
||||
|
|
@ -5962,7 +6000,7 @@ class BaseLLMHTTPHandler:
|
|||
if error_headers is None and error_response:
|
||||
error_headers = getattr(error_response, "headers", None)
|
||||
if error_response and hasattr(error_response, "text"):
|
||||
error_text = getattr(error_response, "text", error_text)
|
||||
error_text = getattr(error_response, "text", None) or error_text
|
||||
if error_headers:
|
||||
error_headers = dict(error_headers)
|
||||
else:
|
||||
|
|
@ -7337,6 +7375,7 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
# Transform the response using the provider config
|
||||
self._raise_for_provider_error_status(response=response, provider_config=video_content_provider_config)
|
||||
return video_content_provider_config.transform_video_content_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -7415,6 +7454,7 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
# Transform the response using the provider config
|
||||
self._raise_for_provider_error_status(response=response, provider_config=video_content_provider_config)
|
||||
return await video_content_provider_config.async_transform_video_content_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -8388,6 +8428,7 @@ class BaseLLMHTTPHandler:
|
|||
params=params,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=video_list_provider_config)
|
||||
return video_list_provider_config.transform_video_list_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -8569,6 +8610,7 @@ class BaseLLMHTTPHandler:
|
|||
headers=headers,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=video_status_provider_config)
|
||||
return video_status_provider_config.transform_video_status_retrieve_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -8659,6 +8701,7 @@ class BaseLLMHTTPHandler:
|
|||
url=url,
|
||||
headers=headers,
|
||||
)
|
||||
self._raise_for_provider_error_status(response=response, provider_config=video_status_provider_config)
|
||||
return await video_status_provider_config.async_transform_video_status_retrieve_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -10156,6 +10199,7 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=vector_store_provider_config)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=vector_store_provider_config)
|
||||
return vector_store_provider_config.transform_create_vector_store_response(
|
||||
response=response,
|
||||
)
|
||||
|
|
@ -10220,6 +10264,7 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=vector_store_provider_config)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=vector_store_provider_config)
|
||||
return vector_store_provider_config.transform_create_vector_store_response(
|
||||
response=response,
|
||||
)
|
||||
|
|
@ -10286,6 +10331,7 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=vector_store_provider_config)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=vector_store_provider_config)
|
||||
return response.json()
|
||||
|
||||
def vector_store_list_handler(
|
||||
|
|
@ -10364,6 +10410,7 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=vector_store_provider_config)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=vector_store_provider_config)
|
||||
return response.json()
|
||||
|
||||
async def async_vector_store_update_handler(
|
||||
|
|
@ -10832,6 +10879,7 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=vector_store_files_provider_config)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=vector_store_files_provider_config)
|
||||
return vector_store_files_provider_config.transform_list_vector_store_files_response(response=response)
|
||||
|
||||
def vector_store_file_list_handler(
|
||||
|
|
@ -10908,6 +10956,7 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=vector_store_files_provider_config)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=vector_store_files_provider_config)
|
||||
return vector_store_files_provider_config.transform_list_vector_store_files_response(response=response)
|
||||
|
||||
async def async_vector_store_file_retrieve_handler(
|
||||
|
|
@ -10967,6 +11016,7 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=vector_store_files_provider_config)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=vector_store_files_provider_config)
|
||||
return vector_store_files_provider_config.transform_retrieve_vector_store_file_response(response=response)
|
||||
|
||||
def vector_store_file_retrieve_handler(
|
||||
|
|
@ -11037,6 +11087,7 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=vector_store_files_provider_config)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=vector_store_files_provider_config)
|
||||
return vector_store_files_provider_config.transform_retrieve_vector_store_file_response(response=response)
|
||||
|
||||
async def async_vector_store_file_content_handler(
|
||||
|
|
@ -11096,6 +11147,7 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=vector_store_files_provider_config)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=vector_store_files_provider_config)
|
||||
return vector_store_files_provider_config.transform_retrieve_vector_store_file_content_response(
|
||||
response=response
|
||||
)
|
||||
|
|
@ -11168,6 +11220,7 @@ class BaseLLMHTTPHandler:
|
|||
except Exception as e:
|
||||
raise self._handle_error(e=e, provider_config=vector_store_files_provider_config)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=vector_store_files_provider_config)
|
||||
return vector_store_files_provider_config.transform_retrieve_vector_store_file_content_response(
|
||||
response=response
|
||||
)
|
||||
|
|
@ -12128,6 +12181,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=skills_api_provider_config,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=skills_api_provider_config)
|
||||
return skills_api_provider_config.transform_list_skills_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -12175,6 +12229,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=skills_api_provider_config,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=skills_api_provider_config)
|
||||
return skills_api_provider_config.transform_list_skills_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -12231,6 +12286,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=skills_api_provider_config,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=skills_api_provider_config)
|
||||
return skills_api_provider_config.transform_get_skill_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -12276,6 +12332,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=skills_api_provider_config,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=skills_api_provider_config)
|
||||
return skills_api_provider_config.transform_get_skill_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -12546,6 +12603,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=evals_api_provider_config,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config)
|
||||
return evals_api_provider_config.transform_list_evals_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -12593,6 +12651,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=evals_api_provider_config,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config)
|
||||
return evals_api_provider_config.transform_list_evals_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -12649,6 +12708,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=evals_api_provider_config,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config)
|
||||
return evals_api_provider_config.transform_get_eval_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -12694,6 +12754,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=evals_api_provider_config,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config)
|
||||
return evals_api_provider_config.transform_get_eval_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -13171,6 +13232,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=evals_api_provider_config,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config)
|
||||
return evals_api_provider_config.transform_list_runs_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -13218,6 +13280,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=evals_api_provider_config,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config)
|
||||
return evals_api_provider_config.transform_list_runs_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -13274,6 +13337,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=evals_api_provider_config,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config)
|
||||
return evals_api_provider_config.transform_get_run_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
@ -13319,6 +13383,7 @@ class BaseLLMHTTPHandler:
|
|||
provider_config=evals_api_provider_config,
|
||||
)
|
||||
|
||||
self._raise_for_provider_error_status(response=response, provider_config=evals_api_provider_config)
|
||||
return evals_api_provider_config.transform_get_run_response(
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -1220,6 +1220,12 @@ async def get_daily_activity(
|
|||
detail={"error": "Please provide start_date and end_date"},
|
||||
)
|
||||
|
||||
if page < 1 or page_size < 1:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"page and page_size must be >= 1, got page={page}, page_size={page_size}",
|
||||
)
|
||||
|
||||
try:
|
||||
where_conditions: Final = _build_where_conditions(
|
||||
entity_id_field=entity_id_field,
|
||||
|
|
|
|||
|
|
@ -12216,6 +12216,8 @@ async def completion(
|
|||
)
|
||||
litellm_call_id: Final = request_litellm_call_id(data)
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
if isinstance(e, ProxyException):
|
||||
raise with_litellm_call_id(e, litellm_call_id)
|
||||
error_msg: Final = f"{e}"
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", error_msg),
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import asyncio
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
|
||||
import httpx
|
||||
|
|
@ -162,8 +163,38 @@ REQUIRED_BODY_PARAMS_BY_ROUTE: Final[Mapping[str, tuple[str, ...]]] = {
|
|||
"aembedding": ("input",),
|
||||
"aresponses": ("input",),
|
||||
"acreate_batch": ("input_file_id", "endpoint", "completion_window"),
|
||||
"aspeech": ("input",),
|
||||
"amoderation": ("input",),
|
||||
"aimage_generation": ("prompt",),
|
||||
"asearch": ("query",),
|
||||
"atext_completion": ("prompt",),
|
||||
"atranscription": ("file",),
|
||||
"arerank": ("query", "documents"),
|
||||
"acompact_responses": ("input",),
|
||||
"anthropic_messages": ("messages", "max_tokens"),
|
||||
"agenerate_content": ("contents",),
|
||||
"aocr": ("document",),
|
||||
"acreate_fine_tuning_job": ("training_file",),
|
||||
"avector_store_search": ("query",),
|
||||
"avector_store_file_create": ("file_id",),
|
||||
"avector_store_file_update": ("attributes",),
|
||||
"avideo_generation": ("prompt",),
|
||||
"avideo_remix": ("prompt",),
|
||||
"avideo_edit": ("prompt",),
|
||||
"avideo_extension": ("prompt", "seconds"),
|
||||
"avideo_create_character": ("name", "video"),
|
||||
"acreate_container": ("name",),
|
||||
"aupload_container_file": ("file",),
|
||||
"acreate_agent": ("name",),
|
||||
"acreate_interaction": ("input",),
|
||||
"acreate_eval": ("data_source_config", "testing_criteria"),
|
||||
"acreate_run": ("data_source",),
|
||||
}
|
||||
|
||||
REQUIRED_ONE_OF_BODY_PARAMS_BY_ROUTE: Final[Mapping[str, tuple[str, str]]] = MappingProxyType(
|
||||
{"acreate_interaction": ("model", "agent")}
|
||||
)
|
||||
|
||||
|
||||
class ProxyMissingRequiredParamError(ProxyException):
|
||||
def __init__(self, route: str, param: str):
|
||||
|
|
@ -175,11 +206,18 @@ class ProxyMissingRequiredParamError(ProxyException):
|
|||
)
|
||||
|
||||
|
||||
def raise_if_required_body_param_missing(route_type: str, data: Mapping[str, object]) -> None:
|
||||
missing_param: Final = next(
|
||||
def _find_missing_required_body_param(route_type: str, data: Mapping[str, object]) -> str | None:
|
||||
one_of_params: Final = REQUIRED_ONE_OF_BODY_PARAMS_BY_ROUTE.get(route_type)
|
||||
if one_of_params is not None and all(data.get(param) is None for param in one_of_params):
|
||||
return one_of_params[0]
|
||||
return next(
|
||||
(param for param in REQUIRED_BODY_PARAMS_BY_ROUTE.get(route_type, ()) if data.get(param) is None),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def raise_if_required_body_param_missing(route_type: str, data: Mapping[str, object]) -> None:
|
||||
missing_param: Final = _find_missing_required_body_param(route_type, data)
|
||||
if missing_param is None:
|
||||
return
|
||||
raise ProxyMissingRequiredParamError(
|
||||
|
|
@ -631,6 +669,11 @@ async def _route_request_single_attempt( # noqa: ANN202 # returns unawaited pr
|
|||
# These endpoints don't need a model, use custom_llm_provider directly
|
||||
return getattr(litellm, f"{route_type}")(**data)
|
||||
|
||||
if "model" not in data:
|
||||
raise ProxyMissingRequiredParamError(
|
||||
route=ROUTE_ENDPOINT_MAPPING.get(route_type, route_type),
|
||||
param="model",
|
||||
)
|
||||
team_model_name: Final = llm_router.map_team_model(data["model"], team_id) if team_id is not None else None
|
||||
if team_model_name is not None:
|
||||
data["model"] = team_model_name
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -369,6 +369,11 @@ async def list_vector_stores(
|
|||
- page: int - Page number for pagination (default: 1)
|
||||
- page_size: int - Number of items per page (default: 100)
|
||||
"""
|
||||
if page < 1 or page_size < 1:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"page and page_size must be >= 1, got page={page}, page_size={page_size}",
|
||||
)
|
||||
await check_feature_access_for_user(user_api_key_dict, "vector_stores")
|
||||
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
|
|
|||
|
|
@ -1205,7 +1205,7 @@ def function_setup(
|
|||
elif call_type == CallTypes.moderation.value or call_type == CallTypes.amoderation.value:
|
||||
messages = args[1] if len(args) > 1 else kwargs["input"]
|
||||
elif call_type == CallTypes.atext_completion.value or call_type == CallTypes.text_completion.value:
|
||||
messages = args[0] if len(args) > 0 else kwargs["prompt"]
|
||||
messages = args[0] if len(args) > 0 else kwargs.get("prompt")
|
||||
elif call_type == CallTypes.rerank.value or call_type == CallTypes.arerank.value:
|
||||
messages = kwargs.get("query")
|
||||
elif call_type in (CallTypes.search.value, CallTypes.asearch.value):
|
||||
|
|
|
|||
|
|
@ -1034,6 +1034,7 @@ def _raw_batches_request(body: Dict[str, Any]) -> MagicMock:
|
|||
request.url.__str__.return_value = "http://localhost/v1/batches"
|
||||
request.url.path = "/v1/batches"
|
||||
request.method = "POST"
|
||||
request.scope = {"type": "http", "path": "/v1/batches", "method": "POST"}
|
||||
request.query_params = {}
|
||||
request.headers = {"Content-Type": "application/json"}
|
||||
request.client = MagicMock()
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -115,6 +115,36 @@ async def test_get_daily_activity_order_has_id_tiebreaker():
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("page, page_size", [(0, 10), (-1, 10), (1, 0), (1, -5)])
|
||||
async def test_get_daily_activity_rejects_non_positive_pagination_with_400(page, page_size):
|
||||
from fastapi import HTTPException
|
||||
|
||||
mock_prisma = MagicMock()
|
||||
mock_table = MagicMock()
|
||||
mock_table.count = AsyncMock(return_value=0)
|
||||
mock_table.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.db.litellm_dailyteamspend = mock_table
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await get_daily_activity(
|
||||
prisma_client=mock_prisma,
|
||||
table_name="litellm_dailyteamspend",
|
||||
entity_id_field="team_id",
|
||||
entity_id=None,
|
||||
entity_metadata_field=None,
|
||||
start_date="2026-09-18",
|
||||
end_date="2026-09-25",
|
||||
model=None,
|
||||
api_key=None,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400, exc_info.value.detail
|
||||
mock_table.find_many.assert_not_called()
|
||||
|
||||
|
||||
def test_is_user_agent_tag():
|
||||
"""Test _is_user_agent_tag function."""
|
||||
# Test None and empty string
|
||||
|
|
|
|||
0
tests/test_litellm/proxy/search_endpoints/__init__.py
Normal file
0
tests/test_litellm/proxy/search_endpoints/__init__.py
Normal file
54
tests/test_litellm/proxy/search_endpoints/test_endpoints.py
Normal file
54
tests/test_litellm/proxy/search_endpoints/test_endpoints.py
Normal file
|
|
@ -0,0 +1,54 @@
|
|||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import orjson
|
||||
import pytest
|
||||
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
|
||||
from litellm.proxy.search_endpoints.endpoints import search
|
||||
|
||||
|
||||
def _json_request(body: dict[str, object]) -> MagicMock:
|
||||
request = MagicMock()
|
||||
request.body = AsyncMock(return_value=orjson.dumps(body))
|
||||
return request
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("body", [{"query": "litellm"}, {"query": "litellm", "search_tool_name": ""}])
|
||||
async def test_search_without_search_tool_name_or_model_is_a_400(body):
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
await search(
|
||||
request=_json_request(body),
|
||||
fastapi_response=MagicMock(),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
|
||||
)
|
||||
|
||||
assert exc_info.value.code == "400"
|
||||
assert exc_info.value.param == "search_tool_name"
|
||||
assert exc_info.value.message == "/search: Missing required parameter: 'search_tool_name'."
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("default_source", ["cli_model", "completion_model"])
|
||||
async def test_search_with_only_a_query_falls_back_to_the_proxy_default_model(monkeypatch, default_source):
|
||||
if default_source == "cli_model":
|
||||
monkeypatch.setattr(proxy_server, "user_model", "perplexity-search")
|
||||
else:
|
||||
monkeypatch.setitem(proxy_server.general_settings, "completion_model", "perplexity-search")
|
||||
search_result = {"object": "search", "results": []}
|
||||
router = MagicMock()
|
||||
router.asearch = AsyncMock(return_value=search_result)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
response = await search(
|
||||
request=_json_request({"query": "litellm"}),
|
||||
fastapi_response=MagicMock(),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
|
||||
)
|
||||
|
||||
assert response == search_result, response
|
||||
router.asearch.assert_awaited_once()
|
||||
assert router.asearch.await_args.kwargs["query"] == "litellm"
|
||||
assert router.asearch.await_args.kwargs["model"] == "perplexity-search"
|
||||
|
|
@ -601,6 +601,13 @@ def test_fallback_login_has_no_deprecation_banner(client_no_auth):
|
|||
assert "<form" in html
|
||||
|
||||
|
||||
def test_text_completion_without_a_prompt_returns_400_naming_prompt(client_no_auth):
|
||||
response = client_no_auth.post("/v1/completions", json={"model": "vllm_embed_model"})
|
||||
|
||||
assert response.status_code == 400, response.text
|
||||
assert response.json()["error"]["param"] == "prompt", response.text
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"ui_logo_path",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -1,8 +1,6 @@
|
|||
|
||||
import pytest
|
||||
|
||||
|
||||
|
||||
from typing import Final
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
|
@ -14,14 +12,14 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError, route_reque
|
|||
@pytest.mark.parametrize(
|
||||
"route_type, required_body_params",
|
||||
[
|
||||
("atext_completion", {}),
|
||||
("atext_completion", {"prompt": "Hello"}),
|
||||
("acompletion", {"messages": [{"role": "user", "content": "Hello"}]}),
|
||||
("aembedding", {"input": "Hello"}),
|
||||
("aimage_generation", {}),
|
||||
("aspeech", {}),
|
||||
("atranscription", {}),
|
||||
("amoderation", {}),
|
||||
("arerank", {}),
|
||||
("aimage_generation", {"prompt": "a cat"}),
|
||||
("aspeech", {"input": "Hello"}),
|
||||
("atranscription", {"file": b"audio"}),
|
||||
("amoderation", {"input": "Hello"}),
|
||||
("arerank", {"query": "Hello", "documents": ["hi"]}),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -253,7 +251,7 @@ async def test_route_request_no_model_required():
|
|||
|
||||
for route_type in test_cases:
|
||||
# Test data without model parameter
|
||||
data = {"input": "test input", "api_key": "test-key"}
|
||||
data = {"input": "test input", "query": "test query", "api_key": "test-key"}
|
||||
|
||||
llm_router = MagicMock()
|
||||
getattr(llm_router, route_type).return_value = "fake_response"
|
||||
|
|
@ -284,6 +282,7 @@ async def test_route_request_no_model_required_with_router_settings():
|
|||
# Test data with model parameter (it will be ignored for these route types)
|
||||
data = {
|
||||
"input": "test input",
|
||||
"query": "test query",
|
||||
"model": "test-model", # Include dummy model to avoid KeyError
|
||||
}
|
||||
|
||||
|
|
@ -1045,9 +1044,51 @@ async def test_route_request_override_enable_tag_filtering_beats_body_value():
|
|||
("aembedding", "input", "/embeddings"),
|
||||
("aresponses", "input", "/responses"),
|
||||
("acreate_batch", "input_file_id", "/batches"),
|
||||
("aspeech", "input", "/audio/speech"),
|
||||
("amoderation", "input", "/moderations"),
|
||||
("aimage_generation", "prompt", "/image/generations"),
|
||||
("asearch", "query", "/search"),
|
||||
("atext_completion", "prompt", "/completions"),
|
||||
("atranscription", "file", "/audio/transcriptions"),
|
||||
("arerank", "query", "/rerank"),
|
||||
("acompact_responses", "input", "/responses/compact"),
|
||||
("anthropic_messages", "messages", "anthropic_messages"),
|
||||
("agenerate_content", "contents", "agenerate_content"),
|
||||
("aocr", "document", "/ocr"),
|
||||
("acreate_fine_tuning_job", "training_file", "acreate_fine_tuning_job"),
|
||||
("avector_store_search", "query", "avector_store_search"),
|
||||
("avector_store_file_create", "file_id", "avector_store_file_create"),
|
||||
("avector_store_file_update", "attributes", "avector_store_file_update"),
|
||||
("avideo_generation", "prompt", "/videos"),
|
||||
("avideo_remix", "prompt", "/videos/{video_id}/remix"),
|
||||
("avideo_edit", "prompt", "/videos/edits"),
|
||||
("avideo_extension", "prompt", "/videos/extensions"),
|
||||
("avideo_create_character", "name", "/videos/characters"),
|
||||
("acreate_container", "name", "/containers"),
|
||||
("aupload_container_file", "file", "/containers/{container_id}/files"),
|
||||
("acreate_agent", "name", "/v1beta/agents"),
|
||||
("acreate_eval", "data_source_config", "/evals"),
|
||||
("acreate_run", "data_source", "/evals/{eval_id}/runs"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"data_extra",
|
||||
[
|
||||
{},
|
||||
{
|
||||
"messages": None,
|
||||
"input": None,
|
||||
"input_file_id": None,
|
||||
"prompt": None,
|
||||
"query": None,
|
||||
"file": None,
|
||||
"image": None,
|
||||
"contents": None,
|
||||
"document": None,
|
||||
"name": None,
|
||||
},
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize("data_extra", [{}, {"messages": None, "input": None, "input_file_id": None}])
|
||||
def test_raise_if_required_body_param_missing_rejects_missing_param(route_type, param, route, data_extra):
|
||||
from litellm.proxy.route_llm_request import (
|
||||
ProxyMissingRequiredParamError,
|
||||
|
|
@ -1084,17 +1125,54 @@ def test_raise_if_required_body_param_missing_names_first_missing_batch_param(da
|
|||
assert exc_info.value.param == param
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route_type, data, param",
|
||||
[
|
||||
("arerank", {"model": "rerank-model", "query": "hi"}, "documents"),
|
||||
("anthropic_messages", {"model": "claude", "messages": []}, "max_tokens"),
|
||||
("avideo_extension", {"model": "sora-2", "prompt": "longer"}, "seconds"),
|
||||
("avideo_create_character", {"name": "hero"}, "video"),
|
||||
("acreate_eval", {"data_source_config": {"type": "custom"}}, "testing_criteria"),
|
||||
("acreate_interaction", {"input": "hi"}, "model"),
|
||||
("acreate_interaction", {"model": None, "agent": None, "input": "hi"}, "model"),
|
||||
("acreate_interaction", {"model": "gemini-3-pro-preview"}, "input"),
|
||||
],
|
||||
)
|
||||
def test_raise_if_required_body_param_missing_names_each_missing_param(route_type, data, param):
|
||||
from litellm.proxy.route_llm_request import (
|
||||
ProxyMissingRequiredParamError,
|
||||
raise_if_required_body_param_missing,
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
raise_if_required_body_param_missing(route_type=route_type, data=data)
|
||||
|
||||
assert exc_info.value.code == "400"
|
||||
assert exc_info.value.param == param
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route_type, data",
|
||||
[
|
||||
("acompletion", {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}),
|
||||
("acompletion", {"model": "gpt-4o", "messages": []}),
|
||||
("atext_completion", {"model": "gpt-4o"}),
|
||||
("atext_completion", {"model": "gpt-4o", "prompt": "hi"}),
|
||||
("aembedding", {"model": "text-embedding-3-small", "input": "hi"}),
|
||||
("aresponses", {"model": "gpt-4o", "input": "hi"}),
|
||||
("aresponses", {"model": "gpt-4o", "input": []}),
|
||||
("arerank", {"model": "rerank-model"}),
|
||||
("aimage_generation", {"model": "dall-e-3"}),
|
||||
("arerank", {"model": "rerank-model", "query": "hi", "documents": ["hello"]}),
|
||||
("aimage_edit", {"model": "gpt-image-1", "image": b"png", "prompt": "a hat"}),
|
||||
("aimage_edit", {"model": "stability.stable-image-remove-background-v1:0", "image": b"png"}),
|
||||
("aimage_edit", {"model": "stability.stable-style-transfer-v1:0", "init_image": b"png"}),
|
||||
("anthropic_messages", {"model": "claude", "messages": [], "max_tokens": 16}),
|
||||
("avideo_extension", {"model": "sora-2", "prompt": "longer", "seconds": "4"}),
|
||||
("acreate_eval", {"data_source_config": {"type": "custom"}, "testing_criteria": []}),
|
||||
("acreate_interaction", {"model": "gemini-3-pro-preview", "input": "hi"}),
|
||||
("acreate_interaction", {"agent": "deep-research", "input": "hi"}),
|
||||
("aimage_generation", {"model": "gpt-image-1", "prompt": "a cat"}),
|
||||
("aspeech", {"model": "gpt-4o-mini-tts", "input": "hi", "voice": "alloy"}),
|
||||
("amoderation", {"model": "omni-moderation-latest", "input": ""}),
|
||||
("asearch", {"model": "perplexity-search", "query": "litellm"}),
|
||||
(
|
||||
"acreate_batch",
|
||||
{"input_file_id": "file-abc", "endpoint": "/v1/chat/completions", "completion_window": "24h"},
|
||||
|
|
@ -1257,6 +1335,7 @@ async def test_route_request_read_through_disabled_without_store_model_in_db(mon
|
|||
|
||||
assert table.find_many_wheres == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_routing_group_name_passes_model_gate():
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
|
@ -1325,3 +1404,23 @@ def test_proxy_model_not_found_error_keeps_the_raw_model_only_in_the_client_resp
|
|||
assert raw_model in error.detail["error"]
|
||||
assert raw_model not in error.spend_log_error_message
|
||||
assert error.spend_log_error_message.startswith("/chat/completions: Invalid model name passed in")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_without_model_on_model_routed_endpoint_is_a_400():
|
||||
import litellm
|
||||
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{"model_name": "rerank-model", "litellm_params": {"model": "cohere/rerank-v3.5", "api_key": "fake"}}
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
await route_request(
|
||||
data={"query": "hi", "documents": ["hello"]}, llm_router=router, user_model=None, route_type="arerank"
|
||||
)
|
||||
|
||||
assert exc_info.value.code == "400"
|
||||
assert exc_info.value.param == "model"
|
||||
|
|
|
|||
|
|
@ -135,3 +135,17 @@ async def test_list_vector_stores_admin_not_blocked():
|
|||
):
|
||||
# Must not raise any HTTPException — admin is always allowed.
|
||||
await list_vector_stores(user_api_key_dict=admin)
|
||||
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("page, page_size", [(0, 10), (-1, 10), (1, 0), (1, -5)])
|
||||
async def test_list_vector_stores_rejects_non_positive_pagination_with_400(page, page_size):
|
||||
from litellm.proxy.vector_store_endpoints.management_endpoints import (
|
||||
list_vector_stores,
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await list_vector_stores(user_api_key_dict=_make_internal_user(), page=page, page_size=page_size)
|
||||
|
||||
assert exc_info.value.status_code == 400, exc_info.value.detail
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -6479,6 +6479,11 @@ def test_function_setup_logs_the_search_query_edit_prompt_and_ocr_document_summa
|
|||
assert _logged_request_messages(original_function, *args, **kwargs) == [{"role": "user", "content": expected}]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("original_function", ("atext_completion", "text_completion"))
|
||||
def test_function_setup_without_a_prompt_leaves_the_missing_prompt_to_request_validation(original_function: str) -> None:
|
||||
assert _logged_request_messages(original_function, model="gpt-4o") is None
|
||||
|
||||
|
||||
def test_search_with_a_mixed_type_query_list_still_reaches_its_own_validation_error() -> None:
|
||||
mixed_query: Final = cast(list[str], ["Eiffel Tower", 7]) # cast-ok: the invalid list is the point of the test
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue