From 564bb3a7b19e22b2c77544369d03bf412cb18767 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Mon, 23 Feb 2026 20:54:02 -0800 Subject: [PATCH] fix: address Greptile review comments - Forward user_api_key_dict through realtime_api/main.py (_arealtime) so it actually reaches RealTimeStreaming instead of always being None - Run guardrail interception in provider_config path too (e.g. Gemini), not only the OpenAI direct path - Narrow exception catch to HTTPException/ValueError only; re-raise unexpected errors so programming bugs surface in logs rather than silently appearing as guardrail blocks - Update tests: mock apply_guardrail directly (hook method was removed), replace session.update client-rewrite test with session.created injection test matching the new server-side approach --- .../litellm_core_utils/realtime_streaming.py | 41 ++++-- litellm/realtime_api/main.py | 1 + .../test_realtime_streaming.py | 128 ++++++++++-------- 3 files changed, 102 insertions(+), 68 deletions(-) diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 1599420f20a..89087ec150c 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -160,9 +160,18 @@ class RealTimeStreaming: request_data={"user_api_key_dict": self.user_api_key_dict}, input_type="request", ) - except Exception as e: + except (ValueError, Exception) as e: + # Only treat HTTPException (guardrail block) and ValueError as intentional blocks. + # Re-raise unexpected programming errors so they surface in logs. + from fastapi import HTTPException + + if not isinstance(e, (HTTPException, ValueError)): + verbose_logger.exception( + "[realtime guardrail] unexpected error in apply_guardrail: %s", e + ) + raise # Extract the human-readable error from HTTPException detail dict, - # falling back to str(e) for other exception types. + # falling back to str(e) for ValueError. try: detail = e.detail # type: ignore[attr-defined] safe_msg = ( @@ -239,14 +248,30 @@ class RealTimeStreaming: self.session_configuration_request = returned_object[ "session_configuration_request" ] - if isinstance(transformed_response, list): - for event in transformed_response: - event_str = json.dumps(event) - ## LOGGING + events = ( + transformed_response + if isinstance(transformed_response, list) + else [transformed_response] + ) + for event in events: + event_str = json.dumps(event) + ## GUARDRAIL: run on transcription events in provider_config path too + if ( + isinstance(event, dict) + and event.get("type") + == "conversation.item.input_audio_transcription.completed" + ): + transcript = event.get("transcript", "") self.store_message(event_str) await self.websocket.send_text(event_str) - else: - event_str = json.dumps(transformed_response) + blocked = await self.run_realtime_guardrails( + transcript, item_id=event.get("item_id") + ) + if not blocked: + await self.backend_ws.send( + json.dumps({"type": "response.create"}) + ) + continue ## LOGGING self.store_message(event_str) await self.websocket.send_text(event_str) diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 01b83067650..f1cc9b1d977 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -156,6 +156,7 @@ async def _arealtime( client=None, timeout=timeout, query_params=query_params, + user_api_key_dict=kwargs.get("user_api_key_dict"), ) elif _custom_llm_provider == "bedrock": # Extract AWS parameters from kwargs diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index dab8e16810e..6b10369604f 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -79,16 +79,13 @@ async def test_realtime_guardrail_blocks_prompt_injection(): # Simple guardrail that blocks anything with "system update" class PromptInjectionGuardrail(CustomGuardrail): - async def async_realtime_input_transcription_hook( - self, - transcription, - user_api_key_dict, - session_id=None, - ): - if "system update" in transcription.lower(): - raise ValueError( - "⚠️ Prompt injection detected. Request blocked by guardrail." - ) + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + for text in inputs.get("texts", []): + if "system update" in text.lower(): + raise ValueError( + "⚠️ Prompt injection detected. Request blocked by guardrail." + ) + return inputs guardrail = PromptInjectionGuardrail( guardrail_name="test_injection_guard", @@ -125,31 +122,32 @@ async def test_realtime_guardrail_blocks_prompt_injection(): streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj) await streaming.backend_to_client_send_messages() - # ASSERT 1: response.create was NOT sent to backend (injection blocked) + # ASSERT 1: no bare response.create was sent to backend (injection blocked). + # The only response.create allowed is the warning one (has "instructions" field). sent_to_backend = [ json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args ] - response_creates = [ - e for e in sent_to_backend if e.get("type") == "response.create" + bare_response_creates = [ + e for e in sent_to_backend + if e.get("type") == "response.create" + and "instructions" not in e.get("response", {}) ] - assert len(response_creates) == 0, ( - f"Guardrail should prevent response.create for injected content, " - f"but got: {response_creates}" + assert len(bare_response_creates) == 0, ( + f"Guardrail should prevent bare response.create for injected content, " + f"but got: {bare_response_creates}" ) - # ASSERT 2: a warning response was sent to the client - sent_to_client = [ - json.loads(c.args[0]) for c in client_ws.send_text.call_args_list + # ASSERT 2: warning response.create was sent to backend (to speak the block message) + warning_creates = [ + e for e in sent_to_backend + if e.get("type") == "response.create" + and "instructions" in e.get("response", {}) ] - warning_events = [ - e - for e in sent_to_client - if e.get("type") == "response.text.delta" and "⚠️" in e.get("delta", "") - ] - assert len(warning_events) > 0, ( - f"Client should receive a guardrail warning, but got: {sent_to_client}" + assert len(warning_creates) > 0, ( + f"Backend should receive a response.create with warning instructions, " + f"but got: {sent_to_backend}" ) litellm.callbacks = [] # cleanup @@ -166,14 +164,11 @@ async def test_realtime_guardrail_allows_clean_transcript(): from litellm.types.guardrails import GuardrailEventHooks class PromptInjectionGuardrail(CustomGuardrail): - async def async_realtime_input_transcription_hook( - self, - transcription, - user_api_key_dict, - session_id=None, - ): - if "system update" in transcription.lower(): - raise ValueError("⚠️ Prompt injection detected.") + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + for text in inputs.get("texts", []): + if "system update" in text.lower(): + raise ValueError("⚠️ Prompt injection detected.") + return inputs guardrail = PromptInjectionGuardrail( guardrail_name="test_injection_guard", @@ -225,45 +220,58 @@ async def test_realtime_guardrail_allows_clean_transcript(): @pytest.mark.asyncio -async def test_realtime_session_update_forces_create_response_false(): +async def test_realtime_session_created_injects_create_response_false(): """ - Test that session.update with create_response=True is rewritten to - create_response=False so the guardrail controls when LLM responds. + Test that when session.created arrives from the backend and realtime guardrails + are registered, the proxy injects a session.update with create_response=False + so the LLM never auto-responds before the guardrail runs. """ import litellm + from litellm.integrations.custom_guardrail import CustomGuardrail + from litellm.types.guardrails import GuardrailEventHooks + + class DummyGuardrail(CustomGuardrail): + async def apply_guardrail(self, inputs, request_data, input_type, logging_obj=None): + return inputs + + guardrail = DummyGuardrail( + guardrail_name="dummy", + event_hook=GuardrailEventHooks.realtime_input_transcription, + default_on=True, + ) + litellm.callbacks = [guardrail] client_ws = MagicMock() client_ws.send_text = AsyncMock() - client_ws.receive_text = AsyncMock( + + session_created_event = json.dumps({"type": "session.created"}).encode() + + backend_ws = MagicMock() + backend_ws.recv = AsyncMock( side_effect=[ - json.dumps( - { - "type": "session.update", - "session": { - "turn_detection": { - "type": "server_vad", - "create_response": True, - "threshold": 0.5, - } - }, - } - ), + session_created_event, ConnectionClosed(None, None), ] ) - - backend_ws = MagicMock() backend_ws.send = AsyncMock() logging_obj = MagicMock() + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj) - await streaming.client_ack_messages() + await streaming.backend_to_client_send_messages() - # ASSERT: forwarded session.update has create_response=False - sent_to_backend = backend_ws.send.call_args_list - assert len(sent_to_backend) == 1 - forwarded = json.loads(sent_to_backend[0].args[0]) - assert forwarded["session"]["turn_detection"]["create_response"] is False, ( - f"session.update should have create_response rewritten to False, " - f"got: {forwarded['session']['turn_detection']}" + # ASSERT: proxy injected session.update with create_response=False to backend + sent_to_backend = [ + json.loads(c.args[0]) for c in backend_ws.send.call_args_list if c.args + ] + session_updates = [e for e in sent_to_backend if e.get("type") == "session.update"] + assert len(session_updates) == 1, ( + f"Expected proxy to inject session.update, got: {sent_to_backend}" ) + td = session_updates[0]["session"]["turn_detection"] + assert td["create_response"] is False, ( + f"Expected create_response=False, got: {td}" + ) + + litellm.callbacks = [] # cleanup