Add EmbeddingClient and OpenAI implementation
This commit is contained in:
@@ -27,20 +27,22 @@ public abstract class AbstractChain implements Chain {
|
||||
|
||||
@Override
|
||||
public Map<String, Object> apply(Map<String, Object> inputMap) {
|
||||
Map<String, Object> inputMapToUse = processBeforeApply(inputMap);
|
||||
Map<String, Object> inputMapToUse = preProcess(inputMap);
|
||||
Map<String, Object> outputMap = doApply(inputMapToUse);
|
||||
Map<String, Object> outputMapToUse = processAfterApply(inputMapToUse, outputMap);
|
||||
Map<String, Object> outputMapToUse = postProcess(inputMapToUse, outputMap);
|
||||
return outputMapToUse;
|
||||
}
|
||||
|
||||
protected Map<String, Object> processBeforeApply(Map<String, Object> inputMap) {
|
||||
@Override
|
||||
public Map<String, Object> preProcess(Map<String, Object> inputMap) {
|
||||
validateInputs(inputMap);
|
||||
return inputMap;
|
||||
}
|
||||
|
||||
protected abstract Map<String, Object> doApply(Map<String, Object> inputMap);
|
||||
|
||||
private Map<String, Object> processAfterApply(Map<String, Object> inputMap, Map<String, Object> outputMap) {
|
||||
@Override
|
||||
public Map<String, Object> postProcess(Map<String, Object> inputMap, Map<String, Object> outputMap) {
|
||||
validateOutputs(outputMap);
|
||||
Map<String, Object> combindedMap = new HashMap<>();
|
||||
combindedMap.putAll(inputMap);
|
||||
|
||||
@@ -10,4 +10,8 @@ public interface Chain extends Function<Map<String, Object>, Map<String, Object>
|
||||
|
||||
List<String> getOutputKeys();
|
||||
|
||||
Map<String, Object> preProcess(Map<String, Object> inputMap);
|
||||
|
||||
Map<String, Object> postProcess(Map<String, Object> inputMap, Map<String, Object> outputMap);
|
||||
|
||||
}
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
package org.springframework.ai.core.embedding;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Objects;
|
||||
|
||||
public class Embedding {
|
||||
|
||||
private List<Double> embedding;
|
||||
|
||||
private Integer index;
|
||||
|
||||
public Embedding(List<Double> embedding, Integer index) {
|
||||
this.embedding = embedding;
|
||||
this.index = index;
|
||||
}
|
||||
|
||||
public List<Double> 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 ? "<empty>" : "<has data>";
|
||||
return "Embedding{" + "embedding=" + messsage + ", index=" + index + '}';
|
||||
}
|
||||
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
package org.springframework.ai.core.embedding;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
public interface EmbeddingClient {
|
||||
|
||||
List<Double> embed(String text);
|
||||
|
||||
EmbeddingResult embed(List<String> texts);
|
||||
|
||||
}
|
||||
@@ -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<Embedding> data;
|
||||
|
||||
private Map<String, Object> metadata = new HashMap<>();
|
||||
|
||||
public EmbeddingResult(List<Embedding> data, Map<String, Object> metadata) {
|
||||
this.data = data;
|
||||
this.metadata = metadata;
|
||||
}
|
||||
|
||||
public List<Embedding> getData() {
|
||||
return data;
|
||||
}
|
||||
|
||||
public Map<String, Object> 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 + '}';
|
||||
}
|
||||
|
||||
}
|
||||
@@ -50,6 +50,7 @@
|
||||
<artifactId>spring-boot-starter-test</artifactId>
|
||||
<scope>test</scope>
|
||||
</dependency>
|
||||
|
||||
</dependencies>
|
||||
|
||||
</project>
|
||||
|
||||
@@ -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<String> 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<Double> 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<Embedding> data = generateEmbeddingList(nativeEmbeddingResult.getData());
|
||||
Map<String, Object> metadata = generateMetadata(nativeEmbeddingResult.getModel(),
|
||||
nativeEmbeddingResult.getUsage());
|
||||
return new EmbeddingResult(data, metadata);
|
||||
}
|
||||
|
||||
private Map<String, Object> generateMetadata(String model, Usage usage) {
|
||||
Map<String, Object> 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<Embedding> generateEmbeddingList(List<com.theokanning.openai.embedding.Embedding> nativeData) {
|
||||
List<Embedding> data = new ArrayList<>();
|
||||
for (com.theokanning.openai.embedding.Embedding nativeDatum : nativeData) {
|
||||
List<Double> nativeDatumEmbedding = nativeDatum.getEmbedding();
|
||||
int nativeIndex = nativeDatum.getIndex();
|
||||
Embedding embedding = new Embedding(nativeDatumEmbedding, nativeIndex);
|
||||
data.add(embedding);
|
||||
}
|
||||
return data;
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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);
|
||||
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user