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] 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