feat: Support interactions API for calling models

Co-authored-by: Xiang (Sean) Zhou <seanzhougoogle@google.com>
PiperOrigin-RevId: 843032402
This commit is contained in:
Xiang (Sean) Zhou
2025-12-10 21:21:14 -08:00
committed by Copybara-Service
parent f0bdcaba44
commit c6320caaa5
8 changed files with 2238 additions and 1 deletions
@@ -0,0 +1,140 @@
# Copyright 2025 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.
"""Interactions API processor for LLM requests."""
from __future__ import annotations
import logging
from typing import AsyncGenerator
from typing import Optional
from typing import TYPE_CHECKING
from ...events.event import Event
from ._base_llm_processor import BaseLlmRequestProcessor
if TYPE_CHECKING:
from ...agents.invocation_context import InvocationContext
from ...models.llm_request import LlmRequest
logger = logging.getLogger('google_adk.' + __name__)
class InteractionsRequestProcessor(BaseLlmRequestProcessor):
"""Request processor for Interactions API stateful conversations.
This processor extracts the previous_interaction_id from session events
to enable stateful conversation chaining via the Interactions API.
The actual content filtering (retaining only latest user messages) is
done in the Gemini class when using the Interactions API.
"""
async def run_async(
self, invocation_context: 'InvocationContext', llm_request: 'LlmRequest'
) -> AsyncGenerator[Event, None]:
"""Process LLM request to extract previous_interaction_id.
Args:
invocation_context: Invocation context containing agent and session info
llm_request: Request to process
Yields:
Event: No events are yielded by this processor
"""
from ...agents.llm_agent import LlmAgent
from ...models.google_llm import Gemini
agent = invocation_context.agent
# Only process if using Gemini with interactions API
if not isinstance(agent, LlmAgent):
return
if not isinstance(agent.model, Gemini):
return
if not agent.model.use_interactions_api:
return
# Extract previous interaction ID from session events
previous_interaction_id = self._find_previous_interaction_id(
invocation_context
)
if previous_interaction_id:
llm_request.previous_interaction_id = previous_interaction_id
logger.debug(
'Found previous_interaction_id for interactions API: %s',
previous_interaction_id,
)
# Don't yield any events - this is just a preprocessing step
return
yield # Required for AsyncGenerator
def _find_previous_interaction_id(
self, invocation_context: 'InvocationContext'
) -> Optional[str]:
"""Find the previous interaction ID from session events.
For interactions API stateful mode, we need to find the most recent
interaction_id from model responses to chain interactions.
Args:
invocation_context: The invocation context containing session events.
Returns:
The previous interaction ID if found, None otherwise.
"""
events = invocation_context.session.events
current_branch = invocation_context.branch
agent_name = invocation_context.agent.name
logger.debug(
'Finding previous_interaction_id: agent=%s, branch=%s, num_events=%d',
agent_name,
current_branch,
len(events),
)
# Iterate backwards through events to find the most recent interaction_id
for event in reversed(events):
# Skip events not in current branch
if not self._is_event_in_branch(current_branch, event):
logger.debug(
'Skipping event not in branch: author=%s, branch=%s, current=%s',
event.author,
event.branch,
current_branch,
)
continue
# Look for model responses with interaction_id from this agent
logger.debug(
'Checking event: author=%s, interaction_id=%s, branch=%s',
event.author,
event.interaction_id,
event.branch,
)
# Only consider events from this agent (skip sub-agent events)
if event.author == agent_name and event.interaction_id:
logger.debug(
'Found interaction_id from agent %s: %s',
agent_name,
event.interaction_id,
)
return event.interaction_id
return None
def _is_event_in_branch(
self, current_branch: Optional[str], event: Event
) -> bool:
"""Check if an event belongs to the current branch.
Args:
current_branch: The current branch name.
event: The event to check.
Returns:
True if the event belongs to the current branch.
"""
if not current_branch:
# No branch means we're at the root, include all events without branch
return not event.branch
# Event must be in the same branch or have no branch (root level)
return event.branch == current_branch or not event.branch
# Module-level processor instance for use in flow configuration
request_processor = InteractionsRequestProcessor()
@@ -26,6 +26,7 @@ from . import contents
from . import context_cache_processor
from . import identity
from . import instructions
from . import interactions_processor
from . import request_confirmation
from ...auth import auth_preprocessor
from .base_llm_flow import BaseLlmFlow
@@ -51,6 +52,9 @@ class SingleFlow(BaseLlmFlow):
contents.request_processor,
# Context cache processor sets up cache config and finds existing cache metadata
context_cache_processor.request_processor,
# Interactions processor extracts previous_interaction_id for stateful
# conversations via the Interactions API
interactions_processor.request_processor,
# Some implementations of NL Planning mark planning contents as thoughts
# in the post processor. Since these need to be unmarked, NL Planning
# should be after contents.
+61 -1
View File
@@ -19,6 +19,7 @@ import contextlib
import copy
from functools import cached_property
import logging
from typing import Any
from typing import AsyncGenerator
from typing import cast
from typing import Optional
@@ -83,12 +84,32 @@ class Gemini(BaseLlm):
Attributes:
model: The name of the Gemini model.
use_interactions_api: Whether to use the interactions API for model
invocation.
"""
model: str = 'gemini-2.5-flash'
speech_config: Optional[types.SpeechConfig] = None
use_interactions_api: bool = False
"""Whether to use the interactions API for model invocation.
When enabled, uses the interactions API (client.aio.interactions.create())
instead of the traditional generate_content API. The interactions API
provides stateful conversation capabilities, allowing you to chain
interactions using previous_interaction_id instead of sending full history.
The response format will be converted to match the existing LlmResponse
structure for compatibility.
Sample:
```python
agent = Agent(
model=Gemini(use_interactions_api=True)
)
```
"""
retry_options: Optional[types.HttpRetryOptions] = None
"""Allow Gemini to retry failed responses.
@@ -163,7 +184,6 @@ class Gemini(BaseLlm):
self._api_backend,
stream,
)
logger.debug(_build_request_log(llm_request))
# Always add tracking headers to custom headers given it will override
# the headers set in the api client constructor to avoid tracking headers
@@ -176,6 +196,16 @@ class Gemini(BaseLlm):
)
try:
# Use interactions API if enabled
if self.use_interactions_api:
async for llm_response in self._generate_content_via_interactions(
llm_request, stream
):
yield llm_response
return
logger.debug(_build_request_log(llm_request))
if stream:
responses = await self.api_client.aio.models.generate_content_stream(
model=llm_request.model,
@@ -231,6 +261,36 @@ class Gemini(BaseLlm):
raise ce
async def _generate_content_via_interactions(
self,
llm_request: LlmRequest,
stream: bool,
) -> AsyncGenerator[LlmResponse, None]:
"""Generate content using the interactions API.
The interactions API provides stateful conversation capabilities. When
previous_interaction_id is set in the request, the API chains interactions
instead of requiring full conversation history.
Note: Context caching is not used with the Interactions API since it
maintains conversation state via previous_interaction_id.
Args:
llm_request: The LLM request to send.
stream: Whether to stream the response.
Yields:
LlmResponse objects converted from interaction responses.
"""
from .interactions_utils import generate_content_via_interactions
async for llm_response in generate_content_via_interactions(
api_client=self.api_client,
llm_request=llm_request,
stream=stream,
):
yield llm_response
@cached_property
def api_client(self) -> Client:
"""Provides the api client.
File diff suppressed because it is too large Load Diff
+8
View File
@@ -91,6 +91,14 @@ class LlmRequest(BaseModel):
cacheable_contents_token_count: Optional[int] = None
"""Token count from previous request's prompt, used for cache size validation."""
previous_interaction_id: Optional[str] = None
"""The ID of the previous interaction for stateful conversations.
When using the interactions API, this ID is used to chain interactions
together, allowing the API to maintain conversation state without sending
the full history.
"""
def append_instructions(
self, instructions: Union[list[str], types.Content]
) -> list[types.Content]:
+7
View File
@@ -135,6 +135,13 @@ class LlmResponse(BaseModel):
This field is automatically populated when citation is enabled.
"""
interaction_id: Optional[str] = None
"""The interaction ID from the interactions API.
This field is populated when using the interactions API for model invocation.
It can be used to identify and chain interactions for stateful conversations.
"""
@staticmethod
def create(
generate_content_response: types.GenerateContentResponse,
@@ -0,0 +1,223 @@
# Copyright 2025 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 interactions processor."""
from unittest.mock import MagicMock
from google.adk.events.event import Event
from google.adk.flows.llm_flows import interactions_processor
from google.genai import types
import pytest
class TestInteractionsRequestProcessor:
"""Tests for InteractionsRequestProcessor."""
def test_find_previous_interaction_id_empty_events(self):
"""Test that None is returned when there are no events."""
processor = interactions_processor.InteractionsRequestProcessor()
invocation_context = MagicMock()
invocation_context.session.events = []
invocation_context.branch = None
invocation_context.agent.name = "test_agent"
result = processor._find_previous_interaction_id(invocation_context)
assert result is None
def test_find_previous_interaction_id_user_only_events(self):
"""Test that None is returned when only user events exist."""
processor = interactions_processor.InteractionsRequestProcessor()
events = [
Event(
invocation_id="inv1",
author="user",
content=types.UserContent("Hello"),
),
Event(
invocation_id="inv2",
author="user",
content=types.UserContent("World"),
),
]
invocation_context = MagicMock()
invocation_context.session.events = events
invocation_context.branch = None
invocation_context.agent.name = "test_agent"
result = processor._find_previous_interaction_id(invocation_context)
assert result is None
def test_find_previous_interaction_id_no_interaction_id(self):
"""Test that None is returned when model events have no interaction_id."""
processor = interactions_processor.InteractionsRequestProcessor()
events = [
Event(
invocation_id="inv1",
author="user",
content=types.UserContent("Hello"),
),
Event(
invocation_id="inv2",
author="test_agent",
content=types.ModelContent("Response without interaction_id"),
),
]
invocation_context = MagicMock()
invocation_context.session.events = events
invocation_context.branch = None
invocation_context.agent.name = "test_agent"
result = processor._find_previous_interaction_id(invocation_context)
assert result is None
def test_find_previous_interaction_id_from_model_event(self):
"""Test that interaction_id is returned from model event."""
processor = interactions_processor.InteractionsRequestProcessor()
events = [
Event(
invocation_id="inv1",
author="user",
content=types.UserContent("Hello"),
),
Event(
invocation_id="inv2",
author="test_agent",
content=types.ModelContent("Response"),
interaction_id="interaction_123",
),
]
invocation_context = MagicMock()
invocation_context.session.events = events
invocation_context.branch = None
invocation_context.agent.name = "test_agent"
result = processor._find_previous_interaction_id(invocation_context)
assert result == "interaction_123"
def test_find_previous_interaction_id_returns_most_recent(self):
"""Test that the most recent interaction_id is returned."""
processor = interactions_processor.InteractionsRequestProcessor()
events = [
Event(
invocation_id="inv1",
author="user",
content=types.UserContent("Hello"),
),
Event(
invocation_id="inv2",
author="test_agent",
content=types.ModelContent("First response"),
interaction_id="interaction_first",
),
Event(
invocation_id="inv3",
author="user",
content=types.UserContent("Second message"),
),
Event(
invocation_id="inv4",
author="test_agent",
content=types.ModelContent("Second response"),
interaction_id="interaction_second",
),
]
invocation_context = MagicMock()
invocation_context.session.events = events
invocation_context.branch = None
invocation_context.agent.name = "test_agent"
result = processor._find_previous_interaction_id(invocation_context)
assert result == "interaction_second"
def test_find_previous_interaction_id_skips_user_events(self):
"""Test that user events with interaction_id are skipped."""
processor = interactions_processor.InteractionsRequestProcessor()
events = [
Event(
invocation_id="inv1",
author="test_agent",
content=types.ModelContent("Model response"),
interaction_id="interaction_model",
),
Event(
invocation_id="inv2",
author="user",
content=types.UserContent("User message"),
interaction_id="interaction_user", # This should be skipped
),
]
invocation_context = MagicMock()
invocation_context.session.events = events
invocation_context.branch = None
invocation_context.agent.name = "test_agent"
result = processor._find_previous_interaction_id(invocation_context)
assert result == "interaction_model"
def test_is_event_in_branch_no_branch(self):
"""Test branch filtering with no current branch."""
processor = interactions_processor.InteractionsRequestProcessor()
# Event without branch should be included when no current branch
event = Event(
invocation_id="inv1",
author="test",
content=types.ModelContent("test"),
)
assert processor._is_event_in_branch(None, event) is True
# Event with branch should be excluded when no current branch
event_with_branch = Event(
invocation_id="inv2",
author="test",
content=types.ModelContent("test"),
branch="some_branch",
)
assert processor._is_event_in_branch(None, event_with_branch) is False
def test_is_event_in_branch_same_branch(self):
"""Test that events in the same branch are included."""
processor = interactions_processor.InteractionsRequestProcessor()
event = Event(
invocation_id="inv1",
author="test",
content=types.ModelContent("test"),
branch="root.child",
)
assert processor._is_event_in_branch("root.child", event) is True
def test_is_event_in_branch_different_branch(self):
"""Test that events in different branches are excluded."""
processor = interactions_processor.InteractionsRequestProcessor()
event = Event(
invocation_id="inv1",
author="test",
content=types.ModelContent("test"),
branch="root.other",
)
assert processor._is_event_in_branch("root.child", event) is False
def test_is_event_in_branch_root_events_included(self):
"""Test that root events (no branch) are included in child branches."""
processor = interactions_processor.InteractionsRequestProcessor()
event = Event(
invocation_id="inv1",
author="test",
content=types.ModelContent("test"),
)
assert processor._is_event_in_branch("root.child", event) is True
File diff suppressed because it is too large Load Diff