Use pre-configured RestClient.Builder by Spring Boot

- Add annotation dependency on RestClientAutoConfiguration
This commit is contained in:
Toshiaki Maki
2024-02-08 21:14:25 +09:00
committed by Christian Tzolov
parent 577a605cdc
commit 7b58f426ec
11 changed files with 55 additions and 36 deletions

View File

@@ -22,9 +22,11 @@ import org.springframework.ai.ollama.api.OllamaApi;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.ImportRuntimeHints;
import org.springframework.web.client.RestClient;
/**
* {@link AutoConfiguration Auto-configuration} for Ollama Chat Client.
@@ -32,7 +34,7 @@ import org.springframework.context.annotation.ImportRuntimeHints;
* @author Christian Tzolov
* @since 0.8.0
*/
@AutoConfiguration
@AutoConfiguration(after = RestClientAutoConfiguration.class)
@ConditionalOnClass(OllamaApi.class)
@EnableConfigurationProperties({ OllamaChatProperties.class, OllamaEmbeddingProperties.class,
OllamaConnectionProperties.class })
@@ -41,8 +43,8 @@ public class OllamaAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public OllamaApi ollamaApi(OllamaConnectionProperties properties) {
return new OllamaApi(properties.getBaseUrl());
public OllamaApi ollamaApi(OllamaConnectionProperties properties, RestClient.Builder restClientBuilder) {
return new OllamaApi(properties.getBaseUrl(), restClientBuilder);
}
@Bean

View File

@@ -26,6 +26,7 @@ import org.springframework.ai.openai.api.OpenAiImageApi;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.ImportRuntimeHints;
@@ -33,7 +34,7 @@ import org.springframework.util.Assert;
import org.springframework.util.StringUtils;
import org.springframework.web.client.RestClient;
@AutoConfiguration
@AutoConfiguration(after = RestClientAutoConfiguration.class)
@ConditionalOnClass(OpenAiApi.class)
@EnableConfigurationProperties({ OpenAiConnectionProperties.class, OpenAiChatProperties.class,
OpenAiEmbeddingProperties.class, OpenAiImageProperties.class })
@@ -46,7 +47,7 @@ public class OpenAiAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public OpenAiChatClient openAiChatClient(OpenAiConnectionProperties commonProperties,
OpenAiChatProperties chatProperties) {
OpenAiChatProperties chatProperties, RestClient.Builder restClientBuilder) {
String apiKey = StringUtils.hasText(chatProperties.getApiKey()) ? chatProperties.getApiKey()
: commonProperties.getApiKey();
@@ -57,7 +58,7 @@ public class OpenAiAutoConfiguration {
Assert.hasText(apiKey, "OpenAI API key must be set");
Assert.hasText(baseUrl, "OpenAI base URL must be set");
var openAiApi = new OpenAiApi(baseUrl, apiKey, RestClient.builder());
var openAiApi = new OpenAiApi(baseUrl, apiKey, restClientBuilder);
OpenAiChatClient openAiChatClient = new OpenAiChatClient(openAiApi)
.withDefaultOptions(chatProperties.getOptions());
@@ -68,7 +69,7 @@ public class OpenAiAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public EmbeddingClient openAiEmbeddingClient(OpenAiConnectionProperties commonProperties,
OpenAiEmbeddingProperties embeddingProperties) {
OpenAiEmbeddingProperties embeddingProperties, RestClient.Builder restClientBuilder) {
String apiKey = StringUtils.hasText(embeddingProperties.getApiKey()) ? embeddingProperties.getApiKey()
: commonProperties.getApiKey();
@@ -78,7 +79,7 @@ public class OpenAiAutoConfiguration {
Assert.hasText(apiKey, "OpenAI API key must be set");
Assert.hasText(baseUrl, "OpenAI base URL must be set");
var openAiApi = new OpenAiApi(baseUrl, apiKey, RestClient.builder());
var openAiApi = new OpenAiApi(baseUrl, apiKey, restClientBuilder);
return new OpenAiEmbeddingClient(openAiApi).withDefaultOptions(embeddingProperties.getOptions());
}
@@ -86,7 +87,7 @@ public class OpenAiAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public OpenAiImageClient openAiImageClient(OpenAiConnectionProperties commonProperties,
OpenAiImageProperties imageProperties) {
OpenAiImageProperties imageProperties, RestClient.Builder restClientBuilder) {
String apiKey = StringUtils.hasText(imageProperties.getApiKey()) ? imageProperties.getApiKey()
: commonProperties.getApiKey();
@@ -96,7 +97,7 @@ public class OpenAiAutoConfiguration {
Assert.hasText(apiKey, "OpenAI API key must be set");
Assert.hasText(baseUrl, "OpenAI base URL must be set");
var openAiImageApi = new OpenAiImageApi(baseUrl, apiKey, RestClient.builder());
var openAiImageApi = new OpenAiImageApi(baseUrl, apiKey, restClientBuilder);
return new OpenAiImageClient(openAiImageApi).withDefaultOptions(imageProperties.getOptions());
}

View File

@@ -21,15 +21,17 @@ import org.springframework.ai.stabilityai.api.StabilityAiApi;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.ImportRuntimeHints;
import org.springframework.web.client.RestClient;
/**
* @author Mark Pollack
* @since 0.8.0
*/
@AutoConfiguration
@AutoConfiguration(after = RestClientAutoConfiguration.class)
@ConditionalOnClass(StabilityAiApi.class)
@EnableConfigurationProperties({ StabilityAiProperties.class })
@ImportRuntimeHints(NativeHints.class)
@@ -37,9 +39,10 @@ public class StabilityAiImageAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public StabilityAiApi stabilityAiApi(StabilityAiProperties stabilityAiProperties) {
public StabilityAiApi stabilityAiApi(StabilityAiProperties stabilityAiProperties,
RestClient.Builder restClientBuilder) {
return new StabilityAiApi(stabilityAiProperties.getApiKey(), stabilityAiProperties.getBaseUrl(),
stabilityAiProperties.getOptions().getModel());
stabilityAiProperties.getOptions().getModel(), restClientBuilder);
}
@Bean

View File

@@ -23,12 +23,13 @@ import org.springframework.ai.vertex.VertexAiChatClient;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.ImportRuntimeHints;
import org.springframework.web.client.RestClient;
@AutoConfiguration
@AutoConfiguration(after = RestClientAutoConfiguration.class)
@ConditionalOnClass(VertexAiApi.class)
@ImportRuntimeHints(NativeHints.class)
@EnableConfigurationProperties({ VertexAiConnectionProperties.class, VertexAiChatProperties.class,
@@ -56,10 +57,11 @@ public class VertexAiAutoConfiguration {
@Bean
@ConditionalOnMissingBean
public VertexAiApi vertexAiApi(VertexAiConnectionProperties connectionProperties,
VertexAiEmbeddingProperties embeddingAiProperties, VertexAiChatProperties chatProperties) {
VertexAiEmbeddingProperties embeddingAiProperties, VertexAiChatProperties chatProperties,
RestClient.Builder restClientBuilder) {
return new VertexAiApi(connectionProperties.getBaseUrl(), connectionProperties.getApiKey(),
chatProperties.getModel(), embeddingAiProperties.getModel(), RestClient.builder());
chatProperties.getModel(), embeddingAiProperties.getModel(), restClientBuilder);
}
}

View File

@@ -30,6 +30,7 @@ import org.springframework.ai.chat.prompt.SystemPromptTemplate;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.testcontainers.containers.GenericContainer;
import org.testcontainers.junit.jupiter.Container;
@@ -71,13 +72,12 @@ public class OllamaAutoConfigurationIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner().withPropertyValues(
// @formatter:off
"spring.ai.ollama.chat.enabled=true",
"spring.ai.ollama.chat.options.model=" + MODEL_NAME,
"spring.ai.ollama.baseUrl=" + baseUrl,
"spring.ai.ollama.chat.options.model=" + MODEL_NAME,
"spring.ai.ollama.chat.options.temperature=0.5",
"spring.ai.ollama.chat.options.topK=10")
// @formatter:on
.withConfiguration(AutoConfigurations.of(OllamaAutoConfiguration.class));
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OllamaAutoConfiguration.class));
private final Message systemMessage = new SystemPromptTemplate("""
You are a helpful AI assistant. Your name is {name}.

View File

@@ -19,6 +19,7 @@ package org.springframework.ai.autoconfigure.ollama;
import org.junit.jupiter.api.Test;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import static org.assertj.core.api.Assertions.assertThat;
@@ -40,7 +41,7 @@ public class OllamaAutoConfigurationTests {
"spring.ai.ollama.chat.options.topP=0.56",
"spring.ai.ollama.chat.options.topK=123")
// @formatter:on
.withConfiguration(AutoConfigurations.of(OllamaAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OllamaAutoConfiguration.class))
.run(context -> {
var chatProperties = context.getBean(OllamaChatProperties.class);
var connectionProperties = context.getBean(OllamaConnectionProperties.class);

View File

@@ -24,6 +24,7 @@ import org.junit.jupiter.api.Test;
import org.springframework.ai.embedding.EmbeddingResponse;
import org.springframework.ai.ollama.OllamaEmbeddingClient;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.testcontainers.containers.GenericContainer;
import org.testcontainers.junit.jupiter.Container;
@@ -63,7 +64,7 @@ public class OllamaEmbeddingAutoConfigurationIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.ollama.embedding.options.model=" + MODEL_NAME,
"spring.ai.ollama.base-url=" + baseUrl)
.withConfiguration(AutoConfigurations.of(OllamaAutoConfiguration.class));
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OllamaAutoConfiguration.class));
@Test
public void singleTextEmbedding() {

View File

@@ -19,6 +19,7 @@ package org.springframework.ai.autoconfigure.ollama;
import org.junit.jupiter.api.Test;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import static org.assertj.core.api.Assertions.assertThat;
@@ -32,10 +33,15 @@ public class OllamaEmbeddingAutoConfigurationTests {
@Test
public void propertiesTest() {
new ApplicationContextRunner()
.withPropertyValues("spring.ai.ollama.base-url=TEST_BASE_URL", "spring.ai.ollama.embedding.model=MODEL_XYZ",
"spring.ai.ollama.embedding.options.temperature=0.13", "spring.ai.ollama.embedding.options.topK=13")
.withConfiguration(AutoConfigurations.of(OllamaAutoConfiguration.class))
new ApplicationContextRunner().withPropertyValues(
// @formatter:off
"spring.ai.ollama.base-url=TEST_BASE_URL",
"spring.ai.ollama.embedding.options.model=MODEL_XYZ",
"spring.ai.ollama.embedding.options.temperature=0.13",
"spring.ai.ollama.embedding.options.topK=13"
// @formatter:on
)
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OllamaAutoConfiguration.class))
.run(context -> {
var embeddingProperties = context.getBean(OllamaEmbeddingProperties.class);
var connectionProperties = context.getBean(OllamaConnectionProperties.class);

View File

@@ -35,6 +35,7 @@ import org.springframework.ai.embedding.EmbeddingResponse;
import org.springframework.ai.openai.OpenAiChatClient;
import org.springframework.ai.openai.OpenAiEmbeddingClient;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import static org.assertj.core.api.Assertions.assertThat;
@@ -46,7 +47,7 @@ public class OpenAiAutoConfigurationIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY"))
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class));
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class));
@Test
void generate() {

View File

@@ -24,6 +24,7 @@ import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionRequest.Respons
import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionRequest.ToolChoice;
import org.springframework.ai.openai.api.OpenAiApi.FunctionTool.Type;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import static org.assertj.core.api.Assertions.assertThat;
@@ -48,7 +49,7 @@ public class OpenAiPropertiesTests {
"spring.ai.openai.chat.options.model=MODEL_XYZ",
"spring.ai.openai.chat.options.temperature=0.55")
// @formatter:on
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class))
.run(context -> {
var chatProperties = context.getBean(OpenAiChatProperties.class);
var connectionProperties = context.getBean(OpenAiConnectionProperties.class);
@@ -76,7 +77,7 @@ public class OpenAiPropertiesTests {
"spring.ai.openai.chat.options.model=MODEL_XYZ",
"spring.ai.openai.chat.options.temperature=0.55")
// @formatter:on
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class))
.run(context -> {
var chatProperties = context.getBean(OpenAiChatProperties.class);
var connectionProperties = context.getBean(OpenAiConnectionProperties.class);
@@ -101,7 +102,7 @@ public class OpenAiPropertiesTests {
"spring.ai.openai.api-key=abc123",
"spring.ai.openai.embedding.options.model=MODEL_XYZ")
// @formatter:on
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class))
.run(context -> {
var embeddingProperties = context.getBean(OpenAiEmbeddingProperties.class);
var connectionProperties = context.getBean(OpenAiConnectionProperties.class);
@@ -127,7 +128,7 @@ public class OpenAiPropertiesTests {
"spring.ai.openai.embedding.api-key=456",
"spring.ai.openai.embedding.options.model=MODEL_XYZ")
// @formatter:on
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class))
.run(context -> {
var embeddingProperties = context.getBean(OpenAiEmbeddingProperties.class);
var connectionProperties = context.getBean(OpenAiConnectionProperties.class);
@@ -151,7 +152,7 @@ public class OpenAiPropertiesTests {
"spring.ai.openai.image.options.model=MODEL_XYZ",
"spring.ai.openai.image.options.n=3")
// @formatter:on
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class))
.run(context -> {
var imageProperties = context.getBean(OpenAiImageProperties.class);
var connectionProperties = context.getBean(OpenAiConnectionProperties.class);
@@ -178,7 +179,7 @@ public class OpenAiPropertiesTests {
"spring.ai.openai.image.options.model=MODEL_XYZ",
"spring.ai.openai.image.options.n=3")
// @formatter:on
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class))
.run(context -> {
var imageProperties = context.getBean(OpenAiImageProperties.class);
var connectionProperties = context.getBean(OpenAiConnectionProperties.class);
@@ -245,7 +246,7 @@ public class OpenAiPropertiesTests {
"spring.ai.openai.chat.options.user=userXYZ"
)
// @formatter:on
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class))
.run(context -> {
var chatProperties = context.getBean(OpenAiChatProperties.class);
var connectionProperties = context.getBean(OpenAiConnectionProperties.class);
@@ -295,7 +296,7 @@ public class OpenAiPropertiesTests {
"spring.ai.openai.embedding.options.user=userXYZ"
)
// @formatter:on
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class))
.run(context -> {
var connectionProperties = context.getBean(OpenAiConnectionProperties.class);
var embeddingProperties = context.getBean(OpenAiEmbeddingProperties.class);
@@ -327,7 +328,7 @@ public class OpenAiPropertiesTests {
"spring.ai.openai.image.options.user=userXYZ"
)
// @formatter:on
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class))
.run(context -> {
var imageProperties = context.getBean(OpenAiImageProperties.class);
var connectionProperties = context.getBean(OpenAiConnectionProperties.class);

View File

@@ -27,6 +27,7 @@ import org.springframework.ai.embedding.EmbeddingResponse;
import org.springframework.ai.vertex.VertexAiEmbeddingClient;
import org.springframework.ai.vertex.VertexAiChatClient;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import static org.assertj.core.api.Assertions.assertThat;
@@ -41,7 +42,7 @@ public class VertexAiAutoConfigurationIT {
"spring.ai.vertex.ai.apiKey=" + System.getenv("PALM_API_KEY"),
"spring.ai.vertex.ai.chat.model=chat-bison-001",
"spring.ai.vertex.ai.embedding.model=embedding-gecko-001")
.withConfiguration(AutoConfigurations.of(VertexAiAutoConfiguration.class));
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, VertexAiAutoConfiguration.class));
@Test
void generate() {