diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/agent/OpenAiDefaultChatAgentIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/agent/OpenAiDefaultChatAgentIT.java index 16d3c4905..1daf55a78 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/agent/OpenAiDefaultChatAgentIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/agent/OpenAiDefaultChatAgentIT.java @@ -3,12 +3,13 @@ package org.springframework.ai.openai.chat.agent; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.agent.ChatAgent; import org.springframework.ai.chat.agent.DefaultChatAgent; -import org.springframework.ai.chat.agent.PromptContext; -import org.springframework.ai.chat.agent.transformer.QAPromptContextTransformer; -import org.springframework.ai.chat.agent.transformer.VectorStorePromptContextTransformer; +import org.springframework.ai.chat.transformer.QuestionContextAugmentor; +import org.springframework.ai.chat.transformer.VectorStoreRetriever; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.chat.transformer.PromptContext; import org.springframework.ai.embedding.EmbeddingClient; import org.springframework.ai.evaluation.EvaluationRequest; import org.springframework.ai.evaluation.EvaluationResponse; @@ -41,32 +42,29 @@ public class OpenAiDefaultChatAgentIT { @Value("classpath:/data/acme/bikes.json") private Resource bikesResource; + private ChatAgent chatAgent; + @Autowired - public OpenAiDefaultChatAgentIT(ChatClient chatClient, VectorStore vectorStore) { + public OpenAiDefaultChatAgentIT(ChatClient chatClient, ChatAgent chatAgent, VectorStore vectorStore) { this.chatClient = chatClient; + this.chatAgent = chatAgent; this.vectorStore = vectorStore; } @Test void simpleChat() { loadData(); - var chatAgent = DefaultChatAgent.builder(chatClient) - .withRetrievers(List.of(new VectorStorePromptContextTransformer(vectorStore, SearchRequest.defaults()))) - .withAugmentors(List.of(new QAPromptContextTransformer())) - .build(); - Prompt prompt = new Prompt(new UserMessage("What bike is good for city commuting?")); + var prompt = new Prompt(new UserMessage("What bike is good for city commuting?")); PromptContext promptContext = new PromptContext(prompt); - var agentResponse = chatAgent.call(promptContext); + var agentResponse = this.chatAgent.call(promptContext); System.out.println(agentResponse.getChatResponse().getResult().getOutput().getContent()); - RelevancyEvaluator relevancyEvaluator = new RelevancyEvaluator(this.chatClient); - EvaluationRequest evaluationRequest = new EvaluationRequest( - agentResponse.getPromptContext().getPromptHistory().get(0), agentResponse.getPromptContext().getNodes(), - agentResponse.getChatResponse()); - + var relevancyEvaluator = new RelevancyEvaluator(this.chatClient); + EvaluationRequest evaluationRequest = new EvaluationRequest(agentResponse); EvaluationResponse evaluationResponse = relevancyEvaluator.evaluate(evaluationRequest); + System.out.println(evaluationResponse); } @@ -100,6 +98,15 @@ public class OpenAiDefaultChatAgentIT { return new SimpleVectorStore(embeddingClient); } + @Bean + public ChatAgent chatagent(ChatClient chatClient, VectorStore vectorStore) { + return DefaultChatAgent.builder(chatClient) + .withRetrievers(List.of(new VectorStoreRetriever(vectorStore, SearchRequest.defaults()))) + .withAugmentors(List.of(new QuestionContextAugmentor())) + .build(); + + } + } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/AgentResponse.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/AgentResponse.java index 0098f79b4..112caaf51 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/AgentResponse.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/AgentResponse.java @@ -1,6 +1,7 @@ package org.springframework.ai.chat.agent; import org.springframework.ai.chat.ChatResponse; +import org.springframework.ai.chat.transformer.PromptContext; import java.util.Objects; diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/ChatAgent.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/ChatAgent.java index c68166b22..20100fb11 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/ChatAgent.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/ChatAgent.java @@ -1,5 +1,7 @@ package org.springframework.ai.chat.agent; +import org.springframework.ai.chat.transformer.PromptContext; + /** * A ChatAgent encapsulates common AI workflows such as Retrieval Augmented Generation. * diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/DefaultChatAgent.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/DefaultChatAgent.java index 860706711..59cd26a68 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/DefaultChatAgent.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/DefaultChatAgent.java @@ -2,7 +2,8 @@ package org.springframework.ai.chat.agent; import org.springframework.ai.chat.ChatClient; import org.springframework.ai.chat.ChatResponse; -import org.springframework.ai.chat.agent.transformer.PromptContextTransformer; +import org.springframework.ai.chat.transformer.PromptContext; +import org.springframework.ai.chat.transformer.PromptTransformer; import java.util.ArrayList; import java.util.List; @@ -12,16 +13,16 @@ public class DefaultChatAgent implements ChatAgent { private ChatClient chatClient; - private List retrievers; + private List retrievers; - private List documentPostProcessors; + private List documentPostProcessors; - private List augmentors; + private List augmentors; private List chatAgentListeners; - public DefaultChatAgent(ChatClient chatClient, List retrievers, - List documentPostProcessors, List augmentors, + public DefaultChatAgent(ChatClient chatClient, List retrievers, + List documentPostProcessors, List augmentors, List chatAgentListeners) { Objects.requireNonNull(chatClient, "chatClient must not be null"); this.chatClient = chatClient; @@ -39,17 +40,17 @@ public class DefaultChatAgent implements ChatAgent { public AgentResponse call(PromptContext promptContext) { // Perform retrieval of documents and messages - for (PromptContextTransformer retriever : retrievers) { + for (PromptTransformer retriever : retrievers) { promptContext = retriever.transform(promptContext); } // Perform post procesing of all retrieved documents and messages - for (PromptContextTransformer documentPostProcessor : documentPostProcessors) { + for (PromptTransformer documentPostProcessor : documentPostProcessors) { promptContext = documentPostProcessor.transform(promptContext); } // Perform prompt augmentation - for (PromptContextTransformer augmentor : augmentors) { + for (PromptTransformer augmentor : augmentors) { promptContext = augmentor.transform(promptContext); } @@ -68,11 +69,11 @@ public class DefaultChatAgent implements ChatAgent { private ChatClient chatClient; - private List retrievers = new ArrayList<>(); + private List retrievers = new ArrayList<>(); - private List documentPostProcessors = new ArrayList<>(); + private List documentPostProcessors = new ArrayList<>(); - private List augmentors = new ArrayList<>(); + private List augmentors = new ArrayList<>(); private List chatAgentListeners = new ArrayList<>(); @@ -81,18 +82,17 @@ public class DefaultChatAgent implements ChatAgent { return this; } - public DefaultChatAgentBuilder withRetrievers(List retrievers) { + public DefaultChatAgentBuilder withRetrievers(List retrievers) { this.retrievers = retrievers; return this; } - public DefaultChatAgentBuilder withDocumentPostProcessors( - List documentPostProcessors) { + public DefaultChatAgentBuilder withDocumentPostProcessors(List documentPostProcessors) { this.documentPostProcessors = documentPostProcessors; return this; } - public DefaultChatAgentBuilder withAugmentors(List augmentors) { + public DefaultChatAgentBuilder withAugmentors(List augmentors) { this.augmentors = augmentors; return this; } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/transformer/PromptContextTransformer.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/transformer/PromptContextTransformer.java deleted file mode 100644 index 2c76f089f..000000000 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/transformer/PromptContextTransformer.java +++ /dev/null @@ -1,22 +0,0 @@ -package org.springframework.ai.chat.agent.transformer; - -import org.springframework.ai.chat.agent.PromptContext; - -/** - * Transforms the PromptContext. Implementations may retrieve data and modify the Prompt - * as needed. - * - * @author Mark Pollack - * @since 1.0 M1 - */ -@FunctionalInterface -public interface PromptContextTransformer { - - /** - * Transforms the given PromptContext. - * @param context the PromptContext to transform - * @return the transformed PromptContext - */ - PromptContext transform(PromptContext context); - -} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/PromptContext.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/transformer/PromptContext.java similarity index 97% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/agent/PromptContext.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/transformer/PromptContext.java index 0e5e19acc..706ffcdbf 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/PromptContext.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/transformer/PromptContext.java @@ -1,4 +1,4 @@ -package org.springframework.ai.chat.agent; +package org.springframework.ai.chat.transformer; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.node.Node; diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/transformer/PromptTransformer.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/transformer/PromptTransformer.java new file mode 100644 index 000000000..c508bc6e1 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/transformer/PromptTransformer.java @@ -0,0 +1,23 @@ +package org.springframework.ai.chat.transformer; + +/** + * Responsible for transforming a Prompt. The PromptContext contains the necessary data to + * make the transformation + * + * Implementations may retrieve data and modify the Prompt object in the PromptContext as + * needed. + * + * @author Mark Pollack + * @since 1.0 M1 + */ +@FunctionalInterface +public interface PromptTransformer { + + /** + * Transforms the given PromptContext. + * @param context the PromptContext to transform + * @return the transformed PromptContext + */ + PromptContext transform(PromptContext context); + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/transformer/QAPromptContextTransformer.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/transformer/QuestionContextAugmentor.java similarity index 79% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/agent/transformer/QAPromptContextTransformer.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/transformer/QuestionContextAugmentor.java index a7edebc46..6712344ad 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/transformer/QAPromptContextTransformer.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/transformer/QuestionContextAugmentor.java @@ -1,11 +1,12 @@ -package org.springframework.ai.chat.agent.transformer; +package org.springframework.ai.chat.transformer; -import org.springframework.ai.chat.agent.PromptContext; +import org.springframework.ai.chat.transformer.PromptContext; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.chat.prompt.PromptTemplate; +import org.springframework.ai.chat.transformer.PromptTransformer; import org.springframework.ai.document.Document; import org.springframework.ai.node.Node; @@ -14,10 +15,13 @@ import java.util.Map; import java.util.stream.Collectors; /** - * Transforms the PromptContext by adding to the prompt a Question and Answer text that - * contains the placeholder names "context" and "question". + * Transforms the Prompt by taking to the current prompt in the Prompt Context and adding + * additional context to create a new prompt. The default user text contains the + * placeholder names "question" and "context". The "question" placeholder is filled using + * the value of the current UserMessage and the "context" placeholder is filled with + * Documents contained in the PromptContext's Nodes. */ -public class QAPromptContextTransformer implements PromptContextTransformer { +public class QuestionContextAugmentor implements PromptTransformer { private static final String DEFAULT_USER_PROMPT_TEXT = """ "Context information is below.\\n" diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/transformer/VectorStorePromptContextTransformer.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/transformer/VectorStoreRetriever.java similarity index 76% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/agent/transformer/VectorStorePromptContextTransformer.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/transformer/VectorStoreRetriever.java index c8717e3ff..625d53594 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/agent/transformer/VectorStorePromptContextTransformer.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/transformer/VectorStoreRetriever.java @@ -1,6 +1,5 @@ -package org.springframework.ai.chat.agent.transformer; +package org.springframework.ai.chat.transformer; -import org.springframework.ai.chat.agent.PromptContext; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.document.Document; @@ -14,13 +13,13 @@ import java.util.stream.Collectors; /** * Transforms the PromptContext by retrieving documents from a VectorStore */ -public class VectorStorePromptContextTransformer implements PromptContextTransformer { +public class VectorStoreRetriever implements PromptTransformer { private final VectorStore vectorStore; private final SearchRequest searchRequest; - public VectorStorePromptContextTransformer(VectorStore vectorStore, SearchRequest searchRequest) { + public VectorStoreRetriever(VectorStore vectorStore, SearchRequest searchRequest) { this.vectorStore = vectorStore; this.searchRequest = searchRequest; } @@ -49,15 +48,14 @@ public class VectorStorePromptContextTransformer implements PromptContextTransfo @Override public String toString() { - return "VectorStorePromptContextTransformer{" + "vectorStore=" + vectorStore + ", searchRequest=" - + searchRequest + '}'; + return "VectorStoreRetriever{" + "vectorStore=" + vectorStore + ", searchRequest=" + searchRequest + '}'; } @Override public boolean equals(Object o) { if (this == o) return true; - if (!(o instanceof VectorStorePromptContextTransformer that)) + if (!(o instanceof VectorStoreRetriever that)) return false; return Objects.equals(vectorStore, that.vectorStore) && Objects.equals(searchRequest, that.searchRequest); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/evaluation/EvaluationRequest.java b/spring-ai-core/src/main/java/org/springframework/ai/evaluation/EvaluationRequest.java index 1081f37c6..9eec5e244 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/evaluation/EvaluationRequest.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/evaluation/EvaluationRequest.java @@ -1,20 +1,25 @@ package org.springframework.ai.evaluation; import org.springframework.ai.chat.ChatResponse; +import org.springframework.ai.chat.agent.AgentResponse; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.node.Node; -import java.util.ArrayList; import java.util.List; import java.util.Objects; public class EvaluationRequest { - private Prompt prompt; + private final Prompt prompt; - private List> dataList; + private final List> dataList; - private ChatResponse chatResponse; + private final ChatResponse chatResponse; + + public EvaluationRequest(AgentResponse agentResponse) { + this(agentResponse.getPromptContext().getPromptHistory().get(0), agentResponse.getPromptContext().getNodes(), + agentResponse.getChatResponse()); + } public EvaluationRequest(Prompt prompt, List> dataList, ChatResponse chatResponse) { this.prompt = prompt;