From abb29137c36a352d7dd348b6b18e59717a3f2382 Mon Sep 17 00:00:00 2001 From: Mark Pollack Date: Sat, 12 Aug 2023 17:58:58 -0400 Subject: [PATCH] Add EmbeddingClient and OpenAI implementation --- .../ai/core/chain/AbstractChain.java | 10 ++- .../springframework/ai/core/chain/Chain.java | 4 + .../ai/core/embedding/Embedding.java | 46 ++++++++++++ .../ai/core/embedding/EmbeddingClient.java | 11 +++ .../ai/core/embedding/EmbeddingResult.java | 47 ++++++++++++ spring-ai-openai/pom.xml | 1 + .../embedding/OpenAiEmbeddingClient.java | 75 +++++++++++++++++++ .../ai/openai/OpenAiTestConfiguration.java | 39 ++++++++++ .../embedding/EmbeddingIntegrationTest.java | 33 ++++++++ .../openai/OpenAiAutoConfiguration.java | 14 +++- 10 files changed, 274 insertions(+), 6 deletions(-) create mode 100644 spring-ai-core/src/main/java/org/springframework/ai/core/embedding/Embedding.java create mode 100644 spring-ai-core/src/main/java/org/springframework/ai/core/embedding/EmbeddingClient.java create mode 100644 spring-ai-core/src/main/java/org/springframework/ai/core/embedding/EmbeddingResult.java create mode 100644 spring-ai-openai/src/main/java/org/springframework/ai/openai/embedding/OpenAiEmbeddingClient.java create mode 100644 spring-ai-openai/src/test/java/org/springframework/ai/openai/OpenAiTestConfiguration.java create mode 100644 spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/EmbeddingIntegrationTest.java diff --git a/spring-ai-core/src/main/java/org/springframework/ai/core/chain/AbstractChain.java b/spring-ai-core/src/main/java/org/springframework/ai/core/chain/AbstractChain.java index 35e63f5c8..5499c1f9e 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/core/chain/AbstractChain.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/core/chain/AbstractChain.java @@ -27,20 +27,22 @@ public abstract class AbstractChain implements Chain { @Override public Map apply(Map inputMap) { - Map inputMapToUse = processBeforeApply(inputMap); + Map inputMapToUse = preProcess(inputMap); Map outputMap = doApply(inputMapToUse); - Map outputMapToUse = processAfterApply(inputMapToUse, outputMap); + Map outputMapToUse = postProcess(inputMapToUse, outputMap); return outputMapToUse; } - protected Map processBeforeApply(Map inputMap) { + @Override + public Map preProcess(Map inputMap) { validateInputs(inputMap); return inputMap; } protected abstract Map doApply(Map inputMap); - private Map processAfterApply(Map inputMap, Map outputMap) { + @Override + public Map postProcess(Map inputMap, Map outputMap) { validateOutputs(outputMap); Map combindedMap = new HashMap<>(); combindedMap.putAll(inputMap); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/core/chain/Chain.java b/spring-ai-core/src/main/java/org/springframework/ai/core/chain/Chain.java index ef59ebd30..277149e04 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/core/chain/Chain.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/core/chain/Chain.java @@ -10,4 +10,8 @@ public interface Chain extends Function, Map List getOutputKeys(); + Map preProcess(Map inputMap); + + Map postProcess(Map inputMap, Map outputMap); + } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/core/embedding/Embedding.java b/spring-ai-core/src/main/java/org/springframework/ai/core/embedding/Embedding.java new file mode 100644 index 000000000..65cba2a65 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/core/embedding/Embedding.java @@ -0,0 +1,46 @@ +package org.springframework.ai.core.embedding; + +import java.util.List; +import java.util.Objects; + +public class Embedding { + + private List embedding; + + private Integer index; + + public Embedding(List embedding, Integer index) { + this.embedding = embedding; + this.index = index; + } + + public List getEmbedding() { + return embedding; + } + + public Integer getIndex() { + return index; + } + + @Override + public boolean equals(Object o) { + if (this == o) + return true; + if (o == null || getClass() != o.getClass()) + return false; + Embedding embedding1 = (Embedding) o; + return Objects.equals(embedding, embedding1.embedding) && Objects.equals(index, embedding1.index); + } + + @Override + public int hashCode() { + return Objects.hash(embedding, index); + } + + @Override + public String toString() { + String messsage = this.embedding.size() == 0 ? "" : ""; + return "Embedding{" + "embedding=" + messsage + ", index=" + index + '}'; + } + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/core/embedding/EmbeddingClient.java b/spring-ai-core/src/main/java/org/springframework/ai/core/embedding/EmbeddingClient.java new file mode 100644 index 000000000..654b71545 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/core/embedding/EmbeddingClient.java @@ -0,0 +1,11 @@ +package org.springframework.ai.core.embedding; + +import java.util.List; + +public interface EmbeddingClient { + + List embed(String text); + + EmbeddingResult embed(List texts); + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/core/embedding/EmbeddingResult.java b/spring-ai-core/src/main/java/org/springframework/ai/core/embedding/EmbeddingResult.java new file mode 100644 index 000000000..9a8a87c4b --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/core/embedding/EmbeddingResult.java @@ -0,0 +1,47 @@ +package org.springframework.ai.core.embedding; + +import java.util.HashMap; +import java.util.List; +import java.util.Map; +import java.util.Objects; + +public class EmbeddingResult { + + private List data; + + private Map metadata = new HashMap<>(); + + public EmbeddingResult(List data, Map metadata) { + this.data = data; + this.metadata = metadata; + } + + public List getData() { + return data; + } + + public Map getMetadata() { + return metadata; + } + + @Override + public boolean equals(Object o) { + if (this == o) + return true; + if (o == null || getClass() != o.getClass()) + return false; + EmbeddingResult that = (EmbeddingResult) o; + return Objects.equals(data, that.data) && Objects.equals(metadata, that.metadata); + } + + @Override + public int hashCode() { + return Objects.hash(data, metadata); + } + + @Override + public String toString() { + return "EmbeddingResult{" + "data=" + data + ", metadata=" + metadata + '}'; + } + +} diff --git a/spring-ai-openai/pom.xml b/spring-ai-openai/pom.xml index b498007b5..a88ce2adb 100644 --- a/spring-ai-openai/pom.xml +++ b/spring-ai-openai/pom.xml @@ -50,6 +50,7 @@ spring-boot-starter-test test + diff --git a/spring-ai-openai/src/main/java/org/springframework/ai/openai/embedding/OpenAiEmbeddingClient.java b/spring-ai-openai/src/main/java/org/springframework/ai/openai/embedding/OpenAiEmbeddingClient.java new file mode 100644 index 000000000..7d323f393 --- /dev/null +++ b/spring-ai-openai/src/main/java/org/springframework/ai/openai/embedding/OpenAiEmbeddingClient.java @@ -0,0 +1,75 @@ +package org.springframework.ai.openai.embedding; + +import com.theokanning.openai.Usage; +import com.theokanning.openai.embedding.EmbeddingRequest; +import com.theokanning.openai.service.OpenAiService; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; +import org.springframework.ai.core.embedding.Embedding; +import org.springframework.ai.core.embedding.EmbeddingClient; +import org.springframework.ai.core.embedding.EmbeddingResult; +import org.springframework.util.Assert; + +import java.util.ArrayList; +import java.util.HashMap; +import java.util.List; +import java.util.Map; + +public class OpenAiEmbeddingClient implements EmbeddingClient { + + private static final Logger logger = LoggerFactory.getLogger(OpenAiEmbeddingClient.class); + + private final OpenAiService openAiService; + + private String model = "text-embedding-ada-002"; + + public OpenAiEmbeddingClient(OpenAiService openAiService) { + Assert.notNull(openAiService, "OpenAiService must not be null"); + this.openAiService = openAiService; + } + + @Override + public EmbeddingResult embed(List texts) { + EmbeddingRequest embeddingRequest = EmbeddingRequest.builder().input(texts).model(this.model).build(); + com.theokanning.openai.embedding.EmbeddingResult nativeEmbeddingResult = this.openAiService + .createEmbeddings(embeddingRequest); + return generateEmbeddingResult(nativeEmbeddingResult); + } + + @Override + public List embed(String text) { + EmbeddingRequest embeddingRequest = EmbeddingRequest.builder().input(List.of(text)).model(this.model).build(); + com.theokanning.openai.embedding.EmbeddingResult nativeEmbeddingResult = this.openAiService + .createEmbeddings(embeddingRequest); + return generateEmbeddingResult(nativeEmbeddingResult).getData().get(0).getEmbedding(); + } + + private EmbeddingResult generateEmbeddingResult( + com.theokanning.openai.embedding.EmbeddingResult nativeEmbeddingResult) { + List data = generateEmbeddingList(nativeEmbeddingResult.getData()); + Map metadata = generateMetadata(nativeEmbeddingResult.getModel(), + nativeEmbeddingResult.getUsage()); + return new EmbeddingResult(data, metadata); + } + + private Map generateMetadata(String model, Usage usage) { + Map metadata = new HashMap<>(); + metadata.put("model", model); + metadata.put("prompt-tokens", usage.getPromptTokens()); + metadata.put("completion-tokens", usage.getCompletionTokens()); + metadata.put("total-tokens", usage.getTotalTokens()); + return metadata; + } + + private List generateEmbeddingList(List nativeData) { + List data = new ArrayList<>(); + for (com.theokanning.openai.embedding.Embedding nativeDatum : nativeData) { + List nativeDatumEmbedding = nativeDatum.getEmbedding(); + int nativeIndex = nativeDatum.getIndex(); + Embedding embedding = new Embedding(nativeDatumEmbedding, nativeIndex); + data.add(embedding); + } + return data; + } + +} diff --git a/spring-ai-openai/src/test/java/org/springframework/ai/openai/OpenAiTestConfiguration.java b/spring-ai-openai/src/test/java/org/springframework/ai/openai/OpenAiTestConfiguration.java new file mode 100644 index 000000000..b64abb456 --- /dev/null +++ b/spring-ai-openai/src/test/java/org/springframework/ai/openai/OpenAiTestConfiguration.java @@ -0,0 +1,39 @@ +package org.springframework.ai.openai; + +import com.theokanning.openai.service.OpenAiService; +import org.springframework.ai.openai.embedding.OpenAiEmbeddingClient; +import org.springframework.ai.openai.llm.OpenAiClient; +import org.springframework.boot.SpringBootConfiguration; +import org.springframework.context.annotation.Bean; +import org.springframework.util.StringUtils; + +import java.io.IOException; +import java.util.Properties; + +@SpringBootConfiguration +public class OpenAiTestConfiguration { + + @Bean + public OpenAiService theoOpenAiService() throws IOException { + // get api token in file ~/.openai + String apiKey = System.getenv("OPENAI_API_KEY"); + + if (!StringUtils.hasText(apiKey)) { + throw new IllegalArgumentException( + "You must provide an API key. Put it in an environment variable under the name OPENAI_API_KEY"); + } + return new OpenAiService(apiKey); + } + + @Bean + public OpenAiClient openAiClient(OpenAiService theoOpenAiService) { + OpenAiClient openAiClient = new OpenAiClient(theoOpenAiService); + return openAiClient; + } + + @Bean + public OpenAiEmbeddingClient openAiEmbeddingClient(OpenAiService theoOpenAiService) { + return new OpenAiEmbeddingClient(theoOpenAiService); + } + +} diff --git a/spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/EmbeddingIntegrationTest.java b/spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/EmbeddingIntegrationTest.java new file mode 100644 index 000000000..a4ed08015 --- /dev/null +++ b/spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/EmbeddingIntegrationTest.java @@ -0,0 +1,33 @@ +package org.springframework.ai.openai.embedding; + +import org.junit.jupiter.api.Test; +import org.springframework.ai.core.embedding.EmbeddingResult; +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.boot.test.context.SpringBootTest; + +import java.util.List; + +import static org.assertj.core.api.Assertions.assertThat; + +@SpringBootTest +class EmbeddingIntegrationTest { + + @Autowired + private OpenAiEmbeddingClient embeddingClient; + + @Test + void simpleEmbedding() { + assertThat(embeddingClient).isNotNull(); + + EmbeddingResult embeddingResult = embeddingClient.embed(List.of("Hello World")); + System.out.println(embeddingResult); + assertThat(embeddingResult.getData()).hasSize(1); + assertThat(embeddingResult.getData().get(0).getEmbedding()).isNotEmpty(); + assertThat(embeddingResult.getMetadata()).containsEntry("model", "text-embedding-ada-002-v2"); + assertThat(embeddingResult.getMetadata()).containsEntry("completion-tokens", 0L); + assertThat(embeddingResult.getMetadata()).containsEntry("total-tokens", 2L); + assertThat(embeddingResult.getMetadata()).containsEntry("prompt-tokens", 2L); + + } + +} diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfiguration.java index 23138f984..cdafe04d3 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfiguration.java @@ -18,6 +18,7 @@ package org.springframework.ai.autoconfigure.openai; import com.theokanning.openai.service.OpenAiService; +import org.springframework.ai.openai.embedding.OpenAiEmbeddingClient; import org.springframework.ai.openai.llm.OpenAiClient; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; @@ -39,16 +40,25 @@ public class OpenAiAutoConfiguration { } @Bean - public OpenAiClient openAiClient(OpenAiProperties openAiProperties) { + public OpenAiService theoOpenAiService(OpenAiProperties openAiProperties) { if (!StringUtils.hasText(openAiProperties.getApiKey())) { throw new IllegalArgumentException( "You must provide an API key with the property name " + CONFIG_PREFIX + ".api-key"); } - OpenAiService theoOpenAiService = new OpenAiService(openAiProperties.getApiKey()); + return new OpenAiService(openAiProperties.getApiKey()); + } + + @Bean + public OpenAiClient openAiClient(OpenAiProperties openAiProperties, OpenAiService theoOpenAiService) { OpenAiClient openAiClient = new OpenAiClient(theoOpenAiService); openAiClient.setTemperature(openAiProperties.getTemperature()); openAiClient.setModel(openAiProperties.getModel()); return openAiClient; } + @Bean + public OpenAiEmbeddingClient openAiEmbeddingClient(OpenAiService theoOpenAiService) { + return new OpenAiEmbeddingClient(theoOpenAiService); + } + }