mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-28 01:32:17 +00:00
Merge f865286f9b into deab92408e
This commit is contained in:
commit
fea11cd136
6 changed files with 176 additions and 12 deletions
|
|
@ -207,6 +207,8 @@ async def image_generation(
|
|||
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
|
||||
)
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
if isinstance(e, ProxyException):
|
||||
raise
|
||||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e)),
|
||||
|
|
|
|||
|
|
@ -117,6 +117,8 @@ async def rerank(
|
|||
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
|
||||
)
|
||||
log_llm_api_exception(e, litellm_call_id)
|
||||
if isinstance(e, ProxyException):
|
||||
raise
|
||||
if isinstance(e, HTTPException):
|
||||
raise ProxyException(
|
||||
message=getattr(e, "message", str(e)),
|
||||
|
|
|
|||
|
|
@ -77,7 +77,7 @@ ROUTE_ENDPOINT_MAPPING: Final = {
|
|||
"acompletion": "/chat/completions",
|
||||
"atext_completion": "/completions",
|
||||
"aembedding": "/embeddings",
|
||||
"aimage_generation": "/image/generations",
|
||||
"aimage_generation": "/images/generations",
|
||||
"aspeech": "/audio/speech",
|
||||
"atranscription": "/audio/transcriptions",
|
||||
"amoderation": "/moderations",
|
||||
|
|
@ -161,6 +161,11 @@ REQUIRED_BODY_PARAMS_BY_ROUTE: Final[Mapping[str, tuple[str, ...]]] = {
|
|||
"aembedding": ("input",),
|
||||
"aresponses": ("input",),
|
||||
"acreate_batch": ("input_file_id", "endpoint", "completion_window"),
|
||||
"aresponses": ("input",),
|
||||
"arerank": ("query", "documents"),
|
||||
"aspeech": ("input",),
|
||||
"amoderation": ("input",),
|
||||
"aimage_generation": ("prompt",),
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.image_endpoints import endpoints
|
||||
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -119,6 +120,69 @@ async def test_image_generation_prompt_rerouting(monkeypatch):
|
|||
assert response.headers.get("x-callback-test") == "value"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_image_generation__missing_required_param_is_400(monkeypatch):
|
||||
"""image_generation()'s except block only special-cased HTTPException, so a
|
||||
ProxyMissingRequiredParamError (a ProxyException with code=400) fell into the
|
||||
`else` branch's `getattr(e, "status_code", 500)` and surfaced as a 500."""
|
||||
|
||||
async def fake_add_litellm_data_to_request(**kwargs):
|
||||
return kwargs["data"]
|
||||
|
||||
async def fake_pre_call_hook(*, user_api_key_dict, data, call_type):
|
||||
return data
|
||||
|
||||
async def fake_post_call_failure_hook(**_: Any) -> None:
|
||||
return None
|
||||
|
||||
async def fake_route_request(*, data, **kwargs):
|
||||
raise ProxyMissingRequiredParamError(route="/images/generations", param="prompt")
|
||||
|
||||
fake_proxy_logger = SimpleNamespace(
|
||||
pre_call_hook=fake_pre_call_hook,
|
||||
post_call_failure_hook=fake_post_call_failure_hook,
|
||||
)
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": "/v1/images/generations",
|
||||
"headers": [],
|
||||
}
|
||||
body = orjson.dumps({"model": "dall-e-3"})
|
||||
|
||||
async def receive():
|
||||
return {"type": "http.request", "body": body, "more_body": False}
|
||||
|
||||
request = Request(scope, receive)
|
||||
response = Response()
|
||||
user_api_key = UserAPIKeyAuth()
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.proxy_server.add_litellm_data_to_request",
|
||||
fake_add_litellm_data_to_request,
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", {})
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", fake_proxy_logger)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.user_model", None)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.version", "test-version")
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.image_endpoints.endpoints.route_request", fake_route_request
|
||||
)
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await endpoints.image_generation(
|
||||
request=request,
|
||||
fastapi_response=response,
|
||||
user_api_key_dict=user_api_key,
|
||||
)
|
||||
|
||||
assert exc_info.value.code == "400"
|
||||
assert exc_info.value.param == "prompt"
|
||||
|
||||
|
||||
def _image_edit_client(monkeypatch, captured: Dict[str, Any]) -> TestClient:
|
||||
class CaptureProcessing:
|
||||
def __init__(self, data: Dict[str, Any]) -> None:
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ import litellm.proxy.proxy_server as proxy_server_mod
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.rerank_endpoints.endpoints import rerank
|
||||
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
|
||||
from litellm.types.utils import RerankResponse
|
||||
|
||||
HIDDEN_PARAMS = {
|
||||
|
|
@ -209,3 +210,37 @@ async def test_a_rejection_raised_before_routing_keeps_its_own_status(monkeypatc
|
|||
error = await _rerank_failure(rejection, raised_before_routing=True, monkeypatch=monkeypatch)
|
||||
|
||||
assert (error.type, error.param, error.code) == ("bad_request_error", "session_id", "400")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rerank__missing_required_param_is_400():
|
||||
"""/rerank's except block only special-cased HTTPException, so a
|
||||
ProxyMissingRequiredParamError (a ProxyException with code=400) fell into the
|
||||
`else` branch's `getattr(e, "status_code", 500)` and surfaced as a 500."""
|
||||
fastapi_response = Response()
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=lambda **kwargs: kwargs["data"])
|
||||
proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||||
|
||||
async def fake_add_litellm_data_to_request(**kwargs):
|
||||
return kwargs["data"]
|
||||
|
||||
async def fake_route_request(**kwargs):
|
||||
raise ProxyMissingRequiredParamError(route="/rerank", param="query")
|
||||
|
||||
with (
|
||||
patch.object(proxy_server_mod, "add_litellm_data_to_request", fake_add_litellm_data_to_request), # test-quality-ok: the rerank route reads these proxy_server module globals; no injection seam on the FastAPI handler
|
||||
patch.object(proxy_server_mod, "route_request", fake_route_request), # test-quality-ok: the rerank route reads these proxy_server module globals; no injection seam on the FastAPI handler
|
||||
patch.object(proxy_server_mod, "proxy_logging_obj", proxy_logging_obj), # test-quality-ok: the rerank route reads these proxy_server module globals; no injection seam on the FastAPI handler
|
||||
patch.object(proxy_server_mod, "llm_router", MagicMock()), # test-quality-ok: the rerank route reads these proxy_server module globals; no injection seam on the FastAPI handler
|
||||
patch.object(proxy_server_mod, "version", "1.2.3"), # test-quality-ok: the rerank route reads these proxy_server module globals; no injection seam on the FastAPI handler
|
||||
):
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await rerank(
|
||||
request=_build_request(),
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
|
||||
)
|
||||
|
||||
assert exc_info.value.code == "400"
|
||||
assert exc_info.value.param == "query"
|
||||
|
|
|
|||
|
|
@ -17,11 +17,12 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError, route_reque
|
|||
("atext_completion", {}),
|
||||
("acompletion", {"messages": [{"role": "user", "content": "Hello"}]}),
|
||||
("aembedding", {"input": "Hello"}),
|
||||
("aimage_generation", {}),
|
||||
("aspeech", {}),
|
||||
("aimage_generation", {"prompt": "a cat"}),
|
||||
("aspeech", {"input": "Hello"}),
|
||||
("atranscription", {}),
|
||||
("amoderation", {}),
|
||||
("arerank", {}),
|
||||
("amoderation", {"input": "Hello"}),
|
||||
("arerank", {"query": "hi", "documents": ["doc1", "doc2"]}),
|
||||
("aresponses", {"input": "Hello"}),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1045,9 +1046,16 @@ async def test_route_request_override_enable_tag_filtering_beats_body_value():
|
|||
("aembedding", "input", "/embeddings"),
|
||||
("aresponses", "input", "/responses"),
|
||||
("acreate_batch", "input_file_id", "/batches"),
|
||||
("aresponses", "input", "/responses"),
|
||||
("aspeech", "input", "/audio/speech"),
|
||||
("amoderation", "input", "/moderations"),
|
||||
("aimage_generation", "prompt", "/images/generations"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize("data_extra", [{}, {"messages": None, "input": None, "input_file_id": None}])
|
||||
@pytest.mark.parametrize(
|
||||
"data_extra",
|
||||
[{}, {"messages": None, "input": None, "input_file_id": None, "prompt": 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,6 +1092,26 @@ def test_raise_if_required_body_param_missing_names_first_missing_batch_param(da
|
|||
assert exc_info.value.param == param
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"data, param",
|
||||
[
|
||||
({"documents": ["doc1"]}, "query"),
|
||||
({"query": "hi"}, "documents"),
|
||||
({}, "query"),
|
||||
],
|
||||
)
|
||||
def test_raise_if_required_body_param_missing_names_first_missing_rerank_param(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="arerank", data={"model": "rerank-model", **data})
|
||||
|
||||
assert exc_info.value.param == param
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"route_type, data",
|
||||
[
|
||||
|
|
@ -1093,8 +1121,10 @@ def test_raise_if_required_body_param_missing_names_first_missing_batch_param(da
|
|||
("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": ["doc1", "doc2"]}),
|
||||
("aimage_generation", {"model": "dall-e-3", "prompt": "a cat"}),
|
||||
("aspeech", {"model": "tts-1", "input": "hi"}),
|
||||
("amoderation", {"model": "omni-moderation-latest", "input": "hi"}),
|
||||
(
|
||||
"acreate_batch",
|
||||
{"input_file_id": "file-abc", "endpoint": "/v1/chat/completions", "completion_window": "24h"},
|
||||
|
|
@ -1124,17 +1154,43 @@ async def test_route_request_rejects_chat_completion_without_messages():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_request_rejects_responses_without_input():
|
||||
@pytest.mark.parametrize(
|
||||
"route_type, param",
|
||||
[
|
||||
("aresponses", "input"),
|
||||
("aspeech", "input"),
|
||||
("amoderation", "input"),
|
||||
("aimage_generation", "prompt"),
|
||||
],
|
||||
)
|
||||
async def test_route_request_rejects_missing_required_param(route_type, param):
|
||||
"""A body missing a required param used to splat into the router method and
|
||||
surface the resulting TypeError/KeyError as a 500."""
|
||||
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
|
||||
|
||||
llm_router = MagicMock()
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError) as exc_info:
|
||||
await route_request({"model": "gpt-4o"}, llm_router, None, "aresponses")
|
||||
await route_request({"model": "gpt-4o"}, llm_router, None, route_type)
|
||||
|
||||
assert exc_info.value.code == "400"
|
||||
assert exc_info.value.param == "input"
|
||||
llm_router.aresponses.assert_not_called()
|
||||
assert exc_info.value.param == param
|
||||
getattr(llm_router, route_type).assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("data", [{"model": "rerank-model"}, {"model": "rerank-model", "query": "hi"}])
|
||||
async def test_route_request_rejects_rerank_missing_required_params(data):
|
||||
"""A /rerank body missing `query`/`documents` used to splat into
|
||||
Router.arerank() and surface the resulting TypeError as a 500."""
|
||||
from litellm.proxy.route_llm_request import ProxyMissingRequiredParamError
|
||||
|
||||
llm_router = MagicMock()
|
||||
|
||||
with pytest.raises(ProxyMissingRequiredParamError):
|
||||
await route_request(data, llm_router, None, "arerank")
|
||||
|
||||
llm_router.arerank.assert_not_called()
|
||||
|
||||
|
||||
class FakeProxyModelTable:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue