From 3fb91d2c374cef57d197ff831ec22a34085eef09 Mon Sep 17 00:00:00 2001 From: silencedoctor <33445544+silencedoctor@users.noreply.github.com> Date: Wed, 29 Jul 2026 16:55:52 +0800 Subject: [PATCH 1/2] fix(proxy): preserve failure spend log attribution --- litellm/proxy/common_request_processing.py | 7 +- .../proxy/hooks/proxy_track_cost_callback.py | 91 ++++- litellm/proxy/litellm_pre_call_utils.py | 1 + litellm/proxy/proxy_server.py | 6 +- .../proxy_server/test_routes_embeddings.py | 45 ++- .../proxy/test_common_request_processing.py | 372 ++++++++---------- 6 files changed, 292 insertions(+), 230 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index dbbf9cb673e..8e6db15d103 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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, diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 99d0c94d11b..ed51d696462 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -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, diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index cb5002e431b..df08a72bd98 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -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] = {} diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7dced4e26b6..bb13c576585 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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 [] diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py b/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py index 98249cb5ad5..b319afcff1f 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py @@ -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 diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 58714a5e319..de247fd46b6 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -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, From caa07aa3f50108131e726c2ff821160f8b9221e3 Mon Sep 17 00:00:00 2001 From: silencedoctor <33445544+silencedoctor@users.noreply.github.com> Date: Wed, 29 Jul 2026 17:51:50 +0800 Subject: [PATCH 2/2] fix(proxy): reuse embeddings processor on failures --- .../proxy/hooks/proxy_track_cost_callback.py | 10 ++- .../proxy_server/test_routes_embeddings.py | 41 ++++++++++ .../proxy/test_chat_completion_metadata.py | 80 +++++++++---------- 3 files changed, 84 insertions(+), 47 deletions(-) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index ed51d696462..f8bdd32c762 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -8,8 +8,8 @@ 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, + 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 @@ -207,8 +207,6 @@ class _ProxyDBLogger(CustomLogger): ) -> 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, *, @@ -229,7 +227,11 @@ class _ProxyDBLogger(CustomLogger): merged_metadata[key] = value merge_metadata( - request_data.get(trusted_metadata_key, {}), + 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, ) diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py b/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py index b319afcff1f..c56b048946f 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_embeddings.py @@ -96,6 +96,36 @@ def embedding_body_read_raises(monkeypatch): 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", @@ -146,3 +176,14 @@ def test_embeddings_body_read_error_preserves_call_type(client, auth_as, embeddi 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 diff --git a/tests/test_litellm/proxy/test_chat_completion_metadata.py b/tests/test_litellm/proxy/test_chat_completion_metadata.py index 38dcdc13c50..5016992deef 100644 --- a/tests/test_litellm/proxy/test_chat_completion_metadata.py +++ b/tests/test_litellm/proxy/test_chat_completion_metadata.py @@ -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