fix: Flush pending transcriptions on turn/generation complete or interrupt for Gemini API

The Gemini API may not always send an explicit transcription finished signal. This change ensures that any buffered input or output transcription text is yielded as a finished transcription when a turn is completed, generation is complete, or the session is interrupted.

Also, refined the check for `event.partial` in runners.py to be more explicit.

Co-authored-by: Hangfei Lin <hangfei@google.com>
PiperOrigin-RevId: 839008606
This commit is contained in:
Hangfei Lin
2025-12-01 18:09:29 -08:00
committed by Copybara-Service
parent 98d82935e6
commit 8da61be45a
5 changed files with 359 additions and 7 deletions
@@ -65,8 +65,8 @@ async def check_prime(nums: list[int]) -> str:
root_agent = Agent(
# model='gemini-live-2.5-flash-preview-native-audio-09-2025', # vertex
model='gemini-2.5-flash-native-audio-preview-09-2025', # for AI studio
model='gemini-live-2.5-flash-preview-native-audio-09-2025', # vertex
# model='gemini-2.5-flash-native-audio-preview-09-2025', # for AI studio
# key
name='roll_dice_agent',
description=(
+36 -1
View File
@@ -21,6 +21,7 @@ from typing import Union
from google.genai import types
from ..utils.context_utils import Aclosing
from ..utils.variant_utils import GoogleLLMVariant
from .base_llm_connection import BaseLlmConnection
from .llm_response import LlmResponse
@@ -36,10 +37,15 @@ if TYPE_CHECKING:
class GeminiLlmConnection(BaseLlmConnection):
"""The Gemini model connection."""
def __init__(self, gemini_session: live.AsyncSession):
def __init__(
self,
gemini_session: live.AsyncSession,
api_backend: GoogleLLMVariant = GoogleLLMVariant.VERTEX_AI,
):
self._gemini_session = gemini_session
self._input_transcription_text: str = ''
self._output_transcription_text: str = ''
self._api_backend = api_backend
async def send_history(self, history: list[types.Content]):
"""Sends the conversation history to the gemini model.
@@ -171,6 +177,9 @@ class GeminiLlmConnection(BaseLlmConnection):
yield self.__build_full_text_response(text)
text = ''
yield llm_response
# Note: in some cases, tool_call may arrive before
# generation_complete, causing transcription to appear after
# tool_call in the session log.
if message.server_content.input_transcription:
if message.server_content.input_transcription.text:
self._input_transcription_text += (
@@ -215,6 +224,32 @@ class GeminiLlmConnection(BaseLlmConnection):
partial=False,
)
self._output_transcription_text = ''
# The Gemini API might not send a transcription finished signal.
# Instead, we rely on generation_complete, turn_complete or
# interrupted signals to flush any pending transcriptions.
if self._api_backend == GoogleLLMVariant.GEMINI_API and (
message.server_content.interrupted
or message.server_content.turn_complete
or message.server_content.generation_complete
):
if self._input_transcription_text:
yield LlmResponse(
input_transcription=types.Transcription(
text=self._input_transcription_text,
finished=True,
),
partial=False,
)
self._input_transcription_text = ''
if self._output_transcription_text:
yield LlmResponse(
output_transcription=types.Transcription(
text=self._output_transcription_text,
finished=True,
),
partial=False,
)
self._output_transcription_text = ''
if message.server_content.turn_complete:
if text:
yield self.__build_full_text_response(text)
+1 -1
View File
@@ -342,7 +342,7 @@ class Gemini(BaseLlm):
async with self._live_api_client.aio.live.connect(
model=llm_request.model, config=llm_request.live_connect_config
) as live_session:
yield GeminiLlmConnection(live_session)
yield GeminiLlmConnection(live_session, api_backend=self._api_backend)
async def _adapt_computer_use_tool(self, llm_request: LlmRequest) -> None:
"""Adapt the google computer use predefined functions to the adk computer use toolset."""
+81 -2
View File
@@ -67,6 +67,23 @@ from .utils.context_utils import Aclosing
logger = logging.getLogger('google_adk.' + __name__)
def _is_tool_call_or_response(event: Event) -> bool:
return bool(event.get_function_calls() or event.get_function_responses())
def _is_transcription(event: Event) -> bool:
return (
event.input_transcription is not None
or event.output_transcription is not None
)
def _has_non_empty_transcription_text(transcription) -> bool:
return bool(
transcription and transcription.text and transcription.text.strip()
)
class Runner:
"""The Runner class is used to run agents.
@@ -626,6 +643,7 @@ class Runner:
invocation_context: The invocation context
session: The current session
execute_fn: A callable that returns an AsyncGenerator of Events
is_live_call: Whether this is a live call
Yields:
Events from the execution, including any generated by plugins
@@ -651,13 +669,74 @@ class Runner:
yield early_exit_event
else:
# Step 2: Otherwise continue with normal execution
# Note for live/bidi:
# the transcription may arrive later then the action(function call
# event and thus function response event). In this case, the order of
# transcription and function call event will be wrong if we just
# append as it arrives. To address this, we should check if there is
# transcription going on. If there is transcription going on, we
# should hold on appending the function call event until the
# transcription is finished. The transcription in progress can be
# identified by checking if the transcription event is partial. When
# the next transcription event is not partial, it means the previous
# transcription is finished. Then if there is any buffered function
# call event, we should append them after this finished(non-parital)
# transcription event.
buffered_events: list[Event] = []
is_transcribing: bool = False
async with Aclosing(execute_fn(invocation_context)) as agen:
async for event in agen:
if not event.partial:
if self._should_append_event(event, is_live_call):
if is_live_call:
if event.partial and _is_transcription(event):
is_transcribing = True
if is_transcribing and _is_tool_call_or_response(event):
# only buffer function call and function response event which is
# non-partial
buffered_events.append(event)
continue
# Note for live/bidi: for audio response, it's considered as
# non-paritla event(event.partial=None)
# event.partial=False and event.partial=None are considered as
# non-partial event; event.partial=True is considered as partial
# event.
if event.partial is not True:
if _is_transcription(event) and (
_has_non_empty_transcription_text(event.input_transcription)
or _has_non_empty_transcription_text(
event.output_transcription
)
):
# transcription end signal, append buffered events
is_transcribing = False
logger.debug(
'Appending transcription finished event: %s', event
)
if self._should_append_event(event, is_live_call):
await self.session_service.append_event(
session=session, event=event
)
for buffered_event in buffered_events:
logger.debug('Appending buffered event: %s', buffered_event)
await self.session_service.append_event(
session=session, event=buffered_event
)
buffered_events = []
else:
# non-transcription event or empty transcription event, for
# example, event that stores blob reference, should be appended.
if self._should_append_event(event, is_live_call):
logger.debug('Appending non-buffered event: %s', event)
await self.session_service.append_event(
session=session, event=event
)
else:
if event.partial is not True:
await self.session_service.append_event(
session=session, event=event
)
# Step 3: Run the on_event callbacks to optionally modify the event.
modified_event = await plugin_manager.run_on_event_callback(
invocation_context=invocation_context, event=event
@@ -15,6 +15,7 @@
from unittest import mock
from google.adk.models.gemini_llm_connection import GeminiLlmConnection
from google.adk.utils.variant_utils import GoogleLLMVariant
from google.genai import types
import pytest
@@ -28,7 +29,17 @@ def mock_gemini_session():
@pytest.fixture
def gemini_connection(mock_gemini_session):
"""GeminiLlmConnection instance with mocked session."""
return GeminiLlmConnection(mock_gemini_session)
return GeminiLlmConnection(
mock_gemini_session, api_backend=GoogleLLMVariant.VERTEX_AI
)
@pytest.fixture
def gemini_api_connection(mock_gemini_session):
"""GeminiLlmConnection instance with mocked session for Gemini API."""
return GeminiLlmConnection(
mock_gemini_session, api_backend=GoogleLLMVariant.GEMINI_API
)
@pytest.fixture
@@ -225,6 +236,227 @@ async def test_receive_usage_metadata_and_server_content(
assert content_response.content == mock_content
@pytest.mark.asyncio
async def test_receive_transcript_finished_on_interrupt(
gemini_api_connection,
mock_gemini_session,
):
"""Test receive finishes transcription on interrupt signal."""
message1 = mock.Mock()
message1.usage_metadata = None
message1.server_content = mock.Mock()
message1.server_content.model_turn = None
message1.server_content.interrupted = False
message1.server_content.input_transcription = types.Transcription(
text='Hello', finished=False
)
message1.server_content.output_transcription = None
message1.server_content.turn_complete = False
message1.server_content.generation_complete = False
message1.tool_call = None
message1.session_resumption_update = None
message2 = mock.Mock()
message2.usage_metadata = None
message2.server_content = mock.Mock()
message2.server_content.model_turn = None
message2.server_content.interrupted = False
message2.server_content.input_transcription = None
message2.server_content.output_transcription = types.Transcription(
text='How can', finished=False
)
message2.server_content.turn_complete = False
message2.server_content.generation_complete = False
message2.tool_call = None
message2.session_resumption_update = None
message3 = mock.Mock()
message3.usage_metadata = None
message3.server_content = mock.Mock()
message3.server_content.model_turn = None
message3.server_content.interrupted = True
message3.server_content.input_transcription = None
message3.server_content.output_transcription = None
message3.server_content.turn_complete = False
message3.server_content.generation_complete = False
message3.tool_call = None
message3.session_resumption_update = None
async def mock_receive_generator():
yield message1
yield message2
yield message3
receive_mock = mock.Mock(return_value=mock_receive_generator())
mock_gemini_session.receive = receive_mock
responses = [resp async for resp in gemini_api_connection.receive()]
assert len(responses) == 5
assert responses[4].interrupted is True
assert responses[0].input_transcription.text == 'Hello'
assert responses[0].input_transcription.finished is False
assert responses[0].partial is True
assert responses[1].output_transcription.text == 'How can'
assert responses[1].output_transcription.finished is False
assert responses[1].partial is True
assert responses[2].input_transcription.text == 'Hello'
assert responses[2].input_transcription.finished is True
assert responses[2].partial is False
assert responses[3].output_transcription.text == 'How can'
assert responses[3].output_transcription.finished is True
assert responses[3].partial is False
@pytest.mark.asyncio
async def test_receive_transcript_finished_on_generation_complete(
gemini_api_connection,
mock_gemini_session,
):
"""Test receive finishes transcription on generation_complete signal."""
message1 = mock.Mock()
message1.usage_metadata = None
message1.server_content = mock.Mock()
message1.server_content.model_turn = None
message1.server_content.interrupted = False
message1.server_content.input_transcription = types.Transcription(
text='Hello', finished=False
)
message1.server_content.output_transcription = None
message1.server_content.turn_complete = False
message1.server_content.generation_complete = False
message1.tool_call = None
message1.session_resumption_update = None
message2 = mock.Mock()
message2.usage_metadata = None
message2.server_content = mock.Mock()
message2.server_content.model_turn = None
message2.server_content.interrupted = False
message2.server_content.input_transcription = None
message2.server_content.output_transcription = types.Transcription(
text='How can', finished=False
)
message2.server_content.turn_complete = False
message2.server_content.generation_complete = False
message2.tool_call = None
message2.session_resumption_update = None
message3 = mock.Mock()
message3.usage_metadata = None
message3.server_content = mock.Mock()
message3.server_content.model_turn = None
message3.server_content.interrupted = False
message3.server_content.input_transcription = None
message3.server_content.output_transcription = None
message3.server_content.turn_complete = False
message3.server_content.generation_complete = True
message3.tool_call = None
message3.session_resumption_update = None
async def mock_receive_generator():
yield message1
yield message2
yield message3
receive_mock = mock.Mock(return_value=mock_receive_generator())
mock_gemini_session.receive = receive_mock
responses = [resp async for resp in gemini_api_connection.receive()]
assert len(responses) == 4
assert responses[0].input_transcription.text == 'Hello'
assert responses[0].input_transcription.finished is False
assert responses[0].partial is True
assert responses[1].output_transcription.text == 'How can'
assert responses[1].output_transcription.finished is False
assert responses[1].partial is True
assert responses[2].input_transcription.text == 'Hello'
assert responses[2].input_transcription.finished is True
assert responses[2].partial is False
assert responses[3].output_transcription.text == 'How can'
assert responses[3].output_transcription.finished is True
assert responses[3].partial is False
@pytest.mark.asyncio
async def test_receive_transcript_finished_on_turn_complete(
gemini_api_connection,
mock_gemini_session,
):
"""Test receive finishes transcription on interrupt or complete signals."""
message1 = mock.Mock()
message1.usage_metadata = None
message1.server_content = mock.Mock()
message1.server_content.model_turn = None
message1.server_content.interrupted = False
message1.server_content.input_transcription = types.Transcription(
text='Hello', finished=False
)
message1.server_content.output_transcription = None
message1.server_content.turn_complete = False
message1.server_content.generation_complete = False
message1.tool_call = None
message1.session_resumption_update = None
message2 = mock.Mock()
message2.usage_metadata = None
message2.server_content = mock.Mock()
message2.server_content.model_turn = None
message2.server_content.interrupted = False
message2.server_content.input_transcription = None
message2.server_content.output_transcription = types.Transcription(
text='How can', finished=False
)
message2.server_content.turn_complete = False
message2.server_content.generation_complete = False
message2.tool_call = None
message2.session_resumption_update = None
message3 = mock.Mock()
message3.usage_metadata = None
message3.server_content = mock.Mock()
message3.server_content.model_turn = None
message3.server_content.interrupted = False
message3.server_content.input_transcription = None
message3.server_content.output_transcription = None
message3.server_content.turn_complete = True
message3.server_content.generation_complete = False
message3.tool_call = None
message3.session_resumption_update = None
async def mock_receive_generator():
yield message1
yield message2
yield message3
receive_mock = mock.Mock(return_value=mock_receive_generator())
mock_gemini_session.receive = receive_mock
responses = [resp async for resp in gemini_api_connection.receive()]
assert len(responses) == 5
assert responses[4].turn_complete is True
assert responses[0].input_transcription.text == 'Hello'
assert responses[0].input_transcription.finished is False
assert responses[0].partial is True
assert responses[1].output_transcription.text == 'How can'
assert responses[1].output_transcription.finished is False
assert responses[1].partial is True
assert responses[2].input_transcription.text == 'Hello'
assert responses[2].input_transcription.finished is True
assert responses[2].partial is False
assert responses[3].output_transcription.text == 'How can'
assert responses[3].output_transcription.finished is True
assert responses[3].partial is False
@pytest.mark.asyncio
async def test_receive_handles_input_transcription_fragments(
gemini_connection, mock_gemini_session
@@ -240,6 +472,7 @@ async def test_receive_handles_input_transcription_fragments(
)
message1.server_content.output_transcription = None
message1.server_content.turn_complete = False
message1.server_content.generation_complete = False
message1.tool_call = None
message1.session_resumption_update = None
@@ -253,6 +486,7 @@ async def test_receive_handles_input_transcription_fragments(
)
message2.server_content.output_transcription = None
message2.server_content.turn_complete = False
message2.server_content.generation_complete = False
message2.tool_call = None
message2.session_resumption_update = None
@@ -266,6 +500,7 @@ async def test_receive_handles_input_transcription_fragments(
)
message3.server_content.output_transcription = None
message3.server_content.turn_complete = False
message3.server_content.generation_complete = False
message3.tool_call = None
message3.session_resumption_update = None
@@ -306,6 +541,7 @@ async def test_receive_handles_output_transcription_fragments(
text='How can', finished=False
)
message1.server_content.turn_complete = False
message1.server_content.generation_complete = False
message1.tool_call = None
message1.session_resumption_update = None
@@ -319,6 +555,7 @@ async def test_receive_handles_output_transcription_fragments(
text=' I help?', finished=False
)
message2.server_content.turn_complete = False
message2.server_content.generation_complete = False
message2.tool_call = None
message2.session_resumption_update = None
@@ -332,6 +569,7 @@ async def test_receive_handles_output_transcription_fragments(
text=None, finished=True
)
message3.server_content.turn_complete = False
message3.server_content.generation_complete = False
message3.tool_call = None
message3.session_resumption_update = None