diff --git a/src/google/adk/agents/invocation_context.py b/src/google/adk/agents/invocation_context.py index 7ec2818799..bfb40c72ee 100644 --- a/src/google/adk/agents/invocation_context.py +++ b/src/google/adk/agents/invocation_context.py @@ -33,7 +33,7 @@ from ..events._branch_path import _BranchPath from ..events.event import Event from ..live._active_streaming_tool import ActiveStreamingTool -from ..live._audio_cache_manager import RealtimeCacheEntry as RealtimeCacheEntry +from ..live._cache_manager import RealtimeCacheEntry from ..live._transcription_entry import TranscriptionEntry from ..live.live_request_queue import LiveRequestQueue from ..memory.base_memory_service import BaseMemoryService @@ -235,6 +235,12 @@ class InvocationContext(BaseModel): output_realtime_cache: list[RealtimeCacheEntry] | None = None """Caches output audio chunks before flushing to session and artifact services.""" + input_media_realtime_cache: list[RealtimeCacheEntry] | None = None + """Caches input media (video/image) frames before flushing to session and artifact services.""" + + output_media_realtime_cache: list[RealtimeCacheEntry] | None = None + """Caches output media (video/image) frames before flushing to session and artifact services.""" + run_config: RunConfig | None = None """Configurations for live agents under this invocation.""" diff --git a/src/google/adk/flows/llm_flows/audio_cache_manager.py b/src/google/adk/flows/llm_flows/audio_cache_manager.py index 23761b6a4d..52101ce6eb 100644 --- a/src/google/adk/flows/llm_flows/audio_cache_manager.py +++ b/src/google/adk/flows/llm_flows/audio_cache_manager.py @@ -15,7 +15,7 @@ """Backward compatibility module for AudioCacheManager. AudioCacheManager and AudioCacheConfig are no longer public; they live in -``google.adk.live._audio_cache_manager`` and this module only keeps +``google.adk.live._cache_manager`` and this module only keeps existing imports working. RealtimeCacheEntry is public as ``google.adk.agents.invocation_context.RealtimeCacheEntry``. """ @@ -24,14 +24,25 @@ import warnings -from ...live._audio_cache_manager import AudioCacheConfig as AudioCacheConfig -from ...live._audio_cache_manager import AudioCacheManager as AudioCacheManager -from ...live._audio_cache_manager import logger as logger -from ...live._audio_cache_manager import RealtimeCacheEntry as RealtimeCacheEntry +from ...live._cache_manager import AudioCacheConfig +from ...live._cache_manager import AudioCacheManager +from ...live._cache_manager import CacheConfig +from ...live._cache_manager import CacheManager +from ...live._cache_manager import logger +from ...live._cache_manager import RealtimeCacheEntry warnings.warn( 'google.adk.flows.llm_flows.audio_cache_manager is deprecated; use' - ' google.adk.live._audio_cache_manager instead.', + ' google.adk.live._cache_manager instead.', DeprecationWarning, stacklevel=2, ) + +__all__ = [ + 'AudioCacheConfig', + 'AudioCacheManager', + 'CacheConfig', + 'CacheManager', + 'RealtimeCacheEntry', + 'logger', +] diff --git a/src/google/adk/flows/llm_flows/base_llm_flow.py b/src/google/adk/flows/llm_flows/base_llm_flow.py index d4513e2272..65c7ea5d57 100644 --- a/src/google/adk/flows/llm_flows/base_llm_flow.py +++ b/src/google/adk/flows/llm_flows/base_llm_flow.py @@ -17,10 +17,12 @@ from abc import ABC from collections.abc import Iterator import logging +from typing import Any from typing import AsyncGenerator from typing import cast from typing import Optional from typing import TYPE_CHECKING +import warnings from google.adk.platform import time as platform_time from google.genai import types @@ -31,7 +33,7 @@ from ...agents.invocation_context import InvocationContext from ...events.event import Event from ...live import _live_llm_flow -from ...live._audio_cache_manager import AudioCacheManager +from ...live._cache_manager import CacheManager from ...live._flow_utils import DEFAULT_ENABLE_CACHE_STATISTICS as DEFAULT_ENABLE_CACHE_STATISTICS from ...live._flow_utils import DEFAULT_MAX_RECONNECT_ATTEMPTS as DEFAULT_MAX_RECONNECT_ATTEMPTS from ...live._flow_utils import DEFAULT_TASK_COMPLETION_DELAY as DEFAULT_TASK_COMPLETION_DELAY @@ -126,7 +128,34 @@ def __init__(self) -> None: self.response_processors: list[BaseLlmResponseProcessor] = [] # Initialize configuration and managers - self.audio_cache_manager = AudioCacheManager() + self.cache_manager = CacheManager() + self.audio_cache_manager = self.cache_manager + + def __getattribute__(self, name: str) -> Any: + if name == 'audio_cache_manager': + warnings.warn( + 'audio_cache_manager is deprecated; use cache_manager instead.', + DeprecationWarning, + stacklevel=2, + ) + return super().__getattribute__('cache_manager') + return super().__getattribute__(name) + + def __setattr__(self, name: str, value: Any) -> None: + if name == 'audio_cache_manager': + if 'audio_cache_manager' in super().__getattribute__('__dict__'): + warnings.warn( + 'audio_cache_manager is deprecated; use cache_manager instead.', + DeprecationWarning, + stacklevel=2, + ) + super().__setattr__('cache_manager', value) + elif ( + name == 'cache_manager' + and 'audio_cache_manager' in super().__getattribute__('__dict__') + ): + super().__setattr__('audio_cache_manager', value) + super().__setattr__(name, value) def _request_processor_lists( self, diff --git a/src/google/adk/live/_audio_cache_manager.py b/src/google/adk/live/_audio_cache_manager.py index 83859605e9..3215a424e8 100644 --- a/src/google/adk/live/_audio_cache_manager.py +++ b/src/google/adk/live/_audio_cache_manager.py @@ -12,299 +12,23 @@ # See the License for the specific language governing permissions and # limitations under the License. -from __future__ import annotations - -import logging -from typing import TYPE_CHECKING - -from google.adk.platform import time as platform_time -from google.genai import types -from pydantic import BaseModel -from pydantic import ConfigDict - -from ..events.event import Event - -if TYPE_CHECKING: - from ..agents.invocation_context import InvocationContext - -logger = logging.getLogger('google_adk.' + __name__) - - -class RealtimeCacheEntry(BaseModel): - """Store audio data chunks for caching before flushing.""" - - model_config = ConfigDict( - arbitrary_types_allowed=True, - extra='forbid', - ) - """The pydantic model config.""" - - role: str - """The role that created this audio data, typically "user" or "model".""" - - data: types.Blob - """The audio data chunk.""" - - timestamp: float - """Timestamp when the audio chunk was received.""" - - -def _require_audio_data(blob: types.Blob) -> bytes: - data = blob.data - if not isinstance(data, bytes): - raise ValueError('Audio blobs must contain byte data.') - return data - - -# Deliberately duplicated from `flows.llm_flows.core._utils` -# rather than imported: `live` sits below `flows.llm_flows` in the -# layering, and importing upward would put a cycle back in. -def _require_agent_name(invocation_context: InvocationContext) -> str: - agent = invocation_context.agent - if agent is None: - raise TypeError('Live audio requires an agent in InvocationContext.') - return agent.name - - -class AudioCacheManager: - """Manages audio caching and flushing for live streaming flows.""" - - def __init__(self, config: AudioCacheConfig | None = None) -> None: - """Initialize the audio cache manager. - - Args: - config: Configuration for audio caching behavior. - """ - self.config = config or AudioCacheConfig() - - def cache_audio( - self, - invocation_context: InvocationContext, - audio_blob: types.Blob, - cache_type: str, - ) -> None: - """Cache incoming user or outgoing model audio data. - - Args: - invocation_context: The current invocation context. - audio_blob: The audio data to cache. - cache_type: Type of audio to cache, either 'input' or 'output'. - - Raises: - ValueError: If cache_type is not 'input' or 'output'. - """ - audio_data = _require_audio_data(audio_blob) - if cache_type == 'input': - if not invocation_context.input_realtime_cache: - invocation_context.input_realtime_cache = [] - cache = invocation_context.input_realtime_cache - role = 'user' - elif cache_type == 'output': - if not invocation_context.output_realtime_cache: - invocation_context.output_realtime_cache = [] - cache = invocation_context.output_realtime_cache - role = 'model' - else: - raise ValueError("cache_type must be either 'input' or 'output'") - - audio_entry = RealtimeCacheEntry( - role=role, data=audio_blob, timestamp=platform_time.get_time() - ) - cache.append(audio_entry) - - logger.debug( - 'Cached %s audio chunk: %d bytes, cache size: %d', - cache_type, - len(audio_data), - len(cache), - ) - - async def flush_caches( - self, - invocation_context: InvocationContext, - flush_user_audio: bool = True, - flush_model_audio: bool = True, - ) -> list[Event]: - """Flush audio caches to artifact services. +"""Backward compatibility module for the audio-only cache manager names. - The multimodality data is saved in artifact service in the format of - audio file. The file data reference is added to the session as an event. - The audio file follows the naming convention: artifact_ref = - f"artifact://{invocation_context.app_name}/{invocation_context.user_id}/ - {invocation_context.session.id}/_adk_live/{filename}#{revision_id}" +The cache manager now lives in ``_cache_manager`` under generalized names. +This module keeps the former names importable; it holds no implementation of +its own. +""" - Note: video data is not supported yet. - - Args: - invocation_context: The invocation context containing audio caches. - flush_user_audio: Whether to flush the input (user) audio cache. - flush_model_audio: Whether to flush the output (model) audio cache. - - Returns: - A list of Event objects created from the flushed caches. - """ - flushed_events: list[Event] = [] - if flush_user_audio and invocation_context.input_realtime_cache: - audio_event = await self._flush_cache_to_services( - invocation_context, - invocation_context.input_realtime_cache, - 'input_audio', - ) - if audio_event: - flushed_events.append(audio_event) - invocation_context.input_realtime_cache = [] - - if flush_model_audio and invocation_context.output_realtime_cache: - logger.debug('Flushed output audio cache') - audio_event = await self._flush_cache_to_services( - invocation_context, - invocation_context.output_realtime_cache, - 'output_audio', - ) - if audio_event: - flushed_events.append(audio_event) - invocation_context.output_realtime_cache = [] - - return flushed_events - - async def _flush_cache_to_services( - self, - invocation_context: InvocationContext, - audio_cache: list[RealtimeCacheEntry], - cache_type: str, - ) -> Event | None: - """Flush a list of audio cache entries to artifact services. - - The artifact service stores the actual blob. The session stores the - reference to the stored blob. - - Args: - invocation_context: The invocation context. - audio_cache: The audio cache to flush. - cache_type: Type identifier for the cache ('input_audio' or - 'output_audio'). - - Returns: - The created Event if the cache was successfully flushed, None otherwise. - """ - if not invocation_context.artifact_service or not audio_cache: - logger.debug('Skipping cache flush: no artifact service or empty cache') - return None - - try: - # Combine audio chunks into a single file. Use join rather than repeated - # `+=`, which is O(n^2) over the total audio size. - mime_type = audio_cache[0].data.mime_type or 'audio/pcm' - combined_audio_data = b''.join( - entry.data.data or b'' for entry in audio_cache - ) - - # Generate filename with timestamp from first audio chunk (when recording started) - timestamp = int(audio_cache[0].timestamp * 1000) # milliseconds - filename = f"adk_live_audio_storage_{cache_type}_{timestamp}.{mime_type.split('/')[-1]}" - - # Save to artifact service - combined_audio_part = types.Part( - inline_data=types.Blob(data=combined_audio_data, mime_type=mime_type) - ) - - revision_id = await invocation_context.artifact_service.save_artifact( - app_name=invocation_context.app_name, - user_id=invocation_context.user_id, - session_id=invocation_context.session.id, - filename=filename, - artifact=combined_audio_part, - ) - - # Create artifact reference for session service - artifact_ref = f'artifact://{invocation_context.app_name}/{invocation_context.user_id}/{invocation_context.session.id}/_adk_live/{filename}#{revision_id}' - - # Create event with file data reference to add to session - # For model events, author should be the agent name, not the role - author = ( - _require_agent_name(invocation_context) - if audio_cache[0].role == 'model' - else audio_cache[0].role - ) - audio_event = Event( - id=Event.new_id(), - invocation_id=invocation_context.invocation_id, - author=author, - content=types.Content( - role=audio_cache[0].role, - parts=[ - types.Part( - file_data=types.FileData( - file_uri=artifact_ref, mime_type=mime_type - ) - ) - ], - ), - timestamp=audio_cache[0].timestamp, - ) - - logger.debug( - 'Successfully flushed %s cache: %d chunks, %d bytes, saved as %s', - cache_type, - len(audio_cache), - len(combined_audio_data), - filename, - ) - return audio_event - - except Exception as e: - logger.error('Failed to flush %s cache: %s', cache_type, e) - return None - - def get_cache_stats( - self, invocation_context: InvocationContext - ) -> dict[str, int]: - """Get statistics about current cache state. - - Args: - invocation_context: The invocation context. - - Returns: - Dictionary containing cache statistics. - """ - input_count = len(invocation_context.input_realtime_cache or []) - output_count = len(invocation_context.output_realtime_cache or []) - - input_bytes = sum( - len(_require_audio_data(entry.data)) - for entry in invocation_context.input_realtime_cache or [] - ) - output_bytes = sum( - len(_require_audio_data(entry.data)) - for entry in invocation_context.output_realtime_cache or [] - ) - - return { - 'input_chunks': input_count, - 'output_chunks': output_count, - 'input_bytes': input_bytes, - 'output_bytes': output_bytes, - 'total_chunks': input_count + output_count, - 'total_bytes': input_bytes + output_bytes, - } - - -class AudioCacheConfig: - """Configuration for audio caching behavior.""" - - def __init__( - self, - max_cache_size_bytes: int = 10 * 1024 * 1024, # 10MB - max_cache_duration_seconds: float = 300.0, # 5 minutes - auto_flush_threshold: int = 100, # Number of chunks - ) -> None: - """Initialize audio cache configuration. +from __future__ import annotations - Args: - max_cache_size_bytes: Maximum cache size in bytes before auto-flush. - max_cache_duration_seconds: Maximum duration to keep data in cache. - auto_flush_threshold: Number of chunks that triggers auto-flush. - """ - self.max_cache_size_bytes = max_cache_size_bytes - self.max_cache_duration_seconds = max_cache_duration_seconds - self.auto_flush_threshold = auto_flush_threshold +from ._cache_manager import AudioCacheConfig +from ._cache_manager import AudioCacheManager +from ._cache_manager import logger +from ._cache_manager import RealtimeCacheEntry + +__all__ = [ + 'AudioCacheConfig', + 'AudioCacheManager', + 'RealtimeCacheEntry', + 'logger', +] diff --git a/src/google/adk/live/_cache_manager.py b/src/google/adk/live/_cache_manager.py new file mode 100644 index 0000000000..a896cf7544 --- /dev/null +++ b/src/google/adk/live/_cache_manager.py @@ -0,0 +1,306 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Multimodal cache manager for live streaming flows.""" + +from __future__ import annotations + +import asyncio +import logging +from typing import TYPE_CHECKING + +from google.adk.platform import time as platform_time +from google.genai import types +from pydantic import BaseModel +from pydantic import ConfigDict + +from ..events.event import Event + +if TYPE_CHECKING: + from ..agents.invocation_context import InvocationContext # pylint: disable=g-import-not-at-top + +logger = logging.getLogger('google_adk.' + __name__) + + +class RealtimeCacheEntry(BaseModel): + """Store raw realtime chunk or frame data for caching before flushing.""" + + model_config = ConfigDict(arbitrary_types_allowed=True, extra='forbid') + role: str + data: types.Blob + timestamp: float + + +def _require_blob_data(blob: types.Blob) -> bytes: + data = blob.data + if not isinstance(data, bytes): + raise ValueError('Blobs must contain byte data.') + return data + + +_require_audio_data = _require_blob_data + + +def _require_agent_name(invocation_context: InvocationContext) -> str: + agent = invocation_context.agent + if agent is None: + raise TypeError('Live streaming requires an agent in InvocationContext.') + return agent.name + + +def _normalize_mime_type( + mime_type: str | None, default: str = 'application/octet-stream' +) -> str: + """Returns the lowercase base MIME type with parameters stripped.""" + base = (mime_type or '').split(';', 1)[0].strip().lower() + return base or default + + +def _get_cache_list( + invocation_context: InvocationContext, attr_name: str +) -> list[RealtimeCacheEntry]: + """Returns the cache list on ``invocation_context`` if initialized, else ``[]``.""" + cache = getattr(invocation_context, attr_name, None) + return cache if isinstance(cache, list) else [] + + +class CacheConfig: + """Configuration for multimodal streaming cache behavior.""" + + def __init__( + self, + max_cache_size_bytes: int = 10 * 1024 * 1024, + max_cache_duration_seconds: float = 300.0, + auto_flush_threshold: int = 100, + ) -> None: + self.max_cache_size_bytes = max_cache_size_bytes + self.max_cache_duration_seconds = max_cache_duration_seconds + self.auto_flush_threshold = auto_flush_threshold + + +class CacheManager: + """Manages multimodal caching and flushing for live streaming flows.""" + + def __init__(self, config: CacheConfig | None = None) -> None: + self.config = config if config is not None else CacheConfig() + + def _resolve_cache_target( + self, invocation_context: InvocationContext, cache_type: str + ) -> tuple[list[RealtimeCacheEntry], str]: + """Validates ``cache_type`` and returns ``(cache_list, role)``.""" + if cache_type == 'input': + attr, role = 'input_realtime_cache', 'user' + elif cache_type == 'output': + attr, role = 'output_realtime_cache', 'model' + else: + raise ValueError("cache_type must be either 'input' or 'output'") + + cache = getattr(invocation_context, attr, None) + if not isinstance(cache, list): + cache = [] + setattr(invocation_context, attr, cache) + return cache, role + + def cache_audio( + self, + invocation_context: InvocationContext, + audio_blob: types.Blob, + cache_type: str, + ) -> None: + """Cache incoming user or outgoing model audio data.""" + audio_data = _require_blob_data(audio_blob) + cache, role = self._resolve_cache_target(invocation_context, cache_type) + cache.append( + RealtimeCacheEntry( + role=role, data=audio_blob, timestamp=platform_time.get_time() + ) + ) + logger.debug( + 'Cached %s audio chunk: %d bytes, cache size: %d', + cache_type, + len(audio_data), + len(cache), + ) + + def cache_blob( + self, + invocation_context: InvocationContext, + blob: types.Blob, + cache_type: str, + ) -> None: + """Routes an incoming or outgoing blob to the appropriate cache by MIME type.""" + mime_type = _normalize_mime_type(blob.mime_type, default='') + if mime_type.startswith('audio/'): + self.cache_audio(invocation_context, blob, cache_type) + else: + logger.debug( + 'Skipping non-audio blob in cache_blob: %s (cache_type=%s)', + mime_type, + cache_type, + ) + + async def flush_caches( + self, + invocation_context: InvocationContext, + flush_user_audio: bool = True, + flush_model_audio: bool = True, + ) -> list[Event]: + """Flush caches concurrently to artifact services.""" + if not invocation_context.artifact_service: + logger.debug('Skipping cache flush: no artifact service or empty cache') + return [] + + tasks = [] + targets: list[str] = [] + + input_audio = _get_cache_list(invocation_context, 'input_realtime_cache') + if flush_user_audio and input_audio: + tasks.append( + self._flush_audio_cache_to_services( + invocation_context, input_audio, 'input_audio' + ) + ) + targets.append('input_realtime_cache') + + output_audio = _get_cache_list(invocation_context, 'output_realtime_cache') + if flush_model_audio and output_audio: + tasks.append( + self._flush_audio_cache_to_services( + invocation_context, output_audio, 'output_audio' + ) + ) + targets.append('output_realtime_cache') + + if not tasks: + return [] + + results = await asyncio.gather(*tasks, return_exceptions=True) + flushed_events: list[Event] = [] + for attr_name, result in zip(targets, results): + if isinstance(result, Exception): + logger.error('Failed to flush %s: %s', attr_name, result) + elif isinstance(result, Event): + flushed_events.append(result) + setattr(invocation_context, attr_name, []) + return flushed_events + + async def _save_artifact_and_build_event( + self, + invocation_context: InvocationContext, + *, + filename: str, + data: bytes, + mime_type: str, + role: str, + timestamp: float, + ) -> Event: + """Saves an artifact and returns the corresponding ``_adk_live`` Event.""" + assert invocation_context.artifact_service is not None + artifact = types.Part( + inline_data=types.Blob(data=data, mime_type=mime_type) + ) + revision_id = await invocation_context.artifact_service.save_artifact( + app_name=invocation_context.app_name, + user_id=invocation_context.user_id, + session_id=invocation_context.session.id, + filename=filename, + artifact=artifact, + ) + artifact_ref = ( + f'artifact://{invocation_context.app_name}/' + f'{invocation_context.user_id}/{invocation_context.session.id}/' + f'_adk_live/{filename}#{revision_id}' + ) + author = ( + _require_agent_name(invocation_context) if role == 'model' else role + ) + return Event( + id=Event.new_id(), + invocation_id=invocation_context.invocation_id, + author=author, + content=types.Content( + role=role, + parts=[ + types.Part( + file_data=types.FileData( + file_uri=artifact_ref, mime_type=mime_type + ) + ) + ], + ), + timestamp=timestamp, + ) + + async def _flush_audio_cache_to_services( + self, + invocation_context: InvocationContext, + audio_cache: list[RealtimeCacheEntry], + cache_type: str, + ) -> Event | None: + """Flush a list of audio cache entries to artifact services.""" + if not invocation_context.artifact_service or not audio_cache: + logger.debug( + 'Skipping audio cache flush: no artifact service or empty cache' + ) + return None + + try: + mime_type = audio_cache[0].data.mime_type or 'audio/pcm' + combined_audio_data = b''.join( + entry.data.data or b'' for entry in audio_cache + ) + timestamp = int(audio_cache[0].timestamp * 1000) + raw_mime = _normalize_mime_type(mime_type, default='audio/pcm') + ext = raw_mime.split('/')[-1] if '/' in raw_mime else 'pcm' + filename = f'adk_live_audio_storage_{cache_type}_{timestamp}.{ext}' + return await self._save_artifact_and_build_event( + invocation_context, + filename=filename, + data=combined_audio_data, + mime_type=mime_type, + role=audio_cache[0].role, + timestamp=audio_cache[0].timestamp, + ) + except Exception as e: # pylint: disable=broad-exception-caught + logger.error('Failed to flush %s cache: %s', cache_type, e) + return None + + _flush_cache_to_services = _flush_audio_cache_to_services + + def get_cache_stats( + self, invocation_context: InvocationContext + ) -> dict[str, int]: + """Get statistics about current cache state.""" + input_audio = _get_cache_list(invocation_context, 'input_realtime_cache') + output_audio = _get_cache_list(invocation_context, 'output_realtime_cache') + input_count = len(input_audio) + output_count = len(output_audio) + input_bytes = sum( + len(_require_blob_data(entry.data)) for entry in input_audio + ) + output_bytes = sum( + len(_require_blob_data(entry.data)) for entry in output_audio + ) + return { + 'input_chunks': input_count, + 'output_chunks': output_count, + 'input_bytes': input_bytes, + 'output_bytes': output_bytes, + 'total_chunks': input_count + output_count, + 'total_bytes': input_bytes + output_bytes, + } + + +AudioCacheConfig = CacheConfig +AudioCacheManager = CacheManager diff --git a/src/google/adk/live/_flow_utils.py b/src/google/adk/live/_flow_utils.py index 2b994d4a50..50b61aec0e 100644 --- a/src/google/adk/live/_flow_utils.py +++ b/src/google/adk/live/_flow_utils.py @@ -175,23 +175,23 @@ async def handle_control_event_flush( Returns: A list of Event objects created from the flushed caches. """ - audio_cache_manager = flow.audio_cache_manager + cache_manager = flow.cache_manager # Log cache statistics if enabled if DEFAULT_ENABLE_CACHE_STATISTICS: - stats = audio_cache_manager.get_cache_stats(invocation_context) + stats = cache_manager.get_cache_stats(invocation_context) logger.debug('Audio cache stats: %s', stats) if llm_response.interrupted: # user interrupts so the model will stop. we can flush model audio here - return await audio_cache_manager.flush_caches( + return await cache_manager.flush_caches( invocation_context, flush_user_audio=False, flush_model_audio=True, ) elif llm_response.turn_complete: # turn completes so we can flush both user and model - return await audio_cache_manager.flush_caches( + return await cache_manager.flush_caches( invocation_context, flush_user_audio=True, flush_model_audio=True, diff --git a/src/google/adk/live/_live_llm_flow.py b/src/google/adk/live/_live_llm_flow.py index 20d82eac49..119220e4c5 100644 --- a/src/google/adk/live/_live_llm_flow.py +++ b/src/google/adk/live/_live_llm_flow.py @@ -96,7 +96,7 @@ async def send_to_model( ) -> None: """Sends data to model.""" run_config = _require_run_config(invocation_context) - audio_cache_manager = flow.audio_cache_manager + cache_manager = flow.cache_manager while True: live_request_queue = invocation_context.live_request_queue assert live_request_queue is not None @@ -152,15 +152,8 @@ async def send_to_model( types.LiveClientRealtimeInput(audio_stream_end=True) # type: ignore[arg-type] ) elif live_request.blob: - # Cache input audio chunks before flushing. The cache concatenates - # every chunk into one audio file, so other blobs (e.g. video frames) - # must stay out of it. - if ( - run_config.save_live_blob - and live_request.blob.mime_type - and live_request.blob.mime_type.startswith('audio/') - ): - audio_cache_manager.cache_audio( + if run_config.save_live_blob: + cache_manager.cache_blob( invocation_context, live_request.blob, cache_type='input' ) @@ -222,7 +215,7 @@ async def receive_from_model( ) -> AsyncGenerator[Event, None]: """Receive data from model and process events using BaseLlmConnection.""" run_config = _require_run_config(invocation_context) - audio_cache_manager = flow.audio_cache_manager + cache_manager = flow.cache_manager def get_author_for_event(llm_response: LlmResponse) -> str: """Get the author of the event. @@ -314,25 +307,17 @@ def get_author_for_event(llm_response: LlmResponse) -> str: ) ) as postprocess_agen: async for event in postprocess_agen: - # Cache output audio chunks from model responses - # TODO: support video data if ( run_config.save_live_blob and event.content and event.content.parts ): for part in event.content.parts: - if ( - part.inline_data - and part.inline_data.mime_type - and part.inline_data.mime_type.startswith('audio/') - ): - audio_blob = types.Blob( - data=part.inline_data.data, - mime_type=part.inline_data.mime_type, - ) - audio_cache_manager.cache_audio( - invocation_context, audio_blob, cache_type='output' + if part.inline_data: + cache_manager.cache_blob( + invocation_context, + part.inline_data, + cache_type='output', ) yield event diff --git a/tests/unittests/flows/llm_flows/test_audio_cache_manager.py b/tests/unittests/flows/llm_flows/test_audio_cache_manager.py index dab29afcd2..54781eaef2 100644 --- a/tests/unittests/flows/llm_flows/test_audio_cache_manager.py +++ b/tests/unittests/flows/llm_flows/test_audio_cache_manager.py @@ -23,12 +23,16 @@ from google.adk.flows.llm_flows.audio_cache_manager import AudioCacheManager as LegacyAudioCacheManager from google.adk.live._audio_cache_manager import AudioCacheConfig from google.adk.live._audio_cache_manager import AudioCacheManager +from google.adk.live._cache_manager import CacheConfig +from google.adk.live._cache_manager import CacheManager import pytest def test_audio_cache_manager_reexport(): assert LegacyAudioCacheManager is AudioCacheManager + assert LegacyAudioCacheManager is CacheManager assert LegacyAudioCacheConfig is AudioCacheConfig + assert LegacyAudioCacheConfig is CacheConfig assert getattr(legacy_module, 'AudioCacheManager') is AudioCacheManager assert getattr(legacy_module, 'AudioCacheConfig') is AudioCacheConfig @@ -36,6 +40,6 @@ def test_audio_cache_manager_reexport(): def test_audio_cache_manager_deprecation_warning(): with pytest.warns( DeprecationWarning, - match='use google.adk.live._audio_cache_manager instead', + match='use google.adk.live._cache_manager instead', ): importlib.reload(legacy_module) diff --git a/tests/unittests/live/test_audio_cache_manager.py b/tests/unittests/live/test_audio_cache_manager.py index d80bc833d8..7501e54239 100644 --- a/tests/unittests/live/test_audio_cache_manager.py +++ b/tests/unittests/live/test_audio_cache_manager.py @@ -18,9 +18,14 @@ from unittest.mock import AsyncMock from unittest.mock import Mock +from google.adk.live import _audio_cache_manager as compat_module from google.adk.live._audio_cache_manager import AudioCacheConfig from google.adk.live._audio_cache_manager import AudioCacheManager from google.adk.live._audio_cache_manager import RealtimeCacheEntry +from google.adk.live._cache_manager import CacheConfig +from google.adk.live._cache_manager import CacheManager +from google.adk.live._cache_manager import logger as canonical_logger +from google.adk.live._cache_manager import RealtimeCacheEntry as CanonicalRealtimeCacheEntry from google.genai import types import pydantic import pytest @@ -481,3 +486,17 @@ async def test_flush_event_author_for_model_audio(self): assert len(events) == 1 assert events[0].author == 'my_test_agent' # Agent name, not 'model' assert events[0].content.role == 'model' # Role is still 'model' + + +def test_audio_cache_manager_reexports_cache_manager_symbols(): + """The legacy `_audio_cache_manager` module must alias `_cache_manager`.""" + assert AudioCacheManager is CacheManager + assert AudioCacheConfig is CacheConfig + assert RealtimeCacheEntry is CanonicalRealtimeCacheEntry + assert compat_module.logger is canonical_logger + assert set(compat_module.__all__) == { + 'AudioCacheConfig', + 'AudioCacheManager', + 'RealtimeCacheEntry', + 'logger', + } diff --git a/tests/unittests/live/test_cache_manager.py b/tests/unittests/live/test_cache_manager.py new file mode 100644 index 0000000000..d79d831b0f --- /dev/null +++ b/tests/unittests/live/test_cache_manager.py @@ -0,0 +1,133 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for the generalized live cache manager (`CacheManager`).""" + +from __future__ import annotations + +from unittest import mock + +from google.adk.live._cache_manager import AudioCacheConfig +from google.adk.live._cache_manager import AudioCacheManager +from google.adk.live._cache_manager import CacheConfig +from google.adk.live._cache_manager import CacheManager +from google.genai import types +import pytest + +from .. import testing_utils + + +def test_cache_config_defaults_and_custom_values(): + """Verifies CacheConfig defaults, custom parameters, and AudioCacheConfig alias.""" + default_config = CacheConfig() + assert default_config.max_cache_size_bytes == 10 * 1024 * 1024 + assert default_config.max_cache_duration_seconds == 300.0 + assert default_config.auto_flush_threshold == 100 + custom_config = CacheConfig(5 * 1024 * 1024, 120.0, 50) + assert custom_config.max_cache_size_bytes == 5 * 1024 * 1024 + assert custom_config.max_cache_duration_seconds == 120.0 + assert custom_config.auto_flush_threshold == 50 + assert AudioCacheConfig is CacheConfig + assert AudioCacheManager is CacheManager + + +@pytest.mark.asyncio +async def test_invocation_context_initializes_media_cache_fields_to_none(): + """Verifies input_media_realtime_cache and output_media_realtime_cache default to None.""" + ctx = await testing_utils.create_invocation_context( + testing_utils.create_test_agent() + ) + assert ctx.input_media_realtime_cache is None + assert ctx.output_media_realtime_cache is None + + +@pytest.mark.asyncio +async def test_cache_blob_routes_audio_and_skips_non_audio(): + """Verifies cache_blob routes audio/* blobs and skips non-audio blobs.""" + manager = CacheManager() + ctx = await testing_utils.create_invocation_context( + testing_utils.create_test_agent() + ) + in_audio = types.Blob(data=b'user_pcm', mime_type='audio/pcm') + out_audio = types.Blob(data=b'model_wav', mime_type='AUDIO/WAV;rate=24000') + manager.cache_blob(ctx, in_audio, 'input') + manager.cache_blob(ctx, out_audio, 'output') + manager.cache_blob( + ctx, types.Blob(data=b'jpg', mime_type='image/jpeg'), 'input' + ) + manager.cache_blob( + ctx, types.Blob(data=b'mp4', mime_type='video/mp4'), 'output' + ) + + assert [e.data for e in ctx.input_realtime_cache] == [in_audio] + assert [e.data for e in ctx.output_realtime_cache] == [out_audio] + assert ctx.input_media_realtime_cache is None + assert ctx.output_media_realtime_cache is None + + +@pytest.mark.asyncio +async def test_flush_normalizes_parameterized_audio_mime_extension(): + """Verifies parameterized MIME types like audio/pcm;rate=16000 produce a .pcm extension.""" + manager = CacheManager() + ctx = await testing_utils.create_invocation_context( + testing_utils.create_test_agent() + ) + ctx.artifact_service = mock.AsyncMock( + save_artifact=mock.AsyncMock(return_value=7) + ) + manager.cache_blob( + ctx, types.Blob(data=b'pcm', mime_type='audio/pcm;rate=16000'), 'input' + ) + events = await manager.flush_caches(ctx) + + assert len(events) == 1 + saved_name = ctx.artifact_service.save_artifact.call_args.kwargs['filename'] + assert saved_name.endswith('.pcm') and ';' not in saved_name + + +@pytest.mark.asyncio +async def test_failed_input_flush_does_not_block_output_flush(): + """Verifies concurrent flushing keeps a failed cache while clearing the succeeded cache.""" + manager = CacheManager() + ctx = await testing_utils.create_invocation_context( + testing_utils.create_test_agent() + ) + + async def fail_on_input(**kwargs): + if 'input_audio' in kwargs['filename']: + raise RuntimeError('input storage failure') + return 9 + + ctx.artifact_service = mock.AsyncMock( + save_artifact=mock.AsyncMock(side_effect=fail_on_input) + ) + manager.cache_audio( + ctx, types.Blob(data=b'in', mime_type='audio/pcm'), 'input' + ) + manager.cache_audio( + ctx, types.Blob(data=b'out', mime_type='audio/pcm'), 'output' + ) + events = await manager.flush_caches(ctx) + + assert len(events) == 1 and events[0].content.role == 'model' + assert len(ctx.input_realtime_cache) == 1 + assert not ctx.output_realtime_cache + + +def test_get_cache_stats_handles_bare_mock_invocation_context(): + """Verifies get_cache_stats returns zero counts on a bare Mock context.""" + stats = CacheManager().get_cache_stats( + mock.Mock(input_realtime_cache=None, output_realtime_cache=None) + ) + assert stats['total_chunks'] == 0 and stats['total_bytes'] == 0 diff --git a/tests/unittests/live/test_flow_utils.py b/tests/unittests/live/test_flow_utils.py index 13957da6d4..612046db93 100644 --- a/tests/unittests/live/test_flow_utils.py +++ b/tests/unittests/live/test_flow_utils.py @@ -47,11 +47,9 @@ def test_require_live_request_queue_raises_when_missing(): async def test_handle_control_event_flush_flushes_model_audio_on_interrupt(): """Flushes only model audio when the response indicates an interruption.""" flushed_event = Event(author='model') - audio_cache_manager = mock.Mock() - audio_cache_manager.flush_caches = mock.AsyncMock( - return_value=[flushed_event] - ) - flow = mock.Mock(audio_cache_manager=audio_cache_manager) + cache_manager = mock.Mock() + cache_manager.flush_caches = mock.AsyncMock(return_value=[flushed_event]) + flow = mock.Mock(cache_manager=cache_manager) invocation_context = mock.Mock() llm_response = LlmResponse(interrupted=True) @@ -60,7 +58,7 @@ async def test_handle_control_event_flush_flushes_model_audio_on_interrupt(): ) assert events == [flushed_event] - audio_cache_manager.flush_caches.assert_awaited_once_with( + cache_manager.flush_caches.assert_awaited_once_with( invocation_context, flush_user_audio=False, flush_model_audio=True, @@ -70,11 +68,9 @@ async def test_handle_control_event_flush_flushes_model_audio_on_interrupt(): async def test_handle_control_event_flush_flushes_both_on_turn_complete(): """Flushes both user and model audio when the turn completes.""" flushed_event = Event(author='model') - audio_cache_manager = mock.Mock() - audio_cache_manager.flush_caches = mock.AsyncMock( - return_value=[flushed_event] - ) - flow = mock.Mock(audio_cache_manager=audio_cache_manager) + cache_manager = mock.Mock() + cache_manager.flush_caches = mock.AsyncMock(return_value=[flushed_event]) + flow = mock.Mock(cache_manager=cache_manager) invocation_context = mock.Mock() llm_response = LlmResponse(turn_complete=True) @@ -83,7 +79,7 @@ async def test_handle_control_event_flush_flushes_both_on_turn_complete(): ) assert events == [flushed_event] - audio_cache_manager.flush_caches.assert_awaited_once_with( + cache_manager.flush_caches.assert_awaited_once_with( invocation_context, flush_user_audio=True, flush_model_audio=True, diff --git a/tests/unittests/live/test_live_llm_flow.py b/tests/unittests/live/test_live_llm_flow.py index 753a157e65..58d35a0eaa 100644 --- a/tests/unittests/live/test_live_llm_flow.py +++ b/tests/unittests/live/test_live_llm_flow.py @@ -25,6 +25,7 @@ from google.adk.events.event import Event from google.adk.flows.llm_flows.base_llm_flow import BaseLlmFlow from google.adk.live import _live_llm_flow +from google.adk.live._cache_manager import CacheManager from google.adk.live.live_request_queue import LiveRequestQueue from google.adk.models.llm_request import LlmRequest from google.adk.models.llm_response import LlmResponse @@ -181,7 +182,7 @@ async def test_handle_control_event_flush_on_interrupted(): response = LlmResponse(interrupted=True) with mock.patch.object( - flow.audio_cache_manager, 'flush_caches', new_callable=mock.AsyncMock + flow.cache_manager, 'flush_caches', new_callable=mock.AsyncMock ) as mock_flush: mock_flush.return_value = [Event(id='flushed-event')] events = await _live_llm_flow.handle_control_event_flush( @@ -201,7 +202,7 @@ async def test_handle_control_event_flush_on_turn_complete(): response = LlmResponse(turn_complete=True) with mock.patch.object( - flow.audio_cache_manager, 'flush_caches', new_callable=mock.AsyncMock + flow.cache_manager, 'flush_caches', new_callable=mock.AsyncMock ) as mock_flush: mock_flush.return_value = [Event(id='flushed-event')] events = await _live_llm_flow.handle_control_event_flush( @@ -308,48 +309,50 @@ async def test_handle_control_event_flush_logs_stats_when_enabled(): with ( mock.patch.object(_flow_utils, 'DEFAULT_ENABLE_CACHE_STATISTICS', True), mock.patch.object( - flow.audio_cache_manager, 'get_cache_stats' + flow.cache_manager, 'get_cache_stats' ) as mock_get_stats, - mock.patch.object( - flow.audio_cache_manager, 'flush_caches', return_value=[] - ), + mock.patch.object(flow.cache_manager, 'flush_caches', return_value=[]), ): await _live_llm_flow.handle_control_event_flush(flow, context, response) mock_get_stats.assert_called_once_with(context) -async def test_send_to_model_uses_flow_audio_cache_manager(): - """send_to_model accesses the audio cache manager directly from the flow instance.""" +async def _run_send_to_model_once(flow, context): + """Drains the queued requests deterministically by appending a close signal.""" + mock_connection = mock.AsyncMock() + context.live_request_queue.close() + await _live_llm_flow.send_to_model( + flow, mock_connection, context, LlmRequest() + ) + return mock_connection + + +async def test_send_to_model_uses_flow_cache_manager(): + """send_to_model delegates blob caching to flow.cache_manager.cache_blob.""" flow = _TestBaseLlmFlow() queue = LiveRequestQueue() - queue.send_realtime(types.Blob(mime_type='audio/pcm', data=b'audio_bytes')) + audio_blob = types.Blob(mime_type='audio/pcm', data=b'audio_bytes') + queue.send_realtime(audio_blob) context = _create_test_context( live_request_queue=queue, run_config=RunConfig(save_live_blob=True) ) - mock_connection = mock.AsyncMock() with mock.patch.object( - flow.audio_cache_manager, 'cache_audio' - ) as mock_cache_audio: - # Run send_to_model briefly and cancel it after processing the queued item - send_task = asyncio.create_task( - _live_llm_flow.send_to_model( - flow, mock_connection, context, LlmRequest() - ) - ) - await asyncio.sleep(0.01) - send_task.cancel() - try: - await send_task - except asyncio.CancelledError: - pass + flow.cache_manager, 'cache_blob', wraps=flow.cache_manager.cache_blob + ) as mock_cache_blob: + mock_connection = await _run_send_to_model_once(flow, context) - mock_cache_audio.assert_called_once() + mock_cache_blob.assert_called_once_with( + context, audio_blob, cache_type='input' + ) + assert len(context.input_realtime_cache) == 1 + assert context.input_realtime_cache[0].data == audio_blob + mock_connection.send_realtime.assert_awaited_once_with(audio_blob) async def test_send_to_model_caches_only_audio_blobs(): - """Non-audio blobs such as video frames are sent but not cached.""" + """Non-audio blobs such as video frames are sent but not cached in CL 2B.""" flow = _TestBaseLlmFlow() queue = LiveRequestQueue() audio_blob = types.Blob(mime_type='audio/pcm', data=b'audio_bytes') @@ -359,20 +362,62 @@ async def test_send_to_model_caches_only_audio_blobs(): context = _create_test_context( live_request_queue=queue, run_config=RunConfig(save_live_blob=True) ) - mock_connection = mock.AsyncMock() - send_task = asyncio.create_task( - _live_llm_flow.send_to_model(flow, mock_connection, context, LlmRequest()) - ) - await asyncio.sleep(0.01) - send_task.cancel() - try: - await send_task - except asyncio.CancelledError: - pass + mock_connection = await _run_send_to_model_once(flow, context) assert [entry.data for entry in context.input_realtime_cache] == [audio_blob] + assert context.input_media_realtime_cache is None assert mock_connection.send_realtime.await_args_list == [ mock.call(video_blob), mock.call(audio_blob), ] + + +async def test_audio_cache_manager_alias_still_resolves_and_warns(): + """Reading or setting `flow.audio_cache_manager` forwards to `cache_manager` and warns.""" + flow = _TestBaseLlmFlow() + + with pytest.warns(DeprecationWarning, match='use cache_manager instead'): + alias = flow.audio_cache_manager + assert alias is flow.cache_manager + + replacement = CacheManager() + with pytest.warns(DeprecationWarning, match='use cache_manager instead'): + flow.audio_cache_manager = replacement + assert flow.cache_manager is replacement + + +async def test_receive_from_model_routes_audio_via_cache_blob(): + """Model-generated audio/* inline_data blobs are routed via flow.cache_manager.cache_blob.""" + flow = _TestBaseLlmFlow() + context = _create_test_context(run_config=RunConfig(save_live_blob=True)) + calls = 0 + + async def _receive(): + nonlocal calls + calls += 1 + if calls == 1: + yield LlmResponse( + content=types.Content( + role='model', + parts=[ + types.Part( + inline_data=types.Blob( + mime_type='audio/pcm', data=b'model_audio' + ) + ) + ], + ) + ) + + connection = mock.MagicMock(receive=_receive) + events = [ + event + async for event in _live_llm_flow.receive_from_model( + flow, connection, context, LlmRequest() + ) + ] + assert len(events) == 1 + assert len(context.output_realtime_cache) == 1 + assert context.output_realtime_cache[0].data.data == b'model_audio' + assert context.output_media_realtime_cache is None