mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
3efd803cd1
commit
615a7da9ba
4 changed files with 137 additions and 37 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue