From 1309d3841b5609943aa5c88b8b8d4b986646e9c6 Mon Sep 17 00:00:00 2001 From: Mateusz Krawiec Date: Thu, 8 Oct 2026 04:28:04 -0700 Subject: [PATCH] feat: add Context as the common base of CallbackContext and ToolContext PiperOrigin-RevId: 995744753 --- .../google/adk/agents/CallbackContext.java | 129 ++-------- .../java/com/google/adk/agents/Context.java | 226 ++++++++++++++++++ .../com/google/adk/tools/ToolContext.java | 84 +------ .../com/google/adk/agents/ContextTest.java | 179 ++++++++++++++ 4 files changed, 435 insertions(+), 183 deletions(-) create mode 100644 core/src/main/java/com/google/adk/agents/Context.java create mode 100644 core/src/test/java/com/google/adk/agents/ContextTest.java diff --git a/core/src/main/java/com/google/adk/agents/CallbackContext.java b/core/src/main/java/com/google/adk/agents/CallbackContext.java index da5b0d794..49d5addb5 100644 --- a/core/src/main/java/com/google/adk/agents/CallbackContext.java +++ b/core/src/main/java/com/google/adk/agents/CallbackContext.java @@ -16,133 +16,38 @@ package com.google.adk.agents; -import com.google.adk.artifacts.ListArtifactsResponse; import com.google.adk.events.EventActions; -import com.google.adk.sessions.State; -import com.google.common.base.Preconditions; -import com.google.genai.types.Part; -import io.reactivex.rxjava3.core.Completable; -import io.reactivex.rxjava3.core.Maybe; -import io.reactivex.rxjava3.core.Single; -import java.util.List; +import org.jspecify.annotations.Nullable; -/** The context of various callbacks for an agent invocation. */ -public class CallbackContext extends ReadonlyContext { - - protected EventActions eventActions; - private final State state; - private final String eventId; +/** + * The context of various callbacks for an agent invocation. + * + *

Extends {@link Context} for backward compatibility; agent and model callback signatures still + * use this type. + */ +public class CallbackContext extends Context { /** * Initializes callback context. * * @param invocationContext Current invocation context. - * @param eventActions Callback event actions. + * @param eventActions Callback event actions, or null for new empty ones. */ - public CallbackContext(InvocationContext invocationContext, EventActions eventActions) { - this(invocationContext, eventActions, null); + public CallbackContext(InvocationContext invocationContext, @Nullable EventActions eventActions) { + super(invocationContext, eventActions, /* eventId= */ null); } /** * Initializes callback context. * * @param invocationContext Current invocation context. - * @param eventActions Callback event actions. - * @param eventId The ID of the event associated with this context. + * @param eventActions Callback event actions, or null for new empty ones. + * @param eventId The ID of the event associated with this context, or null if there is none. */ public CallbackContext( - InvocationContext invocationContext, EventActions eventActions, String eventId) { - super(invocationContext); - this.eventActions = eventActions != null ? eventActions : EventActions.builder().build(); - this.state = new State(invocationContext.session().state(), this.eventActions.stateDelta()); - this.eventId = eventId; - } - - /** Returns the delta-aware state of the current callback. */ - @Override - public State state() { - return state; - } - - /** Returns the EventActions associated with this context. */ - public EventActions eventActions() { - return eventActions; - } - - /** Returns the ID of the event associated with this context. */ - public String eventId() { - return eventId; - } - - /** - * Lists the filenames of the artifacts attached to the current session. - * - * @return the list of artifact filenames - */ - public Single> listArtifacts() { - if (invocationContext.artifactService() == null) { - throw new IllegalStateException("Artifact service is not initialized."); - } - return invocationContext - .artifactService() - .listArtifactKeys( - invocationContext.session().appName(), - invocationContext.session().userId(), - invocationContext.session().id()) - .map(ListArtifactsResponse::filenames); - } - - /** Loads the latest version of an artifact from the service. */ - public Maybe loadArtifact(String filename) { - checkArtifactServiceInitialized(); - return invocationContext - .artifactService() - .loadArtifact( - invocationContext.appName(), - invocationContext.userId(), - invocationContext.session().id(), - filename); - } - - /** Loads a specific version of an artifact from the service. */ - public Maybe loadArtifact(String filename, int version) { - checkArtifactServiceInitialized(); - return invocationContext - .artifactService() - .loadArtifact( - invocationContext.appName(), - invocationContext.userId(), - invocationContext.session().id(), - filename, - version); - } - - private void checkArtifactServiceInitialized() { - Preconditions.checkState( - invocationContext.artifactService() != null, "Artifact service is not initialized."); - } - - /** - * Saves an artifact and records it as a delta for the current session. - * - * @param filename Artifact file name. - * @param artifact Artifact content to save. - * @return a {@link Completable} that completes when the artifact is saved. - * @throws IllegalStateException if the artifact service is not initialized. - */ - public Completable saveArtifact(String filename, Part artifact) { - if (invocationContext.artifactService() == null) { - throw new IllegalStateException("Artifact service is not initialized."); - } - return invocationContext - .artifactService() - .saveArtifact( - invocationContext.appName(), - invocationContext.userId(), - invocationContext.session().id(), - filename, - artifact) - .doOnSuccess(version -> this.eventActions.artifactDelta().put(filename, version)) - .ignoreElement(); + InvocationContext invocationContext, + @Nullable EventActions eventActions, + @Nullable String eventId) { + super(invocationContext, eventActions, eventId); } } diff --git a/core/src/main/java/com/google/adk/agents/Context.java b/core/src/main/java/com/google/adk/agents/Context.java new file mode 100644 index 000000000..0432c97c9 --- /dev/null +++ b/core/src/main/java/com/google/adk/agents/Context.java @@ -0,0 +1,226 @@ +/* + * 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. + */ + +package com.google.adk.agents; + +import com.google.adk.artifacts.ListArtifactsResponse; +import com.google.adk.events.EventActions; +import com.google.adk.events.ToolConfirmation; +import com.google.adk.memory.SearchMemoryResponse; +import com.google.adk.sessions.State; +import com.google.common.base.Preconditions; +import com.google.genai.types.Part; +import io.reactivex.rxjava3.core.Completable; +import io.reactivex.rxjava3.core.Maybe; +import io.reactivex.rxjava3.core.Single; +import java.util.List; +import java.util.Optional; +import org.jspecify.annotations.Nullable; + +/** + * The context passed to callbacks and tools during an invocation. + * + *

Agent and model callbacks receive a {@link CallbackContext}, and tools and tool callbacks a + * {@link com.google.adk.tools.ToolContext}; both extend this class, so a helper that serves both, + * or a function tool's {@code toolContext} parameter, can take a {@code Context}. {@link + * #requestConfirmation()} needs a function call ID, which only a tool call has. + */ +public class Context extends ReadonlyContext { + + /** The event actions that record this context's changes. */ + protected EventActions eventActions; + + private final State state; + private final @Nullable String eventId; + private Optional functionCallId = Optional.empty(); + private Optional toolConfirmation = Optional.empty(); + + /** + * Initializes a context. + * + * @param invocationContext Current invocation context. + * @param eventActions Event actions to record changes in, or null for new empty ones. + * @param eventId The ID of the event associated with this context, or null if there is none. + */ + Context( + InvocationContext invocationContext, + @Nullable EventActions eventActions, + @Nullable String eventId) { + super(invocationContext); + this.eventActions = eventActions != null ? eventActions : EventActions.builder().build(); + this.state = new State(invocationContext.session().state(), this.eventActions.stateDelta()); + this.eventId = eventId; + } + + /** Returns the delta-aware state of the current context. */ + @Override + public State state() { + return state; + } + + /** Returns the {@link EventActions} associated with this context. */ + public EventActions eventActions() { + return eventActions; + } + + /** Returns the same {@link EventActions} as {@link #eventActions()}. */ + public EventActions actions() { + return this.eventActions; + } + + /** Returns the ID of the event associated with this context, or null if there is none. */ + public String eventId() { + return eventId; + } + + /** Returns the ID of the function call that invoked the current tool, if any. */ + public Optional functionCallId() { + return functionCallId; + } + + /** Sets the ID of the function call that invoked the current tool, or clears it if null. */ + public void functionCallId(@Nullable String functionCallId) { + this.functionCallId = Optional.ofNullable(functionCallId); + } + + /** Returns the confirmation of the current tool call, if any. */ + public Optional toolConfirmation() { + return toolConfirmation; + } + + /** Sets the confirmation of the current tool call, or clears it if null. */ + public void toolConfirmation(@Nullable ToolConfirmation toolConfirmation) { + this.toolConfirmation = Optional.ofNullable(toolConfirmation); + } + + /** + * Lists the filenames of the artifacts attached to the current session. + * + * @return the list of artifact filenames + */ + public Single> listArtifacts() { + if (invocationContext.artifactService() == null) { + throw new IllegalStateException("Artifact service is not initialized."); + } + return invocationContext + .artifactService() + .listArtifactKeys( + invocationContext.session().appName(), + invocationContext.session().userId(), + invocationContext.session().id()) + .map(ListArtifactsResponse::filenames); + } + + /** Loads the latest version of an artifact from the service. */ + public Maybe loadArtifact(String filename) { + checkArtifactServiceInitialized(); + return invocationContext + .artifactService() + .loadArtifact( + invocationContext.appName(), + invocationContext.userId(), + invocationContext.session().id(), + filename); + } + + /** Loads a specific version of an artifact from the service. */ + public Maybe loadArtifact(String filename, int version) { + checkArtifactServiceInitialized(); + return invocationContext + .artifactService() + .loadArtifact( + invocationContext.appName(), + invocationContext.userId(), + invocationContext.session().id(), + filename, + version); + } + + private void checkArtifactServiceInitialized() { + Preconditions.checkState( + invocationContext.artifactService() != null, "Artifact service is not initialized."); + } + + /** + * Saves an artifact and records it as a delta for the current session. + * + * @param filename Artifact file name. + * @param artifact Artifact content to save. + * @return a {@link Completable} that completes when the artifact is saved. + * @throws IllegalStateException if the artifact service is not initialized. + */ + public Completable saveArtifact(String filename, Part artifact) { + if (invocationContext.artifactService() == null) { + throw new IllegalStateException("Artifact service is not initialized."); + } + return invocationContext + .artifactService() + .saveArtifact( + invocationContext.appName(), + invocationContext.userId(), + invocationContext.session().id(), + filename, + artifact) + .doOnSuccess(version -> this.eventActions.artifactDelta().put(filename, version)) + .ignoreElement(); + } + + /** + * Requests confirmation for the current function call. + * + * @param hint A hint to the user on how to confirm the tool call. + * @param payload The payload used to confirm the tool call. + * @throws IllegalStateException if this context has no function call ID. + */ + public void requestConfirmation(@Nullable String hint, @Nullable Object payload) { + if (functionCallId.isEmpty()) { + throw new IllegalStateException("function_call_id is not set."); + } + this.eventActions + .requestedToolConfirmations() + .put(functionCallId.get(), ToolConfirmation.builder().hint(hint).payload(payload).build()); + } + + /** + * Requests confirmation for the current function call. + * + * @param hint A hint to the user on how to confirm the tool call. + * @throws IllegalStateException if this context has no function call ID. + */ + public void requestConfirmation(@Nullable String hint) { + requestConfirmation(hint, null); + } + + /** + * Requests confirmation for the current function call. + * + * @throws IllegalStateException if this context has no function call ID. + */ + public void requestConfirmation() { + requestConfirmation(null, null); + } + + /** Searches the memory of the current user. */ + public Single searchMemory(String query) { + if (invocationContext.memoryService() == null) { + throw new IllegalStateException("Memory service is not initialized."); + } + return invocationContext + .memoryService() + .searchMemory( + invocationContext.session().appName(), invocationContext.session().userId(), query); + } +} diff --git a/core/src/main/java/com/google/adk/tools/ToolContext.java b/core/src/main/java/com/google/adk/tools/ToolContext.java index 9e0465145..13de09ba6 100644 --- a/core/src/main/java/com/google/adk/tools/ToolContext.java +++ b/core/src/main/java/com/google/adk/tools/ToolContext.java @@ -17,19 +17,21 @@ package com.google.adk.tools; import com.google.adk.agents.CallbackContext; +import com.google.adk.agents.Context; import com.google.adk.agents.InvocationContext; import com.google.adk.events.EventActions; import com.google.adk.events.ToolConfirmation; -import com.google.adk.memory.SearchMemoryResponse; import com.google.errorprone.annotations.CanIgnoreReturnValue; -import io.reactivex.rxjava3.core.Single; import java.util.Optional; import org.jspecify.annotations.Nullable; -/** ToolContext object provides a structured context for executing tools or functions. */ +/** + * ToolContext object provides a structured context for executing tools or functions. + * + *

Extends {@link CallbackContext} (and through it {@link Context}); {@link BaseTool#runAsync} + * and tool callbacks take this type. + */ public class ToolContext extends CallbackContext { - private Optional functionCallId = Optional.empty(); - private Optional toolConfirmation = Optional.empty(); private ToolContext( InvocationContext invocationContext, @@ -38,34 +40,14 @@ private ToolContext( Optional toolConfirmation, @Nullable String eventId) { super(invocationContext, eventActions, eventId); - this.functionCallId = functionCallId; - this.toolConfirmation = toolConfirmation; - } - - public EventActions actions() { - return this.eventActions; + functionCallId(functionCallId.orElse(null)); + toolConfirmation(toolConfirmation.orElse(null)); } public void setActions(EventActions actions) { this.eventActions = actions; } - public Optional functionCallId() { - return functionCallId; - } - - public void functionCallId(String functionCallId) { - this.functionCallId = Optional.ofNullable(functionCallId); - } - - public Optional toolConfirmation() { - return toolConfirmation; - } - - public void toolConfirmation(ToolConfirmation toolConfirmation) { - this.toolConfirmation = Optional.ofNullable(toolConfirmation); - } - @SuppressWarnings("unused") private void requestCredential() { throw new UnsupportedOperationException("Credential request not implemented yet."); @@ -76,46 +58,6 @@ private void getAuthResponse() { throw new UnsupportedOperationException("Auth response retrieval not implemented yet."); } - /** - * Requests confirmation for the given function call. - * - * @param hint A hint to the user on how to confirm the tool call. - * @param payload The payload used to confirm the tool call. - */ - public void requestConfirmation(@Nullable String hint, @Nullable Object payload) { - if (functionCallId.isEmpty()) { - throw new IllegalStateException("function_call_id is not set."); - } - this.eventActions - .requestedToolConfirmations() - .put(functionCallId.get(), ToolConfirmation.builder().hint(hint).payload(payload).build()); - } - - /** - * Requests confirmation for the given function call. - * - * @param hint A hint to the user on how to confirm the tool call. - */ - public void requestConfirmation(@Nullable String hint) { - requestConfirmation(hint, null); - } - - /** Requests confirmation for the given function call. */ - public void requestConfirmation() { - requestConfirmation(null, null); - } - - /** Searches the memory of the current user. */ - public Single searchMemory(String query) { - if (invocationContext.memoryService() == null) { - throw new IllegalStateException("Memory service is not initialized."); - } - return invocationContext - .memoryService() - .searchMemory( - invocationContext.session().appName(), invocationContext.session().userId(), query); - } - public static Builder builder(InvocationContext invocationContext) { return new Builder(invocationContext); } @@ -123,8 +65,8 @@ public static Builder builder(InvocationContext invocationContext) { public Builder toBuilder() { return new Builder(invocationContext) .actions(eventActions) - .functionCallId(functionCallId.orElse(null)) - .toolConfirmation(toolConfirmation.orElse(null)) + .functionCallId(functionCallId().orElse(null)) + .toolConfirmation(toolConfirmation().orElse(null)) .eventId(eventId()); } @@ -136,9 +78,9 @@ public String toString() { + ", eventActions=" + eventActions + ", functionCallId=" - + functionCallId + + functionCallId() + ", toolConfirmation=" - + toolConfirmation + + toolConfirmation() + '}'; } diff --git a/core/src/test/java/com/google/adk/agents/ContextTest.java b/core/src/test/java/com/google/adk/agents/ContextTest.java new file mode 100644 index 000000000..6ff2b4365 --- /dev/null +++ b/core/src/test/java/com/google/adk/agents/ContextTest.java @@ -0,0 +1,179 @@ +/* + * 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. + */ + +package com.google.adk.agents; + +import static com.google.adk.testing.TestUtils.createFunctionCallLlmResponse; +import static com.google.adk.testing.TestUtils.createInvocationContext; +import static com.google.adk.testing.TestUtils.createTestAgent; +import static com.google.adk.testing.TestUtils.createTestAgentBuilder; +import static com.google.adk.testing.TestUtils.createTestLlm; +import static com.google.adk.testing.TestUtils.createTextLlmResponse; +import static com.google.common.truth.Truth.assertThat; +import static org.junit.Assert.assertThrows; + +import com.google.adk.events.Event; +import com.google.adk.memory.InMemoryMemoryService; +import com.google.adk.memory.SearchMemoryResponse; +import com.google.adk.runner.Runner; +import com.google.adk.sessions.Session; +import com.google.adk.sessions.State; +import com.google.adk.tools.FunctionTool; +import com.google.adk.tools.ToolContext; +import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableMap; +import com.google.genai.types.Content; +import com.google.genai.types.FunctionResponse; +import com.google.genai.types.Part; +import io.reactivex.rxjava3.core.Maybe; +import java.util.Optional; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.junit.runners.JUnit4; + +/** Tests for {@link Context}. */ +@RunWith(JUnit4.class) +public final class ContextTest { + + @Test + public void toolContext_usedAsCallbackContextOrContext_keepsFunctionCallId() { + ToolContext toolContext = + ToolContext.builder(newInvocationContext()).functionCallId("call-1").build(); + + // Compile-time guard: code written against CallbackContext keeps accepting a ToolContext. + CallbackContext callbackContext = toolContext; + Context context = callbackContext; + + assertThat(context.functionCallId()).hasValue("call-1"); + } + + @Test + public void callbackContextSubclassOverridingState_recordsWritesInEventActions() { + LegacyCallbackContext context = new LegacyCallbackContext(newInvocationContext()); + + context.state().put("color", "blue"); + + assertThat(context.eventActions().stateDelta()).containsExactly("color", "blue"); + } + + @Test + public void requestConfirmation_onCallbackContext_throwsIllegalStateException() { + CallbackContext callbackContext = new CallbackContext(newInvocationContext(), null); + + assertThrows(IllegalStateException.class, callbackContext::requestConfirmation); + } + + @Test + public void searchMemory_onCallbackContext_findsRememberedEvents() { + InvocationContext invocationContext = newInvocationContext(); + Event rememberedEvent = + Event.builder() + .id("e1") + .author("user") + .content(Content.fromParts(Part.fromText("My favorite color is teal."))) + .build(); + InMemoryMemoryService memoryService = new InMemoryMemoryService(); + memoryService + .addSessionToMemory( + Session.builder("earlier") + .appName(invocationContext.appName()) + .userId(invocationContext.userId()) + .events(ImmutableList.of(rememberedEvent)) + .build()) + .blockingAwait(); + CallbackContext callbackContext = + new CallbackContext( + invocationContext.toBuilder().memoryService(memoryService).build(), null); + + SearchMemoryResponse response = callbackContext.searchMemory("teal").blockingGet(); + + assertThat(response.memories()).hasSize(1); + } + + @Test + public void runAsync_functionToolWithContextParameter_receivesTheToolContext() { + LlmAgent agent = + createTestAgentBuilder( + createTestLlm( + createFunctionCallLlmResponse( + "call-1", "rememberColor", ImmutableMap.of("color", "blue")), + createTextLlmResponse("done"))) + .tools(FunctionTool.create(ContextTest.class, "rememberColor")) + .build(); + + Session session = runOnce(agent); + + assertThat(session.state()).containsEntry("color", "blue"); + FunctionResponse functionResponse = + session.immutableEvents().stream() + .flatMap(e -> e.functionResponses().stream()) + .findFirst() + .orElseThrow(); + assertThat(functionResponse.response().orElseThrow()).containsEntry("functionCallId", "call-1"); + } + + @Test + public void runAsync_agentCallbackTakingContext_persistsItsStateWrite() { + LlmAgent agent = + createTestAgentBuilder(createTestLlm(createTextLlmResponse("done"))) + .beforeAgentCallback(ContextTest::recordAgentName) + .build(); + + Session session = runOnce(agent); + + assertThat(session.state()).containsEntry("before_agent", agent.name()); + } + + // FunctionTool needs a public method; it injects the context by the parameter name toolContext. + public static ImmutableMap rememberColor(String color, Context toolContext) { + toolContext.state().put("color", color); + return ImmutableMap.of("functionCallId", toolContext.functionCallId().orElse("")); + } + + private static Maybe recordAgentName(Context context) { + context.state().put("before_agent", context.agentName()); + return Maybe.empty(); + } + + private static InvocationContext newInvocationContext() { + return createInvocationContext(createTestAgent(createTestLlm(createTextLlmResponse("unused")))); + } + + private static Session runOnce(LlmAgent agent) { + Runner runner = Runner.builder().agent(agent).appName("test_app").build(); + Session session = runner.sessionService().createSession(runner.appName(), "user").blockingGet(); + runner + .runAsync("user", session.id(), Content.fromParts(Part.fromText("hi"))) + .blockingSubscribe(); + return runner + .sessionService() + .getSession(runner.appName(), "user", session.id(), Optional.empty()) + .blockingGet(); + } + + /** A subclass in the style of code written before Context existed. */ + private static final class LegacyCallbackContext extends CallbackContext { + LegacyCallbackContext(InvocationContext invocationContext) { + super(invocationContext, /* eventActions= */ null); + } + + /** Compile-time guard: this override stops compiling if {@code Context.state()} turns final. */ + @Override + public State state() { + return super.state(); + } + } +}