Skip to content

Commit 7c5b59b

Browse files
committed
Cache citation documents in Anthropic prompt caching
Citation documents never got a cache breakpoint. With SYSTEM_ONLY or SYSTEM_AND_TOOLS they sat after the last breakpoint and were sent uncached on every request, which is the opposite of what you want for a large PDF. Attach the documents to the first user message only. They were repeated on every user message, so multi-turn prompts re-sent the PDF each turn. Put a breakpoint on the last document block under every strategy that caches system content, using the SYSTEM TTL. TOOLS_ONLY is left out since a breakpoint there would also cache the system prompt. Fixes #5106 Signed-off-by: Soby Chacko <soby.chacko@broadcom.com>
1 parent 2c6f79d commit 7c5b59b

6 files changed

Lines changed: 312 additions & 8 deletions

File tree

‎models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java‎

Lines changed: 26 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -714,6 +714,17 @@ else if (requestOptions.getCacheOptions().isMultiBlockSystemCaching() && systemT
714714
}
715715
}
716716

717+
// Citation documents are attached once, to the first user message. Repeating
718+
// them on every user turn would re-send the source material on each request
719+
// and place it after any cache breakpoint set on the previous turn.
720+
int firstUserIndex = -1;
721+
for (int i = 0; i < nonSystemMessages.size(); i++) {
722+
if (nonSystemMessages.get(i).getMessageType() == MessageType.USER) {
723+
firstUserIndex = i;
724+
break;
725+
}
726+
}
727+
717728
// Pre-compute last user message index for CONVERSATION_HISTORY strategy
718729
int lastUserIndex = -1;
719730
if (cacheResolver.isCachingEnabled()) {
@@ -744,7 +755,7 @@ else if (requestOptions.getCacheOptions().isMultiBlockSystemCaching() && systemT
744755

745756
if (message.getMessageType() == MessageType.USER) {
746757
UserMessage userMessage = (UserMessage) message;
747-
boolean hasCitationDocs = !CollectionUtils.isEmpty(citationDocuments);
758+
boolean hasCitationDocs = !CollectionUtils.isEmpty(citationDocuments) && i == firstUserIndex;
748759
boolean hasMedia = !CollectionUtils.isEmpty(userMessage.getMedia());
749760
boolean isLastUserMessage = (i == lastUserIndex);
750761
boolean applyCacheToUser = isLastUserMessage && cacheResolver.isCachingEnabled();
@@ -759,10 +770,21 @@ else if (requestOptions.getCacheOptions().isMultiBlockSystemCaching() && systemT
759770
if (hasCitationDocs || hasMedia || userCacheControl != null) {
760771
List<ContentBlockParam> contentBlocks = new ArrayList<>();
761772

762-
// Prepend citation document blocks to the first user message
773+
// Prepend citation documents to the first user message. The cache
774+
// breakpoint goes on the last document to cover the whole set.
763775
if (hasCitationDocs) {
764-
for (AnthropicCitationDocument doc : Objects.requireNonNull(citationDocuments)) {
765-
contentBlocks.add(ContentBlockParam.ofDocument(doc.toDocumentBlockParam()));
776+
List<AnthropicCitationDocument> documents = Objects.requireNonNull(citationDocuments);
777+
CacheControlEphemeral documentCacheControl = cacheResolver
778+
.resolveCitationDocumentCacheControl();
779+
for (int d = 0; d < documents.size(); d++) {
780+
DocumentBlockParam.Builder documentBuilder = documents.get(d)
781+
.toDocumentBlockParam()
782+
.toBuilder();
783+
if (documentCacheControl != null && d == documents.size() - 1) {
784+
documentBuilder.cacheControl(documentCacheControl);
785+
cacheResolver.useCacheBlock();
786+
}
787+
contentBlocks.add(ContentBlockParam.ofDocument(documentBuilder.build()));
766788
}
767789
}
768790

‎models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/CacheEligibilityResolver.java‎

Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -133,6 +133,44 @@ private static Set<MessageType> extractEligibleMessageTypes(AnthropicCacheStrate
133133
return CacheControlEphemeral.builder().ttl(cacheTtl.getSdkTtl()).build();
134134
}
135135

136+
/**
137+
* Resolves the cache control for citation documents. Documents are stable reference
138+
* material in the same way as the system prompt, so they receive a breakpoint under
139+
* every strategy that caches system content, using the {@code SYSTEM} TTL. Under
140+
* {@link AnthropicCacheStrategy#TOOLS_ONLY} documents are not cached, because a
141+
* breakpoint on a document would also cache the system prompt that precedes it.
142+
* @return the cache control to apply to the last document block, or {@code null} if
143+
* documents are not cached under the current strategy or all breakpoints are used
144+
* @since 2.0.2
145+
*/
146+
public @Nullable CacheControlEphemeral resolveCitationDocumentCacheControl() {
147+
if (this.cacheStrategy != AnthropicCacheStrategy.SYSTEM_ONLY
148+
&& this.cacheStrategy != AnthropicCacheStrategy.SYSTEM_AND_TOOLS
149+
&& this.cacheStrategy != AnthropicCacheStrategy.CONVERSATION_HISTORY) {
150+
if (logger.isDebugEnabled()) {
151+
logger.debug("Caching not enabled for citation documents, cacheStrategy=" + this.cacheStrategy);
152+
}
153+
return null;
154+
}
155+
156+
if (this.cacheBreakpointTracker.allBreakpointsAreUsed()) {
157+
if (logger.isDebugEnabled()) {
158+
logger.debug("Caching not enabled for citation documents, usedBreakpoints="
159+
+ this.cacheBreakpointTracker.getCount());
160+
}
161+
return null;
162+
}
163+
164+
AnthropicCacheTtl cacheTtl = this.messageTypeTtl.get(MessageType.SYSTEM);
165+
Assert.state(cacheTtl != null, "messageTypeTtl must contain a 'system' entry");
166+
167+
if (logger.isDebugEnabled()) {
168+
logger.debug("Caching enabled for citation documents, ttl=" + cacheTtl);
169+
}
170+
171+
return CacheControlEphemeral.builder().ttl(cacheTtl.getSdkTtl()).build();
172+
}
173+
136174
public boolean isCachingEnabled() {
137175
return this.cacheStrategy != AnthropicCacheStrategy.NONE;
138176
}

‎models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatModelTests.java‎

Lines changed: 136 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -51,6 +51,7 @@
5151
import com.anthropic.models.messages.TextBlock;
5252
import com.anthropic.models.messages.ThinkingBlock;
5353
import com.anthropic.models.messages.ToolResultBlockParam;
54+
import com.anthropic.models.messages.ToolUnion;
5455
import com.anthropic.models.messages.ToolUseBlock;
5556
import com.anthropic.models.messages.Usage;
5657
import com.anthropic.services.async.MessageServiceAsync;
@@ -994,6 +995,141 @@ void streamingAttachesRateLimitHeadersToResponse() {
994995
assertThat(rateLimit.getTokensRemaining()).isEqualTo(49000L);
995996
}
996997

998+
@Test
999+
void citationDocumentsAreSentOnlyInFirstUserMessage() {
1000+
Message mockResponse = createMockMessage("Answer", StopReason.END_TURN);
1001+
given(this.messageService.create(any(MessageCreateParams.class))).willReturn(mockResponse);
1002+
1003+
AnthropicCitationDocument document = AnthropicCitationDocument.builder()
1004+
.plainText("Reference material")
1005+
.title("Reference")
1006+
.citationsEnabled(true)
1007+
.build();
1008+
AnthropicChatOptions options = AnthropicChatOptions.builder().citationDocuments(document).build();
1009+
1010+
UserMessage user1 = new UserMessage("First question");
1011+
AssistantMessage assistant1 = new AssistantMessage("First answer");
1012+
UserMessage user2 = new UserMessage("Second question");
1013+
1014+
this.chatModel.call(new Prompt(List.of(user1, assistant1, user2), options));
1015+
1016+
ArgumentCaptor<MessageCreateParams> captor = ArgumentCaptor.forClass(MessageCreateParams.class);
1017+
verify(this.messageService).create(captor.capture());
1018+
1019+
List<MessageParam> messages = captor.getValue().messages();
1020+
assertThat(messages).hasSize(3);
1021+
1022+
List<ContentBlockParam> firstUserBlocks = messages.get(0).content().blockParams().orElseThrow();
1023+
assertThat(firstUserBlocks).hasSize(2);
1024+
assertThat(firstUserBlocks.get(0).isDocument()).isTrue();
1025+
assertThat(firstUserBlocks.get(1).asText().text()).isEqualTo("First question");
1026+
1027+
assertThat(messages.get(2).content().string()).contains("Second question");
1028+
assertThat(messages.get(2).content().blockParams()).isEmpty();
1029+
1030+
long documentBlocks = messages.stream()
1031+
.flatMap(message -> message.content().blockParams().stream().flatMap(List::stream))
1032+
.filter(ContentBlockParam::isDocument)
1033+
.count();
1034+
assertThat(documentBlocks).isEqualTo(1);
1035+
}
1036+
1037+
@Test
1038+
void citationDocumentCacheBreakpointPerStrategy() {
1039+
assertThat(lastDocumentHasCacheControl(AnthropicCacheStrategy.NONE)).isFalse();
1040+
assertThat(lastDocumentHasCacheControl(AnthropicCacheStrategy.TOOLS_ONLY)).isFalse();
1041+
assertThat(lastDocumentHasCacheControl(AnthropicCacheStrategy.SYSTEM_ONLY)).isTrue();
1042+
assertThat(lastDocumentHasCacheControl(AnthropicCacheStrategy.SYSTEM_AND_TOOLS)).isTrue();
1043+
assertThat(lastDocumentHasCacheControl(AnthropicCacheStrategy.CONVERSATION_HISTORY)).isTrue();
1044+
}
1045+
1046+
@Test
1047+
void citationDocumentCacheBreakpointGoesOnLastDocumentOnly() {
1048+
AnthropicCitationDocument first = AnthropicCitationDocument.builder().plainText("First").build();
1049+
AnthropicCitationDocument second = AnthropicCitationDocument.builder().plainText("Second").build();
1050+
AnthropicChatOptions options = AnthropicChatOptions.builder()
1051+
.citationDocuments(first, second)
1052+
.cacheOptions(AnthropicCacheOptions.builder().strategy(AnthropicCacheStrategy.SYSTEM_ONLY).build())
1053+
.build();
1054+
1055+
MessageCreateParams request = this.chatModel.createRequest(new Prompt("Question", options), false);
1056+
1057+
List<ContentBlockParam> blocks = request.messages().get(0).content().blockParams().orElseThrow();
1058+
assertThat(blocks).hasSize(3);
1059+
assertThat(blocks.get(0).asDocument().cacheControl()).isEmpty();
1060+
assertThat(blocks.get(1).asDocument().cacheControl()).isPresent();
1061+
// SYSTEM_ONLY does not cache the user text
1062+
assertThat(blocks.get(2).asText().cacheControl()).isEmpty();
1063+
}
1064+
1065+
@Test
1066+
void citationDocumentAndUserTextBothCachedUnderConversationHistory() {
1067+
// Batch Q&A over the same document: each request is a fresh single-turn
1068+
// prompt, so the document breakpoint is what produces cache hits across
1069+
// requests while the user text breakpoint changes every time.
1070+
AnthropicCitationDocument document = AnthropicCitationDocument.builder().plainText("Reference").build();
1071+
AnthropicChatOptions options = AnthropicChatOptions.builder()
1072+
.citationDocuments(document)
1073+
.cacheOptions(AnthropicCacheOptions.builder().strategy(AnthropicCacheStrategy.CONVERSATION_HISTORY).build())
1074+
.build();
1075+
1076+
MessageCreateParams request = this.chatModel.createRequest(new Prompt("Question", options), false);
1077+
1078+
List<ContentBlockParam> blocks = request.messages().get(0).content().blockParams().orElseThrow();
1079+
assertThat(blocks).hasSize(2);
1080+
assertThat(blocks.get(0).asDocument().cacheControl()).isPresent();
1081+
assertThat(blocks.get(1).asText().cacheControl()).isPresent();
1082+
}
1083+
1084+
@Test
1085+
void citationDocumentBreakpointTakesPrecedenceOverToolDefinitions() {
1086+
// System, document, last user text, and last tool result each take a
1087+
// breakpoint, which is all four. Tool definitions are resolved last and get
1088+
// none. That is harmless: tools precede the system prompt in the request and
1089+
// are inside the prefix cached by the system breakpoint.
1090+
AnthropicCitationDocument document = AnthropicCitationDocument.builder().plainText("Reference").build();
1091+
AnthropicCacheOptions cacheOptions = AnthropicCacheOptions.builder()
1092+
.strategy(AnthropicCacheStrategy.CONVERSATION_HISTORY)
1093+
.cacheToolResults(true)
1094+
.build();
1095+
AnthropicChatOptions options = AnthropicChatOptions.builder()
1096+
.citationDocuments(document)
1097+
.cacheOptions(cacheOptions)
1098+
.toolCallbacks(List.of(new TestToolCallback("getWeather")))
1099+
.build();
1100+
1101+
List<org.springframework.ai.chat.messages.Message> messages = new java.util.ArrayList<>();
1102+
messages.add(new SystemMessage("You are a helpful assistant."));
1103+
messages.addAll(toolCallingConversation());
1104+
1105+
MessageCreateParams request = this.chatModel.createRequest(new Prompt(messages, options), false);
1106+
1107+
assertThat(request.system().orElseThrow().asTextBlockParams().get(0).cacheControl()).isPresent();
1108+
1109+
List<ContentBlockParam> firstUserBlocks = request.messages().get(0).content().blockParams().orElseThrow();
1110+
assertThat(firstUserBlocks.get(0).asDocument().cacheControl()).isPresent();
1111+
assertThat(firstUserBlocks.get(1).asText().cacheControl()).isPresent();
1112+
1113+
assertThat(lastToolResultBlock(request).cacheControl()).isPresent();
1114+
1115+
List<ToolUnion> tools = request.tools().orElseThrow();
1116+
assertThat(tools).hasSize(1);
1117+
assertThat(tools.get(0).asTool().cacheControl()).isEmpty();
1118+
}
1119+
1120+
private boolean lastDocumentHasCacheControl(AnthropicCacheStrategy strategy) {
1121+
AnthropicCitationDocument document = AnthropicCitationDocument.builder().plainText("Reference").build();
1122+
AnthropicChatOptions options = AnthropicChatOptions.builder()
1123+
.citationDocuments(document)
1124+
.cacheOptions(AnthropicCacheOptions.builder().strategy(strategy).build())
1125+
.build();
1126+
1127+
MessageCreateParams request = this.chatModel.createRequest(new Prompt("Question", options), false);
1128+
1129+
List<ContentBlockParam> blocks = request.messages().get(0).content().blockParams().orElseThrow();
1130+
return blocks.get(0).asDocument().cacheControl().isPresent();
1131+
}
1132+
9971133
static class TestToolCallback implements ToolCallback {
9981134

9991135
private final ToolDefinition toolDefinition;

‎models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/CacheEligibilityResolverTests.java‎

Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -255,4 +255,44 @@ void oneHourTtlReturnedForConfiguredMessageType() {
255255
assertThat(cc.ttl().get()).isEqualTo(CacheControlEphemeral.Ttl.TTL_1H);
256256
}
257257

258+
@Test
259+
void citationDocumentCacheControlRespectsStrategy() {
260+
CacheEligibilityResolver none = CacheEligibilityResolver
261+
.from(AnthropicCacheOptions.builder().strategy(AnthropicCacheStrategy.NONE).build());
262+
assertThat(none.resolveCitationDocumentCacheControl()).isNull();
263+
264+
// TOOLS_ONLY must not cache documents: a document breakpoint would also cache
265+
// the system prompt that precedes it.
266+
CacheEligibilityResolver toolsOnly = CacheEligibilityResolver
267+
.from(AnthropicCacheOptions.builder().strategy(AnthropicCacheStrategy.TOOLS_ONLY).build());
268+
assertThat(toolsOnly.resolveCitationDocumentCacheControl()).isNull();
269+
270+
CacheEligibilityResolver systemOnly = CacheEligibilityResolver
271+
.from(AnthropicCacheOptions.builder().strategy(AnthropicCacheStrategy.SYSTEM_ONLY).build());
272+
assertThat(systemOnly.resolveCitationDocumentCacheControl()).isNotNull();
273+
274+
// Documents use the SYSTEM TTL
275+
CacheEligibilityResolver sysAndTools = CacheEligibilityResolver.from(AnthropicCacheOptions.builder()
276+
.strategy(AnthropicCacheStrategy.SYSTEM_AND_TOOLS)
277+
.messageTypeTtl(MessageType.SYSTEM, AnthropicCacheTtl.ONE_HOUR)
278+
.build());
279+
CacheControlEphemeral cc = sysAndTools.resolveCitationDocumentCacheControl();
280+
assertThat(cc).isNotNull();
281+
assertThat(cc.ttl()).contains(CacheControlEphemeral.Ttl.TTL_1H);
282+
283+
CacheEligibilityResolver history = CacheEligibilityResolver
284+
.from(AnthropicCacheOptions.builder().strategy(AnthropicCacheStrategy.CONVERSATION_HISTORY).build());
285+
assertThat(history.resolveCitationDocumentCacheControl()).isNotNull();
286+
}
287+
288+
@Test
289+
void citationDocumentCacheControlRespectsBreakpointLimit() {
290+
CacheEligibilityResolver resolver = CacheEligibilityResolver
291+
.from(AnthropicCacheOptions.builder().strategy(AnthropicCacheStrategy.CONVERSATION_HISTORY).build());
292+
for (int i = 0; i < 4; i++) {
293+
resolver.useCacheBlock();
294+
}
295+
assertThat(resolver.resolveCitationDocumentCacheControl()).isNull();
296+
}
297+
258298
}

‎models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/chat/AnthropicPromptCachingIT.java‎

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@
3131
import org.springframework.ai.anthropic.AnthropicCacheTtl;
3232
import org.springframework.ai.anthropic.AnthropicChatModel;
3333
import org.springframework.ai.anthropic.AnthropicChatOptions;
34+
import org.springframework.ai.anthropic.AnthropicCitationDocument;
3435
import org.springframework.ai.anthropic.AnthropicTestConfiguration;
3536
import org.springframework.ai.chat.messages.AssistantMessage;
3637
import org.springframework.ai.chat.messages.Message;
@@ -460,4 +461,42 @@ void shouldCacheStaticPrefixWithMultiBlockSystemCaching() {
460461
.isTrue();
461462
}
462463

464+
@Test
465+
void shouldCachePdfCitationDocumentAcrossRequests() throws IOException {
466+
// Batch Q&A over one PDF: every request is a fresh single-turn prompt, so the
467+
// only breakpoint that can produce cache reads is the one on the document.
468+
AnthropicCitationDocument document = AnthropicCitationDocument.builder()
469+
.pdfFile("src/test/resources/spring-ai-reference-overview.pdf")
470+
.title("Spring AI Reference")
471+
.citationsEnabled(true)
472+
.build();
473+
474+
AnthropicChatOptions options = AnthropicChatOptions.builder()
475+
.model(Model.CLAUDE_SONNET_4_5.asString())
476+
.citationDocuments(document)
477+
.cacheOptions(AnthropicCacheOptions.builder().strategy(AnthropicCacheStrategy.SYSTEM_ONLY).build())
478+
.maxTokens(150)
479+
.temperature(0.0)
480+
.build();
481+
482+
ChatResponse first = this.chatModel
483+
.call(new Prompt(List.of(new UserMessage("Based only on the document, what is Spring AI?")), options));
484+
Usage firstUsage = getSdkUsage(first);
485+
assertThat(firstUsage).isNotNull();
486+
long firstCreation = firstUsage.cacheCreationInputTokens().orElse(0L);
487+
long firstRead = firstUsage.cacheReadInputTokens().orElse(0L);
488+
assertThat(firstCreation > 0 || firstRead > 0)
489+
.withFailMessage("Expected the document to be written to or read from cache, but got creation=%d, read=%d",
490+
firstCreation, firstRead)
491+
.isTrue();
492+
493+
ChatResponse second = this.chatModel.call(new Prompt(
494+
List.of(new UserMessage("Based only on the document, which models does it support?")), options));
495+
Usage secondUsage = getSdkUsage(second);
496+
assertThat(secondUsage).isNotNull();
497+
assertThat(secondUsage.cacheReadInputTokens().orElse(0L))
498+
.as("Second request should read the document from cache")
499+
.isGreaterThan(0);
500+
}
501+
463502
}

0 commit comments

Comments
 (0)