Clean and fix the openai ITs

This commit is contained in:
Christian Tzolov
2023-11-23 12:30:28 +01:00
parent c7f22f9a99
commit 58a6735a1b
3 changed files with 17 additions and 12 deletions

View File

@@ -1,6 +1,8 @@
package org.springframework.ai.openai;
import com.theokanning.openai.service.OpenAiService;
import org.springframework.ai.client.AiClient;
import org.springframework.ai.embedding.EmbeddingClient;
import org.springframework.ai.openai.client.OpenAiClient;
import org.springframework.ai.openai.embedding.OpenAiEmbeddingClient;
@@ -25,7 +27,7 @@ public class OpenAiTestConfiguration {
}
@Bean
public OpenAiClient openAiClient(OpenAiService theoOpenAiService) {
public AiClient openAiClient(OpenAiService theoOpenAiService) {
OpenAiClient openAiClient = new OpenAiClient(theoOpenAiService);
openAiClient.setTemperature(0.3);
return openAiClient;

View File

@@ -1,11 +1,13 @@
package org.springframework.ai.openai.acme;
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.client.AiClient;
import org.springframework.ai.client.AiResponse;
import org.springframework.ai.document.Document;
import org.springframework.ai.openai.OpenAiTestConfiguration;
import org.springframework.ai.openai.embedding.OpenAiEmbeddingClient;
import org.springframework.ai.openai.testutils.AbstractIT;
import org.springframework.ai.prompt.Prompt;
@@ -28,7 +30,8 @@ import java.util.stream.Collectors;
import static org.assertj.core.api.Assertions.assertThat;
@SpringBootTest
@SpringBootTest(classes = OpenAiTestConfiguration.class)
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+")
public class AcmeIT extends AbstractIT {
private static final Logger logger = LoggerFactory.getLogger(AcmeIT.class);

View File

@@ -1,9 +1,15 @@
package org.springframework.ai.openai.client;
import org.junit.jupiter.api.Disabled;
import java.util.Arrays;
import java.util.List;
import java.util.Map;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.client.AiResponse;
import org.springframework.ai.client.Generation;
import org.springframework.ai.openai.OpenAiTestConfiguration;
import org.springframework.ai.openai.testutils.AbstractIT;
import org.springframework.ai.parser.BeanOutputParser;
import org.springframework.ai.parser.ListOutputParser;
@@ -18,14 +24,11 @@ import org.springframework.boot.test.context.SpringBootTest;
import org.springframework.core.convert.support.DefaultConversionService;
import org.springframework.core.io.Resource;
import java.util.Arrays;
import java.util.List;
import java.util.Map;
import static org.assertj.core.api.Assertions.assertThat;
@SpringBootTest
class ClientIT extends AbstractIT {
@SpringBootTest(classes = OpenAiTestConfiguration.class)
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+")
class OpenAiClientIT extends AbstractIT {
@Value("classpath:/prompts/system-message.st")
private Resource systemResource;
@@ -59,7 +62,6 @@ class ClientIT extends AbstractIT {
Generation generation = this.openAiClient.generate(prompt).getGeneration();
List<String> list = outputParser.parse(generation.getText());
System.out.println(list);
assertThat(list).hasSize(5);
}
@@ -98,7 +100,6 @@ class ClientIT extends AbstractIT {
Generation generation = openAiClient.generate(prompt).getGeneration();
ActorsFilms actorsFilms = outputParser.parse(generation.getText());
System.out.println(actorsFilms);
}
record ActorsFilmsRecord(String actor, List<String> movies) {
@@ -119,7 +120,6 @@ class ClientIT extends AbstractIT {
Generation generation = openAiClient.generate(prompt).getGeneration();
ActorsFilmsRecord actorsFilms = outputParser.parse(generation.getText());
System.out.println(actorsFilms);
assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks");
assertThat(actorsFilms.movies()).hasSize(5);
}