fix(gemini/realtime): reset response IDs after tool-call response.done

After closing a tool-call response, clear current_output_item_id and
current_response_id so post-tool model turns emit a fresh response.created
preamble. Add regression tests and align guardrail turn_detection test with
GA session shape; apply Black formatting.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Sameer Kankute 2026-05-22 18:23:02 +05:30
parent 165fedd83a
commit f5ecff22ed
No known key found for this signature in database
5 changed files with 251 additions and 146 deletions

View file

@ -5217,7 +5217,7 @@ class BaseLLMHTTPHandler:
)
if _session_config:
realtime_streaming.session_configuration_request = _session_config
# For providers that defer setup until client session.update, optionally
# send synthetic session.created to unblock clients waiting on connect.
if not provider_config.requires_session_configuration():
@ -5232,7 +5232,7 @@ class BaseLLMHTTPHandler:
verbose_logger.debug(
"Sent synthetic session.created to client to unblock connection"
)
await realtime_streaming.bidirectional_forward()
except websockets.exceptions.InvalidStatusCode as e: # type: ignore

View file

@ -58,7 +58,9 @@ from litellm.utils import get_empty_usage
from ..common_utils import encode_unserializable_types, get_api_key_from_env
MAP_GEMINI_FIELD_TO_OPENAI_EVENT: Dict[str, Union[OpenAIRealtimeEventTypes, ResponsesAPIStreamEvents]] = {
MAP_GEMINI_FIELD_TO_OPENAI_EVENT: Dict[
str, Union[OpenAIRealtimeEventTypes, ResponsesAPIStreamEvents]
] = {
"setupComplete": OpenAIRealtimeEventTypes.SESSION_CREATED,
"serverContent.generationComplete": OpenAIRealtimeEventTypes.RESPONSE_TEXT_DONE,
"serverContent.turnComplete": OpenAIRealtimeEventTypes.RESPONSE_DONE,
@ -72,7 +74,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
super().__init__()
# Store call_id → function_name mapping for tool call round-trip
self._tool_call_id_to_name: Dict[str, str] = {}
def validate_environment(
self, headers: dict, model: str, api_key: Optional[str] = None
) -> dict:
@ -199,10 +201,8 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
vertex_gemini_config = VertexGeminiConfig()
# Tools should be at the top level of setup, not inside generationConfig
optional_params["tools"] = (
vertex_gemini_config._map_function(
value=value, optional_params=optional_params
)
optional_params["tools"] = vertex_gemini_config._map_function(
value=value, optional_params=optional_params
)
elif key == "input_audio_transcription" and value is not None:
optional_params["inputAudioTranscription"] = {}
@ -231,7 +231,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
) -> List[str]:
"""
Handle session.update by sending setup to Gemini.
Only sends setup on the FIRST session.update (when session_configuration_request is None).
Subsequent session.update messages are ignored because Gemini doesn't support dynamic updates.
"""
@ -244,9 +244,13 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
"generationConfig", {}
)
generation_config.setdefault("responseModalities", ["AUDIO"])
client_session_configuration_request.setdefault("inputAudioTranscription", {})
client_session_configuration_request.setdefault(
"inputAudioTranscription", {}
)
client_session_configuration_request["model"] = f"models/{model}"
gemini_setup_msg = json.dumps({"setup": client_session_configuration_request})
gemini_setup_msg = json.dumps(
{"setup": client_session_configuration_request}
)
verbose_logger.debug(
f"Gemini Realtime: Sending initial setup with tools to backend"
)
@ -261,17 +265,17 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
def _handle_conversation_item(self, json_message: dict) -> List[str]:
"""
Handle conversation.item.create for user text or function call output.
Converts OpenAI format to Gemini's clientContent (for user text) or
Converts OpenAI format to Gemini's clientContent (for user text) or
toolResponse (for function outputs).
"""
item = json_message.get("item", {})
item_type = item.get("type")
# Handle function call output (tool response)
if item_type == "function_call_output":
return self._handle_function_call_output(item)
# Handle regular text content
return self._handle_user_text_content(item)
@ -279,15 +283,17 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
"""Transform function_call_output to Gemini toolResponse format."""
call_id = item.get("call_id", "")
output = item.get("output", "{}")
verbose_logger.debug(f"Gemini Realtime: Transforming function_call_output for call_id={call_id}")
verbose_logger.debug(
f"Gemini Realtime: Transforming function_call_output for call_id={call_id}"
)
# Parse the output to get the result
try:
output_dict = json.loads(output) if isinstance(output, str) else output
except json.JSONDecodeError:
output_dict = {"result": output}
# Look up the function name from stored mapping
function_name = self._tool_call_id_to_name.get(call_id)
if not function_name:
@ -295,7 +301,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
f"Gemini Realtime: Function name not found for call_id={call_id}. "
"This may cause Gemini to reject the response."
)
# Build Gemini toolResponse format
function_response = {
"id": call_id,
@ -303,13 +309,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
}
if function_name:
function_response["name"] = function_name
tool_response_message = {
"toolResponse": {
"functionResponses": [function_response]
}
"toolResponse": {"functionResponses": [function_response]}
}
return [json.dumps(tool_response_message)]
def _handle_user_text_content(self, item: dict) -> List[str]:
@ -323,20 +327,15 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
text = " ".join(filter(None, text_parts))
if not text:
return []
# Build clientContent message with turns (proper Gemini Live API format)
client_content_message = {
"clientContent": {
"turns": [
{
"role": "user",
"parts": [{"text": text}]
}
],
"turnComplete": True
"turns": [{"role": "user", "parts": [{"text": text}]}],
"turnComplete": True,
}
}
return [json.dumps(client_content_message)]
def transform_realtime_request(
@ -371,20 +370,24 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
## HANDLE conversation.item.create — extract user text or function call output ##
if msg_type == "conversation.item.create":
return self._handle_conversation_item(json_message)
## HANDLE INPUT AUDIO BUFFER - use realtimeInput for audio streaming ##
if msg_type == "input_audio_buffer.append":
realtime_input_dict["audio"] = HttpxBlobType(
mimeType=self.get_audio_mime_type(), data=json_message["audio"]
)
realtime_input_dict = cast(
BidiGenerateContentRealtimeInput,
encode_unserializable_types(cast(Dict[str, object], realtime_input_dict)),
encode_unserializable_types(
cast(Dict[str, object], realtime_input_dict)
),
)
gemini_msg = json.dumps({"realtimeInput": realtime_input_dict})
verbose_logger.debug("Gemini Realtime: Sending audio realtimeInput to backend")
verbose_logger.debug(
"Gemini Realtime: Sending audio realtimeInput to backend"
)
messages.append(gemini_msg)
return messages
# Unknown/unsupported OpenAI event type — drop silently rather than
@ -692,38 +695,40 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
) -> List[Dict[str, Any]]:
"""
Transform Gemini toolCall message to OpenAI function call events.
Converts Gemini's functionCalls format to OpenAI's response.function_call_arguments.done events.
Also stores call_id name mapping for later use in function_call_output responses.
"""
function_calls = tool_call_message.get("functionCalls", [])
resolved_response_id = response_id or f"resp_{uuid.uuid4()}"
resolved_output_item_id = output_item_id or f"item_{uuid.uuid4()}"
verbose_logger.debug(
f"Gemini Realtime: Transforming {len(function_calls)} tool call(s) to OpenAI format"
)
events = []
for idx, fc in enumerate(function_calls):
call_id = fc.get("id", "")
name = fc.get("name", "")
# Store call_id → name mapping for round-trip
if call_id and name:
self._tool_call_id_to_name[call_id] = name
events.append({
"type": "response.function_call_arguments.done",
"event_id": f"event_{uuid.uuid4()}",
"response_id": resolved_response_id,
"item_id": f"{resolved_output_item_id}_tool_{idx}",
"output_index": idx,
"call_id": call_id,
"name": name,
"arguments": json.dumps(fc.get("args", {})),
})
events.append(
{
"type": "response.function_call_arguments.done",
"event_id": f"event_{uuid.uuid4()}",
"response_id": resolved_response_id,
"item_id": f"{resolved_output_item_id}_tool_{idx}",
"output_index": idx,
"call_id": call_id,
"name": name,
"arguments": json.dumps(fc.get("args", {})),
}
)
return events
@staticmethod
@ -964,7 +969,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
) -> Union[OpenAIRealtimeEventTypes, ResponsesAPIStreamEvents]:
model_turn_event = value.get("modelTurn")
generation_complete_event = value.get("generationComplete")
openai_event: Optional[Union[OpenAIRealtimeEventTypes, ResponsesAPIStreamEvents]] = None
openai_event: Optional[
Union[OpenAIRealtimeEventTypes, ResponsesAPIStreamEvents]
] = None
if model_turn_event: # check if model turn event
openai_event = self.map_model_turn_event(model_turn_event)
elif generation_complete_event:
@ -1006,7 +1013,9 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
message_str = str(message)
raise ValueError(f"Invalid JSON message: {message_str}")
verbose_logger.debug(f"Realtime Response Transform: Gemini message={json.dumps(json_message)[:500]}")
verbose_logger.debug(
f"Realtime Response Transform: Gemini message={json.dumps(json_message)[:500]}"
)
logging_session_id = logging_obj.litellm_trace_id
@ -1108,21 +1117,25 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
if current_response_id is None:
current_response_id = f"resp_{uuid.uuid4()}"
current_output_item_id = f"item_{uuid.uuid4()}"
current_conversation_id = current_conversation_id or f"conv_{uuid.uuid4()}"
current_conversation_id = (
current_conversation_id or f"conv_{uuid.uuid4()}"
)
# Emit response.created
returned_message.append({
"type": "response.created",
"event_id": f"event_{uuid.uuid4()}",
"response": {
"object": "realtime.response",
"id": current_response_id,
"status": "in_progress",
"output": [],
"conversation_id": current_conversation_id,
},
})
returned_message.append(
{
"type": "response.created",
"event_id": f"event_{uuid.uuid4()}",
"response": {
"object": "realtime.response",
"id": current_response_id,
"status": "in_progress",
"output": [],
"conversation_id": current_conversation_id,
},
}
)
tool_call_events = self.transform_tool_call_events(
value,
response_id=current_response_id,
@ -1132,77 +1145,89 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
for idx, tool_call in enumerate(tool_call_events):
item_id = tool_call["item_id"]
# response.output_item.added
returned_message.append({
"type": "response.output_item.added",
"event_id": f"event_{uuid.uuid4()}",
"response_id": current_response_id,
"output_index": idx,
"item": {
"id": item_id,
"object": "realtime.item",
"type": "function_call",
"status": "in_progress",
"call_id": tool_call["call_id"],
"name": tool_call["name"],
"arguments": "",
},
})
returned_message.append(
{
"type": "response.output_item.added",
"event_id": f"event_{uuid.uuid4()}",
"response_id": current_response_id,
"output_index": idx,
"item": {
"id": item_id,
"object": "realtime.item",
"type": "function_call",
"status": "in_progress",
"call_id": tool_call["call_id"],
"name": tool_call["name"],
"arguments": "",
},
}
)
# response.function_call_arguments.done
returned_message.append(tool_call)
# response.output_item.done
returned_message.append({
"type": "response.output_item.done",
"event_id": f"event_{uuid.uuid4()}",
"response_id": current_response_id,
"output_index": idx,
"item": {
"id": item_id,
"object": "realtime.item",
"type": "function_call",
"status": "completed",
"call_id": tool_call["call_id"],
"name": tool_call["name"],
"arguments": tool_call["arguments"],
},
})
# conversation.item.created
returned_message.append({
"type": "conversation.item.created",
"event_id": f"event_{uuid.uuid4()}",
"item": {
"id": item_id,
"object": "realtime.item",
"type": "function_call",
"status": "completed",
"call_id": tool_call["call_id"],
"name": tool_call["name"],
"arguments": tool_call["arguments"],
},
})
# response.done - close the response so clients can submit tool results
returned_message.append({
"type": "response.done",
"event_id": f"event_{uuid.uuid4()}",
"response": {
"id": current_response_id,
"object": "realtime.response",
"status": "completed",
"output": [
{
"id": te["item_id"],
returned_message.append(
{
"type": "response.output_item.done",
"event_id": f"event_{uuid.uuid4()}",
"response_id": current_response_id,
"output_index": idx,
"item": {
"id": item_id,
"object": "realtime.item",
"type": "function_call",
"status": "completed",
"call_id": te["call_id"],
"name": te["name"],
"arguments": te["arguments"],
}
for te in tool_call_events
],
"usage": None,
},
})
"call_id": tool_call["call_id"],
"name": tool_call["name"],
"arguments": tool_call["arguments"],
},
}
)
# conversation.item.created
returned_message.append(
{
"type": "conversation.item.created",
"event_id": f"event_{uuid.uuid4()}",
"item": {
"id": item_id,
"object": "realtime.item",
"type": "function_call",
"status": "completed",
"call_id": tool_call["call_id"],
"name": tool_call["name"],
"arguments": tool_call["arguments"],
},
}
)
# response.done - close the response so clients can submit tool results
returned_message.append(
{
"type": "response.done",
"event_id": f"event_{uuid.uuid4()}",
"response": {
"id": current_response_id,
"object": "realtime.response",
"status": "completed",
"output": [
{
"id": te["item_id"],
"object": "realtime.item",
"type": "function_call",
"status": "completed",
"call_id": te["call_id"],
"name": te["name"],
"arguments": te["arguments"],
}
for te in tool_call_events
],
"usage": None,
},
}
)
# Reset IDs so the next model turn (after tool results) starts a
# fresh response with its own response.created preamble.
current_output_item_id = None
current_response_id = None
elif openai_event == OpenAIRealtimeEventTypes.RESPONSE_DONE:
transformed_response_done_event = self.transform_response_done_event(
message=BidiGenerateContentServerMessage(**json_message), # type: ignore
@ -1247,12 +1272,14 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
transformed_message=returned_message,
current_item_chunks=current_item_chunks,
)
# Log the transformed events
for msg in returned_message:
event_type = msg.get("type") if isinstance(msg, dict) else "unknown"
verbose_logger.debug(f"Realtime Response Transform: OpenAI event={event_type}, data={json.dumps(msg)[:500] if isinstance(msg, dict) else str(msg)[:500]}")
verbose_logger.debug(
f"Realtime Response Transform: OpenAI event={event_type}, data={json.dumps(msg)[:500] if isinstance(msg, dict) else str(msg)[:500]}"
)
return {
"response": returned_message,
"current_output_item_id": current_output_item_id,

View file

@ -145,14 +145,14 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig):
setup_config = self.map_openai_params(
optional_params={}, non_default_params=session_params
)
# Use full Vertex AI model path
setup_config["model"] = (
f"projects/{self._project}"
f"/locations/{self._location}"
f"/publishers/google/models/{model}"
)
# Add Vertex AI specific defaults if not provided
generation_config = setup_config.setdefault("generationConfig", {})
generation_config.setdefault("responseModalities", ["AUDIO"])
@ -167,7 +167,7 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig):
)
setup_config.setdefault("inputAudioTranscription", {})
setup_config.setdefault("outputAudioTranscription", {})
return setup_config
def transform_realtime_request(
@ -178,12 +178,12 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig):
) -> List[str]:
"""
Translate OpenAI realtime client messages to Vertex AI format.
Handles session.update by sending setup with proper Vertex AI model path.
"""
json_message = json.loads(message)
msg_type = json_message.get("type")
# Handle session.update with Vertex AI specific model path
if msg_type == "session.update":
if session_configuration_request is None:
@ -192,7 +192,7 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig):
model, json_message["session"]
)
gemini_setup_msg = json.dumps({"setup": setup_config})
verbose_logger.debug(
"Vertex AI Realtime: Sending initial setup with tools to backend"
)
@ -203,7 +203,7 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig):
"Vertex AI Realtime: Ignoring session.update (setup already sent)"
)
return []
# For other message types, use parent's logic
return super().transform_realtime_request(
message, model, session_configuration_request

View file

@ -1265,7 +1265,12 @@ async def test_duplicate_session_created_still_triggers_guardrail_turn_detection
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
injected_session = sent_update["session"]
assert injected_session["type"] == "realtime"
assert (
injected_session["audio"]["input"]["turn_detection"]["create_response"]
is False
)
@pytest.mark.asyncio

View file

@ -674,6 +674,79 @@ def test_gemini_tool_call_emits_response_created_preamble():
assert responses[5]["response"]["status"] == "completed"
assert len(responses[5]["response"]["output"]) == 1
assert responses[5]["response"]["output"][0]["type"] == "function_call"
assert result["current_output_item_id"] is None
assert result["current_response_id"] is None
def test_gemini_tool_call_resets_ids_for_post_tool_model_turn():
"""After tool-call response.done, a subsequent modelTurn must emit response.created."""
config = GeminiRealtimeConfig()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_123"
session_configuration_request = json.dumps(
{
"model": "gemini-1.5-flash",
"generationConfig": {"responseModalities": ["TEXT"]},
}
)
tool_result = config.transform_realtime_response(
json.dumps(
{
"toolCall": {
"functionCalls": [
{
"id": "call_123",
"name": "get_weather",
"args": {"location": "San Francisco"},
}
]
}
}
),
"gemini-2.5-flash",
logging_obj,
realtime_response_transform_input={
"session_configuration_request": session_configuration_request,
"current_output_item_id": None,
"current_response_id": None,
"current_conversation_id": None,
"current_delta_chunks": [],
"current_item_chunks": [],
"current_delta_type": None,
},
)
tool_response_id = tool_result["response"][0]["response"]["id"]
assert tool_result["current_output_item_id"] is None
assert tool_result["current_response_id"] is None
post_tool_result = config.transform_realtime_response(
json.dumps(
{
"serverContent": {
"modelTurn": {"parts": [{"text": "The weather is sunny."}]}
}
}
),
"gemini-2.5-flash",
logging_obj,
realtime_response_transform_input={
"session_configuration_request": session_configuration_request,
"current_output_item_id": tool_result["current_output_item_id"],
"current_response_id": tool_result["current_response_id"],
"current_conversation_id": tool_result["current_conversation_id"],
"current_delta_chunks": tool_result["current_delta_chunks"],
"current_item_chunks": tool_result["current_item_chunks"],
"current_delta_type": tool_result["current_delta_type"],
},
)
post_tool_events = post_tool_result["response"]
assert post_tool_events[0]["type"] == "response.created"
assert post_tool_events[0]["response"]["id"] != tool_response_id
assert post_tool_result["current_response_id"] == post_tool_events[0]["response"]["id"]
def test_gemini_function_call_output_includes_name():