diff --git a/src/google/adk/agents/live_request_queue.py b/src/google/adk/agents/live_request_queue.py index c9c7b82c685..e5fc0079dda 100644 --- a/src/google/adk/agents/live_request_queue.py +++ b/src/google/adk/agents/live_request_queue.py @@ -24,39 +24,33 @@ class LiveRequest(BaseModel): - """Request send to live agents.""" - - model_config = ConfigDict(ser_json_bytes='base64', val_json_bytes='base64') - """The pydantic model config.""" - - content: Optional[types.Content] = None - """If set, send the content to the model in turn-by-turn mode. + """Request send to live agents. When multiple fields are set, they are processed by priority (highest first): activity_start > activity_end > blob > content. state_delta, if set, is always applied regardless of the other fields. """ - blob: Optional[types.Blob] = None - """If set, send the blob to the model in realtime mode. - When multiple fields are set, they are processed by priority (highest first): - activity_start > activity_end > blob > content. state_delta, if set, is always - applied regardless of the other fields. - """ - activity_start: Optional[types.ActivityStart] = None - """If set, signal the start of user activity to the model. + model_config = ConfigDict(ser_json_bytes='base64', val_json_bytes='base64') + """The pydantic model config.""" - When multiple fields are set, they are processed by priority (highest first): - activity_start > activity_end > blob > content. state_delta, if set, is always - applied regardless of the other fields. - """ + content: Optional[types.Content] = None + """If set, send the content to the model in turn-by-turn mode.""" + + blob: Optional[types.Blob] = None + """If set, send the blob to the model in realtime mode.""" + + activity_start: Optional[types.ActivityStart] = None + """If set, signal the start of user activity to the model.""" + activity_end: Optional[types.ActivityEnd] = None - """If set, signal the end of user activity to the model. - - When multiple fields are set, they are processed by priority (highest first): - activity_start > activity_end > blob > content. state_delta, if set, is always - applied regardless of the other fields. + """If set, signal the end of user activity to the model.""" + + audio_stream_end: bool = False + """If set, signal the end of the audio stream to the model. + This is only used when Voice Activity Detection is enabled. """ + close: bool = False """If set, close the queue. queue.shutdown() is only supported in Python 3.13+.""" @@ -92,6 +86,10 @@ def send_activity_end(self) -> None: """Sends an activity end signal to mark the end of user input.""" self._queue.put_nowait(LiveRequest(activity_end=types.ActivityEnd())) + def send_audio_stream_end(self) -> None: + """Sends an audio stream end signal to force flush audio.""" + self._queue.put_nowait(LiveRequest(audio_stream_end=True)) + def send(self, req: LiveRequest) -> None: self._queue.put_nowait(req) diff --git a/src/google/adk/flows/llm_flows/base_llm_flow.py b/src/google/adk/flows/llm_flows/base_llm_flow.py index 2a799d54ff3..b66536f0790 100644 --- a/src/google/adk/flows/llm_flows/base_llm_flow.py +++ b/src/google/adk/flows/llm_flows/base_llm_flow.py @@ -827,6 +827,10 @@ async def _send_to_model( await llm_connection.send_realtime(types.ActivityStart()) elif live_request.activity_end: await llm_connection.send_realtime(types.ActivityEnd()) + elif live_request.audio_stream_end: + await llm_connection.send_realtime( + types.LiveClientRealtimeInput(audio_stream_end=True) + ) elif live_request.blob: # Cache input audio chunks before flushing self.audio_cache_manager.cache_audio( diff --git a/src/google/adk/models/gemini_llm_connection.py b/src/google/adk/models/gemini_llm_connection.py index 66534b39fdc..5b1fa2a48ff 100644 --- a/src/google/adk/models/gemini_llm_connection.py +++ b/src/google/adk/models/gemini_llm_connection.py @@ -29,7 +29,12 @@ logger = logging.getLogger('google_adk.' + __name__) -RealtimeInput = Union[types.Blob, types.ActivityStart, types.ActivityEnd] +RealtimeInput = Union[ + types.Blob, + types.ActivityStart, + types.ActivityEnd, + types.LiveClientRealtimeInput, +] from typing import TYPE_CHECKING if TYPE_CHECKING: @@ -173,6 +178,13 @@ async def send_realtime(self, input: RealtimeInput) -> None: elif isinstance(input, types.ActivityEnd): logger.debug('Sending LLM activity end signal.') await self._gemini_session.send_realtime_input(activity_end=input) + + elif isinstance(input, types.LiveClientRealtimeInput): + if input.audio_stream_end: + logger.debug('Sending LLM audio stream end signal.') + await self._gemini_session.send_realtime_input(audio_stream_end=True) + else: + logger.warning('Unary LiveClientRealtimeInput not fully supported yet.') else: raise ValueError('Unsupported input type: %s' % type(input)) diff --git a/tests/unittests/models/test_gemini_llm_connection.py b/tests/unittests/models/test_gemini_llm_connection.py index 5bded80a266..561fe5885b6 100644 --- a/tests/unittests/models/test_gemini_llm_connection.py +++ b/tests/unittests/models/test_gemini_llm_connection.py @@ -72,6 +72,38 @@ async def test_send_realtime_default_behavior( @pytest.mark.asyncio +async def test_send_realtime_audio_stream_end( + gemini_connection, mock_gemini_session +): + """Test send_realtime with LiveClientRealtimeInput(audio_stream_end=True).""" + input_signal = types.LiveClientRealtimeInput(audio_stream_end=True) + await gemini_connection.send_realtime(input_signal) + + # Should call send_realtime_input with audio_stream_end=True + mock_gemini_session.send_realtime_input.assert_called_once_with( + audio_stream_end=True + ) + # Should not call .send function + mock_gemini_session.send.assert_not_called() + + +@pytest.mark.asyncio +async def test_send_realtime_unsupported_liveClientRealtimeInput( + gemini_connection, mock_gemini_session, caplog +): + """Test send_realtime with unsupported LiveClientRealtimeInput.""" + input_signal = types.LiveClientRealtimeInput() + + with caplog.at_level('WARNING'): + await gemini_connection.send_realtime(input_signal) + + # Should log a warning + assert 'Unary LiveClientRealtimeInput not fully supported yet.' in caplog.text + # Should not call send_realtime_input or send + mock_gemini_session.send_realtime_input.assert_not_called() + mock_gemini_session.send.assert_not_called() + + async def test_send_realtime_audio_uses_audio_channel_for_live_translate( mock_gemini_session, test_blob ):