From b5d1fb7298ef4141ae73831e145e674f483b4db0 Mon Sep 17 00:00:00 2001 From: Ultronen <1594778171@qq.com> Date: Sat, 5 Sep 2026 17:54:01 +0800 Subject: [PATCH 1/7] fix(proxy): preserve tool calls through streaming hooks --- litellm/proxy/utils.py | 42 ++++++- .../proxy_server/test_streaming_helpers.py | 77 ++++++++++++- .../proxy_logging/test_streaming_hooks.py | 107 ++++++++++++++++++ 3 files changed, 220 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index c095586b6c9..b53b6bf04f2 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -892,6 +892,27 @@ def _call_type_for_route(route: str | None) -> str | None: return call_types[0].value if len(operations) == 1 else None +class _StreamingHookResponseText(str): + """Marks the exact text object passed to a per-chunk streaming hook.""" + + +def _streaming_hook_response_text(*, response_str: str, str_so_far: str | None, response: object) -> str: + complete_response = str_so_far + response_str if str_so_far is not None else response_str + if complete_response == "" and isinstance(response, (ModelResponse, ModelResponseStream)): + return _StreamingHookResponseText(complete_response) + return complete_response + + +def _is_unchanged_structured_streaming_hook_response( + *, callback_response: object, complete_response: str, response_str: str, response: object +) -> bool: + if response_str != "" or not isinstance(response, (ModelResponse, ModelResponseStream)): + return False + if isinstance(complete_response, _StreamingHookResponseText): + return callback_response is complete_response + return callback_response == complete_response + + def _failure_fields_to_lift(request_data: Mapping[str, object]) -> Mapping[str, object]: """Failure-path callbacks run after ``litellm_logging_obj`` is popped from request_data (it is not serialisable), so the caller merges these fields @@ -3513,18 +3534,29 @@ class ProxyLogging: else: _callback = callback if _callback is not None and isinstance(_callback, CustomLogger): - if str_so_far is not None: - complete_response = str_so_far + response_str - else: - complete_response = response_str + complete_response = _streaming_hook_response_text( + response_str=response_str, + str_so_far=str_so_far, + response=response, + ) callback_response: ( - ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream | None + str | ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream | None ) callback_response = await _callback.async_post_call_streaming_hook( user_api_key_dict=user_api_key_dict, response=complete_response, ) if callback_response is not None: + # A text result cannot represent a structured empty-text + # chunk such as a tool-call delta. Preserve the chunk + # only when the callback returned its input unchanged. + if _is_unchanged_structured_streaming_hook_response( + callback_response=callback_response, + complete_response=complete_response, + response_str=response_str, + response=response, + ): + continue response = callback_response except Exception as e: raise e diff --git a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py index 87e10ce7e8d..2cf17c57f03 100644 --- a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py +++ b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py @@ -24,8 +24,9 @@ from fastapi import Response from fastapi.responses import StreamingResponse import litellm -from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY import litellm.proxy.proxy_server as ps +from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY +from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.proxy_server import ( _apply_streaming_chunk_hooks, @@ -682,6 +683,80 @@ async def test_async_data_generator_yields_sse_chunks_and_done(monkeypatch): } +@pytest.mark.asyncio +async def test_async_data_generator_preserves_tool_calls_through_per_chunk_hook( + monkeypatch, +): + class _PassThrough(CustomLogger): + async def async_post_call_streaming_hook(self, user_api_key_dict, response): + return response + + callback = _PassThrough() + monkeypatch.setattr(litellm, "callbacks", [callback]) + _patch_logging_flags(monkeypatch, needs_per_chunk=True) + + chunk = ModelResponseStream( + id="chatcmpl-tools", + choices=[ + { + "index": 0, + "delta": { + "tool_calls": [ + { + "index": 0, + "id": "call-weather", + "type": "function", + "function": { + "name": "get_weather", + "arguments": "", + }, + }, + { + "index": 1, + "id": "call-time", + "type": "function", + "function": {"name": "get_time", "arguments": ""}, + }, + ] + }, + "finish_reason": None, + } + ], + created=0, + model="gpt-4o-mini", + object="chat.completion.chunk", + ) + + out = [] + async for line in async_data_generator( + response=_async_iter([chunk]), + user_api_key_dict=_user_auth(), + request_data={"model": "gpt-4o-mini"}, + ): + out.append(line) + + first = out[0] + assert isinstance(first, (str, bytes)) + first_text = first.decode() if isinstance(first, bytes) else first + assert first_text.startswith("data: {") + payload = json.loads(first_text.removeprefix("data: ").removesuffix("\n\n")) + assert payload["choices"][0]["delta"]["tool_calls"] == [ + { + "index": 0, + "id": "call-weather", + "type": "function", + "function": {"name": "get_weather", "arguments": ""}, + }, + { + "index": 1, + "id": "call-time", + "type": "function", + "function": {"name": "get_time", "arguments": ""}, + }, + ] + assert out[-1] == "data: [DONE]\n\n" + + @pytest.mark.asyncio async def test_async_data_generator_uses_response_fallback_metadata(monkeypatch): _patch_logging_flags(monkeypatch) diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py index ec5b994f147..5478bac75bf 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py @@ -257,6 +257,113 @@ async def test_async_post_call_streaming_hook_invokes_per_chunk_callback(proxy_l assert out.startswith("modified-") +@pytest.mark.parametrize("str_so_far", [None, "I will check that. "]) +@pytest.mark.asyncio +async def test_async_post_call_streaming_hook_preserves_tool_calls_when_callback_returns_unmodified_text( + proxy_logging, make_user_api_key_auth, monkeypatch, str_so_far +): + class _PassThrough(CustomLogger): + async def async_post_call_streaming_hook(self, user_api_key_dict, response): + return response + + monkeypatch.setattr(litellm, "callbacks", [_PassThrough()]) + + response = litellm.ModelResponseStream( + id="chatcmpl-tools", + choices=[ + { + "index": 0, + "delta": { + "tool_calls": [ + { + "index": 0, + "id": "call-weather", + "type": "function", + "function": {"name": "get_weather", "arguments": ""}, + }, + { + "index": 1, + "id": "call-time", + "type": "function", + "function": {"name": "get_time", "arguments": ""}, + }, + ] + }, + "finish_reason": None, + } + ], + created=0, + model="gpt-4o-mini", + object="chat.completion.chunk", + ) + + out = await proxy_logging.async_post_call_streaming_hook( + data={}, + response=response, + user_api_key_dict=make_user_api_key_auth(), + str_so_far=str_so_far, + ) + + assert out is response + assert [ + tool_call.model_dump(exclude_none=True) + for tool_call in out.choices[0].delta.tool_calls + ] == [ + { + "id": "call-weather", + "function": {"arguments": "", "name": "get_weather"}, + "type": "function", + "index": 0, + }, + { + "id": "call-time", + "function": {"arguments": "", "name": "get_time"}, + "type": "function", + "index": 1, + }, + ] + + class _Replace(CustomLogger): + async def async_post_call_streaming_hook(self, user_api_key_dict, response): + return "replacement" + + monkeypatch.setattr(litellm, "callbacks", [_PassThrough(), _Replace()]) + replaced = await proxy_logging.async_post_call_streaming_hook( + data={}, + response=response, + user_api_key_dict=make_user_api_key_auth(), + str_so_far=str_so_far, + ) + assert replaced == "replacement" + + if str_so_far is not None: + class _EquivalentCopy(CustomLogger): + async def async_post_call_streaming_hook(self, user_api_key_dict, response): + return response.encode().decode() + + monkeypatch.setattr(litellm, "callbacks", [_EquivalentCopy()]) + equivalent = await proxy_logging.async_post_call_streaming_hook( + data={}, + response=response, + user_api_key_dict=make_user_api_key_auth(), + str_so_far=str_so_far, + ) + assert equivalent is response + + class _Suppress(CustomLogger): + async def async_post_call_streaming_hook(self, user_api_key_dict, response): + return "" + + monkeypatch.setattr(litellm, "callbacks", [_Suppress()]) + suppressed = await proxy_logging.async_post_call_streaming_hook( + data={}, + response=response, + user_api_key_dict=make_user_api_key_auth(), + str_so_far=str_so_far, + ) + assert suppressed == "" + + @pytest.mark.asyncio async def test_async_post_call_streaming_hook_callback_error_raises(proxy_logging, make_user_api_key_auth, monkeypatch): class _Per(CustomLogger): From 9bccb5bdc02d1e7555e63a01e361c98fdef8473a Mon Sep 17 00:00:00 2001 From: Ultronen <82553854+Ultronen@users.noreply.github.com> Date: Sun, 6 Sep 2026 18:16:42 +0800 Subject: [PATCH 2/7] fix(proxy): scan streaming tool calls with guardrails --- litellm/proxy/utils.py | 32 ++++++++++- .../proxy_logging/test_streaming_hooks.py | 55 +++++++++++++++++++ 2 files changed, 86 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index b53b6bf04f2..e1ec9444a6e 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -903,13 +903,38 @@ def _streaming_hook_response_text(*, response_str: str, str_so_far: str | None, return complete_response +def _streaming_guardrail_response_text(*, complete_response: str, response: object) -> str: + if not isinstance(response, (ModelResponse, ModelResponseStream)): + return complete_response + response_dict: Final = response.model_dump(mode="json", exclude_none=True) + structured_choices: Final = tuple( + {key: delta[key] for key in ("tool_calls", "function_call") if delta.get(key) is not None} + for choice in response_dict.get("choices", []) + if isinstance(choice, dict) + for delta in (choice.get("delta"),) + if isinstance(delta, dict) and any(delta.get(key) is not None for key in ("tool_calls", "function_call")) + ) + if not structured_choices: + return complete_response + return _StreamingHookResponseText( + json.dumps( + {"content": str(complete_response), "choices": structured_choices}, + ensure_ascii=False, + separators=(",", ":"), + sort_keys=True, + ) + ) + + def _is_unchanged_structured_streaming_hook_response( *, callback_response: object, complete_response: str, response_str: str, response: object ) -> bool: - if response_str != "" or not isinstance(response, (ModelResponse, ModelResponseStream)): + if not isinstance(response, (ModelResponse, ModelResponseStream)): return False if isinstance(complete_response, _StreamingHookResponseText): return callback_response is complete_response + if response_str != "": + return False return callback_response == complete_response @@ -3539,6 +3564,11 @@ class ProxyLogging: str_so_far=str_so_far, response=response, ) + if isinstance(_callback, CustomGuardrail): + complete_response = _streaming_guardrail_response_text( + complete_response=complete_response, + response=response, + ) callback_response: ( str | ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream | None ) diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py index 5478bac75bf..97573e62b7e 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py @@ -18,6 +18,7 @@ import pytest from fastapi import HTTPException import litellm +from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( @@ -364,6 +365,60 @@ async def test_async_post_call_streaming_hook_preserves_tool_calls_when_callback assert suppressed == "" +@pytest.mark.asyncio +async def test_async_post_call_streaming_hook_exposes_tool_calls_to_guardrails( + proxy_logging, make_user_api_key_auth, monkeypatch +): + class _BlockingGuardrail(CustomGuardrail): + def should_run_guardrail(self, data, event_type): + return True + + async def async_post_call_streaming_hook(self, user_api_key_dict, response): + if "blocked-command" in response: + return "data: blocked\n\n" + return response + + monkeypatch.setattr( + litellm, + "callbacks", + [_BlockingGuardrail(guardrail_name="tool-call-scanner")], + ) + + response = litellm.ModelResponseStream( + id="chatcmpl-guardrail-tools", + choices=[ + { + "index": 0, + "delta": { + "tool_calls": [ + { + "index": 0, + "id": "call-shell", + "type": "function", + "function": { + "name": "run_command", + "arguments": '{"command":"blocked-command"}', + }, + } + ] + }, + "finish_reason": None, + } + ], + created=0, + model="gpt-4o-mini", + object="chat.completion.chunk", + ) + + out = await proxy_logging.async_post_call_streaming_hook( + data={}, + response=response, + user_api_key_dict=make_user_api_key_auth(), + ) + + assert out == "data: blocked\n\n" + + @pytest.mark.asyncio async def test_async_post_call_streaming_hook_callback_error_raises(proxy_logging, make_user_api_key_auth, monkeypatch): class _Per(CustomLogger): From 9f231af27036f5f8ea74a8cf0b99610ecf46993e Mon Sep 17 00:00:00 2001 From: Ultronen <82553854+Ultronen@users.noreply.github.com> Date: Sun, 6 Sep 2026 18:22:18 +0800 Subject: [PATCH 3/7] refactor(proxy): keep guardrail projection immutable --- litellm/proxy/utils.py | 21 ++++++++------------- 1 file changed, 8 insertions(+), 13 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index e1ec9444a6e..9c9f643d626 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -907,23 +907,18 @@ def _streaming_guardrail_response_text(*, complete_response: str, response: obje if not isinstance(response, (ModelResponse, ModelResponseStream)): return complete_response response_dict: Final = response.model_dump(mode="json", exclude_none=True) - structured_choices: Final = tuple( - {key: delta[key] for key in ("tool_calls", "function_call") if delta.get(key) is not None} - for choice in response_dict.get("choices", []) + structured_fields: Final = tuple( + f"{key}:{json.dumps(delta[key], ensure_ascii=False, separators=(',', ':'), sort_keys=True)}" + for choice in response_dict.get("choices", ()) if isinstance(choice, dict) for delta in (choice.get("delta"),) - if isinstance(delta, dict) and any(delta.get(key) is not None for key in ("tool_calls", "function_call")) + if isinstance(delta, dict) + for key in ("tool_calls", "function_call") + if delta.get(key) is not None ) - if not structured_choices: + if not structured_fields: return complete_response - return _StreamingHookResponseText( - json.dumps( - {"content": str(complete_response), "choices": structured_choices}, - ensure_ascii=False, - separators=(",", ":"), - sort_keys=True, - ) - ) + return _StreamingHookResponseText("\n".join((str(complete_response), *structured_fields))) def _is_unchanged_structured_streaming_hook_response( From 02e4307523ed17cb75340df290d973c15ba64ba9 Mon Sep 17 00:00:00 2001 From: Ultronen <82553854+Ultronen@users.noreply.github.com> Date: Sun, 6 Sep 2026 19:03:59 +0800 Subject: [PATCH 4/7] fix(proxy): scan accumulated streaming tool calls --- litellm/proxy/common_request_processing.py | 32 ++++- litellm/proxy/proxy_server.py | 13 +- litellm/proxy/utils.py | 116 ++++++++++++++++-- .../proxy_server/test_streaming_helpers.py | 6 +- .../proxy_logging/test_streaming_hooks.py | 78 +++++++++++- 5 files changed, 221 insertions(+), 24 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index fc603701e44..460d3a5a051 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -70,7 +70,7 @@ from litellm.proxy.common_utils.sse_keepalive import ( from litellm.proxy.dd_span_tagger import DDSpanTagger from litellm.proxy.guardrails.auto_router_compression import arm_pre_call as _arm_auto_router_compression from litellm.proxy.route_llm_request import route_request -from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails +from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails, stream_chunks_with_response from litellm.router import Router from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict from litellm.router_utils.common_utils import resolve_model_group_alias @@ -3621,6 +3621,27 @@ class ProxyBaseLLMRequestProcessing: e, ) + @staticmethod + async def _apply_streaming_chunk_hook( + *, + chunk: Any, # noqa: ANN401 # streaming callbacks may replace chunks with provider-specific response types + proxy_logging_obj: ProxyLogging, + user_api_key_dict: UserAPIKeyAuth, + request_data: dict, # mutable-ok: ProxyLogging requires the existing mutable request payload + str_so_far: str, + stream_chunks_so_far: Sequence[ModelResponseStream], + ) -> tuple[Any, tuple[ModelResponseStream, ...]]: + return ( + await proxy_logging_obj.async_post_call_streaming_hook( + user_api_key_dict=user_api_key_dict, + response=chunk, + data=request_data, + str_so_far=str_so_far, + stream_chunks_so_far=stream_chunks_so_far, + ), + stream_chunks_with_response(stream_chunks_so_far, chunk), + ) + @staticmethod async def async_streaming_data_generator( response: Any, @@ -3659,6 +3680,7 @@ class ProxyBaseLLMRequestProcessing: delivered_chunk = False try: str_so_far = "" + stream_chunks_so_far: tuple[ModelResponseStream, ...] = () async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook( user_api_key_dict=user_api_key_dict, response=response, @@ -3670,11 +3692,13 @@ class ProxyBaseLLMRequestProcessing: verbose_proxy_logger.debug("async_data_generator: received streaming chunk - %s", chunk) if not fast_path: - chunk = await proxy_logging_obj.async_post_call_streaming_hook( + chunk, stream_chunks_so_far = await ProxyBaseLLMRequestProcessing._apply_streaming_chunk_hook( + chunk=chunk, + proxy_logging_obj=proxy_logging_obj, user_api_key_dict=user_api_key_dict, - response=chunk, - data=request_data, + request_data=request_data, str_so_far=str_so_far, + stream_chunks_so_far=stream_chunks_so_far, ) if isinstance(chunk, (ModelResponse, ModelResponseStream)): diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index cdf41ed55bb..97af25d960e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -706,6 +706,7 @@ from litellm.proxy.utils import ( migrate_passwords_to_scrypt_async, model_dump_with_preserved_fields, prefetch_config_params, + stream_chunks_with_response, update_spend, ) from litellm.proxy.video_endpoints.endpoints import router as video_router @@ -8777,19 +8778,23 @@ async def _apply_streaming_chunk_hooks( user_api_key_dict: UserAPIKeyAuth, request_data: dict, str_so_far: str, -) -> tuple[Any, str]: + stream_chunks_so_far: Sequence[ModelResponseStream] = (), +) -> tuple[Any, str, tuple[ModelResponseStream, ...]]: + stream_chunk: Final = chunk chunk = await proxy_logging_obj.async_post_call_streaming_hook( user_api_key_dict=user_api_key_dict, response=chunk, data=request_data, str_so_far=str_so_far if str_so_far else None, + stream_chunks_so_far=stream_chunks_so_far, ) if isinstance(chunk, (ModelResponse, ModelResponseStream)): response_str: Final = litellm.get_response_string(response_obj=chunk) str_so_far += response_str - return chunk, str_so_far + updated_stream_chunks: Final = stream_chunks_with_response(stream_chunks_so_far, stream_chunk) + return chunk, str_so_far, updated_stream_chunks def _format_streaming_sse_chunk(chunk: str | bytes) -> str | bytes: @@ -9042,6 +9047,7 @@ async def async_data_generator( # Previously "".join(str_so_far_parts) was called every chunk, re-joining # the entire accumulated response. String += is O(n) amortized total. _str_so_far: str = "" + _stream_chunks_so_far: tuple[ModelResponseStream, ...] = () # Separate iterator-level vs per-chunk hook decisions. The iterator # wrap is needed when any callback overrides # ``async_post_call_streaming_iterator_hook`` or has @@ -9091,11 +9097,12 @@ async def async_data_generator( chunk = cast(Any, item) # cast-ok: sentinel already handled above, item is a real chunk here if needs_per_chunk_hook: ### CALL HOOKS ### - modify outgoing data - chunk, _str_so_far = await _apply_streaming_chunk_hooks( + chunk, _str_so_far, _stream_chunks_so_far = await _apply_streaming_chunk_hooks( chunk=chunk, user_api_key_dict=user_api_key_dict, request_data=request_data, str_so_far=_str_so_far, + stream_chunks_so_far=_stream_chunks_so_far, ) # Mid-stream fallbacks surface metadata on individual chunks rather than diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 9c9f643d626..c898a5c626c 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -196,7 +196,12 @@ from litellm.types.mcp import ( MCPPreCallResponseObject, ) from litellm.types.proxy.policy_engine.pipeline_types import PipelineExecutionResult -from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams +from litellm.types.utils import ( + ChatCompletionDeltaCustomToolCall, + ChatCompletionDeltaToolCall, + LLMResponseTypes, + LoggedLiteLLMParams, +) from litellm.utils import ( _add_custom_logger_callback_to_specific_event, # pyright: ignore[reportPrivateUsage] # only string-to-logger helper ) @@ -896,6 +901,22 @@ class _StreamingHookResponseText(str): """Marks the exact text object passed to a per-chunk streaming hook.""" +class _StructuredStreamingGuardrailText(_StreamingHookResponseText): + pass + + +@dataclass(frozen=True, slots=True) +class _StreamingToolCallFragment: + choice_index: int + tool_index: int + name: str + arguments: str + + @property + def key(self) -> tuple[int, int]: + return self.choice_index, self.tool_index + + def _streaming_hook_response_text(*, response_str: str, str_so_far: str | None, response: object) -> str: complete_response = str_so_far + response_str if str_so_far is not None else response_str if complete_response == "" and isinstance(response, (ModelResponse, ModelResponseStream)): @@ -903,22 +924,89 @@ def _streaming_hook_response_text(*, response_str: str, str_so_far: str | None, return complete_response -def _streaming_guardrail_response_text(*, complete_response: str, response: object) -> str: - if not isinstance(response, (ModelResponse, ModelResponseStream)): +def _streaming_tool_call_fragments(response: ModelResponseStream) -> tuple[_StreamingToolCallFragment, ...]: + tool_call_fragments: Final = tuple( + _StreamingToolCallFragment( + choice_index=choice.index, + tool_index=tool_call.index, + name=function.name or "", + arguments=function.arguments or "", + ) + for choice in response.choices + for tool_call in choice.delta.tool_calls or () + if isinstance(tool_call, ChatCompletionDeltaToolCall) + for function in (tool_call.function,) + ) + custom_tool_call_fragments: Final = tuple( + _StreamingToolCallFragment( + choice_index=choice.index, + tool_index=tool_call.index, + name=custom.name or "", + arguments=custom.input or "", + ) + for choice in response.choices + for tool_call in choice.delta.tool_calls or () + if isinstance(tool_call, ChatCompletionDeltaCustomToolCall) + for custom in (tool_call.custom,) + ) + function_call_fragments: Final = tuple( + _StreamingToolCallFragment( + choice_index=choice.index, + tool_index=-1, + name=function_call.name or "", + arguments=function_call.arguments or "", + ) + for choice in response.choices + for function_call in (choice.delta.function_call,) + if function_call is not None + ) + return (*tool_call_fragments, *custom_tool_call_fragments, *function_call_fragments) + + +def _assembled_streaming_tool_calls( + stream_chunks: Sequence[ModelResponseStream], +) -> tuple[_StreamingToolCallFragment, ...]: + fragments: Final = tuple(fragment for chunk in stream_chunks for fragment in _streaming_tool_call_fragments(chunk)) + keys: Final = tuple( + fragment.key + for position, fragment in enumerate(fragments) + if fragment.key not in tuple(previous.key for previous in fragments[:position]) + ) + return tuple( + _StreamingToolCallFragment( + choice_index=choice_index, + tool_index=tool_index, + name="".join(fragment.name for fragment in fragments if fragment.key == key), + arguments="".join(fragment.arguments for fragment in fragments if fragment.key == key), + ) + for key in keys + for choice_index, tool_index in (key,) + ) + + +def stream_chunks_with_response( + stream_chunks: Sequence[ModelResponseStream], response: object +) -> tuple[ModelResponseStream, ...]: + return (*stream_chunks, response) if isinstance(response, ModelResponseStream) else tuple(stream_chunks) + + +def _streaming_guardrail_response_text( + *, + complete_response: str, + response: object, + stream_chunks_so_far: Sequence[ModelResponseStream], +) -> str: + if not isinstance(response, ModelResponseStream): return complete_response - response_dict: Final = response.model_dump(mode="json", exclude_none=True) + stream_chunks: Final = stream_chunks_with_response(stream_chunks_so_far, response) + tool_calls: Final = _assembled_streaming_tool_calls(stream_chunks) structured_fields: Final = tuple( - f"{key}:{json.dumps(delta[key], ensure_ascii=False, separators=(',', ':'), sort_keys=True)}" - for choice in response_dict.get("choices", ()) - if isinstance(choice, dict) - for delta in (choice.get("delta"),) - if isinstance(delta, dict) - for key in ("tool_calls", "function_call") - if delta.get(key) is not None + f"tool_call:{json.dumps((tool_call.choice_index, tool_call.tool_index, tool_call.name, tool_call.arguments), ensure_ascii=False, separators=(',', ':'))}" + for tool_call in tool_calls ) if not structured_fields: return complete_response - return _StreamingHookResponseText("\n".join((str(complete_response), *structured_fields))) + return _StructuredStreamingGuardrailText("\n".join((str(complete_response), *structured_fields))) def _is_unchanged_structured_streaming_hook_response( @@ -926,6 +1014,8 @@ def _is_unchanged_structured_streaming_hook_response( ) -> bool: if not isinstance(response, (ModelResponse, ModelResponseStream)): return False + if isinstance(complete_response, _StructuredStreamingGuardrailText): + return callback_response == complete_response if isinstance(complete_response, _StreamingHookResponseText): return callback_response is complete_response if response_str != "": @@ -3488,6 +3578,7 @@ class ProxyLogging: response: ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream, user_api_key_dict: UserAPIKeyAuth, str_so_far: str | None = None, + stream_chunks_so_far: Sequence[ModelResponseStream] = (), ): """ Allow user to modify outgoing streaming data -> per chunk @@ -3563,6 +3654,7 @@ class ProxyLogging: complete_response = _streaming_guardrail_response_text( complete_response=complete_response, response=response, + stream_chunks_so_far=stream_chunks_so_far, ) callback_response: ( str | ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream | None diff --git a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py index 2cf17c57f03..e262b21ffc8 100644 --- a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py +++ b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py @@ -557,12 +557,12 @@ def test_serialize_streaming_chunk_invalid_input_raises_attribute_error(): async def test_apply_streaming_chunk_hooks_appends_to_str_so_far(monkeypatch): chunk = _simple_chunk(content="abc") - async def _passthrough(*, user_api_key_dict, response, data, str_so_far=None): + async def _passthrough(*, user_api_key_dict, response, data, str_so_far=None, stream_chunks_so_far=()): return response monkeypatch.setattr(ps.proxy_logging_obj, "async_post_call_streaming_hook", _passthrough) - new_chunk, new_str = await _apply_streaming_chunk_hooks( + new_chunk, new_str, stream_chunks = await _apply_streaming_chunk_hooks( chunk=chunk, user_api_key_dict=_user_auth(), request_data={}, @@ -573,11 +573,13 @@ async def test_apply_streaming_chunk_hooks_appends_to_str_so_far(monkeypatch): "chunk_is_basemodel": isinstance(new_chunk, ModelResponseStream), "str_so_far": new_str, "grew": len(new_str) > len("prior:"), + "stream_chunks": stream_chunks, } assert observed == { "chunk_is_basemodel": True, "str_so_far": "prior:abc", "grew": True, + "stream_chunks": (chunk,), } diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py index 97573e62b7e..758c2e05a15 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py @@ -384,7 +384,7 @@ async def test_async_post_call_streaming_hook_exposes_tool_calls_to_guardrails( [_BlockingGuardrail(guardrail_name="tool-call-scanner")], ) - response = litellm.ModelResponseStream( + first_chunk = litellm.ModelResponseStream( id="chatcmpl-guardrail-tools", choices=[ { @@ -397,7 +397,79 @@ async def test_async_post_call_streaming_hook_exposes_tool_calls_to_guardrails( "type": "function", "function": { "name": "run_command", - "arguments": '{"command":"blocked-command"}', + "arguments": '{"command":"blocked-', + }, + } + ] + }, + "finish_reason": None, + } + ], + created=0, + model="gpt-4o-mini", + object="chat.completion.chunk", + ) + second_chunk = litellm.ModelResponseStream( + id="chatcmpl-guardrail-tools", + choices=[ + { + "index": 0, + "delta": { + "tool_calls": [ + { + "index": 0, + "function": {"arguments": 'command"}'}, + } + ] + }, + "finish_reason": None, + } + ], + created=0, + model="gpt-4o-mini", + object="chat.completion.chunk", + ) + + out = await proxy_logging.async_post_call_streaming_hook( + data={}, + response=second_chunk, + user_api_key_dict=make_user_api_key_auth(), + stream_chunks_so_far=(first_chunk,), + ) + + assert out == "data: blocked\n\n" + + +@pytest.mark.asyncio +async def test_async_post_call_streaming_hook_preserves_equal_guardrail_projection( + proxy_logging, make_user_api_key_auth, monkeypatch +): + class _CopyingGuardrail(CustomGuardrail): + def should_run_guardrail(self, data, event_type): + return True + + async def async_post_call_streaming_hook(self, user_api_key_dict, response): + return response.encode().decode() + + monkeypatch.setattr( + litellm, + "callbacks", + [_CopyingGuardrail(guardrail_name="tool-call-scanner")], + ) + response = litellm.ModelResponseStream( + id="chatcmpl-guardrail-copy", + choices=[ + { + "index": 0, + "delta": { + "tool_calls": [ + { + "index": 0, + "id": "call-weather", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"city":"Paris"}', }, } ] @@ -416,7 +488,7 @@ async def test_async_post_call_streaming_hook_exposes_tool_calls_to_guardrails( user_api_key_dict=make_user_api_key_auth(), ) - assert out == "data: blocked\n\n" + assert out is response @pytest.mark.asyncio From 0768daadd7b2856b73d576fe7254ce8903f4d437 Mon Sep 17 00:00:00 2001 From: Ultronen <82553854+Ultronen@users.noreply.github.com> Date: Sun, 6 Sep 2026 19:28:52 +0800 Subject: [PATCH 5/7] perf(proxy): compact streaming tool-call state --- litellm/proxy/common_request_processing.py | 24 ++++++--- litellm/proxy/proxy_server.py | 19 +++---- litellm/proxy/utils.py | 50 ++++++++++++----- .../proxy_server/test_streaming_helpers.py | 53 +++++++++++++++++-- .../proxy/test_common_request_processing.py | 42 +++++++++++++++ .../proxy_logging/test_streaming_hooks.py | 4 +- 6 files changed, 155 insertions(+), 37 deletions(-) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 460d3a5a051..84f53fc998a 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -70,7 +70,12 @@ from litellm.proxy.common_utils.sse_keepalive import ( from litellm.proxy.dd_span_tagger import DDSpanTagger from litellm.proxy.guardrails.auto_router_compression import arm_pre_call as _arm_auto_router_compression from litellm.proxy.route_llm_request import route_request -from litellm.proxy.utils import ProxyLogging, _check_and_merge_model_level_guardrails, stream_chunks_with_response +from litellm.proxy.utils import ( + ProxyLogging, + StreamingToolCallState, + _check_and_merge_model_level_guardrails, + streaming_tool_calls_with_response, +) from litellm.router import Router from litellm.router_utils.add_retry_fallback_headers import get_hidden_params_dict from litellm.router_utils.common_utils import resolve_model_group_alias @@ -3629,17 +3634,17 @@ class ProxyBaseLLMRequestProcessing: user_api_key_dict: UserAPIKeyAuth, request_data: dict, # mutable-ok: ProxyLogging requires the existing mutable request payload str_so_far: str, - stream_chunks_so_far: Sequence[ModelResponseStream], - ) -> tuple[Any, tuple[ModelResponseStream, ...]]: + streaming_tool_calls_so_far: StreamingToolCallState, + ) -> tuple[Any, StreamingToolCallState]: return ( await proxy_logging_obj.async_post_call_streaming_hook( user_api_key_dict=user_api_key_dict, response=chunk, data=request_data, str_so_far=str_so_far, - stream_chunks_so_far=stream_chunks_so_far, + streaming_tool_calls_so_far=streaming_tool_calls_so_far, ), - stream_chunks_with_response(stream_chunks_so_far, chunk), + streaming_tool_calls_with_response(streaming_tool_calls_so_far, chunk), ) @staticmethod @@ -3680,7 +3685,7 @@ class ProxyBaseLLMRequestProcessing: delivered_chunk = False try: str_so_far = "" - stream_chunks_so_far: tuple[ModelResponseStream, ...] = () + streaming_tool_calls_so_far: StreamingToolCallState = () async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook( user_api_key_dict=user_api_key_dict, response=response, @@ -3692,13 +3697,16 @@ class ProxyBaseLLMRequestProcessing: verbose_proxy_logger.debug("async_data_generator: received streaming chunk - %s", chunk) if not fast_path: - chunk, stream_chunks_so_far = await ProxyBaseLLMRequestProcessing._apply_streaming_chunk_hook( + ( + chunk, + streaming_tool_calls_so_far, + ) = await ProxyBaseLLMRequestProcessing._apply_streaming_chunk_hook( chunk=chunk, proxy_logging_obj=proxy_logging_obj, user_api_key_dict=user_api_key_dict, request_data=request_data, str_so_far=str_so_far, - stream_chunks_so_far=stream_chunks_so_far, + streaming_tool_calls_so_far=streaming_tool_calls_so_far, ) if isinstance(chunk, (ModelResponse, ModelResponseStream)): diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 97af25d960e..1d379eb3211 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -686,6 +686,7 @@ from litellm.proxy.utils import ( PrismaClient, ProxyLogging, ProxyUpdateSpend, + StreamingToolCallState, _cache_user_row, _get_docs_url, _get_openapi_url, @@ -706,7 +707,7 @@ from litellm.proxy.utils import ( migrate_passwords_to_scrypt_async, model_dump_with_preserved_fields, prefetch_config_params, - stream_chunks_with_response, + streaming_tool_calls_with_response, update_spend, ) from litellm.proxy.video_endpoints.endpoints import router as video_router @@ -8778,23 +8779,23 @@ async def _apply_streaming_chunk_hooks( user_api_key_dict: UserAPIKeyAuth, request_data: dict, str_so_far: str, - stream_chunks_so_far: Sequence[ModelResponseStream] = (), -) -> tuple[Any, str, tuple[ModelResponseStream, ...]]: + streaming_tool_calls_so_far: StreamingToolCallState = (), +) -> tuple[Any, str, StreamingToolCallState]: stream_chunk: Final = chunk chunk = await proxy_logging_obj.async_post_call_streaming_hook( user_api_key_dict=user_api_key_dict, response=chunk, data=request_data, str_so_far=str_so_far if str_so_far else None, - stream_chunks_so_far=stream_chunks_so_far, + streaming_tool_calls_so_far=streaming_tool_calls_so_far, ) if isinstance(chunk, (ModelResponse, ModelResponseStream)): response_str: Final = litellm.get_response_string(response_obj=chunk) str_so_far += response_str - updated_stream_chunks: Final = stream_chunks_with_response(stream_chunks_so_far, stream_chunk) - return chunk, str_so_far, updated_stream_chunks + updated_tool_calls: Final = streaming_tool_calls_with_response(streaming_tool_calls_so_far, stream_chunk) + return chunk, str_so_far, updated_tool_calls def _format_streaming_sse_chunk(chunk: str | bytes) -> str | bytes: @@ -9047,7 +9048,7 @@ async def async_data_generator( # Previously "".join(str_so_far_parts) was called every chunk, re-joining # the entire accumulated response. String += is O(n) amortized total. _str_so_far: str = "" - _stream_chunks_so_far: tuple[ModelResponseStream, ...] = () + _streaming_tool_calls_so_far: StreamingToolCallState = () # Separate iterator-level vs per-chunk hook decisions. The iterator # wrap is needed when any callback overrides # ``async_post_call_streaming_iterator_hook`` or has @@ -9097,12 +9098,12 @@ async def async_data_generator( chunk = cast(Any, item) # cast-ok: sentinel already handled above, item is a real chunk here if needs_per_chunk_hook: ### CALL HOOKS ### - modify outgoing data - chunk, _str_so_far, _stream_chunks_so_far = await _apply_streaming_chunk_hooks( + chunk, _str_so_far, _streaming_tool_calls_so_far = await _apply_streaming_chunk_hooks( chunk=chunk, user_api_key_dict=user_api_key_dict, request_data=request_data, str_so_far=_str_so_far, - stream_chunks_so_far=_stream_chunks_so_far, + streaming_tool_calls_so_far=_streaming_tool_calls_so_far, ) # Mid-stream fallbacks surface metadata on individual chunks rather than diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index c898a5c626c..80157b2768b 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -17,7 +17,20 @@ from datetime import date, datetime, timedelta, timezone from email.mime.multipart import MIMEMultipart from email.mime.text import MIMEText from types import MappingProxyType -from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, Protocol, TypeVar, Union, cast, overload +from typing import ( + TYPE_CHECKING, + Any, + ClassVar, + Final, + Literal, + Optional, + Protocol, + TypeAlias, + TypeVar, + Union, + cast, + overload, +) from typing_extensions import ReadOnly, TypedDict @@ -917,6 +930,9 @@ class _StreamingToolCallFragment: return self.choice_index, self.tool_index +StreamingToolCallState: TypeAlias = tuple[_StreamingToolCallFragment, ...] + + def _streaming_hook_response_text(*, response_str: str, str_so_far: str | None, response: object) -> str: complete_response = str_so_far + response_str if str_so_far is not None else response_str if complete_response == "" and isinstance(response, (ModelResponse, ModelResponseStream)): @@ -964,9 +980,8 @@ def _streaming_tool_call_fragments(response: ModelResponseStream) -> tuple[_Stre def _assembled_streaming_tool_calls( - stream_chunks: Sequence[ModelResponseStream], -) -> tuple[_StreamingToolCallFragment, ...]: - fragments: Final = tuple(fragment for chunk in stream_chunks for fragment in _streaming_tool_call_fragments(chunk)) + fragments: Sequence[_StreamingToolCallFragment], +) -> StreamingToolCallState: keys: Final = tuple( fragment.key for position, fragment in enumerate(fragments) @@ -984,24 +999,31 @@ def _assembled_streaming_tool_calls( ) -def stream_chunks_with_response( - stream_chunks: Sequence[ModelResponseStream], response: object -) -> tuple[ModelResponseStream, ...]: - return (*stream_chunks, response) if isinstance(response, ModelResponseStream) else tuple(stream_chunks) +def streaming_tool_calls_with_response( + tool_calls: Sequence[_StreamingToolCallFragment], response: object +) -> StreamingToolCallState: + if not isinstance(response, ModelResponseStream): + return tuple(tool_calls) + fragments: Final = _streaming_tool_call_fragments(response) + return _assembled_streaming_tool_calls((*tool_calls, *fragments)) def _streaming_guardrail_response_text( *, complete_response: str, response: object, - stream_chunks_so_far: Sequence[ModelResponseStream], + streaming_tool_calls_so_far: Sequence[_StreamingToolCallFragment], ) -> str: if not isinstance(response, ModelResponseStream): return complete_response - stream_chunks: Final = stream_chunks_with_response(stream_chunks_so_far, response) - tool_calls: Final = _assembled_streaming_tool_calls(stream_chunks) + tool_calls: Final = streaming_tool_calls_with_response(streaming_tool_calls_so_far, response) structured_fields: Final = tuple( - f"tool_call:{json.dumps((tool_call.choice_index, tool_call.tool_index, tool_call.name, tool_call.arguments), ensure_ascii=False, separators=(',', ':'))}" + "tool_call:" + + json.dumps( + (tool_call.choice_index, tool_call.tool_index, tool_call.name, tool_call.arguments), + ensure_ascii=False, + separators=(",", ":"), + ) for tool_call in tool_calls ) if not structured_fields: @@ -3578,7 +3600,7 @@ class ProxyLogging: response: ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream, user_api_key_dict: UserAPIKeyAuth, str_so_far: str | None = None, - stream_chunks_so_far: Sequence[ModelResponseStream] = (), + streaming_tool_calls_so_far: Sequence[_StreamingToolCallFragment] = (), ): """ Allow user to modify outgoing streaming data -> per chunk @@ -3654,7 +3676,7 @@ class ProxyLogging: complete_response = _streaming_guardrail_response_text( complete_response=complete_response, response=response, - stream_chunks_so_far=stream_chunks_so_far, + streaming_tool_calls_so_far=streaming_tool_calls_so_far, ) callback_response: ( str | ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream | None diff --git a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py index e262b21ffc8..2fbb42853dc 100644 --- a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py +++ b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py @@ -68,6 +68,30 @@ def _simple_chunk(model: str = "gpt-4", content: str = "hi") -> ModelResponseStr ) +def _tool_call_chunk(arguments: str) -> ModelResponseStream: + return ModelResponseStream( + id="chatcmpl-tool-call", + choices=[ + { + "index": 0, + "delta": { + "tool_calls": [ + { + "index": 0, + "type": "function", + "function": {"name": "run_command", "arguments": arguments}, + } + ] + }, + "finish_reason": None, + } + ], + created=0, + model="gpt-4", + object="chat.completion.chunk", + ) + + async def _async_iter(items): for it in items: yield it @@ -557,12 +581,12 @@ def test_serialize_streaming_chunk_invalid_input_raises_attribute_error(): async def test_apply_streaming_chunk_hooks_appends_to_str_so_far(monkeypatch): chunk = _simple_chunk(content="abc") - async def _passthrough(*, user_api_key_dict, response, data, str_so_far=None, stream_chunks_so_far=()): + async def _passthrough(*, user_api_key_dict, response, data, str_so_far=None, streaming_tool_calls_so_far=()): return response monkeypatch.setattr(ps.proxy_logging_obj, "async_post_call_streaming_hook", _passthrough) - new_chunk, new_str, stream_chunks = await _apply_streaming_chunk_hooks( + new_chunk, new_str, tool_calls = await _apply_streaming_chunk_hooks( chunk=chunk, user_api_key_dict=_user_auth(), request_data={}, @@ -573,16 +597,37 @@ async def test_apply_streaming_chunk_hooks_appends_to_str_so_far(monkeypatch): "chunk_is_basemodel": isinstance(new_chunk, ModelResponseStream), "str_so_far": new_str, "grew": len(new_str) > len("prior:"), - "stream_chunks": stream_chunks, + "tool_calls": tool_calls, } assert observed == { "chunk_is_basemodel": True, "str_so_far": "prior:abc", "grew": True, - "stream_chunks": (chunk,), + "tool_calls": (), } +@pytest.mark.asyncio +async def test_apply_streaming_chunk_hooks_compacts_tool_call_history(monkeypatch): + async def _passthrough(*, user_api_key_dict, response, data, str_so_far=None, streaming_tool_calls_so_far=()): + return response + + monkeypatch.setattr(ps.proxy_logging_obj, "async_post_call_streaming_hook", _passthrough) + stream_state = () + + for arguments in ('{"command":"', "blocked-", "command", '"}'): + _, _, stream_state = await _apply_streaming_chunk_hooks( + chunk=_tool_call_chunk(arguments), + user_api_key_dict=_user_auth(), + request_data={}, + str_so_far="", + streaming_tool_calls_so_far=stream_state, + ) + + assert len(stream_state) == 1 + assert stream_state[0].arguments == '{"command":"blocked-command"}' + + @pytest.mark.asyncio async def test_apply_streaming_chunk_hooks_hook_raises_exception(monkeypatch): async def _boom(*args, **kwargs): diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 69e89d1c604..f506f09e6fe 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -4055,6 +4055,48 @@ class TestAsyncStreamingDataGeneratorFastPath: ProxyLogging._callback_capabilities_cache.clear() + @pytest.mark.asyncio + async def test_apply_streaming_chunk_hook_compacts_tool_call_history(self): + proxy_logging_obj = MagicMock(spec=ProxyLogging) + proxy_logging_obj.async_post_call_streaming_hook = AsyncMock(side_effect=lambda **kwargs: kwargs["response"]) + + def _chunk(arguments: str): + return litellm.ModelResponseStream( + id="chatcmpl-tool-call", + choices=[ + { + "index": 0, + "delta": { + "tool_calls": [ + { + "index": 0, + "type": "function", + "function": {"name": "run_command", "arguments": arguments}, + } + ] + }, + "finish_reason": None, + } + ], + created=0, + model="gpt-4", + object="chat.completion.chunk", + ) + + stream_state = () + for arguments in ("blocked-", "command"): + _, stream_state = await ProxyBaseLLMRequestProcessing._apply_streaming_chunk_hook( + chunk=_chunk(arguments), + proxy_logging_obj=proxy_logging_obj, + user_api_key_dict=MagicMock(spec=ProxyUserAPIKeyAuth), + request_data={}, + str_so_far="", + streaming_tool_calls_so_far=stream_state, + ) + + assert len(stream_state) == 1 + assert stream_state[0].arguments == "blocked-command" + class TestDisconnectGatherCleanup: def _disconnect_request(self) -> Request: diff --git a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py index 758c2e05a15..6c786190bfb 100644 --- a/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py +++ b/tests/test_litellm/proxy/utils/proxy_logging/test_streaming_hooks.py @@ -24,7 +24,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( BaseAnthropicMessagesStreamingIterator, ) -from litellm.proxy.utils import ProxyLogging +from litellm.proxy.utils import ProxyLogging, streaming_tool_calls_with_response @pytest.fixture(autouse=True) @@ -434,7 +434,7 @@ async def test_async_post_call_streaming_hook_exposes_tool_calls_to_guardrails( data={}, response=second_chunk, user_api_key_dict=make_user_api_key_auth(), - stream_chunks_so_far=(first_chunk,), + streaming_tool_calls_so_far=streaming_tool_calls_with_response((), first_chunk), ) assert out == "data: blocked\n\n" From ce3612e316ca9dec46566e42726cc1dcdf0894ca Mon Sep 17 00:00:00 2001 From: Ultronen <82553854+Ultronen@users.noreply.github.com> Date: Wed, 9 Sep 2026 11:32:09 +0800 Subject: [PATCH 6/7] refactor(proxy): simplify streaming hook response guard --- litellm/proxy/utils.py | 17 ++++++----------- 1 file changed, 6 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 80157b2768b..13f1071f1fa 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -3685,17 +3685,12 @@ class ProxyLogging: user_api_key_dict=user_api_key_dict, response=complete_response, ) - if callback_response is not None: - # A text result cannot represent a structured empty-text - # chunk such as a tool-call delta. Preserve the chunk - # only when the callback returned its input unchanged. - if _is_unchanged_structured_streaming_hook_response( - callback_response=callback_response, - complete_response=complete_response, - response_str=response_str, - response=response, - ): - continue + if callback_response is not None and not _is_unchanged_structured_streaming_hook_response( + callback_response=callback_response, + complete_response=complete_response, + response_str=response_str, + response=response, + ): response = callback_response except Exception as e: raise e From 3674fb5145941f0b53e70dbf8ff2bfaedd1693ba Mon Sep 17 00:00:00 2001 From: Ultronen <82553854+Ultronen@users.noreply.github.com> Date: Mon, 14 Sep 2026 17:13:30 +0800 Subject: [PATCH 7/7] chore(ui): refresh generated API types --- ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 -- 1 file changed, 2 deletions(-) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index f5a1733d76b..fe38ca2a81a 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -16781,7 +16781,6 @@ export interface paths { * - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking. * - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } * - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. - * - soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests. * - model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys) * - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. * - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys) @@ -16918,7 +16917,6 @@ export interface paths { * - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking. * - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } * - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. - * - soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests. * - model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys) * - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. * - model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)