renaming and package refactoring

This commit is contained in:
Mark Pollack
2024-04-18 22:55:57 -04:00
parent d7b2028f01
commit 71c41341ef
10 changed files with 88 additions and 70 deletions

View File

@@ -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();
}
}
}

View File

@@ -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;

View File

@@ -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.
*

View File

@@ -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;
}

View File

@@ -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);
}

View File

@@ -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;

View File

@@ -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);
}

View File

@@ -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"

View File

@@ -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);
}

View File

@@ -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;