fix: commit of buffered streaming - still not working

This commit is contained in:
Krrish Dholakia 2026-02-02 16:04:12 -08:00
parent e96415b802
commit 1cefc73c90
7 changed files with 88 additions and 2 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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