diff --git a/core/src/main/java/com/google/adk/agents/ContextCacheConfig.java b/core/src/main/java/com/google/adk/agents/ContextCacheConfig.java index 084700d54..49e73401f 100644 --- a/core/src/main/java/com/google/adk/agents/ContextCacheConfig.java +++ b/core/src/main/java/com/google/adk/agents/ContextCacheConfig.java @@ -15,7 +15,10 @@ */ package com.google.adk.agents; +import com.google.errorprone.annotations.InlineMe; +import com.google.genai.types.HttpOptions; import java.time.Duration; +import org.jspecify.annotations.Nullable; /** * Configuration for context caching across all agents in an app. @@ -27,33 +30,55 @@ *

Context caching can significantly reduce costs and improve response times by reusing * previously processed context across multiple requests. * - * @param maxInvocations Maximum number of invocations to reuse the same cache before refreshing it. + * @param cacheIntervals Maximum number of invocations to reuse the same cache before refreshing it. * Defaults to 10. * @param ttl Time-to-live for cache. Defaults to 1800 seconds (30 minutes). * @param minTokens Minimum estimated request tokens required to enable caching. This compares * against the estimated total tokens of the request (system instruction + tools + contents). * Context cache storage may have cost. Set higher to avoid caching small requests where * overhead may exceed benefits. Defaults to 0. + * @param createHttpOptions HTTP options, such as a timeout, for the call that creates a cache; null + * uses the client's defaults. Defaults to null. */ -public record ContextCacheConfig(int maxInvocations, Duration ttl, int minTokens) { +public record ContextCacheConfig( + int cacheIntervals, Duration ttl, int minTokens, @Nullable HttpOptions createHttpOptions) { public ContextCacheConfig() { - this(10, Duration.ofSeconds(1800), 0); + this(10, Duration.ofMinutes(30), 0); + } + + /** Creates a config that creates caches with the client's default HTTP options. */ + public ContextCacheConfig(int cacheIntervals, Duration ttl, int minTokens) { + this(cacheIntervals, ttl, minTokens, /* createHttpOptions= */ null); + } + + /** + * Returns {@link #cacheIntervals()}. + * + * @deprecated Use {@link #cacheIntervals()}, the name ADK Python and ADK Kotlin use. + */ + @Deprecated + @InlineMe(replacement = "this.cacheIntervals()") + public int maxInvocations() { + return cacheIntervals(); } /** Returns TTL as string format for cache creation. */ public String getTtlString() { - return ttl.getSeconds() + "s"; + return ttl.toSeconds() + "s"; } @Override public String toString() { - return "ContextCacheConfig(maxInvocations=" - + maxInvocations + // Says only whether HTTP options are set, since their headers can carry credentials. + return "ContextCacheConfig(cacheIntervals=" + + cacheIntervals + ", ttl=" - + ttl.getSeconds() + + ttl.toSeconds() + "s, minTokens=" + minTokens + + ", createHttpOptions=" + + (createHttpOptions == null ? "null" : "set") + ")"; } } diff --git a/core/src/main/java/com/google/adk/events/Event.java b/core/src/main/java/com/google/adk/events/Event.java index 5b63be76f..eb266f459 100644 --- a/core/src/main/java/com/google/adk/events/Event.java +++ b/core/src/main/java/com/google/adk/events/Event.java @@ -24,6 +24,7 @@ import com.fasterxml.jackson.databind.annotation.JsonDeserialize; import com.google.adk.JsonBaseModel; import com.google.adk.annotations.Experimental; +import com.google.adk.models.CacheMetadata; import com.google.adk.platform.UuidProvider; import com.google.adk.workflow.NodeInfo; import com.google.common.collect.ImmutableList; @@ -70,6 +71,7 @@ public class Event extends JsonBaseModel { private @Nullable Transcription outputTranscription; private @Nullable Object output; private @Nullable NodeInfo nodeInfo; + private @Nullable CacheMetadata cacheMetadata; private long timestamp; @@ -342,6 +344,19 @@ public void setNodeInfo(@Nullable NodeInfo nodeInfo) { this.nodeInfo = nodeInfo; } + /** + * Context cache state of the LLM response this event carries. The next request of the same agent + * reads it to reuse or refresh the cache. + */ + @JsonProperty("cacheMetadata") + public Optional cacheMetadata() { + return Optional.ofNullable(cacheMetadata); + } + + public void setCacheMetadata(@Nullable CacheMetadata cacheMetadata) { + this.cacheMetadata = cacheMetadata; + } + /** The timestamp of the event. */ @JsonProperty("timestamp") public long timestamp() { @@ -453,6 +468,7 @@ public static class Builder { private @Nullable Transcription outputTranscription; private @Nullable Object output; private @Nullable NodeInfo nodeInfo; + private @Nullable CacheMetadata cacheMetadata; private @Nullable Long timestamp; @JsonCreator @@ -646,6 +662,13 @@ public Builder nodeInfo(@Nullable NodeInfo value) { return this; } + @CanIgnoreReturnValue + @JsonProperty("cacheMetadata") + public Builder cacheMetadata(@Nullable CacheMetadata value) { + this.cacheMetadata = value; + return this; + } + public Event build() { Event event = new Event(); event.setId(id); @@ -672,6 +695,7 @@ public Event build() { event.setOutputTranscription(outputTranscription); event.setOutput(output); event.setNodeInfo(nodeInfo); + event.setCacheMetadata(cacheMetadata); return event; } } @@ -710,6 +734,7 @@ public Builder toBuilder() { .outputTranscription(this.outputTranscription) .output(this.output) .nodeInfo(this.nodeInfo) + .cacheMetadata(this.cacheMetadata) .timestamp(this.timestamp); } @@ -743,7 +768,8 @@ public boolean equals(Object obj) { && Objects.equals(inputTranscription, other.inputTranscription) && Objects.equals(outputTranscription, other.outputTranscription) && Objects.equals(output, other.output) - && Objects.equals(nodeInfo, other.nodeInfo); + && Objects.equals(nodeInfo, other.nodeInfo) + && Objects.equals(cacheMetadata, other.cacheMetadata); } @Override @@ -776,6 +802,7 @@ public int hashCode() { outputTranscription, output, nodeInfo, + cacheMetadata, timestamp); } } diff --git a/core/src/main/java/com/google/adk/flows/llmflows/BaseLlmFlow.java b/core/src/main/java/com/google/adk/flows/llmflows/BaseLlmFlow.java index b576a299e..2b921ab22 100644 --- a/core/src/main/java/com/google/adk/flows/llmflows/BaseLlmFlow.java +++ b/core/src/main/java/com/google/adk/flows/llmflows/BaseLlmFlow.java @@ -885,7 +885,8 @@ private Event buildModelResponseEvent( .usageMetadata(llmResponse.usageMetadata().orElse(null)) .modelVersion(llmResponse.modelVersion().orElse(null)) .inputTranscription(llmResponse.inputTranscription().orElse(null)) - .outputTranscription(llmResponse.outputTranscription().orElse(null)); + .outputTranscription(llmResponse.outputTranscription().orElse(null)) + .cacheMetadata(llmResponse.cacheMetadata().orElse(null)); Event event = eventBuilder.build(); diff --git a/core/src/main/java/com/google/adk/flows/llmflows/ContextCacheRequestProcessor.java b/core/src/main/java/com/google/adk/flows/llmflows/ContextCacheRequestProcessor.java new file mode 100644 index 000000000..bfef398f1 --- /dev/null +++ b/core/src/main/java/com/google/adk/flows/llmflows/ContextCacheRequestProcessor.java @@ -0,0 +1,101 @@ +/* + * 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.flows.llmflows; + +import com.google.adk.agents.ContextCacheConfig; +import com.google.adk.agents.InvocationContext; +import com.google.adk.events.Event; +import com.google.adk.models.CacheMetadata; +import com.google.adk.models.LlmRequest; +import com.google.common.base.Strings; +import com.google.common.collect.ImmutableList; +import com.google.genai.types.GenerateContentResponseUsageMetadata; +import io.reactivex.rxjava3.core.Single; +import java.util.Optional; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +/** + * {@link RequestProcessor} that enables context caching when the app configures it. It puts the + * config, the agent's latest cache metadata and its previous prompt token count on the request; the + * model creates, reuses and deletes the caches. + */ +final class ContextCacheRequestProcessor implements RequestProcessor { + + private static final Logger logger = LoggerFactory.getLogger(ContextCacheRequestProcessor.class); + + @Override + public Single processRequest( + InvocationContext context, LlmRequest request) { + Optional cacheConfig = context.contextCacheConfig(); + if (cacheConfig.isEmpty()) { + return Single.just(RequestProcessingResult.create(request, ImmutableList.of())); + } + + String agentName = context.agent().name(); + CacheMetadata cacheMetadata = null; + Integer previousTokenCount = null; + for (Event event : context.session().immutableEvents().reverse()) { + if (!agentName.equals(event.author())) { + continue; + } + if (cacheMetadata == null && event.cacheMetadata().isPresent()) { + cacheMetadata = countInvocation(event, context.invocationId()); + } + if (previousTokenCount == null) { + previousTokenCount = + event + .usageMetadata() + .flatMap(GenerateContentResponseUsageMetadata::promptTokenCount) + .orElse(null); + } + if (cacheMetadata != null && previousTokenCount != null) { + break; + } + } + if (cacheMetadata != null) { + logger.debug("Found cache metadata for agent {}: {}", agentName, cacheMetadata); + } + if (previousTokenCount != null) { + logger.debug( + "Found previous prompt token count for agent {}: {}", agentName, previousTokenCount); + } + logger.debug("Context caching enabled for agent {}", agentName); + + LlmRequest updatedRequest = + request.toBuilder() + .cacheConfig(cacheConfig.get()) + .cacheMetadata(cacheMetadata) + .cacheableContentsTokenCount(previousTokenCount) + .build(); + return Single.just(RequestProcessingResult.create(updatedRequest, ImmutableList.of())); + } + + /** + * Returns the event's cache metadata, counting one more use when it names an active cache from an + * earlier invocation. + */ + private static CacheMetadata countInvocation(Event event, String invocationId) { + CacheMetadata metadata = event.cacheMetadata().get(); + if (Strings.isNullOrEmpty(event.invocationId()) + || event.invocationId().equals(invocationId) + || metadata.cacheName().isEmpty()) { + return metadata; + } + return metadata.toBuilder().invocationsUsed(metadata.invocationsUsed().get() + 1).build(); + } +} diff --git a/core/src/main/java/com/google/adk/flows/llmflows/SingleFlow.java b/core/src/main/java/com/google/adk/flows/llmflows/SingleFlow.java index 41dff3b96..cf4382e28 100644 --- a/core/src/main/java/com/google/adk/flows/llmflows/SingleFlow.java +++ b/core/src/main/java/com/google/adk/flows/llmflows/SingleFlow.java @@ -33,6 +33,7 @@ public class SingleFlow extends BaseLlmFlow { new Identity(), new Compaction(), new Contents(), + new ContextCacheRequestProcessor(), CodeExecution.requestProcessor); protected static final ImmutableList RESPONSE_PROCESSORS = diff --git a/core/src/main/java/com/google/adk/models/CacheMetadata.java b/core/src/main/java/com/google/adk/models/CacheMetadata.java new file mode 100644 index 000000000..576714a4b --- /dev/null +++ b/core/src/main/java/com/google/adk/models/CacheMetadata.java @@ -0,0 +1,202 @@ +/* + * 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.models; + +import static com.google.common.base.Preconditions.checkState; + +import com.fasterxml.jackson.annotation.JsonAlias; +import com.fasterxml.jackson.annotation.JsonCreator; +import com.fasterxml.jackson.annotation.JsonProperty; +import com.fasterxml.jackson.core.JsonParser; +import com.fasterxml.jackson.core.JsonToken; +import com.fasterxml.jackson.databind.DeserializationContext; +import com.fasterxml.jackson.databind.JsonDeserializer; +import com.fasterxml.jackson.databind.annotation.JsonDeserialize; +import com.fasterxml.jackson.databind.annotation.JsonPOJOBuilder; +import com.google.adk.JsonBaseModel; +import com.google.auto.value.AutoValue; +import java.io.IOException; +import java.math.BigDecimal; +import java.time.Duration; +import java.time.Instant; +import java.time.format.DateTimeParseException; +import java.util.Locale; +import java.util.Optional; +import org.jspecify.annotations.Nullable; + +/** + * Context cache state carried on an LLM response, and on its event, from one request to the next. + * + *

Metadata for an active cache has a {@link #cacheName()}, {@link #expireTime()} and {@link + * #invocationsUsed()}. Fingerprint-only metadata has none of them: it records the cacheable prefix + * so that a later request can tell whether the prefix is still the same. + */ +@AutoValue +@JsonDeserialize(builder = CacheMetadata.Builder.class) +public abstract class CacheMetadata extends JsonBaseModel { + + CacheMetadata() {} + + /** Hash of the cacheable state: system instruction, tools and the cached leading contents. */ + @JsonProperty("fingerprint") + public abstract String fingerprint(); + + /** Number of leading contents in the cache, or in the fingerprint when no cache is active. */ + @JsonProperty("contentsCount") + public abstract int contentsCount(); + + /** Full resource name of the cached content; empty when no cache is active. */ + @JsonProperty("cacheName") + public abstract Optional cacheName(); + + /** When the cache expires; empty when no cache is active. */ + @JsonProperty("expireTime") + public abstract Optional expireTime(); + + /** Number of invocations that have used the cache; empty when no cache is active. */ + @JsonProperty("invocationsUsed") + public abstract Optional invocationsUsed(); + + /** When the cache was created; empty when no cache is active. */ + @JsonProperty("createdAt") + public abstract Optional createdAt(); + + public abstract Builder toBuilder(); + + public static Builder builder() { + return new AutoValue_CacheMetadata.Builder(); + } + + /** Returns a short description for logs, in the same shape as ADK Python's. */ + @Override + public final String toString() { + if (cacheName().isEmpty()) { + String shortFingerprint = fingerprint().substring(0, Math.min(8, fingerprint().length())); + return "Fingerprint-only: " + + contentsCount() + + " contents, fingerprint=" + + shortFingerprint + + "..."; + } + String cacheId = cacheName().get().substring(cacheName().get().lastIndexOf('/') + 1); + double minutesToExpiry = + Duration.between(Instant.now(), expireTime().get()).toMillis() / 60_000.0; + return String.format( + Locale.ROOT, + "Cache %s: used %d invocations, cached %d contents, expires in %.1fmin", + cacheId, + invocationsUsed().get(), + contentsCount(), + minutesToExpiry); + } + + /** + * Builder for {@link CacheMetadata}. The snake_case aliases read metadata that ADK Python wrote. + */ + @AutoValue.Builder + @JsonPOJOBuilder(buildMethodName = "build", withPrefix = "") + public abstract static class Builder { + + @JsonCreator + static Builder jacksonBuilder() { + return CacheMetadata.builder(); + } + + @JsonProperty("fingerprint") + public abstract Builder fingerprint(String fingerprint); + + @JsonProperty("contentsCount") + @JsonAlias("contents_count") + public abstract Builder contentsCount(int contentsCount); + + @JsonProperty("cacheName") + @JsonAlias("cache_name") + public abstract Builder cacheName(@Nullable String cacheName); + + @JsonProperty("expireTime") + @JsonAlias("expire_time") + @JsonDeserialize(using = LenientEpochDeserializer.class) + public abstract Builder expireTime(@Nullable Instant expireTime); + + @JsonProperty("invocationsUsed") + @JsonAlias("invocations_used") + public abstract Builder invocationsUsed(@Nullable Integer invocationsUsed); + + @JsonProperty("createdAt") + @JsonAlias("created_at") + @JsonDeserialize(using = LenientEpochDeserializer.class) + public abstract Builder createdAt(@Nullable Instant createdAt); + + abstract CacheMetadata autoBuild(); + + /** + * Builds the metadata. + * + * @throws IllegalStateException if a count is negative, or if only some of {@code cacheName}, + * {@code expireTime} and {@code invocationsUsed} are set + */ + public CacheMetadata build() { + CacheMetadata metadata = autoBuild(); + checkState(metadata.contentsCount() >= 0, "contentsCount must not be negative."); + checkState( + metadata.invocationsUsed().orElse(0) >= 0, "invocationsUsed must not be negative."); + boolean active = metadata.cacheName().isPresent(); + checkState( + metadata.expireTime().isPresent() == active + && metadata.invocationsUsed().isPresent() == active, + "cacheName, expireTime and invocationsUsed must be all set or all unset."); + return metadata; + } + } + + /** + * Reads a numeric timestamp as epoch seconds when its magnitude is below 1e11 and as epoch + * milliseconds otherwise, since ADK Python writes seconds and ADK Kotlin writes milliseconds. + * Also reads ISO-8601 text. + */ + static final class LenientEpochDeserializer extends JsonDeserializer { + private static final BigDecimal SECONDS_MILLIS_BOUNDARY = new BigDecimal("1e11"); + // Years 1 to 9999, as ADK Kotlin accepts; checked before any arithmetic that could be costly. + private static final BigDecimal MIN_SECONDS = BigDecimal.valueOf(-62_135_596_800L); + private static final BigDecimal MAX_SECONDS = BigDecimal.valueOf(253_402_300_799L); + + @Override + public Instant deserialize(JsonParser parser, DeserializationContext context) + throws IOException { + if (parser.currentToken() == JsonToken.VALUE_STRING) { + try { + return Instant.parse(parser.getText()); + } catch (DateTimeParseException e) { + return (Instant) + context.handleWeirdStringValue( + Instant.class, parser.getText(), "not an ISO-8601 time"); + } + } + if (!parser.currentToken().isNumeric()) { + return (Instant) context.handleUnexpectedToken(Instant.class, parser); + } + BigDecimal value = parser.getDecimalValue(); + BigDecimal seconds = + value.abs().compareTo(SECONDS_MILLIS_BOUNDARY) < 0 ? value : value.movePointLeft(3); + if (seconds.compareTo(MIN_SECONDS) < 0 || seconds.compareTo(MAX_SECONDS) > 0) { + return (Instant) context.handleWeirdNumberValue(Instant.class, value, "out of range"); + } + return Instant.ofEpochSecond( + seconds.longValue(), seconds.remainder(BigDecimal.ONE).movePointRight(9).longValue()); + } + } +} diff --git a/core/src/main/java/com/google/adk/models/LlmRequest.java b/core/src/main/java/com/google/adk/models/LlmRequest.java index d6b56acdb..07c06abc6 100644 --- a/core/src/main/java/com/google/adk/models/LlmRequest.java +++ b/core/src/main/java/com/google/adk/models/LlmRequest.java @@ -25,6 +25,7 @@ import com.fasterxml.jackson.annotation.JsonProperty; import com.fasterxml.jackson.databind.annotation.JsonDeserialize; import com.google.adk.JsonBaseModel; +import com.google.adk.agents.ContextCacheConfig; import com.google.adk.agents.Role; import com.google.adk.tools.BaseTool; import com.google.auto.value.AutoValue; @@ -40,6 +41,7 @@ import java.util.Map; import java.util.Optional; import java.util.stream.Stream; +import org.jspecify.annotations.Nullable; /** Represents a request to be sent to the LLM. */ @AutoValue @@ -88,6 +90,18 @@ public abstract class LlmRequest extends JsonBaseModel { @JsonIgnore public abstract Map tools(); + /** Context cache configuration for this request; empty when context caching is disabled. */ + @JsonIgnore + public abstract Optional cacheConfig(); + + /** Cache state from the agent's latest response, used to reuse or refresh the cache. */ + @JsonIgnore + public abstract Optional cacheMetadata(); + + /** Prompt token count of the agent's previous request, used to decide whether to cache. */ + @JsonIgnore + public abstract Optional cacheableContentsTokenCount(); + /** returns the first system instruction text from the request if present. */ @JsonIgnore public Optional getFirstSystemInstruction() { @@ -155,6 +169,16 @@ private static Builder create() { abstract Map tools(); + @CanIgnoreReturnValue + public abstract Builder cacheConfig(@Nullable ContextCacheConfig cacheConfig); + + @CanIgnoreReturnValue + public abstract Builder cacheMetadata(@Nullable CacheMetadata cacheMetadata); + + @CanIgnoreReturnValue + public abstract Builder cacheableContentsTokenCount( + @Nullable Integer cacheableContentsTokenCount); + @CanIgnoreReturnValue public final Builder appendInstructions(List instructions) { if (instructions.isEmpty()) { diff --git a/core/src/main/java/com/google/adk/models/LlmResponse.java b/core/src/main/java/com/google/adk/models/LlmResponse.java index 5e97eafc2..424e9e570 100644 --- a/core/src/main/java/com/google/adk/models/LlmResponse.java +++ b/core/src/main/java/com/google/adk/models/LlmResponse.java @@ -130,6 +130,10 @@ public abstract class LlmResponse extends JsonBaseModel { @JsonProperty("outputTranscription") public abstract Optional outputTranscription(); + /** Context cache state for this response; set only when context caching is enabled. */ + @JsonProperty("cacheMetadata") + public abstract Optional cacheMetadata(); + public abstract Builder toBuilder(); /** Builder for constructing {@link LlmResponse} instances. */ @@ -185,6 +189,9 @@ public abstract Builder usageMetadata( @JsonProperty("outputTranscription") public abstract Builder outputTranscription(@Nullable Transcription outputTranscription); + @JsonProperty("cacheMetadata") + public abstract Builder cacheMetadata(@Nullable CacheMetadata cacheMetadata); + @CanIgnoreReturnValue public final Builder response(GenerateContentResponse response) { Optional> candidatesOpt = response.candidates(); diff --git a/core/src/test/java/com/google/adk/agents/ContextCacheConfigTest.java b/core/src/test/java/com/google/adk/agents/ContextCacheConfigTest.java new file mode 100644 index 000000000..855cdcd8d --- /dev/null +++ b/core/src/test/java/com/google/adk/agents/ContextCacheConfigTest.java @@ -0,0 +1,82 @@ +/* + * 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.common.truth.Truth.assertThat; + +import com.google.common.collect.ImmutableMap; +import com.google.genai.types.HttpOptions; +import java.time.Duration; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.junit.runners.JUnit4; + +@RunWith(JUnit4.class) +public final class ContextCacheConfigTest { + + private static final Duration TTL = Duration.ofMinutes(30); + + @Test + public void defaultConstructor_usesDefaults() { + ContextCacheConfig config = new ContextCacheConfig(); + + assertThat(config.cacheIntervals()).isEqualTo(10); + assertThat(config.ttl()).isEqualTo(Duration.ofMinutes(30)); + assertThat(config.minTokens()).isEqualTo(0); + assertThat(config.createHttpOptions()).isNull(); + } + + @Test + public void threeArgConstructor_leavesHttpOptionsUnset() { + ContextCacheConfig config = new ContextCacheConfig(5, TTL, 100); + + assertThat(config.cacheIntervals()).isEqualTo(5); + assertThat(config.createHttpOptions()).isNull(); + } + + @Test + public void constructor_keepsHttpOptions() { + HttpOptions httpOptions = HttpOptions.builder().timeout(10_000).build(); + + assertThat(new ContextCacheConfig(5, TTL, 100, httpOptions).createHttpOptions()) + .isEqualTo(httpOptions); + } + + @Test + @SuppressWarnings({"deprecation", "InlineMeInliner"}) // Covers the deprecated alias. + public void maxInvocations_returnsCacheIntervals() { + assertThat(new ContextCacheConfig(7, TTL, 0).maxInvocations()).isEqualTo(7); + } + + @Test + public void toString_withoutHttpOptions_saysNull() { + assertThat(new ContextCacheConfig().toString()) + .isEqualTo( + "ContextCacheConfig(cacheIntervals=10, ttl=1800s, minTokens=0," + + " createHttpOptions=null)"); + } + + @Test + public void toString_hidesHttpOptionValues() { + HttpOptions httpOptions = + HttpOptions.builder().headers(ImmutableMap.of("x-goog-api-key", "secret")).build(); + + assertThat(new ContextCacheConfig(5, TTL, 100, httpOptions).toString()) + .isEqualTo( + "ContextCacheConfig(cacheIntervals=5, ttl=1800s, minTokens=100," + + " createHttpOptions=set)"); + } +} diff --git a/core/src/test/java/com/google/adk/events/EventTest.java b/core/src/test/java/com/google/adk/events/EventTest.java index 7b91da2c7..7eb3f82e0 100644 --- a/core/src/test/java/com/google/adk/events/EventTest.java +++ b/core/src/test/java/com/google/adk/events/EventTest.java @@ -18,6 +18,7 @@ import static com.google.common.truth.Truth.assertThat; +import com.google.adk.models.CacheMetadata; import com.google.adk.workflow.NodeInfo; import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableMap; @@ -349,6 +350,32 @@ public void event_json_serialization_with_transcriptions_works() throws Exceptio assertThat(deserialized.outputTranscription()).hasValue(outputTranscription); } + @Test + public void event_cacheMetadata_survivesJsonAndToBuilder() { + CacheMetadata cacheMetadata = + CacheMetadata.builder() + .fingerprint("abc123") + .contentsCount(2) + .cacheName("cachedContents/42") + .expireTime(Instant.ofEpochSecond(2_000_000_000L, 500_000_000)) + .invocationsUsed(3) + .createdAt(Instant.ofEpochSecond(1_999_998_200L)) + .build(); + Event event = EVENT.toBuilder().cacheMetadata(cacheMetadata).build(); + + assertThat(event.cacheMetadata()).hasValue(cacheMetadata); + assertThat(event.toBuilder().build()).isEqualTo(event); + assertThat(Event.fromJson(event.toJson()).cacheMetadata()).hasValue(cacheMetadata); + assertThat(event).isNotEqualTo(EVENT); + assertThat(event.hashCode()).isNotEqualTo(EVENT.hashCode()); + } + + @Test + public void event_cacheMetadata_emptyByDefault() { + assertThat(EVENT.cacheMetadata()).isEmpty(); + assertThat(EVENT.toJson()).doesNotContain("cacheMetadata"); + } + @Test public void finalResponse_returnsTrueIfNoToolCalls() { Event event = diff --git a/core/src/test/java/com/google/adk/flows/llmflows/BaseLlmFlowTest.java b/core/src/test/java/com/google/adk/flows/llmflows/BaseLlmFlowTest.java index 7da99cef0..91bbb8ab1 100644 --- a/core/src/test/java/com/google/adk/flows/llmflows/BaseLlmFlowTest.java +++ b/core/src/test/java/com/google/adk/flows/llmflows/BaseLlmFlowTest.java @@ -35,6 +35,7 @@ import com.google.adk.events.Event; import com.google.adk.flows.llmflows.RequestProcessor.RequestProcessingResult; import com.google.adk.flows.llmflows.ResponseProcessor.ResponseProcessingResult; +import com.google.adk.models.CacheMetadata; import com.google.adk.models.LlmRequest; import com.google.adk.models.LlmResponse; import com.google.adk.testing.TestLlm; @@ -59,6 +60,7 @@ import io.reactivex.rxjava3.core.Maybe; import io.reactivex.rxjava3.core.Single; import io.reactivex.rxjava3.schedulers.Schedulers; +import java.time.Instant; import java.util.List; import java.util.Map; import java.util.Optional; @@ -122,6 +124,31 @@ public void run_singleTextResponse_withMetadata_returnsSingleEventWithMetadata() .build()); } + @Test + public void run_responseWithCacheMetadata_copiesItToEvent() { + CacheMetadata cacheMetadata = + CacheMetadata.builder() + .fingerprint("abc123") + .contentsCount(2) + .cacheName("cachedContents/42") + .expireTime(Instant.ofEpochSecond(2_000_000_000L)) + .invocationsUsed(1) + .createdAt(Instant.ofEpochSecond(1_999_998_200L)) + .build(); + TestLlm testLlm = + createTestLlm( + LlmResponse.builder() + .content(Content.fromParts(Part.fromText("LLM response"))) + .cacheMetadata(cacheMetadata) + .build()); + InvocationContext invocationContext = createInvocationContext(createTestAgent(testLlm)); + BaseLlmFlow baseLlmFlow = createBaseLlmFlowWithoutProcessors(); + + List events = baseLlmFlow.run(invocationContext).toList().blockingGet(); + + assertThat(getOnlyElement(events).cacheMetadata()).hasValue(cacheMetadata); + } + @Test public void run_withFunctionCall_returnsCorrectEvents() { Content firstContent = diff --git a/core/src/test/java/com/google/adk/flows/llmflows/ContextCacheRequestProcessorTest.java b/core/src/test/java/com/google/adk/flows/llmflows/ContextCacheRequestProcessorTest.java new file mode 100644 index 000000000..60c233a32 --- /dev/null +++ b/core/src/test/java/com/google/adk/flows/llmflows/ContextCacheRequestProcessorTest.java @@ -0,0 +1,206 @@ +/* + * 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.flows.llmflows; + +import static com.google.adk.testing.TestUtils.createInvocationContext; +import static com.google.adk.testing.TestUtils.createTestAgent; +import static com.google.adk.testing.TestUtils.createTestLlm; +import static com.google.common.truth.Truth.assertThat; + +import com.google.adk.agents.ContextCacheConfig; +import com.google.adk.agents.InvocationContext; +import com.google.adk.events.Event; +import com.google.adk.flows.llmflows.RequestProcessor.RequestProcessingResult; +import com.google.adk.models.CacheMetadata; +import com.google.adk.models.LlmRequest; +import com.google.adk.models.LlmResponse; +import com.google.genai.types.GenerateContentResponseUsageMetadata; +import java.time.Duration; +import java.time.Instant; +import org.jspecify.annotations.Nullable; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.junit.runners.JUnit4; + +@RunWith(JUnit4.class) +public final class ContextCacheRequestProcessorTest { + + private static final ContextCacheConfig CACHE_CONFIG = + new ContextCacheConfig(5, Duration.ofMinutes(10), 1024); + private static final CacheMetadata ACTIVE_CACHE = + CacheMetadata.builder() + .fingerprint("active") + .contentsCount(3) + .cacheName("cachedContents/42") + .expireTime(Instant.ofEpochSecond(2_000_000_000L)) + .invocationsUsed(2) + .createdAt(Instant.ofEpochSecond(1_999_999_000L)) + .build(); + private static final CacheMetadata FINGERPRINT_ONLY = + CacheMetadata.builder().fingerprint("prefix").contentsCount(1).build(); + private static final LlmRequest REQUEST = LlmRequest.builder().model("gemini-2.5-flash").build(); + + private final ContextCacheRequestProcessor processor = new ContextCacheRequestProcessor(); + private final InvocationContext context = + createInvocationContext(createTestAgent(createTestLlm(LlmResponse.builder().build()))) + .toBuilder() + .contextCacheConfig(CACHE_CONFIG) + .build(); + private final String agentName = context.agent().name(); + + @Test + public void processRequest_withoutCacheConfig_returnsRequestUnchanged() { + InvocationContext uncached = + createInvocationContext(createTestAgent(createTestLlm(LlmResponse.builder().build()))); + uncached.session().addEvent(event(agentName, "earlier", ACTIVE_CACHE, 4000)); + + RequestProcessingResult result = processor.processRequest(uncached, REQUEST).blockingGet(); + + assertThat(result.updatedRequest()).isEqualTo(REQUEST); + assertThat(result.events()).isEmpty(); + } + + @Test + public void processRequest_withoutEvents_setsOnlyCacheConfig() { + LlmRequest request = process(); + + assertThat(request.cacheConfig()).hasValue(CACHE_CONFIG); + assertThat(request.cacheMetadata()).isEmpty(); + assertThat(request.cacheableContentsTokenCount()).isEmpty(); + } + + @Test + public void processRequest_activeCacheFromEarlierInvocation_countsOneMoreUse() { + addEvent(event(agentName, "earlier", ACTIVE_CACHE, 4000)); + + LlmRequest request = process(); + + assertThat(request.cacheMetadata()) + .hasValue(ACTIVE_CACHE.toBuilder().invocationsUsed(3).build()); + assertThat(request.cacheableContentsTokenCount()).hasValue(4000); + } + + @Test + public void processRequest_activeCacheFromCurrentInvocation_keepsUseCount() { + addEvent(event(agentName, context.invocationId(), ACTIVE_CACHE, 4000)); + + LlmRequest request = process(); + + assertThat(request.cacheMetadata()).hasValue(ACTIVE_CACHE); + } + + @Test + public void processRequest_eventWithoutInvocationId_keepsUseCount() { + addEvent(event(agentName, /* invocationId= */ null, ACTIVE_CACHE, 4000)); + + LlmRequest request = process(); + + assertThat(request.cacheMetadata()).hasValue(ACTIVE_CACHE); + } + + @Test + public void processRequest_eventWithEmptyInvocationId_keepsUseCount() { + addEvent(event(agentName, "", ACTIVE_CACHE, 4000)); + + LlmRequest request = process(); + + assertThat(request.cacheMetadata()).hasValue(ACTIVE_CACHE); + } + + @Test + public void processRequest_fingerprintOnlyFromEarlierInvocation_isReturnedAsIs() { + addEvent(event(agentName, "earlier", FINGERPRINT_ONLY, 4000)); + + LlmRequest request = process(); + + assertThat(request.cacheMetadata()).hasValue(FINGERPRINT_ONLY); + } + + @Test + public void processRequest_otherAgentsEvents_areIgnored() { + addEvent(event(agentName, "earlier", FINGERPRINT_ONLY, 1000)); + addEvent(event("other agent", "earlier", ACTIVE_CACHE, 9000)); + + LlmRequest request = process(); + + assertThat(request.cacheMetadata()).hasValue(FINGERPRINT_ONLY); + assertThat(request.cacheableContentsTokenCount()).hasValue(1000); + } + + @Test + public void processRequest_severalEvents_usesTheLatest() { + addEvent(event(agentName, "first", FINGERPRINT_ONLY, 1000)); + addEvent(event(agentName, "second", ACTIVE_CACHE, 2000)); + + LlmRequest request = process(); + + assertThat(request.cacheMetadata()) + .hasValue(ACTIVE_CACHE.toBuilder().invocationsUsed(3).build()); + assertThat(request.cacheableContentsTokenCount()).hasValue(2000); + } + + @Test + public void processRequest_latestEventsMissingOneValue_takesItFromOlderEvents() { + addEvent(event(agentName, "first", FINGERPRINT_ONLY, /* promptTokenCount= */ null)); + addEvent(event(agentName, "second", /* cacheMetadata= */ null, 3000)); + + LlmRequest request = process(); + + assertThat(request.cacheMetadata()).hasValue(FINGERPRINT_ONLY); + assertThat(request.cacheableContentsTokenCount()).hasValue(3000); + } + + @Test + public void processRequest_latestEventWithoutUsage_keepsItsMetadataOverOlderMetadata() { + addEvent(event(agentName, "first", ACTIVE_CACHE, 1000)); + addEvent(event(agentName, "second", FINGERPRINT_ONLY, /* promptTokenCount= */ null)); + + LlmRequest request = process(); + + assertThat(request.cacheMetadata()).hasValue(FINGERPRINT_ONLY); + assertThat(request.cacheableContentsTokenCount()).hasValue(1000); + } + + private LlmRequest process() { + RequestProcessingResult result = processor.processRequest(context, REQUEST).blockingGet(); + assertThat(result.events()).isEmpty(); + return result.updatedRequest(); + } + + private void addEvent(Event event) { + context.session().addEvent(event); + } + + private static Event event( + String author, + @Nullable String invocationId, + @Nullable CacheMetadata cacheMetadata, + @Nullable Integer promptTokenCount) { + Event.Builder event = + Event.builder().id(Event.generateEventId()).author(author).cacheMetadata(cacheMetadata); + if (invocationId != null) { + event.invocationId(invocationId); + } + if (promptTokenCount != null) { + event.usageMetadata( + GenerateContentResponseUsageMetadata.builder() + .promptTokenCount(promptTokenCount) + .build()); + } + return event.build(); + } +} diff --git a/core/src/test/java/com/google/adk/flows/llmflows/SingleFlowTest.java b/core/src/test/java/com/google/adk/flows/llmflows/SingleFlowTest.java index ccb10a3a7..e26ffe366 100644 --- a/core/src/test/java/com/google/adk/flows/llmflows/SingleFlowTest.java +++ b/core/src/test/java/com/google/adk/flows/llmflows/SingleFlowTest.java @@ -16,8 +16,16 @@ package com.google.adk.flows.llmflows; +import static com.google.adk.testing.TestUtils.createInvocationContext; +import static com.google.adk.testing.TestUtils.createTestAgent; +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 com.google.adk.agents.ContextCacheConfig; +import com.google.adk.agents.InvocationContext; +import com.google.adk.testing.TestLlm; +import java.time.Duration; import org.junit.Test; import org.junit.runner.RunWith; import org.junit.runners.JUnit4; @@ -32,4 +40,25 @@ public void requestProcessors_containsCompaction() { .anyMatch(processor -> processor instanceof Compaction); assertThat(hasCompaction).isTrue(); } + + @Test + public void requestProcessors_runContextCacheAfterContents() { + assertThat(SingleFlow.REQUEST_PROCESSORS.stream().map(Object::getClass)) + .containsAtLeast(Contents.class, ContextCacheRequestProcessor.class) + .inOrder(); + } + + @Test + public void run_withContextCacheConfig_passesItToTheModel() { + ContextCacheConfig cacheConfig = new ContextCacheConfig(5, Duration.ofMinutes(10), 1024); + TestLlm testLlm = createTestLlm(createTextLlmResponse("hi")); + InvocationContext context = + createInvocationContext(createTestAgent(testLlm)).toBuilder() + .contextCacheConfig(cacheConfig) + .build(); + + var unused = new SingleFlow().run(context).toList().blockingGet(); + + assertThat(testLlm.getLastRequest().cacheConfig()).hasValue(cacheConfig); + } } diff --git a/core/src/test/java/com/google/adk/models/CacheMetadataTest.java b/core/src/test/java/com/google/adk/models/CacheMetadataTest.java new file mode 100644 index 000000000..4a99c973d --- /dev/null +++ b/core/src/test/java/com/google/adk/models/CacheMetadataTest.java @@ -0,0 +1,217 @@ +/* + * 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.models; + +import static com.google.common.truth.Truth.assertThat; +import static org.junit.Assert.assertThrows; + +import com.google.adk.JsonBaseModel; +import com.google.common.collect.ImmutableList; +import java.time.Duration; +import java.time.Instant; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.junit.runners.JUnit4; + +@RunWith(JUnit4.class) +public final class CacheMetadataTest { + + private static final CacheMetadata ACTIVE = + CacheMetadata.builder() + .fingerprint("abc123") + .contentsCount(4) + .cacheName("projects/1/locations/us-central1/cachedContents/42") + .expireTime(Instant.ofEpochSecond(2_000_000_000L, 250_000_000)) + .invocationsUsed(3) + .createdAt(Instant.ofEpochSecond(1_999_998_200L)) + .build(); + + @Test + public void build_fingerprintOnly_hasNoCacheFields() { + CacheMetadata metadata = CacheMetadata.builder().fingerprint("abc123").contentsCount(0).build(); + + assertThat(metadata.fingerprint()).isEqualTo("abc123"); + assertThat(metadata.contentsCount()).isEqualTo(0); + assertThat(metadata.cacheName()).isEmpty(); + assertThat(metadata.expireTime()).isEmpty(); + assertThat(metadata.invocationsUsed()).isEmpty(); + assertThat(metadata.createdAt()).isEmpty(); + } + + @Test + public void build_negativeContentsCount_throws() { + CacheMetadata.Builder builder = CacheMetadata.builder().fingerprint("abc123").contentsCount(-1); + + assertThrows(IllegalStateException.class, builder::build); + } + + @Test + public void build_negativeInvocationsUsed_throws() { + CacheMetadata.Builder builder = ACTIVE.toBuilder().invocationsUsed(-1); + + assertThrows(IllegalStateException.class, builder::build); + } + + @Test + public void build_cacheNameWithoutExpireTime_throws() { + CacheMetadata.Builder builder = ACTIVE.toBuilder().expireTime(null); + + assertThrows(IllegalStateException.class, builder::build); + } + + @Test + public void build_cacheNameWithoutInvocationsUsed_throws() { + CacheMetadata.Builder builder = ACTIVE.toBuilder().invocationsUsed(null); + + assertThrows(IllegalStateException.class, builder::build); + } + + @Test + public void build_expireTimeWithoutCacheName_throws() { + CacheMetadata.Builder builder = ACTIVE.toBuilder().cacheName(null).invocationsUsed(null); + + assertThrows(IllegalStateException.class, builder::build); + } + + @Test + public void build_withoutContentsCount_throws() { + CacheMetadata.Builder builder = CacheMetadata.builder().fingerprint("abc123"); + + assertThrows(IllegalStateException.class, builder::build); + } + + @Test + public void toJson_activeCache_roundTrips() { + String json = ACTIVE.toJson(); + + assertThat(json).contains("\"contentsCount\":4"); + assertThat(JsonBaseModel.fromJsonString(json, CacheMetadata.class)).isEqualTo(ACTIVE); + } + + @Test + public void toJson_fingerprintOnly_omitsCacheFields() throws Exception { + CacheMetadata metadata = CacheMetadata.builder().fingerprint("abc123").contentsCount(2).build(); + + String json = metadata.toJson(); + + assertThat(ImmutableList.copyOf(JsonBaseModel.getMapper().readTree(json).fieldNames())) + .containsExactly("fingerprint", "contentsCount"); + assertThat(JsonBaseModel.fromJsonString(json, CacheMetadata.class)).isEqualTo(metadata); + } + + @Test + public void fromJson_writtenByAdkPython_readsSnakeCaseAndEpochSeconds() { + String json = + "{\"cache_name\": \"projects/1/locations/us-central1/cachedContents/42\"," + + " \"expire_time\": 2000000000.25, \"fingerprint\": \"abc123\"," + + " \"invocations_used\": 3, \"contents_count\": 4, \"created_at\": 1999998200.0}"; + + assertThat(JsonBaseModel.fromJsonString(json, CacheMetadata.class)).isEqualTo(ACTIVE); + } + + @Test + public void fromJson_writtenByAdkKotlin_readsEpochMilliseconds() { + String json = + "{\"cacheName\": \"projects/1/locations/us-central1/cachedContents/42\"," + + " \"expireTime\": 2000000000250, \"fingerprint\": \"abc123\"," + + " \"invocationsUsed\": 3, \"contentsCount\": 4, \"createdAt\": 1999998200000}"; + + assertThat(JsonBaseModel.fromJsonString(json, CacheMetadata.class)).isEqualTo(ACTIVE); + } + + @Test + public void fromJson_isoTimestamps_reads() { + String json = + "{\"cacheName\": \"projects/1/locations/us-central1/cachedContents/42\"," + + " \"expireTime\": \"2033-05-18T03:33:20.250Z\", \"fingerprint\": \"abc123\"," + + " \"invocationsUsed\": 3, \"contentsCount\": 4," + + " \"createdAt\": \"2033-05-18T03:03:20Z\"}"; + + assertThat(JsonBaseModel.fromJsonString(json, CacheMetadata.class)).isEqualTo(ACTIVE); + } + + @Test + public void fromJson_epochAtSecondsMillisBoundary_switchesUnit() { + CacheMetadata below = fingerprintOnlyWithCreatedAt("99999999999"); + CacheMetadata atBoundary = fingerprintOnlyWithCreatedAt("100000000000"); + + assertThat(below.createdAt()).hasValue(Instant.ofEpochSecond(99_999_999_999L)); + assertThat(atBoundary.createdAt()).hasValue(Instant.ofEpochSecond(100_000_000L)); + } + + @Test + public void fromJson_negativeFractionalSeconds_reads() { + assertThat(fingerprintOnlyWithCreatedAt("-1.5").createdAt()) + .hasValue(Instant.ofEpochSecond(-2, 500_000_000)); + } + + @Test + public void fromJson_hugeExponent_throwsWithoutExpandingIt() { + assertThrows(IllegalStateException.class, () -> fingerprintOnlyWithCreatedAt("1e20000000")); + } + + @Test + public void fromJson_negativeEpochMilliseconds_readsByMagnitude() { + assertThat(fingerprintOnlyWithCreatedAt("-1e12").createdAt()) + .hasValue(Instant.ofEpochSecond(-1_000_000_000)); + } + + @Test + public void fromJson_beforeYearOne_throws() { + // -7e13 ms is -7e10 s, before year one. + assertThrows(IllegalStateException.class, () -> fingerprintOnlyWithCreatedAt("-7e13")); + } + + @Test + public void fromJson_malformedTimestampText_throws() { + assertThrows(IllegalStateException.class, () -> fingerprintOnlyWithCreatedAt("\"not-a-time\"")); + } + + @Test + public void fromJson_booleanTimestamp_throws() { + String json = + "{\"cacheName\": \"cachedContents/42\", \"expireTime\": true," + + " \"fingerprint\": \"abc123\", \"invocationsUsed\": 3, \"contentsCount\": 4}"; + + assertThrows( + IllegalStateException.class, () -> JsonBaseModel.fromJsonString(json, CacheMetadata.class)); + } + + private static CacheMetadata fingerprintOnlyWithCreatedAt(String createdAtJson) { + return JsonBaseModel.fromJsonString( + "{\"fingerprint\": \"abc123\", \"contentsCount\": 0, \"createdAt\": " + createdAtJson + "}", + CacheMetadata.class); + } + + @Test + public void toString_fingerprintOnly_showsShortFingerprint() { + CacheMetadata metadata = + CacheMetadata.builder().fingerprint("abc123456789abcd").contentsCount(2).build(); + + assertThat(metadata.toString()) + .isEqualTo("Fingerprint-only: 2 contents, fingerprint=abc12345..."); + } + + @Test + public void toString_activeCache_showsCacheIdAndUse() { + CacheMetadata metadata = + ACTIVE.toBuilder().expireTime(Instant.now().plus(Duration.ofMinutes(10))).build(); + + assertThat(metadata.toString()) + .matches("Cache 42: used 3 invocations, cached 4 contents, expires in (9\\.9|10\\.0)min"); + } +} diff --git a/core/src/test/java/com/google/adk/models/FakeGeminiApi.java b/core/src/test/java/com/google/adk/models/FakeGeminiApi.java new file mode 100644 index 000000000..fc8ff6301 --- /dev/null +++ b/core/src/test/java/com/google/adk/models/FakeGeminiApi.java @@ -0,0 +1,236 @@ +/* + * 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.models; + +import static com.google.common.collect.ImmutableList.toImmutableList; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.fasterxml.jackson.databind.node.MissingNode; +import com.google.common.collect.ImmutableList; +import com.google.genai.Client; +import com.google.genai.types.Candidate; +import com.google.genai.types.ClientOptions; +import com.google.genai.types.Content; +import com.google.genai.types.FinishReason; +import com.google.genai.types.GenerateContentResponse; +import com.google.genai.types.GenerateContentResponseUsageMetadata; +import com.google.genai.types.Part; +import java.io.IOException; +import java.time.Duration; +import java.time.Instant; +import java.util.ArrayDeque; +import java.util.ArrayList; +import java.util.Deque; +import java.util.List; +import okhttp3.Headers; +import okhttp3.Interceptor; +import okhttp3.MediaType; +import okhttp3.OkHttpClient; +import okhttp3.Protocol; +import okhttp3.Request; +import okhttp3.Response; +import okhttp3.ResponseBody; +import okio.Buffer; +import org.jspecify.annotations.Nullable; + +/** + * A fake Gemini API that a real genai {@link Client} reaches over HTTP. Each model call gets the + * next numbered answer, and each cache create returns a new numbered cache that expires in an hour, + * unless told to fail or to omit the expiry. + */ +final class FakeGeminiApi implements Interceptor { + + /** The kind of an API call the fake has served. */ + enum Kind { + GENERATE, + CREATE_CACHE, + DELETE_CACHE + } + + private static final ObjectMapper objectMapper = new ObjectMapper(); + + private final List calls = new ArrayList<>(); + private final List expireTimes = new ArrayList<>(); + private final Deque createFailures = new ArrayDeque<>(); + private final Deque deleteFailures = new ArrayDeque<>(); + private @Nullable String nextCreateBody = null; + private boolean omitExpireTime = false; + private Duration cacheCallDelay = Duration.ZERO; + private int answers = 0; + + Client client() { + OkHttpClient httpClient = new OkHttpClient.Builder().addInterceptor(this).build(); + return Client.builder() + .apiKey("test-api-key") + .vertexAI(false) + .clientOptions(ClientOptions.builder().customHttpClient(httpClient).build()) + .build(); + } + + synchronized void failNextCreate(int code) { + createFailures.add(code); + } + + synchronized void failNextDelete(int code) { + deleteFailures.add(code); + } + + synchronized void answerNextCreateWith(String body) { + nextCreateBody = body; + } + + synchronized void omitExpireTime() { + omitExpireTime = true; + } + + synchronized void delayCacheCalls(Duration delay) { + cacheCallDelay = delay; + } + + /** Returns the expiry reported for the {@code n}-th created cache, counting from 1. */ + synchronized Instant expireTime(int n) { + return expireTimes.get(n - 1); + } + + synchronized ImmutableList kinds() { + return calls.stream().map(Call::kind).collect(toImmutableList()); + } + + synchronized ImmutableList paths(Kind kind) { + return calls.stream() + .filter(call -> call.kind() == kind) + .map(Call::path) + .collect(toImmutableList()); + } + + synchronized ImmutableList bodies(Kind kind) { + return calls.stream() + .filter(call -> call.kind() == kind) + .map(Call::body) + .collect(toImmutableList()); + } + + synchronized ImmutableList headers(Kind kind) { + return calls.stream() + .filter(call -> call.kind() == kind) + .map(Call::headers) + .collect(toImmutableList()); + } + + @Override + public synchronized Response intercept(Chain chain) throws IOException { + Request request = chain.request(); + String path = request.url().encodedPath(); + JsonNode body = MissingNode.getInstance(); + if (request.body() != null) { + Buffer buffer = new Buffer(); + request.body().writeTo(buffer); + body = objectMapper.readTree(buffer.readUtf8()); + } + Kind kind; + int code = 200; + String responseBody; + if (path.endsWith(":generateContent")) { + kind = Kind.GENERATE; + answers++; + responseBody = answer("Answer " + answers).toJson(); + } else if (path.endsWith(":streamGenerateContent")) { + kind = Kind.GENERATE; + answers++; + responseBody = + "data: " + + partialAnswer("Answer ").toJson() + + "\n\ndata: " + + answer(String.valueOf(answers)).toJson() + + "\n\n"; + } else if (request.method().equals("POST") && path.endsWith("/cachedContents")) { + kind = Kind.CREATE_CACHE; + sleep(cacheCallDelay); + if (!createFailures.isEmpty()) { + code = createFailures.remove(); + responseBody = error(code); + } else if (nextCreateBody != null) { + responseBody = nextCreateBody; + nextCreateBody = null; + } else { + Instant expireTime = Instant.now().plus(Duration.ofHours(1)); + expireTimes.add(expireTime); + String name = "cachedContents/cache-" + expireTimes.size(); + responseBody = + omitExpireTime + ? String.format("{\"name\": \"%s\"}", name) + : String.format("{\"name\": \"%s\", \"expireTime\": \"%s\"}", name, expireTime); + } + } else if (request.method().equals("DELETE") && path.contains("/cachedContents/")) { + kind = Kind.DELETE_CACHE; + sleep(cacheCallDelay); + code = deleteFailures.isEmpty() ? 200 : deleteFailures.remove(); + responseBody = code == 200 ? "{}" : error(code); + } else { + throw new IOException("Unexpected request: " + request.method() + " " + path); + } + calls.add(new Call(kind, path, body, request.headers())); + return new Response.Builder() + .request(request) + .protocol(Protocol.HTTP_1_1) + .code(code) + .message("fake") + .body(ResponseBody.create(responseBody, MediaType.get("application/json"))) + .build(); + } + + private static GenerateContentResponse answer(String text) { + return GenerateContentResponse.builder() + .candidates( + Candidate.builder() + .content(modelText(text)) + .finishReason(new FinishReason(FinishReason.Known.STOP)) + .build()) + .usageMetadata( + GenerateContentResponseUsageMetadata.builder() + .promptTokenCount(5000) + .candidatesTokenCount(10) + .totalTokenCount(5010) + .build()) + .build(); + } + + private static GenerateContentResponse partialAnswer(String text) { + return GenerateContentResponse.builder() + .candidates(Candidate.builder().content(modelText(text)).build()) + .build(); + } + + private static void sleep(Duration duration) throws IOException { + try { + Thread.sleep(duration.toMillis()); + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + throw new IOException(e); + } + } + + private static Content modelText(String text) { + return Content.builder().role("model").parts(Part.fromText(text)).build(); + } + + private static String error(int code) { + return "{\"error\": {\"code\": " + code + ", \"message\": \"fake\"}}"; + } + + private record Call(Kind kind, String path, JsonNode body, Headers headers) {} +} diff --git a/core/src/test/java/com/google/adk/models/FakeGeminiApiTest.java b/core/src/test/java/com/google/adk/models/FakeGeminiApiTest.java new file mode 100644 index 000000000..cc7a3a0a6 --- /dev/null +++ b/core/src/test/java/com/google/adk/models/FakeGeminiApiTest.java @@ -0,0 +1,137 @@ +/* + * 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.models; + +import static com.google.common.collect.ImmutableList.toImmutableList; +import static com.google.common.truth.Truth.assertThat; +import static org.junit.Assert.assertThrows; + +import com.google.adk.models.FakeGeminiApi.Kind; +import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableMap; +import com.google.genai.Client; +import com.google.genai.errors.ApiException; +import com.google.genai.types.CachedContent; +import com.google.genai.types.CreateCachedContentConfig; +import com.google.genai.types.DeleteCachedContentConfig; +import com.google.genai.types.GenerateContentResponse; +import com.google.genai.types.HttpOptions; +import java.time.Duration; +import java.time.Instant; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.junit.runners.JUnit4; + +@RunWith(JUnit4.class) +public final class FakeGeminiApiTest { + + private static final String MODEL = "gemini-test-model"; + + private final FakeGeminiApi api = new FakeGeminiApi(); + private final Client client = api.client(); + + @Test + public void generateContent_returnsNumberedAnswers() { + String first = client.models.generateContent(MODEL, "Hi", null).text(); + String second = client.models.generateContent(MODEL, "Hi again", null).text(); + + assertThat(first).isEqualTo("Answer 1"); + assertThat(second).isEqualTo("Answer 2"); + assertThat(api.kinds()).containsExactly(Kind.GENERATE, Kind.GENERATE); + assertThat(api.bodies(Kind.GENERATE).get(0).at("/contents/0/parts/0/text").asText()) + .isEqualTo("Hi"); + } + + @Test + public void generateContentStream_returnsPartialThenFinalAnswer() { + ImmutableList chunks = + ImmutableList.copyOf(client.models.generateContentStream(MODEL, "Hi", null)).stream() + .map(GenerateContentResponse::text) + .collect(toImmutableList()); + + assertThat(chunks).containsExactly("Answer ", "1").inOrder(); + assertThat(api.kinds()).containsExactly(Kind.GENERATE); + } + + @Test + public void createCache_returnsNumberedCacheThatExpiresInAnHour() { + Instant before = Instant.now(); + + CachedContent first = createCache(); + CachedContent second = createCache(); + + assertThat(first.name()).hasValue("cachedContents/cache-1"); + assertThat(second.name()).hasValue("cachedContents/cache-2"); + assertThat(first.expireTime()).hasValue(api.expireTime(1)); + assertThat(api.expireTime(1)).isAtLeast(before.plus(Duration.ofHours(1))); + assertThat(api.kinds()).containsExactly(Kind.CREATE_CACHE, Kind.CREATE_CACHE); + } + + @Test + public void failNextCreate_failsOnlyTheNextCreate() { + api.failNextCreate(400); + + ApiException error = assertThrows(ApiException.class, this::createCache); + CachedContent next = createCache(); + + assertThat(error.code()).isEqualTo(400); + assertThat(next.name()).hasValue("cachedContents/cache-1"); + } + + @Test + public void failNextDelete_failsOnlyTheNextDelete() { + api.failNextDelete(500); + DeleteCachedContentConfig config = DeleteCachedContentConfig.builder().build(); + + ApiException error = + assertThrows(ApiException.class, () -> client.caches.delete("cachedContents/a", config)); + var unused = client.caches.delete("cachedContents/b", config); + + assertThat(error.code()).isEqualTo(500); + assertThat(api.paths(Kind.DELETE_CACHE).get(1)).endsWith("/cachedContents/b"); + } + + @Test + public void omitExpireTime_returnsCacheWithoutExpireTime() { + api.omitExpireTime(); + + assertThat(createCache().expireTime()).isEmpty(); + } + + @Test + public void answerNextCreateWith_returnsThatBodyOnce() { + api.answerNextCreateWith("{\"name\": \"cachedContents/custom\"}"); + + assertThat(createCache().name()).hasValue("cachedContents/custom"); + assertThat(createCache().name()).hasValue("cachedContents/cache-1"); + } + + @Test + public void headers_recordsRequestHeaders() { + CreateCachedContentConfig config = + CreateCachedContentConfig.builder() + .httpOptions(HttpOptions.builder().headers(ImmutableMap.of("x-test", "yes")).build()) + .build(); + + var unused = client.caches.create(MODEL, config); + + assertThat(api.headers(Kind.CREATE_CACHE).get(0).get("x-test")).isEqualTo("yes"); + } + + private CachedContent createCache() { + return client.caches.create(MODEL, CreateCachedContentConfig.builder().build()); + } +} diff --git a/core/src/test/java/com/google/adk/models/LlmRequestTest.java b/core/src/test/java/com/google/adk/models/LlmRequestTest.java index 56c2debd7..953406845 100644 --- a/core/src/test/java/com/google/adk/models/LlmRequestTest.java +++ b/core/src/test/java/com/google/adk/models/LlmRequestTest.java @@ -19,6 +19,7 @@ import static com.google.common.truth.Truth.assertThat; import static org.junit.Assert.assertThrows; +import com.google.adk.agents.ContextCacheConfig; import com.google.adk.tools.BaseTool; import com.google.common.collect.ImmutableList; import com.google.common.collect.ImmutableMap; @@ -27,6 +28,7 @@ import com.google.genai.types.LiveConnectConfig; import com.google.genai.types.Part; import com.google.genai.types.Schema; +import java.time.Duration; import java.util.Optional; import org.junit.Test; import org.junit.runner.RunWith; @@ -308,4 +310,35 @@ public void getSystemInstructions_whenPresent_returnsList() { .containsExactly(instruction1 + "\n\n" + instruction2) .inOrder(); } + + @Test + public void cacheFields_surviveToBuilderAndStayOutOfJson() { + ContextCacheConfig cacheConfig = new ContextCacheConfig(5, Duration.ofMinutes(10), 1024); + CacheMetadata cacheMetadata = + CacheMetadata.builder().fingerprint("abc123").contentsCount(2).build(); + + LlmRequest request = + LlmRequest.builder() + .model("gemini-2.5-flash") + .cacheConfig(cacheConfig) + .cacheMetadata(cacheMetadata) + .cacheableContentsTokenCount(4096) + .build() + .toBuilder() + .build(); + + assertThat(request.cacheConfig()).hasValue(cacheConfig); + assertThat(request.cacheMetadata()).hasValue(cacheMetadata); + assertThat(request.cacheableContentsTokenCount()).hasValue(4096); + assertThat(request.toJson()).doesNotContain("cache"); + } + + @Test + public void cacheFields_emptyByDefault() { + LlmRequest request = LlmRequest.builder().build(); + + assertThat(request.cacheConfig()).isEmpty(); + assertThat(request.cacheMetadata()).isEmpty(); + assertThat(request.cacheableContentsTokenCount()).isEmpty(); + } } diff --git a/core/src/test/java/com/google/adk/models/LlmResponseTest.java b/core/src/test/java/com/google/adk/models/LlmResponseTest.java index 44b16854b..c930f73b5 100644 --- a/core/src/test/java/com/google/adk/models/LlmResponseTest.java +++ b/core/src/test/java/com/google/adk/models/LlmResponseTest.java @@ -144,6 +144,26 @@ public void testSerializationAndDeserialization_optionalFieldsEmpty() assertThat(deserializedResponse.usageMetadata()).isEmpty(); } + @Test + public void testSerializationAndDeserialization_withCacheMetadata() + throws JsonProcessingException { + CacheMetadata cacheMetadata = + CacheMetadata.builder().fingerprint("abc123").contentsCount(2).build(); + LlmResponse originalResponse = + LlmResponse.builder() + .content(createSampleContent("hello")) + .cacheMetadata(cacheMetadata) + .build(); + + String json = originalResponse.toJson(); + + assertThat(objectMapper.readTree(json).get("cacheMetadata").get("fingerprint").asText()) + .isEqualTo("abc123"); + LlmResponse deserializedResponse = LlmResponse.fromJsonString(json, LlmResponse.class); + assertThat(deserializedResponse).isEqualTo(originalResponse); + assertThat(deserializedResponse.cacheMetadata()).hasValue(cacheMetadata); + } + @Test public void testSerializationAndDeserialization_withTranscriptions() throws JsonProcessingException {