mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge caa07aa3f5 into 3746ba58d7
This commit is contained in:
commit
d73eb81a4f
7 changed files with 372 additions and 273 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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] = {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue