mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
feat: Granular Per Agent Speech Configuration
Merge https://github.com/google/adk-python/pull/3170 Addresses Feature Request: #3116 This PR adds a `speech_config` to the **LLM Agent configuration** for the **live use case**. When an **asynchronous LLM** call is made to the **Gemini Live API**, it prioritizes the most specific agent configuration's speech_config. If that is null, it then uses the run configuration's speech_config. Unit tests have been added to verify this behavior. COPYBARA_INTEGRATE_REVIEW=https://github.com/google/adk-python/pull/3170 from qyuo:bidi_agent_speech_config af1bd277d4f95c4a7d9aa0b16828ba3de826ce08 PiperOrigin-RevId: 822305427
This commit is contained in:
committed by
Copybara-Service
parent
2a901d12f4
commit
409df1378f
@@ -16,6 +16,7 @@ import random
|
||||
|
||||
from google.adk.agents.llm_agent import Agent
|
||||
from google.adk.examples.example import Example
|
||||
from google.adk.models.google_llm import Gemini
|
||||
from google.adk.tools.example_tool import ExampleTool
|
||||
from google.genai import types
|
||||
|
||||
@@ -28,6 +29,17 @@ def roll_die(sides: int) -> int:
|
||||
|
||||
roll_agent = Agent(
|
||||
name="roll_agent",
|
||||
model=Gemini(
|
||||
# model="gemini-2.0-flash-live-preview-04-09", # for Vertex project
|
||||
model="gemini-live-2.5-flash-preview", # for AI studio key
|
||||
speech_config=types.SpeechConfig(
|
||||
voice_config=types.VoiceConfig(
|
||||
prebuilt_voice_config=types.PrebuiltVoiceConfig(
|
||||
voice_name="Kore",
|
||||
)
|
||||
)
|
||||
),
|
||||
),
|
||||
description="Handles rolling dice of different sizes.",
|
||||
instruction="""
|
||||
You are responsible for rolling dice based on the user's request.
|
||||
@@ -69,6 +81,17 @@ def check_prime(nums: list[int]) -> str:
|
||||
|
||||
prime_agent = Agent(
|
||||
name="prime_agent",
|
||||
model=Gemini(
|
||||
# model="gemini-2.0-flash-live-preview-04-09", # for Vertex project
|
||||
model="gemini-live-2.5-flash-preview", # for AI studio key
|
||||
speech_config=types.SpeechConfig(
|
||||
voice_config=types.VoiceConfig(
|
||||
prebuilt_voice_config=types.PrebuiltVoiceConfig(
|
||||
voice_name="Puck",
|
||||
)
|
||||
)
|
||||
),
|
||||
),
|
||||
description="Handles checking if numbers are prime.",
|
||||
instruction="""
|
||||
You are responsible for checking whether numbers are prime.
|
||||
@@ -100,8 +123,17 @@ def get_current_weather(location: str):
|
||||
|
||||
root_agent = Agent(
|
||||
# find supported models here: https://google.github.io/adk-docs/get-started/streaming/quickstart-streaming/
|
||||
model="gemini-2.0-flash-live-preview-04-09", # for Vertex project
|
||||
# model="gemini-live-2.5-flash-preview", # for AI studio key
|
||||
model=Gemini(
|
||||
# model="gemini-2.0-flash-live-preview-04-09", # for Vertex project
|
||||
model="gemini-live-2.5-flash-preview", # for AI studio key
|
||||
speech_config=types.SpeechConfig(
|
||||
voice_config=types.VoiceConfig(
|
||||
prebuilt_voice_config=types.PrebuiltVoiceConfig(
|
||||
voice_name="Zephyr",
|
||||
)
|
||||
)
|
||||
),
|
||||
),
|
||||
name="root_agent",
|
||||
instruction="""
|
||||
You are a helpful assistant that can check time, roll dice and check if numbers are prime.
|
||||
|
||||
@@ -35,7 +35,10 @@ class StreamingMode(Enum):
|
||||
|
||||
|
||||
class RunConfig(BaseModel):
|
||||
"""Configs for runtime behavior of agents."""
|
||||
"""Configs for runtime behavior of agents.
|
||||
|
||||
The configs here will be overriden by agent-specific configurations.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(
|
||||
extra='forbid',
|
||||
|
||||
@@ -60,6 +60,8 @@ class Gemini(BaseLlm):
|
||||
|
||||
model: str = 'gemini-2.5-flash'
|
||||
|
||||
speech_config: Optional[types.SpeechConfig] = None
|
||||
|
||||
retry_options: Optional[types.HttpRetryOptions] = None
|
||||
"""Allow Gemini to retry failed responses.
|
||||
|
||||
@@ -269,6 +271,9 @@ class Gemini(BaseLlm):
|
||||
self._live_api_version
|
||||
)
|
||||
|
||||
if self.speech_config is not None:
|
||||
llm_request.live_connect_config.speech_config = self.speech_config
|
||||
|
||||
llm_request.live_connect_config.system_instruction = types.Content(
|
||||
role='system',
|
||||
parts=[
|
||||
|
||||
@@ -1858,3 +1858,189 @@ def test_build_request_log_fallback_to_repr_on_all_failures(monkeypatch):
|
||||
# Should still succeed using repr()
|
||||
assert "Config:" in log_output
|
||||
assert "GenerateContentConfig" in log_output
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connect_uses_gemini_speech_config_when_request_is_none(
|
||||
gemini_llm, llm_request
|
||||
):
|
||||
"""Tests that Gemini's speech_config is used when live_connect_config's is None."""
|
||||
# Arrange: Set a speech_config on the Gemini instance with the voice "Kore"
|
||||
gemini_llm.speech_config = types.SpeechConfig(
|
||||
voice_config=types.VoiceConfig(
|
||||
prebuilt_voice_config=types.PrebuiltVoiceConfig(
|
||||
voice_name="Kore",
|
||||
)
|
||||
)
|
||||
)
|
||||
llm_request.live_connect_config = (
|
||||
types.LiveConnectConfig()
|
||||
) # speech_config is None
|
||||
|
||||
mock_live_session = mock.AsyncMock()
|
||||
|
||||
with mock.patch.object(gemini_llm, "_live_api_client") as mock_live_client:
|
||||
|
||||
class MockLiveConnect:
|
||||
|
||||
async def __aenter__(self):
|
||||
return mock_live_session
|
||||
|
||||
async def __aexit__(self, *args):
|
||||
pass
|
||||
|
||||
mock_live_client.aio.live.connect.return_value = MockLiveConnect()
|
||||
|
||||
# Act
|
||||
async with gemini_llm.connect(llm_request) as connection:
|
||||
# Assert
|
||||
mock_live_client.aio.live.connect.assert_called_once()
|
||||
call_args = mock_live_client.aio.live.connect.call_args
|
||||
config_arg = call_args.kwargs["config"]
|
||||
|
||||
# Verify the speech_config from the Gemini instance was used
|
||||
assert config_arg.speech_config is not None
|
||||
assert (
|
||||
config_arg.speech_config.voice_config.prebuilt_voice_config.voice_name
|
||||
== "Kore"
|
||||
)
|
||||
assert isinstance(connection, GeminiLlmConnection)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connect_uses_request_speech_config_when_gemini_is_none(
|
||||
gemini_llm, llm_request
|
||||
):
|
||||
"""Tests that request's speech_config is used when Gemini's is None."""
|
||||
# Arrange: Set a speech_config on the request instance with the voice "Kore"
|
||||
gemini_llm.speech_config = None
|
||||
request_speech_config = types.SpeechConfig(
|
||||
voice_config=types.VoiceConfig(
|
||||
prebuilt_voice_config=types.PrebuiltVoiceConfig(
|
||||
voice_name="Kore",
|
||||
)
|
||||
)
|
||||
)
|
||||
llm_request.live_connect_config = types.LiveConnectConfig(
|
||||
speech_config=request_speech_config
|
||||
)
|
||||
|
||||
mock_live_session = mock.AsyncMock()
|
||||
|
||||
with mock.patch.object(gemini_llm, "_live_api_client") as mock_live_client:
|
||||
|
||||
class MockLiveConnect:
|
||||
|
||||
async def __aenter__(self):
|
||||
return mock_live_session
|
||||
|
||||
async def __aexit__(self, *args):
|
||||
pass
|
||||
|
||||
mock_live_client.aio.live.connect.return_value = MockLiveConnect()
|
||||
|
||||
# Act
|
||||
async with gemini_llm.connect(llm_request) as connection:
|
||||
# Assert
|
||||
mock_live_client.aio.live.connect.assert_called_once()
|
||||
call_args = mock_live_client.aio.live.connect.call_args
|
||||
config_arg = call_args.kwargs["config"]
|
||||
|
||||
# Verify the speech_config from the request instance was used
|
||||
assert config_arg.speech_config is not None
|
||||
assert (
|
||||
config_arg.speech_config.voice_config.prebuilt_voice_config.voice_name
|
||||
== "Kore"
|
||||
)
|
||||
assert isinstance(connection, GeminiLlmConnection)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connect_request_gemini_config_overrides_speech_config(
|
||||
gemini_llm, llm_request
|
||||
):
|
||||
"""Tests that live_connect_config's speech_config is preserved even if Gemini has one."""
|
||||
# Arrange: Set different speech_configs on both the Gemini instance ("Puck") and the request ("Zephyr")
|
||||
gemini_llm.speech_config = types.SpeechConfig(
|
||||
voice_config=types.VoiceConfig(
|
||||
prebuilt_voice_config=types.PrebuiltVoiceConfig(
|
||||
voice_name="Puck",
|
||||
)
|
||||
)
|
||||
)
|
||||
request_speech_config = types.SpeechConfig(
|
||||
voice_config=types.VoiceConfig(
|
||||
prebuilt_voice_config=types.PrebuiltVoiceConfig(
|
||||
voice_name="Zephyr",
|
||||
)
|
||||
)
|
||||
)
|
||||
llm_request.live_connect_config = types.LiveConnectConfig(
|
||||
speech_config=request_speech_config
|
||||
)
|
||||
|
||||
mock_live_session = mock.AsyncMock()
|
||||
|
||||
with mock.patch.object(gemini_llm, "_live_api_client") as mock_live_client:
|
||||
|
||||
class MockLiveConnect:
|
||||
|
||||
async def __aenter__(self):
|
||||
return mock_live_session
|
||||
|
||||
async def __aexit__(self, *args):
|
||||
pass
|
||||
|
||||
mock_live_client.aio.live.connect.return_value = MockLiveConnect()
|
||||
|
||||
# Act
|
||||
async with gemini_llm.connect(llm_request) as connection:
|
||||
# Assert
|
||||
mock_live_client.aio.live.connect.assert_called_once()
|
||||
call_args = mock_live_client.aio.live.connect.call_args
|
||||
config_arg = call_args.kwargs["config"]
|
||||
|
||||
# Verify the speech_config from the request ("Zephyr") was overwritten by Gemini's speech_config ("Puck")
|
||||
assert config_arg.speech_config is not None
|
||||
assert (
|
||||
config_arg.speech_config.voice_config.prebuilt_voice_config.voice_name
|
||||
== "Puck"
|
||||
)
|
||||
assert isinstance(connection, GeminiLlmConnection)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connect_speech_config_remains_none_when_both_are_none(
|
||||
gemini_llm, llm_request
|
||||
):
|
||||
"""Tests that speech_config is None when neither Gemini nor the request has it."""
|
||||
# Arrange: Ensure both Gemini instance and request have no speech_config
|
||||
gemini_llm.speech_config = None
|
||||
llm_request.live_connect_config = (
|
||||
types.LiveConnectConfig()
|
||||
) # speech_config is None
|
||||
|
||||
mock_live_session = mock.AsyncMock()
|
||||
|
||||
with mock.patch.object(gemini_llm, "_live_api_client") as mock_live_client:
|
||||
|
||||
class MockLiveConnect:
|
||||
|
||||
async def __aenter__(self):
|
||||
return mock_live_session
|
||||
|
||||
async def __aexit__(self, *args):
|
||||
pass
|
||||
|
||||
mock_live_client.aio.live.connect.return_value = MockLiveConnect()
|
||||
|
||||
# Act
|
||||
async with gemini_llm.connect(llm_request) as connection:
|
||||
# Assert
|
||||
mock_live_client.aio.live.connect.assert_called_once()
|
||||
call_args = mock_live_client.aio.live.connect.call_args
|
||||
config_arg = call_args.kwargs["config"]
|
||||
|
||||
# Verify the final speech_config is still None
|
||||
assert config_arg.speech_config is None
|
||||
assert isinstance(connection, GeminiLlmConnection)
|
||||
|
||||
Reference in New Issue
Block a user