mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
0d28a173db
commit
c6112657d3
4 changed files with 23 additions and 20 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue