diff --git a/src/strands/experimental/bidirectional_streaming/models/gemini_live.py b/src/strands/experimental/bidirectional_streaming/models/gemini_live.py index 5d1763bcd4..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, @@ -59,7 +60,7 @@ class BidiGeminiLiveModel(BidiModel): 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 @@ -75,7 +76,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 = {} @@ -161,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 @@ -178,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: @@ -199,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: @@ -207,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: @@ -221,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: @@ -237,50 +252,65 @@ 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 text output from model - if message.text: - role = getattr(message, 'role', 'assistant') - logger.debug(f"Text output as transcript: {message.text}") - return BidiTranscriptStreamEvent( - delta={"text": message.text}, - text=message.text, - role=role.lower() if isinstance(role, str) else "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') - 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: + 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('_')} + + # 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) + 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: @@ -315,22 +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 None + # Return ErrorEvent in list so caller can handle it + return [BidiErrorEvent(error=e)] async def send( self, diff --git a/src/strands/experimental/bidirectional_streaming/scripts/test_gemini_live.py b/src/strands/experimental/bidirectional_streaming/scripts/test_gemini_live.py index ba0d9edf78..814586de1c 100644 --- a/src/strands/experimental/bidirectional_streaming/scripts/test_gemini_live.py +++ b/src/strands/experimental/bidirectional_streaming/scripts/test_gemini_live.py @@ -40,8 +40,10 @@ from strands.experimental.bidirectional_streaming.agent.agent import BidiAgent from strands.experimental.bidirectional_streaming.models.gemini_live import BidiGeminiLiveModel -# Configure logging -logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s') +# Configure logging - debug only for Gemini Live, info for everything else +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.WARN) logger = logging.getLogger(__name__) @@ -310,17 +312,10 @@ async def main(duration=180): # Initialize Gemini Live model with proper configuration logger.info("Initializing Gemini Live model with API key") - model = BidiGeminiLiveModel( - 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 = BidiGeminiLiveModel(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 = BidiAgent( 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 5e8e7a80d1..272314272d 100644 --- a/tests/strands/experimental/bidirectional_streaming/models/test_gemini_live.py +++ b/tests/strands/experimental/bidirectional_streaming/models/test_gemini_live.py @@ -96,20 +96,28 @@ def test_model_initialization(mock_genai_client, model_id, api_key): # Test default config model_default = BidiGeminiLiveModel() - 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 = BidiGeminiLiveModel(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 = BidiGeminiLiveModel(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 @@ -288,14 +296,29 @@ async def test_event_conversion(mock_genai_client, model): _, _, _ = mock_genai_client await model.start() - # 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 - text_event = model._convert_gemini_live_event(mock_text) + # 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_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" @@ -304,6 +327,33 @@ 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_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 + # Test audio output (base64 encoded) mock_audio = unittest.mock.Mock() mock_audio.text = None @@ -311,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 @@ -319,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" @@ -334,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 @@ -353,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" diff --git a/tests_integ/bidirectional_streaming/test_bidirectional_agent.py b/tests_integ/bidirectional_streaming/test_bidirectional_agent.py index 53bb0f2e3b..e93a267a06 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": BidiGeminiLiveModel, - # "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": BidiGeminiLiveModel, + "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", + }, }