mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
feat: support realtime input config
Merge https://github.com/google/adk-python/pull/981 issue: https://github.com/google/adk-python/issues/982 This pull request introduces a new configuration option, `realtime_input_config`, to the `RunConfig` class. **Reason for this change:** Currently, there is no direct way to configure real-time audio input behaviors, such as Voice Activity Detection (VAD), for live agents through the `RunConfig`. The Gemini API documentation (specifically [Configure automatic VAD](https://ai.google.dev/gemini-api/docs/live#configure-automatic-vad)) outlines parameters for VAD that users may want to customize. This change enables users to pass these real-time input configurations, providing more granular control over the audio input for live agents. **Changes made:** - Added a new optional field `realtime_input_config: Optional[types.RealtimeInputConfig]` to the `RunConfig` class. - The docstring for `realtime_input_config` has been added to explain its purpose. **Example Usage (Conceptual):** While the specific structure of `types.RealtimeInputConfig` would define the exact parameters, a user might configure it like this: ```python # (Assuming types.RealtimeInputConfig and types.VadConfig are defined elsewhere) # import your_project.types as types run_config = RunConfig( # ... other configurations ... realtime_input_config=types.RealtimeInputConfig( automatic_activity_detection =types.AutomaticActivityDetection( # VAD specific parameters like sensitivity, endpoint_duration_millis etc. # based on https://ai.google.dev/gemini-api/docs/live#configure-automatic-vad ) # Potentially other real-time input settings could be added here in the future ) ) COPYBARA_INTEGRATE_REVIEW=https://github.com/google/adk-python/pull/981 from ammmr:patch-add-realtime-input-config b2e17fbf5742d264029ad49bf632422b5c5b1e0a PiperOrigin-RevId: 770797640
This commit is contained in:
committed by
Copybara-Service
parent
2ff9b1f639
commit
d22920bd7f
@@ -12,6 +12,8 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import Enum
|
||||
import logging
|
||||
import sys
|
||||
@@ -68,6 +70,9 @@ class RunConfig(BaseModel):
|
||||
input_audio_transcription: Optional[types.AudioTranscriptionConfig] = None
|
||||
"""Input transcription for live agents with audio input from user."""
|
||||
|
||||
realtime_input_config: Optional[types.RealtimeInputConfig] = None
|
||||
"""Realtime input config for live agents with audio input from user."""
|
||||
|
||||
max_llm_calls: int = 500
|
||||
"""
|
||||
A limit on the total number of llm calls for a given run.
|
||||
|
||||
@@ -65,6 +65,9 @@ class _BasicLlmRequestProcessor(BaseLlmRequestProcessor):
|
||||
llm_request.live_connect_config.input_audio_transcription = (
|
||||
invocation_context.run_config.input_audio_transcription
|
||||
)
|
||||
llm_request.live_connect_config.realtime_input_config = (
|
||||
invocation_context.run_config.realtime_input_config
|
||||
)
|
||||
|
||||
# TODO: handle tool append here, instead of in BaseTool.process_llm_request.
|
||||
|
||||
|
||||
@@ -12,6 +12,8 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import AsyncGenerator
|
||||
|
||||
|
||||
@@ -0,0 +1,201 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from unittest import mock
|
||||
|
||||
from google.adk.agents import Agent
|
||||
from google.adk.agents.live_request_queue import LiveRequest
|
||||
from google.adk.agents.live_request_queue import LiveRequestQueue
|
||||
from google.adk.agents.run_config import RunConfig
|
||||
from google.adk.flows.llm_flows.base_llm_flow import BaseLlmFlow
|
||||
from google.adk.models.llm_request import LlmRequest
|
||||
from google.genai import types
|
||||
import pytest
|
||||
|
||||
from ... import testing_utils
|
||||
|
||||
|
||||
class TestBaseLlmFlow(BaseLlmFlow):
|
||||
"""Test implementation of BaseLlmFlow for testing purposes."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def test_blob():
|
||||
"""Test blob for audio data."""
|
||||
return types.Blob(data=b'\x00\xFF\x00\xFF', mime_type='audio/pcm')
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_llm_connection():
|
||||
"""Mock LLM connection for testing."""
|
||||
connection = mock.AsyncMock()
|
||||
connection.send_realtime = mock.AsyncMock()
|
||||
return connection
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_to_model_with_disabled_vad(test_blob, mock_llm_connection):
|
||||
"""Test _send_to_model with automatic_activity_detection.disabled=True."""
|
||||
# Create LlmRequest with disabled VAD
|
||||
realtime_input_config = types.RealtimeInputConfig(
|
||||
automatic_activity_detection=types.AutomaticActivityDetection(
|
||||
disabled=True
|
||||
)
|
||||
)
|
||||
|
||||
# Create invocation context with live request queue
|
||||
agent = Agent(name='test_agent', model='mock')
|
||||
invocation_context = await testing_utils.create_invocation_context(
|
||||
agent=agent,
|
||||
user_content='',
|
||||
run_config=RunConfig(realtime_input_config=realtime_input_config),
|
||||
)
|
||||
invocation_context.live_request_queue = LiveRequestQueue()
|
||||
|
||||
# Create flow and start _send_to_model task
|
||||
flow = TestBaseLlmFlow()
|
||||
|
||||
# Send a blob to the queue
|
||||
live_request = LiveRequest(blob=test_blob)
|
||||
invocation_context.live_request_queue.send(live_request)
|
||||
invocation_context.live_request_queue.close()
|
||||
|
||||
# Run _send_to_model
|
||||
await flow._send_to_model(mock_llm_connection, invocation_context)
|
||||
|
||||
mock_llm_connection.send_realtime.assert_called_once_with(test_blob)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_to_model_with_enabled_vad(test_blob, mock_llm_connection):
|
||||
"""Test _send_to_model with automatic_activity_detection.disabled=False.
|
||||
|
||||
Custom VAD activity signal is not supported so we should still disable it.
|
||||
"""
|
||||
# Create LlmRequest with enabled VAD
|
||||
realtime_input_config = types.RealtimeInputConfig(
|
||||
automatic_activity_detection=types.AutomaticActivityDetection(
|
||||
disabled=False
|
||||
)
|
||||
)
|
||||
|
||||
# Create invocation context with live request queue
|
||||
agent = Agent(name='test_agent', model='mock')
|
||||
invocation_context = await testing_utils.create_invocation_context(
|
||||
agent=agent, user_content=''
|
||||
)
|
||||
invocation_context.live_request_queue = LiveRequestQueue()
|
||||
|
||||
# Create flow and start _send_to_model task
|
||||
flow = TestBaseLlmFlow()
|
||||
|
||||
# Send a blob to the queue
|
||||
live_request = LiveRequest(blob=test_blob)
|
||||
invocation_context.live_request_queue.send(live_request)
|
||||
invocation_context.live_request_queue.close()
|
||||
|
||||
# Run _send_to_model
|
||||
await flow._send_to_model(mock_llm_connection, invocation_context)
|
||||
|
||||
mock_llm_connection.send_realtime.assert_called_once_with(test_blob)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_to_model_without_realtime_config(
|
||||
test_blob, mock_llm_connection
|
||||
):
|
||||
"""Test _send_to_model without realtime_input_config (default behavior)."""
|
||||
# Create invocation context with live request queue
|
||||
agent = Agent(name='test_agent', model='mock')
|
||||
invocation_context = await testing_utils.create_invocation_context(
|
||||
agent=agent, user_content=''
|
||||
)
|
||||
invocation_context.live_request_queue = LiveRequestQueue()
|
||||
|
||||
# Create flow and start _send_to_model task
|
||||
flow = TestBaseLlmFlow()
|
||||
|
||||
# Send a blob to the queue
|
||||
live_request = LiveRequest(blob=test_blob)
|
||||
invocation_context.live_request_queue.send(live_request)
|
||||
invocation_context.live_request_queue.close()
|
||||
|
||||
# Run _send_to_model
|
||||
await flow._send_to_model(mock_llm_connection, invocation_context)
|
||||
|
||||
mock_llm_connection.send_realtime.assert_called_once_with(test_blob)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_to_model_with_none_automatic_activity_detection(
|
||||
test_blob, mock_llm_connection
|
||||
):
|
||||
"""Test _send_to_model with automatic_activity_detection=None."""
|
||||
# Create LlmRequest with None automatic_activity_detection
|
||||
realtime_input_config = types.RealtimeInputConfig(
|
||||
automatic_activity_detection=None
|
||||
)
|
||||
|
||||
# Create invocation context with live request queue
|
||||
agent = Agent(name='test_agent', model='mock')
|
||||
invocation_context = await testing_utils.create_invocation_context(
|
||||
agent=agent,
|
||||
user_content='',
|
||||
run_config=RunConfig(realtime_input_config=realtime_input_config),
|
||||
)
|
||||
invocation_context.live_request_queue = LiveRequestQueue()
|
||||
|
||||
# Create flow and start _send_to_model task
|
||||
flow = TestBaseLlmFlow()
|
||||
|
||||
# Send a blob to the queue
|
||||
live_request = LiveRequest(blob=test_blob)
|
||||
invocation_context.live_request_queue.send(live_request)
|
||||
invocation_context.live_request_queue.close()
|
||||
|
||||
# Run _send_to_model
|
||||
await flow._send_to_model(mock_llm_connection, invocation_context)
|
||||
|
||||
mock_llm_connection.send_realtime.assert_called_once_with(test_blob)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_to_model_with_text_content(mock_llm_connection):
|
||||
"""Test _send_to_model with text content (not blob)."""
|
||||
# Create invocation context with live request queue
|
||||
agent = Agent(name='test_agent', model='mock')
|
||||
invocation_context = await testing_utils.create_invocation_context(
|
||||
agent=agent, user_content=''
|
||||
)
|
||||
invocation_context.live_request_queue = LiveRequestQueue()
|
||||
|
||||
# Create flow and start _send_to_model task
|
||||
flow = TestBaseLlmFlow()
|
||||
|
||||
# Send text content to the queue
|
||||
content = types.Content(
|
||||
role='user', parts=[types.Part.from_text(text='Hello')]
|
||||
)
|
||||
live_request = LiveRequest(content=content)
|
||||
invocation_context.live_request_queue.send(live_request)
|
||||
invocation_context.live_request_queue.close()
|
||||
|
||||
# Run _send_to_model
|
||||
await flow._send_to_model(mock_llm_connection, invocation_context)
|
||||
|
||||
# Verify send_content was called instead of send_realtime
|
||||
mock_llm_connection.send_content.assert_called_once_with(content)
|
||||
mock_llm_connection.send_realtime.assert_not_called()
|
||||
@@ -0,0 +1,111 @@
|
||||
# Copyright 2025 Google LLC
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from unittest import mock
|
||||
|
||||
from google.adk.models.gemini_llm_connection import GeminiLlmConnection
|
||||
from google.genai import types
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_gemini_session():
|
||||
"""Mock Gemini session for testing."""
|
||||
return mock.AsyncMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def gemini_connection(mock_gemini_session):
|
||||
"""GeminiLlmConnection instance with mocked session."""
|
||||
return GeminiLlmConnection(mock_gemini_session)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def test_blob():
|
||||
"""Test blob for audio data."""
|
||||
return types.Blob(data=b'\x00\xFF\x00\xFF', mime_type='audio/pcm')
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_realtime_default_behavior(
|
||||
gemini_connection, mock_gemini_session, test_blob
|
||||
):
|
||||
"""Test send_realtime with default automatic_activity_detection value (True)."""
|
||||
await gemini_connection.send_realtime(test_blob)
|
||||
|
||||
# Should call send once
|
||||
mock_gemini_session.send.assert_called_once_with(input=test_blob.model_dump())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_history(gemini_connection, mock_gemini_session):
|
||||
"""Test send_history method."""
|
||||
history = [
|
||||
types.Content(role='user', parts=[types.Part.from_text(text='Hello')]),
|
||||
types.Content(
|
||||
role='model', parts=[types.Part.from_text(text='Hi there!')]
|
||||
),
|
||||
]
|
||||
|
||||
await gemini_connection.send_history(history)
|
||||
|
||||
mock_gemini_session.send.assert_called_once()
|
||||
call_args = mock_gemini_session.send.call_args[1]
|
||||
assert 'input' in call_args
|
||||
assert call_args['input'].turns == history
|
||||
assert call_args['input'].turn_complete is False # Last message is from model
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_content_text(gemini_connection, mock_gemini_session):
|
||||
"""Test send_content with text content."""
|
||||
content = types.Content(
|
||||
role='user', parts=[types.Part.from_text(text='Hello')]
|
||||
)
|
||||
|
||||
await gemini_connection.send_content(content)
|
||||
|
||||
mock_gemini_session.send.assert_called_once()
|
||||
call_args = mock_gemini_session.send.call_args[1]
|
||||
assert 'input' in call_args
|
||||
assert call_args['input'].turns == [content]
|
||||
assert call_args['input'].turn_complete is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_content_function_response(
|
||||
gemini_connection, mock_gemini_session
|
||||
):
|
||||
"""Test send_content with function response."""
|
||||
function_response = types.FunctionResponse(
|
||||
name='test_function', response={'result': 'success'}
|
||||
)
|
||||
content = types.Content(
|
||||
role='user', parts=[types.Part(function_response=function_response)]
|
||||
)
|
||||
|
||||
await gemini_connection.send_content(content)
|
||||
|
||||
mock_gemini_session.send.assert_called_once()
|
||||
call_args = mock_gemini_session.send.call_args[1]
|
||||
assert 'input' in call_args
|
||||
assert call_args['input'].function_responses == [function_response]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_close(gemini_connection, mock_gemini_session):
|
||||
"""Test close method."""
|
||||
await gemini_connection.close()
|
||||
|
||||
mock_gemini_session.close.assert_called_once()
|
||||
@@ -56,7 +56,9 @@ class ModelContent(types.Content):
|
||||
super().__init__(role='model', parts=parts)
|
||||
|
||||
|
||||
async def create_invocation_context(agent: Agent, user_content: str = ''):
|
||||
async def create_invocation_context(
|
||||
agent: Agent, user_content: str = '', run_config: RunConfig = None
|
||||
):
|
||||
invocation_id = 'test_id'
|
||||
artifact_service = InMemoryArtifactService()
|
||||
session_service = InMemorySessionService()
|
||||
@@ -73,7 +75,7 @@ async def create_invocation_context(agent: Agent, user_content: str = ''):
|
||||
user_content=types.Content(
|
||||
role='user', parts=[types.Part.from_text(text=user_content)]
|
||||
),
|
||||
run_config=RunConfig(),
|
||||
run_config=run_config or RunConfig(),
|
||||
)
|
||||
if user_content:
|
||||
append_user_content(
|
||||
@@ -205,13 +207,16 @@ class InMemoryRunner:
|
||||
events.append(event)
|
||||
return events
|
||||
|
||||
def run_live(self, live_request_queue: LiveRequestQueue) -> list[Event]:
|
||||
def run_live(
|
||||
self, live_request_queue: LiveRequestQueue, run_config: RunConfig = None
|
||||
) -> list[Event]:
|
||||
collected_responses = []
|
||||
|
||||
async def consume_responses(session: Session):
|
||||
run_res = self.runner.run_live(
|
||||
session=session,
|
||||
live_request_queue=live_request_queue,
|
||||
run_config=run_config or RunConfig(),
|
||||
)
|
||||
|
||||
async for response in run_res:
|
||||
|
||||
Reference in New Issue
Block a user