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();
+ }
+ }
+}