From a9784f043c4dae5916c26eaa89bdc5e803cdfafe Mon Sep 17 00:00:00 2001 From: Murat Kaan Meral Date: Sun, 9 Nov 2025 19:02:33 +0300 Subject: [PATCH 1/4] temp commit message, review the changes --- .../models/gemini_live.py | 57 ++++++++++++++----- .../tests/test_gemini_live.py | 19 ++----- .../models/test_gemini_live.py | 56 ++++++++++++++++-- .../test_bidirectional_agent.py | 27 +++------ 4 files changed, 108 insertions(+), 51 deletions(-) diff --git a/src/strands/experimental/bidirectional_streaming/models/gemini_live.py b/src/strands/experimental/bidirectional_streaming/models/gemini_live.py index 4337a6cfa6..465e8c9ebc 100644 --- a/src/strands/experimental/bidirectional_streaming/models/gemini_live.py +++ b/src/strands/experimental/bidirectional_streaming/models/gemini_live.py @@ -58,7 +58,7 @@ class GeminiLiveModel(BidirectionalModel): def __init__( self, - model_id: str = "models/gemini-2.0-flash-live-preview-04-09", + model_id: str = "gemini-2.5-flash-native-audio-preview-09-2025", api_key: Optional[str] = None, live_config: Optional[Dict[str, Any]] = None, **kwargs @@ -74,7 +74,19 @@ def __init__( # Model configuration self.model_id = model_id self.api_key = api_key - self.live_config = live_config or {} + + # Set default live_config with transcription enabled + default_config = { + "response_modalities": ["AUDIO"], + "outputAudioTranscription": {}, # Enable output transcription by default + "inputAudioTranscription": {} # Enable input transcription by default + } + + # Merge user config with defaults (user config takes precedence) + if live_config: + default_config.update(live_config) + + self.live_config = default_config # Create Gemini client with proper API version client_kwargs = {} @@ -242,18 +254,8 @@ def _convert_gemini_live_event(self, message: LiveServerMessage) -> Optional[Dic current_transcript=transcription_text ) - # Handle text output from model - if message.text: - logger.debug(f"Text output as transcript: {message.text}") - return TranscriptStreamEvent( - delta={"text": message.text}, - text=message.text, - role="assistant", - is_final=True, - current_transcript=message.text - ) - # Handle audio output using SDK's built-in data property + # Check this BEFORE text to avoid triggering warning on mixed content if message.data: # Convert bytes to base64 string for JSON serializability audio_b64 = base64.b64encode(message.data).decode('utf-8') @@ -264,6 +266,32 @@ def _convert_gemini_live_event(self, message: LiveServerMessage) -> Optional[Dic channels=GEMINI_CHANNELS ) + # Handle text output from model_turn (avoids warning by checking parts directly) + if message.server_content and message.server_content.model_turn: + model_turn = message.server_content.model_turn + if model_turn.parts: + # Concatenate all text parts (Gemini may send multiple parts) + text_parts = [] + for part in model_turn.parts: + # Log all part types for debugging + part_attrs = {attr: getattr(part, attr, None) for attr in dir(part) if not attr.startswith('_')} + logger.debug(f"Model turn part attributes: {part_attrs}") + + # Check if part has text attribute and it's not empty + if hasattr(part, 'text') and part.text: + text_parts.append(part.text) + + if text_parts: + full_text = " ".join(text_parts) + logger.debug(f"Text output as transcript ({len(text_parts)} parts): {full_text}") + return TranscriptStreamEvent( + delta={"text": full_text}, + text=full_text, + role="assistant", + is_final=True, + current_transcript=full_text + ) + # Handle tool calls if message.tool_call and message.tool_call.function_calls: for func_call in message.tool_call.function_calls: @@ -326,7 +354,8 @@ def _convert_gemini_live_event(self, message: LiveServerMessage) -> Optional[Dic logger.error("Error converting Gemini Live event: %s", e) logger.error("Message type: %s", type(message).__name__) logger.error("Message attributes: %s", [attr for attr in dir(message) if not attr.startswith('_')]) - return None + # Return ErrorEvent instead of None so caller can handle it + return ErrorEvent(error=e) async def send( self, diff --git a/src/strands/experimental/bidirectional_streaming/tests/test_gemini_live.py b/src/strands/experimental/bidirectional_streaming/tests/test_gemini_live.py index 38791d9ede..e9d715b499 100644 --- a/src/strands/experimental/bidirectional_streaming/tests/test_gemini_live.py +++ b/src/strands/experimental/bidirectional_streaming/tests/test_gemini_live.py @@ -167,13 +167,13 @@ async def receive(agent, context): # Handle transcript events (bidirectional_transcript_stream) elif event_type == "bidirectional_transcript_stream": transcript_text = event.get("text", "") - transcript_source = event.get("source", "unknown") + transcript_role = event.get("role", "unknown") is_final = event.get("is_final", False) # Print transcripts with special formatting - if transcript_source == "user": + if transcript_role == "user": print(f"🎤 User: {transcript_text}") - elif transcript_source == "assistant": + elif transcript_role == "assistant": print(f"🔊 Assistant: {transcript_text}") # Handle turn complete events (bidirectional_turn_complete) @@ -313,17 +313,10 @@ async def main(duration=180): # Initialize Gemini Live model with proper configuration logger.info("Initializing Gemini Live model with API key") - model = GeminiLiveModel( - model_id="gemini-2.5-flash-native-audio-preview-09-2025", - api_key=api_key, - live_config={ - "response_modalities": ["AUDIO"], - "output_audio_transcription": {}, # Enable output transcription - "input_audio_transcription": {} # Enable input transcription - } - ) + # Use default model and config (includes transcription enabled by default) + model = GeminiLiveModel(api_key=api_key) logger.info("Gemini Live model initialized successfully") - print("Using Gemini Live model") + print("Using Gemini Live model with default config (audio output + transcription enabled)") agent = BidirectionalAgent( model=model, diff --git a/tests/strands/experimental/bidirectional_streaming/models/test_gemini_live.py b/tests/strands/experimental/bidirectional_streaming/models/test_gemini_live.py index 107a8a84a7..e890ddcbba 100644 --- a/tests/strands/experimental/bidirectional_streaming/models/test_gemini_live.py +++ b/tests/strands/experimental/bidirectional_streaming/models/test_gemini_live.py @@ -89,20 +89,28 @@ def test_model_initialization(mock_genai_client, model_id, api_key): # Test default config model_default = GeminiLiveModel() - assert model_default.model_id == "models/gemini-2.0-flash-live-preview-04-09" + assert model_default.model_id == "gemini-2.5-flash-native-audio-preview-09-2025" assert model_default.api_key is None assert model_default._active is False assert model_default.live_session is None + # Check default config includes transcription + assert model_default.live_config["response_modalities"] == ["AUDIO"] + assert "outputAudioTranscription" in model_default.live_config + assert "inputAudioTranscription" in model_default.live_config # Test with API key model_with_key = GeminiLiveModel(model_id=model_id, api_key=api_key) assert model_with_key.model_id == model_id assert model_with_key.api_key == api_key - # Test with custom config + # Test with custom config (merges with defaults) live_config = {"temperature": 0.7, "top_p": 0.9} model_custom = GeminiLiveModel(model_id=model_id, live_config=live_config) - assert model_custom.live_config == live_config + # Custom config should be merged with defaults + assert model_custom.live_config["temperature"] == 0.7 + assert model_custom.live_config["top_p"] == 0.9 + # Defaults should still be present + assert "response_modalities" in model_custom.live_config # Connection Tests @@ -292,12 +300,24 @@ async def test_event_conversion(mock_genai_client, model): _, _, _ = mock_genai_client await model.connect() - # Test text output (converted to transcript) + # Test text output (converted to transcript via model_turn.parts) mock_text = unittest.mock.Mock() - mock_text.text = "Hello from Gemini" mock_text.data = None mock_text.tool_call = None - mock_text.server_content = None + + # Create proper server_content structure with model_turn + mock_server_content = unittest.mock.Mock() + mock_server_content.interrupted = False + mock_server_content.input_transcription = None + mock_server_content.output_transcription = None + + mock_model_turn = unittest.mock.Mock() + mock_part = unittest.mock.Mock() + mock_part.text = "Hello from Gemini" + mock_model_turn.parts = [mock_part] + mock_server_content.model_turn = mock_model_turn + + mock_text.server_content = mock_server_content text_event = model._convert_gemini_live_event(mock_text) assert isinstance(text_event, TranscriptStreamEvent) @@ -307,6 +327,30 @@ async def test_event_conversion(mock_genai_client, model): assert text_event.delta == {"text": "Hello from Gemini"} assert text_event.current_transcript == "Hello from Gemini" + # Test multiple text parts (should concatenate) + mock_multi_text = unittest.mock.Mock() + mock_multi_text.data = None + mock_multi_text.tool_call = None + + mock_server_content_multi = unittest.mock.Mock() + mock_server_content_multi.interrupted = False + mock_server_content_multi.input_transcription = None + mock_server_content_multi.output_transcription = None + + mock_model_turn_multi = unittest.mock.Mock() + mock_part1 = unittest.mock.Mock() + mock_part1.text = "Hello" + mock_part2 = unittest.mock.Mock() + mock_part2.text = "from Gemini" + mock_model_turn_multi.parts = [mock_part1, mock_part2] + mock_server_content_multi.model_turn = mock_model_turn_multi + + mock_multi_text.server_content = mock_server_content_multi + + multi_text_event = model._convert_gemini_live_event(mock_multi_text) + assert isinstance(multi_text_event, TranscriptStreamEvent) + assert multi_text_event.text == "Hello from Gemini" # Concatenated with space + # Test audio output (base64 encoded) import base64 mock_audio = unittest.mock.Mock() diff --git a/tests_integ/bidirectional_streaming/test_bidirectional_agent.py b/tests_integ/bidirectional_streaming/test_bidirectional_agent.py index 80b32b1782..f23e6b84f6 100644 --- a/tests_integ/bidirectional_streaming/test_bidirectional_agent.py +++ b/tests_integ/bidirectional_streaming/test_bidirectional_agent.py @@ -82,24 +82,15 @@ def calculator(operation: str, x: float, y: float) -> float: "env_vars": ["OPENAI_API_KEY"], "skip_reason": "OPENAI_API_KEY not available", }, - # NOTE: Gemini Live is temporarily disabled in parameterized tests - # Issue: Transcript events are not being properly emitted alongside audio events - # The model responds with audio but the test infrastructure expects text/transcripts - # TODO: Fix Gemini Live event emission to yield both transcript and audio events - # "gemini_live": { - # "model_class": GeminiLiveModel, - # "model_kwargs": { - # "model_id": "gemini-2.5-flash-native-audio-preview-09-2025", - # "params": { - # "response_modalities": ["AUDIO"], - # "output_audio_transcription": {}, - # "input_audio_transcription": {}, - # }, - # }, - # "silence_duration": 3.0, - # "env_vars": ["GOOGLE_AI_API_KEY"], - # "skip_reason": "GOOGLE_AI_API_KEY not available", - # }, + "gemini_live": { + "model_class": GeminiLiveModel, + "model_kwargs": { + # Uses default model and config (audio output + transcription enabled) + }, + "silence_duration": 1.5, # Gemini has good VAD, similar to OpenAI + "env_vars": ["GOOGLE_AI_API_KEY"], + "skip_reason": "GOOGLE_AI_API_KEY not available", + }, } From f7c18d46ca03d77bb048d7d45ab48e7ad46a1a01 Mon Sep 17 00:00:00 2001 From: Murat Kaan Meral Date: Mon, 10 Nov 2025 16:57:54 +0300 Subject: [PATCH 2/4] fix: improve gemini test script to display interrupts and remove excessive logging --- .../bidirectional_streaming/tests/test_gemini_live.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/strands/experimental/bidirectional_streaming/tests/test_gemini_live.py b/src/strands/experimental/bidirectional_streaming/tests/test_gemini_live.py index e9d715b499..bf427812cd 100644 --- a/src/strands/experimental/bidirectional_streaming/tests/test_gemini_live.py +++ b/src/strands/experimental/bidirectional_streaming/tests/test_gemini_live.py @@ -41,9 +41,9 @@ from strands.experimental.bidirectional_streaming.models.gemini_live import GeminiLiveModel # Configure logging - debug only for Gemini Live, info for everything else -logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s') +logging.basicConfig(level=logging.WARN, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s') gemini_logger = logging.getLogger('strands.experimental.bidirectional_streaming.models.gemini_live') -gemini_logger.setLevel(logging.DEBUG) +gemini_logger.setLevel(logging.WARN) logger = logging.getLogger(__name__) @@ -162,7 +162,7 @@ async def receive(agent, context): # Handle interruption events (bidirectional_interruption) elif event_type == "bidirectional_interruption": context["interrupted"] = True - logger.info("Interruption detected") + print("⚠️ Interruption detected") # Handle transcript events (bidirectional_transcript_stream) elif event_type == "bidirectional_transcript_stream": From c06ec6f930a812fd879fa28f45db8a757ad8c1c5 Mon Sep 17 00:00:00 2001 From: Murat Kaan Meral Date: Tue, 11 Nov 2025 13:34:48 +0300 Subject: [PATCH 3/4] remove text logging --- .../experimental/bidirectional_streaming/models/gemini_live.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/src/strands/experimental/bidirectional_streaming/models/gemini_live.py b/src/strands/experimental/bidirectional_streaming/models/gemini_live.py index 6be0275f8d..47b38e0eb2 100644 --- a/src/strands/experimental/bidirectional_streaming/models/gemini_live.py +++ b/src/strands/experimental/bidirectional_streaming/models/gemini_live.py @@ -278,7 +278,6 @@ def _convert_gemini_live_event(self, message: LiveServerMessage) -> Optional[Dic for part in model_turn.parts: # Log all part types for debugging part_attrs = {attr: getattr(part, attr, None) for attr in dir(part) if not attr.startswith('_')} - logger.debug(f"Model turn part attributes: {part_attrs}") # Check if part has text attribute and it's not empty if hasattr(part, 'text') and part.text: @@ -286,7 +285,6 @@ def _convert_gemini_live_event(self, message: LiveServerMessage) -> Optional[Dic if text_parts: full_text = " ".join(text_parts) - logger.debug(f"Text output as transcript ({len(text_parts)} parts): {full_text}") return BidiTranscriptStreamEvent( delta={"text": full_text}, text=full_text, From 490ce1e29a4bec20495bb339840e6b65c17bcad5 Mon Sep 17 00:00:00 2001 From: Murat Kaan Meral Date: Tue, 11 Nov 2025 18:17:40 +0300 Subject: [PATCH 4/4] fix(gemini): return multiple tool use events --- .../models/gemini_live.py | 53 +++++++++------- .../models/test_gemini_live.py | 63 +++++++++++++++++-- 2 files changed, 86 insertions(+), 30 deletions(-) diff --git a/src/strands/experimental/bidirectional_streaming/models/gemini_live.py b/src/strands/experimental/bidirectional_streaming/models/gemini_live.py index 47b38e0eb2..9bb5bba774 100644 --- a/src/strands/experimental/bidirectional_streaming/models/gemini_live.py +++ b/src/strands/experimental/bidirectional_streaming/models/gemini_live.py @@ -33,6 +33,7 @@ BidiImageInputEvent, BidiInputEvent, BidiInterruptionEvent, + BidiOutputEvent, BidiUsageEvent, BidiTextInputEvent, BidiTranscriptStreamEvent, @@ -173,7 +174,7 @@ async def _send_message_history(self, messages: Messages) -> None: content = genai_types.Content(role=role, parts=content_parts) await self.live_session.send_client_content(turns=content) - async def receive(self) -> AsyncIterable[Dict[str, Any]]: + async def receive(self) -> AsyncIterable[BidiOutputEvent]: """Receive Gemini Live API events and convert to provider-agnostic format.""" # Emit connection start event @@ -190,10 +191,9 @@ async def receive(self) -> AsyncIterable[Dict[str, Any]]: if not self._active: break - # Convert to provider-agnostic format - provider_event = self._convert_gemini_live_event(message) - if provider_event: - yield provider_event + # Convert to provider-agnostic format (always returns list) + for event in self._convert_gemini_live_event(message): + yield event # SDK exits receive loop after turn_complete - restart automatically if self._active: @@ -211,7 +211,7 @@ async def receive(self) -> AsyncIterable[Dict[str, Any]]: # Emit connection close event when exiting yield BidiConnectionCloseEvent(connection_id=self.connection_id, reason="complete") - def _convert_gemini_live_event(self, message: LiveServerMessage) -> Optional[Dict[str, Any]]: + def _convert_gemini_live_event(self, message: LiveServerMessage) -> List[BidiOutputEvent]: """Convert Gemini Live API events to provider-agnostic format. Handles different types of content: @@ -219,11 +219,14 @@ def _convert_gemini_live_event(self, message: LiveServerMessage) -> Optional[Dic - outputTranscription: Model's audio transcribed to text - modelTurn text: Text response from the model - usageMetadata: Token usage information + + Returns: + List of event dicts (empty list if no events to emit). """ try: # Handle interruption first (from server_content) if message.server_content and message.server_content.interrupted: - return BidiInterruptionEvent(reason="user_speech") + return [BidiInterruptionEvent(reason="user_speech")] # Handle input transcription (user's speech) - emit as transcript event if message.server_content and message.server_content.input_transcription: @@ -233,13 +236,13 @@ def _convert_gemini_live_event(self, message: LiveServerMessage) -> Optional[Dic transcription_text = input_transcript.text role = getattr(input_transcript, 'role', 'user') logger.debug(f"Input transcription detected: {transcription_text}") - return BidiTranscriptStreamEvent( + return [BidiTranscriptStreamEvent( delta={"text": transcription_text}, text=transcription_text, role=role.lower() if isinstance(role, str) else "user", is_final=True, current_transcript=transcription_text - ) + )] # Handle output transcription (model's audio) - emit as transcript event if message.server_content and message.server_content.output_transcription: @@ -249,25 +252,25 @@ def _convert_gemini_live_event(self, message: LiveServerMessage) -> Optional[Dic transcription_text = output_transcript.text role = getattr(output_transcript, 'role', 'assistant') logger.debug(f"Output transcription detected: {transcription_text}") - return BidiTranscriptStreamEvent( + return [BidiTranscriptStreamEvent( delta={"text": transcription_text}, text=transcription_text, role=role.lower() if isinstance(role, str) else "assistant", is_final=True, current_transcript=transcription_text - ) + )] # Handle audio output using SDK's built-in data property # Check this BEFORE text to avoid triggering warning on mixed content if message.data: # Convert bytes to base64 string for JSON serializability audio_b64 = base64.b64encode(message.data).decode('utf-8') - return BidiAudioStreamEvent( + return [BidiAudioStreamEvent( audio=audio_b64, format="pcm", sample_rate=GEMINI_OUTPUT_SAMPLE_RATE, channels=GEMINI_CHANNELS - ) + )] # Handle text output from model_turn (avoids warning by checking parts directly) if message.server_content and message.server_content.model_turn: @@ -285,27 +288,29 @@ def _convert_gemini_live_event(self, message: LiveServerMessage) -> Optional[Dic if text_parts: full_text = " ".join(text_parts) - return BidiTranscriptStreamEvent( + return [BidiTranscriptStreamEvent( delta={"text": full_text}, text=full_text, role="assistant", is_final=True, current_transcript=full_text - ) + )] - # Handle tool calls + # Handle tool calls - return list to support multiple tool calls if message.tool_call and message.tool_call.function_calls: + tool_events = [] for func_call in message.tool_call.function_calls: tool_use_event: ToolUse = { "toolUseId": func_call.id, "name": func_call.name, "input": func_call.args or {} } - # Return ToolUseStreamEvent for consistency with standard agent - return ToolUseStreamEvent( + # Create ToolUseStreamEvent for consistency with standard agent + tool_events.append(ToolUseStreamEvent( delta={"toolUse": tool_use_event}, current_tool_use=tool_use_event - ) + )) + return tool_events # Handle usage metadata if hasattr(message, 'usage_metadata') and message.usage_metadata: @@ -340,23 +345,23 @@ def _convert_gemini_live_event(self, message: LiveServerMessage) -> Optional[Dic "output_tokens": detail.token_count }) - return BidiUsageEvent( + return [BidiUsageEvent( input_tokens=usage.prompt_token_count or 0, output_tokens=usage.response_token_count or 0, total_tokens=usage.total_token_count or 0, modality_details=modality_details if modality_details else None, cache_read_input_tokens=usage.cached_content_token_count if usage.cached_content_token_count else None - ) + )] # Silently ignore setup_complete and generation_complete messages - return None + return [] except Exception as e: logger.error("Error converting Gemini Live event: %s", e) logger.error("Message type: %s", type(message).__name__) logger.error("Message attributes: %s", [attr for attr in dir(message) if not attr.startswith('_')]) - # Return ErrorEvent instead of None so caller can handle it - return BidiErrorEvent(error=e) + # Return ErrorEvent in list so caller can handle it + return [BidiErrorEvent(error=e)] async def send( self, diff --git a/tests/strands/experimental/bidirectional_streaming/models/test_gemini_live.py b/tests/strands/experimental/bidirectional_streaming/models/test_gemini_live.py index 85e416164c..272314272d 100644 --- a/tests/strands/experimental/bidirectional_streaming/models/test_gemini_live.py +++ b/tests/strands/experimental/bidirectional_streaming/models/test_gemini_live.py @@ -315,7 +315,10 @@ async def test_event_conversion(mock_genai_client, model): mock_text.server_content = mock_server_content - text_event = model._convert_gemini_live_event(mock_text) + text_events = model._convert_gemini_live_event(mock_text) + assert isinstance(text_events, list) + assert len(text_events) == 1 + text_event = text_events[0] assert isinstance(text_event, BidiTranscriptStreamEvent) assert text_event.get("type") == "bidi_transcript_stream" assert text_event.text == "Hello from Gemini" @@ -344,7 +347,10 @@ async def test_event_conversion(mock_genai_client, model): mock_multi_text.server_content = mock_server_content_multi - multi_text_event = model._convert_gemini_live_event(mock_multi_text) + multi_text_events = model._convert_gemini_live_event(mock_multi_text) + assert isinstance(multi_text_events, list) + assert len(multi_text_events) == 1 + multi_text_event = multi_text_events[0] assert isinstance(multi_text_event, BidiTranscriptStreamEvent) assert multi_text_event.text == "Hello from Gemini" # Concatenated with space @@ -355,7 +361,10 @@ async def test_event_conversion(mock_genai_client, model): mock_audio.tool_call = None mock_audio.server_content = None - audio_event = model._convert_gemini_live_event(mock_audio) + audio_events = model._convert_gemini_live_event(mock_audio) + assert isinstance(audio_events, list) + assert len(audio_events) == 1 + audio_event = audio_events[0] assert isinstance(audio_event, BidiAudioStreamEvent) assert audio_event.get("type") == "bidi_audio_stream" # Audio is now base64 encoded @@ -363,7 +372,7 @@ async def test_event_conversion(mock_genai_client, model): assert audio_event.audio == expected_b64 assert audio_event.format == "pcm" - # Test tool call + # Test single tool call (returns list with one event) mock_func_call = unittest.mock.Mock() mock_func_call.id = "tool-123" mock_func_call.name = "calculator" @@ -378,13 +387,52 @@ async def test_event_conversion(mock_genai_client, model): mock_tool.tool_call = mock_tool_call mock_tool.server_content = None - tool_event = model._convert_gemini_live_event(mock_tool) + tool_events = model._convert_gemini_live_event(mock_tool) + # Should return a list of ToolUseStreamEvent + assert isinstance(tool_events, list) + assert len(tool_events) == 1 + tool_event = tool_events[0] # ToolUseStreamEvent has delta and current_tool_use, not a "type" field assert "delta" in tool_event assert "toolUse" in tool_event["delta"] assert tool_event["delta"]["toolUse"]["toolUseId"] == "tool-123" assert tool_event["delta"]["toolUse"]["name"] == "calculator" + # Test multiple tool calls (returns list with multiple events) + mock_func_call_1 = unittest.mock.Mock() + mock_func_call_1.id = "tool-123" + mock_func_call_1.name = "calculator" + mock_func_call_1.args = {"expression": "2+2"} + + mock_func_call_2 = unittest.mock.Mock() + mock_func_call_2.id = "tool-456" + mock_func_call_2.name = "weather" + mock_func_call_2.args = {"location": "Seattle"} + + mock_tool_call_multi = unittest.mock.Mock() + mock_tool_call_multi.function_calls = [mock_func_call_1, mock_func_call_2] + + mock_tool_multi = unittest.mock.Mock() + mock_tool_multi.text = None + mock_tool_multi.data = None + mock_tool_multi.tool_call = mock_tool_call_multi + mock_tool_multi.server_content = None + + tool_events_multi = model._convert_gemini_live_event(mock_tool_multi) + # Should return a list with two ToolUseStreamEvent + assert isinstance(tool_events_multi, list) + assert len(tool_events_multi) == 2 + + # Verify first tool call + assert tool_events_multi[0]["delta"]["toolUse"]["toolUseId"] == "tool-123" + assert tool_events_multi[0]["delta"]["toolUse"]["name"] == "calculator" + assert tool_events_multi[0]["delta"]["toolUse"]["input"] == {"expression": "2+2"} + + # Verify second tool call + assert tool_events_multi[1]["delta"]["toolUse"]["toolUseId"] == "tool-456" + assert tool_events_multi[1]["delta"]["toolUse"]["name"] == "weather" + assert tool_events_multi[1]["delta"]["toolUse"]["input"] == {"location": "Seattle"} + # Test interruption mock_server_content = unittest.mock.Mock() mock_server_content.interrupted = True @@ -397,7 +445,10 @@ async def test_event_conversion(mock_genai_client, model): mock_interrupt.tool_call = None mock_interrupt.server_content = mock_server_content - interrupt_event = model._convert_gemini_live_event(mock_interrupt) + interrupt_events = model._convert_gemini_live_event(mock_interrupt) + assert isinstance(interrupt_events, list) + assert len(interrupt_events) == 1 + interrupt_event = interrupt_events[0] assert isinstance(interrupt_event, BidiInterruptionEvent) assert interrupt_event.get("type") == "bidi_interruption" assert interrupt_event.reason == "user_speech"