fix(realtime): deep-merge generationConfig and refresh cache on follow-up setup

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.
This commit is contained in:
mateo-berri 2026-05-22 21:35:00 +00:00 • committed by Claude
parent 3efd803cd1
commit 615a7da9ba
No known key found for this signature in database
4 changed files with 137 additions and 37 deletions

View file

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

View file

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

View file

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

View file

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