predicate) {
+ Assert.hasText(value, "Value [%s] must not be null or empty".formatted(value));
+ StringBuilder builder = new StringBuilder();
+ for (char character : value.toCharArray()) {
+ if (predicate.test(character)) {
+ builder.append(character);
+ }
+ }
+ return builder.toString();
+ }
+
+ private static String parseSymbol(String value) {
+ return parse(value, Character::isLetter);
+ }
+
+ private static Long parseTime(String value) {
+ return Long.parseLong(parse(value, Character::isDigit));
+ }
+
+ public String getName() {
+ return this.name;
+ }
+
+ public String getSymbol() {
+ return this.symbol;
+ }
+
+ public ChronoUnit getUnit() {
+ return this.unit;
+ }
+
+ public Duration toDuration(String value) {
+ return Duration.of(parseTime(value), getUnit());
+ }
+
+ }
+
+ }
+
+}
diff --git a/spring-ai-openai/src/test/java/org/springframework/ai/openai/MockOpenAiTestConfiguration.java b/spring-ai-openai/src/test/java/org/springframework/ai/openai/MockOpenAiTestConfiguration.java
deleted file mode 100644
index 0822eb3d7..000000000
--- a/spring-ai-openai/src/test/java/org/springframework/ai/openai/MockOpenAiTestConfiguration.java
+++ /dev/null
@@ -1,96 +0,0 @@
-/*
- * Copyright 2023 the original author or authors.
- *
- * Licensed under the Apache License, Version 2.0 (the "License");
- * you may not use this file except in compliance with the License.
- * You may obtain a copy of the License at
- *
- * https://www.apache.org/licenses/LICENSE-2.0
- *
- * Unless required by applicable law or agreed to in writing, software
- * distributed under the License is distributed on an "AS IS" BASIS,
- * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
- * See the License for the specific language governing permissions and
- * limitations under the License.
- */
-
-package org.springframework.ai.openai;
-
-import static org.springframework.ai.test.config.MockAiTestConfiguration.SPRING_AI_API_PATH;
-
-import java.time.Duration;
-import java.util.UUID;
-
-import com.fasterxml.jackson.databind.ObjectMapper;
-import com.theokanning.openai.client.OpenAiApi;
-import com.theokanning.openai.service.OpenAiService;
-
-import org.springframework.ai.openai.client.OpenAiClient;
-import org.springframework.ai.openai.metadata.support.OpenAiHttpResponseHeadersInterceptor;
-import org.springframework.ai.test.config.MockAiTestConfiguration;
-import org.springframework.boot.SpringBootConfiguration;
-import org.springframework.context.annotation.Bean;
-import org.springframework.context.annotation.Import;
-import org.springframework.context.annotation.Profile;
-import org.springframework.test.web.servlet.MockMvc;
-
-import okhttp3.HttpUrl;
-import okhttp3.OkHttpClient;
-import okhttp3.mockwebserver.Dispatcher;
-import okhttp3.mockwebserver.MockWebServer;
-import retrofit2.Retrofit;
-import retrofit2.adapter.rxjava2.RxJava2CallAdapterFactory;
-import retrofit2.converter.jackson.JacksonConverterFactory;
-
-/**
- * {@link SpringBootConfiguration} for testing {@literal OpenAI's} API using mock objects.
- *
- * This test configuration allows Spring AI framework developers to mock OpenAI's API with
- * Spring {@link MockMvc} and a test provided Spring Web MVC
- * {@link org.springframework.web.bind.annotation.RestController}.
- *
- * This test configuration makes use of the OkHttp3 {@link MockWebServer} and
- * {@link Dispatcher} to integrate with Spring {@link MockMvc}.
- *
- * @author John Blum
- * @see org.springframework.boot.SpringBootConfiguration
- * @see org.springframework.ai.test.config.MockAiTestConfiguration
- * @since 0.7.0
- */
-@SpringBootConfiguration
-@Profile("spring-ai-openai-mocks")
-@Import(MockAiTestConfiguration.class)
-@SuppressWarnings("unused")
-public class MockOpenAiTestConfiguration {
-
- @Bean
- OpenAiService theoOpenAiService(MockWebServer webServer) {
-
- String apiKey = UUID.randomUUID().toString();
- Duration timeout = Duration.ofSeconds(60);
-
- ObjectMapper objectMapper = OpenAiService.defaultObjectMapper();
-
- OkHttpClient httpClient = new OkHttpClient.Builder(OpenAiService.defaultClient(apiKey, timeout))
- .addInterceptor(new OpenAiHttpResponseHeadersInterceptor())
- .build();
-
- HttpUrl baseUrl = webServer.url(SPRING_AI_API_PATH.concat("/"));
-
- Retrofit retrofit = new Retrofit.Builder().baseUrl(baseUrl)
- .addConverterFactory(JacksonConverterFactory.create(objectMapper))
- .addCallAdapterFactory(RxJava2CallAdapterFactory.create())
- .client(httpClient)
- .build();
-
- OpenAiApi api = retrofit.create(OpenAiApi.class);
-
- return new OpenAiService(api);
- }
-
- @Bean
- OpenAiClient apiClient(OpenAiService openAiService) {
- return new OpenAiClient(openAiService);
- }
-
-}
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
index 0db246958..1f8767a4d 100644
--- 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
@@ -1,12 +1,8 @@
package org.springframework.ai.openai;
-import java.time.Duration;
-
-import com.theokanning.openai.service.OpenAiService;
-
import org.springframework.ai.embedding.EmbeddingClient;
+import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.openai.client.OpenAiClient;
-import org.springframework.ai.openai.client.OpenAiStreamClient;
import org.springframework.ai.openai.embedding.OpenAiEmbeddingClient;
import org.springframework.boot.SpringBootConfiguration;
import org.springframework.context.annotation.Bean;
@@ -16,9 +12,9 @@ import org.springframework.util.StringUtils;
public class OpenAiTestConfiguration {
@Bean
- public OpenAiService theoOpenAiService() {
+ public OpenAiApi openAiApi() {
String apiKey = getApiKey();
- OpenAiService openAiService = new OpenAiService(apiKey, Duration.ofSeconds(60));
+ OpenAiApi openAiService = new OpenAiApi(apiKey);
return openAiService;
}
@@ -32,20 +28,15 @@ public class OpenAiTestConfiguration {
}
@Bean
- public OpenAiClient openAiClient(OpenAiService theoOpenAiService) {
- OpenAiClient openAiClient = new OpenAiClient(theoOpenAiService);
+ public OpenAiClient openAiClient(OpenAiApi api) {
+ OpenAiClient openAiClient = new OpenAiClient(api);
openAiClient.setTemperature(0.3);
return openAiClient;
}
@Bean
- public EmbeddingClient openAiEmbeddingClient(OpenAiService theoOpenAiService) {
- return new OpenAiEmbeddingClient(theoOpenAiService);
- }
-
- @Bean
- public OpenAiStreamClient openAiStreamClient() {
- return new OpenAiStreamClient(getApiKey());
+ public EmbeddingClient openAiEmbeddingClient(OpenAiApi api) {
+ return new OpenAiEmbeddingClient(api);
}
}
diff --git a/spring-ai-openai/src/test/java/org/springframework/ai/openai/client/OpenAiClientIT.java b/spring-ai-openai/src/test/java/org/springframework/ai/openai/client/OpenAiClientIT.java
index c28a2ef6c..5baf05229 100644
--- a/spring-ai-openai/src/test/java/org/springframework/ai/openai/client/OpenAiClientIT.java
+++ b/spring-ai-openai/src/test/java/org/springframework/ai/openai/client/OpenAiClientIT.java
@@ -140,11 +140,12 @@ class OpenAiClientIT extends AbstractIT {
Prompt prompt = new Prompt(promptTemplate.createMessage());
String generationTextFromStream = openAiStreamClient.generateStream(prompt)
- .map(OpenAiSseResponse::choices)
- .toStream()
+ .collectList()
+ .block()
+ .stream()
+ .map(AiResponse::getGenerations)
.flatMap(List::stream)
- .map(OpenAiSseResponse.Choice::delta)
- .map(OpenAiSseResponse.Choice.Delta::content)
+ .map(Generation::getContent)
.collect(Collectors.joining());
ActorsFilmsRecord actorsFilms = outputParser.parse(generationTextFromStream);
diff --git a/spring-ai-openai/src/test/java/org/springframework/ai/openai/client/OpenAiClientWithGenerationMetadataTests.java b/spring-ai-openai/src/test/java/org/springframework/ai/openai/client/OpenAiClientWithGenerationMetadata2Tests.java
similarity index 50%
rename from spring-ai-openai/src/test/java/org/springframework/ai/openai/client/OpenAiClientWithGenerationMetadataTests.java
rename to spring-ai-openai/src/test/java/org/springframework/ai/openai/client/OpenAiClientWithGenerationMetadata2Tests.java
index 30ce86faf..83011d74b 100644
--- a/spring-ai-openai/src/test/java/org/springframework/ai/openai/client/OpenAiClientWithGenerationMetadataTests.java
+++ b/spring-ai-openai/src/test/java/org/springframework/ai/openai/client/OpenAiClientWithGenerationMetadata2Tests.java
@@ -1,5 +1,5 @@
/*
- * Copyright 2023 the original author or authors.
+ * Copyright 2023-2023 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
@@ -16,12 +16,9 @@
package org.springframework.ai.openai.client;
-import static org.assertj.core.api.Assertions.assertThat;
-import static org.springframework.ai.test.config.MockAiTestConfiguration.SPRING_AI_API_PATH;
-
-import java.nio.charset.StandardCharsets;
import java.time.Duration;
+import org.junit.jupiter.api.AfterEach;
import org.junit.jupiter.api.Test;
import org.springframework.ai.client.AiResponse;
@@ -30,49 +27,54 @@ import org.springframework.ai.metadata.GenerationMetadata;
import org.springframework.ai.metadata.PromptMetadata;
import org.springframework.ai.metadata.RateLimit;
import org.springframework.ai.metadata.Usage;
-import org.springframework.ai.openai.MockOpenAiTestConfiguration;
+import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.openai.metadata.support.OpenAiApiResponseHeaders;
import org.springframework.ai.prompt.Prompt;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.SpringBootConfiguration;
-import org.springframework.boot.test.context.SpringBootTest;
+import org.springframework.boot.test.autoconfigure.web.client.RestClientTest;
import org.springframework.context.annotation.Bean;
-import org.springframework.context.annotation.Import;
-import org.springframework.http.HttpStatusCode;
+import org.springframework.http.HttpHeaders;
+import org.springframework.http.HttpMethod;
import org.springframework.http.MediaType;
-import org.springframework.http.ResponseEntity;
-import org.springframework.test.context.ActiveProfiles;
-import org.springframework.test.context.ContextConfiguration;
-import org.springframework.test.web.servlet.MockMvc;
-import org.springframework.test.web.servlet.setup.MockMvcBuilders;
-import org.springframework.web.bind.annotation.PostMapping;
-import org.springframework.web.bind.annotation.RequestMapping;
-import org.springframework.web.bind.annotation.RestController;
-import org.springframework.web.context.request.WebRequest;
+import org.springframework.test.web.client.MockRestServiceServer;
+import org.springframework.web.client.RestClient;
+
+import static org.assertj.core.api.Assertions.assertThat;
+import static org.springframework.test.web.client.match.MockRestRequestMatchers.header;
+import static org.springframework.test.web.client.match.MockRestRequestMatchers.method;
+import static org.springframework.test.web.client.match.MockRestRequestMatchers.requestTo;
+import static org.springframework.test.web.client.response.MockRestResponseCreators.withSuccess;
/**
- * Tests using the {@link OpenAiClient} to send an {@literal OpenAI} API request (chat
- * completion) to test the presence of {@link GenerationMetadata} in the
- * {@link AiResponse}.
- *
* @author John Blum
+ * @author Christian Tzolov
* @since 0.7.0
*/
-@SpringBootTest
-@ActiveProfiles("spring-ai-openai-mocks")
-@ContextConfiguration(classes = OpenAiClientWithGenerationMetadataTests.TestConfiguration.class)
-@SuppressWarnings("unused")
-class OpenAiClientWithGenerationMetadataTests {
+@RestClientTest(OpenAiClientWithGenerationMetadata2Tests.Config.class)
+public class OpenAiClientWithGenerationMetadata2Tests {
+
+ private static String TEST_API_KEY = "sk-1234567890";
@Autowired
- private OpenAiClient aiClient;
+ private OpenAiClient openAiClient;
+
+ @Autowired
+ private MockRestServiceServer server;
+
+ @AfterEach
+ void resetMockServer() {
+ server.reset();
+ }
@Test
void aiResponseContainsAiMetadata() {
+ prepareMock();
+
Prompt prompt = new Prompt("Reach for the sky.");
- AiResponse response = this.aiClient.generate(prompt);
+ AiResponse response = this.openAiClient.generate(prompt);
assertThat(response).isNotNull();
@@ -119,65 +121,58 @@ class OpenAiClientWithGenerationMetadataTests {
});
}
- @SpringBootConfiguration
- @Import(MockOpenAiTestConfiguration.class)
- static class TestConfiguration {
+ private void prepareMock() {
- @Bean
- MockMvc mockMvc() {
- return MockMvcBuilders.standaloneSetup(new SpringOpenAiChatCompletionsController()).build();
- }
+ HttpHeaders httpHeaders = new HttpHeaders();
+ httpHeaders.set(OpenAiApiResponseHeaders.REQUESTS_LIMIT_HEADER.getName(), "4000");
+ httpHeaders.set(OpenAiApiResponseHeaders.REQUESTS_REMAINING_HEADER.getName(), "999");
+ httpHeaders.set(OpenAiApiResponseHeaders.REQUESTS_RESET_HEADER.getName(), "2d16h15m29s");
+ httpHeaders.set(OpenAiApiResponseHeaders.TOKENS_LIMIT_HEADER.getName(), "725000");
+ httpHeaders.set(OpenAiApiResponseHeaders.TOKENS_REMAINING_HEADER.getName(), "112358");
+ httpHeaders.set(OpenAiApiResponseHeaders.TOKENS_RESET_HEADER.getName(), "27h55s451ms");
+
+ server.expect(requestTo("/v1/chat/completions"))
+ .andExpect(method(HttpMethod.POST))
+ .andExpect(header(HttpHeaders.AUTHORIZATION, "Bearer " + TEST_API_KEY))
+ .andRespond(withSuccess(getJson(), MediaType.APPLICATION_JSON).headers(httpHeaders));
}
- @RestController
- @RequestMapping(SPRING_AI_API_PATH)
- @SuppressWarnings("all")
- static class SpringOpenAiChatCompletionsController {
+ private String getJson() {
+ return """
+ {
+ "id": "chatcmpl-123",
+ "object": "chat.completion",
+ "created": 1677652288,
+ "model": "gpt-3.5-turbo-0613",
+ "choices": [{
+ "index": 0,
+ "message": {
+ "role": "assistant",
+ "content": "I surrender!"
+ },
+ "finish_reason": "stop"
+ }],
+ "usage": {
+ "prompt_tokens": 9,
+ "completion_tokens": 12,
+ "total_tokens": 21
+ }
+ }
+ """;
+ }
- @PostMapping("/v1/chat/completions")
- ResponseEntity> chatCompletions(WebRequest request) {
+ @SpringBootConfiguration
+ static class Config {
- String json = getJson();
-
- ResponseEntity> response = ResponseEntity.status(HttpStatusCode.valueOf(200))
- .contentType(MediaType.APPLICATION_JSON)
- .contentLength(json.getBytes(StandardCharsets.UTF_8).length)
- .headers(httpHeaders -> {
- httpHeaders.set(OpenAiApiResponseHeaders.REQUESTS_LIMIT_HEADER.getName(), "4000");
- httpHeaders.set(OpenAiApiResponseHeaders.REQUESTS_REMAINING_HEADER.getName(), "999");
- httpHeaders.set(OpenAiApiResponseHeaders.REQUESTS_RESET_HEADER.getName(), "2d16h15m29s");
- httpHeaders.set(OpenAiApiResponseHeaders.TOKENS_LIMIT_HEADER.getName(), "725000");
- httpHeaders.set(OpenAiApiResponseHeaders.TOKENS_REMAINING_HEADER.getName(), "112358");
- httpHeaders.set(OpenAiApiResponseHeaders.TOKENS_RESET_HEADER.getName(), "27h55s451ms");
- })
- .body(getJson());
-
- return response;
+ @Bean
+ public OpenAiApi chatCompletionApi(RestClient.Builder builder) {
+ return new OpenAiApi("", TEST_API_KEY, builder);
}
- private String getJson() {
- return """
- {
- "id": "chatcmpl-123",
- "object": "chat.completion",
- "created": 1677652288,
- "model": "gpt-3.5-turbo-0613",
- "choices": [{
- "index": 0,
- "message": {
- "role": "assistant",
- "content": "I surrender!"
- },
- "finish_reason": "stop"
- }],
- "usage": {
- "prompt_tokens": 9,
- "completion_tokens": 12,
- "total_tokens": 21
- }
- }
- """;
+ @Bean
+ public OpenAiClient openAiClient(OpenAiApi openAiApi) {
+ return new OpenAiClient(openAiApi);
}
}
diff --git a/spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/EmbeddingIT.java b/spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/EmbeddingIT.java
index 368af4a57..2440b998c 100644
--- a/spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/EmbeddingIT.java
+++ b/spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/EmbeddingIT.java
@@ -1,3 +1,18 @@
+/*
+ * Copyright 2023-2023 the original author or authors.
+ *
+ * Licensed under the Apache License, Version 2.0 (the "License");
+ * you may not use this file except in compliance with the License.
+ * You may obtain a copy of the License at
+ *
+ * https://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
package org.springframework.ai.openai.embedding;
import org.junit.jupiter.api.Test;
@@ -20,13 +35,11 @@ class EmbeddingIT {
assertThat(embeddingClient).isNotNull();
EmbeddingResponse embeddingResponse = embeddingClient.embedForResponse(List.of("Hello World"));
- System.out.println(embeddingResponse);
assertThat(embeddingResponse.getData()).hasSize(1);
assertThat(embeddingResponse.getData().get(0).getEmbedding()).isNotEmpty();
assertThat(embeddingResponse.getMetadata()).containsEntry("model", "text-embedding-ada-002-v2");
- assertThat(embeddingResponse.getMetadata()).containsEntry("completion-tokens", 0L);
- assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 2L);
- assertThat(embeddingResponse.getMetadata()).containsEntry("prompt-tokens", 2L);
+ assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 2);
+ assertThat(embeddingResponse.getMetadata()).containsEntry("prompt-tokens", 2);
assertThat(embeddingClient.dimensions()).isEqualTo(1536);
}
diff --git a/spring-ai-openai/src/test/java/org/springframework/ai/openai/testutils/AbstractIT.java b/spring-ai-openai/src/test/java/org/springframework/ai/openai/testutils/AbstractIT.java
index c8685eab6..8f22e2c92 100644
--- a/spring-ai-openai/src/test/java/org/springframework/ai/openai/testutils/AbstractIT.java
+++ b/spring-ai-openai/src/test/java/org/springframework/ai/openai/testutils/AbstractIT.java
@@ -1,10 +1,14 @@
package org.springframework.ai.openai.testutils;
+import java.util.List;
+import java.util.Map;
+
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
+
import org.springframework.ai.client.AiClient;
import org.springframework.ai.client.AiResponse;
-import org.springframework.ai.openai.client.OpenAiStreamClient;
+import org.springframework.ai.client.AiStreamClient;
import org.springframework.ai.prompt.Prompt;
import org.springframework.ai.prompt.PromptTemplate;
import org.springframework.ai.prompt.messages.Message;
@@ -13,9 +17,6 @@ import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.core.io.Resource;
-import java.util.List;
-import java.util.Map;
-
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.fail;
@@ -27,7 +28,7 @@ public abstract class AbstractIT {
protected AiClient openAiClient;
@Autowired
- protected OpenAiStreamClient openAiStreamClient;
+ protected AiStreamClient openAiStreamClient;
@Value("classpath:/prompts/eval/qa-evaluator-accurate-answer.st")
protected Resource qaEvaluatorAccurateAnswerResource;
diff --git a/spring-ai-openai/src/test/java/org/springframework/ai/openai/transformer/MetadataTransformerIT.java b/spring-ai-openai/src/test/java/org/springframework/ai/openai/transformer/MetadataTransformerIT.java
index 47e94fcfc..0cddec197 100644
--- a/spring-ai-openai/src/test/java/org/springframework/ai/openai/transformer/MetadataTransformerIT.java
+++ b/spring-ai-openai/src/test/java/org/springframework/ai/openai/transformer/MetadataTransformerIT.java
@@ -17,16 +17,15 @@
package org.springframework.ai.openai.transformer;
import java.io.IOException;
-import java.time.Duration;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
-import com.theokanning.openai.service.OpenAiService;
import org.junit.jupiter.api.Test;
import org.springframework.ai.document.DefaultContentFormatter;
import org.springframework.ai.document.Document;
+import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.ai.openai.client.OpenAiClient;
import org.springframework.ai.transformer.ContentFormatTransformer;
import org.springframework.ai.transformer.KeywordMetadataEnricher;
@@ -155,19 +154,18 @@ public class MetadataTransformerIT {
public static class OpenAiTestConfiguration {
@Bean
- public OpenAiService theoOpenAiService() throws IOException {
+ public OpenAiApi openAiApi() throws IOException {
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");
}
- OpenAiService openAiService = new OpenAiService(apiKey, Duration.ofSeconds(60));
- return openAiService;
+ return new OpenAiApi(apiKey);
}
@Bean
- public OpenAiClient openAiClient(OpenAiService theoOpenAiService) {
- OpenAiClient openAiClient = new OpenAiClient(theoOpenAiService);
+ public OpenAiClient openAiClient(OpenAiApi openAiApi) {
+ OpenAiClient openAiClient = new OpenAiClient(openAiApi);
openAiClient.setTemperature(0.3);
return openAiClient;
}
diff --git a/spring-ai-spring-boot-autoconfigure/pom.xml b/spring-ai-spring-boot-autoconfigure/pom.xml
index 00d839dd6..70fe61656 100644
--- a/spring-ai-spring-boot-autoconfigure/pom.xml
+++ b/spring-ai-spring-boot-autoconfigure/pom.xml
@@ -72,6 +72,14 @@
true