diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/agent/TextChatHistoryChatAgent3IT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/agent/TextChatHistoryChatAgent3IT.java index 24b5662a9..4af9361f5 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/agent/TextChatHistoryChatAgent3IT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/agent/TextChatHistoryChatAgent3IT.java @@ -46,6 +46,8 @@ import org.springframework.ai.chat.prompt.transformer.PromptContext; import org.springframework.ai.chat.prompt.transformer.QuestionContextAugmentor; import org.springframework.ai.chat.prompt.transformer.TransformerContentType; import org.springframework.ai.chat.prompt.transformer.VectorStoreRetriever; +import org.springframework.ai.document.Document; +import org.springframework.ai.document.DocumentTransformer; import org.springframework.ai.embedding.EmbeddingClient; import org.springframework.ai.evaluation.EvaluationRequest; import org.springframework.ai.evaluation.EvaluationResponse; @@ -96,9 +98,24 @@ public class TextChatHistoryChatAgent3IT { private Resource bikesResource; void loadData() { + + var metadataEnricher = new DocumentTransformer() { + + @Override + public List apply(List documents) { + documents.forEach(d -> { + Map metadata = d.getMetadata(); + metadata.put(TransformerContentType.EXTERNAL_KNOWLEDGE, "true"); + }); + + return documents; + } + + }; + JsonReader jsonReader = new JsonReader(bikesResource, "name", "price", "shortDescription", "description"); var textSplitter = new TokenTextSplitter(); - vectorStore.accept(textSplitter.apply(jsonReader.get())); + vectorStore.accept(metadataEnricher.apply(textSplitter.apply(jsonReader.get()))); } // @Autowired @@ -116,8 +133,8 @@ public class TextChatHistoryChatAgent3IT { logger.info("Response1: " + agentResponse1.getChatResponse().getResult().getOutput().getContent()); - var agentResponse2 = this.chatAgent.call( - new PromptContext(new Prompt(new String("What is my name and what bike model would suggest for me?")))); + var agentResponse2 = this.chatAgent.call(new PromptContext( + new Prompt(new String("What is my name and what bike model would you suggest for me?")))); logger.info("Response2: " + agentResponse2.getChatResponse().getResult().getOutput().getContent()); logger.info(agentResponse2.getPromptContext().getContents().toString()); @@ -173,13 +190,14 @@ public class TextChatHistoryChatAgent3IT { new ChatMemoryRetriever(chatHistory, Map.of(TransformerContentType.SHORT_TERM_MEMORY, "")), new VectorStoreChatMemoryRetriever(vectorStore, 10, Map.of(TransformerContentType.LONG_TERM_MEMORY, "")))) + .withDocumentPostProcessors(List.of( new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 1000, Set.of(TransformerContentType.SHORT_TERM_MEMORY)), new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 1000, Set.of(TransformerContentType.LONG_TERM_MEMORY)), new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 2000, - Set.of(TransformerContentType.QA)))) + Set.of(TransformerContentType.EXTERNAL_KNOWLEDGE)))) .withAugmentors(List.of(new QuestionContextAugmentor(), new SystemPromptChatMemoryAugmentor( """ @@ -190,6 +208,7 @@ public class TextChatHistoryChatAgent3IT { """, Set.of(TransformerContentType.LONG_TERM_MEMORY)), new SystemPromptChatMemoryAugmentor(Set.of(TransformerContentType.SHORT_TERM_MEMORY)))) + .withChatAgentListeners(List.of(new ChatMemoryAgentListener(chatHistory), new VectorStoreChatMemoryAgentListener(vectorStore, Map.of(TransformerContentType.LONG_TERM_MEMORY, "")))) diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/ChatMemory.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/history/ChatMemory.java index 798b88ada..8714d11c2 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/ChatMemory.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/history/ChatMemory.java @@ -22,6 +22,7 @@ import org.springframework.ai.chat.messages.Message; /** * @author Christian Tzolov + * */ public interface ChatMemory { diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/ChatMemoryRetriever.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/history/ChatMemoryRetriever.java index 6671c41aa..57fe7a1cf 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/ChatMemoryRetriever.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/history/ChatMemoryRetriever.java @@ -23,10 +23,10 @@ import java.util.Map; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.MessageType; -import org.springframework.ai.chat.prompt.transformer.TransformerContentType; -import org.springframework.ai.chat.prompt.transformer.InnerContent; import org.springframework.ai.chat.prompt.transformer.PromptContext; import org.springframework.ai.chat.prompt.transformer.PromptTransformer; +import org.springframework.ai.chat.prompt.transformer.TransformerContentType; +import org.springframework.ai.document.Document; import org.springframework.ai.model.Content; /** @@ -57,7 +57,7 @@ public class ChatMemoryRetriever implements PromptTransformer { List historyContent = (messageHistory != null) ? messageHistory.stream().filter(m -> m.getMessageType() != MessageType.SYSTEM).map(m -> { - Content content = new InnerContent(m.getContent(), new ArrayList<>(m.getMedia()), + Content content = new Document(m.getContent(), new ArrayList<>(m.getMedia()), new HashMap<>(m.getMetadata())); content.getMetadata().putAll(this.additionalMetadata); content.getMetadata().put(TransformerContentType.MEMORY, true); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/LastMaxTokenSizeContentTransformer.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/history/LastMaxTokenSizeContentTransformer.java index 67b579330..8788c64d5 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/LastMaxTokenSizeContentTransformer.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/history/LastMaxTokenSizeContentTransformer.java @@ -56,13 +56,20 @@ public class LastMaxTokenSizeContentTransformer implements PromptTransformer { this.filterTags = filterTags; } - protected List doGetDatum(PromptContext promptContext) { + protected List doGetDatumToModify(PromptContext promptContext) { return promptContext.getContents() .stream() .filter(content -> this.filterTags.stream().allMatch(tag -> content.getMetadata().containsKey(tag))) .toList(); } + protected List doGetDatumNotToModify(PromptContext promptContext) { + return promptContext.getContents() + .stream() + .filter(content -> !this.filterTags.stream().allMatch(tag -> content.getMetadata().containsKey(tag))) + .toList(); + } + protected int doEstimateTokenCount(Content datum) { return this.tokenCountEstimator.estimate(datum); } @@ -74,7 +81,7 @@ public class LastMaxTokenSizeContentTransformer implements PromptTransformer { @Override public PromptContext transform(PromptContext promptContext) { - List datum = this.doGetDatum(promptContext); + List datum = this.doGetDatumToModify(promptContext); // int totalSize = this.tokenCountEstimator.estimate(nonSystemChatMessages) - // retrievalRequest.getTokenRunningTotal(); @@ -84,9 +91,12 @@ public class LastMaxTokenSizeContentTransformer implements PromptTransformer { return promptContext; } - List newSessionMessages = this.purgeExcess(datum, totalSize); + List purgedContent = this.purgeExcess(datum, totalSize); - return PromptContext.from(promptContext).withContents(newSessionMessages).build(); + var updatedContent = new ArrayList<>(doGetDatumNotToModify(promptContext)); + updatedContent.addAll(purgedContent); + + return PromptContext.from(promptContext).withContents(updatedContent).build(); } protected List purgeExcess(List datum, int totalSize) { diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/InnerContent.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/InnerContent.java deleted file mode 100644 index 6f02f7785..000000000 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/InnerContent.java +++ /dev/null @@ -1,70 +0,0 @@ -/* - * Copyright 2024-2024 the original author or authors. - * - * 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 - * - * https://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 org.springframework.ai.chat.prompt.transformer; - -import java.util.List; -import java.util.Map; - -import org.springframework.ai.chat.messages.Media; -import org.springframework.ai.model.Content; - -/** - * @author Christian Tzolov - */ -public class InnerContent implements Content { - - private final String content; - - private final List media; - - private final Map metadata; - - public InnerContent(String content) { - this(content, Map.of()); - } - - public InnerContent(String content, Map metadata) { - this(content, List.of(), metadata); - } - - public InnerContent(String content, List media, Map metadata) { - this.content = content; - this.media = media; - this.metadata = metadata; - } - - @Override - public String getContent() { - return this.content; - } - - @Override - public List getMedia() { - return this.media; - } - - @Override - public Map getMetadata() { - return this.metadata; - } - - @Override - public String toString() { - return "InnerContent [content=" + content + ", media=" + media + ", metadata=" + metadata + "]"; - } - -} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/QuestionContextAugmentor.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/QuestionContextAugmentor.java index 9219ad0dc..7a99d9eb8 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/QuestionContextAugmentor.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/QuestionContextAugmentor.java @@ -62,7 +62,7 @@ public class QuestionContextAugmentor implements PromptTransformer { protected String doCreateContext(List data) { return data.stream() - .filter(content -> content.getMetadata().containsKey(TransformerContentType.QA)) + .filter(content -> content.getMetadata().containsKey(TransformerContentType.EXTERNAL_KNOWLEDGE)) .map(Content::getContent) .collect(Collectors.joining(System.lineSeparator())); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/TransformerContentType.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/TransformerContentType.java index 8b56cbc9b..270e85943 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/TransformerContentType.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/TransformerContentType.java @@ -29,6 +29,6 @@ public class TransformerContentType { public static final String CONVERSATION_ID = "conversationId"; - public static final String QA = "QA"; + public static final String EXTERNAL_KNOWLEDGE = "externalKnowledge"; } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/VectorStoreRetriever.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/VectorStoreRetriever.java index cffca5f50..3a62daa77 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/VectorStoreRetriever.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/VectorStoreRetriever.java @@ -56,20 +56,13 @@ public class VectorStoreRetriever implements PromptTransformer { .map(m -> m.getContent()) .collect(Collectors.joining(System.lineSeparator())); - List documents = vectorStore.similaritySearch(searchRequest.withQuery(userMessage)); + List documents = vectorStore.similaritySearch(searchRequest.withQuery(userMessage) + .withFilterExpression(TransformerContentType.EXTERNAL_KNOWLEDGE + "=='true'")); for (Document document : documents) { - if (!document.getMetadata().containsKey(TransformerContentType.MEMORY)) { // TODO: - // Bad - // coupling - // with - // other - // transformers - // types. - var content = new InnerContent(document.getContent(), document.getMetadata()); - content.getMetadata().put(TransformerContentType.QA, true); - promptContext.addData(content); - } + var content = new Document(document.getContent(), document.getMetadata()); + // content.getMetadata().put(TransformerContentType.DOMAIN_DATA, true); + promptContext.addData(content); } return promptContext; } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/document/Document.java b/spring-ai-core/src/main/java/org/springframework/ai/document/Document.java index 5f2ec9ebc..30c4b479a 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/document/Document.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/document/Document.java @@ -79,6 +79,10 @@ public class Document implements Content { this(content, metadata, new RandomIdGenerator()); } + public Document(String content, List media, Map metadata) { + this(new RandomIdGenerator().generateId(content, metadata), content, media, metadata); + } + public Document(String content, Map metadata, IdGenerator idGenerator) { this(idGenerator.generateId(content, metadata), content, metadata); }