diff --git a/src/google/adk/models/gemini_llm_connection.py b/src/google/adk/models/gemini_llm_connection.py index b1a6992b68e..e0908e52c39 100644 --- a/src/google/adk/models/gemini_llm_connection.py +++ b/src/google/adk/models/gemini_llm_connection.py @@ -23,7 +23,7 @@ from google.genai import types from ..utils import model_name_utils -from ..utils.content_utils import filter_audio_parts +from ..utils.content_utils import _filter_media_parts from ..utils.context_utils import Aclosing from ..utils.variant_utils import GoogleLLMVariant from .base_llm_connection import BaseLlmConnection @@ -82,10 +82,12 @@ async def send_history(self, history: list[types.Content]) -> None: # TODO: Remove this filter and translate unary contents to streaming # contents properly. - # Filter out audio parts from history because: + # Filter out audio and internal `_adk_live` artifact parts because: # 1. audio has already been transcribed. # 2. sending audio via connection.send or connection.send_live_content is # not supported by LIVE API (session will be corrupted). + # 3. `_adk_live` artifacts persisted by `save_live_blob` are internal + # `artifact://` references that the model cannot resolve. # This method is called when: # 1. Agent transfer to a new agent # 2. Establishing a new live connection with previous ADK session history @@ -93,7 +95,7 @@ async def send_history(self, history: list[types.Content]) -> None: contents = [ filtered for content in history - if (filtered := filter_audio_parts(content)) is not None + if (filtered := _filter_media_parts(content)) is not None ] if contents: diff --git a/src/google/adk/utils/content_utils.py b/src/google/adk/utils/content_utils.py index d0bada24015..eaf4959083c 100644 --- a/src/google/adk/utils/content_utils.py +++ b/src/google/adk/utils/content_utils.py @@ -50,6 +50,40 @@ def filter_audio_parts(content: types.Content) -> types.Content | None: return types.Content(role=content.role, parts=filtered_parts) +def _is_adk_live_artifact_part(part: types.Part) -> bool: + """Returns whether `part` references an internal `_adk_live` artifact.""" + if part.file_data is None: + return False + file_uri = part.file_data.file_uri or '' + return file_uri.startswith('artifact://') and '/_adk_live/' in file_uri + + +def _filter_media_parts(content: types.Content) -> types.Content | None: + """Filters media parts out of history replayed to a live connection. + + Removes internal `_adk_live` artifact references (audio and other media + persisted when `save_live_blob` is enabled) and any remaining audio parts. + Inline image and video parts (e.g. sent via `send_content`) are kept. + + Args: + content: The content to filter. It is not modified. + + Returns: + A new content with the same role and the remaining parts, or None if no + parts remain. + """ + if not content.parts: + return None + filtered_parts = [ + part + for part in content.parts + if not is_audio_part(part) and not _is_adk_live_artifact_part(part) + ] + if not filtered_parts: + return None + return types.Content(role=content.role, parts=filtered_parts) + + def extract_text_from_content(content: types.Content | None) -> str: """Extracts text from a Content object, filtering out thoughts.""" if not content or not content.parts: diff --git a/tests/unittests/models/test_gemini_llm_connection.py b/tests/unittests/models/test_gemini_llm_connection.py index dc1310485c1..4660c2aab99 100644 --- a/tests/unittests/models/test_gemini_llm_connection.py +++ b/tests/unittests/models/test_gemini_llm_connection.py @@ -139,6 +139,36 @@ async def test_send_history(gemini_connection, mock_gemini_session): assert call_args['turn_complete'] is False # Last message is from model +@pytest.mark.asyncio +async def test_send_history_filters_out_adk_live_artifact_parts( + gemini_connection, mock_gemini_session +): + """Internal `_adk_live` artifact references are not replayed.""" + history = [ + types.Content(role='user', parts=[types.Part.from_text(text='Hello')]), + types.Content( + role='user', + parts=[ + types.Part( + file_data=types.FileData( + file_uri='artifact://app/u/s/_adk_live/media_1.zip#0', + mime_type='application/zip', + ) + ) + ], + ), + types.Content( + role='model', parts=[types.Part.from_text(text='I saw that.')] + ), + ] + + await gemini_connection.send_history(history) + + mock_gemini_session.send_client_content.assert_called_once_with( + turns=[history[0], history[2]], turn_complete=False + ) + + @pytest.mark.asyncio async def test_send_content_text(gemini_connection, mock_gemini_session): """Test send_content with text content.""" diff --git a/tests/unittests/utils/test_content_utils.py b/tests/unittests/utils/test_content_utils.py index f4f705b720e..570e7ef5c81 100644 --- a/tests/unittests/utils/test_content_utils.py +++ b/tests/unittests/utils/test_content_utils.py @@ -14,6 +14,8 @@ from __future__ import annotations +from google.adk.utils.content_utils import _filter_media_parts +from google.adk.utils.content_utils import _is_adk_live_artifact_part from google.adk.utils.content_utils import extract_text_from_content from google.adk.utils.content_utils import filter_audio_parts from google.adk.utils.content_utils import is_audio_part @@ -21,6 +23,7 @@ from google.adk.utils.content_utils import to_user_content from google.genai import types from pydantic import BaseModel +import pytest def test_skip_thought_signature_validator_wire_value(): @@ -93,37 +96,37 @@ def test_to_user_content_list_input_preserves_non_ascii(): assert '\\u' not in text -def _audio_blob_part(mime_type: str) -> types.Part: +def _blob_part(mime_type: str) -> types.Part: return types.Part( inline_data=types.Blob(mime_type=mime_type, data=b'\x00\x01') ) -def _audio_file_part(mime_type: str) -> types.Part: +def _file_part(mime_type: str) -> types.Part: return types.Part( file_data=types.FileData(file_uri='files/clip', mime_type=mime_type) ) def test_is_audio_part_inline_audio_mime_is_audio(): - assert is_audio_part(_audio_blob_part('audio/pcm')) is True + assert is_audio_part(_blob_part('audio/pcm')) is True def test_is_audio_part_file_data_audio_mime_is_audio(): - assert is_audio_part(_audio_file_part('audio/wav')) is True + assert is_audio_part(_file_part('audio/wav')) is True def test_is_audio_part_non_audio_mime_is_not_audio(): # Only the 'audio/' top-level type counts; video and image blobs must # survive so they still reach the model. - assert is_audio_part(_audio_blob_part('image/png')) is False - assert is_audio_part(_audio_file_part('video/mp4')) is False + assert is_audio_part(_blob_part('image/png')) is False + assert is_audio_part(_file_part('video/mp4')) is False def test_is_audio_part_mime_containing_audio_but_not_prefixed_is_not_audio(): # The check is a prefix match on the top-level type, not a substring # match, so 'application/audio-ish' is not audio. - assert is_audio_part(_audio_blob_part('application/audio-ish')) is False + assert is_audio_part(_blob_part('application/audio-ish')) is False def test_is_audio_part_text_part_is_not_audio(): @@ -141,8 +144,8 @@ def test_filter_audio_parts_drops_audio_and_keeps_role_and_order(): role='user', parts=[ types.Part(text='before'), - _audio_blob_part('audio/pcm'), - _audio_file_part('audio/wav'), + _blob_part('audio/pcm'), + _file_part('audio/wav'), types.Part(text='after'), ], ) @@ -157,7 +160,7 @@ def test_filter_audio_parts_drops_audio_and_keeps_role_and_order(): def test_filter_audio_parts_all_audio_returns_none(): # A content whose every part is audio has nothing left to send, so the # caller is told to drop the whole content rather than send an empty one. - content = types.Content(role='user', parts=[_audio_blob_part('audio/pcm')]) + content = types.Content(role='user', parts=[_blob_part('audio/pcm')]) assert filter_audio_parts(content) is None @@ -168,7 +171,7 @@ def test_filter_audio_parts_empty_parts_returns_none(): def test_filter_audio_parts_does_not_mutate_input(): content = types.Content( role='user', - parts=[types.Part(text='keep'), _audio_blob_part('audio/pcm')], + parts=[types.Part(text='keep'), _blob_part('audio/pcm')], ) filter_audio_parts(content) @@ -177,6 +180,151 @@ def test_filter_audio_parts_does_not_mutate_input(): assert content.parts[1].inline_data.mime_type == 'audio/pcm' +def test_is_adk_live_artifact_part_adk_live_references_match_any_type(): + """`artifact://` references under `_adk_live/` match, media or not.""" + audio_ref = types.Part( + file_data=types.FileData( + file_uri='artifact://app/u/s/_adk_live/input_audio_1.pcm#0', + mime_type='audio/pcm', + ) + ) + zip_ref = types.Part( + file_data=types.FileData( + file_uri='artifact://app/u/s/_adk_live/input_media_1.zip#0', + mime_type='application/zip', + ) + ) + + assert _is_adk_live_artifact_part(audio_ref) + assert _is_adk_live_artifact_part(zip_ref) + + +@pytest.mark.parametrize( + 'part', + [ + pytest.param( + types.Part( + file_data=types.FileData( + file_uri='artifact://app/u/s/custom_bundle.zip#0', + mime_type='APPLICATION/ZIP; charset=binary', + ) + ), + id='artifact_outside_adk_live', + ), + pytest.param( + types.Part( + file_data=types.FileData( + file_uri='gs://bucket/_adk_live/frame.jpeg', + mime_type='image/jpeg', + ) + ), + id='non_artifact_uri', + ), + pytest.param(types.Part(file_data=types.FileData()), id='no_uri'), + pytest.param(types.Part(text='hello'), id='text'), + ], +) +def test_is_adk_live_artifact_part_other_parts_do_not_match(part: types.Part): + """Only `artifact://` URIs with an `_adk_live/` path segment are internal.""" + assert not _is_adk_live_artifact_part(part) + + +def test_filter_media_parts_drops_audio_and_keeps_image_and_video(): + """Audio goes; image, video and other parts stay, in order.""" + text_before = types.Part(text='before') + inline_image = _blob_part('image/png') + inline_video = _blob_part('video/mp4') + user_image_ref = types.Part( + file_data=types.FileData( + file_uri='gs://bucket/photo.jpeg', mime_type='image/jpeg' + ) + ) + video_ref = _file_part('video/webm') + text_after = types.Part(text='after') + content = types.Content( + role='user', + parts=[ + text_before, + _blob_part('audio/pcm'), + inline_image, + inline_video, + _file_part('audio/wav'), + user_image_ref, + video_ref, + text_after, + ], + ) + + filtered = _filter_media_parts(content) + + assert filtered is not None + assert filtered.role == 'user' + assert filtered.parts == [ + text_before, + inline_image, + inline_video, + user_image_ref, + video_ref, + text_after, + ] + + +def test_filter_media_parts_drops_adk_live_references_of_any_type(): + """`_adk_live` references go whatever their MIME type, images included.""" + text_part = types.Part(text='Describe this') + adk_live_zip = types.Part( + file_data=types.FileData( + file_uri='artifact://app/u/s/_adk_live/input_media_1.zip#0', + mime_type='application/zip', + ) + ) + adk_live_image = types.Part( + file_data=types.FileData( + file_uri='artifact://app/u/s/_adk_live/input_frame_1.jpeg#0', + mime_type='image/jpeg', + ) + ) + content = types.Content( + role='user', parts=[text_part, adk_live_zip, adk_live_image] + ) + + filtered = _filter_media_parts(content) + + assert filtered is not None + assert filtered.parts == [text_part] + + +def test_filter_media_parts_all_filtered_returns_none(): + """Content with nothing left is dropped rather than sent empty.""" + adk_live_zip = types.Part( + file_data=types.FileData( + file_uri='artifact://app/u/s/_adk_live/input_media_1.zip#0', + mime_type='application/zip', + ) + ) + content = types.Content( + role='user', parts=[_blob_part('audio/pcm'), adk_live_zip] + ) + + assert _filter_media_parts(content) is None + + +def test_filter_media_parts_empty_parts_returns_none(): + assert _filter_media_parts(types.Content(role='user', parts=[])) is None + + +def test_filter_media_parts_does_not_mutate_input(): + content = types.Content( + role='user', + parts=[types.Part(text='keep'), _blob_part('audio/pcm')], + ) + + _filter_media_parts(content) + + assert len(content.parts) == 2 + assert content.parts[1].inline_data.mime_type == 'audio/pcm' + + def test_extract_text_from_content_concatenates_text_parts_verbatim(): # Parts are joined with no separator: the model emits a single logical # string that is chunked arbitrarily across parts. @@ -203,5 +351,5 @@ def test_extract_text_from_content_none_returns_empty_string(): def test_extract_text_from_content_without_text_parts_returns_empty_string(): - content = types.Content(role='user', parts=[_audio_blob_part('audio/pcm')]) + content = types.Content(role='user', parts=[_blob_part('audio/pcm')]) assert extract_text_from_content(content) == ''