From 549c480489d556c0d61ba895d6248bae33648f3d Mon Sep 17 00:00:00 2001 From: Mark Pollack Date: Thu, 16 May 2024 14:01:20 +0200 Subject: [PATCH] Refactoring * Put creation of EvaluationRequest in ChatServiceResponse * Add string constructor to QuestionContextAugmentor * change vectorStore accept() usage to write() --- .../chat/service/LongShortTermChatMemoryWithRagIT.java | 2 +- .../service/OpenAiPromptTransformingChatServiceIT.java | 8 ++++---- .../chat/prompt/transformer/QuestionContextAugmentor.java | 6 +++++- .../ai/chat/service/ChatServiceResponse.java | 6 ++++++ .../springframework/ai/evaluation/EvaluationRequest.java | 5 ----- .../org/springframework/ai/evaluation/BaseMemoryTest.java | 2 +- 6 files changed, 17 insertions(+), 12 deletions(-) diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/LongShortTermChatMemoryWithRagIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/LongShortTermChatMemoryWithRagIT.java index f8fc8f71f..c2739a018 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/LongShortTermChatMemoryWithRagIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/LongShortTermChatMemoryWithRagIT.java @@ -144,7 +144,7 @@ public class LongShortTermChatMemoryWithRagIT { assertThat(chatServiceResponse2.getChatResponse().getResult().getOutput().getContent()).contains("Christian"); EvaluationResponse evaluationResponse = this.relevancyEvaluator - .evaluate(new EvaluationRequest(chatServiceResponse2)); + .evaluate(chatServiceResponse2.toEvaluationRequest()); assertTrue(evaluationResponse.isPass(), "Response is not relevant to the question"); diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/OpenAiPromptTransformingChatServiceIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/OpenAiPromptTransformingChatServiceIT.java index be45cee3e..a5f1ed415 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/OpenAiPromptTransformingChatServiceIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/OpenAiPromptTransformingChatServiceIT.java @@ -104,8 +104,8 @@ public class OpenAiPromptTransformingChatServiceIT { .withModel(GPT_4_TURBO_PREVIEW.getValue()) .build(); var relevancyEvaluator = new RelevancyEvaluator(this.chatClient, openAiChatOptions); - EvaluationRequest evaluationRequest = new EvaluationRequest(chatServiceResponse); - EvaluationResponse evaluationResponse = relevancyEvaluator.evaluate(evaluationRequest); + + EvaluationResponse evaluationResponse = relevancyEvaluator.evaluate(chatServiceResponse.toEvaluationRequest()); assertTrue(evaluationResponse.isPass(), "Response is not relevant to the question"); } @@ -113,13 +113,13 @@ public class OpenAiPromptTransformingChatServiceIT { void loadData() { JsonReader jsonReader = new JsonReader(bikesResource, "name", "price", "shortDescription", "description"); var textSplitter = new TokenTextSplitter(); - List splitDocuments = textSplitter.split(jsonReader.get()); + List splitDocuments = textSplitter.split(jsonReader.read()); for (Document splitDocument : splitDocuments) { splitDocument.getMetadata().put(TransformerContentType.EXTERNAL_KNOWLEDGE, "true"); } - vectorStore.accept(splitDocuments); + vectorStore.write(splitDocuments); } void loadData2() { 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 05d79438e..bdd4830df 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 @@ -55,7 +55,11 @@ public class QuestionContextAugmentor extends AbstractPromptTransformer { private String userText; public QuestionContextAugmentor() { - this.userText = DEFAULT_USER_TEXT; + this(DEFAULT_USER_TEXT); + } + + public QuestionContextAugmentor(String userText) { + this.userText = userText; this.setName("QuestionContextAugmentor"); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/service/ChatServiceResponse.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/ChatServiceResponse.java index ecd46969d..436bd437a 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/service/ChatServiceResponse.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/ChatServiceResponse.java @@ -18,6 +18,7 @@ package org.springframework.ai.chat.service; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; +import org.springframework.ai.evaluation.EvaluationRequest; import java.util.Objects; @@ -47,6 +48,11 @@ public class ChatServiceResponse { return chatResponse; } + public EvaluationRequest toEvaluationRequest() { + return new EvaluationRequest(getPromptContext().getPromptChanges().get(0).revised(), + getPromptContext().getContents(), getChatResponse()); + } + @Override public String toString() { return "ChatServiceResponse{" + "chatServiceContext=" + chatServiceContext + ", chatResponse=" + chatResponse 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 9d30bdfab..0939a6981 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 @@ -16,11 +16,6 @@ public class EvaluationRequest { private final ChatResponse chatResponse; - public EvaluationRequest(ChatServiceResponse chatServiceResponse) { - this(chatServiceResponse.getPromptContext().getPromptChanges().get(0).revised(), - chatServiceResponse.getPromptContext().getContents(), chatServiceResponse.getChatResponse()); - } - public EvaluationRequest(Prompt prompt, List dataList, ChatResponse chatResponse) { this.prompt = prompt; this.dataList = dataList; 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 6721d2bb6..5fa667481 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 @@ -69,7 +69,7 @@ public class BaseMemoryTest { .contains("John Vincent Atanasoff"); EvaluationResponse evaluationResponse = this.relevancyEvaluator - .evaluate(new EvaluationRequest(chatServiceResponse2)); + .evaluate(chatServiceResponse2.toEvaluationRequest()); logger.info("" + evaluationResponse); }