fix(proxy): scan accumulated streaming tool calls

This commit is contained in:
Ultronen 2026-09-06 19:03:59 +08:00
parent 9f231af270
commit 02e4307523
5 changed files with 221 additions and 24 deletions

View file

@ -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)):

View file

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

View file

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

View file

@ -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,),
}

View file

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