diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/ChatMemoryLongTermSystemPromptIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryLongTermSystemPromptIT.java similarity index 73% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/ChatMemoryLongTermSystemPromptIT.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryLongTermSystemPromptIT.java index 237c1d56e..4516e08ae 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/ChatMemoryLongTermSystemPromptIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryLongTermSystemPromptIT.java @@ -14,25 +14,25 @@ * limitations under the License. */ -package org.springframework.ai.openai.chat.chatbot; +package org.springframework.ai.openai.chat.service; import java.util.List; import io.qdrant.client.QdrantClient; import io.qdrant.client.QdrantGrpcClient; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; -import org.springframework.ai.chat.chatbot.ChatBot; +import org.springframework.ai.chat.service.ChatService; +import org.springframework.ai.chat.service.StreamingChatService; import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.qdrant.QdrantContainer; -import org.springframework.ai.chat.chatbot.DefaultChatBot; -import org.springframework.ai.chat.chatbot.DefaultStreamingChatBot; -import org.springframework.ai.chat.chatbot.StreamingChatBot; -import org.springframework.ai.chat.history.VectorStoreChatMemoryChatBotListener; -import org.springframework.ai.chat.history.VectorStoreChatMemoryRetriever; -import org.springframework.ai.chat.history.LastMaxTokenSizeContentTransformer; -import org.springframework.ai.chat.history.SystemPromptChatMemoryAugmentor; +import org.springframework.ai.chat.service.PromptTransformingChatService; +import org.springframework.ai.chat.service.StreamingPromptTransformingChatService; +import org.springframework.ai.chat.memory.VectorStoreChatMemoryChatServiceListener; +import org.springframework.ai.chat.memory.VectorStoreChatMemoryRetriever; +import org.springframework.ai.chat.memory.LastMaxTokenSizeContentTransformer; +import org.springframework.ai.chat.memory.SystemPromptChatMemoryAugmentor; import org.springframework.ai.embedding.EmbeddingClient; import org.springframework.ai.evaluation.BaseMemoryTest; import org.springframework.ai.evaluation.RelevancyEvaluator; @@ -61,9 +61,9 @@ public class ChatMemoryLongTermSystemPromptIT extends BaseMemoryTest { static QdrantContainer qdrantContainer = new QdrantContainer("qdrant/qdrant:v1.9.2"); @Autowired - public ChatMemoryLongTermSystemPromptIT(RelevancyEvaluator relevancyEvaluator, ChatBot chatBot, - StreamingChatBot streamingChatBot) { - super(relevancyEvaluator, chatBot, streamingChatBot); + public ChatMemoryLongTermSystemPromptIT(RelevancyEvaluator relevancyEvaluator, ChatService chatService, + StreamingChatService streamingChatService) { + super(relevancyEvaluator, chatService, streamingChatService); } @SpringBootConfiguration @@ -98,26 +98,26 @@ public class ChatMemoryLongTermSystemPromptIT extends BaseMemoryTest { } @Bean - public ChatBot memoryChatBot(OpenAiChatClient chatClient, VectorStore vectorStore, + public ChatService memoryChatService(OpenAiChatClient chatClient, VectorStore vectorStore, TokenCountEstimator tokenCountEstimator) { - return DefaultChatBot.builder(chatClient) + return PromptTransformingChatService.builder(chatClient) .withRetrievers(List.of(new VectorStoreChatMemoryRetriever(vectorStore, 10))) .withContentPostProcessors(List.of(new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 1000))) .withAugmentors(List.of(new SystemPromptChatMemoryAugmentor())) - .withChatBotListeners(List.of(new VectorStoreChatMemoryChatBotListener(vectorStore))) + .withChatServiceListeners(List.of(new VectorStoreChatMemoryChatServiceListener(vectorStore))) .build(); } @Bean - public StreamingChatBot memoryStreamingChatBot(OpenAiChatClient streamingChatClient, VectorStore vectorStore, - TokenCountEstimator tokenCountEstimator) { + public StreamingChatService memoryStreamingChatService(OpenAiChatClient streamingChatClient, + VectorStore vectorStore, TokenCountEstimator tokenCountEstimator) { - return DefaultStreamingChatBot.builder(streamingChatClient) + return StreamingPromptTransformingChatService.builder(streamingChatClient) .withRetrievers(List.of(new VectorStoreChatMemoryRetriever(vectorStore, 10))) .withDocumentPostProcessors(List.of(new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 1000))) .withAugmentors(List.of(new SystemPromptChatMemoryAugmentor())) - .withChatBotListeners(List.of(new VectorStoreChatMemoryChatBotListener(vectorStore))) + .withChatServiceListeners(List.of(new VectorStoreChatMemoryChatServiceListener(vectorStore))) .build(); } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/ChatMemoryShortTermMessageListIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryShortTermMessageListIT.java similarity index 65% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/ChatMemoryShortTermMessageListIT.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryShortTermMessageListIT.java index 5d88ce34b..d26f6c563 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/ChatMemoryShortTermMessageListIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryShortTermMessageListIT.java @@ -13,22 +13,22 @@ * See the License for the specific language governing permissions and * limitations under the License. */ -package org.springframework.ai.openai.chat.chatbot; +package org.springframework.ai.openai.chat.service; import java.util.List; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; -import org.springframework.ai.chat.chatbot.ChatBot; -import org.springframework.ai.chat.chatbot.DefaultChatBot; -import org.springframework.ai.chat.chatbot.DefaultStreamingChatBot; -import org.springframework.ai.chat.chatbot.StreamingChatBot; -import org.springframework.ai.chat.history.ChatMemory; -import org.springframework.ai.chat.history.ChatMemoryChatBotListener; -import org.springframework.ai.chat.history.ChatMemoryRetriever; -import org.springframework.ai.chat.history.InMemoryChatMemory; -import org.springframework.ai.chat.history.LastMaxTokenSizeContentTransformer; -import org.springframework.ai.chat.history.MessageChatMemoryAugmentor; +import org.springframework.ai.chat.service.ChatService; +import org.springframework.ai.chat.service.PromptTransformingChatService; +import org.springframework.ai.chat.service.StreamingPromptTransformingChatService; +import org.springframework.ai.chat.service.StreamingChatService; +import org.springframework.ai.chat.memory.ChatMemory; +import org.springframework.ai.chat.memory.ChatMemoryChatServiceListener; +import org.springframework.ai.chat.memory.ChatMemoryRetriever; +import org.springframework.ai.chat.memory.InMemoryChatMemory; +import org.springframework.ai.chat.memory.LastMaxTokenSizeContentTransformer; +import org.springframework.ai.chat.memory.MessageChatMemoryAugmentor; import org.springframework.ai.evaluation.BaseMemoryTest; import org.springframework.ai.evaluation.RelevancyEvaluator; import org.springframework.ai.openai.OpenAiChatClient; @@ -45,9 +45,9 @@ import org.springframework.context.annotation.Bean; public class ChatMemoryShortTermMessageListIT extends BaseMemoryTest { @Autowired - public ChatMemoryShortTermMessageListIT(RelevancyEvaluator relevancyEvaluator, ChatBot chatBot, - StreamingChatBot streamingChatBot) { - super(relevancyEvaluator, chatBot, streamingChatBot); + public ChatMemoryShortTermMessageListIT(RelevancyEvaluator relevancyEvaluator, ChatService chatService, + StreamingChatService streamingChatService) { + super(relevancyEvaluator, chatService, streamingChatService); } @SpringBootConfiguration @@ -74,26 +74,26 @@ public class ChatMemoryShortTermMessageListIT extends BaseMemoryTest { } @Bean - public ChatBot memoryChatBot(OpenAiChatClient chatClient, ChatMemory chatHistory, + public ChatService memoryChatService(OpenAiChatClient chatClient, ChatMemory chatHistory, TokenCountEstimator tokenCountEstimator) { - return DefaultChatBot.builder(chatClient) + return PromptTransformingChatService.builder(chatClient) .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) .withContentPostProcessors(List.of(new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 1000))) .withAugmentors(List.of(new MessageChatMemoryAugmentor())) - .withChatBotListeners(List.of(new ChatMemoryChatBotListener(chatHistory))) + .withChatServiceListeners(List.of(new ChatMemoryChatServiceListener(chatHistory))) .build(); } @Bean - public StreamingChatBot memoryStreamingChatBot(OpenAiChatClient streamingChatClient, ChatMemory chatHistory, - TokenCountEstimator tokenCountEstimator) { + public StreamingChatService memoryStreamingChatService(OpenAiChatClient streamingChatClient, + ChatMemory chatHistory, TokenCountEstimator tokenCountEstimator) { - return DefaultStreamingChatBot.builder(streamingChatClient) + return StreamingPromptTransformingChatService.builder(streamingChatClient) .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) .withDocumentPostProcessors(List.of(new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 1000))) .withAugmentors(List.of(new MessageChatMemoryAugmentor())) - .withChatBotListeners(List.of(new ChatMemoryChatBotListener(chatHistory))) + .withChatServiceListeners(List.of(new ChatMemoryChatServiceListener(chatHistory))) .build(); } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/ChatMemoryShortTermSystemPromptIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryShortTermSystemPromptIT.java similarity index 65% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/ChatMemoryShortTermSystemPromptIT.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryShortTermSystemPromptIT.java index b3a9d43ed..7ca4c795b 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/ChatMemoryShortTermSystemPromptIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryShortTermSystemPromptIT.java @@ -14,22 +14,22 @@ * limitations under the License. */ -package org.springframework.ai.openai.chat.chatbot; +package org.springframework.ai.openai.chat.service; import java.util.List; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; -import org.springframework.ai.chat.chatbot.ChatBot; -import org.springframework.ai.chat.chatbot.DefaultChatBot; -import org.springframework.ai.chat.chatbot.DefaultStreamingChatBot; -import org.springframework.ai.chat.chatbot.StreamingChatBot; -import org.springframework.ai.chat.history.ChatMemory; -import org.springframework.ai.chat.history.ChatMemoryChatBotListener; -import org.springframework.ai.chat.history.ChatMemoryRetriever; -import org.springframework.ai.chat.history.InMemoryChatMemory; -import org.springframework.ai.chat.history.LastMaxTokenSizeContentTransformer; -import org.springframework.ai.chat.history.SystemPromptChatMemoryAugmentor; +import org.springframework.ai.chat.service.ChatService; +import org.springframework.ai.chat.service.PromptTransformingChatService; +import org.springframework.ai.chat.service.StreamingPromptTransformingChatService; +import org.springframework.ai.chat.service.StreamingChatService; +import org.springframework.ai.chat.memory.ChatMemory; +import org.springframework.ai.chat.memory.ChatMemoryChatServiceListener; +import org.springframework.ai.chat.memory.ChatMemoryRetriever; +import org.springframework.ai.chat.memory.InMemoryChatMemory; +import org.springframework.ai.chat.memory.LastMaxTokenSizeContentTransformer; +import org.springframework.ai.chat.memory.SystemPromptChatMemoryAugmentor; import org.springframework.ai.evaluation.BaseMemoryTest; import org.springframework.ai.evaluation.RelevancyEvaluator; import org.springframework.ai.openai.OpenAiChatClient; @@ -46,9 +46,9 @@ import org.springframework.context.annotation.Bean; public class ChatMemoryShortTermSystemPromptIT extends BaseMemoryTest { @Autowired - public ChatMemoryShortTermSystemPromptIT(RelevancyEvaluator relevancyEvaluator, ChatBot chatBot, - StreamingChatBot streamingChatBot) { - super(relevancyEvaluator, chatBot, streamingChatBot); + public ChatMemoryShortTermSystemPromptIT(RelevancyEvaluator relevancyEvaluator, ChatService chatService, + StreamingChatService streamingChatService) { + super(relevancyEvaluator, chatService, streamingChatService); } @SpringBootConfiguration @@ -75,26 +75,26 @@ public class ChatMemoryShortTermSystemPromptIT extends BaseMemoryTest { } @Bean - public ChatBot memoryChatBot(OpenAiChatClient chatClient, ChatMemory chatHistory, + public ChatService memoryChatService(OpenAiChatClient chatClient, ChatMemory chatHistory, TokenCountEstimator tokenCountEstimator) { - return DefaultChatBot.builder(chatClient) + return PromptTransformingChatService.builder(chatClient) .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) .withContentPostProcessors(List.of(new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 1000))) .withAugmentors(List.of(new SystemPromptChatMemoryAugmentor())) - .withChatBotListeners(List.of(new ChatMemoryChatBotListener(chatHistory))) + .withChatServiceListeners(List.of(new ChatMemoryChatServiceListener(chatHistory))) .build(); } @Bean - public StreamingChatBot memoryStreamingChatBot(OpenAiChatClient streamingChatClient, ChatMemory chatHistory, - TokenCountEstimator tokenCountEstimator) { + public StreamingChatService memoryStreamingChatService(OpenAiChatClient streamingChatClient, + ChatMemory chatHistory, TokenCountEstimator tokenCountEstimator) { - return DefaultStreamingChatBot.builder(streamingChatClient) + return StreamingPromptTransformingChatService.builder(streamingChatClient) .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) .withDocumentPostProcessors(List.of(new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 1000))) .withAugmentors(List.of(new SystemPromptChatMemoryAugmentor())) - .withChatBotListeners(List.of(new ChatMemoryChatBotListener(chatHistory))) + .withChatServiceListeners(List.of(new ChatMemoryChatServiceListener(chatHistory))) .build(); } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/LongShortTermChatMemoryWithRagIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/LongShortTermChatMemoryWithRagIT.java similarity index 77% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/LongShortTermChatMemoryWithRagIT.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/LongShortTermChatMemoryWithRagIT.java index 9d89014e6..f8fc8f71f 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/LongShortTermChatMemoryWithRagIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/LongShortTermChatMemoryWithRagIT.java @@ -14,7 +14,7 @@ * limitations under the License. */ -package org.springframework.ai.openai.chat.chatbot; +package org.springframework.ai.openai.chat.service; import java.util.List; import java.util.Map; @@ -26,24 +26,24 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.chat.chatbot.ChatBot; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; +import org.springframework.ai.chat.service.ChatService; +import org.springframework.ai.chat.service.PromptTransformingChatService; import org.springframework.ai.openai.OpenAiChatOptions; import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.qdrant.QdrantContainer; -import org.springframework.ai.chat.chatbot.DefaultChatBot; -import org.springframework.ai.chat.history.ChatMemory; -import org.springframework.ai.chat.history.ChatMemoryChatBotListener; -import org.springframework.ai.chat.history.ChatMemoryRetriever; -import org.springframework.ai.chat.history.InMemoryChatMemory; -import org.springframework.ai.chat.history.LastMaxTokenSizeContentTransformer; -import org.springframework.ai.chat.history.SystemPromptChatMemoryAugmentor; -import org.springframework.ai.chat.history.VectorStoreChatMemoryChatBotListener; -import org.springframework.ai.chat.history.VectorStoreChatMemoryRetriever; +import org.springframework.ai.chat.memory.ChatMemory; +import org.springframework.ai.chat.memory.ChatMemoryChatServiceListener; +import org.springframework.ai.chat.memory.ChatMemoryRetriever; +import org.springframework.ai.chat.memory.InMemoryChatMemory; +import org.springframework.ai.chat.memory.LastMaxTokenSizeContentTransformer; +import org.springframework.ai.chat.memory.SystemPromptChatMemoryAugmentor; +import org.springframework.ai.chat.memory.VectorStoreChatMemoryChatServiceListener; +import org.springframework.ai.chat.memory.VectorStoreChatMemoryRetriever; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -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; @@ -89,7 +89,7 @@ public class LongShortTermChatMemoryWithRagIT { static QdrantContainer qdrantContainer = new QdrantContainer("qdrant/qdrant:v1.9.2"); @Autowired - ChatBot chatBot; + ChatService chatService; @Autowired RelevancyEvaluator relevancyEvaluator; @@ -122,29 +122,29 @@ public class LongShortTermChatMemoryWithRagIT { } // @Autowired - // StreamingChatBot streamingChatBot; + // StreamingChatService streamingChatService; @Test - void memoryChatBot() { + void memoryChatService() { loadData(); var prompt = new Prompt(new UserMessage("My name is Christian and I like mountain bikes.")); - PromptContext promptContext = new PromptContext(prompt); + ChatServiceContext chatServiceContext = new ChatServiceContext(prompt); - var chatBotResponse1 = this.chatBot.call(promptContext); + var chatServiceResponse1 = this.chatService.call(chatServiceContext); - logger.info("Response1: " + chatBotResponse1.getChatResponse().getResult().getOutput().getContent()); + logger.info("Response1: " + chatServiceResponse1.getChatResponse().getResult().getOutput().getContent()); - var chatBotResponse2 = this.chatBot.call(new PromptContext( + var chatServiceResponse2 = this.chatService.call(new ChatServiceContext( new Prompt(new String("What is my name and what bike model would you suggest for me?")))); - logger.info("Response2: " + chatBotResponse2.getChatResponse().getResult().getOutput().getContent()); + logger.info("Response2: " + chatServiceResponse2.getChatResponse().getResult().getOutput().getContent()); - // logger.info(chatBotResponse2.getPromptContext().getContents().toString()); - assertThat(chatBotResponse2.getChatResponse().getResult().getOutput().getContent()).contains("Christian"); + // logger.info(chatServiceResponse2.getPromptContext().getContents().toString()); + assertThat(chatServiceResponse2.getChatResponse().getResult().getOutput().getContent()).contains("Christian"); EvaluationResponse evaluationResponse = this.relevancyEvaluator - .evaluate(new EvaluationRequest(chatBotResponse2)); + .evaluate(new EvaluationRequest(chatServiceResponse2)); assertTrue(evaluationResponse.isPass(), "Response is not relevant to the question"); @@ -187,12 +187,15 @@ public class LongShortTermChatMemoryWithRagIT { } @Bean - public ChatBot memoryChatBot(OpenAiChatClient chatClient, VectorStore vectorStore, + public ChatService memoryChatService(OpenAiChatClient chatClient, VectorStore vectorStore, TokenCountEstimator tokenCountEstimator, ChatMemory chatHistory) { - return DefaultChatBot.builder(chatClient) + return PromptTransformingChatService.builder(chatClient) .withRetrievers(List.of(new VectorStoreRetriever(vectorStore, SearchRequest.defaults()), - new ChatMemoryRetriever(chatHistory, Map.of(TransformerContentType.SHORT_TERM_MEMORY, "")), + ChatMemoryRetriever.builder() + .withChatHistory(chatHistory) + .withMetadata(Map.of(TransformerContentType.SHORT_TERM_MEMORY, "")) + .build(), new VectorStoreChatMemoryRetriever(vectorStore, 10, Map.of(TransformerContentType.LONG_TERM_MEMORY, "")))) @@ -214,19 +217,19 @@ public class LongShortTermChatMemoryWithRagIT { Set.of(TransformerContentType.LONG_TERM_MEMORY)), new SystemPromptChatMemoryAugmentor(Set.of(TransformerContentType.SHORT_TERM_MEMORY)))) - .withChatBotListeners(List.of(new ChatMemoryChatBotListener(chatHistory), - new VectorStoreChatMemoryChatBotListener(vectorStore, + .withChatServiceListeners(List.of(new ChatMemoryChatServiceListener(chatHistory), + new VectorStoreChatMemoryChatServiceListener(vectorStore, Map.of(TransformerContentType.LONG_TERM_MEMORY, "")))) .build(); } // @Bean - // public StreamingChatBot memoryStreamingChatAgent(OpenAiChatClient + // public StreamingChatService memoryStreamingChatAgent(OpenAiChatClient // streamingChatClient, // VectorStore vectorStore, TokenCountEstimator tokenCountEstimator, ChatHistory // chatHistory) { - // return DefaultStreamingChatBot.builder(streamingChatClient) + // return StreamingPromptTransformingChatService.builder(streamingChatClient) // .withRetrievers(List.of(new ChatHistoryRetriever(chatHistory), new // DocumentChatHistoryRetriever(vectorStore, 10))) // .withDocumentPostProcessors(List.of(new diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/OpenAiDefaultChatBotIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/OpenAiPromptTransformingChatServiceIT.java similarity index 85% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/OpenAiDefaultChatBotIT.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/OpenAiPromptTransformingChatServiceIT.java index c22085a21..b164b8343 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/chatbot/OpenAiDefaultChatBotIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/OpenAiPromptTransformingChatServiceIT.java @@ -14,7 +14,7 @@ * limitations under the License. */ -package org.springframework.ai.openai.chat.chatbot; +package org.springframework.ai.openai.chat.service; import java.util.List; @@ -22,20 +22,19 @@ import io.qdrant.client.QdrantClient; import io.qdrant.client.QdrantGrpcClient; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; -import org.springframework.ai.chat.chatbot.ChatBot; +import org.springframework.ai.chat.service.ChatService; import org.springframework.ai.chat.prompt.transformer.TransformerContentType; import org.springframework.ai.document.Document; import org.springframework.ai.openai.OpenAiChatOptions; -import org.testcontainers.containers.wait.strategy.Wait; import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.qdrant.QdrantContainer; import org.springframework.ai.chat.ChatClient; -import org.springframework.ai.chat.chatbot.DefaultChatBot; +import org.springframework.ai.chat.service.PromptTransformingChatService; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.chat.prompt.transformer.PromptContext; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; import org.springframework.ai.chat.prompt.transformer.QuestionContextAugmentor; import org.springframework.ai.chat.prompt.transformer.VectorStoreRetriever; import org.springframework.ai.embedding.EmbeddingClient; @@ -61,9 +60,9 @@ import static org.junit.jupiter.api.Assertions.assertTrue; import static org.springframework.ai.openai.api.OpenAiApi.ChatModel.GPT_4_TURBO_PREVIEW; @Testcontainers -@SpringBootTest(classes = OpenAiDefaultChatBotIT.Config.class) +@SpringBootTest(classes = OpenAiPromptTransformingChatServiceIT.Config.class) @EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+") -public class OpenAiDefaultChatBotIT { +public class OpenAiPromptTransformingChatServiceIT { private static final String COLLECTION_NAME = "test_collection"; @@ -79,12 +78,13 @@ public class OpenAiDefaultChatBotIT { @Value("classpath:/data/acme/bikes.json") private Resource bikesResource; - private ChatBot chatBot; + private ChatService chatService; @Autowired - public OpenAiDefaultChatBotIT(ChatClient chatClient, ChatBot chatBot, VectorStore vectorStore) { + public OpenAiPromptTransformingChatServiceIT(ChatClient chatClient, ChatService chatService, + VectorStore vectorStore) { this.chatClient = chatClient; - this.chatBot = chatBot; + this.chatService = chatService; this.vectorStore = vectorStore; } @@ -93,8 +93,8 @@ public class OpenAiDefaultChatBotIT { loadData(); var prompt = new Prompt(new UserMessage("What reliable road bike?")); - var chatBotResponse = this.chatBot.call(new PromptContext(prompt)); - String answer = chatBotResponse.getChatResponse().getResult().getOutput().getContent(); + var chatServiceResponse = this.chatService.call(new ChatServiceContext(prompt)); + String answer = chatServiceResponse.getChatResponse().getResult().getOutput().getContent(); assertTrue(answer.contains("Celerity"), "Response does not include 'Celerity'"); // Use GPT 4 as a better model for determining relevancy. gpt 3.5 makes basic @@ -103,7 +103,7 @@ public class OpenAiDefaultChatBotIT { .withModel(GPT_4_TURBO_PREVIEW.getValue()) .build(); var relevancyEvaluator = new RelevancyEvaluator(this.chatClient, openAiChatOptions); - EvaluationRequest evaluationRequest = new EvaluationRequest(chatBotResponse); + EvaluationRequest evaluationRequest = new EvaluationRequest(chatServiceResponse); EvaluationResponse evaluationResponse = relevancyEvaluator.evaluate(evaluationRequest); assertTrue(evaluationResponse.isPass(), "Response is not relevant to the question"); @@ -148,8 +148,8 @@ public class OpenAiDefaultChatBotIT { } @Bean - public ChatBot chatBot(ChatClient chatClient, VectorStore vectorStore) { - return DefaultChatBot.builder(chatClient) + public ChatService chatService(ChatClient chatClient, VectorStore vectorStore) { + return PromptTransformingChatService.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/chatbot/StreamingChatBotResponse.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/StreamingChatBotResponse.java deleted file mode 100644 index 5e699278c..000000000 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/StreamingChatBotResponse.java +++ /dev/null @@ -1,72 +0,0 @@ -package org.springframework.ai.chat.chatbot; - -import reactor.core.publisher.Flux; - -import org.springframework.ai.chat.ChatResponse; -import org.springframework.ai.chat.prompt.transformer.PromptContext; - -/** - * Encapsulates the response from the ChatBot. Contains the most up-to-date PromptContext - * and the final ChatResponse - * - * @author Mark Pollack - * @since 1.0 M1 - */ -public class StreamingChatBotResponse { - - private final PromptContext promptContext; - - private final Flux chatResponse; - - public StreamingChatBotResponse(PromptContext promptContext, Flux chatResponse) { - this.promptContext = promptContext; - this.chatResponse = chatResponse; - } - - public PromptContext getPromptContext() { - return promptContext; - } - - public Flux getChatResponse() { - return chatResponse; - } - - @Override - public String toString() { - return "ChatBotResponse{" + "promptContext=" + promptContext + ", chatResponse=" + chatResponse + '}'; - } - - @Override - public int hashCode() { - final int prime = 31; - int result = 1; - result = prime * result + ((promptContext == null) ? 0 : promptContext.hashCode()); - result = prime * result + ((chatResponse == null) ? 0 : chatResponse.hashCode()); - return result; - } - - @Override - public boolean equals(Object obj) { - if (this == obj) - return true; - if (obj == null) - return false; - if (getClass() != obj.getClass()) - return false; - StreamingChatBotResponse other = (StreamingChatBotResponse) obj; - if (promptContext == null) { - if (other.promptContext != null) - return false; - } - else if (!promptContext.equals(other.promptContext)) - return false; - if (chatResponse == null) { - if (other.chatResponse != null) - return false; - } - else if (!chatResponse.equals(other.chatResponse)) - return false; - return true; - } - -} 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/memory/ChatMemory.java similarity index 95% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/history/ChatMemory.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/memory/ChatMemory.java index 26d4f1342..52767072d 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/memory/ChatMemory.java @@ -14,7 +14,7 @@ * limitations under the License. */ -package org.springframework.ai.chat.history; +package org.springframework.ai.chat.memory; import java.util.List; diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/ChatMemoryChatBotListener.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/memory/ChatMemoryChatServiceListener.java similarity index 61% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/history/ChatMemoryChatBotListener.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/memory/ChatMemoryChatServiceListener.java index 2ddf30489..934c3d43d 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/ChatMemoryChatBotListener.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/memory/ChatMemoryChatServiceListener.java @@ -14,47 +14,47 @@ * limitations under the License. */ -package org.springframework.ai.chat.history; +package org.springframework.ai.chat.memory; import java.util.List; -import org.springframework.ai.chat.chatbot.ChatBotResponse; -import org.springframework.ai.chat.chatbot.ChatBotListener; +import org.springframework.ai.chat.service.ChatServiceResponse; +import org.springframework.ai.chat.service.ChatServiceListener; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.MessageType; -import org.springframework.ai.chat.prompt.transformer.PromptContext; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; import org.springframework.ai.chat.prompt.transformer.TransformerContentType; /** * @author Christian Tzolov */ -public class ChatMemoryChatBotListener implements ChatBotListener { +public class ChatMemoryChatServiceListener implements ChatServiceListener { private final ChatMemory chatHistory; - public ChatMemoryChatBotListener(ChatMemory chatHistory) { + public ChatMemoryChatServiceListener(ChatMemory chatHistory) { this.chatHistory = chatHistory; } @Override - public void onStart(PromptContext promptContext) { - var messagesToAdd = promptContext.getPrompt() + public void onStart(ChatServiceContext chatServiceContext) { + var messagesToAdd = chatServiceContext.getPrompt() .getInstructions() .stream() .filter(m -> !m.getMetadata().containsKey(TransformerContentType.MEMORY)) .filter(m -> (m.getMessageType() == MessageType.ASSISTANT || m.getMessageType() == MessageType.USER)) .toList(); - this.chatHistory.add(promptContext.getConversationId(), messagesToAdd); + this.chatHistory.add(chatServiceContext.getConversationId(), messagesToAdd); } @Override - public void onComplete(ChatBotResponse chatBotResponse) { - List assistantMessages = chatBotResponse.getChatResponse() + public void onComplete(ChatServiceResponse chatServiceResponse) { + List assistantMessages = chatServiceResponse.getChatResponse() .getResults() .stream() .map(g -> (Message) g.getOutput()) .toList(); - this.chatHistory.add(chatBotResponse.getPromptContext().getConversationId(), assistantMessages); + this.chatHistory.add(chatServiceResponse.getPromptContext().getConversationId(), assistantMessages); } } 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/memory/ChatMemoryRetriever.java similarity index 52% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/history/ChatMemoryRetriever.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/memory/ChatMemoryRetriever.java index 31d083846..23908cbaa 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/memory/ChatMemoryRetriever.java @@ -14,68 +14,105 @@ * limitations under the License. */ -package org.springframework.ai.chat.history; +package org.springframework.ai.chat.memory; + +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.MessageType; +import org.springframework.ai.chat.prompt.transformer.AbstractPromptTransformer; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; +import org.springframework.ai.chat.prompt.transformer.TransformerContentType; +import org.springframework.ai.document.Document; +import org.springframework.ai.model.Content; import java.util.ArrayList; import java.util.HashMap; import java.util.List; 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.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; - /** * @author Christian Tzolov */ -public class ChatMemoryRetriever implements PromptTransformer { +public class ChatMemoryRetriever extends AbstractPromptTransformer { private final ChatMemory chatHistory; /** * Additional metadata to be assigned to the retrieved history messages. */ - private final Map additionalMetadata; + private final Map metadata; private final int maxHistorySize; public ChatMemoryRetriever(ChatMemory chatHistory) { - this(chatHistory, Map.of()); + this(chatHistory, 1000, Map.of(), "ChatMemoryRetriever"); } - public ChatMemoryRetriever(ChatMemory chatHistory, Map additionalMetadata) { - this(chatHistory, 1000, additionalMetadata); - } - - public ChatMemoryRetriever(ChatMemory chatHistory, int maxHistorySize, Map additionalMetadata) { + public ChatMemoryRetriever(ChatMemory chatHistory, int maxHistorySize, Map metadata, String name) { this.chatHistory = chatHistory; - this.additionalMetadata = additionalMetadata; + this.metadata = metadata; this.maxHistorySize = maxHistorySize; + this.setName(name); } @Override - public PromptContext transform(PromptContext promptContext) { + public ChatServiceContext transform(ChatServiceContext chatServiceContext) { - List messageHistory = this.chatHistory.get(promptContext.getConversationId(), maxHistorySize); + List messageHistory = this.chatHistory.get(chatServiceContext.getConversationId(), maxHistorySize); List historyContent = (messageHistory != null) ? messageHistory.stream().filter(m -> m.getMessageType() != MessageType.SYSTEM).map(m -> { Content content = new Document(m.getContent(), new ArrayList<>(m.getMedia()), new HashMap<>(m.getMetadata())); - content.getMetadata().putAll(this.additionalMetadata); + content.getMetadata().putAll(this.metadata); content.getMetadata().put(TransformerContentType.MEMORY, true); return content; }).toList() : List.of(); List updatedContents = new ArrayList<>( - promptContext.getContents() != null ? promptContext.getContents() : List.of()); + chatServiceContext.getContents() != null ? chatServiceContext.getContents() : List.of()); updatedContents.addAll(historyContent); - return PromptContext.from(promptContext).withContents(updatedContents).build(); + return ChatServiceContext.from(chatServiceContext).withContents(updatedContents).build(); + } + + public static Builder builder() { + return new Builder(); + } + + public static class Builder { + + private ChatMemory chatHistory; + + private Map metadata = Map.of(); + + private int maxHistorySize = 1000; + + private String name = "ChatMemoryRetriever"; + + public Builder withChatHistory(ChatMemory chatHistory) { + this.chatHistory = chatHistory; + return this; + } + + public Builder withMetadata(Map metadata) { + this.metadata = metadata; + return this; + } + + public Builder withMaxHistorySize(int maxHistorySize) { + this.maxHistorySize = maxHistorySize; + return this; + } + + public Builder withName(String name) { + this.name = name; + return this; + } + + public ChatMemoryRetriever build() { + return new ChatMemoryRetriever(this.chatHistory, this.maxHistorySize, this.metadata, this.name); + } + } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/InMemoryChatMemory.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/memory/InMemoryChatMemory.java similarity index 97% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/history/InMemoryChatMemory.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/memory/InMemoryChatMemory.java index 80bf6671c..34b78963e 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/InMemoryChatMemory.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/memory/InMemoryChatMemory.java @@ -14,7 +14,7 @@ * limitations under the License. */ -package org.springframework.ai.chat.history; +package org.springframework.ai.chat.memory; import java.util.ArrayList; import java.util.List; 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/memory/LastMaxTokenSizeContentTransformer.java similarity index 77% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/history/LastMaxTokenSizeContentTransformer.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/memory/LastMaxTokenSizeContentTransformer.java index 01b07ecd4..b2e1626e6 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/memory/LastMaxTokenSizeContentTransformer.java @@ -14,14 +14,14 @@ * limitations under the License. */ -package org.springframework.ai.chat.history; +package org.springframework.ai.chat.memory; import java.util.ArrayList; import java.util.List; import java.util.Set; -import org.springframework.ai.chat.prompt.transformer.PromptContext; -import org.springframework.ai.chat.prompt.transformer.PromptTransformer; +import org.springframework.ai.chat.prompt.transformer.AbstractPromptTransformer; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; import org.springframework.ai.model.Content; import org.springframework.ai.tokenizer.TokenCountEstimator; @@ -33,7 +33,7 @@ import org.springframework.ai.tokenizer.TokenCountEstimator; * * @author Christian Tzolov */ -public class LastMaxTokenSizeContentTransformer implements PromptTransformer { +public class LastMaxTokenSizeContentTransformer extends AbstractPromptTransformer { protected final TokenCountEstimator tokenCountEstimator; @@ -56,15 +56,15 @@ public class LastMaxTokenSizeContentTransformer implements PromptTransformer { this.filterTags = filterTags; } - protected List doGetDatumToModify(PromptContext promptContext) { - return promptContext.getContents() + protected List doGetDatumToModify(ChatServiceContext chatServiceContext) { + return chatServiceContext.getContents() .stream() .filter(content -> this.filterTags.stream().allMatch(tag -> content.getMetadata().containsKey(tag))) .toList(); } - protected List doGetDatumNotToModify(PromptContext promptContext) { - return promptContext.getContents() + protected List doGetDatumNotToModify(ChatServiceContext chatServiceContext) { + return chatServiceContext.getContents() .stream() .filter(content -> !this.filterTags.stream().allMatch(tag -> content.getMetadata().containsKey(tag))) .toList(); @@ -79,22 +79,22 @@ public class LastMaxTokenSizeContentTransformer implements PromptTransformer { } @Override - public PromptContext transform(PromptContext promptContext) { + public ChatServiceContext transform(ChatServiceContext chatServiceContext) { - List datum = this.doGetDatumToModify(promptContext); + List datum = this.doGetDatumToModify(chatServiceContext); int totalSize = this.doEstimateTokenCount(datum); if (totalSize <= this.maxTokenSize) { - return promptContext; + return chatServiceContext; } List purgedContent = this.purgeExcess(datum, totalSize); - var updatedContent = new ArrayList<>(doGetDatumNotToModify(promptContext)); + var updatedContent = new ArrayList<>(doGetDatumNotToModify(chatServiceContext)); updatedContent.addAll(purgedContent); - return PromptContext.from(promptContext).withContents(updatedContent).build(); + return ChatServiceContext.from(chatServiceContext).withContents(updatedContent).build(); } protected List purgeExcess(List datum, int totalSize) { diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/MessageChatMemoryAugmentor.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/memory/MessageChatMemoryAugmentor.java similarity index 72% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/history/MessageChatMemoryAugmentor.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/memory/MessageChatMemoryAugmentor.java index dcf3d6daa..e3b3ce847 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/MessageChatMemoryAugmentor.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/memory/MessageChatMemoryAugmentor.java @@ -14,7 +14,7 @@ * limitations under the License. */ -package org.springframework.ai.chat.history; +package org.springframework.ai.chat.memory; import java.util.ArrayList; import java.util.List; @@ -26,22 +26,23 @@ import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.chat.prompt.transformer.AbstractPromptTransformer; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; +import org.springframework.ai.chat.prompt.transformer.PromptChange; import org.springframework.ai.chat.prompt.transformer.TransformerContentType; -import org.springframework.ai.chat.prompt.transformer.PromptContext; -import org.springframework.ai.chat.prompt.transformer.PromptTransformer; /** * @author Christian Tzolov */ -public class MessageChatMemoryAugmentor implements PromptTransformer { +public class MessageChatMemoryAugmentor extends AbstractPromptTransformer { @Override - public PromptContext transform(PromptContext promptContext) { + public ChatServiceContext transform(ChatServiceContext chatServiceContext) { - var originalPrompt = promptContext.getPrompt(); + var originalPrompt = chatServiceContext.getPrompt(); // Convert the retrieved contents into a list of messages. - List historyMessages = promptContext.getContents() + List historyMessages = chatServiceContext.getContents() .stream() .filter(content -> content.getMetadata().containsKey(TransformerContentType.MEMORY)) .map(content -> { @@ -63,8 +64,10 @@ public class MessageChatMemoryAugmentor implements PromptTransformer { promptMessages.addAll(originalPrompt.getInstructions()); Prompt newPrompt = new Prompt(promptMessages, (ChatOptions) originalPrompt.getOptions()); + PromptChange promptChange = new PromptChange(originalPrompt, newPrompt, this.getName(), + "Added chat memory as individual messages in the prompt"); - return PromptContext.from(promptContext).withPrompt(newPrompt).addPromptHistory(originalPrompt).build(); + return ChatServiceContext.from(chatServiceContext).withPrompt(newPrompt).withPromptChange(promptChange).build(); } } \ No newline at end of file diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/SystemPromptChatMemoryAugmentor.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/memory/SystemPromptChatMemoryAugmentor.java similarity index 81% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/history/SystemPromptChatMemoryAugmentor.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/memory/SystemPromptChatMemoryAugmentor.java index b507faf16..2ed8306ee 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/SystemPromptChatMemoryAugmentor.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/memory/SystemPromptChatMemoryAugmentor.java @@ -14,7 +14,7 @@ * limitations under the License. */ -package org.springframework.ai.chat.history; +package org.springframework.ai.chat.memory; import java.util.ArrayList; import java.util.HashSet; @@ -28,15 +28,16 @@ import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.chat.messages.SystemMessage; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.chat.prompt.transformer.AbstractPromptTransformer; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; +import org.springframework.ai.chat.prompt.transformer.PromptChange; import org.springframework.ai.chat.prompt.transformer.TransformerContentType; -import org.springframework.ai.chat.prompt.transformer.PromptContext; -import org.springframework.ai.chat.prompt.transformer.PromptTransformer; import org.springframework.util.Assert; /** * @author Christian Tzolov */ -public class SystemPromptChatMemoryAugmentor implements PromptTransformer { +public class SystemPromptChatMemoryAugmentor extends AbstractPromptTransformer { public static final String DEFAULT_HISTORY_PROMPT = """ Use the conversation history from the HISTORY section to provide accurate answers. @@ -72,9 +73,9 @@ public class SystemPromptChatMemoryAugmentor implements PromptTransformer { } @Override - public PromptContext transform(PromptContext promptContext) { + public ChatServiceContext transform(ChatServiceContext chatServiceContext) { - var originalPrompt = promptContext.getPrompt(); + var originalPrompt = chatServiceContext.getPrompt(); List systemMessages = (originalPrompt.getInstructions() != null) ? originalPrompt.getInstructions() .stream() @@ -89,7 +90,7 @@ public class SystemPromptChatMemoryAugmentor implements PromptTransformer { SystemMessage originalSystemMessage = (!systemMessages.isEmpty()) ? (SystemMessage) systemMessages.get(0) : new SystemMessage(""); - String historyContext = promptContext.getContents() + String historyContext = chatServiceContext.getContents() .stream() .filter(content -> this.filterTags.stream().allMatch(tag -> content.getMetadata().containsKey(tag))) .map(content -> content.getMetadata().get(AbstractMessage.MESSAGE_TYPE) + ": " + content.getContent()) @@ -103,8 +104,9 @@ public class SystemPromptChatMemoryAugmentor implements PromptTransformer { newPromptMessages.addAll(nonSystemMessages); Prompt newPrompt = new Prompt(newPromptMessages, (ChatOptions) originalPrompt.getOptions()); - - return PromptContext.from(promptContext).withPrompt(newPrompt).addPromptHistory(originalPrompt).build(); + PromptChange promptChange = new PromptChange(originalPrompt, newPrompt, this.getName(), + "Added chat memory into the system prompt"); + return ChatServiceContext.from(chatServiceContext).withPrompt(newPrompt).withPromptChange(promptChange).build(); } } \ No newline at end of file diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/VectorStoreChatMemoryChatBotListener.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/memory/VectorStoreChatMemoryChatServiceListener.java similarity index 68% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/history/VectorStoreChatMemoryChatBotListener.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/memory/VectorStoreChatMemoryChatServiceListener.java index cf42bfb6e..0f5aa0100 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/VectorStoreChatMemoryChatBotListener.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/memory/VectorStoreChatMemoryChatServiceListener.java @@ -14,18 +14,18 @@ * limitations under the License. */ -package org.springframework.ai.chat.history; +package org.springframework.ai.chat.memory; import java.util.HashMap; import java.util.List; import java.util.Map; -import org.springframework.ai.chat.chatbot.ChatBotListener; -import org.springframework.ai.chat.chatbot.ChatBotResponse; +import org.springframework.ai.chat.service.ChatServiceListener; +import org.springframework.ai.chat.service.ChatServiceResponse; 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.PromptContext; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; import org.springframework.ai.document.Document; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.util.CollectionUtils; @@ -33,43 +33,43 @@ import org.springframework.util.CollectionUtils; /** * @author Christian Tzolov */ -public class VectorStoreChatMemoryChatBotListener implements ChatBotListener { +public class VectorStoreChatMemoryChatServiceListener implements ChatServiceListener { private final VectorStore vectorStore; private final Map additionalMetadata; - public VectorStoreChatMemoryChatBotListener(VectorStore vectorStore) { + public VectorStoreChatMemoryChatServiceListener(VectorStore vectorStore) { this(vectorStore, new HashMap<>()); } - public VectorStoreChatMemoryChatBotListener(VectorStore vectorStore, Map additionalMetadata) { + public VectorStoreChatMemoryChatServiceListener(VectorStore vectorStore, Map additionalMetadata) { this.vectorStore = vectorStore; this.additionalMetadata = additionalMetadata; } @Override - public void onStart(PromptContext promptContext) { + public void onStart(ChatServiceContext chatServiceContext) { - if (!CollectionUtils.isEmpty(promptContext.getPrompt().getInstructions())) { - List docs = toDocuments(promptContext.getPrompt().getInstructions(), - promptContext.getConversationId()); + if (!CollectionUtils.isEmpty(chatServiceContext.getPrompt().getInstructions())) { + List docs = toDocuments(chatServiceContext.getPrompt().getInstructions(), + chatServiceContext.getConversationId()); this.vectorStore.add(docs); } } @Override - public void onComplete(ChatBotResponse chatBotResponse) { - if (!CollectionUtils.isEmpty(chatBotResponse.getChatResponse().getResults())) { - List assistantMessages = chatBotResponse.getChatResponse() + public void onComplete(ChatServiceResponse chatServiceResponse) { + if (!CollectionUtils.isEmpty(chatServiceResponse.getChatResponse().getResults())) { + List assistantMessages = chatServiceResponse.getChatResponse() .getResults() .stream() .map(g -> (org.springframework.ai.chat.messages.Message) g.getOutput()) .toList(); List docs = toDocuments(assistantMessages, - chatBotResponse.getPromptContext().getConversationId()); + chatServiceResponse.getPromptContext().getConversationId()); this.vectorStore.add(docs); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/VectorStoreChatMemoryRetriever.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/memory/VectorStoreChatMemoryRetriever.java similarity index 77% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/history/VectorStoreChatMemoryRetriever.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/memory/VectorStoreChatMemoryRetriever.java index aff050179..7b476e704 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/history/VectorStoreChatMemoryRetriever.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/memory/VectorStoreChatMemoryRetriever.java @@ -14,7 +14,7 @@ * limitations under the License. */ -package org.springframework.ai.chat.history; +package org.springframework.ai.chat.memory; import java.util.ArrayList; import java.util.List; @@ -22,9 +22,9 @@ import java.util.Map; import java.util.stream.Collectors; import org.springframework.ai.chat.messages.MessageType; +import org.springframework.ai.chat.prompt.transformer.AbstractPromptTransformer; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; import org.springframework.ai.chat.prompt.transformer.TransformerContentType; -import org.springframework.ai.chat.prompt.transformer.PromptContext; -import org.springframework.ai.chat.prompt.transformer.PromptTransformer; import org.springframework.ai.document.Document; import org.springframework.ai.model.Content; import org.springframework.ai.vectorstore.SearchRequest; @@ -34,7 +34,7 @@ import org.springframework.util.CollectionUtils; /** * @author Christian Tzolov */ -public class VectorStoreChatMemoryRetriever implements PromptTransformer { +public class VectorStoreChatMemoryRetriever extends AbstractPromptTransformer { private final VectorStore vectorStore; @@ -56,11 +56,11 @@ public class VectorStoreChatMemoryRetriever implements PromptTransformer { } @Override - public PromptContext transform(PromptContext promptContext) { + public ChatServiceContext transform(ChatServiceContext chatServiceContext) { List updatedContents = new ArrayList<>( - promptContext.getContents() != null ? promptContext.getContents() : List.of()); + chatServiceContext.getContents() != null ? chatServiceContext.getContents() : List.of()); - String query = promptContext.getPrompt() + String query = chatServiceContext.getPrompt() .getInstructions() .stream() .filter(m -> m.getMessageType() == MessageType.USER) @@ -70,7 +70,7 @@ public class VectorStoreChatMemoryRetriever implements PromptTransformer { var searchRequest = SearchRequest.query(query) .withTopK(this.topK) .withFilterExpression( - TransformerContentType.CONVERSATION_ID + "=='" + promptContext.getConversationId() + "'"); + TransformerContentType.CONVERSATION_ID + "=='" + chatServiceContext.getConversationId() + "'"); List documents = this.vectorStore.similaritySearch(searchRequest); @@ -82,7 +82,7 @@ public class VectorStoreChatMemoryRetriever implements PromptTransformer { updatedContents.addAll(documents); } - return PromptContext.from(promptContext).withContents(updatedContents).build(); + return ChatServiceContext.from(chatServiceContext).withContents(updatedContents).build(); } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/AbstractPromptTransformer.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/AbstractPromptTransformer.java new file mode 100644 index 000000000..0636c9543 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/AbstractPromptTransformer.java @@ -0,0 +1,39 @@ +/* + * 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; + +/** + * AbstractPromptTransformer is an abstract class that provides a base implementation of + * the PromptTransformer interface. It includes a name field and corresponding accessor + * methods, as well as a default implementation for the transform method. + * + * @author Mark Pollack + * @author Christian Tzolov + * @since 1.0.0 M1 + */ +public abstract class AbstractPromptTransformer implements PromptTransformer { + + private String name = getClass().getSimpleName(); + + public String getName() { + return name; + } + + public void setName(String name) { + this.name = name; + } + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/ChatServiceContext.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/ChatServiceContext.java new file mode 100644 index 000000000..23191f8f3 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/ChatServiceContext.java @@ -0,0 +1,195 @@ +/* + * 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 org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.model.Content; + +import java.util.*; +import java.util.concurrent.ConcurrentHashMap; + +/** + * Represents the execution context for the {@link ChatService}. This context is used to + * pass initial parameters to the service and facilitate data sharing between different + * components within the service. + * + *

+ * The {@code ChatServiceContext} includes essential information such as the initial + * prompt and a conversation ID, which are crucial for the correct operation of the chat + * service. + *

+ * + * @author Mark Pollack + * @author Christian Tzolov + * @since 1.0.0 M1 + */ +public class ChatServiceContext { + + private Prompt prompt; // The most up-to-date prompt to use + + private List contents; // The most up-to-date data to use + + private List promptChanges; // The changes make due to transformations + + private String conversationId; + + /** + * Contextual data that can be shared between processing steps in a ChatService + * implementation. + */ + private Map context = new ConcurrentHashMap<>(); + + public ChatServiceContext(Prompt prompt) { + this(prompt, "default", new ArrayList<>()); + } + + public ChatServiceContext(Prompt prompt, String conversationId) { + this(prompt, conversationId, new ArrayList<>()); + } + + public ChatServiceContext(Prompt prompt, String conversationId, List contents) { + this.prompt = prompt; + this.conversationId = conversationId; + this.promptChanges = new ArrayList<>(); + this.promptChanges.add(new PromptChange(null, prompt, "none", "initial prompt")); + this.contents = contents; + } + + public Prompt getPrompt() { + return this.prompt; + } + + public void updatePrompt(Prompt prompt, String transformerName, String description) { + this.promptChanges.add(new PromptChange(this.prompt, prompt, transformerName, description)); + this.prompt = prompt; // set the new prompt as current + } + + public void addData(Content datum) { + this.contents.add(datum); + } + + public List getContents() { + return this.contents; + } + + public void setContents(List contents) { + this.contents = contents; + } + + public List getPromptChanges() { + return this.promptChanges; + } + + public String getConversationId() { + return this.conversationId; + } + + public Map getContext() { + return this.context; + } + + public static Builder from(ChatServiceContext chatServiceContext) { + return ChatServiceContext.builder() + .withContents(new ArrayList<>( + chatServiceContext.getContents() != null ? chatServiceContext.getContents() : List.of())) + .withPrompt(chatServiceContext.getPrompt().copy()) // deep copy + .withMetadata( + new HashMap<>(chatServiceContext.getContext() != null ? chatServiceContext.getContext() : Map.of())) + .withPromptChanges(new ArrayList<>( + chatServiceContext.getPromptChanges() != null ? chatServiceContext.getPromptChanges() : List.of())) + .withConversationId(chatServiceContext.getConversationId()); + } + + public static Builder builder() { + return new Builder(); + } + + public static class Builder { + + private Prompt prompt; + + private List contents; + + private List promptChanges; + + private String conversationId; + + private Map context = new HashMap<>(); + + public Builder withPrompt(Prompt prompt) { + this.prompt = prompt; + return this; + } + + public Builder withContents(List contents) { + this.contents = new ArrayList<>(contents); + return this; + } + + public Builder withPromptChanges(List promptChanges) { + this.promptChanges = new ArrayList<>(promptChanges); + return this; + } + + public Builder withPromptChange(PromptChange promptChange) { + this.promptChanges.add(promptChange); + return this; + } + + public Builder withConversationId(String conversationId) { + this.conversationId = conversationId; + return this; + } + + public Builder withMetadata(Map context) { + this.context = new HashMap<>(context); + return this; + } + + public ChatServiceContext build() { + ChatServiceContext chatServiceContext = new ChatServiceContext(this.prompt, this.conversationId, + this.contents); + chatServiceContext.promptChanges = promptChanges; + chatServiceContext.context = context; + return chatServiceContext; + } + + } + + @Override + public String toString() { + return "ChatServiceContext{" + "prompt=" + prompt + ", contents=" + contents + ", promptHistory=" + + promptChanges + ", conversationId='" + conversationId + '\'' + ", metadata=" + context + '}'; + } + + @Override + public boolean equals(Object o) { + if (this == o) + return true; + if (!(o instanceof ChatServiceContext that)) + return false; + return Objects.equals(prompt, that.prompt) && Objects.equals(contents, that.contents) + && Objects.equals(promptChanges, that.promptChanges) + && Objects.equals(conversationId, that.conversationId) && Objects.equals(context, that.context); + } + + @Override + public int hashCode() { + return Objects.hash(prompt, contents, promptChanges, conversationId, context); + } + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/PromptChange.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/PromptChange.java new file mode 100644 index 000000000..b8e8e1e70 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/PromptChange.java @@ -0,0 +1,31 @@ +/* + * 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 org.springframework.ai.chat.prompt.Prompt; + +/** + * The PromptChange class represents a change made to a Prompt object. It contains + * information about the original prompt, the revised prompt, the name of the transformer + * that made the change, and a description of the change. + * + * @author Mark Pollack + * @author Christian Tzolov + * @since 1.0.0 M1 + */ +public record PromptChange(Prompt original, Prompt revised, String transformerName, String description) { + +} \ No newline at end of file diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/PromptContext.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/PromptContext.java deleted file mode 100644 index 7cfce9f89..000000000 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/PromptContext.java +++ /dev/null @@ -1,189 +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.ArrayList; -import java.util.HashMap; -import java.util.List; -import java.util.Map; -import java.util.Objects; - -import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.model.Content; - -/** - * The shared, at the moment, mutable, data structure that can be used to implement - * ChatBot functionality. - * - * @author Mark Pollack - * @author Christian Tzolov - * @since 1.0.0 - */ -public class PromptContext { - - private Prompt prompt; // The most up-to-date prompt to use - - private List contents; // The most up-to-date data to use - - private List promptHistory; - - private String conversationId = "default"; - - private Map metadata = new HashMap<>(); - - public PromptContext(Prompt prompt) { - this(prompt, new ArrayList<>()); - } - - public PromptContext(Prompt prompt, String conversationId) { - this(prompt, new ArrayList<>()); - this.conversationId = conversationId; - } - - public PromptContext(Prompt prompt, List contents) { - this.prompt = prompt; - this.promptHistory = new ArrayList<>(); - this.promptHistory.add(prompt); - this.contents = contents; - } - - public Prompt getPrompt() { - return prompt; - } - - public void setPrompt(Prompt prompt) { - this.prompt = prompt; - } - - public void addData(Content datum) { - this.contents.add(datum); - } - - public List getContents() { - return contents; - } - - public void setContents(List contents) { - this.contents = contents; - } - - public void addPromptHistory(Prompt prompt) { - this.promptHistory.add(prompt); - } - - public List getPromptHistory() { - return promptHistory; - } - - public String getConversationId() { - return conversationId; - } - - public Map getMetadata() { - return metadata; - } - - public static Builder from(PromptContext promptContext) { - return PromptContext.builder() - .withContents( - new ArrayList<>(promptContext.getContents() != null ? promptContext.getContents() : List.of())) - .withPrompt(promptContext.getPrompt().copy()) // deep copy - .withMetadata(new HashMap<>(promptContext.getMetadata() != null ? promptContext.getMetadata() : Map.of())) - .withPromptHistory(new ArrayList<>( - promptContext.getPromptHistory() != null ? promptContext.getPromptHistory() : List.of())) - .withConversationId(promptContext.getConversationId()); - } - - public static Builder builder() { - return new Builder(); - } - - public static class Builder { - - private Prompt prompt; - - private List contents; - - private List promptHistory; - - private String conversationId; - - private Map metadata = new HashMap<>(); - - public Builder withPrompt(Prompt prompt) { - this.prompt = prompt; - return this; - } - - public Builder withContents(List contents) { - this.contents = new ArrayList<>(contents); - return this; - } - - public Builder withPromptHistory(List promptHistory) { - this.promptHistory = new ArrayList<>(promptHistory); - return this; - } - - public Builder addPromptHistory(Prompt prompt) { - this.promptHistory.add(prompt); - return this; - } - - public Builder withConversationId(String conversationId) { - this.conversationId = conversationId; - return this; - } - - public Builder withMetadata(Map metadata) { - this.metadata = new HashMap<>(metadata); - return this; - } - - public PromptContext build() { - PromptContext promptContext = new PromptContext(prompt, contents); - promptContext.promptHistory = promptHistory; - promptContext.conversationId = conversationId; - promptContext.metadata = metadata; - return promptContext; - } - - } - - @Override - public String toString() { - return "PromptContext{" + "prompt=" + prompt + ", contents=" + contents + ", promptHistory=" + promptHistory - + ", conversationId='" + conversationId + '\'' + ", metadata=" + metadata + '}'; - } - - @Override - public boolean equals(Object o) { - if (this == o) - return true; - if (!(o instanceof PromptContext that)) - return false; - return Objects.equals(prompt, that.prompt) && Objects.equals(contents, that.contents) - && Objects.equals(promptHistory, that.promptHistory) - && Objects.equals(conversationId, that.conversationId) && Objects.equals(metadata, that.metadata); - } - - @Override - public int hashCode() { - return Objects.hash(prompt, contents, promptHistory, conversationId, metadata); - } - -} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/PromptTransformer.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/PromptTransformer.java index 7d596a0e9..67c74c5a3 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/PromptTransformer.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/prompt/transformer/PromptTransformer.java @@ -17,23 +17,24 @@ package org.springframework.ai.chat.prompt.transformer; /** - * Responsible for transforming a Prompt. The PromptContext contains the necessary data to - * make the transformation + * Responsible for transforming a Prompt. The ChatServiceContext contains the necessary + * data to make the transformation * - * Implementations may retrieve data and modify the Prompt object in the PromptContext as - * needed. + * Implementations may retrieve data and modify the Prompt object in the + * ChatServiceContext as needed. * * @author Mark Pollack - * @since 1.0 M1 + * @author Christian Tzolov + * @since 1.0.0 M1 */ @FunctionalInterface public interface PromptTransformer { /** - * Transforms the given PromptContext. - * @param context the PromptContext to transform - * @return the transformed PromptContext + * Transforms the given ChatServiceContext. + * @param context the ChatServiceContext to transform + * @return the transformed ChatServiceContext */ - PromptContext transform(PromptContext context); + ChatServiceContext transform(ChatServiceContext context); } 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 12e4aada2..05d79438e 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 @@ -32,11 +32,15 @@ import org.springframework.ai.model.Content; * 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. + * Documents contained in the ChatServiceContext's Nodes. + * + * @author Mark Pollack + * @author Christian Tzolov + * @since 1.0.0 M1 */ -public class QuestionContextAugmentor implements PromptTransformer { +public class QuestionContextAugmentor extends AbstractPromptTransformer { - private static final String DEFAULT_USER_PROMPT_TEXT = """ + private static final String DEFAULT_USER_TEXT = """ "Context information is below.\\n" "---------------------\\n" "{context}\\n" @@ -48,16 +52,26 @@ public class QuestionContextAugmentor implements PromptTransformer { "Answer: " """; + private String userText; + + public QuestionContextAugmentor() { + this.userText = DEFAULT_USER_TEXT; + this.setName("QuestionContextAugmentor"); + } + + public String getUserText() { + return userText; + } + @Override - public PromptContext transform(PromptContext promptContext) { - String context = doCreateContext(promptContext.getContents()); - Map contextMap = doCreateContextMap(promptContext.getPrompt(), context); - Prompt prompt = doCreatePrompt(promptContext.getPrompt(), contextMap); - promptContext.setPrompt(prompt); - promptContext.addPromptHistory(prompt); // BUG? shouldn't this be original - // promptContext.getPrompt()? + public ChatServiceContext transform(ChatServiceContext chatServiceContext) { + String context = doCreateContext(chatServiceContext.getContents()); + Map contextMap = doCreateContextMap(chatServiceContext.getPrompt(), context); + Prompt prompt = doCreatePrompt(chatServiceContext.getPrompt(), contextMap); + chatServiceContext.updatePrompt(prompt, this.getName(), "Updated prompt with Q/A user text"); + // For now return the modified instance instead of a copy - return promptContext; + return chatServiceContext; } protected String doCreateContext(List data) { @@ -75,7 +89,7 @@ public class QuestionContextAugmentor implements PromptTransformer { } protected Prompt doCreatePrompt(Prompt originalPrompt, Map contextMap) { - PromptTemplate promptTemplate = new PromptTemplate(DEFAULT_USER_PROMPT_TEXT); + PromptTemplate promptTemplate = new PromptTemplate(getUserText()); Message userMessageToAppend = promptTemplate.createMessage(contextMap); List messageList = originalPrompt.getInstructions() .stream() @@ -85,4 +99,33 @@ public class QuestionContextAugmentor implements PromptTransformer { return new Prompt(messageList, (ChatOptions) originalPrompt.getOptions()); } + public static Builder builder() { + return new Builder(); + } + + public static class Builder { + + private String name; + + private String userText; + + public Builder withName(String name) { + this.name = name; + return this; + } + + public Builder withUserText(String userText) { + this.userText = userText; + return this; + } + + public QuestionContextAugmentor build() { + QuestionContextAugmentor instance = new QuestionContextAugmentor(); + instance.userText = this.userText != null ? this.userText : instance.userText; + instance.setName(this.name != null ? this.name : instance.getName()); + return instance; + } + + } + } 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 270e85943..b4839ecf9 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 @@ -17,7 +17,10 @@ package org.springframework.ai.chat.prompt.transformer; /** + * This class provides constants for different content types used by transformers. + * * @author Christian Tzolov + * @since 1.0.0 M1 */ public class TransformerContentType { 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 8a9e58b1c..30851f644 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 @@ -23,16 +23,29 @@ import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.document.Document; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; -import org.springframework.ai.vectorstore.filter.Filter; import java.util.List; import java.util.Objects; import java.util.stream.Collectors; /** - * Transforms the PromptContext by retrieving documents from a VectorStore + * A transformer class that retrieves documents from a {@link VectorStore} + * + *

+ * The {@code VectorStoreRetriever} leverages a {@link SearchRequest} to query the + * {@link VectorStore} and retrieve documents that are semantically similar to the user's + * input. These documents are then added to the {@link ChatServiceContext} for further + * processing. + *

+ * + * @see VectorStore + * @see SearchRequest + * @see ChatServiceContext + * @author Mark Pollack + * @author Christian Tzolov + * @since 1.0.0 M1 */ -public class VectorStoreRetriever implements PromptTransformer { +public class VectorStoreRetriever extends AbstractPromptTransformer { private final Logger logger = LoggerFactory.getLogger(getClass()); @@ -41,8 +54,13 @@ public class VectorStoreRetriever implements PromptTransformer { private final SearchRequest searchRequest; public VectorStoreRetriever(VectorStore vectorStore, SearchRequest searchRequest) { + this(vectorStore, searchRequest, "VectorStoreRetriever"); + } + + public VectorStoreRetriever(VectorStore vectorStore, SearchRequest searchRequest, String name) { this.vectorStore = vectorStore; this.searchRequest = searchRequest; + this.setName(name); } public VectorStore getVectorStore() { @@ -54,8 +72,8 @@ public class VectorStoreRetriever implements PromptTransformer { } @Override - public PromptContext transform(PromptContext promptContext) { - List instructions = promptContext.getPrompt().getInstructions(); + public ChatServiceContext transform(ChatServiceContext chatServiceContext) { + List instructions = chatServiceContext.getPrompt().getInstructions(); String userMessage = instructions.stream() .filter(m -> m.getMessageType() == MessageType.USER) .map(m -> m.getContent()) @@ -67,9 +85,9 @@ public class VectorStoreRetriever implements PromptTransformer { for (Document document : documents) { var content = new Document(document.getContent(), document.getMetadata()); // content.getMetadata().put(TransformerContentType.DOMAIN_DATA, true); - promptContext.addData(content); + chatServiceContext.addData(content); } - return promptContext; + return chatServiceContext; } @Override diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/ChatBot.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/ChatService.java similarity index 54% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/ChatBot.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/service/ChatService.java index ff4eeb510..647ad9c3c 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/ChatBot.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/ChatService.java @@ -14,27 +14,26 @@ * limitations under the License. */ -package org.springframework.ai.chat.chatbot; +package org.springframework.ai.chat.service; -import org.springframework.ai.chat.prompt.transformer.PromptContext; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; /** - * A ChatBot encapsulates the logic to perform common AI use cases such as Retrieval - * Augmented Generation. + * A ChatService encapsulates the logic to implement AI use cases. * * @author Mark Pollack * @since 1.0 M1 */ -public interface ChatBot { +public interface ChatService { /** - * Call the chatbot to execute AI actions - * @param promptContext A shared data structure used by the ChatBot to perform - * processing of the Prompt. It includes the intial Prompt and a conversation ID at + * Call the service to execute AI actions + * @param chatServiceContext A data structure used by the ChatService to perform + * processing of the Prompt. It includes the initial Prompt and a conversation ID at * the start of execution. - * @return the ChatBotResponse that contains the ChatResponse and the latest - * PromptContext + * @return the ChatServiceResponse that contains the ChatResponse and the latest + * ChatServiceContext */ - ChatBotResponse call(PromptContext promptContext); + ChatServiceResponse call(ChatServiceContext chatServiceContext); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/ChatBotListener.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/ChatServiceListener.java similarity index 61% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/ChatBotListener.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/service/ChatServiceListener.java index 7e591b46d..8e8596cc0 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/ChatBotListener.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/ChatServiceListener.java @@ -14,23 +14,23 @@ * limitations under the License. */ -package org.springframework.ai.chat.chatbot; +package org.springframework.ai.chat.service; -import org.springframework.ai.chat.prompt.transformer.PromptContext; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; /** - * The ChatBotListener is a callback interface that can be implemented by classes that - * want to be notified of the completion of a ChatBot execution. + * The ChatServiceListener is a callback interface that can be implemented by classes that + * want to be notified of the completion of a ChatService execution. * * @author Mark Pollack * @author Christian Tzolov */ -public interface ChatBotListener { +public interface ChatServiceListener { - default void onStart(PromptContext promptContext) { + default void onStart(ChatServiceContext chatServiceContext) { } - void onComplete(ChatBotResponse chatBotResponse); + void onComplete(ChatServiceResponse chatServiceResponse); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/ChatBotResponse.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/ChatServiceResponse.java similarity index 54% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/ChatBotResponse.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/service/ChatServiceResponse.java index 344d8d08b..ecd46969d 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/ChatBotResponse.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/ChatServiceResponse.java @@ -14,33 +14,33 @@ * limitations under the License. */ -package org.springframework.ai.chat.chatbot; +package org.springframework.ai.chat.service; import org.springframework.ai.chat.ChatResponse; -import org.springframework.ai.chat.prompt.transformer.PromptContext; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; import java.util.Objects; /** - * Encapsulates the response from the ChatBot. Contains the most up-to-date PromptContext - * and the final ChatResponse + * Encapsulates the response from the ChatService. Contains the most up-to-date + * ChatServiceContext and the final ChatResponse * * @author Mark Pollack * @since 1.0 M1 */ -public class ChatBotResponse { +public class ChatServiceResponse { - private final PromptContext promptContext; + private final ChatServiceContext chatServiceContext; private final ChatResponse chatResponse; - public ChatBotResponse(PromptContext promptContext, ChatResponse chatResponse) { - this.promptContext = promptContext; + public ChatServiceResponse(ChatServiceContext chatServiceContext, ChatResponse chatResponse) { + this.chatServiceContext = chatServiceContext; this.chatResponse = chatResponse; } - public PromptContext getPromptContext() { - return promptContext; + public ChatServiceContext getPromptContext() { + return chatServiceContext; } public ChatResponse getChatResponse() { @@ -49,21 +49,23 @@ public class ChatBotResponse { @Override public String toString() { - return "ChatBotResponse{" + "promptContext=" + promptContext + ", chatResponse=" + chatResponse + '}'; + return "ChatServiceResponse{" + "chatServiceContext=" + chatServiceContext + ", chatResponse=" + chatResponse + + '}'; } @Override public boolean equals(Object o) { if (this == o) return true; - if (!(o instanceof ChatBotResponse that)) + if (!(o instanceof ChatServiceResponse that)) return false; - return Objects.equals(promptContext, that.promptContext) && Objects.equals(chatResponse, that.chatResponse); + return Objects.equals(chatServiceContext, that.chatServiceContext) + && Objects.equals(chatResponse, that.chatResponse); } @Override public int hashCode() { - return Objects.hash(promptContext, chatResponse); + return Objects.hash(chatServiceContext, chatResponse); } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/DefaultChatBot.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/PromptTransformingChatService.java similarity index 50% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/DefaultChatBot.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/service/PromptTransformingChatService.java index ebbac06ca..761ae5f42 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/DefaultChatBot.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/PromptTransformingChatService.java @@ -13,11 +13,11 @@ * See the License for the specific language governing permissions and * limitations under the License. */ -package org.springframework.ai.chat.chatbot; +package org.springframework.ai.chat.service; import org.springframework.ai.chat.ChatClient; import org.springframework.ai.chat.ChatResponse; -import org.springframework.ai.chat.prompt.transformer.PromptContext; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; import org.springframework.ai.chat.prompt.transformer.PromptTransformer; import java.util.ArrayList; @@ -25,10 +25,15 @@ import java.util.List; import java.util.Objects; /** + * A PromptTransformingChatService implements the ChatService interface and performs + * transformation of the prompt using a series of PromptTransformers. It also provides a + * builder class for easier construction of the PromptTransformingChatService instance. + * * @author Mark Pollack * @author Christian Tzolov + * @since 1.0 M1 */ -public class DefaultChatBot implements ChatBot { +public class PromptTransformingChatService implements ChatService { private ChatClient chatClient; @@ -38,60 +43,60 @@ public class DefaultChatBot implements ChatBot { private List augmentors; - private List chatBotListeners; + private List chatServiceListeners; - public DefaultChatBot(ChatClient chatClient, List retrievers, + public PromptTransformingChatService(ChatClient chatClient, List retrievers, List documentPostProcessors, List augmentors, - List chatBotListeners) { + List chatServiceListeners) { Objects.requireNonNull(chatClient, "chatClient must not be null"); this.chatClient = chatClient; this.retrievers = retrievers; this.documentPostProcessors = documentPostProcessors; this.augmentors = augmentors; - this.chatBotListeners = chatBotListeners; + this.chatServiceListeners = chatServiceListeners; } - public static DefaultChatBotBuilder builder(ChatClient chatClient) { - return new DefaultChatBotBuilder().withChatClient(chatClient); + public static Builder builder(ChatClient chatClient) { + return new Builder().withChatClient(chatClient); } @Override - public ChatBotResponse call(PromptContext promptContext) { + public ChatServiceResponse call(ChatServiceContext chatServiceContext) { - PromptContext promptContextOnStart = PromptContext.from(promptContext).build(); + ChatServiceContext chatServiceContextOnStart = ChatServiceContext.from(chatServiceContext).build(); // Perform retrieval of documents and messages for (PromptTransformer retriever : this.retrievers) { - promptContext = retriever.transform(promptContext); + chatServiceContext = retriever.transform(chatServiceContext); } // Perform post processing of all retrieved documents and messages for (PromptTransformer documentPostProcessor : this.documentPostProcessors) { - promptContext = documentPostProcessor.transform(promptContext); + chatServiceContext = documentPostProcessor.transform(chatServiceContext); } // Perform prompt augmentation for (PromptTransformer augmentor : this.augmentors) { - promptContext = augmentor.transform(promptContext); + chatServiceContext = augmentor.transform(chatServiceContext); } // Invoke Listeners onStart - for (ChatBotListener listener : this.chatBotListeners) { - listener.onStart(promptContextOnStart); + for (ChatServiceListener listener : this.chatServiceListeners) { + listener.onStart(chatServiceContextOnStart); } // Perform generation - ChatResponse chatResponse = this.chatClient.call(promptContext.getPrompt()); + ChatResponse chatResponse = this.chatClient.call(chatServiceContext.getPrompt()); // Invoke Listeners onComplete - ChatBotResponse chatBotResponse = new ChatBotResponse(promptContext, chatResponse); - for (ChatBotListener listener : this.chatBotListeners) { - listener.onComplete(chatBotResponse); + ChatServiceResponse chatServiceResponse = new ChatServiceResponse(chatServiceContext, chatResponse); + for (ChatServiceListener listener : this.chatServiceListeners) { + listener.onComplete(chatServiceResponse); } - return chatBotResponse; + return chatServiceResponse; } - public static class DefaultChatBotBuilder { + public static class Builder { private ChatClient chatClient; @@ -101,35 +106,36 @@ public class DefaultChatBot implements ChatBot { private List augmentors = new ArrayList<>(); - private List chatBotListeners = new ArrayList<>(); + private List chatServiceListeners = new ArrayList<>(); - public DefaultChatBotBuilder withChatClient(ChatClient chatClient) { + public Builder withChatClient(ChatClient chatClient) { this.chatClient = chatClient; return this; } - public DefaultChatBotBuilder withRetrievers(List retrievers) { + public Builder withRetrievers(List retrievers) { this.retrievers = retrievers; return this; } - public DefaultChatBotBuilder withContentPostProcessors(List documentPostProcessors) { + public Builder withContentPostProcessors(List documentPostProcessors) { this.documentPostProcessors = documentPostProcessors; return this; } - public DefaultChatBotBuilder withAugmentors(List augmentors) { + public Builder withAugmentors(List augmentors) { this.augmentors = augmentors; return this; } - public DefaultChatBotBuilder withChatBotListeners(List chatBotListeners) { - this.chatBotListeners = chatBotListeners; + public Builder withChatServiceListeners(List chatServiceListeners) { + this.chatServiceListeners = chatServiceListeners; return this; } - public DefaultChatBot build() { - return new DefaultChatBot(chatClient, retrievers, documentPostProcessors, augmentors, chatBotListeners); + public PromptTransformingChatService build() { + return new PromptTransformingChatService(chatClient, retrievers, documentPostProcessors, augmentors, + chatServiceListeners); } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/StreamingChatBot.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/StreamingChatService.java similarity index 51% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/StreamingChatBot.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/service/StreamingChatService.java index d5d4843b1..7abb8e63e 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/StreamingChatBot.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/StreamingChatService.java @@ -13,28 +13,27 @@ * See the License for the specific language governing permissions and * limitations under the License. */ -package org.springframework.ai.chat.chatbot; +package org.springframework.ai.chat.service; -import org.springframework.ai.chat.prompt.transformer.PromptContext; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; /** - * A ChatBot encapsulates the logic to perform common AI use cases such as Retrieval - * Augmented Generation. + * A ChatService encapsulates the logic to implement AI use cases. * * @author Mark Pollack * @author Christian Tzolov * @since 1.0 M1 */ -public interface StreamingChatBot { +public interface StreamingChatService { /** - * Call the chatbot to execute AI actions - * @param promptContext A shared data structure used by the ChatBot to perform - * processing of the Prompt. It includes the intial Prompt and a conversation ID at - * the start of execution. - * @return the StreamingChatBotResponse that contains the ChatResponse and the latest - * PromptContext + * Call the service to execute AI actions + * @param chatServiceContext A shared data structure used by the ChatService to + * perform processing of the Prompt. It includes the intial Prompt and a conversation + * ID at the start of execution. + * @return the StreamingChatServiceResponse that contains the ChatResponse and the + * latest ChatServiceContext */ - StreamingChatBotResponse stream(PromptContext promptContext); + StreamingChatServiceResponse stream(ChatServiceContext chatServiceContext); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/service/StreamingChatServiceResponse.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/StreamingChatServiceResponse.java new file mode 100644 index 000000000..5b97e127e --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/StreamingChatServiceResponse.java @@ -0,0 +1,73 @@ +package org.springframework.ai.chat.service; + +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; +import reactor.core.publisher.Flux; + +import org.springframework.ai.chat.ChatResponse; + +/** + * Encapsulates the response from the ChatService. Contains the most up-to-date + * ChatServiceContext and the final ChatResponse + * + * @author Mark Pollack + * @since 1.0 M1 + */ +public class StreamingChatServiceResponse { + + private final ChatServiceContext chatServiceContext; + + private final Flux chatResponse; + + public StreamingChatServiceResponse(ChatServiceContext chatServiceContext, Flux chatResponse) { + this.chatServiceContext = chatServiceContext; + this.chatResponse = chatResponse; + } + + public ChatServiceContext getPromptContext() { + return chatServiceContext; + } + + public Flux getChatResponse() { + return chatResponse; + } + + @Override + public String toString() { + return "ChatServiceResponse{" + "chatServiceContext=" + chatServiceContext + ", chatResponse=" + chatResponse + + '}'; + } + + @Override + public int hashCode() { + final int prime = 31; + int result = 1; + result = prime * result + ((chatServiceContext == null) ? 0 : chatServiceContext.hashCode()); + result = prime * result + ((chatResponse == null) ? 0 : chatResponse.hashCode()); + return result; + } + + @Override + public boolean equals(Object obj) { + if (this == obj) + return true; + if (obj == null) + return false; + if (getClass() != obj.getClass()) + return false; + StreamingChatServiceResponse other = (StreamingChatServiceResponse) obj; + if (chatServiceContext == null) { + if (other.chatServiceContext != null) + return false; + } + else if (!chatServiceContext.equals(other.chatServiceContext)) + return false; + if (chatResponse == null) { + if (other.chatResponse != null) + return false; + } + else if (!chatResponse.equals(other.chatResponse)) + return false; + return true; + } + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/DefaultStreamingChatBot.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/StreamingPromptTransformingChatService.java similarity index 54% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/DefaultStreamingChatBot.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/service/StreamingPromptTransformingChatService.java index bdc79d0b8..c94241531 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/chatbot/DefaultStreamingChatBot.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/StreamingPromptTransformingChatService.java @@ -13,25 +13,25 @@ * See the License for the specific language governing permissions and * limitations under the License. */ -package org.springframework.ai.chat.chatbot; +package org.springframework.ai.chat.service; import java.util.ArrayList; import java.util.List; import java.util.Objects; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; import reactor.core.publisher.Flux; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.StreamingChatClient; import org.springframework.ai.chat.messages.MessageAggregator; -import org.springframework.ai.chat.prompt.transformer.PromptContext; import org.springframework.ai.chat.prompt.transformer.PromptTransformer; /** * @author Mark Pollack * @author Christian Tzolov */ -public class DefaultStreamingChatBot implements StreamingChatBot { +public class StreamingPromptTransformingChatService implements StreamingChatService { private StreamingChatClient streamingChatClient; @@ -41,63 +41,63 @@ public class DefaultStreamingChatBot implements StreamingChatBot { private List augmentors; - private List chatBotListeners; + private List chatServiceListeners; - public DefaultStreamingChatBot(StreamingChatClient chatClient, List retrievers, + public StreamingPromptTransformingChatService(StreamingChatClient chatClient, List retrievers, List documentPostProcessors, List augmentors, - List chatBotListeners) { + List chatServiceListeners) { Objects.requireNonNull(chatClient, "chatClient must not be null"); this.streamingChatClient = chatClient; this.retrievers = retrievers; this.documentPostProcessors = documentPostProcessors; this.augmentors = augmentors; - this.chatBotListeners = chatBotListeners; + this.chatServiceListeners = chatServiceListeners; } - public static DefaultChatBotBuilder builder(StreamingChatClient chatClient) { - return new DefaultChatBotBuilder().withChatClient(chatClient); + public static Builder builder(StreamingChatClient chatClient) { + return new Builder().withChatClient(chatClient); } @Override - public StreamingChatBotResponse stream(PromptContext promptContext) { + public StreamingChatServiceResponse stream(ChatServiceContext chatServiceContext) { - PromptContext promptContextOnStart = PromptContext.from(promptContext).build(); + ChatServiceContext chatServiceContextOnStart = ChatServiceContext.from(chatServiceContext).build(); // Perform retrieval of documents and messages for (PromptTransformer retriever : this.retrievers) { - promptContext = retriever.transform(promptContext); + chatServiceContext = retriever.transform(chatServiceContext); } // Perform post processing of all retrieved documents and messages for (PromptTransformer documentPostProcessor : this.documentPostProcessors) { - promptContext = documentPostProcessor.transform(promptContext); + chatServiceContext = documentPostProcessor.transform(chatServiceContext); } // Perform prompt augmentation for (PromptTransformer augmentor : this.augmentors) { - promptContext = augmentor.transform(promptContext); + chatServiceContext = augmentor.transform(chatServiceContext); } // Invoke Listeners onStart - for (ChatBotListener listener : this.chatBotListeners) { - listener.onStart(promptContextOnStart); + for (ChatServiceListener listener : this.chatServiceListeners) { + listener.onStart(chatServiceContextOnStart); } // Perform generation - final var promptContext2 = promptContext; + final var promptContext2 = chatServiceContext; Flux fluxChatResponse = new MessageAggregator() - .aggregate(this.streamingChatClient.stream(promptContext.getPrompt()), chatResponse -> { - for (ChatBotListener listener : this.chatBotListeners) { - listener.onComplete(new ChatBotResponse(promptContext2, chatResponse)); + .aggregate(this.streamingChatClient.stream(chatServiceContext.getPrompt()), chatResponse -> { + for (ChatServiceListener listener : this.chatServiceListeners) { + listener.onComplete(new ChatServiceResponse(promptContext2, chatResponse)); } }); // Invoke Listeners onComplete - return new StreamingChatBotResponse(promptContext, fluxChatResponse); + return new StreamingChatServiceResponse(chatServiceContext, fluxChatResponse); } - public static class DefaultChatBotBuilder { + public static class Builder { private StreamingChatClient chatClient; @@ -107,36 +107,36 @@ public class DefaultStreamingChatBot implements StreamingChatBot { private List augmentors = new ArrayList<>(); - private List chatBotListeners = new ArrayList<>(); + private List chatServiceListeners = new ArrayList<>(); - public DefaultChatBotBuilder withChatClient(StreamingChatClient chatClient) { + public Builder withChatClient(StreamingChatClient chatClient) { this.chatClient = chatClient; return this; } - public DefaultChatBotBuilder withRetrievers(List retrievers) { + public Builder withRetrievers(List retrievers) { this.retrievers = retrievers; return this; } - public DefaultChatBotBuilder withDocumentPostProcessors(List documentPostProcessors) { + public Builder withDocumentPostProcessors(List documentPostProcessors) { this.documentPostProcessors = documentPostProcessors; return this; } - public DefaultChatBotBuilder withAugmentors(List augmentors) { + public Builder withAugmentors(List augmentors) { this.augmentors = augmentors; return this; } - public DefaultChatBotBuilder withChatBotListeners(List chatBotListeners) { - this.chatBotListeners = chatBotListeners; + public Builder withChatServiceListeners(List chatServiceListeners) { + this.chatServiceListeners = chatServiceListeners; return this; } - public DefaultStreamingChatBot build() { - return new DefaultStreamingChatBot(chatClient, retrievers, documentPostProcessors, augmentors, - chatBotListeners); + public StreamingPromptTransformingChatService build() { + return new StreamingPromptTransformingChatService(chatClient, retrievers, documentPostProcessors, + augmentors, chatServiceListeners); } } 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 657370e86..9d30bdfab 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,7 +1,7 @@ package org.springframework.ai.evaluation; import org.springframework.ai.chat.ChatResponse; -import org.springframework.ai.chat.chatbot.ChatBotResponse; +import org.springframework.ai.chat.service.ChatServiceResponse; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.Content; @@ -16,9 +16,9 @@ public class EvaluationRequest { private final ChatResponse chatResponse; - public EvaluationRequest(ChatBotResponse chatBotResponse) { - this(chatBotResponse.getPromptContext().getPromptHistory().get(0), - chatBotResponse.getPromptContext().getContents(), chatBotResponse.getChatResponse()); + public EvaluationRequest(ChatServiceResponse chatServiceResponse) { + this(chatServiceResponse.getPromptContext().getPromptChanges().get(0).revised(), + chatServiceResponse.getPromptContext().getContents(), chatServiceResponse.getChatResponse()); } public EvaluationRequest(Prompt prompt, List dataList, ChatResponse chatResponse) { diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/Content.java b/spring-ai-core/src/main/java/org/springframework/ai/model/Content.java index a221d7f47..4ecaf351b 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/Content.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/Content.java @@ -19,7 +19,7 @@ public interface Content { /** * Get the content of the message. */ - String getContent(); + String getContent(); // TODO consider getText /** * Get the media associated with the content. diff --git a/spring-ai-core/src/test/java/org/springframework/ai/chat/history/ChatMemoryTests.java b/spring-ai-core/src/test/java/org/springframework/ai/chat/memory/ChatMemoryTests.java similarity index 78% rename from spring-ai-core/src/test/java/org/springframework/ai/chat/history/ChatMemoryTests.java rename to spring-ai-core/src/test/java/org/springframework/ai/chat/memory/ChatMemoryTests.java index f1b2fca90..741fd5341 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/chat/history/ChatMemoryTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/chat/memory/ChatMemoryTests.java @@ -14,7 +14,7 @@ * limitations under the License. */ -package org.springframework.ai.chat.history; +package org.springframework.ai.chat.memory; import java.util.List; @@ -29,12 +29,12 @@ import org.springframework.ai.chat.ChatClient; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; import org.springframework.ai.chat.StreamingChatClient; -import org.springframework.ai.chat.chatbot.ChatBotResponse; -import org.springframework.ai.chat.chatbot.DefaultChatBot; +import org.springframework.ai.chat.service.ChatServiceResponse; +import org.springframework.ai.chat.service.PromptTransformingChatService; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.chat.prompt.transformer.PromptContext; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; import org.springframework.ai.model.Content; import org.springframework.ai.tokenizer.JTokkitTokenCountEstimator; @@ -61,15 +61,15 @@ public class ChatMemoryTests { ChatMemory chatHistory = new InMemoryChatMemory(); - DefaultChatBot chatBot = DefaultChatBot.builder(chatClient) - .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) + PromptTransformingChatService chatService = PromptTransformingChatService.builder(chatClient) + .withRetrievers(List.of(ChatMemoryRetriever.builder().withChatHistory(chatHistory).build())) .withContentPostProcessors( List.of(new LastMaxTokenSizeContentTransformer(new JTokkitTokenCountEstimator(), 10))) .withAugmentors(List.of(new MessageChatMemoryAugmentor())) - .withChatBotListeners(List.of(new ChatMemoryChatBotListener(chatHistory))) + .withChatServiceListeners(List.of(new ChatMemoryChatServiceListener(chatHistory))) .build(); - chatClientUserMessages(chatBot, chatHistory); + chatClientUserMessages(chatService, chatHistory); } @Test @@ -77,32 +77,32 @@ public class ChatMemoryTests { ChatMemory chatHistory = new InMemoryChatMemory(); - DefaultChatBot chatBot = DefaultChatBot.builder(chatClient) + PromptTransformingChatService chatService = PromptTransformingChatService.builder(chatClient) .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) .withContentPostProcessors( List.of(new LastMaxTokenSizeContentTransformer(new JTokkitTokenCountEstimator(), 10))) .withAugmentors(List.of(new SystemPromptChatMemoryAugmentor())) - .withChatBotListeners(List.of(new ChatMemoryChatBotListener(chatHistory))) + .withChatServiceListeners(List.of(new ChatMemoryChatServiceListener(chatHistory))) .build(); - chatClientUserMessages(chatBot, chatHistory); + chatClientUserMessages(chatService, chatHistory); } - public void chatClientUserMessages(DefaultChatBot chatBot, ChatMemory chatHistory) { + public void chatClientUserMessages(PromptTransformingChatService chatService, ChatMemory chatHistory) { when(chatClient.call(promptCaptor.capture())) .thenReturn(new ChatResponse(List.of(new Generation("assistant:1")))) .thenReturn(new ChatResponse(List.of(new Generation("assistant:2")))) .thenReturn(new ChatResponse(List.of(new Generation("assistant:3")))); - var promptContext = PromptContext.builder() + var promptContext = ChatServiceContext.builder() .withConversationId("test-session-id") .withPrompt(new Prompt( List.of(new UserMessage("user:1"), new UserMessage("user:2"), new UserMessage("user:3"), new UserMessage("user:4"), new UserMessage("user:5")))) .build(); - ChatBotResponse response1 = chatBot.call(promptContext); + ChatServiceResponse response1 = chatService.call(promptContext); assertThat(response1.getChatResponse().getResult().getOutput().getContent()).isEqualTo("assistant:1"); @@ -112,7 +112,7 @@ public class ChatMemoryTests { List history = chatHistory.get("test-session-id", 1000); assertThat(history).hasSize(6); - ChatBotResponse response2 = chatBot.call(PromptContext.builder() + ChatServiceResponse response2 = chatService.call(ChatServiceContext.builder() .withConversationId("test-session-id") .withPrompt(new Prompt( List.of(new UserMessage("user:6"), new UserMessage("user:7"), new UserMessage("user:8")))) @@ -129,7 +129,7 @@ public class ChatMemoryTests { assertThat(contents.get(1).getContent()).isEqualTo("user:5"); assertThat(contents.get(2).getContent()).isEqualTo("assistant:1"); - ChatBotResponse response3 = chatBot.call(PromptContext.builder() + ChatServiceResponse response3 = chatService.call(ChatServiceContext.builder() .withConversationId("test-session-id") .withPrompt(new Prompt(List.of(new UserMessage("user:9")))).build()); assertThat(response3.getChatResponse().getResult().getOutput().getContent()).isEqualTo("assistant:3"); diff --git a/spring-ai-test/src/main/java/org/springframework/ai/evaluation/BaseMemoryTest.java b/spring-ai-test/src/main/java/org/springframework/ai/evaluation/BaseMemoryTest.java index 345c55612..6721d2bb6 100644 --- a/spring-ai-test/src/main/java/org/springframework/ai/evaluation/BaseMemoryTest.java +++ b/spring-ai-test/src/main/java/org/springframework/ai/evaluation/BaseMemoryTest.java @@ -22,11 +22,11 @@ import org.junit.jupiter.api.Test; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.chat.chatbot.ChatBot; -import org.springframework.ai.chat.chatbot.StreamingChatBot; +import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; +import org.springframework.ai.chat.service.ChatService; +import org.springframework.ai.chat.service.StreamingChatService; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.chat.prompt.transformer.PromptContext; import static org.assertj.core.api.Assertions.assertThat; @@ -39,48 +39,49 @@ public class BaseMemoryTest { protected RelevancyEvaluator relevancyEvaluator; - protected ChatBot chatBot; + protected ChatService chatService; - protected StreamingChatBot streamingChatBot; + protected StreamingChatService streamingChatService; - public BaseMemoryTest(RelevancyEvaluator relevancyEvaluator, ChatBot chatBot, - StreamingChatBot streamingChatClient) { + public BaseMemoryTest(RelevancyEvaluator relevancyEvaluator, ChatService chatService, + StreamingChatService streamingChatClient) { this.relevancyEvaluator = relevancyEvaluator; - this.chatBot = chatBot; - this.streamingChatBot = streamingChatClient; + this.chatService = chatService; + this.streamingChatService = streamingChatClient; } @Test - void memoryChatBot() { + void memoryChatService() { var prompt = new Prompt(new UserMessage("my name John Vincent Atanasoff")); - PromptContext promptContext = new PromptContext(prompt); + ChatServiceContext chatServiceContext = new ChatServiceContext(prompt); - var chatBotResponse1 = this.chatBot.call(promptContext); + var chatServiceResponse1 = this.chatService.call(chatServiceContext); - logger.info("Response1: " + chatBotResponse1.getChatResponse().getResult().getOutput().getContent()); + logger.info("Response1: " + chatServiceResponse1.getChatResponse().getResult().getOutput().getContent()); // response varies too much. - // assertThat(chatBotResponse1.getChatResponse().getResult().getOutput().getContent()).contains("John"); + // assertThat(chatServiceResponse1.getChatResponse().getResult().getOutput().getContent()).contains("John"); - var chatBotResponse2 = this.chatBot.call(new PromptContext(new Prompt(new String("What is my name?")))); - logger.info("Response2: " + chatBotResponse2.getChatResponse().getResult().getOutput().getContent()); - assertThat(chatBotResponse2.getChatResponse().getResult().getOutput().getContent()) + var chatServiceResponse2 = this.chatService + .call(new ChatServiceContext(new Prompt(new String("What is my name?")))); + logger.info("Response2: " + chatServiceResponse2.getChatResponse().getResult().getOutput().getContent()); + assertThat(chatServiceResponse2.getChatResponse().getResult().getOutput().getContent()) .contains("John Vincent Atanasoff"); EvaluationResponse evaluationResponse = this.relevancyEvaluator - .evaluate(new EvaluationRequest(chatBotResponse2)); + .evaluate(new EvaluationRequest(chatServiceResponse2)); logger.info("" + evaluationResponse); } @Test - void memoryStreamingChatBot() { + void memoryStreamingChatService() { var prompt = new Prompt(new UserMessage("my name John Vincent Atanasoff")); - PromptContext promptContext = new PromptContext(prompt); + ChatServiceContext chatServiceContext = new ChatServiceContext(prompt); - var fluxChatBotResponse1 = this.streamingChatBot.stream(promptContext); + var fluxChatServiceResponse1 = this.streamingChatService.stream(chatServiceContext); - String chatBotResponse1 = fluxChatBotResponse1.getChatResponse() + String chatServiceResponse1 = fluxChatServiceResponse1.getChatResponse() .collectList() .block() .stream() @@ -88,13 +89,13 @@ public class BaseMemoryTest { .map(response -> response.getResult().getOutput().getContent()) .collect(Collectors.joining()); - logger.info("Response1: " + chatBotResponse1); - // response varies too much assertThat(chatBotResponse1).contains("John"); + logger.info("Response1: " + chatServiceResponse1); + // response varies too much assertThat(chatServiceResponse1).contains("John"); - var fluxChatBotResponse2 = this.streamingChatBot - .stream(new PromptContext(new Prompt(new String("What is my name?")))); + var fluxChatServiceResponse2 = this.streamingChatService + .stream(new ChatServiceContext(new Prompt(new String("What is my name?")))); - String chatBotResponse2 = fluxChatBotResponse2.getChatResponse() + String chatServiceResponse2 = fluxChatServiceResponse2.getChatResponse() .collectList() .block() .stream() @@ -102,8 +103,8 @@ public class BaseMemoryTest { .map(response -> response.getResult().getOutput().getContent()) .collect(Collectors.joining()); - logger.info("Response2: " + chatBotResponse2); - assertThat(chatBotResponse2).contains("John Vincent Atanasoff"); + logger.info("Response2: " + chatServiceResponse2); + assertThat(chatServiceResponse2).contains("John Vincent Atanasoff"); } }