Skip to content

Commit d17a2a3

Browse files
committed
fix: fix input and output transcription finished events for Gemini v3.1
Change-Id: I3c33e84569d5f63f46c99154afa9e9d68a2fdf3c
1 parent 85223e6 commit d17a2a3

2 files changed

Lines changed: 30 additions & 8 deletions

File tree

src/google/adk/models/gemini_llm_connection.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -307,10 +307,10 @@ async def receive(self) -> AsyncGenerator[LlmResponse, None]:
307307
live_session_id=live_session_id,
308308
)
309309
self._output_transcription_text = ''
310-
# The Gemini API might not send a transcription finished signal.
310+
# The Gemini API or Vertex AI might not send a transcription finished signal.
311311
# Instead, we rely on generation_complete, turn_complete or
312312
# interrupted signals to flush any pending transcriptions.
313-
if self._api_backend == GoogleLLMVariant.GEMINI_API and (
313+
if (
314314
message.server_content.interrupted
315315
or message.server_content.turn_complete
316316
or message.server_content.generation_complete

tests/unittests/models/test_gemini_llm_connection.py

Lines changed: 28 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -285,11 +285,17 @@ async def mock_receive_generator():
285285

286286

287287
@pytest.mark.asyncio
288+
@pytest.mark.parametrize(
289+
'conn_fixture',
290+
['gemini_api_connection', 'gemini_connection'],
291+
)
288292
async def test_receive_transcript_finished_on_interrupt(
289-
gemini_api_connection,
293+
conn_fixture,
290294
mock_gemini_session,
295+
request,
291296
):
292297
"""Test receive finishes transcription on interrupt signal."""
298+
connection = request.getfixturevalue(conn_fixture)
293299

294300
message1 = mock.Mock()
295301
message1.usage_metadata = None
@@ -345,7 +351,7 @@ async def mock_receive_generator():
345351
receive_mock = mock.Mock(return_value=mock_receive_generator())
346352
mock_gemini_session.receive = receive_mock
347353

348-
responses = [resp async for resp in gemini_api_connection.receive()]
354+
responses = [resp async for resp in connection.receive()]
349355

350356
assert len(responses) == 5
351357
assert responses[4].interrupted is True
@@ -365,11 +371,17 @@ async def mock_receive_generator():
365371

366372

367373
@pytest.mark.asyncio
374+
@pytest.mark.parametrize(
375+
'conn_fixture',
376+
['gemini_api_connection', 'gemini_connection'],
377+
)
368378
async def test_receive_transcript_finished_on_generation_complete(
369-
gemini_api_connection,
379+
conn_fixture,
370380
mock_gemini_session,
381+
request,
371382
):
372383
"""Test receive finishes transcription on generation_complete signal."""
384+
connection = request.getfixturevalue(conn_fixture)
373385

374386
message1 = mock.Mock()
375387
message1.usage_metadata = None
@@ -425,7 +437,7 @@ async def mock_receive_generator():
425437
receive_mock = mock.Mock(return_value=mock_receive_generator())
426438
mock_gemini_session.receive = receive_mock
427439

428-
responses = [resp async for resp in gemini_api_connection.receive()]
440+
responses = [resp async for resp in connection.receive()]
429441

430442
assert len(responses) == 4
431443

@@ -444,11 +456,17 @@ async def mock_receive_generator():
444456

445457

446458
@pytest.mark.asyncio
459+
@pytest.mark.parametrize(
460+
'conn_fixture',
461+
['gemini_api_connection', 'gemini_connection'],
462+
)
447463
async def test_receive_transcript_finished_on_turn_complete(
448-
gemini_api_connection,
464+
conn_fixture,
449465
mock_gemini_session,
466+
request,
450467
):
451468
"""Test receive finishes transcription on interrupt or complete signals."""
469+
connection = request.getfixturevalue(conn_fixture)
452470

453471
message1 = mock.Mock()
454472
message1.usage_metadata = None
@@ -504,7 +522,7 @@ async def mock_receive_generator():
504522
receive_mock = mock.Mock(return_value=mock_receive_generator())
505523
mock_gemini_session.receive = receive_mock
506524

507-
responses = [resp async for resp in gemini_api_connection.receive()]
525+
responses = [resp async for resp in connection.receive()]
508526

509527
assert len(responses) == 5
510528
assert responses[4].turn_complete is True
@@ -867,6 +885,7 @@ async def test_receive_grounding_metadata_standalone(
867885
mock_server_content.interrupted = False
868886
mock_server_content.input_transcription = None
869887
mock_server_content.output_transcription = None
888+
mock_server_content.generation_complete = False
870889

871890
mock_message = mock.create_autospec(types.LiveServerMessage, instance=True)
872891
mock_message.usage_metadata = None
@@ -911,6 +930,7 @@ async def test_receive_grounding_metadata_with_content(
911930
mock_server_content.interrupted = False
912931
mock_server_content.input_transcription = None
913932
mock_server_content.output_transcription = None
933+
mock_server_content.generation_complete = False
914934

915935
mock_message = mock.create_autospec(types.LiveServerMessage, instance=True)
916936
mock_message.usage_metadata = None
@@ -981,6 +1001,7 @@ async def test_receive_tool_call_and_grounding_metadata_with_native_audio(
9811001
mock_server_content.interrupted = False
9821002
mock_server_content.input_transcription = None
9831003
mock_server_content.output_transcription = None
1004+
mock_server_content.generation_complete = False
9841005

9851006
mock_metadata_msg = mock.create_autospec(
9861007
types.LiveServerMessage, instance=True
@@ -1001,6 +1022,7 @@ async def test_receive_tool_call_and_grounding_metadata_with_native_audio(
10011022
mock_turn_complete_content.interrupted = False
10021023
mock_turn_complete_content.input_transcription = None
10031024
mock_turn_complete_content.output_transcription = None
1025+
mock_turn_complete_content.generation_complete = False
10041026

10051027
mock_turn_complete_msg = mock.create_autospec(
10061028
types.LiveServerMessage, instance=True

0 commit comments

Comments
 (0)