Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions .gemini/styleguide.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
# Gemini Code Assist Customization

_Placeholder - customization for Gemini Code Assist in this repository is coming
soon._
4 changes: 4 additions & 0 deletions AGENTS.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
# AGENTS.md

_Placeholder - guidance for AI coding agents working in this repository is coming
soon._
1 change: 1 addition & 0 deletions CLAUDE.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
See [AGENTS.md](./AGENTS.md) for project context, commands, and contribution guidelines for AI coding agents.
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,9 @@
import com.google.adk.utils.Constants;
import com.google.api.core.ApiFuture;
import com.google.api.core.ApiFutures;
import com.google.api.gax.rpc.AlreadyExistsException;
import com.google.cloud.firestore.CollectionReference;
import com.google.cloud.firestore.DocumentReference;
import com.google.cloud.firestore.DocumentSnapshot;
import com.google.cloud.firestore.Firestore;
import com.google.cloud.firestore.Query;
Expand All @@ -48,9 +50,10 @@
import java.util.UUID;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.ConcurrentMap;
import java.util.concurrent.ExecutionException;
import java.util.concurrent.atomic.AtomicBoolean;
import java.util.regex.Matcher;
import javax.annotation.Nullable;
import org.jspecify.annotations.Nullable;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;

Expand All @@ -73,6 +76,9 @@ public class FirestoreSessionService implements BaseSessionService {
private static final String UPDATE_TIME_KEY = Constants.KEY_UPDATE_TIME;
private static final String TIMESTAMP_KEY = Constants.KEY_TIMESTAMP;

/** Random token each create stores, so a retried create can detect its own write. */
static final String CREATE_TOKEN_KEY = "createToken";

/** Constructor for FirestoreSessionService. */
public FirestoreSessionService(Firestore firestore) {
this.firestore = firestore;
Expand All @@ -96,7 +102,10 @@ public Single<Session> createSession(
return createSession(appName, userId, (Map<String, Object>) state, sessionId);
}

/** Creates a new session in Firestore. */
/**
* Creates a new session in Firestore. Session IDs are unique per user across apps, so creating
* one the user already has under any app fails with {@link SessionException}.
*/
@Override
public Single<Session> createSession(
String appName,
Expand Down Expand Up @@ -139,11 +148,22 @@ public Single<Session> createSession(
sessionData.put(USER_ID_KEY, newSession.userId());
sessionData.put(UPDATE_TIME_KEY, newSession.lastUpdateTime().toString());
sessionData.put(STATE_KEY, newSession.state());

// Asynchronously write to Firestore and wait for the result
ApiFuture<WriteResult> future =
getSessionsCollection(userId).document(resolvedSessionId).set(sessionData);
future.get(); // Block until the write is complete
String createToken = UUID.randomUUID().toString();
sessionData.put(CREATE_TOKEN_KEY, createToken);

// Unlike set(), create() fails if the session already exists instead of replacing it.
DocumentReference sessionDoc = getSessionsCollection(userId).document(resolvedSessionId);
try {
sessionDoc.create(sessionData).get();
} catch (ExecutionException e) {
if (!(e.getCause() instanceof AlreadyExistsException)) {
throw e;
}
// A retry after a lost reply fails on this call's own write; its token means success.
if (!createToken.equals(sessionDoc.get().get().getString(CREATE_TOKEN_KEY))) {
throw new SessionException(SessionException.SESSION_ALREADY_EXISTS, e.getCause());
}
}

return newSession;
});
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -175,7 +175,8 @@ void run_withUserInput_createsSessionAndExecutesAgent() throws Exception {
when(mockUserDocRef.collection(anyString())).thenReturn(mockSessionsCollection);
when(mockSessionsCollection.document(anyString())).thenReturn(mockSessionDocRef);

when(mockSessionDocRef.set(anyMap())).thenReturn(ApiFutures.immediateFuture(mockWriteResult));
when(mockSessionDocRef.create(anyMap()))
.thenReturn(ApiFutures.immediateFuture(mockWriteResult));
when(mockSessionDocRef.update(anyMap()))
.thenReturn(ApiFutures.immediateFuture(mockWriteResult));
// Mock the event sub-collection chain
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,10 @@
import com.google.adk.events.EventActions;
import com.google.adk.utils.Constants;
import com.google.api.core.ApiFutures;
import com.google.api.gax.grpc.GrpcStatusCode;
import com.google.api.gax.rpc.AlreadyExistsException;
import com.google.api.gax.rpc.PermissionDeniedException;
import com.google.api.gax.rpc.UnavailableException;
import com.google.cloud.firestore.CollectionReference;
import com.google.cloud.firestore.DocumentReference;
import com.google.cloud.firestore.DocumentSnapshot;
Expand All @@ -45,17 +49,20 @@
import com.google.common.collect.ImmutableMap;
import com.google.genai.types.Content;
import com.google.genai.types.Part;
import io.grpc.Status;
import io.reactivex.rxjava3.observers.TestObserver;
import java.time.Instant;
import java.util.Collections;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.concurrent.ConcurrentHashMap;
import java.util.concurrent.ExecutionException;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.extension.ExtendWith;
import org.mockito.ArgumentCaptor;
import org.mockito.Captor;
import org.mockito.Mock;
import org.mockito.junit.jupiter.MockitoExtension;

Expand Down Expand Up @@ -89,6 +96,7 @@ public class FirestoreSessionServiceTest {
@Mock private QuerySnapshot mockQuerySnapshot;
@Mock private WriteResult mockWriteResult;
@Mock private WriteBatch mockWriteBatch;
@Captor private ArgumentCaptor<Map<String, Object>> sessionDataCaptor;

private FirestoreSessionService sessionService;

Expand Down Expand Up @@ -130,7 +138,7 @@ public void setup() {

// Default mock for writes
lenient()
.when(mockSessionDocRef.set(anyMap()))
.when(mockSessionDocRef.create(anyMap()))
.thenReturn(ApiFutures.immediateFuture(mockWriteResult));
lenient()
.when(mockSessionDocRef.update(anyMap()))
Expand Down Expand Up @@ -307,7 +315,7 @@ void createSession_withSessionId_returnsNewSession() {
assertThat(session.id()).isEqualTo(SESSION_ID);
return true;
});
verify(mockSessionDocRef).set(anyMap());
verify(mockSessionDocRef).create(anyMap());
}

/** Tests that createSession creates a new session with a generated session ID. */
Expand All @@ -334,7 +342,7 @@ void createSession_withNullSessionId_generatesNewId() {
assertThat(session.id()).isNotEmpty();
return true;
});
verify(mockSessionDocRef).set(anyMap());
verify(mockSessionDocRef).create(anyMap());
}

/** Tests that createSession creates a new session with an empty session ID. */
Expand All @@ -358,7 +366,7 @@ void createSession_withEmptySessionId_generatesNewId() {
assertThat(session.id()).isNotEqualTo(" ");
return true;
});
verify(mockSessionDocRef).set(anyMap());
verify(mockSessionDocRef).create(anyMap());
}

/** Tests that createSession creates a new session with an empty state when null state is */
Expand Down Expand Up @@ -391,6 +399,114 @@ void createSession_withNullAppName_throwsNullPointerException() {
.assertError(NullPointerException.class);
}

/** Tests that createSession rejects a session ID that is already taken. */
@Test
void createSession_withSessionIdAlreadyTaken_failsWithSessionException() {
// Arrange
when(mockSessionsCollection.document(SESSION_ID)).thenReturn(mockSessionDocRef);
when(mockSessionDocRef.create(anyMap()))
.thenReturn(ApiFutures.immediateFailedFuture(alreadyExists()));
when(mockSessionDocRef.get()).thenReturn(ApiFutures.immediateFuture(mockSessionSnapshot));
// Another caller's session, or one written before sessions carried a token.
when(mockSessionSnapshot.getString(FirestoreSessionService.CREATE_TOKEN_KEY)).thenReturn(null);

// Act
TestObserver<Session> testObserver =
sessionService.createSession(APP_NAME, USER_ID, null, SESSION_ID).test();

// Assert
testObserver.assertError(
e -> {
assertThat(e).isInstanceOf(SessionException.class);
assertThat(e).hasMessageThat().isEqualTo(SessionException.SESSION_ALREADY_EXISTS);
assertThat(e).hasCauseThat().isInstanceOf(AlreadyExistsException.class);
return true;
});
}

/** Tests that createSession succeeds when a client retry fails on the session's own write. */
@Test
void createSession_whenRetryHitsItsOwnWrite_returnsSession() {
// Arrange
when(mockSessionsCollection.document(SESSION_ID)).thenReturn(mockSessionDocRef);
when(mockSessionDocRef.create(sessionDataCaptor.capture()))
.thenReturn(ApiFutures.immediateFailedFuture(alreadyExists()));
when(mockSessionDocRef.get()).thenReturn(ApiFutures.immediateFuture(mockSessionSnapshot));
when(mockSessionSnapshot.getString(FirestoreSessionService.CREATE_TOKEN_KEY))
.thenAnswer(
unused -> sessionDataCaptor.getValue().get(FirestoreSessionService.CREATE_TOKEN_KEY));

// Act
TestObserver<Session> testObserver =
sessionService.createSession(APP_NAME, USER_ID, null, SESSION_ID).test();

// Assert
testObserver.assertComplete();
testObserver.assertValue(session -> session.id().equals(SESSION_ID));
}

/**
* Tests that a failed read-back after a duplicate is propagated, since it cannot tell whose
* session exists.
*/
@Test
void createSession_whenReadBackFails_propagatesFailure() {
// Arrange
when(mockSessionsCollection.document(SESSION_ID)).thenReturn(mockSessionDocRef);
when(mockSessionDocRef.create(anyMap()))
.thenReturn(ApiFutures.immediateFailedFuture(alreadyExists()));
when(mockSessionDocRef.get())
.thenReturn(
ApiFutures.immediateFailedFuture(
new UnavailableException(
"Service unavailable",
/* cause= */ null,
GrpcStatusCode.of(Status.Code.UNAVAILABLE),
/* retryable= */ true)));

// Act
TestObserver<Session> testObserver =
sessionService.createSession(APP_NAME, USER_ID, null, SESSION_ID).test();

// Assert
testObserver.assertError(
e -> {
assertThat(e).isInstanceOf(ExecutionException.class);
assertThat(e).hasCauseThat().isInstanceOf(UnavailableException.class);
return true;
});
}

/**
* Tests that a write failure unrelated to a duplicate session ID is propagated rather than
* reported as one.
*/
@Test
void createSession_whenWriteFailsForAnotherReason_propagatesFailure() {
// Arrange
when(mockSessionsCollection.document(SESSION_ID)).thenReturn(mockSessionDocRef);
when(mockSessionDocRef.create(anyMap()))
.thenReturn(
ApiFutures.immediateFailedFuture(
new PermissionDeniedException(
"Missing or insufficient permissions",
/* cause= */ null,
GrpcStatusCode.of(Status.Code.PERMISSION_DENIED),
/* retryable= */ false)));

// Act
TestObserver<Session> testObserver =
sessionService.createSession(APP_NAME, USER_ID, null, SESSION_ID).test();

// Assert
testObserver.assertError(
e -> {
assertThat(e).isInstanceOf(ExecutionException.class);
assertThat(e).hasCauseThat().isInstanceOf(PermissionDeniedException.class);
return true;
});
}

// --- appendEvent Tests ---
/** Tests that appendEvent persists the event and updates the session's updateTime. */
@Test
Expand Down Expand Up @@ -978,4 +1094,12 @@ void deleteSession_appNameMismatch_doesNotDelete() {
verify(mockSessionDocRef, never()).delete();
verify(mockWriteBatch, never()).commit();
}

private static AlreadyExistsException alreadyExists() {
return new AlreadyExistsException(
"Document already exists",
/* cause= */ null,
GrpcStatusCode.of(Status.Code.ALREADY_EXISTS),
/* retryable= */ false);
}
}
47 changes: 43 additions & 4 deletions core/src/main/java/com/google/adk/agents/BaseAgent.java
Original file line number Diff line number Diff line change
Expand Up @@ -16,13 +16,15 @@

package com.google.adk.agents;

import static com.google.common.base.Preconditions.checkArgument;
import static com.google.common.base.Strings.isNullOrEmpty;
import static com.google.common.collect.ImmutableList.toImmutableList;
import static java.lang.String.format;

import com.google.adk.agents.Callbacks.AfterAgentCallback;
import com.google.adk.agents.Callbacks.BeforeAgentCallback;
import com.google.adk.events.Event;
import com.google.adk.events.EventActions;
import com.google.adk.plugins.Plugin;
import com.google.adk.telemetry.Instrumentation;
import com.google.adk.telemetry.Instrumentation.AgentInvocation;
Expand All @@ -39,6 +41,7 @@
import java.util.ArrayList;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.function.Function;
import java.util.regex.Pattern;
Expand Down Expand Up @@ -136,10 +139,8 @@ private static void validateAgentName(String name) {
throw new IllegalArgumentException(
format("Agent name '%s' does not match regex '%s'.", name, IDENTIFIER_REGEX));
}
if (name.equals(Role.USER)) {
throw new IllegalArgumentException(
"Agent name cannot be 'user'; reserved for end-user input.");
}
checkArgument(
!name.equals(Role.USER), "Agent name cannot be 'user'; reserved for end-user input.");
}

/**
Expand Down Expand Up @@ -468,6 +469,44 @@ public Flowable<Event> runLive(InvocationContext parentContext) {
return run(parentContext, this::runLiveImpl);
}

/**
* Records this agent's end-of-agent checkpoint and returns it as a single-event stream. The
* recorded state is cleared and the agent marked finished, so a later run skips it.
*
* @param context Current invocation context.
* @return stream of the single {@code endOfAgent = true} checkpoint event.
*/
final Flowable<Event> endOfAgentAndRecord(InvocationContext context) {
context.setAgentState(name(), /* agentState= */ null, /* endOfAgent= */ true);
return Flowable.just(checkpointEvent(context, EventActions.builder().endOfAgent(true).build()));
}

/**
* Records {@code agentState} for this agent and returns the matching checkpoint event as a
* single-event stream. The agent is left unfinished, so a later run resumes from this checkpoint.
*
* @param context Current invocation context.
* @param agentState The serialized agent state to persist.
* @return stream of the single checkpoint event carrying {@code agentState}.
*/
final Flowable<Event> checkpointAndRecord(
InvocationContext context, Map<String, Object> agentState) {
context.setAgentState(name(), agentState, /* endOfAgent= */ false);
return Flowable.just(
checkpointEvent(context, EventActions.builder().agentState(agentState).build()));
}

/** Builds a resumability checkpoint event authored by this agent carrying {@code actions}. */
private Event checkpointEvent(InvocationContext context, EventActions actions) {
return Event.builder()
.id(Event.generateEventId())
.invocationId(context.invocationId())
.author(name())
.branch(context.branch().orElse(null))
.actions(actions)
.build();
}

/**
* Agent-specific asynchronous logic.
*
Expand Down
Loading
Loading