mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(proxy): preserve failure spend log attribution
This commit is contained in:
parent
f005afa146
commit
3fb91d2c37
6 changed files with 292 additions and 230 deletions
|
|
@ -1421,8 +1421,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(
|
||||
|
|
@ -1706,6 +1707,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(
|
||||
|
|
@ -2198,6 +2200,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
|
||||
)
|
||||
|
|
@ -3161,6 +3164,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,
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from litellm._logging import verbose_proxy_logger
|
|||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
_get_parent_otel_span_from_kwargs,
|
||||
get_metadata_variable_name_from_kwargs,
|
||||
get_litellm_metadata_from_kwargs,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
|
|
@ -123,32 +124,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:
|
||||
|
|
@ -207,6 +200,70 @@ 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 {}
|
||||
trusted_metadata_key = _ProxyDBLogger._get_failure_metadata_variable_name(request_data=request_data)
|
||||
|
||||
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(trusted_metadata_key, {}),
|
||||
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,
|
||||
|
|
|
|||
|
|
@ -1771,6 +1771,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] = {}
|
||||
|
||||
|
|
|
|||
|
|
@ -10397,9 +10397,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 []
|
||||
|
|
|
|||
|
|
@ -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,28 @@ 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
|
||||
|
||||
|
||||
_EMBED_PATHS = [
|
||||
"/v1/embeddings",
|
||||
"/embeddings",
|
||||
|
|
@ -119,3 +137,12 @@ 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
|
||||
|
|
|
|||
|
|
@ -129,16 +129,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",
|
||||
|
|
@ -178,14 +174,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",
|
||||
|
|
@ -227,15 +219,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",
|
||||
|
|
@ -254,9 +242,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,
|
||||
|
|
@ -2759,6 +2745,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,
|
||||
|
|
@ -2866,9 +2909,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()
|
||||
|
|
@ -2899,9 +2940,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()
|
||||
|
|
@ -2972,9 +3011,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()
|
||||
|
|
@ -3170,9 +3207,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,
|
||||
)
|
||||
|
||||
|
|
@ -3188,9 +3223,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()
|
||||
|
|
@ -3206,9 +3239,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(
|
||||
|
|
@ -3260,9 +3291,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"
|
||||
|
||||
|
|
@ -3458,9 +3487,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
|
||||
|
||||
|
|
@ -3489,9 +3516,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(
|
||||
|
|
@ -3509,9 +3534,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
|
||||
|
|
@ -3536,9 +3559,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",
|
||||
|
|
@ -3599,9 +3620,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(
|
||||
|
|
@ -3652,9 +3671,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
|
||||
|
||||
|
|
@ -3683,9 +3700,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)
|
||||
|
|
@ -3722,19 +3737,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"
|
||||
)
|
||||
|
||||
|
|
@ -3748,9 +3757,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"]
|
||||
|
|
@ -3775,22 +3782,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):
|
||||
|
|
@ -3806,15 +3803,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):
|
||||
|
|
@ -3825,9 +3818,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,
|
||||
)
|
||||
|
|
@ -3859,9 +3850,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,
|
||||
)
|
||||
|
|
@ -3891,9 +3880,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,
|
||||
)
|
||||
|
|
@ -3908,9 +3895,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)
|
||||
|
|
@ -3921,9 +3906,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(
|
||||
|
|
@ -3942,6 +3925,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:
|
||||
|
|
@ -3968,23 +3953,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()
|
||||
|
|
@ -4003,9 +3982,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()
|
||||
|
|
@ -4020,9 +3997,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()
|
||||
|
|
@ -4031,9 +4006,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)
|
||||
|
|
@ -4041,9 +4014,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()
|
||||
|
|
@ -4122,9 +4093,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,
|
||||
)
|
||||
|
|
@ -4190,7 +4159,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,
|
||||
|
|
@ -4240,7 +4211,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,
|
||||
|
|
@ -4278,7 +4251,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,
|
||||
|
|
@ -4319,7 +4294,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,
|
||||
|
|
@ -4431,7 +4408,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,
|
||||
|
|
@ -4462,9 +4441,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
|
||||
|
|
@ -4515,14 +4492,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)
|
||||
|
||||
|
|
@ -4531,27 +4511,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)
|
||||
|
|
@ -4567,19 +4547,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)
|
||||
|
|
@ -4975,7 +4959,6 @@ class TestResponseCostHeaderForTypedDictResponses:
|
|||
|
||||
|
||||
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
|
||||
|
|
@ -5127,9 +5110,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,
|
||||
|
|
@ -5160,9 +5141,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(
|
||||
|
|
@ -5288,9 +5267,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}},
|
||||
|
|
@ -5298,10 +5275,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},
|
||||
|
|
@ -5332,9 +5306,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",
|
||||
|
|
@ -5364,9 +5336,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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue