Various fixes and stability improvements

- Fix a bug in Azure streaming response. Ensure that the merge functionality resolves the right
   object constructors
 - drop the @ConditionalOnMissingBean for the ChatClientAutoConfiguration#chatClientBuilder .
   If multiple chat model starters are added to the POM this will fail as the ChatClient.Builder
    auto-config can handle only one chat model. Then the spring.ai.chat.client.enabled=false must be set.
 - Add missing AutoConfiguration imports for SpringAiRetryAutoConfiguration.class, RestClientAutoConfiguration.class,
   and WebClientAutoConfiguration.class to the AnthropicAutoConfiguration, MistralAiAutoConfiguration,
   OllamaAutoConfiguration,VertexAiPalm2AutoConfiguration.
 - change the OpenAi and Azure OpenAi default chat models to  gpt-4o
 - clean and improve the stability of various ITs
This commit is contained in:
Christian Tzolov
2024-06-07 14:26:17 +02:00
parent 64378dfcdd
commit a4412b5d3a
31 changed files with 221 additions and 257 deletions

View File

@@ -24,10 +24,12 @@ import java.util.Objects;
import com.azure.ai.openai.models.AzureChatExtensionsMessageContext;
import com.azure.ai.openai.models.ChatChoice;
import com.azure.ai.openai.models.ChatChoiceLogProbabilityInfo;
import com.azure.ai.openai.models.ChatCompletions;
import com.azure.ai.openai.models.ChatCompletionsFunctionToolCall;
import com.azure.ai.openai.models.ChatCompletionsToolCall;
import com.azure.ai.openai.models.ChatResponseMessage;
import com.azure.ai.openai.models.ChatRole;
import com.azure.ai.openai.models.CompletionsFinishReason;
import com.azure.ai.openai.models.CompletionsUsage;
import com.azure.ai.openai.models.ContentFilterResultsForChoice;
@@ -47,31 +49,20 @@ import org.springframework.util.CollectionUtils;
*/
public class MergeUtils {
/**
* Create a new instance of the given class. Can be used to create instances with
* private constructors.
* @param <T> the type of the class to be created.
* @param clazz the class to create an instance of.
* @param args the arguments to pass to the constructor.
* @return a new instance of the given class.
*/
private static <T> T newInstance(Class<T> clazz, Object... args) {
return newInstance(0, clazz, args);
}
/**
* Create a new instance of the given class using the constructor at the given index.
* Can be used to create instances with private constructors.
* @param <T> the type of the class to be created.
* @param index the index of the constructor to use.
* @param argumentTypes the list of constructor argument types. Used to select the
* right constructor.
* @param clazz the class to create an instance of.
* @param args the arguments to pass to the constructor.
* @return a new instance of the given class.
*/
private static <T> T newInstance(int index, Class<T> clazz, Object... args) {
private static <T> T newInstance(Class<?>[] argumentTypes, Class<T> clazz, Object... args) {
try {
@SuppressWarnings("unchecked")
Constructor<T> constructor = (Constructor<T>) clazz.getDeclaredConstructors()[index];
Constructor<T> constructor = (Constructor<T>) clazz.getDeclaredConstructor(argumentTypes);
constructor.setAccessible(true);
return constructor.newInstance(args);
}
@@ -97,6 +88,9 @@ public class MergeUtils {
}
}
private static final Class<?>[] chatCompletionsConstructorArgumentTypes = new Class<?>[] { String.class, long.class,
List.class, CompletionsUsage.class };
/**
* @return an empty ChatCompletions instance.
*/
@@ -105,7 +99,8 @@ public class MergeUtils {
List<ChatChoice> choices = new ArrayList<>();
CompletionsUsage usage = null;
long createdAt = 0;
ChatCompletions chatCompletionsInstance = newInstance(ChatCompletions.class, id, createdAt, choices, usage);
ChatCompletions chatCompletionsInstance = newInstance(chatCompletionsConstructorArgumentTypes,
ChatCompletions.class, id, createdAt, choices, usage);
List<ContentFilterResultsForPrompt> promptFilterResults = new ArrayList<>();
setField(chatCompletionsInstance, "promptFilterResults", promptFilterResults);
String systemFingerprint = null;
@@ -114,6 +109,9 @@ public class MergeUtils {
return chatCompletionsInstance;
}
private static final Class<?>[] chatCompletionsConstructorArgumentTypes0 = new Class<?>[] { String.class,
OffsetDateTime.class, List.class, CompletionsUsage.class };
/**
* Merge two ChatCompletions instances into a single ChatCompletions instance.
* @param left the left ChatCompletions instance.
@@ -150,7 +148,8 @@ public class MergeUtils {
OffsetDateTime createdAt = left.getCreatedAt().isAfter(right.getCreatedAt()) ? left.getCreatedAt()
: right.getCreatedAt();
ChatCompletions instance = newInstance(1, ChatCompletions.class, id, createdAt, choices, usage);
ChatCompletions instance = newInstance(chatCompletionsConstructorArgumentTypes0, ChatCompletions.class, id,
createdAt, choices, usage);
List<ContentFilterResultsForPrompt> promptFilterResults = right.getPromptFilterResults() == null
? left.getPromptFilterResults() : right.getPromptFilterResults();
@@ -162,6 +161,9 @@ public class MergeUtils {
return instance;
}
private static final Class<?>[] chatChoiceConstructorArgumentTypes = new Class<?>[] {
ChatChoiceLogProbabilityInfo.class, int.class, CompletionsFinishReason.class };
/**
* Merge two ChatChoice instances into a single ChatChoice instance.
* @param left the left ChatChoice instance to merge.
@@ -177,7 +179,8 @@ public class MergeUtils {
var logprobs = left.getLogprobs() != null ? left.getLogprobs() : right.getLogprobs();
final ChatChoice instance = newInstance(ChatChoice.class, logprobs, index, finishReason);
final ChatChoice instance = newInstance(chatChoiceConstructorArgumentTypes, ChatChoice.class, logprobs, index,
finishReason);
ChatResponseMessage message = null;
if (left.getMessage() == null) {
@@ -211,6 +214,9 @@ public class MergeUtils {
return instance;
}
private static final Class<?>[] chatResponseMessageConstructorArgumentTypes = new Class<?>[] { ChatRole.class,
String.class };
/**
* Merge two ChatResponseMessage instances into a single ChatResponseMessage instance.
* @param left the left ChatResponseMessage instance to merge.
@@ -231,7 +237,8 @@ public class MergeUtils {
content = left.getContent();
}
ChatResponseMessage instance = newInstance(ChatResponseMessage.class, role, content);
ChatResponseMessage instance = newInstance(chatResponseMessageConstructorArgumentTypes,
ChatResponseMessage.class, role, content);
List<ChatCompletionsToolCall> toolCalls = new ArrayList<>();
if (left.getToolCalls() == null) {

View File

@@ -188,7 +188,7 @@ class MistralAiChatModelIT {
@Test
void functionCallTest() {
UserMessage userMessage = new UserMessage("What's the weather like in San Francisco?");
UserMessage userMessage = new UserMessage("What's the weather like in San Francisco? Response in Celsius");
List<Message> messages = new ArrayList<>(List.of(userMessage));
@@ -211,7 +211,7 @@ class MistralAiChatModelIT {
@Test
void streamFunctionCallTest() {
UserMessage userMessage = new UserMessage("What's the weather like in Tokyo, Japan?");
UserMessage userMessage = new UserMessage("What's the weather like in Tokyo, Japan? Response in Celsius");
List<Message> messages = new ArrayList<>(List.of(userMessage));

View File

@@ -48,7 +48,7 @@ import org.springframework.web.reactive.function.client.WebClient;
*/
public class OpenAiApi {
public static final String DEFAULT_CHAT_MODEL = ChatModel.GPT_3_5_TURBO.getValue();
public static final OpenAiApi.ChatModel DEFAULT_CHAT_MODEL = ChatModel.GPT_4_O;
public static final String DEFAULT_EMBEDDING_MODEL = EmbeddingModel.TEXT_EMBEDDING_ADA_002.getValue();
private static final Predicate<String> SSE_DONE_PREDICATE = "[DONE]"::equals;

View File

@@ -63,9 +63,8 @@ public class OpenAiApiToolFunctionCallIT {
Role.USER);
var functionTool = new OpenAiApi.FunctionTool(Type.FUNCTION,
new OpenAiApi.FunctionTool.Function(
"Get the weather in location. Return temperature in 30°F or 30°C format.", "getCurrentWeather",
ModelOptionsUtils.jsonToMap("""
new OpenAiApi.FunctionTool.Function("Get the weather in location. Return temperature in Celsius.",
"getCurrentWeather", ModelOptionsUtils.jsonToMap("""
{
"type": "object",
"properties": {
@@ -92,7 +91,7 @@ public class OpenAiApiToolFunctionCallIT {
List<ChatCompletionMessage> messages = new ArrayList<>(List.of(message));
ChatCompletionRequest chatCompletionRequest = new ChatCompletionRequest(messages, "gpt-4-turbo-preview",
ChatCompletionRequest chatCompletionRequest = new ChatCompletionRequest(messages, "gpt-4o",
List.of(functionTool), ToolChoiceBuilder.AUTO);
// List.of(functionTool), ToolChoiceBuilder.FUNCTION("getCurrentWeather"));
@@ -127,7 +126,7 @@ public class OpenAiApiToolFunctionCallIT {
}
}
var functionResponseRequest = new ChatCompletionRequest(messages, "gpt-4-turbo-preview", 0.8f);
var functionResponseRequest = new ChatCompletionRequest(messages, "gpt-4o", 0.5f);
ResponseEntity<ChatCompletion> chatCompletion2 = completionApi
.chatCompletionEntity(functionResponseRequest);

View File

@@ -219,7 +219,7 @@ class OpenAiChatModelIT extends AbstractIT {
List<Message> messages = new ArrayList<>(List.of(userMessage));
var promptOptions = OpenAiChatOptions.builder()
.withModel(OpenAiApi.ChatModel.GPT_4_TURBO_PREVIEW.getValue())
.withModel(OpenAiApi.ChatModel.GPT_4_O.getValue())
.withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService())
.withName("getCurrentWeather")
.withDescription("Get the weather in location")

View File

@@ -23,10 +23,12 @@ import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration
import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.model.function.FunctionCallbackContext;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.ImportAutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.autoconfigure.web.reactive.function.client.WebClientAutoConfiguration;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.ApplicationContext;
import org.springframework.context.annotation.Bean;
@@ -44,6 +46,8 @@ import org.springframework.web.client.RestClient;
@ConditionalOnClass(AnthropicApi.class)
@ConditionalOnProperty(prefix = AnthropicChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true",
matchIfMissing = true)
@ImportAutoConfiguration(classes = { SpringAiRetryAutoConfiguration.class, RestClientAutoConfiguration.class,
WebClientAutoConfiguration.class })
public class AnthropicAutoConfiguration {
@Bean

View File

@@ -40,7 +40,7 @@ import org.springframework.context.annotation.Scope;
* @author Mark Pollack
* @author Josh Long
* @author Arjen Poutsma
* @since 1.0.0 M1
* @since 1.0.0
*/
@AutoConfiguration
@ConditionalOnClass(ChatClient.class)
@@ -59,7 +59,6 @@ public class ChatClientAutoConfiguration {
@Bean
@Scope("prototype")
@ConditionalOnMissingBean
ChatClient.Builder chatClientBuilder(ChatClientBuilderConfigurer chatClientBuilderConfigurer, ChatModel chatModel) {
ChatClient.Builder builder = ChatClient.builder(chatModel);
return chatClientBuilderConfigurer.configure(builder);

View File

@@ -24,10 +24,12 @@ import org.springframework.ai.mistralai.api.MistralAiApi;
import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.model.function.FunctionCallbackContext;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.ImportAutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.autoconfigure.web.reactive.function.client.WebClientAutoConfiguration;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.ApplicationContext;
import org.springframework.context.annotation.Bean;
@@ -47,6 +49,8 @@ import org.springframework.web.client.RestClient;
@EnableConfigurationProperties({ MistralAiEmbeddingProperties.class, MistralAiCommonProperties.class,
MistralAiChatProperties.class })
@ConditionalOnClass(MistralAiApi.class)
@ImportAutoConfiguration(classes = { SpringAiRetryAutoConfiguration.class, RestClientAutoConfiguration.class,
WebClientAutoConfiguration.class })
public class MistralAiAutoConfiguration {
@Bean

View File

@@ -19,10 +19,12 @@ import org.springframework.ai.ollama.OllamaChatModel;
import org.springframework.ai.ollama.OllamaEmbeddingModel;
import org.springframework.ai.ollama.api.OllamaApi;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.ImportAutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.autoconfigure.web.reactive.function.client.WebClientAutoConfiguration;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.annotation.Bean;
import org.springframework.web.client.RestClient;
@@ -38,6 +40,7 @@ import org.springframework.web.client.RestClient;
@ConditionalOnClass(OllamaApi.class)
@EnableConfigurationProperties({ OllamaChatProperties.class, OllamaEmbeddingProperties.class,
OllamaConnectionProperties.class })
@ImportAutoConfiguration(classes = { RestClientAutoConfiguration.class, WebClientAutoConfiguration.class })
public class OllamaAutoConfiguration {
@Bean

View File

@@ -24,7 +24,7 @@ public class OpenAiChatProperties extends OpenAiParentProperties {
public static final String CONFIG_PREFIX = "spring.ai.openai.chat";
public static final String DEFAULT_CHAT_MODEL = "gpt-3.5-turbo";
public static final String DEFAULT_CHAT_MODEL = "gpt-4o";
private static final Double DEFAULT_TEMPERATURE = 0.7;

View File

@@ -15,14 +15,17 @@
*/
package org.springframework.ai.autoconfigure.vertexai.palm2;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.vertexai.palm2.VertexAiPaLm2ChatModel;
import org.springframework.ai.vertexai.palm2.VertexAiPaLm2EmbeddingModel;
import org.springframework.ai.vertexai.palm2.api.VertexAiPaLm2Api;
import org.springframework.boot.autoconfigure.AutoConfiguration;
import org.springframework.boot.autoconfigure.ImportAutoConfiguration;
import org.springframework.boot.autoconfigure.condition.ConditionalOnClass;
import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean;
import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.autoconfigure.web.reactive.function.client.WebClientAutoConfiguration;
import org.springframework.boot.context.properties.EnableConfigurationProperties;
import org.springframework.context.annotation.Bean;
import org.springframework.web.client.RestClient;
@@ -31,6 +34,8 @@ import org.springframework.web.client.RestClient;
@ConditionalOnClass(VertexAiPaLm2Api.class)
@EnableConfigurationProperties({ VertexAiPalm2ConnectionProperties.class, VertexAiPlam2ChatProperties.class,
VertexAiPalm2EmbeddingProperties.class })
@ImportAutoConfiguration(classes = { SpringAiRetryAutoConfiguration.class, RestClientAutoConfiguration.class,
WebClientAutoConfiguration.class })
public class VertexAiPalm2AutoConfiguration {
@Bean

View File

@@ -15,6 +15,8 @@
*/
package org.springframework.ai.autoconfigure.anthropic;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.List;
import java.util.stream.Collectors;
@@ -23,19 +25,15 @@ import org.apache.commons.logging.LogFactory;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.anthropic.AnthropicChatModel;
import reactor.core.publisher.Flux;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.chat.prompt.Prompt;
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;
import reactor.core.publisher.Flux;
@EnabledIfEnvironmentVariable(named = "ANTHROPIC_API_KEY", matches = ".*")
public class AnthropicAutoConfigurationIT {
@@ -44,8 +42,7 @@ public class AnthropicAutoConfigurationIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.anthropic.apiKey=" + System.getenv("ANTHROPIC_API_KEY"))
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, AnthropicAutoConfiguration.class));
.withConfiguration(AutoConfigurations.of(AnthropicAutoConfiguration.class));
@Test
void generate() {

View File

@@ -15,6 +15,8 @@
*/
package org.springframework.ai.autoconfigure.anthropic.tool;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.List;
import java.util.function.Function;
@@ -22,26 +24,21 @@ 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.anthropic.AnthropicChatModel;
import org.springframework.ai.anthropic.AnthropicChatOptions;
import org.springframework.ai.anthropic.api.AnthropicApi;
import org.springframework.ai.autoconfigure.anthropic.AnthropicAutoConfiguration;
import org.springframework.ai.autoconfigure.anthropic.tool.MockWeatherService.Request;
import org.springframework.ai.autoconfigure.anthropic.tool.MockWeatherService.Response;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.context.annotation.Description;
import static org.assertj.core.api.Assertions.assertThat;
@EnabledIfEnvironmentVariable(named = "ANTHROPIC_API_KEY", matches = ".*")
class FunctionCallWithFunctionBeanIT {
@@ -49,8 +46,7 @@ class FunctionCallWithFunctionBeanIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.anthropic.apiKey=" + System.getenv("ANTHROPIC_API_KEY"))
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, AnthropicAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(AnthropicAutoConfiguration.class))
.withUserConfiguration(Config.class);
@Test

View File

@@ -15,28 +15,25 @@
*/
package org.springframework.ai.autoconfigure.anthropic.tool;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.List;
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.anthropic.AnthropicChatModel;
import org.springframework.ai.anthropic.AnthropicChatOptions;
import org.springframework.ai.anthropic.api.AnthropicApi;
import org.springframework.ai.autoconfigure.anthropic.AnthropicAutoConfiguration;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.model.function.FunctionCallbackWrapper;
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;
@EnabledIfEnvironmentVariable(named = "ANTHROPIC_API_KEY", matches = ".*")
public class FunctionCallWithPromptFunctionIT {
@@ -44,8 +41,7 @@ public class FunctionCallWithPromptFunctionIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.anthropic.apiKey=" + System.getenv("ANTHROPIC_API_KEY"))
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, AnthropicAutoConfiguration.class));
.withConfiguration(AutoConfigurations.of(AnthropicAutoConfiguration.class));
@Test
void functionCallTest() {

View File

@@ -47,7 +47,7 @@ import static org.assertj.core.api.Assertions.assertThat;
@EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_ENDPOINT", matches = ".+")
public class AzureOpenAiAutoConfigurationIT {
private static String CHAT_MODEL_NAME = "gpt-35-turbo";
private static String CHAT_MODEL_NAME = "gpt-4o";
private static String EMBEDDING_MODEL_NAME = "text-embedding-ada-002";

View File

@@ -44,7 +44,7 @@ public class ChatClientAutoConfigurationIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY"),
"spring.ai.openai.chat.options.model=gpt-4-turbo")
"spring.ai.openai.chat.options.model=gpt-4o")
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class, ChatClientAutoConfiguration.class));

View File

@@ -15,6 +15,8 @@
*/
package org.springframework.ai.autoconfigure.mistralai;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.List;
import java.util.stream.Collectors;
@@ -22,20 +24,16 @@ import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.mistralai.MistralAiChatModel;
import reactor.core.publisher.Flux;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.embedding.EmbeddingResponse;
import org.springframework.ai.mistralai.MistralAiChatModel;
import org.springframework.ai.mistralai.MistralAiEmbeddingModel;
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;
import reactor.core.publisher.Flux;
/**
* @author Christian Tzolov
@@ -48,8 +46,7 @@ public class MistralAiAutoConfigurationIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.mistralai.apiKey=" + System.getenv("MISTRAL_AI_API_KEY"))
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, MistralAiAutoConfiguration.class));
.withConfiguration(AutoConfigurations.of(MistralAiAutoConfiguration.class));
@Test
void generate() {

View File

@@ -15,32 +15,30 @@
*/
package org.springframework.ai.autoconfigure.mistralai.tool;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.List;
import java.util.Map;
import java.util.function.Function;
import com.fasterxml.jackson.annotation.JsonProperty;
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.autoconfigure.mistralai.MistralAiAutoConfiguration;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.mistralai.MistralAiChatModel;
import org.springframework.ai.mistralai.MistralAiChatOptions;
import org.springframework.ai.mistralai.api.MistralAiApi;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.context.annotation.Description;
import static org.assertj.core.api.Assertions.assertThat;
import com.fasterxml.jackson.annotation.JsonProperty;
@EnabledIfEnvironmentVariable(named = "MISTRAL_AI_API_KEY", matches = ".*")
class PaymentStatusBeanIT {
@@ -49,8 +47,7 @@ class PaymentStatusBeanIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.mistralai.apiKey=" + System.getenv("MISTRAL_AI_API_KEY"))
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, MistralAiAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(MistralAiAutoConfiguration.class))
.withUserConfiguration(Config.class);
@Test

View File

@@ -15,32 +15,30 @@
*/
package org.springframework.ai.autoconfigure.mistralai.tool;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.List;
import java.util.Map;
import java.util.function.Function;
import com.fasterxml.jackson.annotation.JsonProperty;
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.autoconfigure.openai.OpenAiAutoConfiguration;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.mistralai.api.MistralAiApi;
import org.springframework.ai.openai.OpenAiChatModel;
import org.springframework.ai.openai.OpenAiChatOptions;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.context.annotation.Description;
import static org.assertj.core.api.Assertions.assertThat;
import com.fasterxml.jackson.annotation.JsonProperty;
/**
* Same test as {@link PaymentStatusBeanIT.java} but using {@link OpenAiChatModel} for
@@ -56,8 +54,7 @@ class PaymentStatusBeanOpenAiIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("MISTRAL_AI_API_KEY"),
"spring.ai.openai.chat.base-url=https://api.mistral.ai")
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withUserConfiguration(Config.class);
@Test

View File

@@ -15,30 +15,28 @@
*/
package org.springframework.ai.autoconfigure.mistralai.tool;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.List;
import java.util.Map;
import java.util.function.Function;
import com.fasterxml.jackson.annotation.JsonProperty;
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.autoconfigure.mistralai.MistralAiAutoConfiguration;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.mistralai.MistralAiChatModel;
import org.springframework.ai.mistralai.MistralAiChatOptions;
import org.springframework.ai.mistralai.api.MistralAiApi;
import org.springframework.ai.model.function.FunctionCallbackWrapper;
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;
import com.fasterxml.jackson.annotation.JsonProperty;
@EnabledIfEnvironmentVariable(named = "MISTRAL_AI_API_KEY", matches = ".*")
public class PaymentStatusPromptIT {
@@ -47,8 +45,7 @@ public class PaymentStatusPromptIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.mistralai.apiKey=" + System.getenv("MISTRAL_AI_API_KEY"))
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, MistralAiAutoConfiguration.class));
.withConfiguration(AutoConfigurations.of(MistralAiAutoConfiguration.class));
public record Transaction(@JsonProperty(required = true, value = "transaction_id") String id) {
}

View File

@@ -15,23 +15,20 @@
*/
package org.springframework.ai.autoconfigure.mistralai.tool;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.List;
import java.util.function.Function;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonInclude.Include;
import com.fasterxml.jackson.annotation.JsonProperty;
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.autoconfigure.mistralai.MistralAiAutoConfiguration;
import org.springframework.ai.autoconfigure.mistralai.tool.WeatherServicePromptIT.MyWeatherService.Request;
import org.springframework.ai.autoconfigure.mistralai.tool.WeatherServicePromptIT.MyWeatherService.Response;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.mistralai.MistralAiChatModel;
import org.springframework.ai.mistralai.MistralAiChatOptions;
@@ -39,10 +36,11 @@ import org.springframework.ai.mistralai.api.MistralAiApi;
import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionRequest.ToolChoice;
import org.springframework.ai.model.function.FunctionCallbackWrapper;
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;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonInclude.Include;
import com.fasterxml.jackson.annotation.JsonProperty;
/**
* @author Christian Tzolov
@@ -55,8 +53,7 @@ public class WeatherServicePromptIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.mistralai.api-key=" + System.getenv("MISTRAL_AI_API_KEY"))
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, MistralAiAutoConfiguration.class));
.withConfiguration(AutoConfigurations.of(MistralAiAutoConfiguration.class));
@Test
void promptFunctionCall() {

View File

@@ -15,30 +15,7 @@
*/
package org.springframework.ai.autoconfigure.ollama;
import com.github.dockerjava.api.DockerClient;
import com.github.dockerjava.api.command.InspectContainerResponse;
import com.github.dockerjava.api.model.Image;
import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.ollama.OllamaChatModel;
import org.springframework.ai.chat.prompt.Prompt;
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.DockerClientFactory;
import org.testcontainers.containers.GenericContainer;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.utility.DockerImageName;
import reactor.core.publisher.Flux;
import static org.assertj.core.api.Assertions.assertThat;
import java.io.IOException;
import java.util.Collections;
@@ -46,7 +23,31 @@ import java.util.List;
import java.util.Map;
import java.util.stream.Collectors;
import static org.assertj.core.api.Assertions.assertThat;
import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Disabled;
import org.junit.jupiter.api.Test;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.chat.prompt.SystemPromptTemplate;
import org.springframework.ai.ollama.OllamaChatModel;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.testcontainers.DockerClientFactory;
import org.testcontainers.containers.GenericContainer;
import org.testcontainers.junit.jupiter.Testcontainers;
import org.testcontainers.utility.DockerImageName;
import com.github.dockerjava.api.DockerClient;
import com.github.dockerjava.api.command.InspectContainerResponse;
import com.github.dockerjava.api.model.Image;
import reactor.core.publisher.Flux;
/**
* @author Christian Tzolov
@@ -89,7 +90,7 @@ public class OllamaChatAutoConfigurationIT {
"spring.ai.ollama.chat.options.temperature=0.5",
"spring.ai.ollama.chat.options.topK=10")
// @formatter:on
.withConfiguration(AutoConfigurations.of(RestClientAutoConfiguration.class, OllamaAutoConfiguration.class));
.withConfiguration(AutoConfigurations.of(OllamaAutoConfiguration.class));
private final Message systemMessage = new SystemPromptTemplate("""
You are a helpful AI assistant. Your name is {name}.

View File

@@ -15,6 +15,8 @@
*/
package org.springframework.ai.autoconfigure.openai;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.Arrays;
import java.util.List;
import java.util.stream.Collectors;
@@ -24,24 +26,23 @@ import org.apache.commons.logging.LogFactory;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.embedding.EmbeddingResponse;
import org.springframework.ai.image.ImagePrompt;
import org.springframework.ai.image.ImageResponse;
import org.springframework.ai.openai.*;
import org.springframework.ai.openai.OpenAiAudioSpeechModel;
import org.springframework.ai.openai.OpenAiAudioTranscriptionModel;
import org.springframework.ai.openai.OpenAiChatModel;
import org.springframework.ai.openai.OpenAiEmbeddingModel;
import org.springframework.ai.openai.OpenAiImageModel;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.core.io.ClassPathResource;
import org.springframework.core.io.Resource;
import reactor.core.publisher.Flux;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.embedding.EmbeddingResponse;
import org.springframework.ai.openai.OpenAiChatModel;
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;
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".*")
public class OpenAiAutoConfigurationIT {
@@ -49,8 +50,7 @@ public class OpenAiAutoConfigurationIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY"))
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class));
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class));
@Test
void generate() {

View File

@@ -15,6 +15,8 @@
*/
package org.springframework.ai.autoconfigure.openai.tool;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.function.Function;
import java.util.stream.Collectors;
@@ -22,17 +24,12 @@ 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.autoconfigure.openai.OpenAiAutoConfiguration;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.client.ChatClient;
import org.springframework.ai.openai.OpenAiChatModel;
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;
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".*")
public class FunctionCallbackInPrompt2IT {
@@ -40,12 +37,11 @@ public class FunctionCallbackInPrompt2IT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY"))
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class));
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class));
@Test
void functionCallTest() {
contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo").run(context -> {
contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4o").run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
@@ -72,7 +68,7 @@ public class FunctionCallbackInPrompt2IT {
@Test
void functionCallTest2() {
contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo").run(context -> {
contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4o").run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
@@ -97,7 +93,7 @@ public class FunctionCallbackInPrompt2IT {
@Test
void streamingFunctionCallTest() {
contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> {
contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4o").run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);

View File

@@ -15,6 +15,8 @@
*/
package org.springframework.ai.autoconfigure.openai.tool;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.List;
import java.util.stream.Collectors;
@@ -22,23 +24,19 @@ import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import reactor.core.publisher.Flux;
import org.springframework.ai.autoconfigure.openai.OpenAiAutoConfiguration;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.model.function.FunctionCallbackWrapper;
import org.springframework.ai.openai.OpenAiChatModel;
import org.springframework.ai.openai.OpenAiChatOptions;
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;
import reactor.core.publisher.Flux;
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".*")
public class FunctionCallbackInPromptIT {
@@ -47,13 +45,12 @@ public class FunctionCallbackInPromptIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY"))
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class));
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class));
@Test
void functionCallTest() {
contextRunner
.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo",
.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4o",
"spring.ai.openai.chat.options.temperature=0.1")
.run(context -> {
@@ -82,8 +79,8 @@ public class FunctionCallbackInPromptIT {
void streamingFunctionCallTest() {
contextRunner
.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo",
"spring.ai.openai.chat.options.temperature=0.1")
.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4o",
"spring.ai.openai.chat.options.temperature=0.5")
.run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);

View File

@@ -15,6 +15,8 @@
*/
package org.springframework.ai.autoconfigure.openai.tool;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.List;
import java.util.function.Function;
import java.util.stream.Collectors;
@@ -23,26 +25,22 @@ import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import reactor.core.publisher.Flux;
import org.springframework.ai.autoconfigure.openai.OpenAiAutoConfiguration;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.client.ChatClient;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.openai.OpenAiChatOptions;
import org.springframework.ai.openai.OpenAiChatModel;
import org.springframework.ai.openai.OpenAiChatOptions;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.context.annotation.Description;
import static org.assertj.core.api.Assertions.assertThat;
import reactor.core.publisher.Flux;
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".*")
class FunctionCallbackWithPlainFunctionBeanIT {
@@ -51,13 +49,12 @@ class FunctionCallbackWithPlainFunctionBeanIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY"))
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withUserConfiguration(Config.class);
@Test
void functionCallTest() {
contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo").run(context -> {
contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4o").run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
@@ -86,7 +83,7 @@ class FunctionCallbackWithPlainFunctionBeanIT {
@Test
void functionCallWithPortableFunctionCallingOptions() {
contextRunner
.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo",
.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4o",
"spring.ai.openai.chat.options.temperature=0.1")
.run(context -> {
@@ -107,7 +104,7 @@ class FunctionCallbackWithPlainFunctionBeanIT {
@Test
void streamFunctionCallTest() {
contextRunner
.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo",
.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4o",
"spring.ai.openai.chat.options.temperature=0.1")
.run(context -> {

View File

@@ -15,27 +15,24 @@
*/
package org.springframework.ai.autoconfigure.openai.tool;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.stream.Collectors;
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.autoconfigure.openai.OpenAiAutoConfiguration;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.client.ChatClient;
import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.model.function.FunctionCallbackWrapper;
import org.springframework.ai.openai.OpenAiChatModel;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import static org.assertj.core.api.Assertions.assertThat;
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".*")
public class FunctionCallbackWrapper2IT {
@@ -43,18 +40,14 @@ public class FunctionCallbackWrapper2IT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY"))
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withUserConfiguration(Config.class);
@Test
void functionCallTest() {
contextRunner
.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo",
"spring.ai.openai.chat.options.temperature=0.1")
.run(context -> {
contextRunner.withPropertyValues("spring.ai.openai.chat.options.temperature=0.1").run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
// @formatter:off
ChatClient chatClient = ChatClient.builder(chatModel)
@@ -67,22 +60,19 @@ public class FunctionCallbackWrapper2IT {
.call().content();
// @formatter:on
logger.info("Response: {}", content);
logger.info("Response: {}", content);
assertThat(content).containsAnyOf("30.0", "30");
assertThat(content).containsAnyOf("15.0", "15");
assertThat(content).containsAnyOf("10", "10");
});
assertThat(content).containsAnyOf("30.0", "30");
assertThat(content).containsAnyOf("15.0", "15");
assertThat(content).containsAnyOf("10", "10");
});
}
@Test
void streamFunctionCallTest() {
contextRunner
.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo",
"spring.ai.openai.chat.options.temperature=0.1")
.run(context -> {
contextRunner.withPropertyValues("spring.ai.openai.chat.options.temperature=0.1").run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
// @formatter:off
String content = ChatClient.builder(chatModel).build().prompt()
@@ -92,12 +82,12 @@ public class FunctionCallbackWrapper2IT {
.collectList().block().stream().collect(Collectors.joining());
// @formatter:on
logger.info("Response: {}", content);
logger.info("Response: {}", content);
assertThat(content).containsAnyOf("30.0", "30");
assertThat(content).containsAnyOf("10.0", "10");
assertThat(content).containsAnyOf("15.0", "15");
});
assertThat(content).containsAnyOf("30.0", "30");
assertThat(content).containsAnyOf("10.0", "10");
assertThat(content).containsAnyOf("15.0", "15");
});
}
@Configuration

View File

@@ -15,6 +15,8 @@
*/
package org.springframework.ai.autoconfigure.openai.tool;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.List;
import java.util.stream.Collectors;
@@ -22,26 +24,22 @@ 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.openai.OpenAiChatModel;
import reactor.core.publisher.Flux;
import org.springframework.ai.autoconfigure.openai.OpenAiAutoConfiguration;
import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.model.function.FunctionCallbackWrapper;
import org.springframework.ai.model.function.FunctionCallback;
import org.springframework.ai.model.function.FunctionCallbackWrapper;
import org.springframework.ai.openai.OpenAiChatModel;
import org.springframework.ai.openai.OpenAiChatOptions;
import org.springframework.boot.autoconfigure.AutoConfigurations;
import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration;
import org.springframework.boot.test.context.runner.ApplicationContextRunner;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import static org.assertj.core.api.Assertions.assertThat;
import reactor.core.publisher.Flux;
@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".*")
public class FunctionCallbackWrapperIT {
@@ -50,62 +48,54 @@ public class FunctionCallbackWrapperIT {
private final ApplicationContextRunner contextRunner = new ApplicationContextRunner()
.withPropertyValues("spring.ai.openai.apiKey=" + System.getenv("OPENAI_API_KEY"))
.withConfiguration(AutoConfigurations.of(SpringAiRetryAutoConfiguration.class,
RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class))
.withConfiguration(AutoConfigurations.of(OpenAiAutoConfiguration.class))
.withUserConfiguration(Config.class);
@Test
void functionCallTest() {
contextRunner
.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo",
"spring.ai.openai.chat.options.temperature=0.1")
.run(context -> {
contextRunner.withPropertyValues("spring.ai.openai.chat.options.temperature=0.1").run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris?");
UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?");
ChatResponse response = chatModel.call(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder().withFunction("WeatherInfo").build()));
ChatResponse response = chatModel.call(
new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withFunction("WeatherInfo").build()));
logger.info("Response: {}", response);
logger.info("Response: {}", response);
assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15");
assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15");
});
});
}
@Test
void streamFunctionCallTest() {
contextRunner
.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo",
"spring.ai.openai.chat.options.temperature=0.1")
.run(context -> {
contextRunner.withPropertyValues("spring.ai.openai.chat.options.temperature=0.1").run(context -> {
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class);
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris? You can call the following functions 'WeatherInfo'");
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, Tokyo, and Paris? You can call the following functions 'WeatherInfo'");
Flux<ChatResponse> response = chatModel.stream(new Prompt(List.of(userMessage),
OpenAiChatOptions.builder().withFunction("WeatherInfo").build()));
Flux<ChatResponse> response = chatModel.stream(
new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withFunction("WeatherInfo").build()));
String content = response.collectList()
.block()
.stream()
.map(ChatResponse::getResults)
.flatMap(List::stream)
.map(Generation::getOutput)
.map(AssistantMessage::getContent)
.collect(Collectors.joining());
logger.info("Response: {}", content);
String content = response.collectList()
.block()
.stream()
.map(ChatResponse::getResults)
.flatMap(List::stream)
.map(Generation::getOutput)
.map(AssistantMessage::getContent)
.collect(Collectors.joining());
logger.info("Response: {}", content);
assertThat(content).containsAnyOf("30.0", "30");
assertThat(content).containsAnyOf("10.0", "10");
assertThat(content).containsAnyOf("15.0", "15");
assertThat(content).containsAnyOf("30.0", "30");
assertThat(content).containsAnyOf("10.0", "10");
assertThat(content).containsAnyOf("15.0", "15");
});
});
}
@Configuration

View File

@@ -55,7 +55,7 @@ public class FunctionCallWithFunctionWrapperIT {
void functionCallTest() {
contextRunner
.withPropertyValues("spring.ai.vertex.ai.gemini.chat.options.model="
+ VertexAiGeminiChatModel.ChatModel.GEMINI_PRO.getValue())
+ VertexAiGeminiChatModel.ChatModel.GEMINI_PRO_1_5_FLASH.getValue())
.run(context -> {
VertexAiGeminiChatModel chatModel = context.getBean(VertexAiGeminiChatModel.class);
@@ -65,7 +65,8 @@ public class FunctionCallWithFunctionWrapperIT {
Answer for all listed locations.
If the information was not fetched call the function again. Repeat at most 3 times.
""");
var userMessage = new UserMessage("What's the weather like in San Francisco, Paris and in Tokyo?");
var userMessage = new UserMessage(
"What's the weather like in San Francisco, Paris and in Tokyo? Perform multiple funciton execution if necessary. Return the temperature in Celsius.");
ChatResponse response = chatModel.call(new Prompt(List.of(systemMessage, userMessage),
VertexAiGeminiChatOptions.builder().withFunction("WeatherInfo").build()));

View File

@@ -51,7 +51,7 @@ public class FunctionCallWithPromptFunctionIT {
void functionCallTest() {
contextRunner
.withPropertyValues("spring.ai.vertex.ai.gemini.chat.options.model="
+ VertexAiGeminiChatModel.ChatModel.GEMINI_PRO.getValue())
+ VertexAiGeminiChatModel.ChatModel.GEMINI_PRO_1_5_FLASH.getValue())
.run(context -> {
VertexAiGeminiChatModel chatModel = context.getBean(VertexAiGeminiChatModel.class);
@@ -62,7 +62,7 @@ public class FunctionCallWithPromptFunctionIT {
If the information was not fetched call the function again. Repeat at most 3 times.
""");
UserMessage userMessage = new UserMessage(
"What's the weather like in San Francisco, in Paris and in Tokyo?");
"What's the weather like in San Francisco, in Paris and in Tokyo? Perform multiple funciton execution if necessary. Return the temperature in Celsius.");
var promptOptions = VertexAiGeminiChatOptions.builder()
.withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService())

View File

@@ -15,22 +15,20 @@
*/
package org.springframework.ai.autoconfigure.vertexai.palm2;
import static org.assertj.core.api.Assertions.assertThat;
import java.util.List;
import org.apache.commons.logging.Log;
import org.apache.commons.logging.LogFactory;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.embedding.EmbeddingResponse;
import org.springframework.ai.vertexai.palm2.VertexAiPaLm2ChatModel;
import org.springframework.ai.vertexai.palm2.VertexAiPaLm2EmbeddingModel;
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;
// NOTE: works only with US location. Use VPN if you are outside US.
@EnabledIfEnvironmentVariable(named = "PALM_API_KEY", matches = ".*")
public class VertexAiPaLm2AutoConfigurationIT {
@@ -42,8 +40,7 @@ public class VertexAiPaLm2AutoConfigurationIT {
"spring.ai.vertex.ai.apiKey=" + System.getenv("PALM_API_KEY"),
"spring.ai.vertex.ai.chat.model=chat-bison-001", "spring.ai.vertex.ai.chat.options.temperature=0.8",
"spring.ai.vertex.ai.embedding.model=embedding-gecko-001")
.withConfiguration(
AutoConfigurations.of(RestClientAutoConfiguration.class, VertexAiPalm2AutoConfiguration.class));
.withConfiguration(AutoConfigurations.of(VertexAiPalm2AutoConfiguration.class));
@Test
void generate() {