From 0e0c4b9ba77ffa51c803d6d28bc951faab1dcbee Mon Sep 17 00:00:00 2001 From: Mark Pollack Date: Tue, 15 Aug 2023 17:46:59 -0400 Subject: [PATCH] client refactoring --- .../azure/openai/llm/AzureOpenAiClient.java | 12 +++--- spring-ai-core/pom.xml | 2 +- .../llm/{LlmClient.java => AiClient.java} | 14 ++++--- .../llm/{LLMResponse.java => AiResponse.java} | 38 ++++++++++--------- .../ai/core/llm/Generation.java | 9 +++-- .../ai/core/prompt/messages/UserMessage.java | 8 ++-- .../ai/openai/llm/OpenAiClient.java | 12 +++--- .../ai/openai/acme/AcmeIntegrationTest.java | 13 +++---- 8 files changed, 56 insertions(+), 52 deletions(-) rename spring-ai-core/src/main/java/org/springframework/ai/core/llm/{LlmClient.java => AiClient.java} (69%) rename spring-ai-core/src/main/java/org/springframework/ai/core/llm/{LLMResponse.java => AiResponse.java} (54%) diff --git a/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/llm/AzureOpenAiClient.java b/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/llm/AzureOpenAiClient.java index d1ca3c849..a24ecc264 100644 --- a/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/llm/AzureOpenAiClient.java +++ b/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/llm/AzureOpenAiClient.java @@ -20,8 +20,8 @@ import com.azure.ai.openai.OpenAIClient; import com.azure.ai.openai.models.*; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.core.llm.LLMResponse; -import org.springframework.ai.core.llm.LlmClient; +import org.springframework.ai.core.llm.AiClient; +import org.springframework.ai.core.llm.AiResponse; import org.springframework.ai.core.llm.Generation; import org.springframework.ai.core.prompt.Prompt; import org.springframework.ai.core.prompt.messages.Message; @@ -31,9 +31,9 @@ import java.util.ArrayList; import java.util.List; /** - * Implementation of {@link LlmClient} backed by an OpenAiService + * Implementation of {@link AiClient} backed by an OpenAiService */ -public class AzureOpenAiClient implements LlmClient { +public class AzureOpenAiClient implements AiClient { private static final Logger logger = LoggerFactory.getLogger(AzureOpenAiClient.class); @@ -67,7 +67,7 @@ public class AzureOpenAiClient implements LlmClient { } @Override - public LLMResponse generate(Prompt prompt) { + public AiResponse generate(Prompt prompt) { List messages = prompt.getMessages(); List azureMessages = new ArrayList<>(); for (Message message : messages) { @@ -87,7 +87,7 @@ public class AzureOpenAiClient implements LlmClient { Generation generation = new Generation(choiceMessage.getContent()); generations.add(generation); } - return new LLMResponse(generations); + return new AiResponse(generations); } public Double getTemperature() { diff --git a/spring-ai-core/pom.xml b/spring-ai-core/pom.xml index 1bc08515a..aaeec7da7 100644 --- a/spring-ai-core/pom.xml +++ b/spring-ai-core/pom.xml @@ -29,7 +29,7 @@ org.springframework - spring-core + spring-messaging diff --git a/spring-ai-core/src/main/java/org/springframework/ai/core/llm/LlmClient.java b/spring-ai-core/src/main/java/org/springframework/ai/core/llm/AiClient.java similarity index 69% rename from spring-ai-core/src/main/java/org/springframework/ai/core/llm/LlmClient.java rename to spring-ai-core/src/main/java/org/springframework/ai/core/llm/AiClient.java index a9a5e791e..d6d51e855 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/core/llm/LlmClient.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/core/llm/AiClient.java @@ -17,14 +17,16 @@ package org.springframework.ai.core.llm; import org.springframework.ai.core.prompt.Prompt; +import org.springframework.ai.core.prompt.messages.UserMessage; -public interface LlmClient { +@FunctionalInterface +public interface AiClient { - String generate(String text); + default String generate(String message) { + Prompt prompt = new Prompt(new UserMessage(message)); + return generate(prompt).getGenerations().get(0).getText(); + } - // TODO Change to LLMResponse, maybe get rid of the varargs to simplify the response - // object, convenience of batch query isn't worth adding the complexity to the - // response object signature - LLMResponse generate(Prompt prompt); + AiResponse generate(Prompt prompt); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/core/llm/LLMResponse.java b/spring-ai-core/src/main/java/org/springframework/ai/core/llm/AiResponse.java similarity index 54% rename from spring-ai-core/src/main/java/org/springframework/ai/core/llm/LLMResponse.java rename to spring-ai-core/src/main/java/org/springframework/ai/core/llm/AiResponse.java index 76dffec9f..fdac545d2 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/core/llm/LLMResponse.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/core/llm/AiResponse.java @@ -15,54 +15,58 @@ */ package org.springframework.ai.core.llm; +import java.util.Collections; import java.util.HashMap; import java.util.List; import java.util.Map; -public class LLMResponse { +public class AiResponse { private final List generations; - private Map providerOutput = new HashMap<>(); + private Map providerOutput; - private Map runInfo = new HashMap<>(); + private Map runInfo; - public LLMResponse(List generations) { - this.generations = generations; + public AiResponse(List generations) { + this(generations, Collections.emptyMap(), Collections.emptyMap()); } - public LLMResponse(List generations, Map providerOutput) { - this.generations = generations; - this.providerOutput = providerOutput; + public AiResponse(List generations, Map providerOutput) { + this(generations, providerOutput, Collections.emptyMap()); } - public LLMResponse(List generations, Map providerOutput, Map runInfo) { - this.generations = generations; - this.providerOutput = providerOutput; - this.runInfo = runInfo; + public AiResponse(List generations, Map providerOutput, Map runInfo) { + this.generations = List.copyOf(generations); + this.providerOutput = Map.copyOf(providerOutput); + this.runInfo = Map.copyOf(runInfo); } /** - * The list of generated outputs. It is a list of lists because a single input could - * have multiple outputs. + * The list of generated outputs. It is a list of lists because the Prompt could + * request multiple output generations. * @return */ public List getGenerations() { - return this.generations; + return Collections.unmodifiableList(generations); + } + + public Generation getGeneration() { + return this.generations.get(0); } /** * Arbitrary LLM-provider specific output */ public Map getProviderOutput() { - return null; + return Collections.unmodifiableMap(providerOutput); } /** * The run metadata information */ public Map getRunInfo() { - return null; + return Collections.unmodifiableMap(runInfo); } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/core/llm/Generation.java b/spring-ai-core/src/main/java/org/springframework/ai/core/llm/Generation.java index bee820312..4ced8cadd 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/core/llm/Generation.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/core/llm/Generation.java @@ -16,6 +16,7 @@ package org.springframework.ai.core.llm; +import java.util.Collections; import java.util.HashMap; import java.util.Map; @@ -23,10 +24,10 @@ public class Generation { private final String text; - private Map info = new HashMap<>(); + private Map info; public Generation(String text) { - this.text = text; + this(text, Collections.emptyMap()); } public Generation(String text, Map info) { @@ -35,11 +36,11 @@ public class Generation { } public String getText() { - return text; + return this.text; } public Map getInfo() { - return info; + return Collections.unmodifiableMap(this.info); } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/core/prompt/messages/UserMessage.java b/spring-ai-core/src/main/java/org/springframework/ai/core/prompt/messages/UserMessage.java index 9abb2388b..a41df43b3 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/core/prompt/messages/UserMessage.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/core/prompt/messages/UserMessage.java @@ -25,12 +25,12 @@ import java.util.Map; */ public class UserMessage extends AbstractMessage { - public UserMessage(String content) { - super(MessageType.USER, content); + public UserMessage(String message) { + super(MessageType.USER, message); } - public UserMessage(String content, Map properties) { - super(MessageType.USER, content, properties); + public UserMessage(String message, Map properties) { + super(MessageType.USER, message, properties); } } diff --git a/spring-ai-openai/src/main/java/org/springframework/ai/openai/llm/OpenAiClient.java b/spring-ai-openai/src/main/java/org/springframework/ai/openai/llm/OpenAiClient.java index d2e37fb7c..48885f5b0 100644 --- a/spring-ai-openai/src/main/java/org/springframework/ai/openai/llm/OpenAiClient.java +++ b/spring-ai-openai/src/main/java/org/springframework/ai/openai/llm/OpenAiClient.java @@ -25,17 +25,17 @@ import com.theokanning.openai.service.OpenAiService; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.core.llm.LLMResponse; -import org.springframework.ai.core.llm.LlmClient; +import org.springframework.ai.core.llm.AiClient; +import org.springframework.ai.core.llm.AiResponse; import org.springframework.ai.core.prompt.Prompt; import org.springframework.ai.core.prompt.messages.Message; import org.springframework.util.Assert; /** - * Implementation of {@link LlmClient} backed by an OpenAiService + * Implementation of {@link AiClient} backed by an OpenAiService */ -public class OpenAiClient implements LlmClient { +public class OpenAiClient implements AiClient { private static final Logger logger = LoggerFactory.getLogger(OpenAiClient.class); @@ -75,7 +75,7 @@ public class OpenAiClient implements LlmClient { } @Override - public LLMResponse generate(Prompt prompt) { + public AiResponse generate(Prompt prompt) { List chatCompletionRequests = getChatCompletionRequest(prompt); return getLLMResult(chatCompletionRequests); } @@ -99,7 +99,7 @@ public class OpenAiClient implements LlmClient { return response; } - private LLMResponse getLLMResult(List chatCompletionRequest) { + private AiResponse getLLMResult(List chatCompletionRequest) { // TODO throw new RuntimeException("LLMResult getLLMResult not yet implemented"); } diff --git a/spring-ai-openai/src/test/java/org/springframework/ai/openai/acme/AcmeIntegrationTest.java b/spring-ai-openai/src/test/java/org/springframework/ai/openai/acme/AcmeIntegrationTest.java index d28a05fbb..c707ddc75 100644 --- a/spring-ai-openai/src/test/java/org/springframework/ai/openai/acme/AcmeIntegrationTest.java +++ b/spring-ai-openai/src/test/java/org/springframework/ai/openai/acme/AcmeIntegrationTest.java @@ -2,13 +2,10 @@ package org.springframework.ai.openai.acme; import org.junit.jupiter.api.Test; import org.springframework.ai.core.document.Document; -import org.springframework.ai.core.llm.LLMResponse; -import org.springframework.ai.core.llm.LlmClient; +import org.springframework.ai.core.llm.AiResponse; +import org.springframework.ai.core.llm.AiClient; import org.springframework.ai.core.loader.impl.JsonLoader; -import org.springframework.ai.core.prompt.ChatPromptTemplate; import org.springframework.ai.core.prompt.Prompt; -import org.springframework.ai.core.prompt.PromptTemplate; -import org.springframework.ai.core.prompt.messages.ChatMessage; import org.springframework.ai.core.prompt.messages.SystemMessage; import org.springframework.ai.core.prompt.messages.UserMessage; import org.springframework.ai.core.retriever.impl.VectorStoreRetriever; @@ -35,13 +32,13 @@ public class AcmeIntegrationTest { private OpenAiEmbeddingClient embeddingClient; @Autowired - private LlmClient llmClient; + private AiClient aiClient; @Test void beanTest() { assertThat(resource).isNotNull(); assertThat(embeddingClient).isNotNull(); - assertThat(llmClient).isNotNull(); + assertThat(aiClient).isNotNull(); } void acmeChain() { @@ -74,7 +71,7 @@ public class AcmeIntegrationTest { // Create the prompt ad-hoc for now, need to put in system message and user // message via ChatPromptTemplate or some other message building mechanic Prompt prompt = new Prompt(List.of(systemMessage, userMessage)); - LLMResponse response = llmClient.generate(prompt); + AiResponse response = aiClient.generate(prompt); // Chain // qa = new ConversationalRetrievalChain(llmClient, userPromptTemplate,