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