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"