From f5ecff22ed003c68e926126824b10bd5fa70ab8d Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Fri, 22 May 2026 18:23:02 +0530 Subject: [PATCH] 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 --- litellm/llms/custom_httpx/llm_http_handler.py | 4 +- .../llms/gemini/realtime/transformation.py | 299 ++++++++++-------- .../llms/vertex_ai/realtime/transformation.py | 14 +- .../test_realtime_streaming.py | 7 +- .../test_gemini_realtime_transformation.py | 73 +++++ 5 files changed, 251 insertions(+), 146 deletions(-) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 1e5af050e6b..df06bd8d3fe 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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 diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py index 917f13eefb6..6bd2fd57ab9 100644 --- a/litellm/llms/gemini/realtime/transformation.py +++ b/litellm/llms/gemini/realtime/transformation.py @@ -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, diff --git a/litellm/llms/vertex_ai/realtime/transformation.py b/litellm/llms/vertex_ai/realtime/transformation.py index 9bc6da05ed7..4d5b9b7c779 100644 --- a/litellm/llms/vertex_ai/realtime/transformation.py +++ b/litellm/llms/vertex_ai/realtime/transformation.py @@ -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 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 598119137f5..b4a7a05a136 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -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 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 3aefc4d149f..c95a085d17c 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 @@ -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():