mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(proxy): return 4xx instead of 500 for missing required params, invalid pagination and unknown ids (#43787)
* fix(proxy): return 400 instead of 500 for missing required params and invalid pagination Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci: run search_endpoints tests in proxy-endpoints shard Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): return 4xx for missing required params across all LLM routes and propagate provider status on lookups Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(llm_http_handler): keep provider error text when re-raising mapped errors Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): allow promptless image edits and default search models Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): default missing image edit image to None Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(proxy): build image edit defaults without mutating request data Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): keep image edit defaults within type-discipline budget and give request mocks a scope Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): inject a fake router for the search default model test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(llms): cover provider error status on vector store and file lookup handlers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(llms): keep the lookup handler raise block to a single statement Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(llms): cover provider error status on eval, eval run, skill and vector store file content lookups Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): cover missing required body params and provider lookup status codes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci: run tests/unit/proxy/search_endpoints in the proxy-endpoints shard Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): bind spend-row request id with partial to satisfy B023 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): only reject non-positive page_size on vector store list Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): remove unreachable fine-tuning body validation Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): cover streaming anthropic messages reaching the upstream Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): count only provider calls when asserting missing params never reach the upstream Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): preserve merge-base request compatibility Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): preserve interaction completion model defaults Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): retry model read-through before rejecting params a DB-only deployment may default Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: shivam <shivam@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: yucheng <yucheng@berri.ai>
This commit is contained in:
parent
8efb4a21f6
commit
6d8434f940
28 changed files with 2788 additions and 64 deletions
1
.github/workflows/test-unit.yml
vendored
1
.github/workflows/test-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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_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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
1570
tests/integration/compatibility/test_missing_body_param_status.py
Normal file
1570
tests/integration/compatibility/test_missing_body_param_status.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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:
|
||||
|
|
|
|||
274
tests/integration/providers/test_provider_lookup_status.py
Normal file
274
tests/integration/providers/test_provider_lookup_status.py
Normal file
|
|
@ -0,0 +1,274 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Final, Literal
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value
|
||||
from integration._support.process import owned_proxy_process
|
||||
from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario
|
||||
from openai import APIStatusError, AsyncOpenAI, OpenAI
|
||||
from pydantic import JsonValue
|
||||
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import encode_file_id_with_model
|
||||
from litellm.types.videos.utils import encode_video_id_with_provider
|
||||
from tests.integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse
|
||||
|
||||
|
||||
def _provider_error(status: int) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"error": {
|
||||
"message": f"scripted provider status {status}",
|
||||
"type": "rate_limit_error"
|
||||
if status == 429
|
||||
else "server_error"
|
||||
if status >= 500
|
||||
else "invalid_request_error",
|
||||
"code": str(status),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
class _ObservationBuffer:
|
||||
def __init__(self, upstream_url: str) -> None:
|
||||
self._url = upstream_url.rstrip("/")
|
||||
self._items: tuple[dict[str, JsonValue], ...] = ()
|
||||
|
||||
def read(self) -> tuple[dict[str, JsonValue], ...]:
|
||||
with httpx.Client(timeout=10, trust_env=False) as client:
|
||||
payload: Final = JSON_OBJECT.validate_python(client.get(f"{self._url}/__observations").json())
|
||||
requests: Final = payload.get("requests")
|
||||
assert isinstance(requests, list)
|
||||
self._items = (*self._items, *(object_value(item) for item in requests if isinstance(item, dict)))
|
||||
return self._items
|
||||
|
||||
def route(self, scenario_id: str, suffix: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
return tuple(
|
||||
item
|
||||
for item in self._items
|
||||
if f"/{scenario_id}/" in str(item.get("path")) and str(item.get("path")).endswith(suffix)
|
||||
)
|
||||
|
||||
|
||||
def _ready(gateway: Gateway, model: str) -> None:
|
||||
eventually(
|
||||
lambda: gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": [{"role": "user", "content": "model readiness"}]},
|
||||
),
|
||||
lambda response: not (response.status_code == 400 and "Invalid model name" in response.text),
|
||||
seconds=30,
|
||||
)
|
||||
|
||||
|
||||
def _assert_observation(gateway: Gateway, scenario_id: str, suffix: str) -> tuple[dict[str, JsonValue], ...]:
|
||||
buffer: Final = _ObservationBuffer(gateway.upstream_url)
|
||||
result: Final = eventually(
|
||||
buffer.read,
|
||||
lambda _items: len(buffer.route(scenario_id, suffix)) >= 1,
|
||||
seconds=20,
|
||||
)
|
||||
observations: Final = buffer.route(scenario_id, suffix)
|
||||
assert observations, result
|
||||
return observations
|
||||
|
||||
|
||||
def _add_vector_store(
|
||||
gateway: Gateway,
|
||||
scenario: Scenario,
|
||||
vector_store_id: str,
|
||||
alias: str,
|
||||
handle: ScenarioHandle,
|
||||
) -> None:
|
||||
created: Final = gateway.request(
|
||||
"POST",
|
||||
"/vector_store/new",
|
||||
{
|
||||
"vector_store_id": vector_store_id,
|
||||
"custom_llm_provider": "openai",
|
||||
"litellm_params": {"model": alias, "api_base": handle.api_base(), "api_key": handle.scenario_id},
|
||||
},
|
||||
)
|
||||
assert created.status_code == 200, created.text
|
||||
scenario.cleanups.callback(gateway.post, "/vector_store/delete", {"vector_store_id": vector_store_id})
|
||||
|
||||
|
||||
def _assert_provider_error(response: httpx.Response, status: int) -> None:
|
||||
assert response.status_code == status, response.text
|
||||
body: Final = JSON_OBJECT.validate_python(response.json())
|
||||
error: Final = object_value(body["error"])
|
||||
assert str(error.get("code")) == str(status), response.text
|
||||
assert "scripted provider status" in str(error.get("message")), response.text
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"status", (400, 401, 404, 429, 500), ids=("bad-request", "unauthorized", "not-found", "rate-limit", "server-error")
|
||||
)
|
||||
def test_vector_store_lookup_preserves_provider_status(gateway: Gateway, status: int) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
vector_store_id: Final = f"vs-{uuid.uuid4().hex}"
|
||||
handle: Final = register_scenario(
|
||||
f"vector-{uuid.uuid4().hex}",
|
||||
RoutedResponse(
|
||||
content_type="application/x-routed",
|
||||
routes={
|
||||
f"GET /vector_stores/{vector_store_id}": JsonResponse(
|
||||
content_type="application/json",
|
||||
status=status,
|
||||
body=_provider_error(status),
|
||||
)
|
||||
},
|
||||
),
|
||||
)
|
||||
scenario.cleanups.callback(delete_scenario, handle)
|
||||
alias: Final = scenario.model(api_base=handle.api_base(), api_key=handle.scenario_id)
|
||||
_add_vector_store(gateway, scenario, vector_store_id, alias, handle)
|
||||
_ready(gateway, alias)
|
||||
response: Final = gateway.request("GET", f"/v1/vector_stores/{vector_store_id}")
|
||||
_assert_provider_error(response, status)
|
||||
_assert_observation(gateway, handle.scenario_id, f"/vector_stores/{vector_store_id}")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("provider_route", "path", "model_name", "status"),
|
||||
(
|
||||
("GET /videos/video-id", "/v1/videos/video-id", "openai/gpt-4o-mini", 404),
|
||||
("GET /v1/evals/eval-id", "/v1/evals/eval-id", "openai/gpt-4o-mini", 404),
|
||||
("GET /v1/skills/skill-id", "/v1/skills/skill-id?beta=true", "anthropic/claude-3-5-haiku-20241022", 404),
|
||||
("GET /v1/batch/jobs/batch-id", "/v1/batches/batch-id", "mistral/mistral-large-latest", 404),
|
||||
("GET /v1/messages/batches/batch-id", "/v1/batches/batch-id", "anthropic/claude-3-5-haiku-20241022", 500),
|
||||
),
|
||||
ids=("video", "eval", "skill", "mistral-batch", "anthropic-batch-gap"),
|
||||
)
|
||||
def test_model_scoped_lookup_returns_scripted_provider_404(
|
||||
gateway: Gateway,
|
||||
provider_route: str,
|
||||
path: str,
|
||||
model_name: str,
|
||||
status: int,
|
||||
) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
route_path: Final = provider_route.partition(" ")[2]
|
||||
handle: Final = register_scenario(
|
||||
f"scoped-{uuid.uuid4().hex}",
|
||||
RoutedResponse(
|
||||
content_type="application/x-routed",
|
||||
routes={
|
||||
provider_route: JsonResponse(content_type="application/json", status=404, body=_provider_error(404))
|
||||
},
|
||||
),
|
||||
)
|
||||
scenario.cleanups.callback(delete_scenario, handle)
|
||||
alias: Final = scenario.model(model=model_name, api_base=handle.api_base(), api_key=handle.scenario_id)
|
||||
_ready(gateway, alias)
|
||||
request_path: Final = (
|
||||
f"/v1/videos/{encode_video_id_with_provider('video-id', 'openai', model_id=alias)}"
|
||||
if "/videos/" in route_path
|
||||
else f"/v1/batches/{encode_file_id_with_model('batch-id', alias, id_type='batch')}"
|
||||
if "/batches/" in route_path or "/batch/jobs/" in route_path
|
||||
else path
|
||||
)
|
||||
headers: Final = {"x-litellm-model": alias} if "/skills/" in request_path else {}
|
||||
params: Final = {"model": alias} if "/evals/" in request_path else None
|
||||
response: Final = gateway.request("GET", request_path, params=params, headers=headers)
|
||||
assert response.status_code == status, response.text
|
||||
body: Final = JSON_OBJECT.validate_python(response.json())
|
||||
message: Final = str(object_value(body["error"]).get("message"))
|
||||
assert (
|
||||
"Client error '404 Not Found'" in message and "/v1/messages/batches/batch-id" in message
|
||||
if status == 500
|
||||
else "scripted provider status 404" in message
|
||||
), response.text
|
||||
suffix: Final = "/videos/video-id" if "/videos/" in route_path else route_path
|
||||
_assert_observation(gateway, handle.scenario_id, suffix)
|
||||
|
||||
|
||||
def _sdk_file_lookup(gateway: Gateway, file_id: str) -> httpx.Response:
|
||||
with OpenAI(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0) as client:
|
||||
try:
|
||||
client.files.retrieve(file_id)
|
||||
except APIStatusError as error:
|
||||
return error.response
|
||||
pytest.fail("OpenAI SDK file lookup unexpectedly succeeded")
|
||||
|
||||
|
||||
async def _async_sdk_file_lookup(gateway: Gateway, file_id: str) -> httpx.Response:
|
||||
async with AsyncOpenAI(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0) as client:
|
||||
try:
|
||||
await client.files.retrieve(file_id)
|
||||
except APIStatusError as error:
|
||||
return error.response
|
||||
pytest.fail("OpenAI async SDK file lookup unexpectedly succeeded")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("client_kind", ("sync", "async"), ids=("sync", "async"))
|
||||
def test_openai_sdk_file_lookup_returns_head_provider_error(
|
||||
gateway: Gateway,
|
||||
client_kind: Literal["sync", "async"],
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
scenario_id: Final = f"d12-file-lookup-{client_kind}"
|
||||
handle: Final = register_scenario(
|
||||
scenario_id,
|
||||
RoutedResponse(
|
||||
content_type="application/x-routed",
|
||||
routes={
|
||||
"GET /files/file-id": JsonResponse(
|
||||
content_type="application/json",
|
||||
status=404,
|
||||
body=_provider_error(404),
|
||||
)
|
||||
},
|
||||
),
|
||||
)
|
||||
scenario.cleanups.callback(delete_scenario, handle)
|
||||
model: Final = f"audit-file-lookup-{uuid.uuid4().hex}"
|
||||
config: Final = tmp_path / f"d12-{client_kind}.yaml"
|
||||
config.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": model,
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_base": handle.api_base(),
|
||||
"api_key": scenario_id,
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
with owned_proxy_process(
|
||||
gateway,
|
||||
tmp_path,
|
||||
{},
|
||||
config=config,
|
||||
workers=2,
|
||||
) as owned:
|
||||
candidate: Final = Gateway(owned.gateway.client, owned.gateway.key, gateway.upstream_url)
|
||||
file_id: Final = encode_file_id_with_model("file-id", model)
|
||||
response: Final = (
|
||||
_sdk_file_lookup(candidate, file_id)
|
||||
if client_kind == "sync"
|
||||
else asyncio.run(_async_sdk_file_lookup(candidate, file_id))
|
||||
)
|
||||
expected: Final = {
|
||||
"error": {
|
||||
"message": f"Error code: 404 - {_provider_error(404)}",
|
||||
"type": "invalid_request_error",
|
||||
"param": None,
|
||||
"code": "404",
|
||||
}
|
||||
}
|
||||
assert response.status_code == 404, response.text
|
||||
assert response.json() == expected, response.text
|
||||
_assert_observation(candidate, scenario_id, "/files/file-id")
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
0
tests/unit/proxy/search_endpoints/__init__.py
Normal file
0
tests/unit/proxy/search_endpoints/__init__.py
Normal file
54
tests/unit/proxy/search_endpoints/test_endpoints.py
Normal file
54
tests/unit/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"
|
||||
|
|
@ -600,6 +600,13 @@ def test_fallback_login_has_no_deprecation_banner(client_no_auth):
|
|||
assert "<form" in html
|
||||
|
||||
|
||||
def test_text_completion_without_a_prompt_returns_400_naming_prompt(client_no_auth):
|
||||
response = client_no_auth.post("/v1/completions", json={"model": "vllm_embed_model"})
|
||||
|
||||
assert response.status_code == 400, response.text
|
||||
assert response.json()["error"]["param"] == "prompt", response.text
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"ui_logo_path",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -1,8 +1,6 @@
|
|||
|
||||
import pytest
|
||||
|
||||
|
||||
|
||||
from typing import Final
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
|
|
@ -14,14 +12,14 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError, route_reque
|
|||
@pytest.mark.parametrize(
|
||||
"route_type, required_body_params",
|
||||
[
|
||||
("atext_completion", {}),
|
||||
("atext_completion", {"prompt": "Hello"}),
|
||||
("acompletion", {"messages": [{"role": "user", "content": "Hello"}]}),
|
||||
("aembedding", {"input": "Hello"}),
|
||||
("aimage_generation", {}),
|
||||
("aspeech", {}),
|
||||
("atranscription", {}),
|
||||
("amoderation", {}),
|
||||
("arerank", {}),
|
||||
("aimage_generation", {"prompt": "a cat"}),
|
||||
("aspeech", {"input": "Hello"}),
|
||||
("atranscription", {"file": b"audio"}),
|
||||
("amoderation", {"input": "Hello"}),
|
||||
("arerank", {"query": "Hello", "documents": ["hi"]}),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -253,7 +251,7 @@ async def test_route_request_no_model_required():
|
|||
|
||||
for route_type in test_cases:
|
||||
# Test data without model parameter
|
||||
data = {"input": "test input", "api_key": "test-key"}
|
||||
data = {"input": "test input", "query": "test query", "api_key": "test-key"}
|
||||
|
||||
llm_router = MagicMock()
|
||||
getattr(llm_router, route_type).return_value = "fake_response"
|
||||
|
|
@ -284,6 +282,7 @@ async def test_route_request_no_model_required_with_router_settings():
|
|||
# Test data with model parameter (it will be ignored for these route types)
|
||||
data = {
|
||||
"input": "test input",
|
||||
"query": "test query",
|
||||
"model": "test-model", # Include dummy model to avoid KeyError
|
||||
}
|
||||
|
||||
|
|
@ -1045,17 +1044,44 @@ async def test_route_request_override_enable_tag_filtering_beats_body_value():
|
|||
("aembedding", "input", "/embeddings"),
|
||||
("aresponses", "input", "/responses"),
|
||||
("acreate_batch", "input_file_id", "/batches"),
|
||||
("aspeech", "input", "/audio/speech"),
|
||||
("amoderation", "input", "/moderations"),
|
||||
("aimage_generation", "prompt", "/image/generations"),
|
||||
("asearch", "query", "/search"),
|
||||
("atext_completion", "prompt", "/completions"),
|
||||
("atranscription", "file", "/audio/transcriptions"),
|
||||
("arerank", "query", "/rerank"),
|
||||
("acompact_responses", "input", "/responses/compact"),
|
||||
("anthropic_messages", "messages", "anthropic_messages"),
|
||||
("agenerate_content", "contents", "agenerate_content"),
|
||||
("aocr", "document", "/ocr"),
|
||||
("avector_store_search", "query", "avector_store_search"),
|
||||
("avector_store_file_create", "file_id", "avector_store_file_create"),
|
||||
("avector_store_file_update", "attributes", "avector_store_file_update"),
|
||||
("avideo_generation", "prompt", "/videos"),
|
||||
("avideo_remix", "prompt", "/videos/{video_id}/remix"),
|
||||
("avideo_edit", "prompt", "/videos/edits"),
|
||||
("avideo_extension", "prompt", "/videos/extensions"),
|
||||
("avideo_create_character", "name", "/videos/characters"),
|
||||
("acreate_container", "name", "/containers"),
|
||||
("aupload_container_file", "file", "/containers/{container_id}/files"),
|
||||
("acreate_agent", "name", "/v1beta/agents"),
|
||||
("acreate_eval", "data_source_config", "/evals"),
|
||||
("acreate_run", "data_source", "/evals/{eval_id}/runs"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize("data_extra", [{}, {"messages": None, "input": None, "input_file_id": None}])
|
||||
def test_raise_if_required_body_param_missing_rejects_missing_param(route_type, param, route, data_extra):
|
||||
def test_raise_if_required_body_param_missing_rejects_missing_param(route_type, param, route):
|
||||
from litellm.proxy.route_llm_request import (
|
||||
ProxyMissingRequiredParamError,
|
||||
raise_if_required_body_param_missing,
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
raise_if_required_body_param_missing(route_type=route_type, data={"model": "gpt-4o", **data_extra})
|
||||
raise_if_required_body_param_missing(
|
||||
route_type=route_type,
|
||||
data={"model": "gpt-4o"},
|
||||
llm_router=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.code == "400"
|
||||
assert exc_info.value.param == param
|
||||
|
|
@ -1079,22 +1105,149 @@ def test_raise_if_required_body_param_missing_names_first_missing_batch_param(da
|
|||
)
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
raise_if_required_body_param_missing(route_type="acreate_batch", data=data)
|
||||
raise_if_required_body_param_missing(route_type="acreate_batch", data=data, llm_router=None)
|
||||
|
||||
assert exc_info.value.param == param
|
||||
|
||||
|
||||
def test_raise_if_required_body_param_missing_rejects_null_for_merge_base_route() -> None:
|
||||
from litellm.proxy.route_llm_request import (
|
||||
ProxyMissingRequiredParamError,
|
||||
raise_if_required_body_param_missing,
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
raise_if_required_body_param_missing(
|
||||
route_type="aembedding",
|
||||
data={"model": "text-embedding-3-small", "input": None},
|
||||
llm_router=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.param == "input"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("route_type", "data"),
|
||||
(
|
||||
pytest.param(
|
||||
"anthropic_messages",
|
||||
{"model": "claude", "messages": [], "max_tokens": None},
|
||||
id="anthropic-max-tokens",
|
||||
),
|
||||
pytest.param(
|
||||
"aimage_generation",
|
||||
{"model": "gpt-image-1", "prompt": None},
|
||||
id="image-prompt",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_required_present_body_param_accepts_explicit_null(route_type: str, data: dict[str, object]) -> None:
|
||||
from litellm.proxy.route_llm_request import raise_if_required_body_param_missing
|
||||
|
||||
raise_if_required_body_param_missing(route_type=route_type, data=data, llm_router=None)
|
||||
|
||||
|
||||
def test_required_present_body_param_uses_router_deployment_default() -> None:
|
||||
import litellm
|
||||
from litellm.proxy.route_llm_request import raise_if_required_body_param_missing
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "claude-default",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_key": "test-key",
|
||||
"max_tokens": 32,
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
raise_if_required_body_param_missing(
|
||||
route_type="anthropic_messages",
|
||||
data={"model": "claude-default", "messages": []},
|
||||
llm_router=router,
|
||||
)
|
||||
|
||||
|
||||
def test_required_present_body_param_without_router_default_still_raises() -> None:
|
||||
import litellm
|
||||
from litellm.proxy.route_llm_request import (
|
||||
ProxyMissingRequiredParamError,
|
||||
raise_if_required_body_param_missing,
|
||||
)
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "claude-without-default",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_key": "test-key",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
raise_if_required_body_param_missing(
|
||||
route_type="anthropic_messages",
|
||||
data={"model": "claude-without-default", "messages": []},
|
||||
llm_router=router,
|
||||
)
|
||||
|
||||
assert exc_info.value.param == "max_tokens"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route_type, data, param",
|
||||
[
|
||||
("arerank", {"model": "rerank-model", "query": "hi"}, "documents"),
|
||||
("anthropic_messages", {"model": "claude", "messages": []}, "max_tokens"),
|
||||
("avideo_extension", {"model": "sora-2", "prompt": "longer"}, "seconds"),
|
||||
("avideo_create_character", {"name": "hero"}, "video"),
|
||||
("acreate_eval", {"data_source_config": {"type": "custom"}}, "testing_criteria"),
|
||||
("acreate_interaction", {"input": "hi"}, "model"),
|
||||
("acreate_interaction", {"model": None, "agent": None, "input": "hi"}, "model"),
|
||||
("acreate_interaction", {"model": "gemini-3-pro-preview"}, "input"),
|
||||
],
|
||||
)
|
||||
def test_raise_if_required_body_param_missing_names_each_missing_param(route_type, data, param):
|
||||
from litellm.proxy.route_llm_request import (
|
||||
ProxyMissingRequiredParamError,
|
||||
raise_if_required_body_param_missing,
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
raise_if_required_body_param_missing(route_type=route_type, data=data, llm_router=None)
|
||||
|
||||
assert exc_info.value.code == "400"
|
||||
assert exc_info.value.param == param
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route_type, data",
|
||||
[
|
||||
("acompletion", {"model": "gpt-4o", "messages": [{"role": "user", "content": "hi"}]}),
|
||||
("acompletion", {"model": "gpt-4o", "messages": []}),
|
||||
("atext_completion", {"model": "gpt-4o"}),
|
||||
("atext_completion", {"model": "gpt-4o", "prompt": "hi"}),
|
||||
("aembedding", {"model": "text-embedding-3-small", "input": "hi"}),
|
||||
("aresponses", {"model": "gpt-4o", "input": "hi"}),
|
||||
("aresponses", {"model": "gpt-4o", "input": []}),
|
||||
("arerank", {"model": "rerank-model"}),
|
||||
("aimage_generation", {"model": "dall-e-3"}),
|
||||
("arerank", {"model": "rerank-model", "query": "hi", "documents": ["hello"]}),
|
||||
("aimage_edit", {"model": "gpt-image-1", "image": b"png", "prompt": "a hat"}),
|
||||
("aimage_edit", {"model": "stability.stable-image-remove-background-v1:0", "image": b"png"}),
|
||||
("aimage_edit", {"model": "stability.stable-style-transfer-v1:0", "init_image": b"png"}),
|
||||
("anthropic_messages", {"model": "claude", "messages": [], "max_tokens": 16}),
|
||||
("avideo_extension", {"model": "sora-2", "prompt": "longer", "seconds": "4"}),
|
||||
("acreate_eval", {"data_source_config": {"type": "custom"}, "testing_criteria": []}),
|
||||
("acreate_interaction", {"model": "gemini-3-pro-preview", "input": "hi"}),
|
||||
("acreate_interaction", {"agent": "deep-research", "input": "hi"}),
|
||||
("aimage_generation", {"model": "gpt-image-1", "prompt": "a cat"}),
|
||||
("aspeech", {"model": "gpt-4o-mini-tts", "input": "hi", "voice": "alloy"}),
|
||||
("amoderation", {"model": "omni-moderation-latest", "input": ""}),
|
||||
("asearch", {"model": "perplexity-search", "query": "litellm"}),
|
||||
(
|
||||
"acreate_batch",
|
||||
{"input_file_id": "file-abc", "endpoint": "/v1/chat/completions", "completion_window": "24h"},
|
||||
|
|
@ -1104,7 +1257,7 @@ def test_raise_if_required_body_param_missing_names_first_missing_batch_param(da
|
|||
def test_raise_if_required_body_param_missing_allows_valid_requests(route_type, data):
|
||||
from litellm.proxy.route_llm_request import raise_if_required_body_param_missing
|
||||
|
||||
raise_if_required_body_param_missing(route_type=route_type, data=data)
|
||||
raise_if_required_body_param_missing(route_type=route_type, data=data, llm_router=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1257,6 +1410,66 @@ async def test_route_request_read_through_disabled_without_store_model_in_db(mon
|
|||
|
||||
assert table.find_many_wheres == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_read_through_supplies_db_model_default_for_missing_param(monkeypatch):
|
||||
import litellm
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
model_name = "e2e-db-only-max-tokens-default"
|
||||
router = litellm.Router(
|
||||
model_list=[{"model_name": "some-other-model", "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake"}}]
|
||||
)
|
||||
db_row = SimpleNamespace(
|
||||
model_id=f"{model_name}-id",
|
||||
model_name=model_name,
|
||||
litellm_params={"model": "anthropic/claude-sonnet-4-5", "api_key": "fake", "max_tokens": 64},
|
||||
model_info={},
|
||||
blocked=False,
|
||||
)
|
||||
fake_prisma, table = _fake_prisma_client_with_models([db_row])
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", fake_prisma)
|
||||
monkeypatch.setattr(proxy_server, "store_model_in_db", True)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
data = {"model": model_name, "messages": [{"role": "user", "content": "hi"}]}
|
||||
|
||||
with patch.object(router, "anthropic_messages", new=AsyncMock(return_value="db_default_used")) as spy:
|
||||
response = await (await route_request(data, router, None, "anthropic_messages"))
|
||||
|
||||
assert response == "db_default_used"
|
||||
spy.assert_called_once()
|
||||
assert table.find_many_wheres[0] == {"model_name": model_name}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_missing_param_for_unknown_model_still_400s_after_read_through(monkeypatch):
|
||||
import litellm
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
|
||||
|
||||
model_name = "e2e-unknown-model-missing-max-tokens"
|
||||
router = litellm.Router(
|
||||
model_list=[{"model_name": "some-other-model", "litellm_params": {"model": "openai/gpt-4o", "api_key": "fake"}}]
|
||||
)
|
||||
fake_prisma, table = _fake_prisma_client_with_models([])
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", fake_prisma)
|
||||
monkeypatch.setattr(proxy_server, "store_model_in_db", True)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
await route_request(
|
||||
{"model": model_name, "messages": [{"role": "user", "content": "hi"}]},
|
||||
router,
|
||||
None,
|
||||
"anthropic_messages",
|
||||
)
|
||||
|
||||
assert (exc_info.value.code, exc_info.value.param) == ("400", "max_tokens")
|
||||
assert table.find_many_wheres[0] == {"model_name": model_name}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_routing_group_name_passes_model_gate():
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
|
@ -1325,3 +1538,23 @@ def test_proxy_model_not_found_error_keeps_the_raw_model_only_in_the_client_resp
|
|||
assert raw_model in error.detail["error"]
|
||||
assert raw_model not in error.spend_log_error_message
|
||||
assert error.spend_log_error_message.startswith("/chat/completions: Invalid model name passed in")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_without_model_on_model_routed_endpoint_is_a_400():
|
||||
import litellm
|
||||
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
|
||||
|
||||
router = litellm.Router(
|
||||
model_list=[
|
||||
{"model_name": "rerank-model", "litellm_params": {"model": "cohere/rerank-v3.5", "api_key": "fake"}}
|
||||
]
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
await route_request(
|
||||
data={"query": "hi", "documents": ["hello"]}, llm_router=router, user_model=None, route_type="arerank"
|
||||
)
|
||||
|
||||
assert exc_info.value.code == "400"
|
||||
assert exc_info.value.param == "model"
|
||||
|
|
|
|||
|
|
@ -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"] == []
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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