fix(proxy): bill partial usage when a streaming request is cancelled (#30630)

* fix(proxy): bill partial usage when a streaming request is cancelled

On a mid-stream client disconnect the stream never reaches normal completion:
CancelledError / GeneratorExit are BaseException, so neither the success nor the
failure logging path runs, and the assembled-response success logging that
writes SpendLogs never fires. The tokens already produced upstream are billed by
the provider but never recorded on the proxy, so spend undercounts by roughly
the abort rate; the gap is invisible in SpendLogs and only shows up against
provider invoices.

In the shielded streaming cleanup, when a client disconnect is recorded, assemble
the partial usage from the chunks received so far via stream_chunk_builder and
dispatch success logging for it. dispatch_success_handlers de-dupes via
has_dispatched_final_stream_success, so it is a no-op when normal completion
already logged, and it only runs on the cancellation path (the exception path
sets stream_completed and already emits a failure log).

Reported and root-caused in production by @jingyu-lin. Complements #30522, which
releases the budget reservation on the same cancellation path.

* fix(proxy): close upstream stream before billing partial usage on disconnect

Release the provider connection before running partial-usage success
logging on a client disconnect, so slow or external success callbacks
can no longer keep the upstream stream held open. aclose() only closes
completion_stream and leaves response.chunks intact, so the partial
billing still assembles usage from the chunks already received.

---------

Co-authored-by: Bytechoreographer <Bytechoreographer@users.noreply.github.com>
This commit is contained in:
Rick 2026-06-23 20:58:47 +08:00 • committed by Sameer Kankute
parent 54426c193b
commit 5c525eb086
No known key found for this signature in database
2 changed files with 525 additions and 113 deletions

View file

@ -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,

View file

@ -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()