diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 4d931a47e9d..1dc3c0819ce 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -2429,6 +2429,76 @@ class ProxyBaseLLMRequestProcessing: e, ) + if recorded_client_disconnect: + await ProxyBaseLLMRequestProcessing._bill_partial_stream_on_disconnect( + response, + request_data, + ) + + @staticmethod + async def _bill_partial_stream_on_disconnect( + response: Any, request_data: dict + ) -> None: + """Record SpendLogs for tokens already produced when a stream is cut off. + + A mid-stream client disconnect surfaces as CancelledError / GeneratorExit + (both BaseException), so the stream never reaches normal completion and + the assembled-response success logging never runs. The provider has + already generated and billed for the chunks received so far, but they are + never written to SpendLogs. Assemble the partial usage from the received + chunks and dispatch success logging for it; dispatch_success_handlers + de-dupes via has_dispatched_final_stream_success, so this is a no-op when + normal completion already logged. + """ + logging_obj = request_data.get("litellm_logging_obj") + chunks = getattr(response, "chunks", None) + if logging_obj is None or not chunks: + return + # Optimization, not a correctness guard: dispatch_success_handlers is the + # authoritative de-dup via has_dispatched_final_stream_success. This just + # skips the stream_chunk_builder assembly when completion already logged. + if logging_obj.model_call_details.get("has_dispatched_final_stream_success"): + return + try: + partial_response = litellm.stream_chunk_builder( + chunks=chunks, + messages=getattr(response, "messages", None), + logging_obj=logging_obj, + ) + except Exception: + verbose_proxy_logger.exception( + "Failed to assemble partial streaming usage on client disconnect" + ) + return + if partial_response is None: + return + # When post-call guardrails are active, normal completion routes the + # assembled response through the deferred guardrail path (guardrails, + # then logging). Mirror that here so a cancelled guarded stream is not + # logged unguarded. + # + # Best-effort: a logging or callback failure here must not propagate out + # of the shielded cleanup. + deferred_complete = getattr(logging_obj, "_on_deferred_stream_complete", None) + try: + if deferred_complete is not None: + await deferred_complete(partial_response, False) + else: + # start_time=None -> logging falls back to self.start_time + # (original request start); end_time=None -> now. Correct for a + # partial record. + await logging_obj.dispatch_success_handlers( + partial_response, + start_time=None, + end_time=None, + cache_hit=False, + prefer_async_handlers=True, + ) + except Exception: + verbose_proxy_logger.exception( + "Failed to record partial streaming usage on client disconnect" + ) + @staticmethod async def async_streaming_data_generator( response: Any, diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index c77ee2b784d..4e56fe68856 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -77,7 +77,9 @@ class TestProxyBaseLLMRequestProcessing: assert result.headers["x-litellm-version"] == "test-version" @pytest.mark.asyncio - async def test_base_passthrough_process_llm_request_returns_fastapi_response_from_guardrails(self, monkeypatch): + async def test_base_passthrough_process_llm_request_returns_fastapi_response_from_guardrails( + self, monkeypatch + ): """Post-call guardrails return a FastAPI Response; must not call httpx aread().""" import json @@ -260,7 +262,9 @@ class TestProxyBaseLLMRequestProcessing: assert kwargs["request_headers"] == {"authorization": "Bearer sk-test"} @pytest.mark.asyncio - async def test_common_processing_pre_call_logic_pre_call_hook_receives_litellm_call_id(self, monkeypatch): + async def test_common_processing_pre_call_logic_pre_call_hook_receives_litellm_call_id( + self, monkeypatch + ): processing_obj = ProxyBaseLLMRequestProcessing(data={}) mock_request = MagicMock(spec=Request) mock_request.headers = {} @@ -268,12 +272,16 @@ class TestProxyBaseLLMRequestProcessing: async def mock_add_litellm_data_to_request(*args, **kwargs): return {} - async def mock_common_processing_pre_call_logic(user_api_key_dict, data, call_type): + async def mock_common_processing_pre_call_logic( + user_api_key_dict, data, call_type + ): data_copy = copy.deepcopy(data) return data_copy mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) - mock_proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=mock_common_processing_pre_call_logic) + mock_proxy_logging_obj.pre_call_hook = AsyncMock( + side_effect=mock_common_processing_pre_call_logic + ) monkeypatch.setattr( litellm.proxy.common_request_processing, "add_litellm_data_to_request", @@ -309,7 +317,9 @@ class TestProxyBaseLLMRequestProcessing: pytest.fail("litellm_call_id is not a valid UUID") assert data_passed["litellm_call_id"] == returned_data["litellm_call_id"] - def test_add_dd_apm_tags_for_litellm_call_id_uses_dd_tracing_helper(self, monkeypatch): + def test_add_dd_apm_tags_for_litellm_call_id_uses_dd_tracing_helper( + self, monkeypatch + ): mock_set_active_span_tag = MagicMock(return_value=True) import litellm.proxy.dd_span_tagger @@ -321,10 +331,14 @@ class TestProxyBaseLLMRequestProcessing: DDSpanTagger.tag_call_id("test-call-id") - mock_set_active_span_tag.assert_called_once_with("litellm.call_id", "test-call-id") + mock_set_active_span_tag.assert_called_once_with( + "litellm.call_id", "test-call-id" + ) @pytest.mark.asyncio - async def test_should_apply_hierarchical_router_settings_as_override(self, monkeypatch): + async def test_should_apply_hierarchical_router_settings_as_override( + self, monkeypatch + ): """ Test that hierarchical router settings are stored as router_settings_override instead of creating a full user_config with model_list. @@ -339,12 +353,16 @@ class TestProxyBaseLLMRequestProcessing: async def mock_add_litellm_data_to_request(*args, **kwargs): return {} - async def mock_common_processing_pre_call_logic(user_api_key_dict, data, call_type): + async def mock_common_processing_pre_call_logic( + user_api_key_dict, data, call_type + ): data_copy = copy.deepcopy(data) return data_copy mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) - mock_proxy_logging_obj.pre_call_hook = AsyncMock(side_effect=mock_common_processing_pre_call_logic) + mock_proxy_logging_obj.pre_call_hook = AsyncMock( + side_effect=mock_common_processing_pre_call_logic + ) monkeypatch.setattr( litellm.proxy.common_request_processing, "add_litellm_data_to_request", @@ -360,7 +378,9 @@ class TestProxyBaseLLMRequestProcessing: "timeout": 30.0, "num_retries": 3, } - mock_proxy_config._get_hierarchical_router_settings = AsyncMock(return_value=mock_router_settings) + mock_proxy_config._get_hierarchical_router_settings = AsyncMock( + return_value=mock_router_settings + ) mock_llm_router = MagicMock() @@ -414,18 +434,24 @@ class TestProxyBaseLLMRequestProcessing: # Test with stream timeout header headers_with_timeout = {"x-litellm-stream-timeout": "30.5"} - result = LiteLLMProxyRequestSetup._get_stream_timeout_from_request(headers_with_timeout) + result = LiteLLMProxyRequestSetup._get_stream_timeout_from_request( + headers_with_timeout + ) assert result == 30.5 # Test without stream timeout header headers_without_timeout = {} - result = LiteLLMProxyRequestSetup._get_stream_timeout_from_request(headers_without_timeout) + result = LiteLLMProxyRequestSetup._get_stream_timeout_from_request( + headers_without_timeout + ) assert result is None # Test with invalid header value (should raise ValueError when converting to float) headers_with_invalid = {"x-litellm-stream-timeout": "invalid"} with pytest.raises(ValueError): - LiteLLMProxyRequestSetup._get_stream_timeout_from_request(headers_with_invalid) + LiteLLMProxyRequestSetup._get_stream_timeout_from_request( + headers_with_invalid + ) @pytest.mark.asyncio async def test_build_litellm_proxy_success_headers_from_llm_response(self): @@ -520,7 +546,9 @@ class TestProxyBaseLLMRequestProcessing: ) assert headers["x-litellm-model-id"] == "stream-model-id" - assert headers["x-litellm-model-api-base"] == ("https://generativelanguage.googleapis.com/v1beta") + assert headers["x-litellm-model-api-base"] == ( + "https://generativelanguage.googleapis.com/v1beta" + ) assert headers["llm_provider-x"] == "y" @pytest.mark.asyncio @@ -960,7 +988,9 @@ class TestProxyBaseLLMRequestProcessing: assert "x-litellm-key-spend" in headers_1 expected_spend_1 = 0.001 + 0.0005 # Initial spend + current request cost - assert float(headers_1["x-litellm-key-spend"]) == pytest.approx(expected_spend_1, abs=1e-10) + assert float(headers_1["x-litellm-key-spend"]) == pytest.approx( + expected_spend_1, abs=1e-10 + ) assert float(headers_1["x-litellm-response-cost"]) == response_cost_1 # Test case 2: response_cost is provided as string @@ -973,7 +1003,9 @@ class TestProxyBaseLLMRequestProcessing: assert "x-litellm-key-spend" in headers_2 expected_spend_2 = 0.001 + 0.0003 # Initial spend + current request cost - assert float(headers_2["x-litellm-key-spend"]) == pytest.approx(expected_spend_2, abs=1e-10) + assert float(headers_2["x-litellm-key-spend"]) == pytest.approx( + expected_spend_2, abs=1e-10 + ) # Test case 3: response_cost is None (should use original spend) headers_3 = ProxyBaseLLMRequestProcessing.get_custom_headers( @@ -983,7 +1015,9 @@ class TestProxyBaseLLMRequestProcessing: ) assert "x-litellm-key-spend" in headers_3 - assert float(headers_3["x-litellm-key-spend"]) == 0.001 # Should use original spend + assert ( + float(headers_3["x-litellm-key-spend"]) == 0.001 + ) # Should use original spend # Test case 4: response_cost is 0 (should not change spend) headers_4 = ProxyBaseLLMRequestProcessing.get_custom_headers( @@ -993,7 +1027,9 @@ class TestProxyBaseLLMRequestProcessing: ) assert "x-litellm-key-spend" in headers_4 - assert float(headers_4["x-litellm-key-spend"]) == 0.001 # Should remain unchanged for 0 cost + assert ( + float(headers_4["x-litellm-key-spend"]) == 0.001 + ) # Should remain unchanged for 0 cost # Test case 5: user_api_key_dict.spend is None (should default to 0.0) mock_user_api_key_dict.spend = None @@ -1015,7 +1051,9 @@ class TestProxyBaseLLMRequestProcessing: ) assert "x-litellm-key-spend" in headers_6 - assert float(headers_6["x-litellm-key-spend"]) == 0.001 # Should use original spend + assert ( + float(headers_6["x-litellm-key-spend"]) == 0.001 + ) # Should use original spend # Test case 7: response_cost is invalid string (should fallback to original spend) headers_7 = ProxyBaseLLMRequestProcessing.get_custom_headers( @@ -1025,7 +1063,9 @@ class TestProxyBaseLLMRequestProcessing: ) assert "x-litellm-key-spend" in headers_7 - assert float(headers_7["x-litellm-key-spend"]) == 0.001 # Should use original spend on error + assert ( + float(headers_7["x-litellm-key-spend"]) == 0.001 + ) # Should use original spend on error @pytest.mark.asyncio async def test_queue_time_seconds_is_set_in_metadata(self, monkeypatch): @@ -1086,10 +1126,12 @@ class TestProxyBaseLLMRequestProcessing: # Verify queue_time_seconds is set and non-negative metadata = returned_data.get("metadata", {}) - assert "queue_time_seconds" in metadata, "queue_time_seconds should be set in metadata" - assert metadata["queue_time_seconds"] >= 0.5, ( - f"queue_time_seconds should be at least 0.5, got {metadata['queue_time_seconds']}" - ) + assert ( + "queue_time_seconds" in metadata + ), "queue_time_seconds should be set in metadata" + assert ( + metadata["queue_time_seconds"] >= 0.5 + ), f"queue_time_seconds should be at least 0.5, got {metadata['queue_time_seconds']}" @pytest.mark.asyncio @@ -1240,7 +1282,9 @@ class TestCommonRequestProcessingHelpers: the original status code instead of hardcoding 500. """ mock_gen = AsyncMock() - mock_gen.__anext__.side_effect = HTTPException(status_code=400, detail="Content blocked by guardrail") + mock_gen.__anext__.side_effect = HTTPException( + status_code=400, detail="Content blocked by guardrail" + ) response = await create_response(mock_gen, "text/event-stream", {}) assert response.status_code == 400 @@ -1334,8 +1378,14 @@ class TestCommonRequestProcessingHelpers: response = await create_response(mock_gen, "text/event-stream", {}) content = await self.consume_stream(response) payload = json.loads(content[0][len("data: ") :].strip()) - assert payload["error"]["message"] == "MCP request blocked: no rewritable argument field present" - assert payload["error"]["provider_specific_fields"]["error"]["code"] == "panw_prisma_airs_blocked" + assert ( + payload["error"]["message"] + == "MCP request blocked: no rewritable argument field present" + ) + assert ( + payload["error"]["provider_specific_fields"]["error"]["code"] + == "panw_prisma_airs_blocked" + ) async def test_serialize_http_exception_detail_helper(self): """Direct unit coverage for the L1 helper across all branches.""" @@ -1346,11 +1396,15 @@ class TestCommonRequestProcessingHelpers: assert _serialize_http_exception_detail("plain") == ("plain", None) - msg, fields = _serialize_http_exception_detail({"error": "Violated", "extra": "x"}) + msg, fields = _serialize_http_exception_detail( + {"error": "Violated", "extra": "x"} + ) assert msg == "Violated" assert fields == {"error": "Violated", "extra": "x"} - msg, fields = _serialize_http_exception_detail({"error": {"message": "blocked", "code": "x"}}) + msg, fields = _serialize_http_exception_detail( + {"error": {"message": "blocked", "code": "x"}} + ) assert msg == "blocked" assert fields == {"error": {"message": "blocked", "code": "x"}} @@ -1390,7 +1444,9 @@ class TestCommonRequestProcessingHelpers: yield "data: [DONE]\n\n" custom_headers = {"X-Custom-Header": "TestValue"} - response = await create_response(mock_generator(), "text/event-stream", custom_headers) + response = await create_response( + mock_generator(), "text/event-stream", custom_headers + ) assert response.headers["x-custom-header"] == "TestValue" async def test_create_streaming_response_disables_proxy_buffering(self): @@ -1410,7 +1466,9 @@ class TestCommonRequestProcessingHelpers: error_stream.__anext__.side_effect = ValueError("boom") for generator in (normal_stream(), empty_stream(), error_stream): - response = await create_response(generator, "text/event-stream", {"X-Custom-Header": "keep"}) + response = await create_response( + generator, "text/event-stream", {"X-Custom-Header": "keep"} + ) assert isinstance(response, StreamingResponse) assert response.headers["x-accel-buffering"] == "no" assert response.headers["cache-control"] == "no-cache" @@ -1509,9 +1567,9 @@ class TestCommonRequestProcessingHelpers: for i, call in enumerate(actual_calls): args, kwargs = call - assert args[0] == "streaming.chunk.yield", ( - f"Call {i} should have operation name 'streaming.chunk.yield', got {args[0]}" - ) + assert ( + args[0] == "streaming.chunk.yield" + ), f"Call {i} should have operation name 'streaming.chunk.yield', got {args[0]}" async def test_create_streaming_response_skips_dd_trace_when_disabled(self): """When DD tracing is disabled (the default), the per-chunk span @@ -1692,7 +1750,9 @@ class TestOverrideOpenAIResponseModel: # _hidden_params is an attribute (not a dict key) accessed via getattr response_obj = MagicMock() response_obj.model = fallback_model - response_obj._hidden_params = {"additional_headers": {"x-litellm-attempted-fallbacks": 1}} + response_obj._hidden_params = { + "additional_headers": {"x-litellm-attempted-fallbacks": 1} + } # Call the function - should preserve fallback model _override_openai_response_model( @@ -1819,7 +1879,9 @@ class TestOverrideOpenAIResponseModel: # Create a mock object response response_obj = MagicMock() response_obj.model = downstream_model - response_obj._hidden_params = {"additional_headers": {"x-litellm-attempted-fallbacks": None}} + response_obj._hidden_params = { + "additional_headers": {"x-litellm-attempted-fallbacks": None} + } # Call the function - should override to requested model _override_openai_response_model( @@ -1864,7 +1926,9 @@ class TestOverrideOpenAIResponseModel: # Create a mock object response response_obj = MagicMock() response_obj.model = fallback_model - response_obj._hidden_params = {"additional_headers": {"x-litellm-attempted-fallbacks": 1}} + response_obj._hidden_params = { + "additional_headers": {"x-litellm-attempted-fallbacks": 1} + } # Call the function with None requested_model _override_openai_response_model( @@ -2076,7 +2140,10 @@ class TestIsAzureModelRouterRequest: def test_detects_model_router_with_underscore(self): assert _is_azure_model_router_request("azure_ai/model_router") is True - assert _is_azure_model_router_request("azure_ai/model_router/my-deployment") is True + assert ( + _is_azure_model_router_request("azure_ai/model_router/my-deployment") + is True + ) def test_detects_model_router_with_hyphen(self): assert _is_azure_model_router_request("azure_ai/model-router") is True @@ -2300,7 +2367,9 @@ class TestDDSpanTaggerTagRequest: def test_tags_key_alias_and_model(self): """key_alias and requested_model are set on the span when present.""" - user_key = self._make_user_api_key_dict(key_alias="my-prod-key", token="hashed123") + user_key = self._make_user_api_key_dict( + key_alias="my-prod-key", token="hashed123" + ) with patch("litellm.proxy.dd_span_tagger.set_active_span_tag") as mock_set_tag: DDSpanTagger.tag_request( @@ -2334,7 +2403,9 @@ class TestDDSpanTaggerTagRequest: requested_model="claude-3-5-sonnet", ) - mock_set_tag.assert_called_once_with("litellm.requested_model", "claude-3-5-sonnet") + mock_set_tag.assert_called_once_with( + "litellm.requested_model", "claude-3-5-sonnet" + ) class TestHasAttributeErrorInChain: @@ -2423,7 +2494,9 @@ class TestHandleLLMApiExceptionDictDetail: ) proxy_exc = await self._invoke(exc) assert proxy_exc.message == "Violated guardrail policy" - assert proxy_exc.provider_specific_fields["guardrail_name"] == "bedrock-pii-guard" + assert ( + proxy_exc.provider_specific_fields["guardrail_name"] == "bedrock-pii-guard" + ) # No Python repr leakage of the dict into the message field. assert "{'error':" not in proxy_exc.message @@ -2937,7 +3010,9 @@ class TestAsyncStreamingDataGeneratorFastPath: proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock()) hook_spy = AsyncMock(side_effect=lambda **kw: kw["response"]) - monkeypatch.setattr(proxy_logging_obj, "async_post_call_streaming_hook", hook_spy) + monkeypatch.setattr( + proxy_logging_obj, "async_post_call_streaming_hook", hook_spy + ) chunks = [b"event: a\ndata: {}\n\n", b"event: b\ndata: {}\n\n"] out = [ @@ -2970,7 +3045,9 @@ class TestAsyncStreamingDataGeneratorFastPath: proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock()) hook_spy = AsyncMock(side_effect=lambda **kw: kw["response"]) - monkeypatch.setattr(proxy_logging_obj, "async_post_call_streaming_hook", hook_spy) + monkeypatch.setattr( + proxy_logging_obj, "async_post_call_streaming_hook", hook_spy + ) out = [ c @@ -3012,7 +3089,9 @@ class TestDisconnectGatherCleanup: import asyncio import litellm.proxy.common_request_processing as cpr - from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) async def slow_llm(): await asyncio.sleep(9999) @@ -3030,7 +3109,9 @@ class TestDisconnectGatherCleanup: monkeypatch.setattr(cpr, "route_request", fake_route_request) - processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "gemini-2.0-flash"}) + processing_obj = ProxyBaseLLMRequestProcessing( + data={"model": "gemini-2.0-flash"} + ) monkeypatch.setattr( processing_obj, "common_processing_pre_call_logic", @@ -3062,7 +3143,9 @@ class TestDisconnectGatherCleanup: import asyncio import litellm.proxy.common_request_processing as cpr - from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) async def fake_gather(*_tasks, **_kwargs): raise asyncio.CancelledError() @@ -3077,7 +3160,9 @@ class TestDisconnectGatherCleanup: monkeypatch.setattr(cpr.asyncio, "gather", fake_gather) - processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "gemini-2.0-flash"}) + processing_obj = ProxyBaseLLMRequestProcessing( + data={"model": "gemini-2.0-flash"} + ) monkeypatch.setattr( processing_obj, "common_processing_pre_call_logic", @@ -3112,7 +3197,9 @@ class TestDisconnectGatherCleanup: import asyncio import litellm.proxy.common_request_processing as cpr - from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) hook_cancelled = False @@ -3140,7 +3227,9 @@ class TestDisconnectGatherCleanup: monkeypatch.setattr(cpr, "route_request", fake_route_request) - processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "gemini-2.0-flash"}) + processing_obj = ProxyBaseLLMRequestProcessing( + data={"model": "gemini-2.0-flash"} + ) monkeypatch.setattr( processing_obj, "common_processing_pre_call_logic", @@ -3205,7 +3294,9 @@ class TestDisconnectGatherCleanup: import asyncio import litellm.proxy.common_request_processing as cpr - from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) async def failing_llm(): raise ValueError("llm api error") @@ -3226,7 +3317,9 @@ class TestDisconnectGatherCleanup: monkeypatch.setattr(cpr, "route_request", fake_route_request) - processing_obj = ProxyBaseLLMRequestProcessing(data={"model": "gemini-2.0-flash"}) + processing_obj = ProxyBaseLLMRequestProcessing( + data={"model": "gemini-2.0-flash"} + ) monkeypatch.setattr( processing_obj, "common_processing_pre_call_logic", @@ -3277,9 +3370,7 @@ class TestStreamingClientDisconnectLogging: 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" @@ -3325,7 +3416,9 @@ class TestStreamingClientDisconnectLogging: request_data = { "metadata": {}, "litellm_params": {"metadata": {}}, - "litellm_logging_obj": MagicMock(model_call_details={"metadata": {}, "litellm_params": {}}), + "litellm_logging_obj": MagicMock( + model_call_details={"metadata": {}, "litellm_params": {}} + ), } await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup( @@ -3370,6 +3463,50 @@ class TestStreamingClientDisconnectLogging: mock_response.aclose.assert_awaited_once() assert "client_disconnected" not in request_data["metadata"] + @pytest.mark.asyncio + async def test_finalize_releases_upstream_before_partial_billing(self, monkeypatch): + """The upstream provider connection must be released before the partial + billing dispatch runs, so slow success-logging callbacks cannot keep the + provider stream held open on a client disconnect.""" + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) + + monkeypatch.setattr( + "litellm.proxy.utils.ProxyLogging._fire_deferred_stream_logging", + MagicMock(), + ) + + order = [] + + logging_obj = MagicMock() + logging_obj.model_call_details = {"metadata": {}, "litellm_params": {}} + logging_obj._on_deferred_stream_complete = None + logging_obj.dispatch_success_handlers = AsyncMock( + side_effect=lambda *a, **k: order.append("bill") + ) + + mock_request = MagicMock(spec=Request) + mock_request.is_disconnected = AsyncMock(return_value=True) + mock_response = MagicMock() + mock_response.chunks = ["chunk-1"] + mock_response.aclose = AsyncMock(side_effect=lambda: order.append("aclose")) + request_data = { + "metadata": {}, + "litellm_params": {"metadata": {}}, + "litellm_logging_obj": logging_obj, + } + + with patch.object(litellm, "stream_chunk_builder", return_value=MagicMock()): + await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup( + request=mock_request, + request_data=request_data, + response=mock_response, + ) + + logging_obj.dispatch_success_handlers.assert_awaited_once() + assert order == ["aclose", "bill"] + @pytest.mark.asyncio async def test_async_streaming_data_generator_records_499_on_early_aclose( self, monkeypatch @@ -3422,6 +3559,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: @@ -3448,9 +3587,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 llm_call.cancelled() assert disconnect_event.is_set() @@ -3483,9 +3620,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() @@ -3521,9 +3656,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() @@ -3627,13 +3760,18 @@ class TestAllmPassthroughRoutePostCallGuardrails: cb = MagicMock(spec=CustomGuardrail) cb.guardrail_name = name - cb.event_hook = [GuardrailEventHooks.pre_call.value, GuardrailEventHooks.post_call.value] + cb.event_hook = [ + GuardrailEventHooks.pre_call.value, + GuardrailEventHooks.post_call.value, + ] cb._event_hook_is_event_type = lambda et: et.value in cb.event_hook cb.should_run_guardrail = MagicMock(return_value=True) return cb @pytest.mark.asyncio - async def test_post_call_hook_receives_parsed_dict_not_httpx_response(self, monkeypatch): + async def test_post_call_hook_receives_parsed_dict_not_httpx_response( + self, monkeypatch + ): """ post_call_success_hook must be called with the parsed JSON dict when the non-streaming allm_passthrough_route response is application/json. @@ -3670,7 +3808,11 @@ 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, @@ -3681,9 +3823,9 @@ class TestAllmPassthroughRoutePostCallGuardrails: ) assert len(received_responses) == 1 - assert isinstance(received_responses[0], dict), ( - "post_call_success_hook must receive parsed dict, not httpx.Response" - ) + assert isinstance( + received_responses[0], dict + ), "post_call_success_hook must receive parsed dict, not httpx.Response" assert received_responses[0]["stopReason"] == "end_turn" assert isinstance(result, Response) body = json.loads(result.body) @@ -3720,7 +3862,11 @@ 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, @@ -3758,7 +3904,11 @@ 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, @@ -3799,7 +3949,11 @@ 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, @@ -3855,7 +4009,9 @@ class TestEventStreamAllmPassthroughRoute: @pytest.mark.asyncio async def test_bedrock_provider_dispatches_to_handler(self): stream_bytes = _build_event_stream_frame("messageStart", {"role": "assistant"}) - expected_bytes = _build_event_stream_frame("messageStart", {"role": "assistant"}) + b"extra" + expected_bytes = ( + _build_event_stream_frame("messageStart", {"role": "assistant"}) + b"extra" + ) proxy_logging_obj = MagicMock() user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) @@ -3864,7 +4020,9 @@ class TestEventStreamAllmPassthroughRoute: "litellm.llms.bedrock.passthrough.guardrail_translation.handler.BedrockPassthroughGuardrailHandler.de_anonymize_event_stream", new=AsyncMock(return_value=expected_bytes), ) as mock_handler: - processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "bedrock"}) + processing_obj = ProxyBaseLLMRequestProcessing( + data={"custom_llm_provider": "bedrock"} + ) result = await processing_obj._handle_event_stream_allm_passthrough_route( body_bytes=stream_bytes, proxy_logging_obj=proxy_logging_obj, @@ -3879,7 +4037,9 @@ class TestEventStreamAllmPassthroughRoute: stream_bytes = _build_event_stream_frame("messageStart", {"role": "assistant"}) proxy_logging_obj = MagicMock() - processing_obj = ProxyBaseLLMRequestProcessing(data={"custom_llm_provider": "anthropic"}) + processing_obj = ProxyBaseLLMRequestProcessing( + data={"custom_llm_provider": "anthropic"} + ) result = await processing_obj._handle_event_stream_allm_passthrough_route( body_bytes=stream_bytes, proxy_logging_obj=proxy_logging_obj, @@ -3892,10 +4052,15 @@ class TestEventStreamAllmPassthroughRoute: async def test_non_streaming_response_includes_custom_headers(self): import json - body = {"output": {"message": {"role": "assistant", "content": [{"text": "hi"}]}}} + body = { + "output": {"message": {"role": "assistant", "content": [{"text": "hi"}]}} + } mock_response = MagicMock() mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json", "content-length": "99"} + mock_response.headers = { + "content-type": "application/json", + "content-length": "99", + } mock_response.aread = AsyncMock(return_value=json.dumps(body).encode()) async def mock_hook(data, user_api_key_dict, response): @@ -3911,7 +4076,11 @@ 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, @@ -3995,14 +4164,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) @@ -4019,19 +4191,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, Response) @@ -4047,22 +4223,188 @@ 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) streamed = [chunk async for chunk in result.body_iterator] assert streamed == chunks mock_handler.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_bill_partial_stream_on_disconnect_dispatches_partial_usage(): + """A stream cut off mid-flight bills the tokens already received.""" + logging_obj = MagicMock() + logging_obj.model_call_details = {} + logging_obj.dispatch_success_handlers = AsyncMock() + logging_obj._on_deferred_stream_complete = None # no post-call guardrails + + response = MagicMock() + response.chunks = ["chunk-1", "chunk-2"] + response.messages = [{"role": "user", "content": "hi"}] + request_data = {"litellm_logging_obj": logging_obj} + + partial_response = MagicMock(name="partial_response") + with patch.object( + litellm, "stream_chunk_builder", return_value=partial_response + ) as scb: + await ProxyBaseLLMRequestProcessing._bill_partial_stream_on_disconnect( + response, request_data + ) + + scb.assert_called_once() + assert scb.call_args.kwargs["chunks"] == ["chunk-1", "chunk-2"] + logging_obj.dispatch_success_handlers.assert_awaited_once() + assert logging_obj.dispatch_success_handlers.call_args.args[0] is partial_response + + +@pytest.mark.asyncio +async def test_bill_partial_stream_on_disconnect_routes_through_guardrails(): + """With post-call guardrails active, the partial response goes through the + deferred guardrail path (guardrails then logging), not raw to success.""" + logging_obj = MagicMock() + logging_obj.model_call_details = {} + logging_obj.dispatch_success_handlers = AsyncMock() + logging_obj._on_deferred_stream_complete = AsyncMock() + + response = MagicMock() + response.chunks = ["chunk-1"] + response.messages = [{"role": "user", "content": "hi"}] + request_data = {"litellm_logging_obj": logging_obj} + + partial_response = MagicMock(name="partial_response") + with patch.object(litellm, "stream_chunk_builder", return_value=partial_response): + await ProxyBaseLLMRequestProcessing._bill_partial_stream_on_disconnect( + response, request_data + ) + + logging_obj._on_deferred_stream_complete.assert_awaited_once() + assert ( + logging_obj._on_deferred_stream_complete.call_args.args[0] is partial_response + ) + # must not bypass guardrails by dispatching success directly + logging_obj.dispatch_success_handlers.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_bill_partial_stream_on_disconnect_skips_when_already_dispatched(): + """Normal completion already logged -> must not double-log.""" + logging_obj = MagicMock() + logging_obj.model_call_details = {"has_dispatched_final_stream_success": True} + logging_obj.dispatch_success_handlers = AsyncMock() + + response = MagicMock() + response.chunks = ["chunk-1"] + request_data = {"litellm_logging_obj": logging_obj} + + with patch.object(litellm, "stream_chunk_builder", return_value=MagicMock()) as scb: + await ProxyBaseLLMRequestProcessing._bill_partial_stream_on_disconnect( + response, request_data + ) + + scb.assert_not_called() + logging_obj.dispatch_success_handlers.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_bill_partial_stream_on_disconnect_skips_when_no_chunks(): + """No chunks received -> nothing to bill.""" + logging_obj = MagicMock() + logging_obj.model_call_details = {} + logging_obj.dispatch_success_handlers = AsyncMock() + + response = MagicMock() + response.chunks = [] + request_data = {"litellm_logging_obj": logging_obj} + + with patch.object(litellm, "stream_chunk_builder", return_value=MagicMock()) as scb: + await ProxyBaseLLMRequestProcessing._bill_partial_stream_on_disconnect( + response, request_data + ) + + scb.assert_not_called() + logging_obj.dispatch_success_handlers.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_bill_partial_stream_on_disconnect_swallows_builder_errors(): + """A failure assembling partial usage must not escape the shielded cleanup.""" + logging_obj = MagicMock() + logging_obj.model_call_details = {} + logging_obj.dispatch_success_handlers = AsyncMock() + logging_obj._on_deferred_stream_complete = None + + response = MagicMock() + response.chunks = ["chunk-1"] + request_data = {"litellm_logging_obj": logging_obj} + + with patch.object( + litellm, "stream_chunk_builder", side_effect=RuntimeError("boom") + ): + # must return without raising + await ProxyBaseLLMRequestProcessing._bill_partial_stream_on_disconnect( + response, request_data + ) + + logging_obj.dispatch_success_handlers.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_bill_partial_stream_on_disconnect_skips_when_builder_returns_none(): + """If nothing assembles into a billable response, do not log.""" + logging_obj = MagicMock() + logging_obj.model_call_details = {} + logging_obj.dispatch_success_handlers = AsyncMock() + logging_obj._on_deferred_stream_complete = None + + response = MagicMock() + response.chunks = ["chunk-1"] + request_data = {"litellm_logging_obj": logging_obj} + + with patch.object(litellm, "stream_chunk_builder", return_value=None): + await ProxyBaseLLMRequestProcessing._bill_partial_stream_on_disconnect( + response, request_data + ) + + logging_obj.dispatch_success_handlers.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_bill_partial_stream_on_disconnect_swallows_dispatch_errors(): + """A logging/callback failure must not escape the shielded cleanup; if it did + it would skip response.aclose() and leak the upstream connection.""" + logging_obj = MagicMock() + logging_obj.model_call_details = {} + logging_obj.dispatch_success_handlers = AsyncMock( + side_effect=RuntimeError("db write failed") + ) + logging_obj._on_deferred_stream_complete = None + + response = MagicMock() + response.chunks = ["chunk-1"] + request_data = {"litellm_logging_obj": logging_obj} + + with patch.object(litellm, "stream_chunk_builder", return_value=MagicMock()): + # must return without raising + await ProxyBaseLLMRequestProcessing._bill_partial_stream_on_disconnect( + response, request_data + ) + + logging_obj.dispatch_success_handlers.assert_awaited_once()