Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 10 additions & 5 deletions src/google/adk/live/_runner_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
import asyncio
from contextlib import aclosing
import logging
from typing import Any
from typing import AsyncGenerator
from typing import Optional
from typing import TYPE_CHECKING
Expand Down Expand Up @@ -52,16 +53,20 @@ def new_invocation_context_for_live(
run_config = run_config or RunConfig()

# For live multi-agents system, we need model's text transcription as
# context for the transferred agent.
# context for the transferred agent. Only fill the defaults when the caller
# did not set the fields explicitly, so an explicit None stays disabled.
if hasattr(runner.agent, "sub_agents") and runner.agent.sub_agents:
updates: dict[str, Any] = {}
if (
run_config.response_modalities
and types.Modality.AUDIO in run_config.response_modalities
and "output_audio_transcription" not in run_config.model_fields_set
):
if not run_config.output_audio_transcription:
run_config.output_audio_transcription = types.AudioTranscriptionConfig()
if not run_config.input_audio_transcription:
run_config.input_audio_transcription = types.AudioTranscriptionConfig()
updates["output_audio_transcription"] = types.AudioTranscriptionConfig()
if "input_audio_transcription" not in run_config.model_fields_set:
updates["input_audio_transcription"] = types.AudioTranscriptionConfig()
if updates:
run_config = run_config.model_copy(update=updates)
return runner._new_invocation_context( # pylint: disable=protected-access
session,
live_request_queue=live_request_queue,
Expand Down
81 changes: 81 additions & 0 deletions tests/unittests/live/test__runner_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,6 +85,87 @@ async def test_new_invocation_context_for_live_subagents_audio_transcription():
assert ic.run_config.input_audio_transcription is not None


def _multi_agent_runner() -> Runner:
parent_agent = _MockLiveAgent(name="parent")
parent_agent.sub_agents = [_MockLiveAgent(name="child")]
return Runner(
app_name="test_app",
agent=parent_agent,
session_service=InMemorySessionService(),
)


@pytest.mark.asyncio
async def test_new_invocation_context_for_live_respects_explicit_opt_out():
runner = _multi_agent_runner()
session = await runner.session_service.create_session(
user_id="u1", session_id="s1", app_name=runner.app_name
)
run_config = RunConfig(
response_modalities=[types.Modality.AUDIO],
output_audio_transcription=None,
input_audio_transcription=None,
)

ic = _runner_utils.new_invocation_context_for_live(
runner,
session,
live_request_queue=LiveRequestQueue(),
run_config=run_config,
)

assert ic.run_config.output_audio_transcription is None
assert ic.run_config.input_audio_transcription is None
# The caller's config object must not be mutated either.
assert run_config.output_audio_transcription is None
assert run_config.input_audio_transcription is None


@pytest.mark.asyncio
async def test_new_invocation_context_for_live_opt_out_without_modalities():
runner = _multi_agent_runner()
session = await runner.session_service.create_session(
user_id="u1", session_id="s1", app_name=runner.app_name
)
run_config = RunConfig(input_audio_transcription=None)

ic = _runner_utils.new_invocation_context_for_live(
runner,
session,
live_request_queue=LiveRequestQueue(),
run_config=run_config,
)

assert ic.run_config.input_audio_transcription is None


@pytest.mark.asyncio
async def test_new_invocation_context_for_live_without_subagents_passthrough():
runner = Runner(
app_name="test_app",
agent=_MockLiveAgent(name="solo"),
session_service=InMemorySessionService(),
)
session = await runner.session_service.create_session(
user_id="u1", session_id="s1", app_name=runner.app_name
)
run_config = RunConfig(
response_modalities=[types.Modality.AUDIO],
output_audio_transcription=None,
input_audio_transcription=None,
)

ic = _runner_utils.new_invocation_context_for_live(
runner,
session,
live_request_queue=LiveRequestQueue(),
run_config=run_config,
)

assert ic.run_config.output_audio_transcription is None
assert ic.run_config.input_audio_transcription is None


@pytest.mark.asyncio
async def test_run_live_validates_required_arguments():
agent = _MockLiveAgent()
Expand Down