diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 97de2488b8d..ea23a3ebda3 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -114,7 +114,12 @@ 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.native_compaction import with_proxy_compaction_executor 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, + 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 @@ -3856,6 +3861,27 @@ class ProxyBaseLLMRequestProcessing: ): await logging_obj.invalidate_baseline_cache_estimate("incomplete_response", completed=True) + @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, + 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, + streaming_tool_calls_so_far=streaming_tool_calls_so_far, + ), + streaming_tool_calls_with_response(streaming_tool_calls_so_far, chunk), + ) + @staticmethod async def async_streaming_data_generator( response: object, @@ -3902,6 +3928,7 @@ class ProxyBaseLLMRequestProcessing: recent_tail = SSE_STREAM_START_TAIL # rebind-ok: rolling window over the yielded bytes try: str_so_far = "" + 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, @@ -3913,11 +3940,16 @@ 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, + 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, - response=chunk, - data=request_data, + request_data=request_data, str_so_far=str_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 7cd116c23d7..7e2d16cf0d2 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -793,6 +793,7 @@ from litellm.proxy.utils import ( PrismaClient, ProxyLogging, ProxyUpdateSpend, + StreamingToolCallState, _cache_user_row, _get_docs_url, _get_openapi_url, @@ -813,6 +814,7 @@ from litellm.proxy.utils import ( migrate_passwords_to_scrypt_async, model_dump_with_preserved_fields, prefetch_config_params, + streaming_tool_calls_with_response, update_spend, ) from litellm.proxy.video_endpoints.endpoints import router as video_router @@ -9338,19 +9340,23 @@ async def _apply_streaming_chunk_hooks( user_api_key_dict: UserAPIKeyAuth, request_data: dict, str_so_far: str, -) -> tuple[Any, str]: + 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, + 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 - return chunk, str_so_far + 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: @@ -9607,6 +9613,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 = "" + _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 @@ -9656,11 +9663,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, _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, + 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 ea294b76e92..19d66137d59 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -19,6 +19,7 @@ from collections.abc import ( Awaitable, Callable, Coroutine, + Iterator, Mapping, Sequence, ) @@ -246,7 +247,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 ) @@ -1084,6 +1090,127 @@ 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.""" + + +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 + + +StreamingToolCallState: TypeAlias = tuple[_StreamingToolCallFragment, ...] + + +def _streaming_hook_response_text(*, response_str: str, str_so_far: str | None, response: object) -> str: + complete_response: Final = 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 _streaming_tool_call_fragments(response: ModelResponseStream) -> Iterator[_StreamingToolCallFragment]: + for choice in response.choices: + for tool_call in choice.delta.tool_calls or (): + if isinstance(tool_call, ChatCompletionDeltaToolCall): + yield _StreamingToolCallFragment( + choice_index=choice.index, + tool_index=tool_call.index, + name=tool_call.function.name or "", + arguments=tool_call.function.arguments or "", + ) + elif isinstance(tool_call, ChatCompletionDeltaCustomToolCall): + yield _StreamingToolCallFragment( + choice_index=choice.index, + tool_index=tool_call.index, + name=tool_call.custom.name or "", + arguments=tool_call.custom.input or "", + ) + if choice.delta.function_call is not None: + yield _StreamingToolCallFragment( + choice_index=choice.index, + tool_index=-1, + name=choice.delta.function_call.name or "", + arguments=choice.delta.function_call.arguments or "", + ) + + +def _assembled_streaming_tool_calls( + fragments: Sequence[_StreamingToolCallFragment], +) -> StreamingToolCallState: + 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=key[0], + tool_index=key[1], + 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 + ) + + +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, + streaming_tool_calls_so_far: Sequence[_StreamingToolCallFragment], +) -> str: + if not isinstance(response, ModelResponseStream): + return complete_response + tool_calls: Final = streaming_tool_calls_with_response(streaming_tool_calls_so_far, response) + structured_fields: Final = tuple( + "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 _StructuredStreamingGuardrailText("\n".join((str(complete_response), *structured_fields))) + + +def _is_unchanged_structured_streaming_hook_response( + *, callback_response: object, complete_response: str, response_str: str, response: object +) -> 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 != "": + return False + return callback_response == complete_response + + _PROXY_ONLY_LLM_API_ERRORS: Final = (HTTPException, ProxyException, GuardrailRaisedException) @@ -3812,6 +3939,7 @@ class ProxyLogging: response: ModelResponse | EmbeddingResponse | ImageResponse | ModelResponseStream, user_api_key_dict: UserAPIKeyAuth, str_so_far: str | None = None, + streaming_tool_calls_so_far: Sequence[_StreamingToolCallFragment] = (), ): """ Allow user to modify outgoing streaming data -> per chunk @@ -3878,18 +4006,30 @@ 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, + ) + if isinstance(_callback, CustomGuardrail): + complete_response = _streaming_guardrail_response_text( + complete_response=complete_response, + response=response, + streaming_tool_calls_so_far=streaming_tool_calls_so_far, + ) 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: + if callback_response is not None and not _is_unchanged_structured_streaming_hook_response( + callback_response=callback_response, + complete_response=complete_response, + response_str=response_str, + response=response, + ): 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 86dd356e5f5..f12a16508c0 100644 --- a/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py +++ b/tests/test_litellm/proxy/proxy_server/test_streaming_helpers.py @@ -31,6 +31,7 @@ from pydantic import BaseModel import litellm 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, @@ -78,6 +79,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 @@ -567,12 +592,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, streaming_tool_calls_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, tool_calls = await _apply_streaming_chunk_hooks( chunk=chunk, user_api_key_dict=_user_auth(), request_data={}, @@ -583,14 +608,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:"), + "tool_calls": tool_calls, } assert observed == { "chunk_is_basemodel": True, "str_so_far": "prior:abc", "grew": True, + "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): @@ -693,6 +741,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/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index c17f41a8b8f..d8d264c4289 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -4756,6 +4756,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 50f50478ad3..414411956db 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 @@ -27,7 +27,7 @@ from litellm.llms.anthropic.pass_through.messages.streaming_iterator import ( BaseAnthropicMessagesStreamingIterator, ) from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.utils import ProxyLogging +from litellm.proxy.utils import ProxyLogging, streaming_tool_calls_with_response from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import Usage @@ -384,6 +384,239 @@ 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_exposes_tool_calls_to_guardrails( + proxy_logging, make_user_api_key_auth, monkeypatch +): + class _BlockingGuardrail(CustomGuardrail): + def should_run_guardrail(self, data, event_type): + return True + + async def async_post_call_streaming_hook(self, user_api_key_dict, response): + if "blocked-command" in response: + return "data: blocked\n\n" + return response + + monkeypatch.setattr( + litellm, + "callbacks", + [_BlockingGuardrail(guardrail_name="tool-call-scanner")], + ) + + first_chunk = litellm.ModelResponseStream( + id="chatcmpl-guardrail-tools", + choices=[ + { + "index": 0, + "delta": { + "tool_calls": [ + { + "index": 0, + "id": "call-shell", + "type": "function", + "function": { + "name": "run_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(), + streaming_tool_calls_so_far=streaming_tool_calls_with_response((), 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"}', + }, + } + ] + }, + "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(), + ) + + assert out is response + + @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):