fix content purging logic
This commit is contained in:
@@ -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<Document> apply(List<Document> documents) {
|
||||
documents.forEach(d -> {
|
||||
Map<String, Object> 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, ""))))
|
||||
|
||||
@@ -22,6 +22,7 @@ import org.springframework.ai.chat.messages.Message;
|
||||
|
||||
/**
|
||||
* @author Christian Tzolov
|
||||
*
|
||||
*/
|
||||
public interface ChatMemory {
|
||||
|
||||
|
||||
@@ -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<Content> 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);
|
||||
|
||||
@@ -56,13 +56,20 @@ public class LastMaxTokenSizeContentTransformer implements PromptTransformer {
|
||||
this.filterTags = filterTags;
|
||||
}
|
||||
|
||||
protected List<Content> doGetDatum(PromptContext promptContext) {
|
||||
protected List<Content> doGetDatumToModify(PromptContext promptContext) {
|
||||
return promptContext.getContents()
|
||||
.stream()
|
||||
.filter(content -> this.filterTags.stream().allMatch(tag -> content.getMetadata().containsKey(tag)))
|
||||
.toList();
|
||||
}
|
||||
|
||||
protected List<Content> 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<Content> datum = this.doGetDatum(promptContext);
|
||||
List<Content> datum = this.doGetDatumToModify(promptContext);
|
||||
|
||||
// int totalSize = this.tokenCountEstimator.estimate(nonSystemChatMessages) -
|
||||
// retrievalRequest.getTokenRunningTotal();
|
||||
@@ -84,9 +91,12 @@ public class LastMaxTokenSizeContentTransformer implements PromptTransformer {
|
||||
return promptContext;
|
||||
}
|
||||
|
||||
List<Content> newSessionMessages = this.purgeExcess(datum, totalSize);
|
||||
List<Content> 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<Content> purgeExcess(List<Content> datum, int totalSize) {
|
||||
|
||||
@@ -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> media;
|
||||
|
||||
private final Map<String, Object> metadata;
|
||||
|
||||
public InnerContent(String content) {
|
||||
this(content, Map.of());
|
||||
}
|
||||
|
||||
public InnerContent(String content, Map<String, Object> metadata) {
|
||||
this(content, List.of(), metadata);
|
||||
}
|
||||
|
||||
public InnerContent(String content, List<Media> media, Map<String, Object> metadata) {
|
||||
this.content = content;
|
||||
this.media = media;
|
||||
this.metadata = metadata;
|
||||
}
|
||||
|
||||
@Override
|
||||
public String getContent() {
|
||||
return this.content;
|
||||
}
|
||||
|
||||
@Override
|
||||
public List<Media> getMedia() {
|
||||
return this.media;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Map<String, Object> getMetadata() {
|
||||
return this.metadata;
|
||||
}
|
||||
|
||||
@Override
|
||||
public String toString() {
|
||||
return "InnerContent [content=" + content + ", media=" + media + ", metadata=" + metadata + "]";
|
||||
}
|
||||
|
||||
}
|
||||
@@ -62,7 +62,7 @@ public class QuestionContextAugmentor implements PromptTransformer {
|
||||
|
||||
protected String doCreateContext(List<Content> 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()));
|
||||
}
|
||||
|
||||
@@ -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";
|
||||
|
||||
}
|
||||
|
||||
@@ -56,20 +56,13 @@ public class VectorStoreRetriever implements PromptTransformer {
|
||||
.map(m -> m.getContent())
|
||||
.collect(Collectors.joining(System.lineSeparator()));
|
||||
|
||||
List<Document> documents = vectorStore.similaritySearch(searchRequest.withQuery(userMessage));
|
||||
List<Document> 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;
|
||||
}
|
||||
|
||||
@@ -79,6 +79,10 @@ public class Document implements Content {
|
||||
this(content, metadata, new RandomIdGenerator());
|
||||
}
|
||||
|
||||
public Document(String content, List<Media> media, Map<String, Object> metadata) {
|
||||
this(new RandomIdGenerator().generateId(content, metadata), content, media, metadata);
|
||||
}
|
||||
|
||||
public Document(String content, Map<String, Object> metadata, IdGenerator idGenerator) {
|
||||
this(idGenerator.generateId(content, metadata), content, metadata);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user