fix(proxy): make the streaming cost-injection gate public

The gate is called from the passthrough streaming handler as well as from
common_request_processing, and a private cross-module call pushes basedpyright's
reportPrivateUsage over its budget. It is a shared decision helper, so make it
public rather than suppressing the rule.

The two budget-reservation slow-path tests used include_cost_in_streaming_usage
alone to force fast_path off. The gate now also needs the caller's
stream_options.include_usage opt-in for that, so they pass one.
This commit is contained in:
Priyansh Nandwana 2026-08-27 16:48:11 +05:30
parent 0d28a173db
commit c6112657d3
4 changed files with 23 additions and 20 deletions

View file

@ -3632,7 +3632,7 @@ class ProxyBaseLLMRequestProcessing:
chunks can emit anything still held at end of stream.
``protocol_supports_stream_options`` says whether this route's protocol gives
callers a ``stream_options.include_usage`` to opt in with, which gates cost
injection; see ``_should_inject_cost_for_request``.
injection; see ``should_inject_cost_for_request``.
"""
verbose_proxy_logger.debug("inside generator")
# Resolve per-stream (not per-chunk) whether the heavy per-chunk path
@ -3643,7 +3643,7 @@ class ProxyBaseLLMRequestProcessing:
# await, response-string materialization, and cost-injection call are
# pure overhead on the streaming hot path (the default config).
caps: Final = ProxyLogging._callback_capabilities()
cost_injection_enabled: Final = ProxyBaseLLMRequestProcessing._should_inject_cost_for_request(
cost_injection_enabled: Final = ProxyBaseLLMRequestProcessing.should_inject_cost_for_request(
request_data,
protocol_supports_stream_options=protocol_supports_stream_options,
)
@ -3798,7 +3798,7 @@ class ProxyBaseLLMRequestProcessing:
)
@staticmethod
def _should_inject_cost_for_request(
def should_inject_cost_for_request(
request_data: Mapping[str, Any] | None,
*,
protocol_supports_stream_options: bool = True,
@ -3855,7 +3855,7 @@ class ProxyBaseLLMRequestProcessing:
chunk: The streaming chunk (dict, str, bytes, or bytearray)
model_name: Model name for cost calculation
litellm_logging_obj: The call's logging object, used for pricing
enabled: Per-stream decision from ``_should_inject_cost_for_request``.
enabled: Per-stream decision from ``should_inject_cost_for_request``.
Falls back to the global flag alone when not passed.
Returns:

View file

@ -162,7 +162,7 @@ class PassThroughStreamingHandler:
)
cost_injection_active: Final = (
ProxyBaseLLMRequestProcessing._should_inject_cost_for_request(
ProxyBaseLLMRequestProcessing.should_inject_cost_for_request(
request_body,
protocol_supports_stream_options=endpoint_type == EndpointType.OPENAI,
)

View file

@ -73,11 +73,12 @@ def spend_counter_state():
ps.prisma_client = original_prisma_client
def _request_body() -> dict:
def _request_body(*, include_usage: bool = False) -> dict:
return {
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": "hello"}],
"max_tokens": 10,
**({"stream_options": {"include_usage": True}} if include_usage else {}),
}
@ -2794,14 +2795,15 @@ async def test_streaming_cancel_in_slow_path_before_yield_refunds(spend_counter_
generator = ProxyBaseLLMRequestProcessing.async_streaming_data_generator(
response=MagicMock(),
user_api_key_dict=valid_token,
request_data=_request_body(),
request_data=_request_body(include_usage=True),
proxy_logging_obj=streaming_logging_obj,
serialize_chunk=lambda chunk: chunk,
serialize_error=lambda exc: str(exc),
)
received = []
# include_cost_in_streaming_usage forces fast_path off, so the hook above runs
# include_cost_in_streaming_usage plus the caller's stream_options.include_usage
# opt-in forces fast_path off, so the hook above runs
with patch.object(litellm, "include_cost_in_streaming_usage", True, create=True):
async def _drain():
async for chunk in generator:
@ -2870,15 +2872,16 @@ async def test_streaming_slow_path_processes_and_yields_chunk(spend_counter_stat
generator = ProxyBaseLLMRequestProcessing.async_streaming_data_generator(
response=MagicMock(),
user_api_key_dict=valid_token,
request_data=_request_body(),
request_data=_request_body(include_usage=True),
proxy_logging_obj=streaming_logging_obj,
serialize_chunk=lambda chunk: chunk,
serialize_error=lambda exc: str(exc),
)
received = []
# include_cost_in_streaming_usage forces the slow path so the per-chunk hook,
# content accumulation, and cost-injection branch all run to a successful yield
# include_cost_in_streaming_usage plus the caller's stream_options.include_usage
# opt-in forces the slow path, so the per-chunk hook, content accumulation, and
# cost-injection branch all run to a successful yield
with patch.object(litellm, "include_cost_in_streaming_usage", True, create=True):
async for chunk in generator:
received.append(chunk)

View file

@ -8507,7 +8507,7 @@ class TestShouldInjectCostForRequest:
def test_global_flag_off_never_injects(self, monkeypatch):
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", False)
assert (
ProxyBaseLLMRequestProcessing._should_inject_cost_for_request(
ProxyBaseLLMRequestProcessing.should_inject_cost_for_request(
{"stream_options": {"include_usage": True}}
)
is False
@ -8516,7 +8516,7 @@ class TestShouldInjectCostForRequest:
def test_openai_protocol_requires_caller_opt_in(self, monkeypatch):
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
assert (
ProxyBaseLLMRequestProcessing._should_inject_cost_for_request(
ProxyBaseLLMRequestProcessing.should_inject_cost_for_request(
{"model": "gpt-4o-mini", "stream": True},
protocol_supports_stream_options=True,
)
@ -8526,7 +8526,7 @@ class TestShouldInjectCostForRequest:
def test_openai_protocol_opted_in_injects(self, monkeypatch):
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
assert (
ProxyBaseLLMRequestProcessing._should_inject_cost_for_request(
ProxyBaseLLMRequestProcessing.should_inject_cost_for_request(
{"stream_options": {"include_usage": True}},
protocol_supports_stream_options=True,
)
@ -8537,7 +8537,7 @@ class TestShouldInjectCostForRequest:
def test_explicit_opt_out_is_honoured_on_every_protocol(self, monkeypatch, protocol_supports_stream_options):
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
assert (
ProxyBaseLLMRequestProcessing._should_inject_cost_for_request(
ProxyBaseLLMRequestProcessing.should_inject_cost_for_request(
{"stream_options": {"include_usage": False}},
protocol_supports_stream_options=protocol_supports_stream_options,
)
@ -8549,7 +8549,7 @@ class TestShouldInjectCostForRequest:
in, so the flag remains always-on there."""
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
assert (
ProxyBaseLLMRequestProcessing._should_inject_cost_for_request(
ProxyBaseLLMRequestProcessing.should_inject_cost_for_request(
{"model": "claude-haiku-4-5", "stream": True},
protocol_supports_stream_options=False,
)
@ -8558,9 +8558,9 @@ class TestShouldInjectCostForRequest:
def test_missing_request_data_falls_back_to_protocol_default(self, monkeypatch):
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
assert ProxyBaseLLMRequestProcessing._should_inject_cost_for_request(None) is False
assert ProxyBaseLLMRequestProcessing.should_inject_cost_for_request(None) is False
assert (
ProxyBaseLLMRequestProcessing._should_inject_cost_for_request(
ProxyBaseLLMRequestProcessing.should_inject_cost_for_request(
None, protocol_supports_stream_options=False
)
is True
@ -8569,7 +8569,7 @@ class TestShouldInjectCostForRequest:
def test_malformed_stream_options_falls_back_to_protocol_default(self, monkeypatch):
monkeypatch.setattr(litellm, "include_cost_in_streaming_usage", True)
assert (
ProxyBaseLLMRequestProcessing._should_inject_cost_for_request(
ProxyBaseLLMRequestProcessing.should_inject_cost_for_request(
{"stream_options": "include_usage"},
protocol_supports_stream_options=False,
)
@ -8579,7 +8579,7 @@ class TestShouldInjectCostForRequest:
class TestProcessChunkCostInjectionGate:
"""``_process_chunk_with_cost_injection`` takes the per-stream decision from
``_should_inject_cost_for_request`` and falls back to the global flag when the
``should_inject_cost_for_request`` and falls back to the global flag when the
caller does not pass one."""
@staticmethod