mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: commit of buffered streaming - still not working
This commit is contained in:
parent
e96415b802
commit
1cefc73c90
7 changed files with 88 additions and 2 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue