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 1802fe7228
commit 9869fd6143
4 changed files with 23 additions and 20 deletions

View file

@ -3887,7 +3887,7 @@ class ProxyBaseLLMRequestProcessing:
``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
@ -3898,7 +3898,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,
)
@ -4060,7 +4060,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,
@ -4117,7 +4117,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

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

@ -78,11 +78,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 {}),
}
@ -3115,14 +3116,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:
@ -3191,15 +3193,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

@ -10236,7 +10236,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
@ -10245,7 +10245,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,
)
@ -10255,7 +10255,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,
)
@ -10266,7 +10266,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,
)
@ -10278,7 +10278,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,
)
@ -10287,9 +10287,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
@ -10298,7 +10298,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,
)
@ -10308,7 +10308,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