This commit is contained in:
kejunleng 2026-08-27 16:38:14 -04:00 • committed by GitHub
commit d73eb81a4f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 372 additions and 273 deletions

View file

@ -1471,8 +1471,9 @@ async def _await_llm_call_cancelling_on_disconnect(
class ProxyBaseLLMRequestProcessing:
def __init__(self, data: dict):
def __init__(self, data: dict, failure_call_type: str | None = None):
self.data = data
self._failure_call_type = failure_call_type
@staticmethod
def get_custom_headers(
@ -1758,6 +1759,7 @@ class ProxyBaseLLMRequestProcessing:
model: str | None = None,
llm_router: Router | None = None,
) -> tuple[dict, LiteLLMLoggingObj]:
self._failure_call_type = route_type
start_time: Final = datetime.now() # start before calling guardrail hooks
self.data = await add_litellm_data_to_request(
@ -2202,6 +2204,7 @@ class ProxyBaseLLMRequestProcessing:
"""
Common request processing logic for both chat completions and responses API endpoints
"""
self._failure_call_type = route_type
requested_model_from_client: Final[str | None] = (
self.data.get("model") if isinstance(self.data.get("model"), str) else None
)
@ -3166,6 +3169,8 @@ class ProxyBaseLLMRequestProcessing:
"""Raises ProxyException (OpenAI API compatible) if an exception is raised"""
_log_llm_api_exception(e)
# Allow callbacks to transform the error response
if self._failure_call_type:
self.data["call_type"] = self._failure_call_type
transformed_exception: Final = await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict,
original_exception=e,

View file

@ -10,6 +10,7 @@ from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import (
_get_parent_otel_span_from_kwargs,
get_litellm_metadata_from_kwargs,
get_metadata_variable_name_from_kwargs,
)
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import guardrail_information_cost
@ -124,32 +125,24 @@ class _ProxyDBLogger(CustomLogger):
metadata=_metadata,
)
existing_metadata: Final[dict] = request_data.get("metadata", None) or {}
existing_metadata.update(_metadata)
litellm_metadata_bucket: Final = request_data.get("litellm_metadata")
if (
isinstance(litellm_metadata_bucket, dict)
and "standard_logging_guardrail_information" not in existing_metadata
):
guardrail_info: Final = litellm_metadata_bucket.get("standard_logging_guardrail_information")
if guardrail_info is not None:
existing_metadata["standard_logging_guardrail_information"] = guardrail_info
if "litellm_params" not in request_data:
request_data["litellm_params"] = {}
existing_litellm_params: Final = request_data.get("litellm_params", {})
existing_litellm_metadata: Final = existing_litellm_params.get("metadata", {}) or {}
# Preserve tags from existing metadata
if existing_litellm_metadata.get("tags"):
existing_metadata["tags"] = existing_litellm_metadata.get("tags")
existing_metadata: Final = _ProxyDBLogger._get_merged_failure_metadata(
request_data=request_data,
failure_metadata=_metadata,
)
request_data["litellm_params"]["proxy_server_request"] = (
request_data.get("proxy_server_request") or existing_litellm_params.get("proxy_server_request") or {}
)
request_data["litellm_params"]["metadata"] = existing_metadata
if _ProxyDBLogger._should_write_failure_litellm_metadata(
request_data=request_data,
litellm_params=existing_litellm_params,
):
request_data["litellm_params"]["litellm_metadata"] = dict(existing_metadata)
# Preserve model name and custom_llm_provider
if "model" not in request_data:
@ -208,6 +201,72 @@ class _ProxyDBLogger(CustomLogger):
org_id=user_api_key_dict.org_id,
)
@staticmethod
def _get_merged_failure_metadata(
request_data: dict,
failure_metadata: dict,
) -> dict:
merged_metadata: dict = {}
existing_litellm_params = request_data.get("litellm_params", {}) or {}
def merge_metadata(
metadata: Any,
*,
overwrite: bool = False,
overwrite_keys: set[str] | None = None,
skip_user_api_key_fields: bool = False,
) -> None:
if not isinstance(metadata, dict):
return
keys_to_overwrite = overwrite_keys or set()
for key, value in metadata.items():
if skip_user_api_key_fields and (key == "user_api_key" or key.startswith("user_api_key_")):
continue
if value in (None, "", {}):
continue
if not overwrite and key not in keys_to_overwrite and key in merged_metadata:
continue
merged_metadata[key] = value
merge_metadata(
request_data.get("litellm_metadata", {}),
skip_user_api_key_fields=True,
)
merge_metadata(
request_data.get("metadata", {}),
overwrite=True,
skip_user_api_key_fields=True,
)
merge_metadata(
existing_litellm_params.get("metadata", {}),
overwrite_keys={"tags"},
skip_user_api_key_fields=True,
)
merge_metadata(
existing_litellm_params.get("litellm_metadata", {}),
skip_user_api_key_fields=True,
)
merge_metadata(failure_metadata, overwrite=True)
return merged_metadata
@staticmethod
def _get_failure_metadata_variable_name(request_data: dict) -> str:
proxy_server_request = request_data.get("proxy_server_request", {}) or {}
metadata_variable_name = proxy_server_request.get("metadata_variable_name")
if metadata_variable_name in ("metadata", "litellm_metadata"):
return metadata_variable_name
return get_metadata_variable_name_from_kwargs(request_data)
@staticmethod
def _should_write_failure_litellm_metadata(
request_data: dict,
litellm_params: dict,
) -> bool:
return (
_ProxyDBLogger._get_failure_metadata_variable_name(request_data) == "litellm_metadata"
or isinstance(request_data.get("litellm_metadata"), dict)
or isinstance(litellm_params.get("litellm_metadata"), dict)
)
@log_db_metrics
async def _PROXY_track_cost_callback(
self,

View file

@ -1772,6 +1772,7 @@ async def add_litellm_data_to_request(
safe_add_api_version_from_query_params(data, request)
_metadata_variable_name: Final = _get_metadata_variable_name(request)
data["proxy_server_request"]["metadata_variable_name"] = _metadata_variable_name
if data.get(_metadata_variable_name, None) is None:
data[_metadata_variable_name] = {}

View file

@ -10531,9 +10531,11 @@ async def embeddings(
"""
global proxy_logging_obj
data: Final = await _read_request_body(request=request)
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
data: Any = {}
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data, failure_call_type="aembedding")
try:
data = await _read_request_body(request=request)
base_llm_response_processor.data = data
### HANDLE TOKEN ARRAY INPUT DECODING ###
# This must happen BEFORE base_process_llm_request() since it modifies the input
router_model_names: Final = llm_router.model_names if llm_router is not None else []

View file

@ -31,9 +31,7 @@ def patched_embedding(monkeypatch):
router.model_names = ["text-embedding-ada-002"]
router.get_deployment_by_model_group_name = MagicMock(return_value=None)
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(
proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock())
)
monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()))
async def _fake_process(self, *args, **kwargs):
return dict(HAPPY_RESPONSE)
@ -51,19 +49,17 @@ def embedding_pipeline_raises(monkeypatch):
router = MagicMock()
router.model_names = []
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(
proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock())
)
monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()))
from litellm.proxy._types import ProxyException
async def _raise(self, *args, **kwargs):
self._failure_call_type = kwargs["route_type"]
raise ValueError("boom")
async def _handler(self, *, e, user_api_key_dict, proxy_logging_obj, version=None):
return ProxyException(
message="boom", type="bad_request_error", param="model", code=400
)
assert self._failure_call_type == "aembedding"
return ProxyException(message="boom", type="bad_request_error", param="model", code=400)
monkeypatch.setattr(
common_request_processing.ProxyBaseLLMRequestProcessing,
@ -78,6 +74,58 @@ def embedding_pipeline_raises(monkeypatch):
yield
@pytest.fixture
def embedding_body_read_raises(monkeypatch):
monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()))
from litellm.proxy._types import ProxyException
async def _raise_read(*args, **kwargs):
raise ValueError("body-read")
async def _handler(self, *, e, user_api_key_dict, proxy_logging_obj, version=None):
assert self._failure_call_type == "aembedding"
return ProxyException(message="body-read", type="bad_request_error", param="body", code=400)
monkeypatch.setattr(proxy_server, "_read_request_body", _raise_read)
monkeypatch.setattr(
common_request_processing.ProxyBaseLLMRequestProcessing,
"_handle_llm_api_exception",
_handler,
)
yield
@pytest.fixture
def embedding_token_decode_raises(monkeypatch):
router = MagicMock()
router.model_names = ["text-embedding-ada-002"]
router.get_deployment_by_model_group_name = MagicMock(
return_value={"litellm_params": {"model": "custom/provider-model"}}
)
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock()))
from litellm.proxy._types import ProxyException
def _raise_decode(*args, **kwargs):
raise ValueError("token-decode")
async def _handler(self, *, e, user_api_key_dict, proxy_logging_obj, version=None):
assert self._failure_call_type == "aembedding"
assert self.data["model"] == "text-embedding-ada-002"
assert self.data["input"] == [[1, 2, 3]]
return ProxyException(message="token-decode", type="bad_request_error", param="input", code=400)
monkeypatch.setattr(proxy_server.litellm, "decode", _raise_decode)
monkeypatch.setattr(
common_request_processing.ProxyBaseLLMRequestProcessing,
"_handle_llm_api_exception",
_handler,
)
yield
_EMBED_PATHS = [
"/v1/embeddings",
"/embeddings",
@ -119,3 +167,23 @@ def test_embeddings_pipeline_error(client, auth_as, embedding_pipeline_raises, p
response = client.post(path, json=payload)
assert response.status_code == 400
assert response.content # non-empty error body
@pytest.mark.parametrize("path", _EMBED_PATHS)
def test_embeddings_body_read_error_preserves_call_type(client, auth_as, embedding_body_read_raises, path):
payload = {"model": "text-embedding-ada-002", "input": "boom"}
with auth_as():
response = client.post(path, json=payload)
assert response.status_code == 400
assert response.content
@pytest.mark.parametrize("path", _EMBED_PATHS)
def test_embeddings_preprocessing_error_preserves_parsed_data(
client, auth_as, embedding_token_decode_raises, path
):
payload = {"model": "text-embedding-ada-002", "input": [[1, 2, 3]]}
with auth_as():
response = client.post(path, json=payload)
assert response.status_code == 400
assert response.content

View file

@ -56,55 +56,49 @@ async def test_embedding_metadata_population():
Test that the embedding endpoint correctly populates metadata
from UserAPIKeyAuth.
"""
captured_data = {}
async def mock_base_process(self, *args, **kwargs):
captured_data.update(self.data)
return {"data": []}
# Setup
with patch(
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.base_process_llm_request"
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.base_process_llm_request",
new=mock_base_process,
):
# Create a mock UserAPIKeyAuth object
mock_user_auth = MagicMock(spec=UserAPIKeyAuth)
mock_user_auth.user_id = "test_user_id_emb"
mock_user_auth.team_id = "test_team_id_emb"
mock_user_auth.org_id = "test_org_id_emb"
# Create a mock Request object
mock_request = MagicMock(spec=Request)
mock_request.json = AsyncMock(
return_value={"model": "gpt-3.5-turbo", "input": "hello"}
)
# Mock _read_request_body to return our data
with patch(
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.__init__",
return_value=None,
) as mock_base_process_init:
# Create a mock UserAPIKeyAuth object
mock_user_auth = MagicMock(spec=UserAPIKeyAuth)
mock_user_auth.user_id = "test_user_id_emb"
mock_user_auth.team_id = "test_team_id_emb"
mock_user_auth.org_id = "test_org_id_emb"
# Create a mock Request object
mock_request = MagicMock(spec=Request)
mock_request.json = AsyncMock(
return_value={"model": "gpt-3.5-turbo", "input": "hello"}
"litellm.proxy.proxy_server._read_request_body",
new=AsyncMock(return_value={"model": "gpt-3.5-turbo", "input": "hello"}),
):
# Call the endpoint function directly
await embeddings(
request=mock_request,
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=mock_user_auth,
)
# Mock _read_request_body to return our data
with patch(
"litellm.proxy.proxy_server._read_request_body",
new=AsyncMock(
return_value={"model": "gpt-3.5-turbo", "input": "hello"}
),
):
# Call the endpoint function directly
await embeddings(
request=mock_request,
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=mock_user_auth,
)
# Check if ProxyBaseLLMRequestProcessing was initialized with the correct metadata
mock_base_process_init.assert_called_once()
call_args = mock_base_process_init.call_args
# handle both positional and keyword args for data
if "data" in call_args.kwargs:
data_arg = call_args.kwargs["data"]
else:
data_arg = call_args.args[0]
assert (
data_arg["metadata"]["user_api_key_user_id"] == "test_user_id_emb"
)
assert (
data_arg["metadata"]["user_api_key_team_id"] == "test_team_id_emb"
)
assert data_arg["metadata"]["user_api_key_org_id"] == "test_org_id_emb"
assert (
captured_data["metadata"]["user_api_key_user_id"]
== "test_user_id_emb"
)
assert (
captured_data["metadata"]["user_api_key_team_id"]
== "test_team_id_emb"
)
assert captured_data["metadata"]["user_api_key_org_id"] == "test_org_id_emb"
@pytest.mark.asyncio

View file

@ -126,16 +126,12 @@ class TestProxyBaseLLMRequestProcessing:
assert json.loads(result.body) == guardrailed_body
@pytest.mark.asyncio
async def test_handle_non_streaming_allm_passthrough_route_forwards_upstream_headers(
self, monkeypatch
):
async def test_handle_non_streaming_allm_passthrough_route_forwards_upstream_headers(self, monkeypatch):
"""The guardrail JSON path must forward upstream response headers (e.g.
x-amzn-requestid) alongside the x-litellm-* headers, matching the
non-guardrail passthrough path, while dropping length headers that no
longer match the rewritten body."""
processing_obj = ProxyBaseLLMRequestProcessing(
data={"custom_llm_provider": "bedrock"}
)
processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "bedrock"})
monkeypatch.setattr(
processing_obj,
"_has_post_call_guardrails_for_passthrough",
@ -175,14 +171,10 @@ class TestProxyBaseLLMRequestProcessing:
assert result.headers["content-length"] == str(len(result.body))
@pytest.mark.asyncio
async def test_handle_event_stream_allm_passthrough_route_forwards_upstream_headers(
self, monkeypatch
):
async def test_handle_event_stream_allm_passthrough_route_forwards_upstream_headers(self, monkeypatch):
"""The guardrail event-stream branch must also forward upstream response
headers alongside the x-litellm-* headers."""
processing_obj = ProxyBaseLLMRequestProcessing(
data={"custom_llm_provider": "bedrock"}
)
processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "bedrock"})
monkeypatch.setattr(
processing_obj,
"_has_post_call_guardrails_for_passthrough",
@ -224,15 +216,11 @@ class TestProxyBaseLLMRequestProcessing:
assert result.headers["x-litellm-call-id"] == "test-call-id"
@pytest.mark.asyncio
async def test_handle_non_streaming_allm_passthrough_route_applies_response_headers_hook(
self, monkeypatch
):
async def test_handle_non_streaming_allm_passthrough_route_applies_response_headers_hook(self, monkeypatch):
"""Guardrailed non-streaming passthrough responses must include headers
injected by post_call_response_headers_hook, matching the headers a
non-guardrailed passthrough response would carry."""
processing_obj = ProxyBaseLLMRequestProcessing(
data={"custom_llm_provider": "bedrock"}
)
processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "bedrock"})
monkeypatch.setattr(
processing_obj,
"_has_post_call_guardrails_for_passthrough",
@ -251,9 +239,7 @@ class TestProxyBaseLLMRequestProcessing:
return kwargs["response"]
proxy_logging_obj.post_call_success_hook = fake_post_call_success_hook
proxy_logging_obj.post_call_response_headers_hook = AsyncMock(
return_value={"x-litellm-custom": "from-hook"}
)
proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={"x-litellm-custom": "from-hook"})
result = await processing_obj._handle_non_streaming_allm_passthrough_route(
response=upstream,
@ -2804,6 +2790,63 @@ class TestHandleLLMApiExceptionDictDetail:
return raised
raise AssertionError("ProxyException was not raised")
async def test_failure_logging_receives_route_call_type(self):
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
processor = ProxyBaseLLMRequestProcessing(data={"call_type": "acompletion"})
processor._failure_call_type = "aresponses"
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test")
proxy_logging_obj = MagicMock()
proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={})
with pytest.raises(ProxyException):
await processor._handle_llm_api_exception(
e=Exception("provider failed"),
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
)
request_data = proxy_logging_obj.post_call_failure_hook.call_args.kwargs["request_data"]
assert request_data["call_type"] == "aresponses"
async def test_direct_pre_call_failure_logging_receives_route_call_type(self):
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
processor = ProxyBaseLLMRequestProcessing(data={"model": "test-model"})
mock_request = MagicMock(spec=Request)
mock_request.headers = {}
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test")
proxy_logging_obj = MagicMock()
proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=Exception("pre-call"))
proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={})
proxy_config = MagicMock(spec=ProxyConfig)
with patch(
"litellm.proxy.common_request_processing.add_litellm_data_to_request",
new=AsyncMock(return_value={"model": "test-model"}),
):
with pytest.raises(Exception, match="pre-call") as exc_info:
await processor.common_processing_pre_call_logic(
request=mock_request,
general_settings={},
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
proxy_config=proxy_config,
route_type="aresponses",
)
with pytest.raises(ProxyException):
await processor._handle_llm_api_exception(
e=exc_info.value,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
)
request_data = proxy_logging_obj.post_call_failure_hook.call_args.kwargs["request_data"]
assert request_data["call_type"] == "aresponses"
async def test_dict_detail_bedrock_shape_preserved(self):
exc = HTTPException(
status_code=400,
@ -2911,9 +2954,7 @@ class TestStreamCloseOnDisconnect:
finally:
closed.set()
response = _UpstreamClosingStreamingResponse(
body(), media_type="text/event-stream"
)
response = _UpstreamClosingStreamingResponse(body(), media_type="text/event-stream")
async def receive():
await asyncio.Event().wait()
@ -2944,9 +2985,7 @@ class TestStreamCloseOnDisconnect:
finally:
closed.set()
response = _UpstreamClosingStreamingResponse(
body(), media_type="text/event-stream"
)
response = _UpstreamClosingStreamingResponse(body(), media_type="text/event-stream")
async def receive():
await disconnected.wait()
@ -3017,9 +3056,7 @@ class TestStreamCloseOnDisconnect:
finally:
inner_closed.set()
response = await create_response(
generator=wrapped(), media_type="text/event-stream", headers={}
)
response = await create_response(generator=wrapped(), media_type="text/event-stream", headers={})
async def receive():
await asyncio.Event().wait()
@ -3215,9 +3252,7 @@ class TestStreamCloseOnDisconnect:
with pytest.raises(_ClientDisconnectedBeforeFirstChunk):
await asyncio.wait_for(
_buffer_first_chunk_honoring_disconnect(
AcloseRaises(), request=self._request_that_disconnects()
),
_buffer_first_chunk_honoring_disconnect(AcloseRaises(), request=self._request_that_disconnects()),
timeout=5,
)
@ -3233,9 +3268,7 @@ class TestStreamCloseOnDisconnect:
with pytest.raises(_ClientDisconnectedBeforeFirstChunk):
await asyncio.wait_for(
_buffer_first_chunk_honoring_disconnect(
blocking_gen(), request=self._request_that_disconnects()
),
_buffer_first_chunk_honoring_disconnect(blocking_gen(), request=self._request_that_disconnects()),
timeout=5,
)
assert closed.is_set()
@ -3251,9 +3284,7 @@ class TestHandleLLMApiExceptionRetryAfter:
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test")
proxy_logging_obj = MagicMock()
proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
proxy_logging_obj.post_call_response_headers_hook = AsyncMock(
return_value=callback_headers or {}
)
proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value=callback_headers or {})
try:
await processor._handle_llm_api_exception(
@ -3305,9 +3336,7 @@ class TestHandleLLMApiExceptionRetryAfter:
enable_pre_call_checks=False,
cooldown_list=[],
)
proxy_exc = await self._invoke(
exc, callback_headers={"retry-after": "", "x-custom": "1"}
)
proxy_exc = await self._invoke(exc, callback_headers={"retry-after": "", "x-custom": "1"})
assert proxy_exc.headers["retry-after"] == "43"
assert proxy_exc.headers["x-custom"] == "1"
@ -3503,9 +3532,7 @@ class TestDisconnectGatherCleanup:
return Request(scope={"type": "http", "headers": []}, receive=receive)
@pytest.mark.asyncio
async def test_base_process_llm_request_raises_499_on_client_disconnect(
self, monkeypatch
):
async def test_base_process_llm_request_raises_499_on_client_disconnect(self, monkeypatch):
"""With cancel_on_disconnect enabled, base_process_llm_request returns 499."""
import asyncio
@ -3534,9 +3561,7 @@ class TestDisconnectGatherCleanup:
"common_processing_pre_call_logic",
AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)),
)
monkeypatch.setattr(
processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)
)
monkeypatch.setattr(processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False))
with pytest.raises(HTTPException) as exc_info:
await processing_obj.base_process_llm_request(
@ -3554,9 +3579,7 @@ class TestDisconnectGatherCleanup:
assert "disconnected" in exc_info.value.detail.lower()
@pytest.mark.asyncio
async def test_base_process_llm_request_reraises_cancelled_error_without_client_disconnect(
self, monkeypatch
):
async def test_base_process_llm_request_reraises_cancelled_error_without_client_disconnect(self, monkeypatch):
import asyncio
import litellm.proxy.common_request_processing as cpr
@ -3581,9 +3604,7 @@ class TestDisconnectGatherCleanup:
"common_processing_pre_call_logic",
AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)),
)
monkeypatch.setattr(
processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)
)
monkeypatch.setattr(processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False))
monkeypatch.setattr(
cpr,
"route_request",
@ -3644,9 +3665,7 @@ class TestDisconnectGatherCleanup:
"common_processing_pre_call_logic",
AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)),
)
monkeypatch.setattr(
processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)
)
monkeypatch.setattr(processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False))
with pytest.raises(HTTPException):
await processing_obj.base_process_llm_request(
@ -3697,9 +3716,7 @@ class TestDisconnectGatherCleanup:
assert task.done()
@pytest.mark.asyncio
async def test_base_process_llm_request_preserves_llm_error_after_gather(
self, monkeypatch
):
async def test_base_process_llm_request_preserves_llm_error_after_gather(self, monkeypatch):
import litellm.proxy.common_request_processing as cpr
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
@ -3728,9 +3745,7 @@ class TestDisconnectGatherCleanup:
"common_processing_pre_call_logic",
AsyncMock(return_value=({"model": "gemini-2.0-flash"}, mock_logging_obj)),
)
monkeypatch.setattr(
processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False)
)
monkeypatch.setattr(processing_obj, "_has_post_call_guardrails", MagicMock(return_value=False))
mock_request = MagicMock(spec=Request)
mock_request.is_disconnected = AsyncMock(return_value=False)
@ -3767,19 +3782,13 @@ class TestStreamingClientDisconnectLogging:
"litellm_params": {"metadata": {}},
}
recorded = await _record_streaming_client_disconnect_if_needed(
mock_request, request_data
)
recorded = await _record_streaming_client_disconnect_if_needed(mock_request, request_data)
assert recorded is True
assert request_data["metadata"]["client_disconnected"] is True
assert request_data["metadata"]["error_information"]["error_code"] == "499"
assert (
request_data["metadata"]["error_information"]["error_code"] == "499"
)
assert (
mock_logging_obj.model_call_details["litellm_params"]["metadata"][
"error_information"
]["error_code"]
mock_logging_obj.model_call_details["litellm_params"]["metadata"]["error_information"]["error_code"]
== "499"
)
@ -3793,9 +3802,7 @@ class TestStreamingClientDisconnectLogging:
mock_request.is_disconnected = AsyncMock(return_value=False)
request_data = {"metadata": {}}
recorded = await _record_streaming_client_disconnect_if_needed(
mock_request, request_data
)
recorded = await _record_streaming_client_disconnect_if_needed(mock_request, request_data)
assert recorded is False
assert "client_disconnected" not in request_data["metadata"]
@ -3820,22 +3827,12 @@ class TestStreamingClientDisconnectLogging:
"litellm_params": {"metadata": {}},
}
recorded = await _record_streaming_client_disconnect_if_needed(
mock_request, request_data
)
recorded = await _record_streaming_client_disconnect_if_needed(mock_request, request_data)
assert recorded is True
assert request_data["metadata"]["client_disconnected"] is True
assert (
mock_logging_obj.model_call_details["litellm_params"]["metadata"][
"client_disconnected"
]
is True
)
assert (
mock_logging_obj.model_call_details["metadata"]["client_disconnected"]
is True
)
assert mock_logging_obj.model_call_details["litellm_params"]["metadata"]["client_disconnected"] is True
assert mock_logging_obj.model_call_details["metadata"]["client_disconnected"] is True
@pytest.mark.asyncio
async def test_record_streaming_client_disconnect_handles_none_request_data_metadata(self):
@ -3851,15 +3848,11 @@ class TestStreamingClientDisconnectLogging:
"litellm_params": {"metadata": None},
}
recorded = await _record_streaming_client_disconnect_if_needed(
mock_request, request_data
)
recorded = await _record_streaming_client_disconnect_if_needed(mock_request, request_data)
assert recorded is True
assert request_data["metadata"]["client_disconnected"] is True
assert (
request_data["litellm_params"]["metadata"]["client_disconnected"] is True
)
assert request_data["litellm_params"]["metadata"]["client_disconnected"] is True
@pytest.mark.asyncio
async def test_apply_client_disconnect_metadata_none_returns_early(self):
@ -3870,9 +3863,7 @@ class TestStreamingClientDisconnectLogging:
_apply_client_disconnect_metadata(None)
@pytest.mark.asyncio
async def test_finalize_streaming_generator_cleanup_fires_deferred_logging(
self, monkeypatch
):
async def test_finalize_streaming_generator_cleanup_fires_deferred_logging(self, monkeypatch):
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
)
@ -3904,9 +3895,7 @@ class TestStreamingClientDisconnectLogging:
assert request_data["metadata"]["error_information"]["error_code"] == "499"
@pytest.mark.asyncio
async def test_finalize_streaming_generator_cleanup_skips_disconnect_after_completion(
self, monkeypatch
):
async def test_finalize_streaming_generator_cleanup_skips_disconnect_after_completion(self, monkeypatch):
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
)
@ -3936,9 +3925,7 @@ class TestStreamingClientDisconnectLogging:
assert "client_disconnected" not in request_data["metadata"]
@pytest.mark.asyncio
async def test_async_streaming_data_generator_records_499_on_early_aclose(
self, monkeypatch
):
async def test_async_streaming_data_generator_records_499_on_early_aclose(self, monkeypatch):
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
)
@ -3953,9 +3940,7 @@ class TestStreamingClientDisconnectLogging:
yield {"choices": [{"delta": {"content": " there"}}]}
mock_proxy_logging = MagicMock(spec=ProxyLogging)
mock_proxy_logging.async_post_call_streaming_iterator_hook = (
mock_streaming_iterator
)
mock_proxy_logging.async_post_call_streaming_iterator_hook = mock_streaming_iterator
ProxyLogging._callback_capabilities_cache.clear()
mock_request = MagicMock(spec=Request)
@ -3966,9 +3951,7 @@ class TestStreamingClientDisconnectLogging:
"model": "gemini-2.0-flash",
"metadata": {},
"litellm_params": {"metadata": {}},
"litellm_logging_obj": MagicMock(
model_call_details={"metadata": {}, "litellm_params": {}}
),
"litellm_logging_obj": MagicMock(model_call_details={"metadata": {}, "litellm_params": {}}),
}
gen = ProxyBaseLLMRequestProcessing.async_streaming_data_generator(
@ -3987,6 +3970,8 @@ class TestStreamingClientDisconnectLogging:
assert request_data["metadata"]["error_information"]["error_code"] == "499"
ProxyLogging._callback_capabilities_cache.clear()
class TestCancelOnDisconnect:
"""
Coverage for the opt-in `general_settings.cancel_on_disconnect` flag:
@ -4013,23 +3998,17 @@ class TestCancelOnDisconnect:
llm_call = asyncio.get_running_loop().create_future()
disconnect_event = asyncio.Event()
await _cancel_llm_call_on_client_disconnect(
request, llm_call, disconnect_event
)
await _cancel_llm_call_on_client_disconnect(request, llm_call, disconnect_event)
assert llm_call.cancelled()
assert disconnect_event.is_set()
async def test_monitor_is_noop_while_client_stays_connected(self):
request = self._request(
[{"type": "http.request", "body": b"", "more_body": False}]
)
request = self._request([{"type": "http.request", "body": b"", "more_body": False}])
llm_call = asyncio.get_running_loop().create_future()
disconnect_event = asyncio.Event()
monitor = asyncio.create_task(
_cancel_llm_call_on_client_disconnect(request, llm_call, disconnect_event)
)
monitor = asyncio.create_task(_cancel_llm_call_on_client_disconnect(request, llm_call, disconnect_event))
await asyncio.sleep(0.01)
assert not monitor.done()
@ -4048,9 +4027,7 @@ class TestCancelOnDisconnect:
llm_call = asyncio.get_running_loop().create_future()
disconnect_event = asyncio.Event()
await _cancel_llm_call_on_client_disconnect(
request, llm_call, disconnect_event
)
await _cancel_llm_call_on_client_disconnect(request, llm_call, disconnect_event)
assert not llm_call.cancelled()
assert not disconnect_event.is_set()
@ -4065,9 +4042,7 @@ class TestCancelOnDisconnect:
with pytest.raises(asyncio.CancelledError):
await _await_llm_call_cancelling_on_disconnect(request, llm_call)
async def _drive_base_process_llm_request(
self, monkeypatch, general_settings: dict, llm_call, request: Request
):
async def _drive_base_process_llm_request(self, monkeypatch, general_settings: dict, llm_call, request: Request):
from litellm.proxy._types import UserAPIKeyAuth
logging_obj = MagicMock()
@ -4076,9 +4051,7 @@ class TestCancelOnDisconnect:
logging_obj._on_deferred_stream_complete = None
logging_obj.cost_breakdown = None
processor = ProxyBaseLLMRequestProcessing(
data={"model": "fake-model", "litellm_logging_obj": logging_obj}
)
processor = ProxyBaseLLMRequestProcessing(data={"model": "fake-model", "litellm_logging_obj": logging_obj})
proxy_logging_obj = MagicMock(spec=ProxyLogging)
proxy_logging_obj.during_call_hook = AsyncMock(return_value=None)
@ -4086,9 +4059,7 @@ class TestCancelOnDisconnect:
proxy_logging_obj.post_call_success_hook = AsyncMock(
side_effect=lambda data, user_api_key_dict, response: response
)
proxy_logging_obj.post_call_response_headers_hook = AsyncMock(
return_value=None
)
proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value=None)
async def fake_route_request(**kwargs):
return llm_call()
@ -4167,9 +4138,7 @@ class TestCancelOnDisconnect:
with pytest.raises(ProxyException) as exc_info:
await processor._handle_llm_api_exception(
e=HTTPException(
status_code=499, detail="Client disconnected the request"
),
e=HTTPException(status_code=499, detail="Client disconnected the request"),
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"),
proxy_logging_obj=proxy_logging_obj,
)
@ -4235,7 +4204,9 @@ class TestAllmPassthroughRoutePostCallGuardrails:
proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock())
monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", capture_hook)
with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True):
with patch.object(
ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True
):
processing_obj = ProxyBaseLLMRequestProcessing(data={})
result = await processing_obj._handle_non_streaming_allm_passthrough_route(
response=httpx_response,
@ -4285,7 +4256,9 @@ class TestAllmPassthroughRoutePostCallGuardrails:
proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock())
monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", non_dict_hook)
with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True):
with patch.object(
ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True
):
processing_obj = ProxyBaseLLMRequestProcessing(data={})
result = await processing_obj._handle_non_streaming_allm_passthrough_route(
response=httpx_response,
@ -4323,7 +4296,9 @@ class TestAllmPassthroughRoutePostCallGuardrails:
hook_spy = AsyncMock()
monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", hook_spy)
with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True):
with patch.object(
ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True
):
processing_obj = ProxyBaseLLMRequestProcessing(data={})
result = await processing_obj._handle_non_streaming_allm_passthrough_route(
response=httpx_response,
@ -4364,7 +4339,9 @@ class TestAllmPassthroughRoutePostCallGuardrails:
hook_spy = AsyncMock()
monkeypatch.setattr(proxy_logging_obj, "post_call_success_hook", hook_spy)
with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=False):
with patch.object(
ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=False
):
processing_obj = ProxyBaseLLMRequestProcessing(data={})
result = await processing_obj._handle_non_streaming_allm_passthrough_route(
response=httpx_response,
@ -4476,7 +4453,9 @@ class TestEventStreamAllmPassthroughRoute:
"content-length": "99",
}
with patch.object(ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True):
with patch.object(
ProxyBaseLLMRequestProcessing, "_has_post_call_guardrails_for_passthrough", return_value=True
):
processing_obj = ProxyBaseLLMRequestProcessing(data={})
result = await processing_obj._handle_non_streaming_allm_passthrough_route(
response=mock_response,
@ -4507,9 +4486,7 @@ class TestAllmPassthroughStreamingProviderGate:
de-anonymized.
"""
def _build_processing_obj(
self, custom_llm_provider: str, endpoint: str = ""
) -> ProxyBaseLLMRequestProcessing:
def _build_processing_obj(self, custom_llm_provider: str, endpoint: str = "") -> ProxyBaseLLMRequestProcessing:
logging_obj = MagicMock()
logging_obj.litellm_call_id = "call-123"
logging_obj.cost_breakdown = None
@ -4560,14 +4537,17 @@ class TestAllmPassthroughStreamingProviderGate:
processing_obj = self._build_processing_obj("anthropic")
chunks = [b"chunk-1", b"chunk-2"]
with patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails",
return_value=False,
), patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails_for_passthrough",
return_value=True,
with (
patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails",
return_value=False,
),
patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails_for_passthrough",
return_value=True,
),
):
result = await self._run(processing_obj, monkeypatch, chunks)
@ -4576,27 +4556,27 @@ class TestAllmPassthroughStreamingProviderGate:
assert streamed == chunks
@pytest.mark.asyncio
async def test_bedrock_converse_stream_is_buffered_through_handler(
self, monkeypatch
):
processing_obj = self._build_processing_obj(
"bedrock", "model/us.amazon.nova-lite-v1:0/converse-stream"
)
async def test_bedrock_converse_stream_is_buffered_through_handler(self, monkeypatch):
processing_obj = self._build_processing_obj("bedrock", "model/us.amazon.nova-lite-v1:0/converse-stream")
chunks = [b"raw-1", b"raw-2"]
with patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails",
return_value=False,
), patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails_for_passthrough",
return_value=True,
), patch(
"litellm.llms.bedrock.passthrough.guardrail_translation.handler."
"BedrockPassthroughGuardrailHandler.de_anonymize_event_stream",
new=AsyncMock(return_value=b"modified-body"),
) as mock_handler:
with (
patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails",
return_value=False,
),
patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails_for_passthrough",
return_value=True,
),
patch(
"litellm.llms.bedrock.passthrough.guardrail_translation.handler."
"BedrockPassthroughGuardrailHandler.de_anonymize_event_stream",
new=AsyncMock(return_value=b"modified-body"),
) as mock_handler,
):
result = await self._run(processing_obj, monkeypatch, chunks)
assert isinstance(result, Response)
@ -4612,19 +4592,23 @@ class TestAllmPassthroughStreamingProviderGate:
)
chunks = [b"raw-1", b"raw-2"]
with patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails",
return_value=False,
), patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails_for_passthrough",
return_value=True,
), patch(
"litellm.llms.bedrock.passthrough.guardrail_translation.handler."
"BedrockPassthroughGuardrailHandler.de_anonymize_event_stream",
new=AsyncMock(return_value=b"modified-body"),
) as mock_handler:
with (
patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails",
return_value=False,
),
patch.object(
ProxyBaseLLMRequestProcessing,
"_has_post_call_guardrails_for_passthrough",
return_value=True,
),
patch(
"litellm.llms.bedrock.passthrough.guardrail_translation.handler."
"BedrockPassthroughGuardrailHandler.de_anonymize_event_stream",
new=AsyncMock(return_value=b"modified-body"),
) as mock_handler,
):
result = await self._run(processing_obj, monkeypatch, chunks)
assert isinstance(result, StreamingResponse)
@ -5183,7 +5167,6 @@ class TestCostHeadersForCallsPricedAtZero:
class TestPreCallWithFallbacksOnLocalRateLimit:
@pytest.mark.asyncio
async def test_fallback_triggered_on_local_rate_limit(self):
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
@ -5335,9 +5318,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
mock_router.fallbacks = [{"gpt-4": ["gpt-3.5-turbo"]}]
user_api_key_dict = MagicMock()
user_api_key_dict.router_settings = {
"fallbacks": [{"gpt-4": ["claude-3-haiku"]}]
}
user_api_key_dict.router_settings = {"fallbacks": [{"gpt-4": ["claude-3-haiku"]}]}
with patch.object(
processor,
@ -5368,9 +5349,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
processor = ProxyBaseLLMRequestProcessing(
data={"model": "gpt-4", "disable_fallbacks": True}
)
processor = ProxyBaseLLMRequestProcessing(data={"model": "gpt-4", "disable_fallbacks": True})
async def mock_pre_call_logic(**kwargs):
raise ProxyRateLimitError(
@ -5496,9 +5475,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
# Real per-key per-model TPM limiter + a key carrying the customer's
# `model_tpm_limit` metadata (only the primary is capped).
limiter = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache())
)
limiter = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-lit3890",
metadata={"model_tpm_limit": {primary_model: 100}},
@ -5506,10 +5483,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
# Pre-seed the primary's per-model token counter at the cap so the very
# next request trips it. The counter key uses the *hashed* api_key.
counter_key = (
f"{user_api_key_dict.api_key}::{primary_model}"
f"::{precise_minute}::request_count"
)
counter_key = f"{user_api_key_dict.api_key}::{primary_model}::{precise_minute}::request_count"
await limiter.internal_usage_cache.async_set_cache(
key=counter_key,
value={"current_requests": 0, "current_tpm": 100, "current_rpm": 0},
@ -5540,9 +5514,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
mock_router = MagicMock()
mock_router.fallbacks = [{primary_model: [fallback_model]}]
with patch(
"litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock
):
with patch("litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock):
with patch.object(
processor,
"common_processing_pre_call_logic",
@ -5572,9 +5544,7 @@ class TestPreCallWithFallbacksOnLocalRateLimit:
# Sanity-check the premise: the limiter genuinely raises a
# ProxyRateLimitError for the capped primary under the frozen clock.
with patch(
"litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock
):
with patch("litellm.proxy.hooks.parallel_request_limiter.datetime", _FrozenClock):
with pytest.raises(ProxyRateLimitError):
await limiter.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,