perf(proxy): compact streaming tool-call state

This commit is contained in:
Ultronen 2026-09-06 19:28:52 +08:00
parent 02e4307523
commit 0768daadd7
6 changed files with 155 additions and 37 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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