From 1cefc73c90c45ef5d33faa346742214662038c1f Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Mon, 2 Feb 2026 16:04:12 -0800 Subject: [PATCH] fix: commit of buffered streaming - still not working --- .../litellm_core_utils/get_litellm_params.py | 1 + .../litellm_core_utils/realtime_streaming.py | 45 ++++++++++++++++++- litellm/proxy/_new_secret_config.yaml | 3 +- litellm/proxy/proxy_server.py | 4 ++ litellm/proxy/utils.py | 33 ++++++++++++++ litellm/types/router.py | 3 ++ litellm/types/utils.py | 1 + 7 files changed, 88 insertions(+), 2 deletions(-) diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index 060e98fd49f..71dd05e5285 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -144,5 +144,6 @@ def get_litellm_params( "aws_bedrock_runtime_endpoint": kwargs.get("aws_bedrock_runtime_endpoint"), "tpm": kwargs.get("tpm"), "rpm": kwargs.get("rpm"), + "has_post_call_guardrails": kwargs.get("has_post_call_guardrails"), } return litellm_params diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 158b318f02f..3cbd0c5c911 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -73,6 +73,11 @@ class RealTimeStreaming: self.current_item_chunks: Optional[List[OpenAIRealtimeOutputItemDone]] = None self.current_delta_type: Optional[ALL_DELTA_TYPES] = None self.session_configuration_request: Optional[str] = None + self.buffered_messages: List = [] + self.is_buffering = False + + def _should_buffer_message(self) -> bool: + return self.logging_obj.litellm_params.get("has_post_call_guardrails", False) def _should_store_message( self, @@ -142,9 +147,43 @@ class RealTimeStreaming: response=message, ) + async def buffer_message(self, message: OpenAIRealtimeEvents): + """ + Only buffer if response.content_part.added has part type "audio" - message_json: {'type': 'response.content_part.added', 'event_id': 'event_D4ugsph441vhnQJ6Wn9KV', 'response_id': 'resp_D4ugsW02oco5ULadjJsQG', 'item_id': 'item_D4ugsRGmxeDXOkpeSHc5r', 'output_index': 0, 'content_index': 0, 'part': {'type': 'audio', 'transcript': ''}} + + Buffer until response.audio_transcript.done is received. - message_json: {'type': 'response.audio_transcript.done', 'event_id': 'event_D4ugu5A33JruyingR1cpZ', 'response_id': 'resp_D4ugsW02oco5ULadjJsQG', 'item_id': 'item_D4ugsRGmxeDXOkpeSHc5r', 'output_index': 0, 'content_index': 0, 'transcript': "That sounds like the famous opening of the Gettysburg Address by Abraham Lincoln. It's a powerful and historic speech. Would you like to discuss it further or explore its meaning?"} + """ + + if isinstance(message, bytes): + message_obj = json.loads(message.decode("utf-8")) + if isinstance(message, dict): + message_obj = message + else: + message_obj = json.loads(message) + if ( + message_obj.get("type") == "response.content_part.added" + and message_obj.get("part", {}).get("type") == "audio" + ): + self.is_buffering = True + + if message_obj.get("type") == "response.audio_transcript.done": + + for buffered_message, original_message_str in self.buffered_messages: + + await self.websocket.send_text(original_message_str) + self.buffered_messages = [] + self.is_buffering = False + await self.websocket.send_text(message) + + if self.is_buffering: + self.buffered_messages.append((message_obj, message)) + else: + await self.websocket.send_text(message) + async def backend_to_client_send_messages(self): import websockets + should_buffer_message = self._should_buffer_message() try: while True: try: @@ -197,8 +236,12 @@ class RealTimeStreaming: await self.websocket.send_text(event_str) else: - ## LOGGING await self.store_and_check_message(raw_response) + ## LOGGING + # if should_buffer_message: + # print("REACHES HERE") + # await self.buffer_message(raw_response) + # else: await self.websocket.send_text(raw_response) except websockets.exceptions.ConnectionClosed as e: # type: ignore diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 2733d98159d..6093a3595bd 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -27,7 +27,8 @@ guardrails: litellm_params: guardrail: litellm_content_filter mode: "post_call" + default_on: true blocked_words: - - keyword: "lincoln" + - keyword: "apple" action: "BLOCK" description: "Do not talk about Lincoln" \ No newline at end of file diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 4450716d476..49ec0f02f6d 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6564,6 +6564,10 @@ async def realtime_websocket_endpoint( # PASS post-call guardrails to the route request for realtime requests data["proxy_logging_obj"] = proxy_logging_obj + if proxy_logging_obj.post_call_guardrail_exists(data=data): + + data["has_post_call_guardrails"] = True + llm_call = await route_request( data=data, route_type="_arealtime", diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 6bbf0df74de..7271b698f78 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1725,6 +1725,39 @@ class ProxyLogging: ), ).start() + def post_call_guardrail_exists(self, data: dict) -> bool: + """ + Return True if post_call_guardrail exists in litellm.callbacks + """ + guardrail_callbacks: List[CustomGuardrail] = [] + for callback in litellm.callbacks: + _callback: Optional[CustomLogger] = None + if isinstance(callback, str): + _callback = litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class( + cast(_custom_logger_compatible_callbacks_literal, callback) + ) + else: + _callback = callback # type: ignore + + if _callback is not None: + if isinstance(_callback, CustomGuardrail): + guardrail_callbacks.append(_callback) + ############## Handle Guardrails ######################################## + ############################################################################# + + for callback in guardrail_callbacks: + # Main - V2 Guardrails implementation + + if ( + callback.should_run_guardrail( + data=data, event_type=GuardrailEventHooks.post_call + ) + is True + ): + return True + + return False + async def post_call_success_hook( self, data: dict, diff --git a/litellm/types/router.py b/litellm/types/router.py index f31c6df3005..93b6d2c6af1 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -213,6 +213,9 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams): vector_store_id: Optional[str] = None milvus_text_field: Optional[str] = None + # Guardrails Params - used for realtime streaming + has_post_call_guardrails: Optional[bool] = False + def __init__( self, custom_llm_provider: Optional[str] = None, diff --git a/litellm/types/utils.py b/litellm/types/utils.py index b396e156973..0fb9379abb1 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2938,6 +2938,7 @@ all_litellm_params = ( "shared_session", "search_tool_name", "order", + "has_post_call_guardrails", ] + list(StandardCallbackDynamicParams.__annotations__.keys()) + list(CustomPricingLiteLLMParams.model_fields.keys())