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:
@@ -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) {
|
||||
|
||||
@@ -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));
|
||||
|
||||
|
||||
@@ -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;
|
||||
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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;
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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";
|
||||
|
||||
|
||||
@@ -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));
|
||||
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
}
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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}.
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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 -> {
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()));
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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() {
|
||||
|
||||
Reference in New Issue
Block a user