renaming and package refactoring
This commit is contained in:
@@ -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();
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
|
||||
|
||||
@@ -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.
|
||||
*
|
||||
|
||||
@@ -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<PromptContextTransformer> retrievers;
|
||||
private List<PromptTransformer> retrievers;
|
||||
|
||||
private List<PromptContextTransformer> documentPostProcessors;
|
||||
private List<PromptTransformer> documentPostProcessors;
|
||||
|
||||
private List<PromptContextTransformer> augmentors;
|
||||
private List<PromptTransformer> augmentors;
|
||||
|
||||
private List<ChatAgentListener> chatAgentListeners;
|
||||
|
||||
public DefaultChatAgent(ChatClient chatClient, List<PromptContextTransformer> retrievers,
|
||||
List<PromptContextTransformer> documentPostProcessors, List<PromptContextTransformer> augmentors,
|
||||
public DefaultChatAgent(ChatClient chatClient, List<PromptTransformer> retrievers,
|
||||
List<PromptTransformer> documentPostProcessors, List<PromptTransformer> augmentors,
|
||||
List<ChatAgentListener> 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<PromptContextTransformer> retrievers = new ArrayList<>();
|
||||
private List<PromptTransformer> retrievers = new ArrayList<>();
|
||||
|
||||
private List<PromptContextTransformer> documentPostProcessors = new ArrayList<>();
|
||||
private List<PromptTransformer> documentPostProcessors = new ArrayList<>();
|
||||
|
||||
private List<PromptContextTransformer> augmentors = new ArrayList<>();
|
||||
private List<PromptTransformer> augmentors = new ArrayList<>();
|
||||
|
||||
private List<ChatAgentListener> chatAgentListeners = new ArrayList<>();
|
||||
|
||||
@@ -81,18 +82,17 @@ public class DefaultChatAgent implements ChatAgent {
|
||||
return this;
|
||||
}
|
||||
|
||||
public DefaultChatAgentBuilder withRetrievers(List<PromptContextTransformer> retrievers) {
|
||||
public DefaultChatAgentBuilder withRetrievers(List<PromptTransformer> retrievers) {
|
||||
this.retrievers = retrievers;
|
||||
return this;
|
||||
}
|
||||
|
||||
public DefaultChatAgentBuilder withDocumentPostProcessors(
|
||||
List<PromptContextTransformer> documentPostProcessors) {
|
||||
public DefaultChatAgentBuilder withDocumentPostProcessors(List<PromptTransformer> documentPostProcessors) {
|
||||
this.documentPostProcessors = documentPostProcessors;
|
||||
return this;
|
||||
}
|
||||
|
||||
public DefaultChatAgentBuilder withAugmentors(List<PromptContextTransformer> augmentors) {
|
||||
public DefaultChatAgentBuilder withAugmentors(List<PromptTransformer> augmentors) {
|
||||
this.augmentors = augmentors;
|
||||
return this;
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
}
|
||||
@@ -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;
|
||||
@@ -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);
|
||||
|
||||
}
|
||||
@@ -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"
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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<Node<?>> dataList;
|
||||
private final List<Node<?>> 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<Node<?>> dataList, ChatResponse chatResponse) {
|
||||
this.prompt = prompt;
|
||||
|
||||
Reference in New Issue
Block a user