From 615a7da9ba049ee6ecae8fd823b2c140424b3a6b Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 22 May 2026 21:35:00 +0000 Subject: [PATCH] fix(realtime): deep-merge generationConfig and refresh cache on follow-up setup MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit A subsequent Gemini session.update that touches any generationConfig sub-field (e.g. just temperature) was clobbering the original generationConfig — silently dropping responseModalities and switching the session to text-only. Deep-merge generationConfig so existing keys (responseModalities, maxOutputTokens, ...) are preserved when the client updates only a subset. Also drop the early-return in _cache_session_configuration_request so the cached payload tracks the latest setup sent to the backend. Without this, downstream readers (transform_session_created_event, modality lookup in return_new_content_delta_events) keep reading stale modalities/system instruction after a follow-up setup. --- .../litellm_core_utils/realtime_streaming.py | 11 +- .../llms/gemini/realtime/transformation.py | 8 +- .../test_realtime_streaming.py | 121 +++++++++++++----- .../test_gemini_realtime_transformation.py | 34 +++++ 4 files changed, 137 insertions(+), 37 deletions(-) diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 1d3528832df..61a3471bdd5 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -286,9 +286,14 @@ class RealTimeStreaming: return True def _cache_session_configuration_request(self, transformed_message: str) -> None: - """Store setup payload once sent to backend.""" - if self.session_configuration_request is not None: - return + """Store setup payload once sent to backend. + + Updates the cached setup on every successful setup send so follow-up + ``session.update`` messages (which produce a merged setup with new + ``generationConfig`` / ``systemInstruction`` / etc.) are reflected in + the cache used by downstream readers (``transform_session_created_event``, + ``return_new_content_delta_events`` modality lookup, ...). + """ try: message_obj = json.loads(transformed_message) if "setup" in message_obj: diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py index 44adc7a54bf..84bdd6051a4 100644 --- a/litellm/llms/gemini/realtime/transformation.py +++ b/litellm/llms/gemini/realtime/transformation.py @@ -344,14 +344,16 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): except (json.JSONDecodeError, AttributeError): original_setup = {} + # Deep-merge ``generationConfig`` and ``realtimeInputConfig`` so a + # partial session.update (e.g. only ``temperature`` or only + # ``modalities``) does not silently drop unrelated sub-keys + # (``responseModalities``, ``maxOutputTokens``, ...) from the original + # setup. follow_up_setup: BidiGenerateContentSetup = { **original_setup, **new_overrides, "model": f"models/{model}", } - # Deep-merge nested config dicts so that a partial session.update - # (e.g. only ``modalities``) does not silently drop unrelated - # sub-keys (e.g. ``temperature``) from the original setup. original_generation_config = original_setup.get("generationConfig") new_generation_config = new_overrides.get("generationConfig") if isinstance(original_generation_config, dict) and isinstance( 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 35e4e232f1a..f9a9897e2b6 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -1268,8 +1268,7 @@ async def test_duplicate_session_created_still_triggers_guardrail_turn_detection injected_session = sent_update["session"] assert injected_session["type"] == "realtime" assert ( - injected_session["audio"]["input"]["turn_detection"]["create_response"] - is False + injected_session["audio"]["input"]["turn_detection"]["create_response"] is False ) @@ -1298,13 +1297,13 @@ async def test_guardrail_update_respects_idempotency_flag(): model="gemini-2.5-flash", ) streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign] - + # First call should send the update assert streaming._guardrail_turn_detection_update_sent is False await streaming._maybe_send_guardrail_turn_detection_update() assert streaming._guardrail_turn_detection_update_sent is True assert backend_ws.send.await_count == 1 - + # Second call should be a no-op (idempotent) await streaming._maybe_send_guardrail_turn_detection_update() assert backend_ws.send.await_count == 1 # Still 1, not 2 @@ -1316,13 +1315,15 @@ async def test_guardrail_turn_detection_injected_into_first_session_update_defer client_ws = AsyncMock() client_ws.receive_text = AsyncMock( side_effect=[ - json.dumps({ - "type": "session.update", - "session": { - "modalities": ["text", "audio"], - "tools": [{"type": "function", "name": "get_weather"}], + json.dumps( + { + "type": "session.update", + "session": { + "modalities": ["text", "audio"], + "tools": [{"type": "function", "name": "get_weather"}], + }, } - }), + ), ConnectionClosed(None, None), ] ) @@ -1336,9 +1337,11 @@ async def test_guardrail_turn_detection_injected_into_first_session_update_defer provider_config = MagicMock() transformed_messages = [] + def mock_transform(msg, model, session_config): transformed_messages.append((msg, session_config)) return [msg] # Pass through for simplicity + provider_config.transform_realtime_request = MagicMock(side_effect=mock_transform) streaming = RealTimeStreaming( @@ -1349,10 +1352,10 @@ async def test_guardrail_turn_detection_injected_into_first_session_update_defer model="gemini-2.5-flash", ) streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign] - + # Simulate first session.update in deferred mode await streaming.client_ack_messages() - + # Verify turn_detection was injected into the session.update. The # injection runs before the GA remap, so the create_response flag ends # up nested under audio.input.turn_detection in the GA-shaped payload. @@ -1361,10 +1364,9 @@ async def test_guardrail_turn_detection_injected_into_first_session_update_defer msg_obj = json.loads(transformed_msg) assert msg_obj["type"] == "session.update" session_obj = msg_obj["session"] - injected_turn_detection = ( - session_obj.get("turn_detection") - or session_obj.get("audio", {}).get("input", {}).get("turn_detection") - ) + injected_turn_detection = session_obj.get("turn_detection") or session_obj.get( + "audio", {} + ).get("input", {}).get("turn_detection") assert injected_turn_detection is not None assert injected_turn_detection["create_response"] is False assert streaming._guardrail_turn_detection_update_sent is True @@ -1379,13 +1381,15 @@ async def test_guardrail_turn_detection_injection_tolerates_non_dict_value( client_ws = AsyncMock() client_ws.receive_text = AsyncMock( side_effect=[ - json.dumps({ - "type": "session.update", - "session": { - "modalities": ["text", "audio"], - "turn_detection": existing_turn_detection, - }, - }), + json.dumps( + { + "type": "session.update", + "session": { + "modalities": ["text", "audio"], + "turn_detection": existing_turn_detection, + }, + } + ), ConnectionClosed(None, None), ] ) @@ -1399,9 +1403,11 @@ async def test_guardrail_turn_detection_injection_tolerates_non_dict_value( provider_config = MagicMock() transformed_messages = [] + def mock_transform(msg, model, session_config): transformed_messages.append((msg, session_config)) return [msg] + provider_config.transform_realtime_request = MagicMock(side_effect=mock_transform) streaming = RealTimeStreaming( @@ -1419,10 +1425,9 @@ async def test_guardrail_turn_detection_injection_tolerates_non_dict_value( transformed_msg, _ = transformed_messages[0] msg_obj = json.loads(transformed_msg) session_obj = msg_obj["session"] - injected_turn_detection = ( - session_obj.get("turn_detection") - or session_obj.get("audio", {}).get("input", {}).get("turn_detection") - ) + injected_turn_detection = session_obj.get("turn_detection") or session_obj.get( + "audio", {} + ).get("input", {}).get("turn_detection") assert isinstance(injected_turn_detection, dict) assert injected_turn_detection["create_response"] is False assert streaming._guardrail_turn_detection_update_sent is True @@ -1492,9 +1497,63 @@ async def test_subsequent_session_update_cannot_reenable_vad_when_guardrails_act forwarded_msg, _ = transformed_messages[0] msg_obj = json.loads(forwarded_msg) session_obj = msg_obj["session"] - forwarded_turn_detection = ( - session_obj.get("turn_detection") - or session_obj.get("audio", {}).get("input", {}).get("turn_detection") - ) + forwarded_turn_detection = session_obj.get("turn_detection") or session_obj.get( + "audio", {} + ).get("input", {}).get("turn_detection") assert isinstance(forwarded_turn_detection, dict) assert forwarded_turn_detection["create_response"] is False + + +@pytest.mark.asyncio +async def test_follow_up_setup_updates_cached_session_configuration_request(): + """A follow-up setup produced by a subsequent session.update must replace + the cached ``session_configuration_request`` so downstream readers + (e.g. modality lookup in ``response.created``) see the latest config.""" + client_ws = AsyncMock() + client_ws.receive_text = AsyncMock( + side_effect=[ + json.dumps({"type": "session.update", "session": {"tools": []}}), + ConnectionClosed(None, None), + ] + ) + backend_ws = MagicMock() + backend_ws.send = AsyncMock() + + logging_obj = MagicMock() + logging_obj.async_success_handler = AsyncMock() + logging_obj.success_handler = MagicMock() + + provider_config = MagicMock() + follow_up_setup = json.dumps( + { + "setup": { + "model": "models/gemini-2.5-flash", + "generationConfig": {"responseModalities": ["TEXT"]}, + "tools": [{"function_declarations": []}], + } + } + ) + provider_config.transform_realtime_request = MagicMock( + return_value=[follow_up_setup] + ) + + streaming = RealTimeStreaming( + websocket=client_ws, + backend_ws=backend_ws, + logging_obj=logging_obj, + provider_config=provider_config, + model="gemini-2.5-flash", + ) + # Simulate that the original auto-setup was already cached. + streaming.session_configuration_request = json.dumps( + { + "setup": { + "model": "models/gemini-2.5-flash", + "generationConfig": {"responseModalities": ["AUDIO"]}, + } + } + ) + + await streaming.client_ack_messages() + + assert streaming.session_configuration_request == follow_up_setup diff --git a/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py b/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py index 3c51c209e92..aac14588c95 100644 --- a/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py +++ b/tests/test_litellm/llms/gemini/realtime/test_gemini_realtime_transformation.py @@ -965,3 +965,37 @@ def test_gemini_subsequent_session_update_with_turn_detection_only_preserves_ori follow_up["realtimeInputConfig"]["automaticActivityDetection"]["disabled"] is True ) + + +def test_gemini_follow_up_session_update_preserves_response_modalities_on_partial_generation_config(): + """A follow-up session.update that only sets `temperature` (or any other + generationConfig sub-field) must not wipe `responseModalities` from the + original setup.""" + config = GeminiRealtimeConfig() + + original_setup = { + "setup": { + "model": "models/gemini-2.5-flash-native-audio", + "generationConfig": { + "responseModalities": ["AUDIO"], + "maxOutputTokens": 2048, + }, + "inputAudioTranscription": {}, + } + } + + session_update = { + "type": "session.update", + "session": {"temperature": 0.7}, + } + + messages = config.transform_realtime_request( + json.dumps(session_update), + "gemini-2.5-flash-native-audio", + session_configuration_request=json.dumps(original_setup), + ) + + follow_up = json.loads(messages[0])["setup"] + assert follow_up["generationConfig"]["responseModalities"] == ["AUDIO"] + assert follow_up["generationConfig"]["maxOutputTokens"] == 2048 + assert follow_up["generationConfig"]["temperature"] == 0.7