Skip to content
Open
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
29 changes: 28 additions & 1 deletion core/src/main/java/com/google/adk/events/Event.java
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
import com.fasterxml.jackson.annotation.JsonProperty;
import com.fasterxml.jackson.databind.annotation.JsonDeserialize;
import com.google.adk.JsonBaseModel;
import com.google.adk.models.CacheMetadata;
import com.google.adk.platform.UuidProvider;
import com.google.common.collect.ImmutableList;
import com.google.common.collect.Iterables;
Expand Down Expand Up @@ -66,6 +67,7 @@ public class Event extends JsonBaseModel {
private @Nullable String modelVersion;
private @Nullable Transcription inputTranscription;
private @Nullable Transcription outputTranscription;
private @Nullable CacheMetadata cacheMetadata;

private long timestamp;

Expand Down Expand Up @@ -306,6 +308,19 @@ public void setOutputTranscription(@Nullable Transcription outputTranscription)
this.outputTranscription = outputTranscription;
}

/**
* 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> cacheMetadata() {
return Optional.ofNullable(cacheMetadata);
}

public void setCacheMetadata(@Nullable CacheMetadata cacheMetadata) {
this.cacheMetadata = cacheMetadata;
}

/** The timestamp of the event. */
@JsonProperty("timestamp")
public long timestamp() {
Expand Down Expand Up @@ -415,6 +430,7 @@ public static class Builder {
private @Nullable String modelVersion;
private @Nullable Transcription inputTranscription;
private @Nullable Transcription outputTranscription;
private @Nullable CacheMetadata cacheMetadata;
private @Nullable Long timestamp;

@JsonCreator
Expand Down Expand Up @@ -592,6 +608,13 @@ public Builder outputTranscription(@Nullable Transcription 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);
Expand All @@ -616,6 +639,7 @@ public Event build() {
timestamp().orElseGet(() -> InstantSource.system().instant().toEpochMilli()));
event.setInputTranscription(inputTranscription);
event.setOutputTranscription(outputTranscription);
event.setCacheMetadata(cacheMetadata);
return event;
}
}
Expand Down Expand Up @@ -653,6 +677,7 @@ public Builder toBuilder() {
.modelVersion(this.modelVersion)
.inputTranscription(this.inputTranscription)
.outputTranscription(this.outputTranscription)
.cacheMetadata(this.cacheMetadata)
.timestamp(this.timestamp);
return builder;
}
Expand Down Expand Up @@ -685,7 +710,8 @@ public boolean equals(Object obj) {
&& Objects.equals(customMetadata, other.customMetadata)
&& Objects.equals(modelVersion, other.modelVersion)
&& Objects.equals(inputTranscription, other.inputTranscription)
&& Objects.equals(outputTranscription, other.outputTranscription);
&& Objects.equals(outputTranscription, other.outputTranscription)
&& Objects.equals(cacheMetadata, other.cacheMetadata);
}

@Override
Expand Down Expand Up @@ -716,6 +742,7 @@ public int hashCode() {
modelVersion,
inputTranscription,
outputTranscription,
cacheMetadata,
timestamp);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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();

Expand Down
Original file line number Diff line number Diff line change
@@ -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<RequestProcessingResult> processRequest(
InvocationContext context, LlmRequest request) {
Optional<ContextCacheConfig> 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();
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ public class SingleFlow extends BaseLlmFlow {
new Identity(),
new Compaction(),
new Contents(),
new ContextCacheRequestProcessor(),
CodeExecution.requestProcessor);

protected static final ImmutableList<ResponseProcessor> RESPONSE_PROCESSORS =
Expand Down
Loading
Loading