mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix lint
This commit is contained in:
parent
b719e35ed9
commit
817792c6cc
3 changed files with 80 additions and 8 deletions
|
|
@ -78,6 +78,9 @@ class RealTimeStreaming:
|
|||
# Track whether session.created has already been sent to the client
|
||||
# (e.g. synthetic event in deferred setup mode).
|
||||
self._session_created_sent_to_client: bool = False
|
||||
# Track whether we have already sent the guardrail turn-detection update
|
||||
# that disables provider auto-response for transcription guardrails.
|
||||
self._guardrail_turn_detection_update_sent: bool = False
|
||||
|
||||
def _should_store_message(
|
||||
self,
|
||||
|
|
@ -245,6 +248,22 @@ class RealTimeStreaming:
|
|||
except (json.JSONDecodeError, TypeError):
|
||||
return
|
||||
|
||||
async def _maybe_send_guardrail_turn_detection_update(self) -> None:
|
||||
"""Disable provider auto-response once when transcription guardrails are enabled."""
|
||||
if self._guardrail_turn_detection_update_sent:
|
||||
return
|
||||
if not self._has_audio_transcription_guardrails():
|
||||
return
|
||||
await self._send_to_backend(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "session.update",
|
||||
"session": {"turn_detection": {"create_response": False}},
|
||||
}
|
||||
)
|
||||
)
|
||||
self._guardrail_turn_detection_update_sent = True
|
||||
|
||||
def _has_realtime_guardrails(self) -> bool:
|
||||
"""Return True if any callback is registered for realtime guardrail event types."""
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
|
@ -446,6 +465,7 @@ class RealTimeStreaming:
|
|||
event_str = json.dumps(event)
|
||||
if isinstance(event, dict) and event.get("type") == "session.created":
|
||||
if self._session_created_sent_to_client:
|
||||
await self._maybe_send_guardrail_turn_detection_update()
|
||||
verbose_logger.debug(
|
||||
"Skipping duplicate session.created from provider stream"
|
||||
)
|
||||
|
|
@ -459,14 +479,7 @@ class RealTimeStreaming:
|
|||
):
|
||||
self.store_message(event_str)
|
||||
await self.websocket.send_text(event_str)
|
||||
await self._send_to_backend(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "session.update",
|
||||
"session": {"turn_detection": {"create_response": False}},
|
||||
}
|
||||
)
|
||||
)
|
||||
await self._maybe_send_guardrail_turn_detection_update()
|
||||
continue
|
||||
## GUARDRAIL: run on transcription events in provider_config path too
|
||||
if (
|
||||
|
|
|
|||
|
|
@ -5167,6 +5167,7 @@ class BaseLLMHTTPHandler:
|
|||
if synthetic_session is not None:
|
||||
await websocket.send_text(json.dumps(synthetic_session))
|
||||
realtime_streaming._session_created_sent_to_client = True
|
||||
await realtime_streaming._maybe_send_guardrail_turn_detection_update()
|
||||
verbose_logger.debug(
|
||||
"Sent synthetic session.created to client to unblock connection"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1022,3 +1022,61 @@ async def test_provider_path_suppresses_duplicate_session_created_after_syntheti
|
|||
assert not any(
|
||||
payload.get("type") == "session.created" for payload in sent_payloads
|
||||
), f"Expected duplicate session.created to be suppressed, got: {sent_payloads}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_session_created_still_triggers_guardrail_turn_detection_update():
|
||||
client_ws = MagicMock()
|
||||
client_ws.send_text = AsyncMock()
|
||||
|
||||
backend_ws = MagicMock()
|
||||
backend_ws.recv = AsyncMock(
|
||||
side_effect=[b'{"setupComplete": {}}', ConnectionClosed(None, None)]
|
||||
)
|
||||
backend_ws.send = AsyncMock()
|
||||
|
||||
provider_config = MagicMock()
|
||||
provider_config.transform_realtime_response = MagicMock(
|
||||
return_value={
|
||||
"response": [
|
||||
{
|
||||
"type": "session.created",
|
||||
"event_id": "event_1",
|
||||
"session": {"id": "sess_1", "modalities": ["audio"]},
|
||||
}
|
||||
],
|
||||
"current_output_item_id": None,
|
||||
"current_response_id": None,
|
||||
"current_delta_chunks": [],
|
||||
"current_conversation_id": None,
|
||||
"current_item_chunks": [],
|
||||
"current_delta_type": None,
|
||||
"session_configuration_request": None,
|
||||
}
|
||||
)
|
||||
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.litellm_trace_id = "trace_1"
|
||||
logging_obj.async_success_handler = AsyncMock()
|
||||
logging_obj.success_handler = MagicMock()
|
||||
|
||||
streaming = RealTimeStreaming(
|
||||
websocket=client_ws,
|
||||
backend_ws=backend_ws,
|
||||
logging_obj=logging_obj,
|
||||
provider_config=provider_config,
|
||||
model="gemini-2.5-flash",
|
||||
)
|
||||
# Synthetic session.created already sent by llm_http_handler.
|
||||
streaming._session_created_sent_to_client = True
|
||||
streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign]
|
||||
streaming._send_to_backend = AsyncMock() # type: ignore[method-assign]
|
||||
|
||||
await streaming.backend_to_client_send_messages()
|
||||
|
||||
# Duplicate session.created should still cause the one-time guardrail
|
||||
# turn_detection update to be sent to backend.
|
||||
assert streaming._send_to_backend.await_count == 1
|
||||
sent_update = json.loads(streaming._send_to_backend.await_args_list[0].args[0])
|
||||
assert sent_update["type"] == "session.update"
|
||||
assert sent_update["session"]["turn_detection"]["create_response"] is False
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue