diff --git a/README.md b/README.md index 91f625092..1800788c8 100644 --- a/README.md +++ b/README.md @@ -8,9 +8,147 @@ Let's make your `@Beans` intelligent! For further information go to our [Spring AI reference documentation](https://docs.spring.io/spring-ai/reference/). -### Breaking changes -(15.05.2024) -On our march to release 1.0 M1 we have made several breaking changes. Apologies, it is for the best! +## Breaking changes + +On our march to release 1.0.0 M1 we have made several breaking changes. Apologies, it is for the best! + +**(22.05.2024)** + +A major change was made that took the 'old' `ChatClient` and moved the functionality into `ChatModel`. The 'new' `ChatClient` now takes an instance of `ChatModel`. This was done do support a fluent API for creating and executing prompts in a style similar to other client classes in the Spring ecosystem, such as `RestClient`, `WebClient`, and `JdbcClient`. Refer to the [JavaDoc](https://docs.spring.io/spring-ai/docs/1.0.0-SNAPSHOT/api/) for more information on the Fluent API, proper reference documentation is coming shortly. + +We renamed the 'old' `ModelClient` to `Model` and renamed implementing classes, for example `ImageClient` was renamed to `ImageModel`. The `Model` implementation represent the portability layer that converts between the Spring AI API and the underlying AI Model API. + +### Adapting to the changes + +#### Approach 1 + +Now, instead of getting an Autoconfigured `ChatClient` instance, you will get a `ChatModel` instance. The `call` method signatures after renaming remain the same. +To adapt your code should refactor you code to change use of the type `ChatClient` to `ChatModel` +Here is an example of existing code before the change + +```java +@RestController +public class OldSimpleAiController { + + private final ChatClient chatClient; + + @Autowired + public OldSimpleAiController(ChatClient chatClient) { + this.chatClient = chatClient; + } + + @GetMapping("/ai/simple") + public Map completion(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { + return Map.of("generation", chatClient.call(message)); + } +} +``` + +Now after the changes this will be + +```java +@RestController +public class SimpleAiController { + + private final ChatModel chatModel; + + @Autowired + public SimpleAiController(ChatModel chatModel) { + this.chatModel = chatModel; + } + + @GetMapping("/ai/simple") + public Map completion(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { + return Map.of("generation", chatModel.call(message)); + } +} +``` + +NOTE: The renaming also applies to the classes +* `StreamingChatClient` -> `StreamingChatModel` +* `EmbeddingClient` -> `EmbeddingModel` +* `ImageClient` -> `ImageModel` +* `SpeechClient` -> `SpeechModel` +* and similar for other `Client` classes + +#### Approach 2 + +In this approach you will use the new fluent API available on the 'new' `ChatClient` + +Here is an example of existing code before the change + +```java +@RestController +public class OldSimpleAiController { + + private final ChatClient chatClient; + + @Autowired + public OldSimpleAiController(ChatClient chatClient) { + this.chatClient = chatClient; + } + + @GetMapping("/ai/simple") + public Map completion(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { + return Map.of( + "generation", + chatClient.call(message) + ); + } +} +``` + + +Now after the changes this will be + +```java +@RestController +public class SimpleAiController { + + private final ChatClient chatClient; + + @Autowired + public SimpleAiController(ChatClient chatClient) { + this.chatClient = chatClient; + } + + @GetMapping("/ai/simple") + public Map completion(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { + return Map.of( + "generation", + chatClient.prompt().user(message).call().content() + ); + } +} +``` + +and in your `@Configuration` class you need to define an `@Bean` as shown below + +```java +@Configuration +public class ApplicationConfiguration { + + @Bean + ChatClient chatClient(ChatModel chatModel) { + return ChatClient.builder(chatModel).build(); + } +} +``` + +NOTE: The `ChatModel` instance is made available to you through autoconfiguration. + +#### Approach 3 + +There is a tag in the GitHub repository called [v1.0.0-SNAPSHOT-before-chatclient-changes](https://github.com/spring-projects/spring-ai/tree/v1.0.0-SNAPSHOT-before-chatclient-changes) that you can checkout and do a local build to avoid updating any of your code until you are ready to migrate your code base. + +```bash +git checkout tags/v1.0.0-SNAPSHOT-before-chatclient-changes + +./mvnw clean install -DskipTests +``` + + +**(15.05.2024)** Renamed POM artifact names: - spring-ai-qdrant -> spring-ai-qdrant-store diff --git a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatClient.java b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java similarity index 94% rename from models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatClient.java rename to models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java index 0f9bcbf44..c8802cbaa 100644 --- a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatClient.java +++ b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java @@ -26,6 +26,7 @@ import java.util.stream.Collectors; import org.slf4j.Logger; import org.slf4j.LoggerFactory; +import org.springframework.ai.chat.ChatModel; import reactor.core.publisher.Flux; import org.springframework.ai.anthropic.api.AnthropicApi; @@ -38,10 +39,9 @@ import org.springframework.ai.anthropic.api.AnthropicApi.Role; import org.springframework.ai.anthropic.api.AnthropicApi.StreamResponse; import org.springframework.ai.anthropic.api.AnthropicApi.Usage; import org.springframework.ai.anthropic.metadata.AnthropicChatResponseMetadata; -import org.springframework.ai.chat.ChatClient; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; import org.springframework.ai.chat.prompt.ChatOptions; @@ -56,16 +56,16 @@ import org.springframework.util.Assert; import org.springframework.util.CollectionUtils; /** - * The {@link ChatClient} implementation for the Anthropic service. + * The {@link ChatModel} implementation for the Anthropic service. * * @author Christian Tzolov * @since 1.0.0 */ -public class AnthropicChatClient extends +public class AnthropicChatModel extends AbstractFunctionCallSupport> - implements ChatClient, StreamingChatClient { + implements ChatModel, StreamingChatModel { - private static final Logger logger = LoggerFactory.getLogger(AnthropicChatClient.class); + private static final Logger logger = LoggerFactory.getLogger(AnthropicChatModel.class); public static final String DEFAULT_MODEL_NAME = AnthropicApi.ChatModel.CLAUDE_3_OPUS.getValue(); @@ -89,10 +89,10 @@ public class AnthropicChatClient extends public final RetryTemplate retryTemplate; /** - * Construct a new {@link AnthropicChatClient} instance. + * Construct a new {@link AnthropicChatModel} instance. * @param anthropicApi the lower-level API for the Anthropic service. */ - public AnthropicChatClient(AnthropicApi anthropicApi) { + public AnthropicChatModel(AnthropicApi anthropicApi) { this(anthropicApi, AnthropicChatOptions.builder() .withModel(DEFAULT_MODEL_NAME) @@ -102,34 +102,34 @@ public class AnthropicChatClient extends } /** - * Construct a new {@link AnthropicChatClient} instance. + * Construct a new {@link AnthropicChatModel} instance. * @param anthropicApi the lower-level API for the Anthropic service. * @param defaultOptions the default options used for the chat completion requests. */ - public AnthropicChatClient(AnthropicApi anthropicApi, AnthropicChatOptions defaultOptions) { + public AnthropicChatModel(AnthropicApi anthropicApi, AnthropicChatOptions defaultOptions) { this(anthropicApi, defaultOptions, RetryUtils.DEFAULT_RETRY_TEMPLATE); } /** - * Construct a new {@link AnthropicChatClient} instance. + * Construct a new {@link AnthropicChatModel} instance. * @param anthropicApi the lower-level API for the Anthropic service. * @param defaultOptions the default options used for the chat completion requests. * @param retryTemplate the retry template used to retry the Anthropic API calls. */ - public AnthropicChatClient(AnthropicApi anthropicApi, AnthropicChatOptions defaultOptions, + public AnthropicChatModel(AnthropicApi anthropicApi, AnthropicChatOptions defaultOptions, RetryTemplate retryTemplate) { this(anthropicApi, defaultOptions, retryTemplate, null); } /** - * Construct a new {@link AnthropicChatClient} instance. + * Construct a new {@link AnthropicChatModel} instance. * @param anthropicApi the lower-level API for the Anthropic service. * @param defaultOptions the default options used for the chat completion requests. * @param retryTemplate the retry template used to retry the Anthropic API calls. * @param functionCallbackContext the function callback context used to store the * state of the function calls. */ - public AnthropicChatClient(AnthropicApi anthropicApi, AnthropicChatOptions defaultOptions, + public AnthropicChatModel(AnthropicApi anthropicApi, AnthropicChatOptions defaultOptions, RetryTemplate retryTemplate, FunctionCallbackContext functionCallbackContext) { super(functionCallbackContext); @@ -457,4 +457,9 @@ public class AnthropicChatClient extends "Streaming (stream=true) is not yet supported. We plan to add streaming support in a future beta version."); } + @Override + public ChatOptions getDefaultOptions() { + return AnthropicChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatOptions.java b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatOptions.java index 7cd7bdb23..6d13f6bae 100644 --- a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatOptions.java +++ b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatOptions.java @@ -51,11 +51,11 @@ public class AnthropicChatOptions implements ChatOptions, FunctionCallingOptions private @JsonProperty("top_k") Integer topK; /** - * Tool Function Callbacks to register with the ChatClient. For Prompt + * Tool Function Callbacks to register with the ChatModel. For Prompt * Options the functionCallbacks are automatically enabled for the duration of the * prompt execution. For Default Options the functionCallbacks are registered but * disabled by default. Use the enableFunctions to set the functions from the registry - * to be used by the ChatClient chat completion requests. + * to be used by the ChatModel chat completion requests. */ @NestedConfigurationProperty @JsonIgnore @@ -223,4 +223,17 @@ public class AnthropicChatOptions implements ChatOptions, FunctionCallingOptions this.functions = functions; } + public static AnthropicChatOptions fromOptions(AnthropicChatOptions fromOptions) { + return builder().withModel(fromOptions.getModel()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withMetadata(fromOptions.getMetadata()) + .withStopSequences(fromOptions.getStopSequences()) + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withTopK(fromOptions.getTopK()) + .withFunctionCallbacks(fromOptions.getFunctionCallbacks()) + .withFunctions(fromOptions.getFunctions()) + .build(); + } + } diff --git a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/api/AnthropicApi.java b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/api/AnthropicApi.java index eb5a96289..76c969a44 100644 --- a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/api/AnthropicApi.java +++ b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/api/AnthropicApi.java @@ -26,6 +26,7 @@ import com.fasterxml.jackson.annotation.JsonProperty; import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; +import org.springframework.ai.model.ModelDescription; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.retry.RetryUtils; import org.springframework.http.HttpHeaders; @@ -116,7 +117,7 @@ public class AnthropicApi { * "https://docs.anthropic.com/claude/docs/models-overview#model-comparison">model * comparison for additional details and options. */ - public enum ChatModel { + public enum ChatModel implements ModelDescription { // @formatter:off CLAUDE_3_OPUS("claude-3-opus-20240229"), @@ -140,6 +141,11 @@ public class AnthropicApi { return this.value; } + @Override + public String getModelName() { + return this.value; + } + } /** diff --git a/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatClientIT.java b/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatModelIT.java similarity index 91% rename from models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatClientIT.java rename to models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatModelIT.java index ce5b45d37..e2eed692c 100644 --- a/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatClientIT.java +++ b/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatModelIT.java @@ -29,10 +29,10 @@ import org.slf4j.LoggerFactory; import org.springframework.ai.anthropic.api.AnthropicApi; import org.springframework.ai.anthropic.api.tool.MockWeatherService; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.Media; import org.springframework.ai.chat.messages.Message; @@ -56,15 +56,15 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest(classes = AnthropicTestConfiguration.class, properties = "spring.ai.retry.on-http-codes=429") @EnabledIfEnvironmentVariable(named = "ANTHROPIC_API_KEY", matches = ".+") -class AnthropicChatClientIT { +class AnthropicChatModelIT { - private static final Logger logger = LoggerFactory.getLogger(AnthropicChatClientIT.class); + private static final Logger logger = LoggerFactory.getLogger(AnthropicChatModelIT.class); @Autowired - protected ChatClient chatClient; + protected ChatModel chatModel; @Autowired - protected StreamingChatClient streamingChatClient; + protected StreamingChatModel streamingChatModel; @Value("classpath:/prompts/system-message.st") private Resource systemResource; @@ -76,7 +76,7 @@ class AnthropicChatClientIT { SystemPromptTemplate systemPromptTemplate = new SystemPromptTemplate(systemResource); Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", "Bob", "voice", "pirate")); Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); - ChatResponse response = chatClient.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResults()).hasSize(1); assertThat(response.getMetadata().getUsage().getGenerationTokens()).isGreaterThan(0); assertThat(response.getMetadata().getUsage().getPromptTokens()).isGreaterThan(0); @@ -102,7 +102,7 @@ class AnthropicChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "ice cream flavors", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = this.chatClient.call(prompt).getResult(); + Generation generation = this.chatModel.call(prompt).getResult(); List list = listOutputConverter.convert(generation.getOutput().getContent()); assertThat(list).hasSize(5); @@ -120,7 +120,7 @@ class AnthropicChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = chatClient.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); Map result = mapOutputConverter.convert(generation.getOutput().getContent()); assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); @@ -142,7 +142,7 @@ class AnthropicChatClientIT { """; PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = chatClient.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); ActorsFilmsRecord actorsFilms = beanOutputConverter.convert(generation.getOutput().getContent()); logger.info("" + actorsFilms); @@ -163,7 +163,7 @@ class AnthropicChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - String generationTextFromStream = streamingChatClient.stream(prompt) + String generationTextFromStream = streamingChatModel.stream(prompt) .collectList() .block() .stream() @@ -187,7 +187,7 @@ class AnthropicChatClientIT { var userMessage = new UserMessage("Explain what do you see on this picture?", List.of(new Media(MimeTypeUtils.IMAGE_PNG, imageData))); - var response = chatClient.call(new Prompt(List.of(userMessage))); + var response = chatModel.call(new Prompt(List.of(userMessage))); logger.info(response.getResult().getOutput().getContent()); assertThat(response.getResult().getOutput().getContent()).contains("bananas", "apple", "basket"); @@ -209,7 +209,7 @@ class AnthropicChatClientIT { .build())) .build(); - ChatResponse response = chatClient.call(new Prompt(messages, promptOptions)); + ChatResponse response = chatModel.call(new Prompt(messages, promptOptions)); logger.info("Response: {}", response); diff --git a/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicTestConfiguration.java b/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicTestConfiguration.java index 649fbef1a..3f4551321 100644 --- a/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicTestConfiguration.java +++ b/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicTestConfiguration.java @@ -38,9 +38,9 @@ public class AnthropicTestConfiguration { } @Bean - public AnthropicChatClient openAiChatClient(AnthropicApi api) { - AnthropicChatClient anthropicChatClient = new AnthropicChatClient(api); - return anthropicChatClient; + public AnthropicChatModel openAiChatModel(AnthropicApi api) { + AnthropicChatModel anthropicChatModel = new AnthropicChatModel(api); + return anthropicChatModel; } } diff --git a/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/ChatCompletionRequestTests.java b/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/ChatCompletionRequestTests.java index bc777a846..b5c47abd0 100644 --- a/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/ChatCompletionRequestTests.java +++ b/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/ChatCompletionRequestTests.java @@ -30,7 +30,7 @@ public class ChatCompletionRequestTests { @Test public void createRequestWithChatOptions() { - var client = new AnthropicChatClient(new AnthropicApi("TEST"), + var client = new AnthropicChatModel(new AnthropicApi("TEST"), AnthropicChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6f).build()); var request = client.createRequest(new Prompt("Test message content"), false); diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatClient.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java similarity index 96% rename from models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatClient.java rename to models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java index a49a42ff5..712f53567 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatClient.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java @@ -38,10 +38,10 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.azure.openai.metadata.AzureOpenAiChatResponseMetadata; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; import org.springframework.ai.chat.metadata.PromptMetadata; @@ -63,7 +63,7 @@ import java.util.Set; import java.util.concurrent.atomic.AtomicBoolean; /** - * {@link ChatClient} implementation for {@literal Microsoft Azure AI} backed by + * {@link ChatModel} implementation for {@literal Microsoft Azure AI} backed by * {@link OpenAIClient}. * * @author Mark Pollack @@ -71,12 +71,12 @@ import java.util.concurrent.atomic.AtomicBoolean; * @author John Blum * @author Christian Tzolov * @author Grogdunn - * @see ChatClient + * @see ChatModel * @see com.azure.ai.openai.OpenAIClient */ -public class AzureOpenAiChatClient +public class AzureOpenAiChatModel extends AbstractFunctionCallSupport - implements ChatClient, StreamingChatClient { + implements ChatModel, StreamingChatModel { private static final String DEFAULT_DEPLOYMENT_NAME = "gpt-35-turbo"; @@ -94,7 +94,7 @@ public class AzureOpenAiChatClient */ private final OpenAIClient openAIClient; - public AzureOpenAiChatClient(OpenAIClient microsoftOpenAiClient) { + public AzureOpenAiChatModel(OpenAIClient microsoftOpenAiClient) { this(microsoftOpenAiClient, AzureOpenAiChatOptions.builder() .withDeploymentName(DEFAULT_DEPLOYMENT_NAME) @@ -102,11 +102,11 @@ public class AzureOpenAiChatClient .build()); } - public AzureOpenAiChatClient(OpenAIClient microsoftOpenAiClient, AzureOpenAiChatOptions options) { + public AzureOpenAiChatModel(OpenAIClient microsoftOpenAiClient, AzureOpenAiChatOptions options) { this(microsoftOpenAiClient, options, null); } - public AzureOpenAiChatClient(OpenAIClient microsoftOpenAiClient, AzureOpenAiChatOptions options, + public AzureOpenAiChatModel(OpenAIClient microsoftOpenAiClient, AzureOpenAiChatOptions options, FunctionCallbackContext functionCallbackContext) { super(functionCallbackContext); Assert.notNull(microsoftOpenAiClient, "com.azure.ai.openai.OpenAIClient must not be null"); @@ -117,17 +117,17 @@ public class AzureOpenAiChatClient /** * @deprecated since 0.8.0, use - * {@link #AzureOpenAiChatClient(OpenAIClient, AzureOpenAiChatOptions)} instead. + * {@link #AzureOpenAiChatModel(OpenAIClient, AzureOpenAiChatOptions)} instead. */ @Deprecated(forRemoval = true, since = "0.8.0") - public AzureOpenAiChatClient withDefaultOptions(AzureOpenAiChatOptions defaultOptions) { + public AzureOpenAiChatModel withDefaultOptions(AzureOpenAiChatOptions defaultOptions) { Assert.notNull(defaultOptions, "DefaultOptions must not be null"); this.defaultOptions = defaultOptions; return this; } public AzureOpenAiChatOptions getDefaultOptions() { - return this.defaultOptions; + return AzureOpenAiChatOptions.fromOptions(this.defaultOptions); } @Override diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatOptions.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatOptions.java index 681285ecc..6e2d4f5eb 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatOptions.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatOptions.java @@ -127,11 +127,11 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio private String deploymentName; /** - * OpenAI Tool Function Callbacks to register with the ChatClient. For Prompt Options + * OpenAI Tool Function Callbacks to register with the ChatModel. For Prompt Options * the functionCallbacks are automatically enabled for the duration of the prompt * execution. For Default Options the functionCallbacks are registered but disabled by * default. Use the enableFunctions to set the functions from the registry to be used - * by the ChatClient chat completion requests. + * by the ChatModel chat completion requests. */ @NestedConfigurationProperty @JsonIgnore @@ -356,4 +356,22 @@ public class AzureOpenAiChatOptions implements FunctionCallingOptions, ChatOptio this.functions = functions; } + public static AzureOpenAiChatOptions fromOptions(AzureOpenAiChatOptions fromOptions) { + return builder().withDeploymentName(fromOptions.getDeploymentName()) + .withFrequencyPenalty( + fromOptions.getFrequencyPenalty() != null ? fromOptions.getFrequencyPenalty().floatValue() : null) + .withLogitBias(fromOptions.getLogitBias()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withN(fromOptions.getN()) + .withPresencePenalty( + fromOptions.getPresencePenalty() != null ? fromOptions.getPresencePenalty().floatValue() : null) + .withStop(fromOptions.getStop()) + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withUser(fromOptions.getUser()) + .withFunctionCallbacks(fromOptions.getFunctionCallbacks()) + .withFunctions(fromOptions.getFunctions()) + .build(); + } + } diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingClient.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModel.java similarity index 91% rename from models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingClient.java rename to models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModel.java index 66add3a9e..e9679bc39 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingClient.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModel.java @@ -28,7 +28,7 @@ import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; import org.springframework.ai.document.MetadataMode; -import org.springframework.ai.embedding.AbstractEmbeddingClient; +import org.springframework.ai.embedding.AbstractEmbeddingModel; import org.springframework.ai.embedding.Embedding; import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingRequest; @@ -37,9 +37,9 @@ import org.springframework.ai.embedding.EmbeddingResponseMetadata; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.util.Assert; -public class AzureOpenAiEmbeddingClient extends AbstractEmbeddingClient { +public class AzureOpenAiEmbeddingModel extends AbstractEmbeddingModel { - private static final Logger logger = LoggerFactory.getLogger(AzureOpenAiEmbeddingClient.class); + private static final Logger logger = LoggerFactory.getLogger(AzureOpenAiEmbeddingModel.class); private final OpenAIClient azureOpenAiClient; @@ -47,16 +47,16 @@ public class AzureOpenAiEmbeddingClient extends AbstractEmbeddingClient { private final MetadataMode metadataMode; - public AzureOpenAiEmbeddingClient(OpenAIClient azureOpenAiClient) { + public AzureOpenAiEmbeddingModel(OpenAIClient azureOpenAiClient) { this(azureOpenAiClient, MetadataMode.EMBED); } - public AzureOpenAiEmbeddingClient(OpenAIClient azureOpenAiClient, MetadataMode metadataMode) { + public AzureOpenAiEmbeddingModel(OpenAIClient azureOpenAiClient, MetadataMode metadataMode) { this(azureOpenAiClient, metadataMode, AzureOpenAiEmbeddingOptions.builder().withDeploymentName("text-embedding-ada-002").build()); } - public AzureOpenAiEmbeddingClient(OpenAIClient azureOpenAiClient, MetadataMode metadataMode, + public AzureOpenAiEmbeddingModel(OpenAIClient azureOpenAiClient, MetadataMode metadataMode, AzureOpenAiEmbeddingOptions options) { Assert.notNull(azureOpenAiClient, "com.azure.ai.openai.OpenAIClient must not be null"); Assert.notNull(metadataMode, "Metadata mode must not be null"); diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureChatCompletionsOptionsTests.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureChatCompletionsOptionsTests.java index 1e0ba2939..1bbcc5e8f 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureChatCompletionsOptionsTests.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureChatCompletionsOptionsTests.java @@ -53,7 +53,7 @@ public class AzureChatCompletionsOptionsTests { .withUser("user") .build(); - var client = new AzureOpenAiChatClient(mockClient, defaultOptions); + var client = new AzureOpenAiChatModel(mockClient, defaultOptions); var requestOptions = client.toAzureChatCompletionsOptions(new Prompt("Test message content")); diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureEmbeddingsOptionsTests.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureEmbeddingsOptionsTests.java index 3a5e194d9..18fe0e56a 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureEmbeddingsOptionsTests.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureEmbeddingsOptionsTests.java @@ -36,7 +36,7 @@ public class AzureEmbeddingsOptionsTests { public void createRequestWithChatOptions() { OpenAIClient mockClient = Mockito.mock(OpenAIClient.class); - var client = new AzureOpenAiEmbeddingClient(mockClient, MetadataMode.EMBED, + var client = new AzureOpenAiEmbeddingModel(mockClient, MetadataMode.EMBED, AzureOpenAiEmbeddingOptions.builder() .withDeploymentName("DEFAULT_MODEL") .withUser("USER_TEST") diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatClientIT.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatModelIT.java similarity index 91% rename from models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatClientIT.java rename to models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatModelIT.java index 253e13d57..f991e334f 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatClientIT.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatModelIT.java @@ -46,13 +46,13 @@ import org.springframework.core.convert.support.DefaultConversionService; import static org.assertj.core.api.Assertions.assertThat; -@SpringBootTest(classes = AzureOpenAiChatClientIT.TestConfiguration.class) +@SpringBootTest(classes = AzureOpenAiChatModelIT.TestConfiguration.class) @EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_API_KEY", matches = ".+") @EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_ENDPOINT", matches = ".+") -class AzureOpenAiChatClientIT { +class AzureOpenAiChatModelIT { @Autowired - private AzureOpenAiChatClient chatClient; + private AzureOpenAiChatModel chatModel; record ActorsFilms(String actor, List movies) { } @@ -69,7 +69,7 @@ class AzureOpenAiChatClientIT { UserMessage userMessage = new UserMessage("Generate the names of 5 famous pirates."); Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); - ChatResponse response = chatClient.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); } @@ -86,7 +86,7 @@ class AzureOpenAiChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "ice cream flavors", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = chatClient.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); List list = outputConverter.convert(generation.getOutput().getContent()); assertThat(list).hasSize(5); @@ -105,7 +105,7 @@ class AzureOpenAiChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = chatClient.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); Map result = outputConverter.convert(generation.getOutput().getContent()); assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); @@ -124,7 +124,7 @@ class AzureOpenAiChatClientIT { """; PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = chatClient.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); ActorsFilms actorsFilms = outputConverter.convert(generation.getOutput().getContent()); assertThat(actorsFilms.actor()).isNotNull(); @@ -145,7 +145,7 @@ class AzureOpenAiChatClientIT { """; PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = chatClient.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); ActorsFilmsRecord actorsFilms = outputConverter.convert(generation.getOutput().getContent()); System.out.println(actorsFilms); @@ -166,7 +166,7 @@ class AzureOpenAiChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - String generationTextFromStream = chatClient.stream(prompt) + String generationTextFromStream = chatModel.stream(prompt) .collectList() .block() .stream() @@ -194,8 +194,8 @@ class AzureOpenAiChatClientIT { } @Bean - public AzureOpenAiChatClient azureOpenAiChatClient(OpenAIClient openAIClient) { - return new AzureOpenAiChatClient(openAIClient, + public AzureOpenAiChatModel azureOpenAiChatModel(OpenAIClient openAIClient) { + return new AzureOpenAiChatModel(openAIClient, AzureOpenAiChatOptions.builder().withDeploymentName("gpt-35-turbo").withMaxTokens(200).build()); } diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingClientIT.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModelIT.java similarity index 79% rename from models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingClientIT.java rename to models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModelIT.java index b973aa19b..7c1710e06 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingClientIT.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiEmbeddingModelIT.java @@ -34,25 +34,25 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_API_KEY", matches = ".+") @EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_ENDPOINT", matches = ".+") -class AzureOpenAiEmbeddingClientIT { +class AzureOpenAiEmbeddingModelIT { @Autowired - private AzureOpenAiEmbeddingClient embeddingClient; + private AzureOpenAiEmbeddingModel embeddingModel; @Test void singleEmbedding() { - assertThat(embeddingClient).isNotNull(); - EmbeddingResponse embeddingResponse = embeddingClient.embedForResponse(List.of("Hello World")); + assertThat(embeddingModel).isNotNull(); + EmbeddingResponse embeddingResponse = embeddingModel.embedForResponse(List.of("Hello World")); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); - System.out.println(embeddingClient.dimensions()); - assertThat(embeddingClient.dimensions()).isEqualTo(1536); + System.out.println(embeddingModel.dimensions()); + assertThat(embeddingModel.dimensions()).isEqualTo(1536); } @Test void batchEmbedding() { - assertThat(embeddingClient).isNotNull(); - EmbeddingResponse embeddingResponse = embeddingClient + assertThat(embeddingModel).isNotNull(); + EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); assertThat(embeddingResponse.getResults()).hasSize(2); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); @@ -60,7 +60,7 @@ class AzureOpenAiEmbeddingClientIT { assertThat(embeddingResponse.getResults().get(1).getOutput()).isNotEmpty(); assertThat(embeddingResponse.getResults().get(1).getIndex()).isEqualTo(1); - assertThat(embeddingClient.dimensions()).isEqualTo(1536); + assertThat(embeddingModel.dimensions()).isEqualTo(1536); } @SpringBootConfiguration @@ -74,8 +74,8 @@ class AzureOpenAiEmbeddingClientIT { } @Bean - public AzureOpenAiEmbeddingClient azureEmbeddingClient(OpenAIClient openAIClient) { - return new AzureOpenAiEmbeddingClient(openAIClient); + public AzureOpenAiEmbeddingModel azureEmbeddingModel(OpenAIClient openAIClient) { + return new AzureOpenAiEmbeddingModel(openAIClient); } } diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/MockAzureOpenAiTestConfiguration.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/MockAzureOpenAiTestConfiguration.java index f5e7b438a..6ae824bad 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/MockAzureOpenAiTestConfiguration.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/MockAzureOpenAiTestConfiguration.java @@ -59,8 +59,8 @@ public class MockAzureOpenAiTestConfiguration { } @Bean - AzureOpenAiChatClient azureOpenAiChatClient(OpenAIClient microsoftAzureOpenAiClient) { - return new AzureOpenAiChatClient(microsoftAzureOpenAiClient); + AzureOpenAiChatModel azureOpenAiChatModel(OpenAIClient microsoftAzureOpenAiClient) { + return new AzureOpenAiChatModel(microsoftAzureOpenAiClient); } } diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/AzureOpenAiChatClientFunctionCallIT.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/AzureOpenAiChatModelFunctionCallIT.java similarity index 89% rename from models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/AzureOpenAiChatClientFunctionCallIT.java rename to models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/AzureOpenAiChatModelFunctionCallIT.java index 08c81ebd1..736c69041 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/AzureOpenAiChatClientFunctionCallIT.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/function/AzureOpenAiChatModelFunctionCallIT.java @@ -29,7 +29,7 @@ import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.azure.openai.AzureOpenAiChatClient; +import org.springframework.ai.azure.openai.AzureOpenAiChatModel; import org.springframework.ai.azure.openai.AzureOpenAiChatOptions; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; @@ -46,18 +46,18 @@ import reactor.core.publisher.Flux; import static org.assertj.core.api.Assertions.assertThat; -@SpringBootTest(classes = AzureOpenAiChatClientFunctionCallIT.TestConfiguration.class) +@SpringBootTest(classes = AzureOpenAiChatModelFunctionCallIT.TestConfiguration.class) @EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_API_KEY", matches = ".+") @EnabledIfEnvironmentVariable(named = "AZURE_OPENAI_ENDPOINT", matches = ".+") -class AzureOpenAiChatClientFunctionCallIT { +class AzureOpenAiChatModelFunctionCallIT { - private static final Logger logger = LoggerFactory.getLogger(AzureOpenAiChatClientFunctionCallIT.class); + private static final Logger logger = LoggerFactory.getLogger(AzureOpenAiChatModelFunctionCallIT.class); @Autowired private String selectedModel; @Autowired - private AzureOpenAiChatClient chatClient; + private AzureOpenAiChatModel chatModel; @Test void functionCallTest() { @@ -75,7 +75,7 @@ class AzureOpenAiChatClientFunctionCallIT { .build())) .build(); - ChatResponse response = chatClient.call(new Prompt(messages, promptOptions)); + ChatResponse response = chatModel.call(new Prompt(messages, promptOptions)); logger.info("Response: {}", response); @@ -99,7 +99,7 @@ class AzureOpenAiChatClientFunctionCallIT { .build())) .build(); - Flux response = chatClient.stream(new Prompt(messages, promptOptions)); + Flux response = chatModel.stream(new Prompt(messages, promptOptions)); final var counter = new AtomicInteger(); String content = response.doOnEach(listSignal -> counter.getAndIncrement()) @@ -129,8 +129,8 @@ class AzureOpenAiChatClientFunctionCallIT { } @Bean - public AzureOpenAiChatClient azureOpenAiChatClient(OpenAIClient openAIClient, String selectedModel) { - return new AzureOpenAiChatClient(openAIClient, + public AzureOpenAiChatModel azureOpenAiChatModel(OpenAIClient openAIClient, String selectedModel) { + return new AzureOpenAiChatModel(openAIClient, AzureOpenAiChatOptions.builder().withDeploymentName(selectedModel).withMaxTokens(500).build()); } diff --git a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiChatClientMetadataTests.java b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiChatModelMetadataTests.java similarity index 96% rename from models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiChatClientMetadataTests.java rename to models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiChatModelMetadataTests.java index cc938c9b9..3397f84a3 100644 --- a/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiChatClientMetadataTests.java +++ b/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/metadata/AzureOpenAiChatModelMetadataTests.java @@ -23,7 +23,7 @@ import com.azure.ai.openai.models.ContentFilterResultsForChoice; import com.azure.ai.openai.models.ContentFilterSeverity; import org.junit.jupiter.api.Test; -import org.springframework.ai.azure.openai.AzureOpenAiChatClient; +import org.springframework.ai.azure.openai.AzureOpenAiChatModel; import org.springframework.ai.azure.openai.MockAzureOpenAiTestConfiguration; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; @@ -55,7 +55,7 @@ import org.springframework.web.context.request.WebRequest; import static org.assertj.core.api.Assertions.assertThat; /** - * Unit Tests for {@link AzureOpenAiChatClient} asserting AI metadata. + * Unit Tests for {@link AzureOpenAiChatModel} asserting AI metadata. * * @author John Blum * @author Christian Tzolov @@ -63,12 +63,12 @@ import static org.assertj.core.api.Assertions.assertThat; */ @SpringBootTest @ActiveProfiles("spring-ai-azure-openai-mocks") -@ContextConfiguration(classes = AzureOpenAiChatClientMetadataTests.TestConfiguration.class) +@ContextConfiguration(classes = AzureOpenAiChatModelMetadataTests.TestConfiguration.class) @SuppressWarnings("unused") -class AzureOpenAiChatClientMetadataTests { +class AzureOpenAiChatModelMetadataTests { @Autowired - private AzureOpenAiChatClient aiClient; + private AzureOpenAiChatModel aiClient; @Test void azureOpenAiMetadataCapturedDuringGeneration() { diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/AnthropicChatOptions.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/AnthropicChatOptions.java index 2daceca21..d0d5a5a2c 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/AnthropicChatOptions.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/AnthropicChatOptions.java @@ -164,4 +164,14 @@ public class AnthropicChatOptions implements ChatOptions { this.anthropicVersion = anthropicVersion; } + public static AnthropicChatOptions fromOptions(AnthropicChatOptions fromOptions) { + return builder().withTemperature(fromOptions.getTemperature()) + .withMaxTokensToSample(fromOptions.getMaxTokensToSample()) + .withTopK(fromOptions.getTopK()) + .withTopP(fromOptions.getTopP()) + .withStopSequences(fromOptions.getStopSequences()) + .withAnthropicVersion(fromOptions.getAnthropicVersion()) + .build(); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatClient.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatModel.java similarity index 87% rename from models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatClient.java rename to models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatModel.java index 26b887344..b5d62da9d 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatClient.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatModel.java @@ -17,7 +17,7 @@ package org.springframework.ai.bedrock.anthropic; import java.util.List; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; @@ -27,25 +27,25 @@ import org.springframework.ai.bedrock.MessageToPromptConverter; import org.springframework.ai.bedrock.anthropic.api.AnthropicChatBedrockApi; import org.springframework.ai.bedrock.anthropic.api.AnthropicChatBedrockApi.AnthropicChatRequest; import org.springframework.ai.bedrock.anthropic.api.AnthropicChatBedrockApi.AnthropicChatResponse; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.Generation; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.ModelOptionsUtils; /** - * Java {@link ChatClient} and {@link StreamingChatClient} for the Bedrock Anthropic chat + * Java {@link ChatModel} and {@link StreamingChatModel} for the Bedrock Anthropic chat * generative. * * @author Christian Tzolov * @since 0.8.0 */ -public class BedrockAnthropicChatClient implements ChatClient, StreamingChatClient { +public class BedrockAnthropicChatModel implements ChatModel, StreamingChatModel { private final AnthropicChatBedrockApi anthropicChatApi; private final AnthropicChatOptions defaultOptions; - public BedrockAnthropicChatClient(AnthropicChatBedrockApi chatApi) { + public BedrockAnthropicChatModel(AnthropicChatBedrockApi chatApi) { this(chatApi, AnthropicChatOptions.builder() .withTemperature(0.8f) @@ -55,7 +55,7 @@ public class BedrockAnthropicChatClient implements ChatClient, StreamingChatClie .build()); } - public BedrockAnthropicChatClient(AnthropicChatBedrockApi chatApi, AnthropicChatOptions options) { + public BedrockAnthropicChatModel(AnthropicChatBedrockApi chatApi, AnthropicChatOptions options) { this.anthropicChatApi = chatApi; this.defaultOptions = options; } @@ -117,4 +117,9 @@ public class BedrockAnthropicChatClient implements ChatClient, StreamingChatClie return request; } + @Override + public ChatOptions getDefaultOptions() { + return AnthropicChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/api/AnthropicChatBedrockApi.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/api/AnthropicChatBedrockApi.java index 55a2d80af..8c800a382 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/api/AnthropicChatBedrockApi.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/api/AnthropicChatBedrockApi.java @@ -29,6 +29,7 @@ import software.amazon.awssdk.regions.Region; import org.springframework.ai.bedrock.anthropic.api.AnthropicChatBedrockApi.AnthropicChatRequest; import org.springframework.ai.bedrock.anthropic.api.AnthropicChatBedrockApi.AnthropicChatResponse; import org.springframework.ai.bedrock.api.AbstractBedrockApi; +import org.springframework.ai.model.ModelDescription; import org.springframework.util.Assert; /** @@ -225,7 +226,7 @@ public class AnthropicChatBedrockApi extends /** * Anthropic models version. */ - public enum AnthropicChatModel { + public enum AnthropicChatModel implements ModelDescription { /** * anthropic.claude-instant-v1 */ @@ -251,6 +252,11 @@ public class AnthropicChatBedrockApi extends AnthropicChatModel(String value) { this.id = value; } + + @Override + public String getModelName() { + return this.id; + } } @Override diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/Anthropic3ChatOptions.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/Anthropic3ChatOptions.java index 2862359fe..b4995683a 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/Anthropic3ChatOptions.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/Anthropic3ChatOptions.java @@ -163,4 +163,14 @@ public class Anthropic3ChatOptions implements ChatOptions { this.anthropicVersion = anthropicVersion; } + public static Anthropic3ChatOptions fromOptions(Anthropic3ChatOptions fromOptions) { + return builder().withTemperature(fromOptions.getTemperature()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withTopK(fromOptions.getTopK()) + .withTopP(fromOptions.getTopP()) + .withStopSequences(fromOptions.getStopSequences()) + .withAnthropicVersion(fromOptions.getAnthropicVersion()) + .build(); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatClient.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatModel.java similarity index 91% rename from models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatClient.java rename to models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatModel.java index 12dba850c..2ab9364e9 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatClient.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatModel.java @@ -15,17 +15,25 @@ */ package org.springframework.ai.bedrock.anthropic3; +import java.util.ArrayList; +import java.util.Base64; +import java.util.List; +import java.util.concurrent.atomic.AtomicReference; +import java.util.stream.Collectors; + +import reactor.core.publisher.Flux; + import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi; import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi.AnthropicChatRequest; import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi.AnthropicChatResponse; import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi.AnthropicChatStreamingResponse.StreamingType; -import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi.MediaContent; import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi.ChatCompletionMessage; import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi.ChatCompletionMessage.Role; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi.MediaContent; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; @@ -34,29 +42,21 @@ import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.util.CollectionUtils; -import reactor.core.publisher.Flux; - -import java.util.ArrayList; -import java.util.Base64; -import java.util.List; -import java.util.concurrent.atomic.AtomicReference; -import java.util.stream.Collectors; - /** - * Java {@link ChatClient} and {@link StreamingChatClient} for the Bedrock Anthropic chat + * Java {@link ChatModel} and {@link StreamingChatModel} for the Bedrock Anthropic chat * generative. * * @author Ben Middleton * @author Christian Tzolov * @since 1.0.0 */ -public class BedrockAnthropic3ChatClient implements ChatClient, StreamingChatClient { +public class BedrockAnthropic3ChatModel implements ChatModel, StreamingChatModel { private final Anthropic3ChatBedrockApi anthropicChatApi; private final Anthropic3ChatOptions defaultOptions; - public BedrockAnthropic3ChatClient(Anthropic3ChatBedrockApi chatApi) { + public BedrockAnthropic3ChatModel(Anthropic3ChatBedrockApi chatApi) { this(chatApi, Anthropic3ChatOptions.builder() .withTemperature(0.8f) @@ -66,7 +66,7 @@ public class BedrockAnthropic3ChatClient implements ChatClient, StreamingChatCli .build()); } - public BedrockAnthropic3ChatClient(Anthropic3ChatBedrockApi chatApi, Anthropic3ChatOptions options) { + public BedrockAnthropic3ChatModel(Anthropic3ChatBedrockApi chatApi, Anthropic3ChatOptions options) { this.anthropicChatApi = chatApi; this.defaultOptions = options; } @@ -187,4 +187,9 @@ public class BedrockAnthropic3ChatClient implements ChatClient, StreamingChatCli } } + @Override + public ChatOptions getDefaultOptions() { + return Anthropic3ChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/api/Anthropic3ChatBedrockApi.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/api/Anthropic3ChatBedrockApi.java index 0148b4983..8b5b29ed1 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/api/Anthropic3ChatBedrockApi.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/api/Anthropic3ChatBedrockApi.java @@ -23,6 +23,7 @@ import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi.An import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi.AnthropicChatResponse; import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi.AnthropicChatStreamingResponse; import org.springframework.ai.bedrock.api.AbstractBedrockApi; +import org.springframework.ai.model.ModelDescription; import org.springframework.util.Assert; import reactor.core.publisher.Flux; import software.amazon.awssdk.auth.credentials.AwsCredentialsProvider; @@ -436,7 +437,7 @@ public class Anthropic3ChatBedrockApi extends /** * Anthropic models version. */ - public enum AnthropicChatModel { + public enum AnthropicChatModel implements ModelDescription { /** * anthropic.claude-instant-v1 @@ -476,6 +477,11 @@ public class Anthropic3ChatBedrockApi extends this.id = value; } + @Override + public String getModelName() { + return this.id; + } + } @Override diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatClient.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatModel.java similarity index 89% rename from models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatClient.java rename to models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatModel.java index 3ff2b2cee..456b7566c 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatClient.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatModel.java @@ -24,11 +24,11 @@ import org.springframework.ai.bedrock.MessageToPromptConverter; import org.springframework.ai.bedrock.cohere.api.CohereChatBedrockApi; import org.springframework.ai.bedrock.cohere.api.CohereChatBedrockApi.CohereChatRequest; import org.springframework.ai.bedrock.cohere.api.CohereChatBedrockApi.CohereChatResponse; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; import org.springframework.ai.chat.metadata.Usage; import org.springframework.ai.chat.prompt.Prompt; @@ -39,17 +39,17 @@ import org.springframework.util.Assert; * @author Christian Tzolov * @since 0.8.0 */ -public class BedrockCohereChatClient implements ChatClient, StreamingChatClient { +public class BedrockCohereChatModel implements ChatModel, StreamingChatModel { private final CohereChatBedrockApi chatApi; private final BedrockCohereChatOptions defaultOptions; - public BedrockCohereChatClient(CohereChatBedrockApi chatApi) { + public BedrockCohereChatModel(CohereChatBedrockApi chatApi) { this(chatApi, BedrockCohereChatOptions.builder().build()); } - public BedrockCohereChatClient(CohereChatBedrockApi chatApi, BedrockCohereChatOptions options) { + public BedrockCohereChatModel(CohereChatBedrockApi chatApi, BedrockCohereChatOptions options) { Assert.notNull(chatApi, "CohereChatBedrockApi must not be null"); Assert.notNull(options, "BedrockCohereChatOptions must not be null"); @@ -114,4 +114,9 @@ public class BedrockCohereChatClient implements ChatClient, StreamingChatClient return request; } + @Override + public ChatOptions getDefaultOptions() { + return BedrockCohereChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatOptions.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatOptions.java index 89f625432..e0ab181cc 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatOptions.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatOptions.java @@ -213,4 +213,17 @@ public class BedrockCohereChatOptions implements ChatOptions { this.truncate = truncate; } + public static BedrockCohereChatOptions fromOptions(BedrockCohereChatOptions fromOptions) { + return builder().withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withTopK(fromOptions.getTopK()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withStopSequences(fromOptions.getStopSequences()) + .withReturnLikelihoods(fromOptions.getReturnLikelihoods()) + .withNumGenerations(fromOptions.getNumGenerations()) + .withLogitBias(fromOptions.getLogitBias()) + .withTruncate(fromOptions.getTruncate()) + .build(); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingClient.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingModel.java similarity index 88% rename from models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingClient.java rename to models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingModel.java index 2c0145059..25e3f35b4 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingClient.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingModel.java @@ -22,7 +22,7 @@ import org.springframework.ai.bedrock.cohere.api.CohereEmbeddingBedrockApi; import org.springframework.ai.bedrock.cohere.api.CohereEmbeddingBedrockApi.CohereEmbeddingRequest; import org.springframework.ai.bedrock.cohere.api.CohereEmbeddingBedrockApi.CohereEmbeddingResponse; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.AbstractEmbeddingClient; +import org.springframework.ai.embedding.AbstractEmbeddingModel; import org.springframework.ai.embedding.Embedding; import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingRequest; @@ -31,14 +31,14 @@ import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.util.Assert; /** - * {@link org.springframework.ai.embedding.EmbeddingClient} implementation that uses the + * {@link org.springframework.ai.embedding.EmbeddingModel} implementation that uses the * Bedrock Cohere Embedding API. Note: The invocation metrics are not exposed by AWS for * this API. If this change in the future we will add it as metadata. * * @author Christian Tzolov * @since 0.8.0 */ -public class BedrockCohereEmbeddingClient extends AbstractEmbeddingClient { +public class BedrockCohereEmbeddingModel extends AbstractEmbeddingModel { private final CohereEmbeddingBedrockApi embeddingApi; @@ -50,7 +50,7 @@ public class BedrockCohereEmbeddingClient extends AbstractEmbeddingClient { // private CohereEmbeddingRequest.Truncate truncate = // CohereEmbeddingRequest.Truncate.NONE; - public BedrockCohereEmbeddingClient(CohereEmbeddingBedrockApi cohereEmbeddingBedrockApi) { + public BedrockCohereEmbeddingModel(CohereEmbeddingBedrockApi cohereEmbeddingBedrockApi) { this(cohereEmbeddingBedrockApi, BedrockCohereEmbeddingOptions.builder() .withInputType(CohereEmbeddingRequest.InputType.SEARCH_DOCUMENT) @@ -58,7 +58,7 @@ public class BedrockCohereEmbeddingClient extends AbstractEmbeddingClient { .build()); } - public BedrockCohereEmbeddingClient(CohereEmbeddingBedrockApi cohereEmbeddingBedrockApi, + public BedrockCohereEmbeddingModel(CohereEmbeddingBedrockApi cohereEmbeddingBedrockApi, BedrockCohereEmbeddingOptions options) { Assert.notNull(cohereEmbeddingBedrockApi, "CohereEmbeddingBedrockApi must not be null"); Assert.notNull(options, "BedrockCohereEmbeddingOptions must not be null"); @@ -71,7 +71,7 @@ public class BedrockCohereEmbeddingClient extends AbstractEmbeddingClient { // * @param inputType the input type to use. // * @return this client. // */ - // public BedrockCohereEmbeddingClient withInputType(CohereEmbeddingRequest.InputType + // public BedrockCohereEmbeddingModel withInputType(CohereEmbeddingRequest.InputType // inputType) { // this.inputType = inputType; // return this; @@ -85,7 +85,7 @@ public class BedrockCohereEmbeddingClient extends AbstractEmbeddingClient { // * @param truncate the truncate option to use. // * @return this client. // */ - // public BedrockCohereEmbeddingClient withTruncate(CohereEmbeddingRequest.Truncate + // public BedrockCohereEmbeddingModel withTruncate(CohereEmbeddingRequest.Truncate // truncate) { // this.truncate = truncate; // return this; diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/api/CohereChatBedrockApi.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/api/CohereChatBedrockApi.java index 5b133a997..766271b87 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/api/CohereChatBedrockApi.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/api/CohereChatBedrockApi.java @@ -30,6 +30,7 @@ import software.amazon.awssdk.regions.Region; import org.springframework.ai.bedrock.api.AbstractBedrockApi; import org.springframework.ai.bedrock.cohere.api.CohereChatBedrockApi.CohereChatRequest; import org.springframework.ai.bedrock.cohere.api.CohereChatBedrockApi.CohereChatResponse; +import org.springframework.ai.model.ModelDescription; import org.springframework.util.Assert; /** @@ -366,7 +367,7 @@ public class CohereChatBedrockApi extends /** * Cohere models version. */ - public enum CohereChatModel { + public enum CohereChatModel implements ModelDescription { /** * cohere.command-light-text-v14 @@ -390,6 +391,11 @@ public class CohereChatBedrockApi extends CohereChatModel(String value) { this.id = value; } + + @Override + public String getModelName() { + return this.id; + } } @Override diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatClient.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatModel.java similarity index 85% rename from models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatClient.java rename to models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatModel.java index 7a11a2524..883e55b0e 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatClient.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatModel.java @@ -19,7 +19,7 @@ package org.springframework.ai.bedrock.jurassic2; import org.springframework.ai.bedrock.MessageToPromptConverter; import org.springframework.ai.bedrock.jurassic2.api.Ai21Jurassic2ChatBedrockApi; import org.springframework.ai.bedrock.jurassic2.api.Ai21Jurassic2ChatBedrockApi.Ai21Jurassic2ChatRequest; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; @@ -29,19 +29,18 @@ import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.util.Assert; /** - * Java {@link ChatClient} for the Bedrock Jurassic2 chat generative model. + * Java {@link ChatModel} for the Bedrock Jurassic2 chat generative model. * * @author Ahmed Yousri * @since 1.0.0 */ -public class BedrockAi21Jurassic2ChatClient implements ChatClient { +public class BedrockAi21Jurassic2ChatModel implements ChatModel { private final Ai21Jurassic2ChatBedrockApi chatApi; private final BedrockAi21Jurassic2ChatOptions defaultOptions; - public BedrockAi21Jurassic2ChatClient(Ai21Jurassic2ChatBedrockApi chatApi, - BedrockAi21Jurassic2ChatOptions options) { + public BedrockAi21Jurassic2ChatModel(Ai21Jurassic2ChatBedrockApi chatApi, BedrockAi21Jurassic2ChatOptions options) { Assert.notNull(chatApi, "Ai21Jurassic2ChatBedrockApi must not be null"); Assert.notNull(options, "BedrockAi21Jurassic2ChatOptions must not be null"); @@ -49,7 +48,7 @@ public class BedrockAi21Jurassic2ChatClient implements ChatClient { this.defaultOptions = options; } - public BedrockAi21Jurassic2ChatClient(Ai21Jurassic2ChatBedrockApi chatApi) { + public BedrockAi21Jurassic2ChatModel(Ai21Jurassic2ChatBedrockApi chatApi) { this(chatApi, BedrockAi21Jurassic2ChatOptions.builder() .withTemperature(0.8f) @@ -114,11 +113,16 @@ public class BedrockAi21Jurassic2ChatClient implements ChatClient { return this; } - public BedrockAi21Jurassic2ChatClient build() { - return new BedrockAi21Jurassic2ChatClient(chatApi, + public BedrockAi21Jurassic2ChatModel build() { + return new BedrockAi21Jurassic2ChatModel(chatApi, options != null ? options : BedrockAi21Jurassic2ChatOptions.builder().build()); } } + @Override + public ChatOptions getDefaultOptions() { + return BedrockAi21Jurassic2ChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatOptions.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatOptions.java index 4b62a5854..c165c61c1 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatOptions.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatOptions.java @@ -413,4 +413,19 @@ public class BedrockAi21Jurassic2ChatOptions implements ChatOptions { } } + public static BedrockAi21Jurassic2ChatOptions fromOptions(BedrockAi21Jurassic2ChatOptions fromOptions) { + return builder().withPrompt(fromOptions.getPrompt()) + .withNumResults(fromOptions.getNumResults()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withMinTokens(fromOptions.getMinTokens()) + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withTopK(fromOptions.getTopK()) + .withStopSequences(fromOptions.getStopSequences()) + .withFrequencyPenalty(fromOptions.getFrequencyPenalty()) + .withPresencePenalty(fromOptions.getPresencePenalty()) + .withCountPenalty(fromOptions.getCountPenalty()) + .build(); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/api/Ai21Jurassic2ChatBedrockApi.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/api/Ai21Jurassic2ChatBedrockApi.java index 0ec58c8bd..fecf70fa4 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/api/Ai21Jurassic2ChatBedrockApi.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/api/Ai21Jurassic2ChatBedrockApi.java @@ -27,6 +27,7 @@ import com.fasterxml.jackson.databind.ObjectMapper; import org.springframework.ai.bedrock.api.AbstractBedrockApi; import org.springframework.ai.bedrock.jurassic2.api.Ai21Jurassic2ChatBedrockApi.Ai21Jurassic2ChatRequest; import org.springframework.ai.bedrock.jurassic2.api.Ai21Jurassic2ChatBedrockApi.Ai21Jurassic2ChatResponse; +import org.springframework.ai.model.ModelDescription; import software.amazon.awssdk.auth.credentials.AwsCredentialsProvider; import software.amazon.awssdk.regions.Region; @@ -371,7 +372,7 @@ public class Ai21Jurassic2ChatBedrockApi extends /** * Ai21 Jurassic2 models version. */ - public enum Ai21Jurassic2ChatModel { + public enum Ai21Jurassic2ChatModel implements ModelDescription { /** * ai21.j2-mid-v1 @@ -395,6 +396,11 @@ public class Ai21Jurassic2ChatBedrockApi extends Ai21Jurassic2ChatModel(String value) { this.id = value; } + + @Override + public String getModelName() { + return this.id; + } } @Override diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatClient.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatModel.java similarity index 88% rename from models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatClient.java rename to models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatModel.java index c1be58e5a..3c0634c53 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatClient.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatModel.java @@ -23,11 +23,11 @@ import org.springframework.ai.bedrock.MessageToPromptConverter; import org.springframework.ai.bedrock.llama.api.LlamaChatBedrockApi; import org.springframework.ai.bedrock.llama.api.LlamaChatBedrockApi.LlamaChatRequest; import org.springframework.ai.bedrock.llama.api.LlamaChatBedrockApi.LlamaChatResponse; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; import org.springframework.ai.chat.metadata.Usage; import org.springframework.ai.chat.prompt.Prompt; @@ -35,25 +35,25 @@ import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.util.Assert; /** - * Java {@link ChatClient} and {@link StreamingChatClient} for the Bedrock Llama chat + * Java {@link ChatModel} and {@link StreamingChatModel} for the Bedrock Llama chat * generative. * * @author Christian Tzolov * @author Wei Jiang * @since 0.8.0 */ -public class BedrockLlamaChatClient implements ChatClient, StreamingChatClient { +public class BedrockLlamaChatModel implements ChatModel, StreamingChatModel { private final LlamaChatBedrockApi chatApi; private final BedrockLlamaChatOptions defaultOptions; - public BedrockLlamaChatClient(LlamaChatBedrockApi chatApi) { + public BedrockLlamaChatModel(LlamaChatBedrockApi chatApi) { this(chatApi, BedrockLlamaChatOptions.builder().withTemperature(0.8f).withTopP(0.9f).withMaxGenLen(100).build()); } - public BedrockLlamaChatClient(LlamaChatBedrockApi chatApi, BedrockLlamaChatOptions options) { + public BedrockLlamaChatModel(LlamaChatBedrockApi chatApi, BedrockLlamaChatOptions options) { Assert.notNull(chatApi, "LlamaChatBedrockApi must not be null"); Assert.notNull(options, "BedrockLlamaChatOptions must not be null"); @@ -130,4 +130,9 @@ public class BedrockLlamaChatClient implements ChatClient, StreamingChatClient { return request; } + @Override + public ChatOptions getDefaultOptions() { + return BedrockLlamaChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatOptions.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatOptions.java index 3502fd4c4..4d6c0a6e0 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatOptions.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatOptions.java @@ -109,4 +109,11 @@ public class BedrockLlamaChatOptions implements ChatOptions { throw new UnsupportedOperationException("Unsupported option: 'TopK'"); } + public static BedrockLlamaChatOptions fromOptions(BedrockLlamaChatOptions fromOptions) { + return builder().withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withMaxGenLen(fromOptions.getMaxGenLen()) + .build(); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/api/LlamaChatBedrockApi.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/api/LlamaChatBedrockApi.java index 25d71aede..16af9735e 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/api/LlamaChatBedrockApi.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/api/LlamaChatBedrockApi.java @@ -26,6 +26,7 @@ import software.amazon.awssdk.regions.Region; import org.springframework.ai.bedrock.api.AbstractBedrockApi; import org.springframework.ai.bedrock.llama.api.LlamaChatBedrockApi.LlamaChatRequest; import org.springframework.ai.bedrock.llama.api.LlamaChatBedrockApi.LlamaChatResponse; +import org.springframework.ai.model.ModelDescription; import java.time.Duration; @@ -204,7 +205,7 @@ public class LlamaChatBedrockApi extends /** * Llama models version. */ - public enum LlamaChatModel { + public enum LlamaChatModel implements ModelDescription { /** * meta.llama2-13b-chat-v1 @@ -238,6 +239,11 @@ public class LlamaChatBedrockApi extends LlamaChatModel(String value) { this.id = value; } + + @Override + public String getModelName() { + return this.id; + } } @Override diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatClient.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatModel.java similarity index 91% rename from models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatClient.java rename to models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatModel.java index e77d8277c..571210178 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatClient.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatModel.java @@ -24,11 +24,11 @@ import org.springframework.ai.bedrock.titan.api.TitanChatBedrockApi; import org.springframework.ai.bedrock.titan.api.TitanChatBedrockApi.TitanChatRequest; import org.springframework.ai.bedrock.titan.api.TitanChatBedrockApi.TitanChatResponse; import org.springframework.ai.bedrock.titan.api.TitanChatBedrockApi.TitanChatResponseChunk; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; import org.springframework.ai.chat.metadata.Usage; import org.springframework.ai.chat.prompt.Prompt; @@ -39,17 +39,17 @@ import org.springframework.util.Assert; * @author Christian Tzolov * @since 0.8.0 */ -public class BedrockTitanChatClient implements ChatClient, StreamingChatClient { +public class BedrockTitanChatModel implements ChatModel, StreamingChatModel { private final TitanChatBedrockApi chatApi; private final BedrockTitanChatOptions defaultOptions; - public BedrockTitanChatClient(TitanChatBedrockApi chatApi) { + public BedrockTitanChatModel(TitanChatBedrockApi chatApi) { this(chatApi, BedrockTitanChatOptions.builder().withTemperature(0.8f).build()); } - public BedrockTitanChatClient(TitanChatBedrockApi chatApi, BedrockTitanChatOptions defaultOptions) { + public BedrockTitanChatModel(TitanChatBedrockApi chatApi, BedrockTitanChatOptions defaultOptions) { Assert.notNull(chatApi, "ChatApi must not be null"); Assert.notNull(defaultOptions, "DefaultOptions must not be null"); this.chatApi = chatApi; @@ -146,4 +146,9 @@ public class BedrockTitanChatClient implements ChatClient, StreamingChatClient { }; } + @Override + public ChatOptions getDefaultOptions() { + return BedrockTitanChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatOptions.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatOptions.java index 61b16d16e..d53126a0b 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatOptions.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatOptions.java @@ -128,4 +128,12 @@ public class BedrockTitanChatOptions implements ChatOptions { throw new UnsupportedOperationException("Bedrock Titan Chat does not support the 'TopK' option.'"); } + public static BedrockTitanChatOptions fromOptions(BedrockTitanChatOptions fromOptions) { + return builder().withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withMaxTokenCount(fromOptions.getMaxTokenCount()) + .withStopSequences(fromOptions.getStopSequences()) + .build(); + } + } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingClient.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingModel.java similarity index 91% rename from models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingClient.java rename to models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingModel.java index 1d64f92ef..e3089eec5 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingClient.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingModel.java @@ -26,7 +26,7 @@ import org.springframework.ai.bedrock.titan.api.TitanEmbeddingBedrockApi; import org.springframework.ai.bedrock.titan.api.TitanEmbeddingBedrockApi.TitanEmbeddingRequest; import org.springframework.ai.bedrock.titan.api.TitanEmbeddingBedrockApi.TitanEmbeddingResponse; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.AbstractEmbeddingClient; +import org.springframework.ai.embedding.AbstractEmbeddingModel; import org.springframework.ai.embedding.Embedding; import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingRequest; @@ -34,7 +34,7 @@ import org.springframework.ai.embedding.EmbeddingResponse; import org.springframework.util.Assert; /** - * {@link org.springframework.ai.embedding.EmbeddingClient} implementation that uses the + * {@link org.springframework.ai.embedding.EmbeddingModel} implementation that uses the * Bedrock Titan Embedding API. Titan Embedding supports text and image (encoded in * base64) inputs. * @@ -44,7 +44,7 @@ import org.springframework.util.Assert; * @author Wei Jiang * @since 0.8.0 */ -public class BedrockTitanEmbeddingClient extends AbstractEmbeddingClient { +public class BedrockTitanEmbeddingModel extends AbstractEmbeddingModel { private final Logger logger = LoggerFactory.getLogger(getClass()); @@ -61,7 +61,7 @@ public class BedrockTitanEmbeddingClient extends AbstractEmbeddingClient { */ private InputType inputType = InputType.TEXT; - public BedrockTitanEmbeddingClient(TitanEmbeddingBedrockApi titanEmbeddingBedrockApi) { + public BedrockTitanEmbeddingModel(TitanEmbeddingBedrockApi titanEmbeddingBedrockApi) { this.embeddingApi = titanEmbeddingBedrockApi; } @@ -69,7 +69,7 @@ public class BedrockTitanEmbeddingClient extends AbstractEmbeddingClient { * Titan Embedding API input types. Could be either text or image (encoded in base64). * @param inputType the input type to use. */ - public BedrockTitanEmbeddingClient withInputType(InputType inputType) { + public BedrockTitanEmbeddingModel withInputType(InputType inputType) { this.inputType = inputType; return this; } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingOptions.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingOptions.java index fd1c609bf..d2770f245 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingOptions.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingOptions.java @@ -18,7 +18,7 @@ package org.springframework.ai.bedrock.titan; import com.fasterxml.jackson.annotation.JsonInclude; import com.fasterxml.jackson.annotation.JsonInclude.Include; -import org.springframework.ai.bedrock.titan.BedrockTitanEmbeddingClient.InputType; +import org.springframework.ai.bedrock.titan.BedrockTitanEmbeddingModel.InputType; import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.util.Assert; diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/api/TitanChatBedrockApi.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/api/TitanChatBedrockApi.java index 78c7cd931..ce1842adf 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/api/TitanChatBedrockApi.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/api/TitanChatBedrockApi.java @@ -31,6 +31,7 @@ import org.springframework.ai.bedrock.titan.api.TitanChatBedrockApi.TitanChatReq import org.springframework.ai.bedrock.titan.api.TitanChatBedrockApi.TitanChatResponse; import org.springframework.ai.bedrock.titan.api.TitanChatBedrockApi.TitanChatResponse.CompletionReason; import org.springframework.ai.bedrock.titan.api.TitanChatBedrockApi.TitanChatResponseChunk; +import org.springframework.ai.model.ModelDescription; /** * Java client for the Bedrock Titan chat model. @@ -265,7 +266,7 @@ public class TitanChatBedrockApi extends /** * Titan models version. */ - public enum TitanChatModel { + public enum TitanChatModel implements ModelDescription { /** * amazon.titan-text-lite-v1 @@ -294,6 +295,11 @@ public class TitanChatBedrockApi extends TitanChatModel(String value) { this.id = value; } + + @Override + public String getModelName() { + return this.id; + } } @Override diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatClientIT.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatModelIT.java similarity index 90% rename from models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatClientIT.java rename to models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatModelIT.java index 43eb7e776..4cdaa92b8 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatClientIT.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatModelIT.java @@ -55,12 +55,12 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @EnabledIfEnvironmentVariable(named = "AWS_ACCESS_KEY_ID", matches = ".*") @EnabledIfEnvironmentVariable(named = "AWS_SECRET_ACCESS_KEY", matches = ".*") -class BedrockAnthropicChatClientIT { +class BedrockAnthropicChatModelIT { - private static final Logger logger = LoggerFactory.getLogger(BedrockAnthropicChatClientIT.class); + private static final Logger logger = LoggerFactory.getLogger(BedrockAnthropicChatModelIT.class); @Autowired - private BedrockAnthropicChatClient client; + private BedrockAnthropicChatModel chatModel; @Value("classpath:/prompts/system-message.st") private Resource systemResource; @@ -68,8 +68,8 @@ class BedrockAnthropicChatClientIT { @Test void multipleStreamAttempts() { - Flux joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?"))); - Flux joke2Stream = client.stream(new Prompt(new UserMessage("Tell me a toy joke?"))); + Flux joke1Stream = chatModel.stream(new Prompt(new UserMessage("Tell me a joke?"))); + Flux joke2Stream = chatModel.stream(new Prompt(new UserMessage("Tell me a toy joke?"))); String joke1 = joke1Stream.collectList() .block() @@ -101,7 +101,7 @@ class BedrockAnthropicChatClientIT { Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); - ChatResponse response = client.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); } @@ -119,7 +119,7 @@ class BedrockAnthropicChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "ice cream flavors.", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = this.client.call(prompt).getResult(); + Generation generation = this.chatModel.call(prompt).getResult(); List list = outputParser.convert(generation.getOutput().getContent()); assertThat(list).hasSize(5); @@ -137,7 +137,7 @@ class BedrockAnthropicChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); Map result = outputConverter.convert(generation.getOutput().getContent()); assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); @@ -161,7 +161,7 @@ class BedrockAnthropicChatClientIT { """; PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); ActorsFilmsRecord actorsFilms = outputConvert.convert(generation.getOutput().getContent()); assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks"); @@ -182,7 +182,7 @@ class BedrockAnthropicChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - String generationTextFromStream = client.stream(prompt) + String generationTextFromStream = chatModel.stream(prompt) .collectList() .block() .stream() @@ -209,8 +209,8 @@ class BedrockAnthropicChatClientIT { } @Bean - public BedrockAnthropicChatClient anthropicChatClient(AnthropicChatBedrockApi anthropicApi) { - return new BedrockAnthropicChatClient(anthropicApi); + public BedrockAnthropicChatModel anthropicChatModel(AnthropicChatBedrockApi anthropicApi) { + return new BedrockAnthropicChatModel(anthropicApi); } } diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicCreateRequestTests.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicCreateRequestTests.java index 928e183ad..c8b5cbe85 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicCreateRequestTests.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicCreateRequestTests.java @@ -38,7 +38,7 @@ public class BedrockAnthropicCreateRequestTests { @Test public void createRequestWithChatOptions() { - var client = new BedrockAnthropicChatClient(anthropicChatApi, + var client = new BedrockAnthropicChatModel(anthropicChatApi, AnthropicChatOptions.builder() .withTemperature(66.6f) .withTopK(66) diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatClientIT.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatModelIT.java similarity index 90% rename from models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatClientIT.java rename to models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatModelIT.java index 9568f4f69..5ff88d2b4 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatClientIT.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatModelIT.java @@ -59,12 +59,12 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @EnabledIfEnvironmentVariable(named = "AWS_ACCESS_KEY_ID", matches = ".*") @EnabledIfEnvironmentVariable(named = "AWS_SECRET_ACCESS_KEY", matches = ".*") -class BedrockAnthropic3ChatClientIT { +class BedrockAnthropic3ChatModelIT { - private static final Logger logger = LoggerFactory.getLogger(BedrockAnthropic3ChatClientIT.class); + private static final Logger logger = LoggerFactory.getLogger(BedrockAnthropic3ChatModelIT.class); @Autowired - private BedrockAnthropic3ChatClient client; + private BedrockAnthropic3ChatModel chatModel; @Value("classpath:/prompts/system-message.st") private Resource systemResource; @@ -72,8 +72,8 @@ class BedrockAnthropic3ChatClientIT { @Test void multipleStreamAttempts() { - Flux joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?"))); - Flux joke2Stream = client.stream(new Prompt(new UserMessage("Tell me a toy joke?"))); + Flux joke1Stream = chatModel.stream(new Prompt(new UserMessage("Tell me a joke?"))); + Flux joke2Stream = chatModel.stream(new Prompt(new UserMessage("Tell me a toy joke?"))); String joke1 = joke1Stream.collectList() .block() @@ -105,7 +105,7 @@ class BedrockAnthropic3ChatClientIT { Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); - ChatResponse response = client.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); } @@ -123,7 +123,7 @@ class BedrockAnthropic3ChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "ice cream flavors.", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = this.client.call(prompt).getResult(); + Generation generation = this.chatModel.call(prompt).getResult(); List list = outputConverter.convert(generation.getOutput().getContent()); assertThat(list).hasSize(5); @@ -142,7 +142,7 @@ class BedrockAnthropic3ChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); Map result = outputConverter.convert(generation.getOutput().getContent()); assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); @@ -166,7 +166,7 @@ class BedrockAnthropic3ChatClientIT { """; PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); ActorsFilmsRecord actorsFilms = outputConverter.convert(generation.getOutput().getContent()); assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks"); @@ -187,7 +187,7 @@ class BedrockAnthropic3ChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - String generationTextFromStream = client.stream(prompt) + String generationTextFromStream = chatModel.stream(prompt) .collectList() .block() .stream() @@ -211,7 +211,7 @@ class BedrockAnthropic3ChatClientIT { var userMessage = new UserMessage("Explain what do you see o this picture?", List.of(new Media(MimeTypeUtils.IMAGE_PNG, imageData))); - var response = client.call(new Prompt(List.of(userMessage))); + var response = chatModel.call(new Prompt(List.of(userMessage))); logger.info(response.getResult().getOutput().getContent()); assertThat(response.getResult().getOutput().getContent()).contains("bananas", "apple", "basket"); @@ -228,8 +228,8 @@ class BedrockAnthropic3ChatClientIT { } @Bean - public BedrockAnthropic3ChatClient anthropicChatClient(Anthropic3ChatBedrockApi anthropicApi) { - return new BedrockAnthropic3ChatClient(anthropicApi); + public BedrockAnthropic3ChatModel anthropicChatModel(Anthropic3ChatBedrockApi anthropicApi) { + return new BedrockAnthropic3ChatModel(anthropicApi); } } diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3CreateRequestTests.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3CreateRequestTests.java index 480f914c3..31486f9e9 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3CreateRequestTests.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3CreateRequestTests.java @@ -37,7 +37,7 @@ public class BedrockAnthropic3CreateRequestTests { @Test public void createRequestWithChatOptions() { - var client = new BedrockAnthropic3ChatClient(anthropicChatApi, + var client = new BedrockAnthropic3ChatModel(anthropicChatApi, Anthropic3ChatOptions.builder() .withTemperature(66.6f) .withTopK(66) diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatCreateRequestTests.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatCreateRequestTests.java index 661c19e32..c757efe04 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatCreateRequestTests.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatCreateRequestTests.java @@ -45,7 +45,7 @@ public class BedrockCohereChatCreateRequestTests { @Test public void createRequestWithChatOptions() { - var client = new BedrockCohereChatClient(chatApi, + var client = new BedrockCohereChatModel(chatApi, BedrockCohereChatOptions.builder() .withTemperature(66.6f) .withTopK(66) diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatClientIT.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatModelIT.java similarity index 90% rename from models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatClientIT.java rename to models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatModelIT.java index 567926810..287d22bce 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatClientIT.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatModelIT.java @@ -54,10 +54,10 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @EnabledIfEnvironmentVariable(named = "AWS_ACCESS_KEY_ID", matches = ".*") @EnabledIfEnvironmentVariable(named = "AWS_SECRET_ACCESS_KEY", matches = ".*") -class BedrockCohereChatClientIT { +class BedrockCohereChatModelIT { @Autowired - private BedrockCohereChatClient client; + private BedrockCohereChatModel chatModel; @Value("classpath:/prompts/system-message.st") private Resource systemResource; @@ -65,8 +65,8 @@ class BedrockCohereChatClientIT { @Test void multipleStreamAttempts() { - Flux joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?"))); - Flux joke2Stream = client.stream(new Prompt(new UserMessage("Tell me a toy joke?"))); + Flux joke1Stream = chatModel.stream(new Prompt(new UserMessage("Tell me a joke?"))); + Flux joke2Stream = chatModel.stream(new Prompt(new UserMessage("Tell me a toy joke?"))); String joke1 = joke1Stream.collectList() .block() @@ -98,7 +98,7 @@ class BedrockCohereChatClientIT { SystemPromptTemplate systemPromptTemplate = new SystemPromptTemplate(systemResource); Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", name, "voice", voice)); Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); - ChatResponse response = client.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); } @@ -115,7 +115,7 @@ class BedrockCohereChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "ice cream flavors.", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = this.client.call(prompt).getResult(); + Generation generation = this.chatModel.call(prompt).getResult(); List list = outputConverter.convert(generation.getOutput().getContent()); assertThat(list).hasSize(5); @@ -134,7 +134,7 @@ class BedrockCohereChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); Map result = outputConverter.convert(generation.getOutput().getContent()); assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); @@ -157,7 +157,7 @@ class BedrockCohereChatClientIT { """; PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); ActorsFilmsRecord actorsFilms = outputConverter.convert(generation.getOutput().getContent()); assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks"); @@ -178,7 +178,7 @@ class BedrockCohereChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - String generationTextFromStream = client.stream(prompt) + String generationTextFromStream = chatModel.stream(prompt) .collectList() .block() .stream() @@ -205,8 +205,8 @@ class BedrockCohereChatClientIT { } @Bean - public BedrockCohereChatClient cohereChatClient(CohereChatBedrockApi cohereApi) { - return new BedrockCohereChatClient(cohereApi); + public BedrockCohereChatModel cohereChatModel(CohereChatBedrockApi cohereApi) { + return new BedrockCohereChatModel(cohereApi); } } diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingClientIT.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingModelIT.java similarity index 81% rename from models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingClientIT.java rename to models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingModelIT.java index 1dab72ce3..194b657ed 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingClientIT.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingModelIT.java @@ -39,24 +39,24 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @EnabledIfEnvironmentVariable(named = "AWS_ACCESS_KEY_ID", matches = ".*") @EnabledIfEnvironmentVariable(named = "AWS_SECRET_ACCESS_KEY", matches = ".*") -class BedrockCohereEmbeddingClientIT { +class BedrockCohereEmbeddingModelIT { @Autowired - private BedrockCohereEmbeddingClient embeddingClient; + private BedrockCohereEmbeddingModel embeddingModel; @Test void singleEmbedding() { - assertThat(embeddingClient).isNotNull(); - EmbeddingResponse embeddingResponse = embeddingClient.embedForResponse(List.of("Hello World")); + assertThat(embeddingModel).isNotNull(); + EmbeddingResponse embeddingResponse = embeddingModel.embedForResponse(List.of("Hello World")); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); - assertThat(embeddingClient.dimensions()).isEqualTo(1024); + assertThat(embeddingModel.dimensions()).isEqualTo(1024); } @Test void batchEmbedding() { - assertThat(embeddingClient).isNotNull(); - EmbeddingResponse embeddingResponse = embeddingClient + assertThat(embeddingModel).isNotNull(); + EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); assertThat(embeddingResponse.getResults()).hasSize(2); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); @@ -64,13 +64,13 @@ class BedrockCohereEmbeddingClientIT { assertThat(embeddingResponse.getResults().get(1).getOutput()).isNotEmpty(); assertThat(embeddingResponse.getResults().get(1).getIndex()).isEqualTo(1); - assertThat(embeddingClient.dimensions()).isEqualTo(1024); + assertThat(embeddingModel.dimensions()).isEqualTo(1024); } @Test void embeddingWthOptions() { - assertThat(embeddingClient).isNotNull(); - EmbeddingResponse embeddingResponse = embeddingClient + assertThat(embeddingModel).isNotNull(); + EmbeddingResponse embeddingResponse = embeddingModel .call(new EmbeddingRequest(List.of("Hello World", "World is big and salvation is near"), BedrockCohereEmbeddingOptions.builder().withInputType(InputType.SEARCH_DOCUMENT).build())); assertThat(embeddingResponse.getResults()).hasSize(2); @@ -79,7 +79,7 @@ class BedrockCohereEmbeddingClientIT { assertThat(embeddingResponse.getResults().get(1).getOutput()).isNotEmpty(); assertThat(embeddingResponse.getResults().get(1).getIndex()).isEqualTo(1); - assertThat(embeddingClient.dimensions()).isEqualTo(1024); + assertThat(embeddingModel.dimensions()).isEqualTo(1024); } @SpringBootConfiguration @@ -93,8 +93,8 @@ class BedrockCohereEmbeddingClientIT { } @Bean - public BedrockCohereEmbeddingClient cohereAiEmbedding(CohereEmbeddingBedrockApi cohereEmbeddingApi) { - return new BedrockCohereEmbeddingClient(cohereEmbeddingApi); + public BedrockCohereEmbeddingModel cohereAiEmbedding(CohereEmbeddingBedrockApi cohereEmbeddingApi) { + return new BedrockCohereEmbeddingModel(cohereEmbeddingApi); } } diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatClientIT.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatModelIT.java similarity index 92% rename from models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatClientIT.java rename to models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatModelIT.java index f6614b852..c80ebdef9 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatClientIT.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatModelIT.java @@ -49,10 +49,10 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @EnabledIfEnvironmentVariable(named = "AWS_ACCESS_KEY_ID", matches = ".*") @EnabledIfEnvironmentVariable(named = "AWS_SECRET_ACCESS_KEY", matches = ".*") -class BedrockAi21Jurassic2ChatClientIT { +class BedrockAi21Jurassic2ChatModelIT { @Autowired - private BedrockAi21Jurassic2ChatClient client; + private BedrockAi21Jurassic2ChatModel chatModel; @Value("classpath:/prompts/system-message.st") private Resource systemResource; @@ -66,7 +66,7 @@ class BedrockAi21Jurassic2ChatClientIT { Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); - ChatResponse response = client.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); } @@ -83,7 +83,7 @@ class BedrockAi21Jurassic2ChatClientIT { UserMessage userMessage = new UserMessage("Can you express happiness using an emoji like 😄 ?"); Prompt prompt = new Prompt(List.of(userMessage), options); - ChatResponse response = client.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResult().getOutput().getContent()).matches(content -> content.contains("😄")); } @@ -103,7 +103,7 @@ class BedrockAi21Jurassic2ChatClientIT { Prompt prompt = new Prompt(List.of(userMessage, systemMessage), options); - ChatResponse response = client.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResult().getOutput().getContent()).doesNotContain("😄"); } @@ -120,7 +120,7 @@ class BedrockAi21Jurassic2ChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); Map result = outputConverter.convert(generation.getOutput().getContent()); assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); @@ -135,7 +135,7 @@ class BedrockAi21Jurassic2ChatClientIT { Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); - ChatResponse response = client.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResult().getOutput().getContent()).contains("AI"); } @@ -152,9 +152,9 @@ class BedrockAi21Jurassic2ChatClientIT { } @Bean - public BedrockAi21Jurassic2ChatClient bedrockAi21Jurassic2ChatClient( + public BedrockAi21Jurassic2ChatModel bedrockAi21Jurassic2ChatModel( Ai21Jurassic2ChatBedrockApi jurassic2ChatBedrockApi) { - return new BedrockAi21Jurassic2ChatClient(jurassic2ChatBedrockApi, + return new BedrockAi21Jurassic2ChatModel(jurassic2ChatBedrockApi, BedrockAi21Jurassic2ChatOptions.builder() .withTemperature(0.5f) .withMaxTokens(100) diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama/BedrockLlamaCreateRequestTests.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama/BedrockLlamaCreateRequestTests.java index 77018af9e..4bd48680d 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama/BedrockLlamaCreateRequestTests.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama/BedrockLlamaCreateRequestTests.java @@ -45,7 +45,7 @@ public class BedrockLlamaCreateRequestTests { @Test public void createRequestWithChatOptions() { - var client = new BedrockLlamaChatClient(api, + var client = new BedrockLlamaChatModel(api, BedrockLlamaChatOptions.builder().withTemperature(66.6f).withMaxGenLen(666).withTopP(0.66f).build()); var request = client.createRequest(new Prompt("Test message content")); diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatClientIT.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama/BedrockLlamaModelCallerIT.java similarity index 91% rename from models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatClientIT.java rename to models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama/BedrockLlamaModelCallerIT.java index 3ecaf8650..699e0acd2 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatClientIT.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama/BedrockLlamaModelCallerIT.java @@ -54,10 +54,10 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @EnabledIfEnvironmentVariable(named = "AWS_ACCESS_KEY_ID", matches = ".*") @EnabledIfEnvironmentVariable(named = "AWS_SECRET_ACCESS_KEY", matches = ".*") -class BedrockLlamaChatClientIT { +class BedrockLlamaChatModelIT { @Autowired - private BedrockLlamaChatClient client; + private BedrockLlamaChatModel chatModel; @Value("classpath:/prompts/system-message.st") private Resource systemResource; @@ -65,8 +65,8 @@ class BedrockLlamaChatClientIT { @Test void multipleStreamAttempts() { - Flux joke2Stream = client.stream(new Prompt(new UserMessage("Tell me a Toy joke?"))); - Flux joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?"))); + Flux joke2Stream = chatModel.stream(new Prompt(new UserMessage("Tell me a Toy joke?"))); + Flux joke1Stream = chatModel.stream(new Prompt(new UserMessage("Tell me a joke?"))); String joke1 = joke1Stream.collectList() .block() @@ -98,7 +98,7 @@ class BedrockLlamaChatClientIT { Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); - ChatResponse response = client.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); } @@ -116,7 +116,7 @@ class BedrockLlamaChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "ice cream flavors.", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = this.client.call(prompt).getResult(); + Generation generation = this.chatModel.call(prompt).getResult(); List list = outputConverter.convert(generation.getOutput().getContent()); assertThat(list).hasSize(5); @@ -134,7 +134,7 @@ class BedrockLlamaChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); Map result = outputConverter.convert(generation.getOutput().getContent()); assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); @@ -158,7 +158,7 @@ class BedrockLlamaChatClientIT { """; PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); ActorsFilmsRecord actorsFilms = outputConverter.convert(generation.getOutput().getContent()); assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks"); @@ -179,7 +179,7 @@ class BedrockLlamaChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - String generationTextFromStream = client.stream(prompt) + String generationTextFromStream = chatModel.stream(prompt) .collectList() .block() .stream() @@ -206,8 +206,8 @@ class BedrockLlamaChatClientIT { } @Bean - public BedrockLlamaChatClient llamaChatClient(LlamaChatBedrockApi llamaApi) { - return new BedrockLlamaChatClient(llamaApi, + public BedrockLlamaChatModel llamaChatModel(LlamaChatBedrockApi llamaApi) { + return new BedrockLlamaChatModel(llamaApi, BedrockLlamaChatOptions.builder().withTemperature(0.5f).withMaxGenLen(100).withTopP(0.9f).build()); } diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanChatCreateRequestTests.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanChatCreateRequestTests.java index 5f8065bc3..a921bb35a 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanChatCreateRequestTests.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanChatCreateRequestTests.java @@ -41,7 +41,7 @@ public class BedrockTitanChatCreateRequestTests { @Test public void createRequestWithChatOptions() { - var client = new BedrockTitanChatClient(api, + var client = new BedrockTitanChatModel(api, BedrockTitanChatOptions.builder() .withTemperature(66.6f) .withTopP(0.66f) diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingClientIT.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingModelIT.java similarity index 83% rename from models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingClientIT.java rename to models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingModelIT.java index 6400a15f0..7670e8db9 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingClientIT.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingModelIT.java @@ -26,7 +26,7 @@ import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import software.amazon.awssdk.auth.credentials.EnvironmentVariableCredentialsProvider; import software.amazon.awssdk.regions.Region; -import org.springframework.ai.bedrock.titan.BedrockTitanEmbeddingClient.InputType; +import org.springframework.ai.bedrock.titan.BedrockTitanEmbeddingModel.InputType; import org.springframework.ai.bedrock.titan.api.TitanEmbeddingBedrockApi; import org.springframework.ai.bedrock.titan.api.TitanEmbeddingBedrockApi.TitanEmbeddingModel; import org.springframework.ai.embedding.EmbeddingRequest; @@ -44,19 +44,19 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @EnabledIfEnvironmentVariable(named = "AWS_ACCESS_KEY_ID", matches = ".*") @EnabledIfEnvironmentVariable(named = "AWS_SECRET_ACCESS_KEY", matches = ".*") -class BedrockTitanEmbeddingClientIT { +class BedrockTitanEmbeddingModelIT { @Autowired - private BedrockTitanEmbeddingClient embeddingClient; + private BedrockTitanEmbeddingModel embeddingModel; @Test void singleEmbedding() { - assertThat(embeddingClient).isNotNull(); - EmbeddingResponse embeddingResponse = embeddingClient.call(new EmbeddingRequest(List.of("Hello World"), + assertThat(embeddingModel).isNotNull(); + EmbeddingResponse embeddingResponse = embeddingModel.call(new EmbeddingRequest(List.of("Hello World"), BedrockTitanEmbeddingOptions.builder().withInputType(InputType.TEXT).build())); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); - assertThat(embeddingClient.dimensions()).isEqualTo(1024); + assertThat(embeddingModel.dimensions()).isEqualTo(1024); } @Test @@ -65,12 +65,12 @@ class BedrockTitanEmbeddingClientIT { byte[] image = new DefaultResourceLoader().getResource("classpath:/spring_framework.png") .getContentAsByteArray(); - EmbeddingResponse embeddingResponse = embeddingClient + EmbeddingResponse embeddingResponse = embeddingModel .call(new EmbeddingRequest(List.of(Base64.getEncoder().encodeToString(image)), BedrockTitanEmbeddingOptions.builder().withInputType(InputType.IMAGE).build())); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); - assertThat(embeddingClient.dimensions()).isEqualTo(1024); + assertThat(embeddingModel.dimensions()).isEqualTo(1024); } @SpringBootConfiguration @@ -84,8 +84,8 @@ class BedrockTitanEmbeddingClientIT { } @Bean - public BedrockTitanEmbeddingClient titanEmbedding(TitanEmbeddingBedrockApi titanEmbeddingApi) { - return new BedrockTitanEmbeddingClient(titanEmbeddingApi); + public BedrockTitanEmbeddingModel titanEmbedding(TitanEmbeddingBedrockApi titanEmbeddingApi) { + return new BedrockTitanEmbeddingModel(titanEmbeddingApi); } } diff --git a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanChatClientIT.java b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanModelCalerlIT.java similarity index 91% rename from models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanChatClientIT.java rename to models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanModelCalerlIT.java index 3f6b611c9..acc876cf1 100644 --- a/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanChatClientIT.java +++ b/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/titan/BedrockTitanModelCalerlIT.java @@ -55,10 +55,10 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @EnabledIfEnvironmentVariable(named = "AWS_ACCESS_KEY_ID", matches = ".*") @EnabledIfEnvironmentVariable(named = "AWS_SECRET_ACCESS_KEY", matches = ".*") -class BedrockTitanChatClientIT { +class BedrockTitanModelCalerlIT { @Autowired - private BedrockTitanChatClient client; + private BedrockTitanChatModel chatModel; @Value("classpath:/prompts/system-message.st") private Resource systemResource; @@ -66,8 +66,8 @@ class BedrockTitanChatClientIT { @Test void multipleStreamAttempts() { - Flux joke1Stream = client.stream(new Prompt(new UserMessage("Tell me a joke?"))); - Flux joke2Stream = client.stream(new Prompt(new UserMessage("Tell me a toy joke?"))); + Flux joke1Stream = chatModel.stream(new Prompt(new UserMessage("Tell me a joke?"))); + Flux joke2Stream = chatModel.stream(new Prompt(new UserMessage("Tell me a toy joke?"))); String joke1 = joke1Stream.collectList() .block() @@ -99,7 +99,7 @@ class BedrockTitanChatClientIT { SystemPromptTemplate systemPromptTemplate = new SystemPromptTemplate(systemResource); Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", name, "voice", voice)); Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); - ChatResponse response = client.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); } @@ -117,7 +117,7 @@ class BedrockTitanChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "ice cream flavors.", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = this.client.call(prompt).getResult(); + Generation generation = this.chatModel.call(prompt).getResult(); List list = outputConverter.convert(generation.getOutput().getContent()); assertThat(list).hasSize(5); @@ -138,7 +138,7 @@ class BedrockTitanChatClientIT { Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); Map result = outputConverter.convert(generation.getOutput().getContent()); assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); @@ -162,7 +162,7 @@ class BedrockTitanChatClientIT { """; PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); ActorsFilmsRecord actorsFilms = outputConverter.convert(generation.getOutput().getContent()); assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks"); @@ -184,7 +184,7 @@ class BedrockTitanChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - String generationTextFromStream = client.stream(prompt) + String generationTextFromStream = chatModel.stream(prompt) .collectList() .block() .stream() @@ -211,8 +211,8 @@ class BedrockTitanChatClientIT { } @Bean - public BedrockTitanChatClient titanChatClient(TitanChatBedrockApi titanApi) { - return new BedrockTitanChatClient(titanApi); + public BedrockTitanChatModel titanChatModel(TitanChatBedrockApi titanApi) { + return new BedrockTitanChatModel(titanApi); } } diff --git a/models/spring-ai-huggingface/src/main/java/org/springframework/ai/huggingface/HuggingfaceChatClient.java b/models/spring-ai-huggingface/src/main/java/org/springframework/ai/huggingface/HuggingfaceChatModel.java similarity index 87% rename from models/spring-ai-huggingface/src/main/java/org/springframework/ai/huggingface/HuggingfaceChatClient.java rename to models/spring-ai-huggingface/src/main/java/org/springframework/ai/huggingface/HuggingfaceChatModel.java index 65b6ecdc9..6a8e8d7de 100644 --- a/models/spring-ai-huggingface/src/main/java/org/springframework/ai/huggingface/HuggingfaceChatClient.java +++ b/models/spring-ai-huggingface/src/main/java/org/springframework/ai/huggingface/HuggingfaceChatModel.java @@ -22,7 +22,7 @@ import java.util.Map; import com.fasterxml.jackson.core.type.TypeReference; import com.fasterxml.jackson.databind.ObjectMapper; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; import org.springframework.ai.huggingface.api.TextGenerationInferenceApi; @@ -31,15 +31,17 @@ import org.springframework.ai.huggingface.model.AllOfGenerateResponseDetails; import org.springframework.ai.huggingface.model.GenerateParameters; import org.springframework.ai.huggingface.model.GenerateRequest; import org.springframework.ai.huggingface.model.GenerateResponse; +import org.springframework.ai.chat.prompt.ChatOptions; +import org.springframework.ai.chat.prompt.ChatOptionsBuilder; import org.springframework.ai.chat.prompt.Prompt; /** - * An implementation of {@link ChatClient} that interfaces with HuggingFace Inference + * An implementation of {@link ChatModel} that interfaces with HuggingFace Inference * Endpoints for text generation. * * @author Mark Pollack */ -public class HuggingfaceChatClient implements ChatClient { +public class HuggingfaceChatModel implements ChatModel { /** * Token required for authenticating with the HuggingFace Inference API. @@ -68,11 +70,11 @@ public class HuggingfaceChatClient implements ChatClient { private int maxNewTokens = 1000; /** - * Constructs a new HuggingfaceChatClient with the specified API token and base path. + * Constructs a new HuggingfaceChatModel with the specified API token and base path. * @param apiToken The API token for HuggingFace. * @param basePath The base path for API requests. */ - public HuggingfaceChatClient(final String apiToken, String basePath) { + public HuggingfaceChatModel(final String apiToken, String basePath) { this.apiToken = apiToken; this.apiClient.setBasePath(basePath); this.apiClient.addDefaultHeader("Authorization", "Bearer " + this.apiToken); @@ -120,4 +122,9 @@ public class HuggingfaceChatClient implements ChatClient { this.maxNewTokens = maxNewTokens; } + @Override + public ChatOptions getDefaultOptions() { + return ChatOptionsBuilder.builder().build(); + } + } diff --git a/models/spring-ai-huggingface/src/test/java/org/springframework/ai/huggingface/HuggingfaceTestConfiguration.java b/models/spring-ai-huggingface/src/test/java/org/springframework/ai/huggingface/HuggingfaceTestConfiguration.java index e4adf3bb8..011650141 100644 --- a/models/spring-ai-huggingface/src/test/java/org/springframework/ai/huggingface/HuggingfaceTestConfiguration.java +++ b/models/spring-ai-huggingface/src/test/java/org/springframework/ai/huggingface/HuggingfaceTestConfiguration.java @@ -23,7 +23,7 @@ import org.springframework.util.StringUtils; public class HuggingfaceTestConfiguration { @Bean - public HuggingfaceChatClient huggingfaceChatClient() { + public HuggingfaceChatModel huggingfaceChatModel() { String apiKey = System.getenv("HUGGINGFACE_API_KEY"); if (!StringUtils.hasText(apiKey)) { throw new IllegalArgumentException( @@ -31,9 +31,9 @@ public class HuggingfaceTestConfiguration { } // Created aws-mistral-7b-instruct-v0-1-805 via // https://ui.endpoints.huggingface.co/ - HuggingfaceChatClient huggingfaceChatClient = new HuggingfaceChatClient(apiKey, + HuggingfaceChatModel huggingfaceChatModel = new HuggingfaceChatModel(apiKey, "https://f6hg7b3cvlmntp5i.us-east-1.aws.endpoints.huggingface.cloud"); - return huggingfaceChatClient; + return huggingfaceChatModel; } } diff --git a/models/spring-ai-huggingface/src/test/java/org/springframework/ai/huggingface/client/ClientIT.java b/models/spring-ai-huggingface/src/test/java/org/springframework/ai/huggingface/client/ClientIT.java index f21c7b9ab..654475896 100644 --- a/models/spring-ai-huggingface/src/test/java/org/springframework/ai/huggingface/client/ClientIT.java +++ b/models/spring-ai-huggingface/src/test/java/org/springframework/ai/huggingface/client/ClientIT.java @@ -20,7 +20,7 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.chat.ChatResponse; -import org.springframework.ai.huggingface.HuggingfaceChatClient; +import org.springframework.ai.huggingface.HuggingfaceChatModel; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; @@ -33,7 +33,7 @@ import static org.assertj.core.api.Assertions.assertThat; public class ClientIT { @Autowired - protected HuggingfaceChatClient huggingfaceChatClient; + protected HuggingfaceChatModel huggingfaceChatModel; @Test void helloWorldCompletion() { @@ -46,7 +46,7 @@ public class ClientIT { [/INST] """; Prompt prompt = new Prompt(mistral7bInstruct); - ChatResponse chatResponse = huggingfaceChatClient.call(prompt); + ChatResponse chatResponse = huggingfaceChatModel.call(prompt); assertThat(chatResponse.getResult().getOutput().getContent()).isNotEmpty(); String expectedResponse = """ ```json diff --git a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatClient.java b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java similarity index 91% rename from models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatClient.java rename to models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java index 34e5e44ad..26637cdc3 100644 --- a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatClient.java +++ b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java @@ -17,10 +17,10 @@ package org.springframework.ai.minimax; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; @@ -38,30 +38,24 @@ import org.springframework.util.Assert; import org.springframework.util.CollectionUtils; import reactor.core.publisher.Flux; -import java.util.Base64; -import java.util.HashMap; -import java.util.HashSet; -import java.util.List; -import java.util.Map; -import java.util.Optional; -import java.util.Set; +import java.util.*; import java.util.concurrent.ConcurrentHashMap; /** - * {@link ChatClient} and {@link StreamingChatClient} implementation for - * {@literal MiniMax} backed by {@link MiniMaxApi}. + * {@link ChatModel} and {@link StreamingChatModel} implementation for {@literal MiniMax} + * backed by {@link MiniMaxApi}. * * @author Geng Rong - * @see ChatClient - * @see StreamingChatClient + * @see ChatModel + * @see StreamingChatModel * @see MiniMaxApi * @since 1.0.0 M1 */ -public class MiniMaxChatClient extends - AbstractFunctionCallSupport> - implements ChatClient, StreamingChatClient { +public class MiniMaxChatModel extends + AbstractFunctionCallSupport> + implements ChatModel, StreamingChatModel { - private static final Logger logger = LoggerFactory.getLogger(MiniMaxChatClient.class); + private static final Logger logger = LoggerFactory.getLogger(MiniMaxChatModel.class); /** * The default options used for the chat completion requests. @@ -79,35 +73,35 @@ public class MiniMaxChatClient extends private final MiniMaxApi miniMaxApi; /** - * Creates an instance of the MiniMaxChatClient. + * Creates an instance of the MiniMaxChatModel. * @param miniMaxApi The MiniMaxApi instance to be used for interacting with the * MiniMax Chat API. * @throws IllegalArgumentException if MiniMaxApi is null */ - public MiniMaxChatClient(MiniMaxApi miniMaxApi) { + public MiniMaxChatModel(MiniMaxApi miniMaxApi) { this(miniMaxApi, MiniMaxChatOptions.builder().withModel(MiniMaxApi.DEFAULT_CHAT_MODEL).withTemperature(0.7f).build()); } /** - * Initializes an instance of the MiniMaxChatClient. + * Initializes an instance of the MiniMaxChatModel. * @param miniMaxApi The MiniMaxApi instance to be used for interacting with the * MiniMax Chat API. - * @param options The MiniMaxChatOptions to configure the chat client. + * @param options The MiniMaxChatOptions to configure the chat model. */ - public MiniMaxChatClient(MiniMaxApi miniMaxApi, MiniMaxChatOptions options) { + public MiniMaxChatModel(MiniMaxApi miniMaxApi, MiniMaxChatOptions options) { this(miniMaxApi, options, null, RetryUtils.DEFAULT_RETRY_TEMPLATE); } /** - * Initializes a new instance of the MiniMaxChatClient. + * Initializes a new instance of the MiniMaxChatModel. * @param miniMaxApi The MiniMaxApi instance to be used for interacting with the * MiniMax Chat API. - * @param options The MiniMaxChatOptions to configure the chat client. + * @param options The MiniMaxChatOptions to configure the chat model. * @param functionCallbackContext The function callback context. * @param retryTemplate The retry template. */ - public MiniMaxChatClient(MiniMaxApi miniMaxApi, MiniMaxChatOptions options, + public MiniMaxChatModel(MiniMaxApi miniMaxApi, MiniMaxChatOptions options, FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate) { super(functionCallbackContext); Assert.notNull(miniMaxApi, "MiniMaxApi must not be null"); @@ -279,7 +273,7 @@ public class MiniMaxChatClient extends return request; } - private List getFunctionTools(Set functionNames) { + private List getFunctionTools(Set functionNames) { return this.resolveFunctionCallbacks(functionNames).stream().map(functionCallback -> { var function = new FunctionTool.Function(functionCallback.getDescription(), functionCallback.getName(), functionCallback.getInputTypeSchema()); @@ -358,4 +352,9 @@ public class MiniMaxChatClient extends && choice.finishReason() == ChatCompletionFinishReason.TOOL_CALLS; } + @Override + public ChatOptions getDefaultOptions() { + return MiniMaxChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatOptions.java b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatOptions.java index 1b46ca6a6..10ed0f51d 100644 --- a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatOptions.java +++ b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatOptions.java @@ -114,10 +114,10 @@ public class MiniMaxChatOptions implements FunctionCallingOptions, ChatOptions { private @JsonProperty("tool_choice") String toolChoice; /** - * MiniMax Tool Function Callbacks to register with the ChatClient. + * MiniMax Tool Function Callbacks to register with the ChatModel. * For Prompt Options the functionCallbacks are automatically enabled for the duration of the prompt execution. * For Default Options the functionCallbacks are registered but disabled by default. Use the enableFunctions to set the functions - * from the registry to be used by the ChatClient chat completion requests. + * from the registry to be used by the ChatModel chat completion requests. */ @NestedConfigurationProperty @JsonIgnore @@ -467,4 +467,22 @@ public class MiniMaxChatOptions implements FunctionCallingOptions, ChatOptions { return true; } + public static MiniMaxChatOptions fromOptions(MiniMaxChatOptions fromOptions) { + return builder().withModel(fromOptions.getModel()) + .withFrequencyPenalty(fromOptions.getFrequencyPenalty()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withN(fromOptions.getN()) + .withPresencePenalty(fromOptions.getPresencePenalty()) + .withResponseFormat(fromOptions.getResponseFormat()) + .withSeed(fromOptions.getSeed()) + .withStop(fromOptions.getStop()) + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withTools(fromOptions.getTools()) + .withToolChoice(fromOptions.getToolChoice()) + .withFunctionCallbacks(fromOptions.getFunctionCallbacks()) + .withFunctions(fromOptions.getFunctions()) + .build(); + } + } diff --git a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxEmbeddingClient.java b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxEmbeddingModel.java similarity index 87% rename from models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxEmbeddingClient.java rename to models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxEmbeddingModel.java index acede92b8..f66eba408 100644 --- a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxEmbeddingClient.java +++ b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxEmbeddingModel.java @@ -19,7 +19,7 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; import org.springframework.ai.document.MetadataMode; -import org.springframework.ai.embedding.AbstractEmbeddingClient; +import org.springframework.ai.embedding.AbstractEmbeddingModel; import org.springframework.ai.embedding.Embedding; import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingRequest; @@ -40,9 +40,9 @@ import java.util.List; * @author Geng Rong * @since 1.0.0 M1 */ -public class MiniMaxEmbeddingClient extends AbstractEmbeddingClient { +public class MiniMaxEmbeddingModel extends AbstractEmbeddingModel { - private static final Logger logger = LoggerFactory.getLogger(MiniMaxEmbeddingClient.class); + private static final Logger logger = LoggerFactory.getLogger(MiniMaxEmbeddingModel.class); private final MiniMaxEmbeddingOptions defaultOptions; @@ -53,43 +53,43 @@ public class MiniMaxEmbeddingClient extends AbstractEmbeddingClient { private final MetadataMode metadataMode; /** - * Constructor for the MiniMaxEmbeddingClient class. + * Constructor for the MiniMaxEmbeddingModel class. * @param miniMaxApi The MiniMaxApi instance to use for making API requests. */ - public MiniMaxEmbeddingClient(MiniMaxApi miniMaxApi) { + public MiniMaxEmbeddingModel(MiniMaxApi miniMaxApi) { this(miniMaxApi, MetadataMode.EMBED); } /** - * Initializes a new instance of the MiniMaxEmbeddingClient class. + * Initializes a new instance of the MiniMaxEmbeddingModel class. * @param miniMaxApi The MiniMaxApi instance to use for making API requests. * @param metadataMode The mode for generating metadata. */ - public MiniMaxEmbeddingClient(MiniMaxApi miniMaxApi, MetadataMode metadataMode) { + public MiniMaxEmbeddingModel(MiniMaxApi miniMaxApi, MetadataMode metadataMode) { this(miniMaxApi, metadataMode, MiniMaxEmbeddingOptions.builder().withModel(MiniMaxApi.DEFAULT_EMBEDDING_MODEL).build(), RetryUtils.DEFAULT_RETRY_TEMPLATE); } /** - * Initializes a new instance of the MiniMaxEmbeddingClient class. + * Initializes a new instance of the MiniMaxEmbeddingModel class. * @param miniMaxApi The MiniMaxApi instance to use for making API requests. * @param metadataMode The mode for generating metadata. * @param miniMaxEmbeddingOptions The options for MiniMax embedding. */ - public MiniMaxEmbeddingClient(MiniMaxApi miniMaxApi, MetadataMode metadataMode, + public MiniMaxEmbeddingModel(MiniMaxApi miniMaxApi, MetadataMode metadataMode, MiniMaxEmbeddingOptions miniMaxEmbeddingOptions) { this(miniMaxApi, metadataMode, miniMaxEmbeddingOptions, RetryUtils.DEFAULT_RETRY_TEMPLATE); } /** - * Initializes a new instance of the MiniMaxEmbeddingClient class. + * Initializes a new instance of the MiniMaxEmbeddingModel class. * @param miniMaxApi - The MiniMaxApi instance to use for making API requests. * @param metadataMode - The mode for generating metadata. * @param options - The options for MiniMax embedding. * @param retryTemplate - The RetryTemplate for retrying failed API requests. */ - public MiniMaxEmbeddingClient(MiniMaxApi miniMaxApi, MetadataMode metadataMode, MiniMaxEmbeddingOptions options, + public MiniMaxEmbeddingModel(MiniMaxApi miniMaxApi, MetadataMode metadataMode, MiniMaxEmbeddingOptions options, RetryTemplate retryTemplate) { Assert.notNull(miniMaxApi, "MiniMaxApi must not be null"); Assert.notNull(metadataMode, "metadataMode must not be null"); diff --git a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/api/MiniMaxApi.java b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/api/MiniMaxApi.java index e75d4892a..5b728b641 100644 --- a/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/api/MiniMaxApi.java +++ b/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/api/MiniMaxApi.java @@ -19,6 +19,8 @@ import com.fasterxml.jackson.annotation.JsonInclude; import com.fasterxml.jackson.annotation.JsonInclude.Include; import com.fasterxml.jackson.annotation.JsonProperty; import com.fasterxml.jackson.annotation.JsonValue; + +import org.springframework.ai.model.ModelDescription; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.retry.RetryUtils; import org.springframework.boot.context.properties.bind.ConstructorBinding; @@ -111,7 +113,7 @@ public class MiniMaxApi { * MiniMax Chat Completion Models: * MiniMax Model. */ - public enum ChatModel { + public enum ChatModel implements ModelDescription { ABAB_6_Chat("abab6-chat"), ABAB_5_5_Chat("abab5.5-chat"), ABAB_5_5_S_Chat("abab5.5s-chat"); @@ -125,6 +127,11 @@ public class MiniMaxApi { public String getValue() { return value; } + + @Override + public String getModelName() { + return this.value; + } } /** diff --git a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/ChatCompletionRequestTests.java b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/ChatCompletionRequestTests.java index 9adf803a4..3232836c5 100644 --- a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/ChatCompletionRequestTests.java +++ b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/ChatCompletionRequestTests.java @@ -33,7 +33,7 @@ public class ChatCompletionRequestTests { @Test public void createRequestWithChatOptions() { - var client = new MiniMaxChatClient(new MiniMaxApi("TEST"), + var client = new MiniMaxChatModel(new MiniMaxApi("TEST"), MiniMaxChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6f).build()); var request = client.createRequest(new Prompt("Test message content"), false); @@ -59,7 +59,7 @@ public class ChatCompletionRequestTests { final String TOOL_FUNCTION_NAME = "CurrentWeather"; - var client = new MiniMaxChatClient(new MiniMaxApi("TEST"), + var client = new MiniMaxChatModel(new MiniMaxApi("TEST"), MiniMaxChatOptions.builder().withModel("DEFAULT_MODEL").build()); var request = client.createRequest(new Prompt("Test message content", @@ -89,7 +89,7 @@ public class ChatCompletionRequestTests { final String TOOL_FUNCTION_NAME = "CurrentWeather"; - var client = new MiniMaxChatClient(new MiniMaxApi("TEST"), + var client = new MiniMaxChatModel(new MiniMaxApi("TEST"), MiniMaxChatOptions.builder() .withModel("DEFAULT_MODEL") .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) diff --git a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/MiniMaxTestConfiguration.java b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/MiniMaxTestConfiguration.java index f544b4896..8a7914da9 100644 --- a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/MiniMaxTestConfiguration.java +++ b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/MiniMaxTestConfiguration.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.minimax; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.minimax.api.MiniMaxApi; import org.springframework.boot.SpringBootConfiguration; import org.springframework.context.annotation.Bean; @@ -42,13 +42,13 @@ public class MiniMaxTestConfiguration { } @Bean - public MiniMaxChatClient miniMaxChatClient(MiniMaxApi api) { - return new MiniMaxChatClient(api); + public MiniMaxChatModel miniMaxChatModel(MiniMaxApi api) { + return new MiniMaxChatModel(api); } @Bean - public EmbeddingClient miniMaxEmbeddingClient(MiniMaxApi api) { - return new MiniMaxEmbeddingClient(api); + public EmbeddingModel miniMaxEmbeddingModel(MiniMaxApi api) { + return new MiniMaxEmbeddingModel(api); } } diff --git a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java index 46d62ae3c..157f7d063 100644 --- a/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java +++ b/models/spring-ai-minimax/src/test/java/org/springframework/ai/minimax/api/MiniMaxRetryTests.java @@ -22,9 +22,9 @@ import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.document.MetadataMode; -import org.springframework.ai.minimax.MiniMaxChatClient; +import org.springframework.ai.minimax.MiniMaxChatModel; import org.springframework.ai.minimax.MiniMaxChatOptions; -import org.springframework.ai.minimax.MiniMaxEmbeddingClient; +import org.springframework.ai.minimax.MiniMaxEmbeddingModel; import org.springframework.ai.minimax.MiniMaxEmbeddingOptions; import org.springframework.ai.minimax.api.MiniMaxApi.ChatCompletion; import org.springframework.ai.minimax.api.MiniMaxApi.ChatCompletionChunk; @@ -83,9 +83,9 @@ public class MiniMaxRetryTests { private @Mock MiniMaxApi miniMaxApi; - private MiniMaxChatClient chatClient; + private MiniMaxChatModel chatModel; - private MiniMaxEmbeddingClient embeddingClient; + private MiniMaxEmbeddingModel embeddingModel; @BeforeEach public void beforeEach() { @@ -93,8 +93,8 @@ public class MiniMaxRetryTests { retryListener = new TestRetryListener(); retryTemplate.registerListener(retryListener); - chatClient = new MiniMaxChatClient(miniMaxApi, MiniMaxChatOptions.builder().build(), null, retryTemplate); - embeddingClient = new MiniMaxEmbeddingClient(miniMaxApi, MetadataMode.EMBED, + chatModel = new MiniMaxChatModel(miniMaxApi, MiniMaxChatOptions.builder().build(), null, retryTemplate); + embeddingModel = new MiniMaxEmbeddingModel(miniMaxApi, MetadataMode.EMBED, MiniMaxEmbeddingOptions.builder().build(), retryTemplate); } @@ -111,7 +111,7 @@ public class MiniMaxRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(ResponseEntity.of(Optional.of(expectedChatCompletion))); - var result = chatClient.call(new Prompt("text")); + var result = chatModel.call(new Prompt("text")); assertThat(result).isNotNull(); assertThat(result.getResult().getOutput().getContent()).isSameAs("Response"); @@ -123,7 +123,7 @@ public class MiniMaxRetryTests { public void miniMaxChatNonTransientError() { when(miniMaxApi.chatCompletionEntity(isA(ChatCompletionRequest.class))) .thenThrow(new RuntimeException("Non Transient Error")); - assertThrows(RuntimeException.class, () -> chatClient.call(new Prompt("text"))); + assertThrows(RuntimeException.class, () -> chatModel.call(new Prompt("text"))); } @Test @@ -139,7 +139,7 @@ public class MiniMaxRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(Flux.just(expectedChatCompletion)); - var result = chatClient.stream(new Prompt("text")); + var result = chatModel.stream(new Prompt("text")); assertThat(result).isNotNull(); assertThat(result.collectList().block().get(0).getResult().getOutput().getContent()).isSameAs("Response"); @@ -151,7 +151,7 @@ public class MiniMaxRetryTests { public void miniMaxChatStreamNonTransientError() { when(miniMaxApi.chatCompletionStream(isA(ChatCompletionRequest.class))) .thenThrow(new RuntimeException("Non Transient Error")); - assertThrows(RuntimeException.class, () -> chatClient.stream(new Prompt("text"))); + assertThrows(RuntimeException.class, () -> chatModel.stream(new Prompt("text"))); } @Test @@ -164,7 +164,7 @@ public class MiniMaxRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(ResponseEntity.of(Optional.of(expectedEmbeddings))); - var result = embeddingClient + var result = embeddingModel .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), null)); assertThat(result).isNotNull(); @@ -177,7 +177,7 @@ public class MiniMaxRetryTests { public void miniMaxEmbeddingNonTransientError() { when(miniMaxApi.embeddings(isA(EmbeddingRequest.class))) .thenThrow(new RuntimeException("Non Transient Error")); - assertThrows(RuntimeException.class, () -> embeddingClient + assertThrows(RuntimeException.class, () -> embeddingModel .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), null))); } diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatClient.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java similarity index 95% rename from models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatClient.java rename to models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java index 17b7f8d93..a4a8a3313 100644 --- a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatClient.java +++ b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java @@ -17,10 +17,10 @@ package org.springframework.ai.mistralai; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; @@ -55,9 +55,9 @@ import java.util.concurrent.ConcurrentHashMap; * @author Grogdunn * @since 0.8.1 */ -public class MistralAiChatClient extends +public class MistralAiChatModel extends AbstractFunctionCallSupport> - implements ChatClient, StreamingChatClient { + implements ChatModel, StreamingChatModel { private final Logger log = LoggerFactory.getLogger(getClass()); @@ -73,7 +73,7 @@ public class MistralAiChatClient extends private final RetryTemplate retryTemplate; - public MistralAiChatClient(MistralAiApi mistralAiApi) { + public MistralAiChatModel(MistralAiApi mistralAiApi) { this(mistralAiApi, MistralAiChatOptions.builder() .withTemperature(0.7f) @@ -83,11 +83,11 @@ public class MistralAiChatClient extends .build()); } - public MistralAiChatClient(MistralAiApi mistralAiApi, MistralAiChatOptions options) { + public MistralAiChatModel(MistralAiApi mistralAiApi, MistralAiChatOptions options) { this(mistralAiApi, options, null, RetryUtils.DEFAULT_RETRY_TEMPLATE); } - public MistralAiChatClient(MistralAiApi mistralAiApi, MistralAiChatOptions options, + public MistralAiChatModel(MistralAiApi mistralAiApi, MistralAiChatOptions options, FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate) { super(functionCallbackContext); Assert.notNull(mistralAiApi, "MistralAiApi must not be null"); @@ -324,4 +324,9 @@ public class MistralAiChatClient extends return !CollectionUtils.isEmpty(choices.get(0).message().toolCalls()); } + @Override + public ChatOptions getDefaultOptions() { + return MistralAiChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatOptions.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatOptions.java index 86c0bda36..aa88700e3 100644 --- a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatOptions.java +++ b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatOptions.java @@ -101,11 +101,11 @@ public class MistralAiChatOptions implements FunctionCallingOptions, ChatOptions private @JsonProperty("tool_choice") ToolChoice toolChoice; /** - * MistralAI Tool Function Callbacks to register with the ChatClient. For Prompt + * MistralAI Tool Function Callbacks to register with the ChatModel. For Prompt * Options the functionCallbacks are automatically enabled for the duration of the * prompt execution. For Default Options the functionCallbacks are registered but * disabled by default. Use the enableFunctions to set the functions from the registry - * to be used by the ChatClient chat completion requests. + * to be used by the ChatModel chat completion requests. */ @NestedConfigurationProperty @JsonIgnore @@ -139,7 +139,7 @@ public class MistralAiChatOptions implements FunctionCallingOptions, ChatOptions return this; } - public Builder withMaxToken(Integer maxTokens) { + public Builder withMaxTokens(Integer maxTokens) { this.options.setMaxTokens(maxTokens); return this; } @@ -309,4 +309,19 @@ public class MistralAiChatOptions implements FunctionCallingOptions, ChatOptions this.functions = functions; } + public static MistralAiChatOptions fromOptions(MistralAiChatOptions fromOptions) { + return builder().withModel(fromOptions.getModel()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withSafePrompt(fromOptions.getSafePrompt()) + .withRandomSeed(fromOptions.getRandomSeed()) + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withResponseFormat(fromOptions.getResponseFormat()) + .withTools(fromOptions.getTools()) + .withToolChoice(fromOptions.getToolChoice()) + .withFunctionCallbacks(fromOptions.getFunctionCallbacks()) + .withFunctions(fromOptions.getFunctions()) + .build(); + } + } diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiEmbeddingClient.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiEmbeddingModel.java similarity index 89% rename from models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiEmbeddingClient.java rename to models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiEmbeddingModel.java index e42908c34..78851d804 100644 --- a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiEmbeddingClient.java +++ b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiEmbeddingModel.java @@ -22,7 +22,7 @@ import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; import org.springframework.ai.document.MetadataMode; -import org.springframework.ai.embedding.AbstractEmbeddingClient; +import org.springframework.ai.embedding.AbstractEmbeddingModel; import org.springframework.ai.embedding.Embedding; import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingRequest; @@ -38,7 +38,7 @@ import org.springframework.util.Assert; * @author Ricken Bazolo * @since 0.8.1 */ -public class MistralAiEmbeddingClient extends AbstractEmbeddingClient { +public class MistralAiEmbeddingModel extends AbstractEmbeddingModel { private final Logger log = LoggerFactory.getLogger(getClass()); @@ -50,21 +50,21 @@ public class MistralAiEmbeddingClient extends AbstractEmbeddingClient { private final RetryTemplate retryTemplate; - public MistralAiEmbeddingClient(MistralAiApi mistralAiApi) { + public MistralAiEmbeddingModel(MistralAiApi mistralAiApi) { this(mistralAiApi, MetadataMode.EMBED); } - public MistralAiEmbeddingClient(MistralAiApi mistralAiApi, MetadataMode metadataMode) { + public MistralAiEmbeddingModel(MistralAiApi mistralAiApi, MetadataMode metadataMode) { this(mistralAiApi, metadataMode, MistralAiEmbeddingOptions.builder().withModel(MistralAiApi.EmbeddingModel.EMBED.getValue()).build(), RetryUtils.DEFAULT_RETRY_TEMPLATE); } - public MistralAiEmbeddingClient(MistralAiApi mistralAiApi, MistralAiEmbeddingOptions options) { + public MistralAiEmbeddingModel(MistralAiApi mistralAiApi, MistralAiEmbeddingOptions options) { this(mistralAiApi, MetadataMode.EMBED, options, RetryUtils.DEFAULT_RETRY_TEMPLATE); } - public MistralAiEmbeddingClient(MistralAiApi mistralAiApi, MetadataMode metadataMode, + public MistralAiEmbeddingModel(MistralAiApi mistralAiApi, MetadataMode metadataMode, MistralAiEmbeddingOptions options, RetryTemplate retryTemplate) { Assert.notNull(mistralAiApi, "MistralAiApi must not be null"); Assert.notNull(metadataMode, "metadataMode must not be null"); diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/api/MistralAiApi.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/api/MistralAiApi.java index 16f2465c5..b2d5230eb 100644 --- a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/api/MistralAiApi.java +++ b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/api/MistralAiApi.java @@ -27,6 +27,7 @@ import com.fasterxml.jackson.annotation.JsonProperty; import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; +import org.springframework.ai.model.ModelDescription; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.retry.RetryUtils; import org.springframework.boot.context.properties.bind.ConstructorBinding; @@ -706,7 +707,7 @@ public class MistralAiApi { *
  • LARGE - mistral-large-latest (aka mistral-large-2402)
  • * */ - public enum ChatModel { + public enum ChatModel implements ModelDescription { // @formatter:off TINY("open-mistral-7b"), @@ -726,6 +727,11 @@ public class MistralAiApi { return this.value; } + @Override + public String getModelName() { + return this.value; + } + } /** diff --git a/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiChatCompletionRequestTest.java b/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiChatCompletionRequestTest.java index 57135164b..b83fdcd42 100644 --- a/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiChatCompletionRequestTest.java +++ b/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiChatCompletionRequestTest.java @@ -32,12 +32,12 @@ import static org.assertj.core.api.Assertions.assertThat; @EnabledIfEnvironmentVariable(named = "MISTRAL_AI_API_KEY", matches = ".+") public class MistralAiChatCompletionRequestTest { - MistralAiChatClient chatClient = new MistralAiChatClient(new MistralAiApi("test")); + MistralAiChatModel chatModel = new MistralAiChatModel(new MistralAiApi("test")); @Test void chatCompletionDefaultRequestTest() { - var request = chatClient.createRequest(new Prompt("test content"), false); + var request = chatModel.createRequest(new Prompt("test content"), false); assertThat(request.messages()).hasSize(1); assertThat(request.topP()).isEqualTo(1); @@ -52,7 +52,7 @@ public class MistralAiChatCompletionRequestTest { var options = MistralAiChatOptions.builder().withTemperature(0.5f).withTopP(0.8f).build(); - var request = chatClient.createRequest(new Prompt("test content", options), true); + var request = chatModel.createRequest(new Prompt("test content", options), true); assertThat(request.messages().size()).isEqualTo(1); assertThat(request.topP()).isEqualTo(0.8f); diff --git a/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiChatClientIT.java b/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiChatModelIT.java similarity index 91% rename from models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiChatClientIT.java rename to models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiChatModelIT.java index 87e6acbdb..069c92f80 100644 --- a/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiChatClientIT.java +++ b/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiChatModelIT.java @@ -27,10 +27,10 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import reactor.core.publisher.Flux; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.UserMessage; @@ -56,15 +56,15 @@ import static org.assertj.core.api.Assertions.assertThat; */ @SpringBootTest(classes = MistralAiTestConfiguration.class) @EnabledIfEnvironmentVariable(named = "MISTRAL_AI_API_KEY", matches = ".+") -class MistralAiChatClientIT { +class MistralAiChatModelIT { - private static final Logger logger = LoggerFactory.getLogger(MistralAiChatClientIT.class); + private static final Logger logger = LoggerFactory.getLogger(MistralAiChatModelIT.class); @Autowired - protected ChatClient chatClient; + protected ChatModel chatModel; @Autowired - protected StreamingChatClient streamingChatClient; + protected StreamingChatModel streamingChatModel; @Value("classpath:/prompts/system-message.st") private Resource systemResource; @@ -90,7 +90,7 @@ class MistralAiChatClientIT { // NOTE: Mistral expects the system message to be before the user message or will // fail with 400 error. Prompt prompt = new Prompt(List.of(systemMessage, userMessage)); - ChatResponse response = chatClient.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResults()).hasSize(1); assertThat(response.getResults().get(0).getOutput().getContent()).contains("Blackbeard"); } @@ -108,7 +108,7 @@ class MistralAiChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "ice cream flavors", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = this.chatClient.call(prompt).getResult(); + Generation generation = this.chatModel.call(prompt).getResult(); List list = outputConverter.convert(generation.getOutput().getContent()); assertThat(list).hasSize(5); @@ -126,7 +126,7 @@ class MistralAiChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = chatClient.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); Map result = outputConverter.convert(generation.getOutput().getContent()); assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); @@ -148,7 +148,7 @@ class MistralAiChatClientIT { """; PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = chatClient.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); ActorsFilmsRecord actorsFilms = outputConverter.convert(generation.getOutput().getContent()); logger.info("" + actorsFilms); @@ -169,7 +169,7 @@ class MistralAiChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - String generationTextFromStream = streamingChatClient.stream(prompt) + String generationTextFromStream = streamingChatModel.stream(prompt) .collectList() .block() .stream() @@ -201,7 +201,7 @@ class MistralAiChatClientIT { .build())) .build(); - ChatResponse response = chatClient.call(new Prompt(messages, promptOptions)); + ChatResponse response = chatModel.call(new Prompt(messages, promptOptions)); logger.info("Response: {}", response); @@ -224,7 +224,7 @@ class MistralAiChatClientIT { .build())) .build(); - Flux response = streamingChatClient.stream(new Prompt(messages, promptOptions)); + Flux response = streamingChatModel.stream(new Prompt(messages, promptOptions)); String content = response.collectList() .block() diff --git a/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiEmbeddingIT.java b/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiEmbeddingIT.java index ba472a097..cd8197e1f 100644 --- a/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiEmbeddingIT.java +++ b/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiEmbeddingIT.java @@ -31,25 +31,25 @@ import static org.assertj.core.api.Assertions.assertThat; class MistralAiEmbeddingIT { @Autowired - private MistralAiEmbeddingClient mistralAiEmbeddingClient; + private MistralAiEmbeddingModel mistralAiEmbeddingModel; @Test void defaultEmbedding() { - assertThat(mistralAiEmbeddingClient).isNotNull(); - var embeddingResponse = mistralAiEmbeddingClient.embedForResponse(List.of("Hello World")); + assertThat(mistralAiEmbeddingModel).isNotNull(); + var embeddingResponse = mistralAiEmbeddingModel.embedForResponse(List.of("Hello World")); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0)).isNotNull(); assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(1024); assertThat(embeddingResponse.getMetadata()).containsEntry("model", "mistral-embed"); assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 4); assertThat(embeddingResponse.getMetadata()).containsEntry("prompt-tokens", 4); - assertThat(mistralAiEmbeddingClient.dimensions()).isEqualTo(1024); + assertThat(mistralAiEmbeddingModel.dimensions()).isEqualTo(1024); } @Test void embeddingTest() { - assertThat(mistralAiEmbeddingClient).isNotNull(); - var embeddingResponse = mistralAiEmbeddingClient.call(new EmbeddingRequest( + assertThat(mistralAiEmbeddingModel).isNotNull(); + var embeddingResponse = mistralAiEmbeddingModel.call(new EmbeddingRequest( List.of("Hello World", "World is big"), MistralAiEmbeddingOptions.builder().withModel("mistral-embed").withEncodingFormat("float").build())); assertThat(embeddingResponse.getResults()).hasSize(2); @@ -58,7 +58,7 @@ class MistralAiEmbeddingIT { assertThat(embeddingResponse.getMetadata()).containsEntry("model", "mistral-embed"); assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 9); assertThat(embeddingResponse.getMetadata()).containsEntry("prompt-tokens", 9); - assertThat(mistralAiEmbeddingClient.dimensions()).isEqualTo(1024); + assertThat(mistralAiEmbeddingModel.dimensions()).isEqualTo(1024); } } diff --git a/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiRetryTests.java b/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiRetryTests.java index 1ca349d21..2c3f0b48d 100644 --- a/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiRetryTests.java +++ b/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiRetryTests.java @@ -82,9 +82,9 @@ public class MistralAiRetryTests { private @Mock MistralAiApi mistralAiApi; - private MistralAiChatClient chatClient; + private MistralAiChatModel chatModel; - private MistralAiEmbeddingClient embeddingClient; + private MistralAiEmbeddingModel embeddingModel; @BeforeEach public void beforeEach() { @@ -92,7 +92,7 @@ public class MistralAiRetryTests { retryListener = new TestRetryListener(); retryTemplate.registerListener(retryListener); - chatClient = new MistralAiChatClient(mistralAiApi, + chatModel = new MistralAiChatModel(mistralAiApi, MistralAiChatOptions.builder() .withTemperature(0.7f) .withTopP(1f) @@ -100,7 +100,7 @@ public class MistralAiRetryTests { .withModel(MistralAiApi.ChatModel.TINY.getValue()) .build(), null, retryTemplate); - embeddingClient = new MistralAiEmbeddingClient(mistralAiApi, MetadataMode.EMBED, + embeddingModel = new MistralAiEmbeddingModel(mistralAiApi, MetadataMode.EMBED, MistralAiEmbeddingOptions.builder().withModel(MistralAiApi.EmbeddingModel.EMBED.getValue()).build(), retryTemplate); } @@ -118,7 +118,7 @@ public class MistralAiRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(ResponseEntity.of(Optional.of(expectedChatCompletion))); - var result = chatClient.call(new Prompt("text")); + var result = chatModel.call(new Prompt("text")); assertThat(result).isNotNull(); assertThat(result.getResult().getOutput().getContent()).isSameAs("Response"); @@ -130,7 +130,7 @@ public class MistralAiRetryTests { public void mistralAiChatNonTransientError() { when(mistralAiApi.chatCompletionEntity(isA(ChatCompletionRequest.class))) .thenThrow(new RuntimeException("Non Transient Error")); - assertThrows(RuntimeException.class, () -> chatClient.call(new Prompt("text"))); + assertThrows(RuntimeException.class, () -> chatModel.call(new Prompt("text"))); } @Test @@ -146,7 +146,7 @@ public class MistralAiRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(Flux.just(expectedChatCompletion)); - var result = chatClient.stream(new Prompt("text")); + var result = chatModel.stream(new Prompt("text")); assertThat(result).isNotNull(); assertThat(result.collectList().block().get(0).getResult().getOutput().getContent()).isSameAs("Response"); @@ -158,7 +158,7 @@ public class MistralAiRetryTests { public void mistralAiChatStreamNonTransientError() { when(mistralAiApi.chatCompletionStream(isA(ChatCompletionRequest.class))) .thenThrow(new RuntimeException("Non Transient Error")); - assertThrows(RuntimeException.class, () -> chatClient.stream(new Prompt("text"))); + assertThrows(RuntimeException.class, () -> chatModel.stream(new Prompt("text"))); } @Test @@ -172,7 +172,7 @@ public class MistralAiRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(ResponseEntity.of(Optional.of(expectedEmbeddings))); - var result = embeddingClient + var result = embeddingModel .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), null)); assertThat(result).isNotNull(); @@ -185,7 +185,7 @@ public class MistralAiRetryTests { public void mistralAiEmbeddingNonTransientError() { when(mistralAiApi.embeddings(isA(EmbeddingRequest.class))) .thenThrow(new RuntimeException("Non Transient Error")); - assertThrows(RuntimeException.class, () -> embeddingClient + assertThrows(RuntimeException.class, () -> embeddingModel .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), null))); } diff --git a/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiTestConfiguration.java b/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiTestConfiguration.java index 7952571d6..64e608a92 100644 --- a/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiTestConfiguration.java +++ b/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiTestConfiguration.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.mistralai; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.mistralai.api.MistralAiApi; import org.springframework.boot.SpringBootConfiguration; import org.springframework.context.annotation.Bean; @@ -35,14 +35,14 @@ public class MistralAiTestConfiguration { } @Bean - public EmbeddingClient mistralAiEmbeddingClient(MistralAiApi api) { - return new MistralAiEmbeddingClient(api, + public EmbeddingModel mistralAiEmbeddingModel(MistralAiApi api) { + return new MistralAiEmbeddingModel(api, MistralAiEmbeddingOptions.builder().withModel(MistralAiApi.EmbeddingModel.EMBED.getValue()).build()); } @Bean - public MistralAiChatClient mistralAiChatClient(MistralAiApi mistralAiApi) { - return new MistralAiChatClient(mistralAiApi, + public MistralAiChatModel mistralAiChatModel(MistralAiApi mistralAiApi) { + return new MistralAiChatModel(mistralAiApi, MistralAiChatOptions.builder().withModel(MistralAiApi.ChatModel.MIXTRAL.getValue()).build()); } diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatClient.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java similarity index 91% rename from models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatClient.java rename to models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java index 273d98866..291cffdf8 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatClient.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java @@ -18,13 +18,13 @@ package org.springframework.ai.ollama; import java.util.Base64; import java.util.List; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.ollama.metadata.OllamaChatResponseMetadata; import reactor.core.publisher.Flux; -import org.springframework.ai.chat.ChatClient; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; @@ -39,7 +39,7 @@ import org.springframework.util.CollectionUtils; import org.springframework.util.StringUtils; /** - * {@link ChatClient} implementation for {@literal Ollama}. + * {@link ChatModel} implementation for {@literal Ollama}. * * Ollama allows developers to run large language models and generate embeddings locally. * It supports open-source models available on [Ollama AI @@ -52,7 +52,7 @@ import org.springframework.util.StringUtils; * @author Christian Tzolov * @since 0.8.0 */ -public class OllamaChatClient implements ChatClient, StreamingChatClient { +public class OllamaChatModel implements ChatModel, StreamingChatModel { /** * Low-level Ollama API library. @@ -64,11 +64,11 @@ public class OllamaChatClient implements ChatClient, StreamingChatClient { */ private OllamaOptions defaultOptions; - public OllamaChatClient(OllamaApi chatApi) { + public OllamaChatModel(OllamaApi chatApi) { this(chatApi, OllamaOptions.create().withModel(OllamaOptions.DEFAULT_MODEL)); } - public OllamaChatClient(OllamaApi chatApi, OllamaOptions defaultOptions) { + public OllamaChatModel(OllamaApi chatApi, OllamaOptions defaultOptions) { Assert.notNull(chatApi, "OllamaApi must not be null"); Assert.notNull(defaultOptions, "DefaultOptions must not be null"); this.chatApi = chatApi; @@ -79,7 +79,7 @@ public class OllamaChatClient implements ChatClient, StreamingChatClient { * @deprecated Use {@link OllamaOptions#setModel} instead. */ @Deprecated - public OllamaChatClient withModel(String model) { + public OllamaChatModel withModel(String model) { this.defaultOptions.setModel(model); return this; } @@ -88,7 +88,7 @@ public class OllamaChatClient implements ChatClient, StreamingChatClient { * @deprecated Use {@link OllamaOptions} constructor instead. */ @Deprecated - public OllamaChatClient withDefaultOptions(OllamaOptions options) { + public OllamaChatModel withDefaultOptions(OllamaOptions options) { this.defaultOptions = options; return this; } @@ -205,4 +205,9 @@ public class OllamaChatClient implements ChatClient, StreamingChatClient { } } + @Override + public ChatOptions getDefaultOptions() { + return OllamaOptions.fromOptions(this.defaultOptions); + } + } \ No newline at end of file diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaEmbeddingClient.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaEmbeddingModel.java similarity index 89% rename from models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaEmbeddingClient.java rename to models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaEmbeddingModel.java index 1748709aa..66be40c5c 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaEmbeddingClient.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaEmbeddingModel.java @@ -23,9 +23,9 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.AbstractEmbeddingClient; +import org.springframework.ai.embedding.AbstractEmbeddingModel; import org.springframework.ai.embedding.Embedding; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingResponse; import org.springframework.ai.model.ModelOptionsUtils; @@ -36,7 +36,7 @@ import org.springframework.util.Assert; import org.springframework.util.StringUtils; /** - * {@link EmbeddingClient} implementation for {@literal Ollama}. + * {@link EmbeddingModel} implementation for {@literal Ollama}. * * Ollama allows developers to run large language models and generate embeddings locally. * It supports open-source models available on [Ollama AI @@ -51,7 +51,7 @@ import org.springframework.util.StringUtils; * @author Christian Tzolov * @since 0.8.0 */ -public class OllamaEmbeddingClient extends AbstractEmbeddingClient { +public class OllamaEmbeddingModel extends AbstractEmbeddingModel { private final Logger logger = LoggerFactory.getLogger(getClass()); @@ -62,11 +62,11 @@ public class OllamaEmbeddingClient extends AbstractEmbeddingClient { */ private OllamaOptions defaultOptions = OllamaOptions.create().withModel(OllamaOptions.DEFAULT_MODEL); - public OllamaEmbeddingClient(OllamaApi ollamaApi) { + public OllamaEmbeddingModel(OllamaApi ollamaApi) { this.ollamaApi = ollamaApi; } - public OllamaEmbeddingClient(OllamaApi ollamaApi, OllamaOptions defaultOptions) { + public OllamaEmbeddingModel(OllamaApi ollamaApi, OllamaOptions defaultOptions) { this.ollamaApi = ollamaApi; this.defaultOptions = defaultOptions; } @@ -75,7 +75,7 @@ public class OllamaEmbeddingClient extends AbstractEmbeddingClient { * @deprecated Use {@link OllamaOptions#setModel} instead. */ @Deprecated - public OllamaEmbeddingClient withModel(String model) { + public OllamaEmbeddingModel withModel(String model) { this.defaultOptions.setModel(model); return this; } @@ -84,7 +84,7 @@ public class OllamaEmbeddingClient extends AbstractEmbeddingClient { * @deprecated Use {@link OllamaOptions} constructor instead. */ @Deprecated - public OllamaEmbeddingClient withDefaultOptions(OllamaOptions options) { + public OllamaEmbeddingModel withDefaultOptions(OllamaOptions options) { this.defaultOptions = options; return this; } diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaModel.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaModel.java index 73d41053c..449bab647 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaModel.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaModel.java @@ -15,13 +15,15 @@ */ package org.springframework.ai.ollama.api; +import org.springframework.ai.model.ModelDescription; + /** * Helper class for common Ollama models. * * @author Siarhei Blashuk * @since 0.8.1 */ -public enum OllamaModel { +public enum OllamaModel implements ModelDescription { /** * Llama 2 is a collection of language models ranging from 7B to 70B parameters. @@ -99,4 +101,9 @@ public enum OllamaModel { return this.id; } + @Override + public String getModelName() { + return this.id; + } + } diff --git a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaOptions.java b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaOptions.java index abc329e81..436631f6e 100644 --- a/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaOptions.java +++ b/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaOptions.java @@ -714,6 +714,43 @@ public class OllamaOptions implements ChatOptions, EmbeddingOptions { .collect(Collectors.toMap(Map.Entry::getKey, Map.Entry::getValue)); } + public static OllamaOptions fromOptions(OllamaOptions fromOptions) { + return new OllamaOptions() + .withModel(fromOptions.getModel()) + .withFormat(fromOptions.getFormat()) + .withKeepAlive(fromOptions.getKeepAlive()) + .withUseNUMA(fromOptions.getUseNUMA()) + .withNumCtx(fromOptions.getNumCtx()) + .withNumBatch(fromOptions.getNumBatch()) + .withNumGQA(fromOptions.getNumGQA()) + .withNumGPU(fromOptions.getNumGPU()) + .withMainGPU(fromOptions.getMainGPU()) + .withLowVRAM(fromOptions.getLowVRAM()) + .withF16KV(fromOptions.getF16KV()) + .withLogitsAll(fromOptions.getLogitsAll()) + .withVocabOnly(fromOptions.getVocabOnly()) + .withUseMMap(fromOptions.getUseMMap()) + .withUseMLock(fromOptions.getUseMLock()) + .withNumThread(fromOptions.getNumThread()) + .withNumKeep(fromOptions.getNumKeep()) + .withSeed(fromOptions.getSeed()) + .withNumPredict(fromOptions.getNumPredict()) + .withTopK(fromOptions.getTopK()) + .withTopP(fromOptions.getTopP()) + .withTfsZ(fromOptions.getTfsZ()) + .withTypicalP(fromOptions.getTypicalP()) + .withRepeatLastN(fromOptions.getRepeatLastN()) + .withTemperature(fromOptions.getTemperature()) + .withRepeatPenalty(fromOptions.getRepeatPenalty()) + .withPresencePenalty(fromOptions.getPresencePenalty()) + .withFrequencyPenalty(fromOptions.getFrequencyPenalty()) + .withMirostat(fromOptions.getMirostat()) + .withMirostatTau(fromOptions.getMirostatTau()) + .withMirostatEta(fromOptions.getMirostatEta()) + .withPenalizeNewline(fromOptions.getPenalizeNewline()) + .withStop(fromOptions.getStop()); + } + // @formatter:on diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatClientIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelIT.java similarity index 90% rename from models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatClientIT.java rename to models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelIT.java index 6f4254fa4..f1d974fc6 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatClientIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelIT.java @@ -56,11 +56,11 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @Testcontainers @Disabled("For manual smoke testing only.") -class OllamaChatClientIT { +class OllamaChatModelIT { private static String MODEL = "mistral"; - private static final Log logger = LogFactory.getLog(OllamaChatClientIT.class); + private static final Log logger = LogFactory.getLog(OllamaChatModelIT.class); @Container static OllamaContainer ollamaContainer = new OllamaContainer("ollama/ollama:0.1.32"); @@ -77,7 +77,7 @@ class OllamaChatClientIT { } @Autowired - private OllamaChatClient client; + private OllamaChatModel chatModel; @Test void roleTest() { @@ -95,13 +95,13 @@ class OllamaChatClientIT { Prompt prompt = new Prompt(List.of(userMessage, systemMessage), portableOptions); - ChatResponse response = client.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); // ollama specific options var ollamaOptions = new OllamaOptions().withLowVRAM(true); - response = client.call(new Prompt(List.of(userMessage, systemMessage), ollamaOptions)); + response = chatModel.call(new Prompt(List.of(userMessage, systemMessage), ollamaOptions)); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); } @@ -109,7 +109,7 @@ class OllamaChatClientIT { @Test void usageTest() { Prompt prompt = new Prompt("Tell me a joke"); - ChatResponse response = client.call(prompt); + ChatResponse response = chatModel.call(prompt); Usage usage = response.getMetadata().getUsage(); assertThat(usage).isNotNull(); @@ -131,7 +131,7 @@ class OllamaChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "ice cream flavors.", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = this.client.call(prompt).getResult(); + Generation generation = this.chatModel.call(prompt).getResult(); List list = outputConverter.convert(generation.getOutput().getContent()); assertThat(list).hasSize(5); @@ -151,7 +151,7 @@ class OllamaChatClientIT { Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); Map result = outputConverter.convert(generation.getOutput().getContent()); assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); @@ -173,7 +173,7 @@ class OllamaChatClientIT { """; PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); ActorsFilmsRecord actorsFilms = outputConverter.convert(generation.getOutput().getContent()); assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks"); @@ -194,7 +194,7 @@ class OllamaChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - String generationTextFromStream = client.stream(prompt) + String generationTextFromStream = chatModel.stream(prompt) .collectList() .block() .stream() @@ -219,8 +219,8 @@ class OllamaChatClientIT { } @Bean - public OllamaChatClient ollamaChat(OllamaApi ollamaApi) { - return new OllamaChatClient(ollamaApi, OllamaOptions.create().withModel(MODEL).withTemperature(0.9f)); + public OllamaChatModel ollamaChat(OllamaApi ollamaApi) { + return new OllamaChatModel(ollamaApi, OllamaOptions.create().withModel(MODEL).withTemperature(0.9f)); } } diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatClientMultimodalIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelMultimodalIT.java similarity index 88% rename from models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatClientMultimodalIT.java rename to models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelMultimodalIT.java index 8599cfe9e..4ef7b7d36 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatClientMultimodalIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelMultimodalIT.java @@ -44,11 +44,11 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @Testcontainers @Disabled("For manual smoke testing only.") -class OllamaChatClientMultimodalIT { +class OllamaChatModelMultimodalIT { private static String MODEL = "llava"; - private static final Log logger = LogFactory.getLog(OllamaChatClientIT.class); + private static final Log logger = LogFactory.getLog(OllamaChatModelIT.class); @Container static OllamaContainer ollamaContainer = new OllamaContainer("ollama/ollama:0.1.32"); @@ -65,7 +65,7 @@ class OllamaChatClientMultimodalIT { } @Autowired - private OllamaChatClient client; + private OllamaChatModel chatModel; @Test void multiModalityTest() throws IOException { @@ -75,7 +75,7 @@ class OllamaChatClientMultimodalIT { var userMessage = new UserMessage("Explain what do you see on this picture?", List.of(new Media(MimeTypeUtils.IMAGE_PNG, imageData))); - var response = client.call(new Prompt(List.of(userMessage))); + var response = chatModel.call(new Prompt(List.of(userMessage))); logger.info(response.getResult().getOutput().getContent()); assertThat(response.getResult().getOutput().getContent()).contains("bananas", "apple", "basket"); @@ -90,8 +90,8 @@ class OllamaChatClientMultimodalIT { } @Bean - public OllamaChatClient ollamaChat(OllamaApi ollamaApi) { - return new OllamaChatClient(ollamaApi, OllamaOptions.create().withModel(MODEL).withTemperature(0.9f)); + public OllamaChatModel ollamaChat(OllamaApi ollamaApi) { + return new OllamaChatModel(ollamaApi, OllamaOptions.create().withModel(MODEL).withTemperature(0.9f)); } } diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatRequestTests.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatRequestTests.java index f78b8f2fc..3ea38e6d0 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatRequestTests.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatRequestTests.java @@ -30,13 +30,13 @@ import static org.assertj.core.api.Assertions.assertThat; */ public class OllamaChatRequestTests { - OllamaChatClient client = new OllamaChatClient(new OllamaApi(), + OllamaChatModel chatModel = new OllamaChatModel(new OllamaApi(), new OllamaOptions().withModel("MODEL_NAME").withTopK(99).withTemperature(66.6f).withNumGPU(1)); @Test public void createRequestWithDefaultOptions() { - var request = client.ollamaChatRequest(new Prompt("Test message content"), false); + var request = chatModel.ollamaChatRequest(new Prompt("Test message content"), false); assertThat(request.messages()).hasSize(1); assertThat(request.stream()).isFalse(); @@ -54,7 +54,7 @@ public class OllamaChatRequestTests { // Runtime options should override the default options. OllamaOptions promptOptions = new OllamaOptions().withTemperature(0.8f).withTopP(0.5f).withNumGPU(2); - var request = client.ollamaChatRequest(new Prompt("Test message content", promptOptions), true); + var request = chatModel.ollamaChatRequest(new Prompt("Test message content", promptOptions), true); assertThat(request.messages()).hasSize(1); assertThat(request.stream()).isTrue(); @@ -79,7 +79,7 @@ public class OllamaChatRequestTests { .withTopP(0.6f) .build(); - var request = client.ollamaChatRequest(new Prompt("Test message content", portablePromptOptions), true); + var request = chatModel.ollamaChatRequest(new Prompt("Test message content", portablePromptOptions), true); assertThat(request.messages()).hasSize(1); assertThat(request.stream()).isTrue(); @@ -97,7 +97,7 @@ public class OllamaChatRequestTests { // Ollama runtime options. OllamaOptions promptOptions = new OllamaOptions().withModel("PROMPT_MODEL"); - var request = client.ollamaChatRequest(new Prompt("Test message content", promptOptions), true); + var request = chatModel.ollamaChatRequest(new Prompt("Test message content", promptOptions), true); assertThat(request.model()).isEqualTo("PROMPT_MODEL"); } @@ -105,17 +105,17 @@ public class OllamaChatRequestTests { @Test public void createRequestWithDefaultOptionsModelOverride() { - OllamaChatClient client2 = new OllamaChatClient(new OllamaApi(), + OllamaChatModel chatModel = new OllamaChatModel(new OllamaApi(), new OllamaOptions().withModel("DEFAULT_OPTIONS_MODEL")); - var request = client2.ollamaChatRequest(new Prompt("Test message content"), true); + var request = chatModel.ollamaChatRequest(new Prompt("Test message content"), true); assertThat(request.model()).isEqualTo("DEFAULT_OPTIONS_MODEL"); // Prompt options should override the default options. OllamaOptions promptOptions = new OllamaOptions().withModel("PROMPT_MODEL"); - request = client2.ollamaChatRequest(new Prompt("Test message content", promptOptions), true); + request = chatModel.ollamaChatRequest(new Prompt("Test message content", promptOptions), true); assertThat(request.model()).isEqualTo("PROMPT_MODEL"); } diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingClientIT.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelIT.java similarity index 84% rename from models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingClientIT.java rename to models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelIT.java index d903aab49..437fba302 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingClientIT.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingModelIT.java @@ -23,9 +23,9 @@ 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.testcontainers.containers.GenericContainer; import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; +import org.testcontainers.ollama.OllamaContainer; import org.springframework.ai.embedding.EmbeddingResponse; import org.springframework.ai.ollama.api.OllamaApi; @@ -34,14 +34,13 @@ import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.test.context.SpringBootTest; import org.springframework.context.annotation.Bean; -import org.testcontainers.ollama.OllamaContainer; import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @Disabled("For manual smoke testing only.") @Testcontainers -class OllamaEmbeddingClientIT { +class OllamaEmbeddingModelIT { private static final Log logger = LogFactory.getLog(OllamaApiIT.class); @@ -60,15 +59,15 @@ class OllamaEmbeddingClientIT { } @Autowired - private OllamaEmbeddingClient embeddingClient; + private OllamaEmbeddingModel embeddingModel; @Test void singleEmbedding() { - assertThat(embeddingClient).isNotNull(); - EmbeddingResponse embeddingResponse = embeddingClient.embedForResponse(List.of("Hello World")); + assertThat(embeddingModel).isNotNull(); + EmbeddingResponse embeddingResponse = embeddingModel.embedForResponse(List.of("Hello World")); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); - assertThat(embeddingClient.dimensions()).isEqualTo(3200); + assertThat(embeddingModel.dimensions()).isEqualTo(3200); } @SpringBootConfiguration @@ -80,8 +79,8 @@ class OllamaEmbeddingClientIT { } @Bean - public OllamaEmbeddingClient ollamaEmbedding(OllamaApi ollamaApi) { - return new OllamaEmbeddingClient(ollamaApi).withModel("orca-mini"); + public OllamaEmbeddingModel ollamaEmbedding(OllamaApi ollamaApi) { + return new OllamaEmbeddingModel(ollamaApi).withModel("orca-mini"); } } diff --git a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingRequestTests.java b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingRequestTests.java index c4987129c..82fcab7d7 100644 --- a/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingRequestTests.java +++ b/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaEmbeddingRequestTests.java @@ -28,13 +28,13 @@ import static org.assertj.core.api.Assertions.assertThat; */ public class OllamaEmbeddingRequestTests { - OllamaEmbeddingClient client = new OllamaEmbeddingClient(new OllamaApi()).withDefaultOptions( + OllamaEmbeddingModel chatModel = new OllamaEmbeddingModel(new OllamaApi()).withDefaultOptions( new OllamaOptions().withModel("DEFAULT_MODEL").withMainGPU(11).withUseMMap(true).withNumGPU(1)); @Test public void ollamaEmbeddingRequestDefaultOptions() { - var request = client.ollamaEmbeddingRequest("Hello", null); + var request = chatModel.ollamaEmbeddingRequest("Hello", null); assertThat(request.model()).isEqualTo("DEFAULT_MODEL"); assertThat(request.options().get("num_gpu")).isEqualTo(1); @@ -51,7 +51,7 @@ public class OllamaEmbeddingRequestTests { .withUseMMap(true) .withNumGPU(2); - var request = client.ollamaEmbeddingRequest("Hello", promptOptions); + var request = chatModel.ollamaEmbeddingRequest("Hello", promptOptions); assertThat(request.model()).isEqualTo("PROMPT_MODEL"); assertThat(request.options().get("num_gpu")).isEqualTo(2); diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioSpeechClient.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioSpeechModel.java similarity index 92% rename from models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioSpeechClient.java rename to models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioSpeechModel.java index 49fb21694..027ee50be 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioSpeechClient.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioSpeechModel.java @@ -24,10 +24,10 @@ import org.springframework.ai.openai.api.OpenAiAudioApi; import org.springframework.ai.openai.api.OpenAiAudioApi.SpeechRequest.AudioResponseFormat; import org.springframework.ai.openai.api.common.OpenAiApiException; import org.springframework.ai.openai.audio.speech.Speech; -import org.springframework.ai.openai.audio.speech.SpeechClient; +import org.springframework.ai.openai.audio.speech.SpeechModel; import org.springframework.ai.openai.audio.speech.SpeechPrompt; import org.springframework.ai.openai.audio.speech.SpeechResponse; -import org.springframework.ai.openai.audio.speech.StreamingSpeechClient; +import org.springframework.ai.openai.audio.speech.StreamingSpeechModel; import org.springframework.ai.openai.metadata.audio.OpenAiAudioSpeechResponseMetadata; import org.springframework.ai.openai.metadata.support.OpenAiResponseHeaderExtractor; import org.springframework.http.ResponseEntity; @@ -44,7 +44,7 @@ import java.time.Duration; * @see OpenAiAudioApi * @since 1.0.0-M1 */ -public class OpenAiAudioSpeechClient implements SpeechClient, StreamingSpeechClient { +public class OpenAiAudioSpeechModel implements SpeechModel, StreamingSpeechModel { private final Logger logger = LoggerFactory.getLogger(getClass()); @@ -61,12 +61,12 @@ public class OpenAiAudioSpeechClient implements SpeechClient, StreamingSpeechCli private final OpenAiAudioApi audioApi; /** - * Initializes a new instance of the OpenAiAudioSpeechClient class with the provided + * Initializes a new instance of the OpenAiAudioSpeechModel class with the provided * OpenAiAudioApi. It uses the model tts-1, response format mp3, voice alloy, and the * default speed of 1.0. * @param audioApi The OpenAiAudioApi to use for speech synthesis. */ - public OpenAiAudioSpeechClient(OpenAiAudioApi audioApi) { + public OpenAiAudioSpeechModel(OpenAiAudioApi audioApi) { this(audioApi, OpenAiAudioSpeechOptions.builder() .withModel(OpenAiAudioApi.TtsModel.TTS_1.getValue()) @@ -77,13 +77,13 @@ public class OpenAiAudioSpeechClient implements SpeechClient, StreamingSpeechCli } /** - * Initializes a new instance of the OpenAiAudioSpeechClient class with the provided + * Initializes a new instance of the OpenAiAudioSpeechModel class with the provided * OpenAiAudioApi and options. * @param audioApi The OpenAiAudioApi to use for speech synthesis. * @param options The OpenAiAudioSpeechOptions containing the speech synthesis * options. */ - public OpenAiAudioSpeechClient(OpenAiAudioApi audioApi, OpenAiAudioSpeechOptions options) { + public OpenAiAudioSpeechModel(OpenAiAudioApi audioApi, OpenAiAudioSpeechOptions options) { Assert.notNull(audioApi, "OpenAiAudioApi must not be null"); Assert.notNull(options, "OpenAiSpeechOptions must not be null"); this.audioApi = audioApi; diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioTranscriptionClient.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioTranscriptionModel.java similarity index 92% rename from models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioTranscriptionClient.java rename to models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioTranscriptionModel.java index 71e8bb8fc..ed5d9709d 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioTranscriptionClient.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiAudioTranscriptionModel.java @@ -35,7 +35,7 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.chat.metadata.RateLimit; -import org.springframework.ai.model.ModelClient; +import org.springframework.ai.model.Model; import org.springframework.ai.openai.api.OpenAiAudioApi; import org.springframework.ai.openai.api.OpenAiAudioApi.StructuredResponse; import org.springframework.ai.openai.audio.transcription.AudioTranscription; @@ -59,8 +59,7 @@ import org.springframework.util.Assert; * @see OpenAiAudioApi * @since 0.8.1 */ -public class OpenAiAudioTranscriptionClient - implements ModelClient { +public class OpenAiAudioTranscriptionModel implements Model { private final Logger logger = LoggerFactory.getLogger(getClass()); @@ -71,11 +70,11 @@ public class OpenAiAudioTranscriptionClient private final OpenAiAudioApi audioApi; /** - * OpenAiAudioTranscriptionClient is a client class used to interact with the OpenAI + * OpenAiAudioTranscriptionModel is a client class used to interact with the OpenAI * Audio Transcription API. * @param audioApi The OpenAiAudioApi instance to be used for making API calls. */ - public OpenAiAudioTranscriptionClient(OpenAiAudioApi audioApi) { + public OpenAiAudioTranscriptionModel(OpenAiAudioApi audioApi) { this(audioApi, OpenAiAudioTranscriptionOptions.builder() .withModel(OpenAiAudioApi.WhisperModel.WHISPER_1.getValue()) @@ -86,25 +85,25 @@ public class OpenAiAudioTranscriptionClient } /** - * OpenAiAudioTranscriptionClient is a client class used to interact with the OpenAI + * OpenAiAudioTranscriptionModel is a client class used to interact with the OpenAI * Audio Transcription API. * @param audioApi The OpenAiAudioApi instance to be used for making API calls. * @param options The OpenAiAudioTranscriptionOptions instance for configuring the * audio transcription. */ - public OpenAiAudioTranscriptionClient(OpenAiAudioApi audioApi, OpenAiAudioTranscriptionOptions options) { + public OpenAiAudioTranscriptionModel(OpenAiAudioApi audioApi, OpenAiAudioTranscriptionOptions options) { this(audioApi, options, RetryUtils.DEFAULT_RETRY_TEMPLATE); } /** - * OpenAiAudioTranscriptionClient is a client class used to interact with the OpenAI + * OpenAiAudioTranscriptionModel is a client class used to interact with the OpenAI * Audio Transcription API. * @param audioApi The OpenAiAudioApi instance to be used for making API calls. * @param options The OpenAiAudioTranscriptionOptions instance for configuring the * audio transcription. * @param retryTemplate The RetryTemplate instance for retrying failed API calls. */ - public OpenAiAudioTranscriptionClient(OpenAiAudioApi audioApi, OpenAiAudioTranscriptionOptions options, + public OpenAiAudioTranscriptionModel(OpenAiAudioApi audioApi, OpenAiAudioTranscriptionOptions options, RetryTemplate retryTemplate) { Assert.notNull(audioApi, "OpenAiAudioApi must not be null"); Assert.notNull(options, "OpenAiTranscriptionOptions must not be null"); diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatClient.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java similarity index 93% rename from models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatClient.java rename to models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java index 6ec6904db..266c475a0 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatClient.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java @@ -17,10 +17,10 @@ package org.springframework.ai.openai; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; import org.springframework.ai.chat.metadata.RateLimit; import org.springframework.ai.chat.prompt.ChatOptions; @@ -58,7 +58,7 @@ import java.util.Set; import java.util.concurrent.ConcurrentHashMap; /** - * {@link ChatClient} and {@link StreamingChatClient} implementation for {@literal OpenAI} + * {@link ChatModel} and {@link StreamingChatModel} implementation for {@literal OpenAI} * backed by {@link OpenAiApi}. * * @author Mark Pollack @@ -68,15 +68,15 @@ import java.util.concurrent.ConcurrentHashMap; * @author Josh Long * @author Jemin Huh * @author Grogdunn - * @see ChatClient - * @see StreamingChatClient + * @see ChatModel + * @see StreamingChatModel * @see OpenAiApi */ -public class OpenAiChatClient extends +public class OpenAiChatModel extends AbstractFunctionCallSupport> - implements ChatClient, StreamingChatClient { + implements ChatModel, StreamingChatModel { - private static final Logger logger = LoggerFactory.getLogger(OpenAiChatClient.class); + private static final Logger logger = LoggerFactory.getLogger(OpenAiChatModel.class); /** * The default options used for the chat completion requests. @@ -94,35 +94,35 @@ public class OpenAiChatClient extends private final OpenAiApi openAiApi; /** - * Creates an instance of the OpenAiChatClient. + * Creates an instance of the OpenAiChatModel. * @param openAiApi The OpenAiApi instance to be used for interacting with the OpenAI * Chat API. * @throws IllegalArgumentException if openAiApi is null */ - public OpenAiChatClient(OpenAiApi openAiApi) { + public OpenAiChatModel(OpenAiApi openAiApi) { this(openAiApi, OpenAiChatOptions.builder().withModel(OpenAiApi.DEFAULT_CHAT_MODEL).withTemperature(0.7f).build()); } /** - * Initializes an instance of the OpenAiChatClient. + * Initializes an instance of the OpenAiChatModel. * @param openAiApi The OpenAiApi instance to be used for interacting with the OpenAI * Chat API. - * @param options The OpenAiChatOptions to configure the chat client. + * @param options The OpenAiChatOptions to configure the chat model. */ - public OpenAiChatClient(OpenAiApi openAiApi, OpenAiChatOptions options) { + public OpenAiChatModel(OpenAiApi openAiApi, OpenAiChatOptions options) { this(openAiApi, options, null, RetryUtils.DEFAULT_RETRY_TEMPLATE); } /** - * Initializes a new instance of the OpenAiChatClient. + * Initializes a new instance of the OpenAiChatModel. * @param openAiApi The OpenAiApi instance to be used for interacting with the OpenAI * Chat API. - * @param options The OpenAiChatOptions to configure the chat client. + * @param options The OpenAiChatOptions to configure the chat model. * @param functionCallbackContext The function callback context. * @param retryTemplate The retry template. */ - public OpenAiChatClient(OpenAiApi openAiApi, OpenAiChatOptions options, + public OpenAiChatModel(OpenAiApi openAiApi, OpenAiChatOptions options, FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate) { super(functionCallbackContext); Assert.notNull(openAiApi, "OpenAiApi must not be null"); @@ -394,4 +394,9 @@ public class OpenAiChatClient extends && choice.finishReason() == ChatCompletionFinishReason.TOOL_CALLS; } + @Override + public ChatOptions getDefaultOptions() { + return OpenAiChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java index d45d4db18..09f86baac 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java @@ -134,10 +134,10 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions { private @JsonProperty("user") String user; /** - * OpenAI Tool Function Callbacks to register with the ChatClient. + * OpenAI Tool Function Callbacks to register with the ChatModel. * For Prompt Options the functionCallbacks are automatically enabled for the duration of the prompt execution. * For Default Options the functionCallbacks are registered but disabled by default. Use the enableFunctions to set the functions - * from the registry to be used by the ChatClient chat completion requests. + * from the registry to be used by the ChatModel chat completion requests. */ @NestedConfigurationProperty @JsonIgnore @@ -567,4 +567,27 @@ public class OpenAiChatOptions implements FunctionCallingOptions, ChatOptions { throw new UnsupportedOperationException("Unimplemented method 'setTopK'"); } + public static OpenAiChatOptions fromOptions(OpenAiChatOptions fromOptions) { + return OpenAiChatOptions.builder() + .withModel(fromOptions.getModel()) + .withFrequencyPenalty(fromOptions.getFrequencyPenalty()) + .withLogitBias(fromOptions.getLogitBias()) + .withLogprobs(fromOptions.getLogprobs()) + .withTopLogprobs(fromOptions.getTopLogprobs()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withN(fromOptions.getN()) + .withPresencePenalty(fromOptions.getPresencePenalty()) + .withResponseFormat(fromOptions.getResponseFormat()) + .withSeed(fromOptions.getSeed()) + .withStop(fromOptions.getStop()) + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withTools(fromOptions.getTools()) + .withToolChoice(fromOptions.getToolChoice()) + .withUser(fromOptions.getUser()) + .withFunctionCallbacks(fromOptions.getFunctionCallbacks()) + .withFunctions(fromOptions.getFunctions()) + .build(); + } + } diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingClient.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingModel.java similarity index 88% rename from models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingClient.java rename to models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingModel.java index 808c5f351..7192d2771 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingClient.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiEmbeddingModel.java @@ -22,7 +22,7 @@ import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; import org.springframework.ai.document.MetadataMode; -import org.springframework.ai.embedding.AbstractEmbeddingClient; +import org.springframework.ai.embedding.AbstractEmbeddingModel; import org.springframework.ai.embedding.Embedding; import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingRequest; @@ -41,9 +41,9 @@ import org.springframework.util.Assert; * * @author Christian Tzolov */ -public class OpenAiEmbeddingClient extends AbstractEmbeddingClient { +public class OpenAiEmbeddingModel extends AbstractEmbeddingModel { - private static final Logger logger = LoggerFactory.getLogger(OpenAiEmbeddingClient.class); + private static final Logger logger = LoggerFactory.getLogger(OpenAiEmbeddingModel.class); private final OpenAiEmbeddingOptions defaultOptions; @@ -54,43 +54,43 @@ public class OpenAiEmbeddingClient extends AbstractEmbeddingClient { private final MetadataMode metadataMode; /** - * Constructor for the OpenAiEmbeddingClient class. + * Constructor for the OpenAiEmbeddingModel class. * @param openAiApi The OpenAiApi instance to use for making API requests. */ - public OpenAiEmbeddingClient(OpenAiApi openAiApi) { + public OpenAiEmbeddingModel(OpenAiApi openAiApi) { this(openAiApi, MetadataMode.EMBED); } /** - * Initializes a new instance of the OpenAiEmbeddingClient class. + * Initializes a new instance of the OpenAiEmbeddingModel class. * @param openAiApi The OpenAiApi instance to use for making API requests. * @param metadataMode The mode for generating metadata. */ - public OpenAiEmbeddingClient(OpenAiApi openAiApi, MetadataMode metadataMode) { + public OpenAiEmbeddingModel(OpenAiApi openAiApi, MetadataMode metadataMode) { this(openAiApi, metadataMode, OpenAiEmbeddingOptions.builder().withModel(OpenAiApi.DEFAULT_EMBEDDING_MODEL).build(), RetryUtils.DEFAULT_RETRY_TEMPLATE); } /** - * Initializes a new instance of the OpenAiEmbeddingClient class. + * Initializes a new instance of the OpenAiEmbeddingModel class. * @param openAiApi The OpenAiApi instance to use for making API requests. * @param metadataMode The mode for generating metadata. * @param openAiEmbeddingOptions The options for OpenAi embedding. */ - public OpenAiEmbeddingClient(OpenAiApi openAiApi, MetadataMode metadataMode, + public OpenAiEmbeddingModel(OpenAiApi openAiApi, MetadataMode metadataMode, OpenAiEmbeddingOptions openAiEmbeddingOptions) { this(openAiApi, metadataMode, openAiEmbeddingOptions, RetryUtils.DEFAULT_RETRY_TEMPLATE); } /** - * Initializes a new instance of the OpenAiEmbeddingClient class. + * Initializes a new instance of the OpenAiEmbeddingModel class. * @param openAiApi - The OpenAiApi instance to use for making API requests. * @param metadataMode - The mode for generating metadata. * @param options - The options for OpenAI embedding. * @param retryTemplate - The RetryTemplate for retrying failed API requests. */ - public OpenAiEmbeddingClient(OpenAiApi openAiApi, MetadataMode metadataMode, OpenAiEmbeddingOptions options, + public OpenAiEmbeddingModel(OpenAiApi openAiApi, MetadataMode metadataMode, OpenAiEmbeddingOptions options, RetryTemplate retryTemplate) { Assert.notNull(openAiApi, "OpenAiService must not be null"); Assert.notNull(metadataMode, "metadataMode must not be null"); diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageClient.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageModel.java similarity index 94% rename from models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageClient.java rename to models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageModel.java index 69863ef9f..32cd93395 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageClient.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageModel.java @@ -21,7 +21,7 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.image.Image; -import org.springframework.ai.image.ImageClient; +import org.springframework.ai.image.ImageModel; import org.springframework.ai.image.ImageGeneration; import org.springframework.ai.image.ImageOptions; import org.springframework.ai.image.ImagePrompt; @@ -37,16 +37,16 @@ import org.springframework.retry.support.RetryTemplate; import org.springframework.util.Assert; /** - * OpenAiImageClient is a class that implements the ImageClient interface. It provides a + * OpenAiImageModel is a class that implements the ImageModel interface. It provides a * client for calling the OpenAI image generation API. * * @author Mark Pollack * @author Christian Tzolov * @since 0.8.0 */ -public class OpenAiImageClient implements ImageClient { +public class OpenAiImageModel implements ImageModel { - private final static Logger logger = LoggerFactory.getLogger(OpenAiImageClient.class); + private final static Logger logger = LoggerFactory.getLogger(OpenAiImageModel.class); private OpenAiImageOptions defaultOptions; @@ -54,11 +54,11 @@ public class OpenAiImageClient implements ImageClient { public final RetryTemplate retryTemplate; - public OpenAiImageClient(OpenAiImageApi openAiImageApi) { + public OpenAiImageModel(OpenAiImageApi openAiImageApi) { this(openAiImageApi, OpenAiImageOptions.builder().build(), RetryUtils.DEFAULT_RETRY_TEMPLATE); } - public OpenAiImageClient(OpenAiImageApi openAiImageApi, OpenAiImageOptions defaultOptions, + public OpenAiImageModel(OpenAiImageApi openAiImageApi, OpenAiImageOptions defaultOptions, RetryTemplate retryTemplate) { Assert.notNull(openAiImageApi, "OpenAiImageApi must not be null"); Assert.notNull(defaultOptions, "defaultOptions must not be null"); diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java index cc4256267..5d1a8b12b 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/api/OpenAiApi.java @@ -26,6 +26,7 @@ import com.fasterxml.jackson.annotation.JsonProperty; import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; +import org.springframework.ai.model.ModelDescription; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.retry.RetryUtils; import org.springframework.boot.context.properties.bind.ConstructorBinding; @@ -113,7 +114,7 @@ public class OpenAiApi { * - GPT-4 and GPT-4 Turbo * - GPT-3.5 Turbo. */ - public enum ChatModel { + public enum ChatModel implements ModelDescription { /** * Multimodal flagship model that’s cheaper and faster than GPT-4 Turbo. * Currently points to gpt-4o-2024-05-13. @@ -199,6 +200,11 @@ public class OpenAiApi { public String getValue() { return value; } + + @Override + public String getModelName() { + return this.value; + } } /** diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/audio/speech/SpeechClient.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/audio/speech/SpeechModel.java similarity index 87% rename from models/spring-ai-openai/src/main/java/org/springframework/ai/openai/audio/speech/SpeechClient.java rename to models/spring-ai-openai/src/main/java/org/springframework/ai/openai/audio/speech/SpeechModel.java index 1dc3876d0..9d976fd75 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/audio/speech/SpeechClient.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/audio/speech/SpeechModel.java @@ -16,10 +16,10 @@ package org.springframework.ai.openai.audio.speech; -import org.springframework.ai.model.ModelClient; +import org.springframework.ai.model.Model; /** - * The {@link SpeechClient} interface provides a way to interact with the OpenAI + * The {@link SpeechModel} interface provides a way to interact with the OpenAI * Text-to-Speech (TTS) API. It allows you to convert text input into lifelike spoken * audio. * @@ -27,7 +27,7 @@ import org.springframework.ai.model.ModelClient; * @since 1.0.0-M1 */ @FunctionalInterface -public interface SpeechClient extends ModelClient { +public interface SpeechModel extends Model { /** * Generates spoken audio from the provided text message. diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/audio/speech/StreamingSpeechClient.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/audio/speech/StreamingSpeechModel.java similarity index 87% rename from models/spring-ai-openai/src/main/java/org/springframework/ai/openai/audio/speech/StreamingSpeechClient.java rename to models/spring-ai-openai/src/main/java/org/springframework/ai/openai/audio/speech/StreamingSpeechModel.java index 209789bfd..a8ae06b07 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/audio/speech/StreamingSpeechClient.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/audio/speech/StreamingSpeechModel.java @@ -16,11 +16,11 @@ package org.springframework.ai.openai.audio.speech; -import org.springframework.ai.model.StreamingModelClient; +import org.springframework.ai.model.StreamingModel; import reactor.core.publisher.Flux; /** - * The {@link StreamingSpeechClient} interface provides a way to interact with the OpenAI + * The {@link StreamingSpeechModel} interface provides a way to interact with the OpenAI * Text-to-Speech (TTS) API using a streaming approach, allowing you to receive the * generated audio in a real-time fashion. * @@ -28,7 +28,7 @@ import reactor.core.publisher.Flux; * @since 1.0.0-M1 */ @FunctionalInterface -public interface StreamingSpeechClient extends StreamingModelClient { +public interface StreamingSpeechModel extends StreamingModel { /** * Generates a stream of audio bytes from the provided text message. diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/ChatCompletionRequestTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/ChatCompletionRequestTests.java index 300cdfafb..a4b117791 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/ChatCompletionRequestTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/ChatCompletionRequestTests.java @@ -34,7 +34,7 @@ public class ChatCompletionRequestTests { @Test public void createRequestWithChatOptions() { - var client = new OpenAiChatClient(new OpenAiApi("TEST"), + var client = new OpenAiChatModel(new OpenAiApi("TEST"), OpenAiChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6f).build()); var request = client.createRequest(new Prompt("Test message content"), false); @@ -60,7 +60,7 @@ public class ChatCompletionRequestTests { final String TOOL_FUNCTION_NAME = "CurrentWeather"; - var client = new OpenAiChatClient(new OpenAiApi("TEST"), + var client = new OpenAiChatModel(new OpenAiApi("TEST"), OpenAiChatOptions.builder().withModel("DEFAULT_MODEL").build()); var request = client.createRequest(new Prompt("Test message content", @@ -90,7 +90,7 @@ public class ChatCompletionRequestTests { final String TOOL_FUNCTION_NAME = "CurrentWeather"; - var client = new OpenAiChatClient(new OpenAiApi("TEST"), + var client = new OpenAiChatModel(new OpenAiApi("TEST"), OpenAiChatOptions.builder() .withModel("DEFAULT_MODEL") .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/OpenAiTestConfiguration.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/OpenAiTestConfiguration.java index 3a235cf7a..5c3f80dbb 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/OpenAiTestConfiguration.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/OpenAiTestConfiguration.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.openai; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.openai.api.OpenAiAudioApi; import org.springframework.ai.openai.api.OpenAiImageApi; @@ -51,33 +51,33 @@ public class OpenAiTestConfiguration { } @Bean - public OpenAiChatClient openAiChatClient(OpenAiApi api) { - OpenAiChatClient openAiChatClient = new OpenAiChatClient(api); - return openAiChatClient; + public OpenAiChatModel openAiChatModel(OpenAiApi api) { + OpenAiChatModel openAiChatModel = new OpenAiChatModel(api); + return openAiChatModel; } @Bean - public OpenAiAudioTranscriptionClient openAiTranscriptionClient(OpenAiAudioApi api) { - OpenAiAudioTranscriptionClient openAiTranscriptionClient = new OpenAiAudioTranscriptionClient(api); - return openAiTranscriptionClient; + public OpenAiAudioTranscriptionModel openAiTranscriptionModel(OpenAiAudioApi api) { + OpenAiAudioTranscriptionModel openAiTranscriptionModel = new OpenAiAudioTranscriptionModel(api); + return openAiTranscriptionModel; } @Bean - public OpenAiAudioSpeechClient openAiAudioSpeechClient(OpenAiAudioApi api) { - OpenAiAudioSpeechClient openAiAudioSpeechClient = new OpenAiAudioSpeechClient(api); - return openAiAudioSpeechClient; + public OpenAiAudioSpeechModel openAiAudioSpeechModel(OpenAiAudioApi api) { + OpenAiAudioSpeechModel openAiAudioSpeechModel = new OpenAiAudioSpeechModel(api); + return openAiAudioSpeechModel; } @Bean - public OpenAiImageClient openAiImageClient(OpenAiImageApi imageApi) { - OpenAiImageClient openAiImageClient = new OpenAiImageClient(imageApi); - // openAiImageClient.setModel("foobar"); - return openAiImageClient; + public OpenAiImageModel openAiImageModel(OpenAiImageApi imageApi) { + OpenAiImageModel openAiImageModel = new OpenAiImageModel(imageApi); + // openAiImageModel.setModel("foobar"); + return openAiImageModel; } @Bean - public EmbeddingClient openAiEmbeddingClient(OpenAiApi api) { - return new OpenAiEmbeddingClient(api); + public EmbeddingModel openAiEmbeddingModel(OpenAiApi api) { + return new OpenAiEmbeddingModel(api); } } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/TranscriptionRequestTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/TranscriptionRequestTests.java index 2f1239654..70530e9d3 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/TranscriptionRequestTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/TranscriptionRequestTests.java @@ -34,7 +34,7 @@ public class TranscriptionRequestTests { @Test public void defaultOptions() { - var client = new OpenAiAudioTranscriptionClient(new OpenAiAudioApi("TEST"), + var client = new OpenAiAudioTranscriptionModel(new OpenAiAudioApi("TEST"), OpenAiAudioTranscriptionOptions.builder() .withModel("DEFAULT_MODEL") .withResponseFormat(TranscriptResponseFormat.TEXT) @@ -58,7 +58,7 @@ public class TranscriptionRequestTests { @Test public void runtimeOptions() { - var client = new OpenAiAudioTranscriptionClient(new OpenAiAudioApi("TEST"), + var client = new OpenAiAudioTranscriptionModel(new OpenAiAudioApi("TEST"), OpenAiAudioTranscriptionOptions.builder() .withModel("DEFAULT_MODEL") .withResponseFormat(TranscriptResponseFormat.TEXT) diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/acme/AcmeIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/acme/AcmeIT.java index e38595b90..5d1e0dc06 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/acme/AcmeIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/acme/AcmeIT.java @@ -26,9 +26,9 @@ import org.slf4j.LoggerFactory; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.document.Document; +import org.springframework.ai.openai.OpenAiChatModel; import org.springframework.ai.openai.OpenAiTestConfiguration; -import org.springframework.ai.openai.OpenAiChatClient; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.ai.openai.testutils.AbstractIT; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.chat.prompt.SystemPromptTemplate; @@ -58,16 +58,16 @@ public class AcmeIT extends AbstractIT { private Resource systemBikePrompt; @Autowired - private OpenAiEmbeddingClient embeddingClient; + private OpenAiEmbeddingModel embeddingModel; @Autowired - private OpenAiChatClient chatClient; + private OpenAiChatModel chatModel; @Test void beanTest() { assertThat(bikesResource).isNotNull(); - assertThat(embeddingClient).isNotNull(); - assertThat(chatClient).isNotNull(); + assertThat(embeddingModel).isNotNull(); + assertThat(chatModel).isNotNull(); } // @Test @@ -81,7 +81,7 @@ public class AcmeIT extends AbstractIT { // Step 2 - Create embeddings and save to vector store logger.info("Creating Embeddings..."); - VectorStore vectorStore = new SimpleVectorStore(embeddingClient); + VectorStore vectorStore = new SimpleVectorStore(embeddingModel); vectorStore.accept(textSplitter.apply(jsonReader.get())); @@ -108,7 +108,7 @@ public class AcmeIT extends AbstractIT { logger.info("Asking AI generative to reply to question."); Prompt prompt = new Prompt(List.of(systemMessage, userMessage)); logger.info("AI responded."); - ChatResponse response = chatClient.call(prompt); + ChatResponse response = chatModel.call(prompt); evaluateQuestionAndAnswer(userQuery, response, true); } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechClientIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechModelIT.java similarity index 90% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechClientIT.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechModelIT.java index 5fdc3e24b..0ff96b259 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechClientIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechModelIT.java @@ -32,13 +32,13 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest(classes = OpenAiTestConfiguration.class) @EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+") -class OpenAiSpeechClientIT extends AbstractIT { +class OpenAiSpeechModelIT extends AbstractIT { private static final Float SPEED = 1.0f; @Test void shouldSuccessfullyStreamAudioBytesForEmptyMessage() { - Flux response = speechClient.stream("Today is a wonderful day to build something people love!"); + Flux response = speechModel.stream("Today is a wonderful day to build something people love!"); assertThat(response).isNotNull(); assertThat(response.collectList().block()).isNotNull(); System.out.println(response.collectList().block()); @@ -46,7 +46,7 @@ class OpenAiSpeechClientIT extends AbstractIT { @Test void shouldProduceAudioBytesDirectlyFromMessage() { - byte[] audioBytes = speechClient.call("Today is a wonderful day to build something people love!"); + byte[] audioBytes = speechModel.call("Today is a wonderful day to build something people love!"); assertThat(audioBytes).hasSizeGreaterThan(0); } @@ -61,7 +61,7 @@ class OpenAiSpeechClientIT extends AbstractIT { .build(); SpeechPrompt speechPrompt = new SpeechPrompt("Today is a wonderful day to build something people love!", speechOptions); - SpeechResponse response = speechClient.call(speechPrompt); + SpeechResponse response = speechModel.call(speechPrompt); byte[] audioBytes = response.getResult().getOutput(); assertThat(response.getResults()).hasSize(1); assertThat(response.getResults().get(0).getOutput()).isNotEmpty(); @@ -79,7 +79,7 @@ class OpenAiSpeechClientIT extends AbstractIT { .build(); SpeechPrompt speechPrompt = new SpeechPrompt("Today is a wonderful day to build something people love!", speechOptions); - SpeechResponse response = speechClient.call(speechPrompt); + SpeechResponse response = speechModel.call(speechPrompt); OpenAiAudioSpeechResponseMetadata metadata = response.getMetadata(); assertThat(metadata).isNotNull(); assertThat(metadata.getRateLimit()).isNotNull(); @@ -100,7 +100,7 @@ class OpenAiSpeechClientIT extends AbstractIT { SpeechPrompt speechPrompt = new SpeechPrompt("Today is a wonderful day to build something people love!", speechOptions); - Flux responseFlux = speechClient.stream(speechPrompt); + Flux responseFlux = speechModel.stream(speechPrompt); assertThat(responseFlux).isNotNull(); List responses = responseFlux.collectList().block(); assertThat(responses).isNotNull(); diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechClientWithSpeechResponseMetadataTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechModelWithSpeechResponseMetadataTests.java similarity index 92% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechClientWithSpeechResponseMetadataTests.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechModelWithSpeechResponseMetadataTests.java index 93f87781b..089c9c824 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechClientWithSpeechResponseMetadataTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechModelWithSpeechResponseMetadataTests.java @@ -18,7 +18,7 @@ package org.springframework.ai.openai.audio.speech; import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.Test; -import org.springframework.ai.openai.OpenAiAudioSpeechClient; +import org.springframework.ai.openai.OpenAiAudioSpeechModel; import org.springframework.ai.openai.OpenAiAudioSpeechOptions; import org.springframework.ai.openai.api.OpenAiAudioApi; import org.springframework.ai.openai.metadata.audio.OpenAiAudioSpeechResponseMetadata; @@ -45,15 +45,15 @@ import static org.springframework.test.web.client.response.MockRestResponseCreat /** * @author Ahmed Yousri */ -@RestClientTest(OpenAiSpeechClientWithSpeechResponseMetadataTests.Config.class) -public class OpenAiSpeechClientWithSpeechResponseMetadataTests { +@RestClientTest(OpenAiSpeechModelWithSpeechResponseMetadataTests.Config.class) +public class OpenAiSpeechModelWithSpeechResponseMetadataTests { private static String TEST_API_KEY = "sk-1234567890"; private static final Float SPEED = 1.0f; @Autowired - private OpenAiAudioSpeechClient openAiSpeechClient; + private OpenAiAudioSpeechModel openAiSpeechClient; @Autowired private MockRestServiceServer server; @@ -121,8 +121,8 @@ public class OpenAiSpeechClientWithSpeechResponseMetadataTests { static class Config { @Bean - public OpenAiAudioSpeechClient openAiAudioSpeechClient(OpenAiAudioApi openAiAudioApi) { - return new OpenAiAudioSpeechClient(openAiAudioApi); + public OpenAiAudioSpeechModel openAiAudioSpeechClient(OpenAiAudioApi openAiAudioApi) { + return new OpenAiAudioSpeechModel(openAiAudioApi); } @Bean diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionClientIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionModelIT.java similarity index 92% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionClientIT.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionModelIT.java index dcb10cd10..851d04c69 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionClientIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionModelIT.java @@ -31,7 +31,7 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest(classes = OpenAiTestConfiguration.class) @EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+") -class OpenAiTranscriptionClientIT extends AbstractIT { +class OpenAiTranscriptionModelIT extends AbstractIT { @Value("classpath:/speech/jfk.flac") private Resource audioFile; @@ -43,7 +43,7 @@ class OpenAiTranscriptionClientIT extends AbstractIT { .withTemperature(0f) .build(); AudioTranscriptionPrompt transcriptionRequest = new AudioTranscriptionPrompt(audioFile, transcriptionOptions); - AudioTranscriptionResponse response = transcriptionClient.call(transcriptionRequest); + AudioTranscriptionResponse response = transcriptionModel.call(transcriptionRequest); assertThat(response.getResults()).hasSize(1); assertThat(response.getResults().get(0).getOutput().toLowerCase().contains("fellow")).isTrue(); } @@ -59,7 +59,7 @@ class OpenAiTranscriptionClientIT extends AbstractIT { .withResponseFormat(responseFormat) .build(); AudioTranscriptionPrompt transcriptionRequest = new AudioTranscriptionPrompt(audioFile, transcriptionOptions); - AudioTranscriptionResponse response = transcriptionClient.call(transcriptionRequest); + AudioTranscriptionResponse response = transcriptionModel.call(transcriptionRequest); assertThat(response.getResults()).hasSize(1); assertThat(response.getResults().get(0).getOutput().toLowerCase().contains("fellow")).isTrue(); } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionClientWithTranscriptionResponseMetadataTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionModelWithTranscriptionResponseMetadataTests.java similarity index 92% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionClientWithTranscriptionResponseMetadataTests.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionModelWithTranscriptionResponseMetadataTests.java index 3b2629714..244af308f 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionClientWithTranscriptionResponseMetadataTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionModelWithTranscriptionResponseMetadataTests.java @@ -21,7 +21,7 @@ import org.junit.jupiter.api.AfterEach; import org.junit.jupiter.api.Test; import org.springframework.ai.chat.metadata.RateLimit; -import org.springframework.ai.openai.OpenAiAudioTranscriptionClient; +import org.springframework.ai.openai.OpenAiAudioTranscriptionModel; import org.springframework.ai.openai.api.OpenAiAudioApi; import org.springframework.ai.openai.metadata.audio.OpenAiAudioTranscriptionMetadata; import org.springframework.ai.openai.metadata.audio.OpenAiAudioTranscriptionResponseMetadata; @@ -48,13 +48,13 @@ import static org.springframework.test.web.client.response.MockRestResponseCreat /** * @author Michael Lavelle */ -@RestClientTest(OpenAiTranscriptionClientWithTranscriptionResponseMetadataTests.Config.class) -public class OpenAiTranscriptionClientWithTranscriptionResponseMetadataTests { +@RestClientTest(OpenAiTranscriptionModelWithTranscriptionResponseMetadataTests.Config.class) +public class OpenAiTranscriptionModelWithTranscriptionResponseMetadataTests { private static String TEST_API_KEY = "sk-1234567890"; @Autowired - private OpenAiAudioTranscriptionClient openAiTranscriptionClient; + private OpenAiAudioTranscriptionModel openAiTranscriptionClient; @Autowired private MockRestServiceServer server; @@ -156,8 +156,8 @@ public class OpenAiTranscriptionClientWithTranscriptionResponseMetadataTests { } @Bean - public OpenAiAudioTranscriptionClient openAiClient(OpenAiAudioApi openAiAudioApi) { - return new OpenAiAudioTranscriptionClient(openAiAudioApi); + public OpenAiAudioTranscriptionModel openAiClient(OpenAiAudioApi openAiAudioApi) { + return new OpenAiAudioTranscriptionModel(openAiAudioApi); } } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/TranscriptionClientTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/TranscriptionModelTests.java similarity index 91% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/TranscriptionClientTests.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/TranscriptionModelTests.java index 2b4f48e1d..96431af0e 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/TranscriptionClientTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/TranscriptionModelTests.java @@ -18,7 +18,7 @@ package org.springframework.ai.openai.audio.transcription; import org.junit.jupiter.api.Test; import org.mockito.Mockito; -import org.springframework.ai.openai.OpenAiAudioTranscriptionClient; +import org.springframework.ai.openai.OpenAiAudioTranscriptionModel; import org.springframework.core.io.Resource; import static org.assertj.core.api.Assertions.assertThat; @@ -33,18 +33,18 @@ import static org.mockito.Mockito.verifyNoMoreInteractions; import static org.mockito.Mockito.when; /** - * Unit Tests for {@link TranscriptionClient}. + * Unit Tests for {@link TranscriptionModel}. * * @author Michael Lavelle */ -class TranscriptionClientTests { +class TranscriptionModelTests { @Test void transcrbeRequestReturnsResponseCorrectly() { Resource mockAudioFile = Mockito.mock(Resource.class); - OpenAiAudioTranscriptionClient mockClient = Mockito.mock(OpenAiAudioTranscriptionClient.class); + OpenAiAudioTranscriptionModel mockClient = Mockito.mock(OpenAiAudioTranscriptionModel.class); String mockTranscription = "All your bases are belong to us"; diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java index 56ac4da9f..35da3bc3a 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java @@ -17,8 +17,8 @@ package org.springframework.ai.openai.chat; import java.io.IOException; import java.net.URL; -import java.util.ArrayList; import java.util.Arrays; +import java.util.Collection; import java.util.List; import java.util.Map; import java.util.stream.Collectors; @@ -31,19 +31,9 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import reactor.core.publisher.Flux; +import org.springframework.ai.chat.ChatClient; import org.springframework.ai.chat.ChatResponse; -import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.messages.AssistantMessage; -import org.springframework.ai.chat.messages.Media; -import org.springframework.ai.chat.messages.Message; -import org.springframework.ai.chat.messages.UserMessage; -import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.chat.prompt.PromptTemplate; -import org.springframework.ai.chat.prompt.SystemPromptTemplate; import org.springframework.ai.converter.BeanOutputConverter; -import org.springframework.ai.converter.ListOutputConverter; -import org.springframework.ai.converter.MapOutputConverter; -import org.springframework.ai.model.function.FunctionCallbackWrapper; import org.springframework.ai.openai.OpenAiChatOptions; import org.springframework.ai.openai.OpenAiTestConfiguration; import org.springframework.ai.openai.api.OpenAiApi; @@ -51,7 +41,7 @@ import org.springframework.ai.openai.api.tool.MockWeatherService; import org.springframework.ai.openai.testutils.AbstractIT; import org.springframework.beans.factory.annotation.Value; import org.springframework.boot.test.context.SpringBootTest; -import org.springframework.core.convert.support.DefaultConversionService; +import org.springframework.core.ParameterizedTypeReference; import org.springframework.core.io.ClassPathResource; import org.springframework.core.io.Resource; import org.springframework.util.MimeTypeUtils; @@ -65,75 +55,81 @@ class OpenAiChatClientIT extends AbstractIT { private static final Logger logger = LoggerFactory.getLogger(OpenAiChatClientIT.class); @Value("classpath:/prompts/system-message.st") - private Resource systemResource; + private Resource systemTextResource; @Test void roleTest() { - UserMessage userMessage = new UserMessage( - "Tell me about 3 famous pirates from the Golden Age of Piracy and what they did."); - SystemPromptTemplate systemPromptTemplate = new SystemPromptTemplate(systemResource); - Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", "Bob", "voice", "pirate")); - Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); - ChatResponse response = chatClient.call(prompt); + + // @formatter:off + ChatResponse response = ChatClient.builder(chatModel).build().prompt() + .system(s -> s.text(systemTextResource) + .param("name", "Bob") + .param("voice", "pirate")) + .user("Tell me about 3 famous pirates from the Golden Age of Piracy and what they did") + .call() + .chatResponse(); + // @formatter:on + + logger.info("" + response); assertThat(response.getResults()).hasSize(1); assertThat(response.getResults().get(0).getOutput().getContent()).contains("Blackbeard"); - // needs fine tuning... evaluateQuestionAndAnswer(request, response, false); } @Test void listOutputConverter() { - DefaultConversionService conversionService = new DefaultConversionService(); - ListOutputConverter outputConverter = new ListOutputConverter(conversionService); + // @formatter:off + Collection collection = ChatClient.builder(chatModel).build().prompt() + .user(u -> u.text("List five {subject}") + .param("subject", "ice cream flavors")) + .call() + .entity(new ParameterizedTypeReference>() {}); + // @formatter:on - String format = outputConverter.getFormat(); - String template = """ - List five {subject} - {format} - """; - PromptTemplate promptTemplate = new PromptTemplate(template, - Map.of("subject", "ice cream flavors", "format", format)); - Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = this.chatClient.call(prompt).getResult(); + assertThat(collection).hasSize(5); + } - List list = outputConverter.convert(generation.getOutput().getContent()); - assertThat(list).hasSize(5); + @Test + void listOutputConverter2() { + + // @formatter:off + List actorsFilms = ChatClient.builder(chatModel).build().prompt() + .user("Generate the filmography of 5 movies for Tom Hanks and Bill Murray.") + .call() + .entity(new ParameterizedTypeReference>() { + }); + // @formatter:on + + logger.info("" + actorsFilms); + assertThat(actorsFilms).hasSize(2); } @Test void mapOutputConverter() { - MapOutputConverter outputConverter = new MapOutputConverter(); + // @formatter:off + Map result = ChatClient.builder(chatModel).build().prompt() + .user(u -> u.text("Provide me a List of {subject}") + .param("subject", "an array of numbers from 1 to 9 under they key name 'numbers'")) + .call() + .entity(new ParameterizedTypeReference>() { + }); + // @formatter:on - String format = outputConverter.getFormat(); - String template = """ - Provide me a List of {subject} - {format} - """; - PromptTemplate promptTemplate = new PromptTemplate(template, - Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); - Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = chatClient.call(prompt).getResult(); - - Map result = outputConverter.convert(generation.getOutput().getContent()); assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); - } @Test void beanOutputConverter() { - BeanOutputConverter outputConverter = new BeanOutputConverter<>(ActorsFilms.class); + // @formatter:off + ActorsFilms actorsFilms = ChatClient.builder(chatModel).build().prompt() + .user("Generate the filmography for a random actor.") + .call() + .entity(ActorsFilms.class); + // @formatter:on - String format = outputConverter.getFormat(); - String template = """ - Generate the filmography for a random actor. - {format} - """; - PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); - Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = chatClient.call(prompt).getResult(); - - ActorsFilms actorsFilms = outputConverter.convert(generation.getOutput().getContent()); + logger.info("" + actorsFilms); + assertThat(actorsFilms.getActor()).isNotBlank(); } record ActorsFilmsRecord(String actor, List movies) { @@ -142,18 +138,13 @@ class OpenAiChatClientIT extends AbstractIT { @Test void beanOutputConverterRecords() { - BeanOutputConverter outputConverter = new BeanOutputConverter<>(ActorsFilmsRecord.class); + // @formatter:off + ActorsFilmsRecord actorsFilms = ChatClient.builder(chatModel).build().prompt() + .user("Generate the filmography of 5 movies for Tom Hanks.") + .call() + .entity(ActorsFilmsRecord.class); + // @formatter:on - String format = outputConverter.getFormat(); - String template = """ - Generate the filmography of 5 movies for Tom Hanks. - {format} - """; - PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); - Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = chatClient.call(prompt).getResult(); - - ActorsFilmsRecord actorsFilms = outputConverter.convert(generation.getOutput().getContent()); logger.info("" + actorsFilms); assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks"); assertThat(actorsFilms.movies()).hasSize(5); @@ -164,25 +155,25 @@ class OpenAiChatClientIT extends AbstractIT { BeanOutputConverter outputConverter = new BeanOutputConverter<>(ActorsFilmsRecord.class); - String format = outputConverter.getFormat(); - String template = """ - Generate the filmography of 5 movies for Tom Hanks. - {format} - """; - PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); - Prompt prompt = new Prompt(promptTemplate.createMessage()); + // @formatter:off + Flux chatResponse = ChatClient.builder(chatModel) + .build() + .prompt() + .user(u -> u + .text("Generate the filmography of 5 movies for Tom Hanks. " + System.lineSeparator() + + "{format}") + .param("format", outputConverter.getFormat())) + .stream() + .content(); - String generationTextFromStream = streamingChatClient.stream(prompt) - .collectList() - .block() - .stream() - .map(ChatResponse::getResults) - .flatMap(List::stream) - .map(Generation::getOutput) - .map(AssistantMessage::getContent) - .collect(Collectors.joining()); + String generationTextFromStream = chatResponse.collectList() + .block() + .stream() + .collect(Collectors.joining()); + // @formatter:on ActorsFilmsRecord actorsFilms = outputConverter.convert(generationTextFromStream); + logger.info("" + actorsFilms); assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks"); assertThat(actorsFilms.movies()).hasSize(5); @@ -191,54 +182,33 @@ class OpenAiChatClientIT extends AbstractIT { @Test void functionCallTest() { - UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - - List messages = new ArrayList<>(List.of(userMessage)); - - var promptOptions = OpenAiChatOptions.builder() - .withModel(OpenAiApi.ChatModel.GPT_4_TURBO_PREVIEW.getValue()) - .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) - .withName("getCurrentWeather") - .withDescription("Get the weather in location") - .withResponseConverter((response) -> "" + response.temp() + response.unit()) - .build())) - .build(); - - ChatResponse response = chatClient.call(new Prompt(messages, promptOptions)); + // @formatter:off + String response = ChatClient.builder(chatModel).build().prompt() + .user(u -> u.text("What's the weather like in San Francisco, Tokyo, and Paris?")) + .function("getCurrentWeather", "Get the weather in location", new MockWeatherService()) + .call() + .content(); + // @formatter:on logger.info("Response: {}", response); - assertThat(response.getResult().getOutput().getContent()).containsAnyOf("30.0", "30"); - assertThat(response.getResult().getOutput().getContent()).containsAnyOf("10.0", "10"); - assertThat(response.getResult().getOutput().getContent()).containsAnyOf("15.0", "15"); + assertThat(response).containsAnyOf("30.0", "30"); + assertThat(response).containsAnyOf("10.0", "10"); + assertThat(response).containsAnyOf("15.0", "15"); } @Test void streamFunctionCallTest() { - UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); + // @formatter:off + Flux response = ChatClient.builder(chatModel).build().prompt() + .user("What's the weather like in San Francisco, Tokyo, and Paris?") + .function("getCurrentWeather", "Get the weather in location", new MockWeatherService()) + .stream() + .content(); + // @formatter:on - List messages = new ArrayList<>(List.of(userMessage)); - - var promptOptions = OpenAiChatOptions.builder() - // .withModel(OpenAiApi.ChatModel.GPT_4_TURBO_PREVIEW.getValue()) - .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) - .withName("getCurrentWeather") - .withDescription("Get the weather in location") - .withResponseConverter((response) -> "" + response.temp() + response.unit()) - .build())) - .build(); - - Flux response = streamingChatClient.stream(new Prompt(messages, promptOptions)); - - String content = response.collectList() - .block() - .stream() - .map(ChatResponse::getResults) - .flatMap(List::stream) - .map(Generation::getOutput) - .map(AssistantMessage::getContent) - .collect(Collectors.joining()); + String content = response.collectList().block().stream().collect(Collectors.joining()); logger.info("Response: {}", content); assertThat(content).containsAnyOf("30.0", "30"); @@ -250,53 +220,62 @@ class OpenAiChatClientIT extends AbstractIT { @ValueSource(strings = { "gpt-4-vision-preview", "gpt-4o" }) void multiModalityEmbeddedImage(String modelName) throws IOException { - var imageData = new ClassPathResource("/test.png"); + // @formatter:off + String response = ChatClient.builder(chatModel).build().prompt() + // TODO consider adding model(...) method to ChatClient as a shortcut to + // OpenAiChatOptions.builder().withModel(modelName).build() + .options(OpenAiChatOptions.builder().withModel(modelName).build()) + .user(u -> u.text("Explain what do you see on this picture?") + .media(MimeTypeUtils.IMAGE_PNG, new ClassPathResource("/test.png"))) + .call() + .content(); + // @formatter:on - var userMessage = new UserMessage("Explain what do you see on this picture?", - List.of(new Media(MimeTypeUtils.IMAGE_PNG, imageData))); - - var response = chatClient - .call(new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withModel(modelName).build())); - - logger.info(response.getResult().getOutput().getContent()); - assertThat(response.getResult().getOutput().getContent()).contains("bananas", "apple"); - assertThat(response.getResult().getOutput().getContent()).containsAnyOf("bowl", "basket"); + logger.info(response); + assertThat(response).contains("bananas", "apple"); + assertThat(response).containsAnyOf("bowl", "basket"); } @ParameterizedTest(name = "{0} : {displayName} ") @ValueSource(strings = { "gpt-4-vision-preview", "gpt-4o" }) void multiModalityImageUrl(String modelName) throws IOException { - var userMessage = new UserMessage("Explain what do you see on this picture?", List - .of(new Media(MimeTypeUtils.IMAGE_PNG, - new URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png")))); + // TODO: add url method that wrapps the checked exception. + URL url = new URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png"); - ChatResponse response = chatClient - .call(new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withModel(modelName).build())); + // @formatter:off + String response = ChatClient.builder(chatModel).build().prompt() + // TODO consider adding model(...) method to ChatClient as a shortcut to + // OpenAiChatOptions.builder().withModel(modelName).build() + .options(OpenAiChatOptions.builder().withModel(modelName).build()) + .user(u -> u.text("Explain what do you see on this picture?").media(MimeTypeUtils.IMAGE_PNG, url)) + .call() + .content(); + // @formatter:on - logger.info(response.getResult().getOutput().getContent()); - assertThat(response.getResult().getOutput().getContent()).contains("bananas", "apple"); - assertThat(response.getResult().getOutput().getContent()).containsAnyOf("bowl", "basket"); + logger.info(response); + assertThat(response).contains("bananas", "apple"); + assertThat(response).containsAnyOf("bowl", "basket"); } @Test void streamingMultiModalityImageUrl() throws IOException { - var userMessage = new UserMessage("Explain what do you see on this picture?", List - .of(new Media(MimeTypeUtils.IMAGE_PNG, - new URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png")))); + // TODO: add url method that wrapps the checked exception. + URL url = new URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png"); - Flux response = streamingChatClient.stream(new Prompt(List.of(userMessage), - OpenAiChatOptions.builder().withModel(OpenAiApi.ChatModel.GPT_4_VISION_PREVIEW.getValue()).build())); + // @formatter:off + Flux response = ChatClient.builder(chatModel).build().prompt() + .options(OpenAiChatOptions.builder().withModel(OpenAiApi.ChatModel.GPT_4_VISION_PREVIEW.getValue()) + .build()) + .user(u -> u.text("Explain what do you see on this picture?") + .media(MimeTypeUtils.IMAGE_PNG, url)) + .stream() + .content(); + // @formatter:on + + String content = response.collectList().block().stream().collect(Collectors.joining()); - 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).contains("bananas", "apple"); assertThat(content).containsAnyOf("bowl", "basket"); diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClient2IT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModel2IT.java similarity index 89% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClient2IT.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModel2IT.java index fc5904be1..6c0127641 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClient2IT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModel2IT.java @@ -27,7 +27,7 @@ import org.slf4j.LoggerFactory; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.openai.OpenAiChatClient; +import org.springframework.ai.openai.OpenAiChatModel; import org.springframework.ai.openai.OpenAiChatOptions; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionRequest; @@ -41,14 +41,14 @@ import static org.assertj.core.api.Assertions.assertThat; /** * @author Christian Tzolov */ -@SpringBootTest(classes = OpenAiChatClient2IT.Config.class) +@SpringBootTest(classes = OpenAiChatModel2IT.Config.class) @EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+") -public class OpenAiChatClient2IT { +public class OpenAiChatModel2IT { private final Logger logger = LoggerFactory.getLogger(getClass()); @Autowired - private OpenAiChatClient openAiChatClient; + private OpenAiChatModel openAiChatModel; @Test void responseFormatTest() throws JsonMappingException, JsonProcessingException { @@ -67,7 +67,7 @@ public class OpenAiChatClient2IT { .withResponseFormat(new ChatCompletionRequest.ResponseFormat("json_object")) .build()); - ChatResponse response = this.openAiChatClient.call(prompt); + ChatResponse response = this.openAiChatModel.call(prompt); assertThat(response).isNotNull(); @@ -99,8 +99,8 @@ public class OpenAiChatClient2IT { } @Bean - public OpenAiChatClient openAiClient(OpenAiApi openAiApi) { - return new OpenAiChatClient(openAiApi); + public OpenAiChatModel openAiClient(OpenAiApi openAiApi) { + return new OpenAiChatModel(openAiApi); } } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelIT.java new file mode 100644 index 000000000..1d811e82e --- /dev/null +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelIT.java @@ -0,0 +1,305 @@ +/* + * Copyright 2023 - 2024 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.ai.openai.chat; + +import java.io.IOException; +import java.net.URL; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; +import java.util.Map; +import java.util.stream.Collectors; + +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; +import reactor.core.publisher.Flux; + +import org.springframework.ai.chat.ChatResponse; +import org.springframework.ai.chat.Generation; +import org.springframework.ai.chat.messages.AssistantMessage; +import org.springframework.ai.chat.messages.Media; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.chat.prompt.PromptTemplate; +import org.springframework.ai.chat.prompt.SystemPromptTemplate; +import org.springframework.ai.converter.BeanOutputConverter; +import org.springframework.ai.converter.ListOutputConverter; +import org.springframework.ai.converter.MapOutputConverter; +import org.springframework.ai.model.function.FunctionCallbackWrapper; +import org.springframework.ai.openai.OpenAiChatOptions; +import org.springframework.ai.openai.OpenAiTestConfiguration; +import org.springframework.ai.openai.api.OpenAiApi; +import org.springframework.ai.openai.api.tool.MockWeatherService; +import org.springframework.ai.openai.testutils.AbstractIT; +import org.springframework.beans.factory.annotation.Value; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.core.convert.support.DefaultConversionService; +import org.springframework.core.io.ClassPathResource; +import org.springframework.core.io.Resource; +import org.springframework.util.MimeTypeUtils; + +import static org.assertj.core.api.Assertions.assertThat; + +@SpringBootTest(classes = OpenAiTestConfiguration.class) +@EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+") +class OpenAiChatModelIT extends AbstractIT { + + private static final Logger logger = LoggerFactory.getLogger(OpenAiChatModelIT.class); + + @Value("classpath:/prompts/system-message.st") + private Resource systemResource; + + @Test + void roleTest() { + UserMessage userMessage = new UserMessage( + "Tell me about 3 famous pirates from the Golden Age of Piracy and what they did."); + SystemPromptTemplate systemPromptTemplate = new SystemPromptTemplate(systemResource); + Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", "Bob", "voice", "pirate")); + Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); + ChatResponse response = chatModel.call(prompt); + assertThat(response.getResults()).hasSize(1); + assertThat(response.getResults().get(0).getOutput().getContent()).contains("Blackbeard"); + // needs fine tuning... evaluateQuestionAndAnswer(request, response, false); + } + + @Test + void listOutputConverter() { + DefaultConversionService conversionService = new DefaultConversionService(); + ListOutputConverter outputConverter = new ListOutputConverter(conversionService); + + String format = outputConverter.getFormat(); + String template = """ + List five {subject} + {format} + """; + PromptTemplate promptTemplate = new PromptTemplate(template, + Map.of("subject", "ice cream flavors", "format", format)); + Prompt prompt = new Prompt(promptTemplate.createMessage()); + Generation generation = this.chatModel.call(prompt).getResult(); + + List list = outputConverter.convert(generation.getOutput().getContent()); + assertThat(list).hasSize(5); + + } + + @Test + void mapOutputConverter() { + MapOutputConverter outputConverter = new MapOutputConverter(); + + String format = outputConverter.getFormat(); + String template = """ + Provide me a List of {subject} + {format} + """; + PromptTemplate promptTemplate = new PromptTemplate(template, + Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); + Prompt prompt = new Prompt(promptTemplate.createMessage()); + Generation generation = chatModel.call(prompt).getResult(); + + Map result = outputConverter.convert(generation.getOutput().getContent()); + assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); + + } + + @Test + void beanOutputConverter() { + + BeanOutputConverter outputConverter = new BeanOutputConverter<>(ActorsFilms.class); + + String format = outputConverter.getFormat(); + String template = """ + Generate the filmography for a random actor. + {format} + """; + PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); + Prompt prompt = new Prompt(promptTemplate.createMessage()); + Generation generation = chatModel.call(prompt).getResult(); + + ActorsFilms actorsFilms = outputConverter.convert(generation.getOutput().getContent()); + } + + record ActorsFilmsRecord(String actor, List movies) { + } + + @Test + void beanOutputConverterRecords() { + + BeanOutputConverter outputConverter = new BeanOutputConverter<>(ActorsFilmsRecord.class); + + String format = outputConverter.getFormat(); + String template = """ + Generate the filmography of 5 movies for Tom Hanks. + {format} + """; + PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); + Prompt prompt = new Prompt(promptTemplate.createMessage()); + Generation generation = chatModel.call(prompt).getResult(); + + ActorsFilmsRecord actorsFilms = outputConverter.convert(generation.getOutput().getContent()); + logger.info("" + actorsFilms); + assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks"); + assertThat(actorsFilms.movies()).hasSize(5); + } + + @Test + void beanStreamOutputConverterRecords() { + + BeanOutputConverter outputConverter = new BeanOutputConverter<>(ActorsFilmsRecord.class); + + String format = outputConverter.getFormat(); + String template = """ + Generate the filmography of 5 movies for Tom Hanks. + {format} + """; + PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); + Prompt prompt = new Prompt(promptTemplate.createMessage()); + + String generationTextFromStream = streamingChatModel.stream(prompt) + .collectList() + .block() + .stream() + .map(ChatResponse::getResults) + .flatMap(List::stream) + .map(Generation::getOutput) + .map(AssistantMessage::getContent) + .collect(Collectors.joining()); + + ActorsFilmsRecord actorsFilms = outputConverter.convert(generationTextFromStream); + logger.info("" + actorsFilms); + assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks"); + assertThat(actorsFilms.movies()).hasSize(5); + } + + @Test + void functionCallTest() { + + UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); + + List messages = new ArrayList<>(List.of(userMessage)); + + var promptOptions = OpenAiChatOptions.builder() + .withModel(OpenAiApi.ChatModel.GPT_4_TURBO_PREVIEW.getValue()) + .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) + .withName("getCurrentWeather") + .withDescription("Get the weather in location") + .withResponseConverter((response) -> "" + response.temp() + response.unit()) + .build())) + .build(); + + ChatResponse response = chatModel.call(new Prompt(messages, promptOptions)); + + logger.info("Response: {}", response); + + assertThat(response.getResult().getOutput().getContent()).containsAnyOf("30.0", "30"); + assertThat(response.getResult().getOutput().getContent()).containsAnyOf("10.0", "10"); + assertThat(response.getResult().getOutput().getContent()).containsAnyOf("15.0", "15"); + } + + @Test + void streamFunctionCallTest() { + + UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); + + List messages = new ArrayList<>(List.of(userMessage)); + + var promptOptions = OpenAiChatOptions.builder() + // .withModel(OpenAiApi.ChatModel.GPT_4_TURBO_PREVIEW.getValue()) + .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) + .withName("getCurrentWeather") + .withDescription("Get the weather in location") + .withResponseConverter((response) -> "" + response.temp() + response.unit()) + .build())) + .build(); + + Flux response = streamingChatModel.stream(new Prompt(messages, promptOptions)); + + 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"); + } + + @ParameterizedTest(name = "{0} : {displayName} ") + @ValueSource(strings = { "gpt-4-vision-preview", "gpt-4o" }) + void multiModalityEmbeddedImage(String modelName) throws IOException { + + var imageData = new ClassPathResource("/test.png"); + + var userMessage = new UserMessage("Explain what do you see on this picture?", + List.of(new Media(MimeTypeUtils.IMAGE_PNG, imageData))); + + var response = chatModel + .call(new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withModel(modelName).build())); + + logger.info(response.getResult().getOutput().getContent()); + assertThat(response.getResult().getOutput().getContent()).contains("bananas", "apple"); + assertThat(response.getResult().getOutput().getContent()).containsAnyOf("bowl", "basket"); + } + + @ParameterizedTest(name = "{0} : {displayName} ") + @ValueSource(strings = { "gpt-4-vision-preview", "gpt-4o" }) + void multiModalityImageUrl(String modelName) throws IOException { + + var userMessage = new UserMessage("Explain what do you see on this picture?", List + .of(new Media(MimeTypeUtils.IMAGE_PNG, + new URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png")))); + + ChatResponse response = chatModel + .call(new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withModel(modelName).build())); + + logger.info(response.getResult().getOutput().getContent()); + assertThat(response.getResult().getOutput().getContent()).contains("bananas", "apple"); + assertThat(response.getResult().getOutput().getContent()).containsAnyOf("bowl", "basket"); + } + + @Test + void streamingMultiModalityImageUrl() throws IOException { + + var userMessage = new UserMessage("Explain what do you see on this picture?", List + .of(new Media(MimeTypeUtils.IMAGE_PNG, + new URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png")))); + + Flux response = streamingChatModel.stream(new Prompt(List.of(userMessage), + OpenAiChatOptions.builder().withModel(OpenAiApi.ChatModel.GPT_4_VISION_PREVIEW.getValue()).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); + assertThat(content).contains("bananas", "apple"); + assertThat(content).containsAnyOf("bowl", "basket"); + } + +} \ No newline at end of file diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientTypeReferenceBeanOutputConverterIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelTypeReferenceBeanOutputConverterIT.java similarity index 93% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientTypeReferenceBeanOutputConverterIT.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelTypeReferenceBeanOutputConverterIT.java index 4626af1fa..a35d31d61 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientTypeReferenceBeanOutputConverterIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelTypeReferenceBeanOutputConverterIT.java @@ -39,10 +39,10 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest(classes = OpenAiTestConfiguration.class) @EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+") -class OpenAiChatClientTypeReferenceBeanOutputConverterIT extends AbstractIT { +class OpenAiChatModelTypeReferenceBeanOutputConverterIT extends AbstractIT { private static final Logger logger = LoggerFactory - .getLogger(OpenAiChatClientTypeReferenceBeanOutputConverterIT.class); + .getLogger(OpenAiChatModelTypeReferenceBeanOutputConverterIT.class); record ActorsFilmsRecord(String actor, List movies) { } @@ -61,7 +61,7 @@ class OpenAiChatClientTypeReferenceBeanOutputConverterIT extends AbstractIT { """; PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = chatClient.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); List actorsFilms = outputConverter.convert(generation.getOutput().getContent()); logger.info("" + actorsFilms); @@ -87,7 +87,7 @@ class OpenAiChatClientTypeReferenceBeanOutputConverterIT extends AbstractIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - String generationTextFromStream = streamingChatClient.stream(prompt) + String generationTextFromStream = streamingChatModel.stream(prompt) .collectList() .block() .stream() diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientWithChatResponseMetadataTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelWithChatResponseMetadataTests.java similarity index 94% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientWithChatResponseMetadataTests.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelWithChatResponseMetadataTests.java index b1ef5e43b..7e4ee28f1 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientWithChatResponseMetadataTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelWithChatResponseMetadataTests.java @@ -27,7 +27,7 @@ import org.springframework.ai.chat.metadata.PromptMetadata; import org.springframework.ai.chat.metadata.RateLimit; import org.springframework.ai.chat.metadata.Usage; import org.springframework.ai.openai.api.OpenAiApi; -import org.springframework.ai.openai.OpenAiChatClient; +import org.springframework.ai.openai.OpenAiChatModel; import org.springframework.ai.openai.metadata.support.OpenAiApiResponseHeaders; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.beans.factory.annotation.Autowired; @@ -51,13 +51,13 @@ import static org.springframework.test.web.client.response.MockRestResponseCreat * @author Christian Tzolov * @since 0.7.0 */ -@RestClientTest(OpenAiChatClientWithChatResponseMetadataTests.Config.class) -public class OpenAiChatClientWithChatResponseMetadataTests { +@RestClientTest(OpenAiChatModelWithChatResponseMetadataTests.Config.class) +public class OpenAiChatModelWithChatResponseMetadataTests { private static String TEST_API_KEY = "sk-1234567890"; @Autowired - private OpenAiChatClient openAiChatClient; + private OpenAiChatModel openAiChatClient; @Autowired private MockRestServiceServer server; @@ -171,8 +171,8 @@ public class OpenAiChatClientWithChatResponseMetadataTests { } @Bean - public OpenAiChatClient openAiClient(OpenAiApi openAiApi) { - return new OpenAiChatClient(openAiApi); + public OpenAiChatModel openAiClient(OpenAiApi openAiApi) { + return new OpenAiChatModel(openAiApi); } } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiRetryTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiRetryTests.java index dcf0f303f..6bcd64408 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiRetryTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiRetryTests.java @@ -29,13 +29,13 @@ import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.document.MetadataMode; import org.springframework.ai.image.ImageMessage; import org.springframework.ai.image.ImagePrompt; -import org.springframework.ai.openai.OpenAiAudioTranscriptionClient; +import org.springframework.ai.openai.OpenAiAudioTranscriptionModel; import org.springframework.ai.openai.OpenAiAudioTranscriptionOptions; -import org.springframework.ai.openai.OpenAiChatClient; +import org.springframework.ai.openai.OpenAiChatModel; import org.springframework.ai.openai.OpenAiChatOptions; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.ai.openai.OpenAiEmbeddingOptions; -import org.springframework.ai.openai.OpenAiImageClient; +import org.springframework.ai.openai.OpenAiImageModel; import org.springframework.ai.openai.OpenAiImageOptions; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletion; @@ -107,13 +107,13 @@ public class OpenAiRetryTests { private @Mock OpenAiImageApi openAiImageApi; - private OpenAiChatClient chatClient; + private OpenAiChatModel chatModel; - private OpenAiEmbeddingClient embeddingClient; + private OpenAiEmbeddingModel embeddingModel; - private OpenAiAudioTranscriptionClient audioTranscriptionClient; + private OpenAiAudioTranscriptionModel audioTranscriptionModel; - private OpenAiImageClient imageClient; + private OpenAiImageModel imageModel; @BeforeEach public void beforeEach() { @@ -121,16 +121,16 @@ public class OpenAiRetryTests { retryListener = new TestRetryListener(); retryTemplate.registerListener(retryListener); - chatClient = new OpenAiChatClient(openAiApi, OpenAiChatOptions.builder().build(), null, retryTemplate); - embeddingClient = new OpenAiEmbeddingClient(openAiApi, MetadataMode.EMBED, + chatModel = new OpenAiChatModel(openAiApi, OpenAiChatOptions.builder().build(), null, retryTemplate); + embeddingModel = new OpenAiEmbeddingModel(openAiApi, MetadataMode.EMBED, OpenAiEmbeddingOptions.builder().build(), retryTemplate); - audioTranscriptionClient = new OpenAiAudioTranscriptionClient(openAiAudioApi, + audioTranscriptionModel = new OpenAiAudioTranscriptionModel(openAiAudioApi, OpenAiAudioTranscriptionOptions.builder() .withModel("model") .withResponseFormat(TranscriptResponseFormat.JSON) .build(), retryTemplate); - imageClient = new OpenAiImageClient(openAiImageApi, OpenAiImageOptions.builder().build(), retryTemplate); + imageModel = new OpenAiImageModel(openAiImageApi, OpenAiImageOptions.builder().build(), retryTemplate); } @Test @@ -146,7 +146,7 @@ public class OpenAiRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(ResponseEntity.of(Optional.of(expectedChatCompletion))); - var result = chatClient.call(new Prompt("text")); + var result = chatModel.call(new Prompt("text")); assertThat(result).isNotNull(); assertThat(result.getResult().getOutput().getContent()).isSameAs("Response"); @@ -158,7 +158,7 @@ public class OpenAiRetryTests { public void openAiChatNonTransientError() { when(openAiApi.chatCompletionEntity(isA(ChatCompletionRequest.class))) .thenThrow(new RuntimeException("Non Transient Error")); - assertThrows(RuntimeException.class, () -> chatClient.call(new Prompt("text"))); + assertThrows(RuntimeException.class, () -> chatModel.call(new Prompt("text"))); } @Test @@ -174,7 +174,7 @@ public class OpenAiRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(Flux.just(expectedChatCompletion)); - var result = chatClient.stream(new Prompt("text")); + var result = chatModel.stream(new Prompt("text")); assertThat(result).isNotNull(); assertThat(result.collectList().block().get(0).getResult().getOutput().getContent()).isSameAs("Response"); @@ -186,7 +186,7 @@ public class OpenAiRetryTests { public void openAiChatStreamNonTransientError() { when(openAiApi.chatCompletionStream(isA(ChatCompletionRequest.class))) .thenThrow(new RuntimeException("Non Transient Error")); - assertThrows(RuntimeException.class, () -> chatClient.stream(new Prompt("text"))); + assertThrows(RuntimeException.class, () -> chatModel.stream(new Prompt("text"))); } @Test @@ -199,7 +199,7 @@ public class OpenAiRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(ResponseEntity.of(Optional.of(expectedEmbeddings))); - var result = embeddingClient + var result = embeddingModel .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), null)); assertThat(result).isNotNull(); @@ -212,7 +212,7 @@ public class OpenAiRetryTests { public void openAiEmbeddingNonTransientError() { when(openAiApi.embeddings(isA(EmbeddingRequest.class))) .thenThrow(new RuntimeException("Non Transient Error")); - assertThrows(RuntimeException.class, () -> embeddingClient + assertThrows(RuntimeException.class, () -> embeddingModel .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), null))); } @@ -226,7 +226,7 @@ public class OpenAiRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(ResponseEntity.of(Optional.of(expectedResponse))); - AudioTranscriptionResponse result = audioTranscriptionClient + AudioTranscriptionResponse result = audioTranscriptionModel .call(new AudioTranscriptionPrompt(new ClassPathResource("speech/jfk.flac"))); assertThat(result).isNotNull(); @@ -239,7 +239,7 @@ public class OpenAiRetryTests { public void openAiAudioTranscriptionNonTransientError() { when(openAiAudioApi.createTranscription(isA(TranscriptionRequest.class), isA(Class.class))) .thenThrow(new RuntimeException("Transient Error 1")); - assertThrows(RuntimeException.class, () -> audioTranscriptionClient + assertThrows(RuntimeException.class, () -> audioTranscriptionModel .call(new AudioTranscriptionPrompt(new ClassPathResource("speech/jfk.flac")))); } @@ -253,7 +253,7 @@ public class OpenAiRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(ResponseEntity.of(Optional.of(expectedResponse))); - var result = imageClient.call(new ImagePrompt(List.of(new ImageMessage("Image Message")))); + var result = imageModel.call(new ImagePrompt(List.of(new ImageMessage("Image Message")))); assertThat(result).isNotNull(); assertThat(result.getResult().getOutput().getUrl()).isEqualTo("url678"); @@ -266,7 +266,7 @@ public class OpenAiRetryTests { when(openAiImageApi.createImage(isA(OpenAiImageRequest.class))) .thenThrow(new RuntimeException("Transient Error 1")); assertThrows(RuntimeException.class, - () -> imageClient.call(new ImagePrompt(List.of(new ImageMessage("Image Message"))))); + () -> imageModel.call(new ImagePrompt(List.of(new ImageMessage("Image Message"))))); } } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryLongTermSystemPromptIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryLongTermSystemPromptIT.java index 4516e08ae..e50fc1b13 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryLongTermSystemPromptIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryLongTermSystemPromptIT.java @@ -33,11 +33,11 @@ import org.springframework.ai.chat.memory.VectorStoreChatMemoryChatServiceListen import org.springframework.ai.chat.memory.VectorStoreChatMemoryRetriever; import org.springframework.ai.chat.memory.LastMaxTokenSizeContentTransformer; import org.springframework.ai.chat.memory.SystemPromptChatMemoryAugmentor; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.evaluation.BaseMemoryTest; import org.springframework.ai.evaluation.RelevancyEvaluator; -import org.springframework.ai.openai.OpenAiChatClient; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.openai.OpenAiChatModel; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.tokenizer.JTokkitTokenCountEstimator; import org.springframework.ai.tokenizer.TokenCountEstimator; @@ -75,21 +75,21 @@ public class ChatMemoryLongTermSystemPromptIT extends BaseMemoryTest { } @Bean - public OpenAiChatClient openAiClient(OpenAiApi openAiApi) { - return new OpenAiChatClient(openAiApi); + public OpenAiChatModel openAiClient(OpenAiApi openAiApi) { + return new OpenAiChatModel(openAiApi); } @Bean - public EmbeddingClient embeddingClient(OpenAiApi openAiApi) { - return new OpenAiEmbeddingClient(openAiApi); + public EmbeddingModel embeddingModel(OpenAiApi openAiApi) { + return new OpenAiEmbeddingModel(openAiApi); } @Bean - public VectorStore qdrantVectorStore(EmbeddingClient embeddingClient) { + public VectorStore qdrantVectorStore(EmbeddingModel embeddingModel) { QdrantClient qdrantClient = new QdrantClient(QdrantGrpcClient .newBuilder(qdrantContainer.getHost(), qdrantContainer.getMappedPort(QDRANT_GRPC_PORT), false) .build()); - return new QdrantVectorStore(qdrantClient, COLLECTION_NAME, embeddingClient); + return new QdrantVectorStore(qdrantClient, COLLECTION_NAME, embeddingModel); } @Bean @@ -98,10 +98,10 @@ public class ChatMemoryLongTermSystemPromptIT extends BaseMemoryTest { } @Bean - public ChatService memoryChatService(OpenAiChatClient chatClient, VectorStore vectorStore, + public ChatService memoryChatService(OpenAiChatModel chatModel, VectorStore vectorStore, TokenCountEstimator tokenCountEstimator) { - return PromptTransformingChatService.builder(chatClient) + return PromptTransformingChatService.builder(chatModel) .withRetrievers(List.of(new VectorStoreChatMemoryRetriever(vectorStore, 10))) .withContentPostProcessors(List.of(new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 1000))) .withAugmentors(List.of(new SystemPromptChatMemoryAugmentor())) @@ -110,10 +110,10 @@ public class ChatMemoryLongTermSystemPromptIT extends BaseMemoryTest { } @Bean - public StreamingChatService memoryStreamingChatService(OpenAiChatClient streamingChatClient, + public StreamingChatService memoryStreamingChatService(OpenAiChatModel streamingChatModel, VectorStore vectorStore, TokenCountEstimator tokenCountEstimator) { - return StreamingPromptTransformingChatService.builder(streamingChatClient) + return StreamingPromptTransformingChatService.builder(streamingChatModel) .withRetrievers(List.of(new VectorStoreChatMemoryRetriever(vectorStore, 10))) .withDocumentPostProcessors(List.of(new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 1000))) .withAugmentors(List.of(new SystemPromptChatMemoryAugmentor())) @@ -122,8 +122,8 @@ public class ChatMemoryLongTermSystemPromptIT extends BaseMemoryTest { } @Bean - public RelevancyEvaluator relevancyEvaluator(OpenAiChatClient chatClient) { - return new RelevancyEvaluator(chatClient); + public RelevancyEvaluator relevancyEvaluator(OpenAiChatModel chatModel) { + return new RelevancyEvaluator(chatModel); } } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryShortTermMessageListIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryShortTermMessageListIT.java index d26f6c563..60489546e 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryShortTermMessageListIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryShortTermMessageListIT.java @@ -31,7 +31,7 @@ import org.springframework.ai.chat.memory.LastMaxTokenSizeContentTransformer; import org.springframework.ai.chat.memory.MessageChatMemoryAugmentor; import org.springframework.ai.evaluation.BaseMemoryTest; import org.springframework.ai.evaluation.RelevancyEvaluator; -import org.springframework.ai.openai.OpenAiChatClient; +import org.springframework.ai.openai.OpenAiChatModel; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.tokenizer.JTokkitTokenCountEstimator; import org.springframework.ai.tokenizer.TokenCountEstimator; @@ -59,8 +59,8 @@ public class ChatMemoryShortTermMessageListIT extends BaseMemoryTest { } @Bean - public OpenAiChatClient openAiClient(OpenAiApi openAiApi) { - return new OpenAiChatClient(openAiApi); + public OpenAiChatModel openAiClient(OpenAiApi openAiApi) { + return new OpenAiChatModel(openAiApi); } @Bean @@ -74,10 +74,10 @@ public class ChatMemoryShortTermMessageListIT extends BaseMemoryTest { } @Bean - public ChatService memoryChatService(OpenAiChatClient chatClient, ChatMemory chatHistory, + public ChatService memoryChatService(OpenAiChatModel chatModel, ChatMemory chatHistory, TokenCountEstimator tokenCountEstimator) { - return PromptTransformingChatService.builder(chatClient) + return PromptTransformingChatService.builder(chatModel) .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) .withContentPostProcessors(List.of(new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 1000))) .withAugmentors(List.of(new MessageChatMemoryAugmentor())) @@ -86,10 +86,10 @@ public class ChatMemoryShortTermMessageListIT extends BaseMemoryTest { } @Bean - public StreamingChatService memoryStreamingChatService(OpenAiChatClient streamingChatClient, + public StreamingChatService memoryStreamingChatService(OpenAiChatModel streamingChatModel, ChatMemory chatHistory, TokenCountEstimator tokenCountEstimator) { - return StreamingPromptTransformingChatService.builder(streamingChatClient) + return StreamingPromptTransformingChatService.builder(streamingChatModel) .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) .withDocumentPostProcessors(List.of(new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 1000))) .withAugmentors(List.of(new MessageChatMemoryAugmentor())) @@ -98,8 +98,8 @@ public class ChatMemoryShortTermMessageListIT extends BaseMemoryTest { } @Bean - public RelevancyEvaluator relevancyEvaluator(OpenAiChatClient chatClient) { - return new RelevancyEvaluator(chatClient); + public RelevancyEvaluator relevancyEvaluator(OpenAiChatModel chatModel) { + return new RelevancyEvaluator(chatModel); } } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryShortTermSystemPromptIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryShortTermSystemPromptIT.java index 7ca4c795b..cc07d6da4 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryShortTermSystemPromptIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/ChatMemoryShortTermSystemPromptIT.java @@ -32,7 +32,7 @@ import org.springframework.ai.chat.memory.LastMaxTokenSizeContentTransformer; import org.springframework.ai.chat.memory.SystemPromptChatMemoryAugmentor; import org.springframework.ai.evaluation.BaseMemoryTest; import org.springframework.ai.evaluation.RelevancyEvaluator; -import org.springframework.ai.openai.OpenAiChatClient; +import org.springframework.ai.openai.OpenAiChatModel; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.tokenizer.JTokkitTokenCountEstimator; import org.springframework.ai.tokenizer.TokenCountEstimator; @@ -60,8 +60,8 @@ public class ChatMemoryShortTermSystemPromptIT extends BaseMemoryTest { } @Bean - public OpenAiChatClient openAiClient(OpenAiApi openAiApi) { - return new OpenAiChatClient(openAiApi); + public OpenAiChatModel openAiClient(OpenAiApi openAiApi) { + return new OpenAiChatModel(openAiApi); } @Bean @@ -75,10 +75,10 @@ public class ChatMemoryShortTermSystemPromptIT extends BaseMemoryTest { } @Bean - public ChatService memoryChatService(OpenAiChatClient chatClient, ChatMemory chatHistory, + public ChatService memoryChatService(OpenAiChatModel chatModel, ChatMemory chatHistory, TokenCountEstimator tokenCountEstimator) { - return PromptTransformingChatService.builder(chatClient) + return PromptTransformingChatService.builder(chatModel) .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) .withContentPostProcessors(List.of(new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 1000))) .withAugmentors(List.of(new SystemPromptChatMemoryAugmentor())) @@ -87,10 +87,10 @@ public class ChatMemoryShortTermSystemPromptIT extends BaseMemoryTest { } @Bean - public StreamingChatService memoryStreamingChatService(OpenAiChatClient streamingChatClient, + public StreamingChatService memoryStreamingChatService(OpenAiChatModel streamingChatModel, ChatMemory chatHistory, TokenCountEstimator tokenCountEstimator) { - return StreamingPromptTransformingChatService.builder(streamingChatClient) + return StreamingPromptTransformingChatService.builder(streamingChatModel) .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) .withDocumentPostProcessors(List.of(new LastMaxTokenSizeContentTransformer(tokenCountEstimator, 1000))) .withAugmentors(List.of(new SystemPromptChatMemoryAugmentor())) @@ -99,8 +99,8 @@ public class ChatMemoryShortTermSystemPromptIT extends BaseMemoryTest { } @Bean - public RelevancyEvaluator relevancyEvaluator(OpenAiChatClient chatClient) { - return new RelevancyEvaluator(chatClient); + public RelevancyEvaluator relevancyEvaluator(OpenAiChatModel chatModel) { + return new RelevancyEvaluator(chatModel); } } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/LongShortTermChatMemoryWithRagIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/LongShortTermChatMemoryWithRagIT.java index c2739a018..1162c9fce 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/LongShortTermChatMemoryWithRagIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/LongShortTermChatMemoryWithRagIT.java @@ -30,6 +30,7 @@ import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; import org.springframework.ai.chat.service.ChatService; import org.springframework.ai.chat.service.PromptTransformingChatService; import org.springframework.ai.openai.OpenAiChatOptions; +import org.springframework.ai.openai.OpenAiChatModel; import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.qdrant.QdrantContainer; @@ -49,12 +50,10 @@ import org.springframework.ai.chat.prompt.transformer.TransformerContentType; import org.springframework.ai.chat.prompt.transformer.VectorStoreRetriever; import org.springframework.ai.document.Document; import org.springframework.ai.document.DocumentTransformer; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.evaluation.EvaluationRequest; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.evaluation.EvaluationResponse; import org.springframework.ai.evaluation.RelevancyEvaluator; -import org.springframework.ai.openai.OpenAiChatClient; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.reader.JsonReader; import org.springframework.ai.tokenizer.JTokkitTokenCountEstimator; @@ -164,21 +163,21 @@ public class LongShortTermChatMemoryWithRagIT { } @Bean - public OpenAiChatClient openAiClient(OpenAiApi openAiApi) { - return new OpenAiChatClient(openAiApi); + public OpenAiChatModel openAiClient(OpenAiApi openAiApi) { + return new OpenAiChatModel(openAiApi); } @Bean - public OpenAiEmbeddingClient embeddingClient(OpenAiApi openAiApi) { - return new OpenAiEmbeddingClient(openAiApi); + public OpenAiEmbeddingModel embeddingModel(OpenAiApi openAiApi) { + return new OpenAiEmbeddingModel(openAiApi); } @Bean - public VectorStore qdrantVectorStore(EmbeddingClient embeddingClient) { + public VectorStore qdrantVectorStore(EmbeddingModel embeddingModel) { QdrantClient qdrantClient = new QdrantClient(QdrantGrpcClient .newBuilder(qdrantContainer.getHost(), qdrantContainer.getMappedPort(QDRANT_GRPC_PORT), false) .build()); - return new QdrantVectorStore(qdrantClient, COLLECTION_NAME, embeddingClient); + return new QdrantVectorStore(qdrantClient, COLLECTION_NAME, embeddingModel); } @Bean @@ -187,10 +186,10 @@ public class LongShortTermChatMemoryWithRagIT { } @Bean - public ChatService memoryChatService(OpenAiChatClient chatClient, VectorStore vectorStore, + public ChatService memoryChatService(OpenAiChatModel chatModel, VectorStore vectorStore, TokenCountEstimator tokenCountEstimator, ChatMemory chatHistory) { - return PromptTransformingChatService.builder(chatClient) + return PromptTransformingChatService.builder(chatModel) .withRetrievers(List.of(new VectorStoreRetriever(vectorStore, SearchRequest.defaults()), ChatMemoryRetriever.builder() .withChatHistory(chatHistory) @@ -224,12 +223,12 @@ public class LongShortTermChatMemoryWithRagIT { } // @Bean - // public StreamingChatService memoryStreamingChatAgent(OpenAiChatClient - // streamingChatClient, + // public StreamingChatService memoryStreamingChatAgent(OpenAiChatModel + // streamingChatModel, // VectorStore vectorStore, TokenCountEstimator tokenCountEstimator, ChatHistory // chatHistory) { - // return StreamingPromptTransformingChatService.builder(streamingChatClient) + // return StreamingPromptTransformingChatService.builder(streamingChatModel) // .withRetrievers(List.of(new ChatHistoryRetriever(chatHistory), new // DocumentChatHistoryRetriever(vectorStore, 10))) // .withDocumentPostProcessors(List.of(new @@ -241,13 +240,13 @@ public class LongShortTermChatMemoryWithRagIT { // } @Bean - public RelevancyEvaluator relevancyEvaluator(OpenAiChatClient chatClient) { + public RelevancyEvaluator relevancyEvaluator(OpenAiChatModel chatModel) { // Use GPT 4 as a better model for determining relevancy. gpt 3.5 makes basic // mistakes OpenAiChatOptions openAiChatOptions = OpenAiChatOptions.builder() .withModel(GPT_4_TURBO_PREVIEW.getValue()) .build(); - return new RelevancyEvaluator(chatClient, openAiChatOptions); + return new RelevancyEvaluator(chatModel, openAiChatOptions); } } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/OpenAiPromptTransformingChatServiceIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/OpenAiPromptTransformingChatServiceIT.java index a5f1ed415..4ed06cbd2 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/OpenAiPromptTransformingChatServiceIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/service/OpenAiPromptTransformingChatServiceIT.java @@ -23,6 +23,7 @@ import io.qdrant.client.QdrantClient; import io.qdrant.client.QdrantGrpcClient; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.service.ChatService; import org.springframework.ai.chat.prompt.transformer.TransformerContentType; import org.springframework.ai.document.Document; @@ -31,19 +32,17 @@ import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.qdrant.QdrantContainer; -import org.springframework.ai.chat.ChatClient; import org.springframework.ai.chat.service.PromptTransformingChatService; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; import org.springframework.ai.chat.prompt.transformer.QuestionContextAugmentor; import org.springframework.ai.chat.prompt.transformer.VectorStoreRetriever; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.evaluation.EvaluationRequest; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.evaluation.EvaluationResponse; import org.springframework.ai.evaluation.RelevancyEvaluator; -import org.springframework.ai.openai.OpenAiChatClient; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.openai.OpenAiChatModel; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.reader.JsonReader; import org.springframework.ai.transformer.splitter.TokenTextSplitter; @@ -72,7 +71,7 @@ public class OpenAiPromptTransformingChatServiceIT { @Container static QdrantContainer qdrantContainer = new QdrantContainer("qdrant/qdrant:v1.9.2"); - private final ChatClient chatClient; + private final ChatModel chatModel; private final VectorStore vectorStore; @@ -82,9 +81,9 @@ public class OpenAiPromptTransformingChatServiceIT { private ChatService chatService; @Autowired - public OpenAiPromptTransformingChatServiceIT(ChatClient chatClient, ChatService chatService, + public OpenAiPromptTransformingChatServiceIT(ChatModel chatModel, ChatService chatService, VectorStore vectorStore) { - this.chatClient = chatClient; + this.chatModel = chatModel; this.chatService = chatService; this.vectorStore = vectorStore; } @@ -103,7 +102,7 @@ public class OpenAiPromptTransformingChatServiceIT { OpenAiChatOptions openAiChatOptions = OpenAiChatOptions.builder() .withModel(GPT_4_TURBO_PREVIEW.getValue()) .build(); - var relevancyEvaluator = new RelevancyEvaluator(this.chatClient, openAiChatOptions); + var relevancyEvaluator = new RelevancyEvaluator(this.chatModel, openAiChatOptions); EvaluationResponse evaluationResponse = relevancyEvaluator.evaluate(chatServiceResponse.toEvaluationRequest()); assertTrue(evaluationResponse.isPass(), "Response is not relevant to the question"); @@ -146,26 +145,26 @@ public class OpenAiPromptTransformingChatServiceIT { } @Bean - public ChatClient openAiClient(OpenAiApi openAiApi) { - return new OpenAiChatClient(openAiApi); + public ChatModel openAiClient(OpenAiApi openAiApi) { + return new OpenAiChatModel(openAiApi); } @Bean - public EmbeddingClient embeddingClient(OpenAiApi openAiApi) { - return new OpenAiEmbeddingClient(openAiApi); + public EmbeddingModel embeddingModel(OpenAiApi openAiApi) { + return new OpenAiEmbeddingModel(openAiApi); } @Bean - public VectorStore qdrantVectorStore(EmbeddingClient embeddingClient) { + public VectorStore qdrantVectorStore(EmbeddingModel embeddingModel) { QdrantClient qdrantClient = new QdrantClient(QdrantGrpcClient .newBuilder(qdrantContainer.getHost(), qdrantContainer.getMappedPort(QDRANT_GRPC_PORT), false) .build()); - return new QdrantVectorStore(qdrantClient, COLLECTION_NAME, embeddingClient); + return new QdrantVectorStore(qdrantClient, COLLECTION_NAME, embeddingModel); } @Bean - public ChatService chatService(ChatClient chatClient, VectorStore vectorStore) { - return PromptTransformingChatService.builder(chatClient) + public ChatService chatService(ChatModel chatModel, VectorStore vectorStore) { + return PromptTransformingChatService.builder(chatModel) .withRetrievers(List.of(new VectorStoreRetriever(vectorStore, SearchRequest.defaults()))) .withAugmentors(List.of(new QuestionContextAugmentor())) .build(); diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/EmbeddingIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/EmbeddingIT.java index 683f54b48..326e9ebce 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/EmbeddingIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/embedding/EmbeddingIT.java @@ -19,7 +19,7 @@ import org.junit.jupiter.api.Test; import org.springframework.ai.embedding.EmbeddingRequest; import org.springframework.ai.embedding.EmbeddingResponse; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.ai.openai.OpenAiEmbeddingOptions; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.test.context.SpringBootTest; @@ -32,13 +32,13 @@ import static org.assertj.core.api.Assertions.assertThat; class EmbeddingIT { @Autowired - private OpenAiEmbeddingClient embeddingClient; + private OpenAiEmbeddingModel embeddingModel; @Test void defaultEmbedding() { - assertThat(embeddingClient).isNotNull(); + assertThat(embeddingModel).isNotNull(); - EmbeddingResponse embeddingResponse = embeddingClient.embedForResponse(List.of("Hello World")); + EmbeddingResponse embeddingResponse = embeddingModel.embedForResponse(List.of("Hello World")); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0)).isNotNull(); assertThat(embeddingResponse.getResults().get(0).getOutput()).hasSize(1536); @@ -46,13 +46,13 @@ class EmbeddingIT { assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 2); assertThat(embeddingResponse.getMetadata()).containsEntry("prompt-tokens", 2); - assertThat(embeddingClient.dimensions()).isEqualTo(1536); + assertThat(embeddingModel.dimensions()).isEqualTo(1536); } @Test void embedding3Large() { - EmbeddingResponse embeddingResponse = embeddingClient.call(new EmbeddingRequest(List.of("Hello World"), + EmbeddingResponse embeddingResponse = embeddingModel.call(new EmbeddingRequest(List.of("Hello World"), OpenAiEmbeddingOptions.builder().withModel("text-embedding-3-large").build())); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0)).isNotNull(); @@ -61,13 +61,13 @@ class EmbeddingIT { assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 2); assertThat(embeddingResponse.getMetadata()).containsEntry("prompt-tokens", 2); - // assertThat(embeddingClient.dimensions()).isEqualTo(3072); + // assertThat(embeddingModel.dimensions()).isEqualTo(3072); } @Test void textEmbeddingAda002() { - EmbeddingResponse embeddingResponse = embeddingClient.call(new EmbeddingRequest(List.of("Hello World"), + EmbeddingResponse embeddingResponse = embeddingModel.call(new EmbeddingRequest(List.of("Hello World"), OpenAiEmbeddingOptions.builder().withModel("text-embedding-3-small").build())); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0)).isNotNull(); @@ -77,7 +77,7 @@ class EmbeddingIT { assertThat(embeddingResponse.getMetadata()).containsEntry("total-tokens", 2); assertThat(embeddingResponse.getMetadata()).containsEntry("prompt-tokens", 2); - // assertThat(embeddingClient.dimensions()).isEqualTo(3072); + // assertThat(embeddingModel.dimensions()).isEqualTo(3072); } } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageClientIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageModelIT.java similarity index 95% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageClientIT.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageModelIT.java index 057e20dae..415760b0a 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageClientIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageModelIT.java @@ -33,7 +33,7 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest(classes = OpenAiTestConfiguration.class) @EnabledIfEnvironmentVariable(named = "OPENAI_API_KEY", matches = ".+") -public class OpenAiImageClientIT extends AbstractIT { +public class OpenAiImageModelIT extends AbstractIT { @Test void imageAsUrlTest() { @@ -44,7 +44,7 @@ public class OpenAiImageClientIT extends AbstractIT { ImagePrompt imagePrompt = new ImagePrompt(instructions, options); - ImageResponse imageResponse = imageClient.call(imagePrompt); + ImageResponse imageResponse = imageModel.call(imagePrompt); assertThat(imageResponse.getResults()).hasSize(1); diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageClientWithImageResponseMetadataTests.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageModelWithImageResponseMetadataTests.java similarity index 91% rename from models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageClientWithImageResponseMetadataTests.java rename to models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageModelWithImageResponseMetadataTests.java index 13d87c174..0133fd973 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageClientWithImageResponseMetadataTests.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/image/OpenAiImageModelWithImageResponseMetadataTests.java @@ -21,7 +21,7 @@ import org.springframework.ai.image.ImageGeneration; import org.springframework.ai.image.ImagePrompt; import org.springframework.ai.image.ImageResponse; import org.springframework.ai.image.ImageResponseMetadata; -import org.springframework.ai.openai.OpenAiImageClient; +import org.springframework.ai.openai.OpenAiImageModel; import org.springframework.ai.openai.api.OpenAiImageApi; import org.springframework.ai.openai.metadata.support.OpenAiApiResponseHeaders; import org.springframework.beans.factory.annotation.Autowired; @@ -47,13 +47,13 @@ import static org.springframework.test.web.client.response.MockRestResponseCreat * @author Christian Tzolov * @since 0.7.0 */ -@RestClientTest(OpenAiImageClientWithImageResponseMetadataTests.Config.class) -public class OpenAiImageClientWithImageResponseMetadataTests { +@RestClientTest(OpenAiImageModelWithImageResponseMetadataTests.Config.class) +public class OpenAiImageModelWithImageResponseMetadataTests { private static String TEST_API_KEY = "sk-1234567890"; @Autowired - private OpenAiImageClient openAiImageClient; + private OpenAiImageModel openAiImageModel; @Autowired private MockRestServiceServer server; @@ -70,7 +70,7 @@ public class OpenAiImageClientWithImageResponseMetadataTests { ImagePrompt prompt = new ImagePrompt("Create an image of a mini golden doodle dog."); - ImageResponse response = this.openAiImageClient.call(prompt); + ImageResponse response = this.openAiImageModel.call(prompt); assertThat(response).isNotNull(); List imageGenerations = response.getResults(); @@ -134,8 +134,8 @@ public class OpenAiImageClientWithImageResponseMetadataTests { } @Bean - public OpenAiImageClient openAiImageClient(OpenAiImageApi openAiImageApi) { - return new OpenAiImageClient(openAiImageApi); + public OpenAiImageModel openAiImageModel(OpenAiImageApi openAiImageApi) { + return new OpenAiImageModel(openAiImageApi); } } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/testutils/AbstractIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/testutils/AbstractIT.java index 89155e9e8..abae0a879 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/testutils/AbstractIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/testutils/AbstractIT.java @@ -21,16 +21,16 @@ import java.util.Map; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.chat.prompt.PromptTemplate; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.SystemMessage; -import org.springframework.ai.image.ImageClient; -import org.springframework.ai.openai.OpenAiAudioSpeechClient; -import org.springframework.ai.openai.OpenAiAudioTranscriptionClient; +import org.springframework.ai.image.ImageModel; +import org.springframework.ai.openai.OpenAiAudioSpeechModel; +import org.springframework.ai.openai.OpenAiAudioTranscriptionModel; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.beans.factory.annotation.Value; import org.springframework.core.io.Resource; @@ -43,19 +43,19 @@ public abstract class AbstractIT { private static final Logger logger = LoggerFactory.getLogger(AbstractIT.class); @Autowired - protected ChatClient chatClient; + protected ChatModel chatModel; @Autowired - protected StreamingChatClient streamingChatClient; + protected StreamingChatModel streamingChatModel; @Autowired - protected OpenAiAudioTranscriptionClient transcriptionClient; + protected OpenAiAudioTranscriptionModel transcriptionModel; @Autowired - protected OpenAiAudioSpeechClient speechClient; + protected OpenAiAudioSpeechModel speechModel; @Autowired - protected ImageClient imageClient; + protected ImageModel imageModel; @Value("classpath:/prompts/eval/qa-evaluator-accurate-answer.st") protected Resource qaEvaluatorAccurateAnswerResource; @@ -85,12 +85,12 @@ public abstract class AbstractIT { } Message userMessage = userPromptTemplate.createMessage(); Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); - String yesOrNo = chatClient.call(prompt).getResult().getOutput().getContent(); + String yesOrNo = chatModel.call(prompt).getResult().getOutput().getContent(); logger.info("Is Answer related to question: " + yesOrNo); if (yesOrNo.equalsIgnoreCase("no")) { SystemMessage notRelatedSystemMessage = new SystemMessage(qaEvaluatorNotRelatedResource); prompt = new Prompt(List.of(userMessage, notRelatedSystemMessage)); - String reasonForFailure = chatClient.call(prompt).getResult().getOutput().getContent(); + String reasonForFailure = chatModel.call(prompt).getResult().getOutput().getContent(); fail(reasonForFailure); } else { diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/transformer/MetadataTransformerIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/transformer/MetadataTransformerIT.java index d13c06612..aca29f27b 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/transformer/MetadataTransformerIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/transformer/MetadataTransformerIT.java @@ -24,8 +24,8 @@ import org.junit.jupiter.api.Test; import org.springframework.ai.document.DefaultContentFormatter; import org.springframework.ai.document.Document; +import org.springframework.ai.openai.OpenAiChatModel; import org.springframework.ai.openai.api.OpenAiApi; -import org.springframework.ai.openai.OpenAiChatClient; import org.springframework.ai.transformer.ContentFormatTransformer; import org.springframework.ai.transformer.KeywordMetadataEnricher; import org.springframework.ai.transformer.SummaryMetadataEnricher; @@ -163,18 +163,18 @@ public class MetadataTransformerIT { } @Bean - public OpenAiChatClient openAiChatClient(OpenAiApi openAiApi) { - OpenAiChatClient openAiChatClient = new OpenAiChatClient(openAiApi); - return openAiChatClient; + public OpenAiChatModel openAiChatModel(OpenAiApi openAiApi) { + OpenAiChatModel openAiChatModel = new OpenAiChatModel(openAiApi); + return openAiChatModel; } @Bean - public KeywordMetadataEnricher keywordMetadata(OpenAiChatClient aiClient) { + public KeywordMetadataEnricher keywordMetadata(OpenAiChatModel aiClient) { return new KeywordMetadataEnricher(aiClient, 5); } @Bean - public SummaryMetadataEnricher summaryMetadata(OpenAiChatClient aiClient) { + public SummaryMetadataEnricher summaryMetadata(OpenAiChatModel aiClient) { return new SummaryMetadataEnricher(aiClient, List.of(SummaryType.PREVIOUS, SummaryType.CURRENT, SummaryType.NEXT)); } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/vectorstore/SimplePersistentVectorStoreIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/vectorstore/SimplePersistentVectorStoreIT.java index 169d6e53d..32ba4572c 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/vectorstore/SimplePersistentVectorStoreIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/vectorstore/SimplePersistentVectorStoreIT.java @@ -19,7 +19,7 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.io.CleanupMode; import org.junit.jupiter.api.io.TempDir; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.reader.JsonReader; import org.springframework.ai.vectorstore.SimpleVectorStore; import org.springframework.ai.reader.JsonMetadataGenerator; @@ -42,21 +42,21 @@ public class SimplePersistentVectorStoreIT { private Resource bikesJsonResource; @Autowired - private EmbeddingClient embeddingClient; + private EmbeddingModel embeddingModel; @Test void persist(@TempDir(cleanup = CleanupMode.ON_SUCCESS) Path workingDir) { JsonReader jsonReader = new JsonReader(bikesJsonResource, new ProductMetadataGenerator(), "price", "name", "shortDescription", "description", "tags"); List documents = jsonReader.get(); - SimpleVectorStore vectorStore = new SimpleVectorStore(this.embeddingClient); + SimpleVectorStore vectorStore = new SimpleVectorStore(this.embeddingModel); vectorStore.add(documents); File tempFile = new File(workingDir.toFile(), "temp.txt"); vectorStore.save(tempFile); assertThat(tempFile).isNotEmpty(); assertThat(tempFile).content().contains("Velo 99 XR1 AXS"); - SimpleVectorStore vectorStore2 = new SimpleVectorStore(this.embeddingClient); + SimpleVectorStore vectorStore2 = new SimpleVectorStore(this.embeddingModel); vectorStore2.load(tempFile); List similaritySearch = vectorStore2.similaritySearch("Velo 99 XR1 AXS"); diff --git a/models/spring-ai-postgresml/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingClient.java b/models/spring-ai-postgresml/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingModel.java similarity index 91% rename from models/spring-ai-postgresml/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingClient.java rename to models/spring-ai-postgresml/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingModel.java index 677a7fba4..2783f6df3 100644 --- a/models/spring-ai-postgresml/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingClient.java +++ b/models/spring-ai-postgresml/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingModel.java @@ -24,7 +24,7 @@ import java.util.Map; import org.springframework.ai.document.Document; import org.springframework.ai.document.MetadataMode; -import org.springframework.ai.embedding.AbstractEmbeddingClient; +import org.springframework.ai.embedding.AbstractEmbeddingModel; import org.springframework.ai.embedding.Embedding; import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingRequest; @@ -39,12 +39,12 @@ import org.springframework.util.CollectionUtils; import org.springframework.util.StringUtils; /** - * PostgresML EmbeddingClient + * PostgresML EmbeddingModel * * @author Toshiaki Maki * @author Christian Tzolov */ -public class PostgresMlEmbeddingClient extends AbstractEmbeddingClient implements InitializingBean { +public class PostgresMlEmbeddingModel extends AbstractEmbeddingModel implements InitializingBean { public static final String DEFAULT_TRANSFORMER_MODEL = "distilbert-base-uncased"; @@ -83,16 +83,16 @@ public class PostgresMlEmbeddingClient extends AbstractEmbeddingClient implement * a constructor * @param jdbcTemplate JdbcTemplate */ - public PostgresMlEmbeddingClient(JdbcTemplate jdbcTemplate) { + public PostgresMlEmbeddingModel(JdbcTemplate jdbcTemplate) { this(jdbcTemplate, PostgresMlEmbeddingOptions.builder().build()); } /** - * a PostgresMlEmbeddingClient constructor + * a PostgresMlEmbeddingModel constructor * @param jdbcTemplate JdbcTemplate to use to interact with the database. * @param options PostgresMlEmbeddingOptions to configure the client. */ - public PostgresMlEmbeddingClient(JdbcTemplate jdbcTemplate, PostgresMlEmbeddingOptions options) { + public PostgresMlEmbeddingModel(JdbcTemplate jdbcTemplate, PostgresMlEmbeddingOptions options) { Assert.notNull(jdbcTemplate, "jdbc template must not be null."); Assert.notNull(options, "options must not be null."); Assert.notNull(options.getTransformer(), "transformer must not be null."); @@ -110,7 +110,7 @@ public class PostgresMlEmbeddingClient extends AbstractEmbeddingClient implement * @param transformer huggingface sentence-transformer name */ @Deprecated(since = "0.8.0", forRemoval = true) - public PostgresMlEmbeddingClient(JdbcTemplate jdbcTemplate, String transformer) { + public PostgresMlEmbeddingModel(JdbcTemplate jdbcTemplate, String transformer) { this(jdbcTemplate, transformer, VectorType.PG_ARRAY); } @@ -122,7 +122,7 @@ public class PostgresMlEmbeddingClient extends AbstractEmbeddingClient implement * @param vectorType vector type in PostgreSQL */ @Deprecated(since = "0.8.0", forRemoval = true) - public PostgresMlEmbeddingClient(JdbcTemplate jdbcTemplate, String transformer, VectorType vectorType) { + public PostgresMlEmbeddingModel(JdbcTemplate jdbcTemplate, String transformer, VectorType vectorType) { this(jdbcTemplate, transformer, vectorType, Map.of(), MetadataMode.EMBED); } @@ -135,7 +135,7 @@ public class PostgresMlEmbeddingClient extends AbstractEmbeddingClient implement * @param kwargs optional arguments */ @Deprecated(since = "0.8.0", forRemoval = true) - public PostgresMlEmbeddingClient(JdbcTemplate jdbcTemplate, String transformer, VectorType vectorType, + public PostgresMlEmbeddingModel(JdbcTemplate jdbcTemplate, String transformer, VectorType vectorType, Map kwargs, MetadataMode metadataMode) { Assert.notNull(jdbcTemplate, "jdbc template must not be null."); Assert.notNull(transformer, "transformer must not be null."); diff --git a/models/spring-ai-postgresml/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingOptions.java b/models/spring-ai-postgresml/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingOptions.java index b80e1f639..220d5d283 100644 --- a/models/spring-ai-postgresml/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingOptions.java +++ b/models/spring-ai-postgresml/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingOptions.java @@ -24,7 +24,7 @@ import com.fasterxml.jackson.annotation.JsonProperty; import org.springframework.ai.document.MetadataMode; import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.model.ModelOptionsUtils; -import org.springframework.ai.postgresml.PostgresMlEmbeddingClient.VectorType; +import org.springframework.ai.postgresml.PostgresMlEmbeddingModel.VectorType; /** * @author Christian Tzolov @@ -36,7 +36,7 @@ public class PostgresMlEmbeddingOptions implements EmbeddingOptions { /** * The Huggingface transformer model to use for the embedding. */ - private @JsonProperty("transformer") String transformer = PostgresMlEmbeddingClient.DEFAULT_TRANSFORMER_MODEL; + private @JsonProperty("transformer") String transformer = PostgresMlEmbeddingModel.DEFAULT_TRANSFORMER_MODEL; /** * PostgresML vector type to use for the embedding. diff --git a/models/spring-ai-postgresml/src/test/java/org/springframework/ai/postgresml/PostgresMlEmbeddingClientIT.java b/models/spring-ai-postgresml/src/test/java/org/springframework/ai/postgresml/PostgresMlEmbeddingModelIT.java similarity index 79% rename from models/spring-ai-postgresml/src/test/java/org/springframework/ai/postgresml/PostgresMlEmbeddingClientIT.java rename to models/spring-ai-postgresml/src/test/java/org/springframework/ai/postgresml/PostgresMlEmbeddingModelIT.java index 7a4539e0e..55c05d06d 100644 --- a/models/spring-ai-postgresml/src/test/java/org/springframework/ai/postgresml/PostgresMlEmbeddingClientIT.java +++ b/models/spring-ai-postgresml/src/test/java/org/springframework/ai/postgresml/PostgresMlEmbeddingModelIT.java @@ -30,7 +30,7 @@ import org.junit.jupiter.params.provider.ValueSource; import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingRequest; import org.springframework.ai.embedding.EmbeddingResponse; -import org.springframework.ai.postgresml.PostgresMlEmbeddingClient.VectorType; +import org.springframework.ai.postgresml.PostgresMlEmbeddingModel.VectorType; import org.testcontainers.containers.PostgreSQLContainer; import org.testcontainers.containers.wait.strategy.LogMessageWaitStrategy; @@ -56,7 +56,7 @@ import static org.assertj.core.api.Assertions.assertThat; @AutoConfigureTestDatabase(replace = AutoConfigureTestDatabase.Replace.NONE) @Testcontainers @Disabled("Disabled from automatic execution, as it requires an excessive amount of memory (over 9GB)!") -class PostgresMlEmbeddingClientIT { +class PostgresMlEmbeddingModelIT { @Container @ServiceConnection @@ -80,51 +80,51 @@ class PostgresMlEmbeddingClientIT { @Test void embed() { - PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.jdbcTemplate); - embeddingClient.afterPropertiesSet(); + PostgresMlEmbeddingModel embeddingModel = new PostgresMlEmbeddingModel(this.jdbcTemplate); + embeddingModel.afterPropertiesSet(); - List embed = embeddingClient.embed("Hello World!"); + List embed = embeddingModel.embed("Hello World!"); assertThat(embed).hasSize(768); } @Test void embedWithPgVector() { - PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.jdbcTemplate, + PostgresMlEmbeddingModel embeddingModel = new PostgresMlEmbeddingModel(this.jdbcTemplate, PostgresMlEmbeddingOptions.builder() .withTransformer("distilbert-base-uncased") - .withVectorType(PostgresMlEmbeddingClient.VectorType.PG_VECTOR) + .withVectorType(PostgresMlEmbeddingModel.VectorType.PG_VECTOR) .build()); - embeddingClient.afterPropertiesSet(); + embeddingModel.afterPropertiesSet(); - List embed = embeddingClient.embed(new Document("Hello World!")); + List embed = embeddingModel.embed(new Document("Hello World!")); assertThat(embed).hasSize(768); } @Test void embedWithDifferentModel() { - PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.jdbcTemplate, + PostgresMlEmbeddingModel embeddingModel = new PostgresMlEmbeddingModel(this.jdbcTemplate, PostgresMlEmbeddingOptions.builder().withTransformer("intfloat/e5-small").build()); - embeddingClient.afterPropertiesSet(); + embeddingModel.afterPropertiesSet(); - List embed = embeddingClient.embed(new Document("Hello World!")); + List embed = embeddingModel.embed(new Document("Hello World!")); assertThat(embed).hasSize(384); } @Test void embedWithKwargs() { - PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.jdbcTemplate, + PostgresMlEmbeddingModel embeddingModel = new PostgresMlEmbeddingModel(this.jdbcTemplate, PostgresMlEmbeddingOptions.builder() .withTransformer("distilbert-base-uncased") - .withVectorType(PostgresMlEmbeddingClient.VectorType.PG_ARRAY) + .withVectorType(PostgresMlEmbeddingModel.VectorType.PG_ARRAY) .withKwargs(Map.of("device", "cpu")) .withMetadataMode(MetadataMode.EMBED) .build()); - embeddingClient.afterPropertiesSet(); + embeddingModel.afterPropertiesSet(); - List embed = embeddingClient.embed(new Document("Hello World!")); + List embed = embeddingModel.embed(new Document("Hello World!")); assertThat(embed).hasSize(768); } @@ -132,14 +132,14 @@ class PostgresMlEmbeddingClientIT { @ParameterizedTest @ValueSource(strings = { "PG_ARRAY", "PG_VECTOR" }) void embedForResponse(String vectorType) { - PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.jdbcTemplate, + PostgresMlEmbeddingModel embeddingModel = new PostgresMlEmbeddingModel(this.jdbcTemplate, PostgresMlEmbeddingOptions.builder() .withTransformer("distilbert-base-uncased") .withVectorType(VectorType.valueOf(vectorType)) .build()); - embeddingClient.afterPropertiesSet(); + embeddingModel.afterPropertiesSet(); - EmbeddingResponse embeddingResponse = embeddingClient + EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World!", "Spring AI!", "LLM!")); assertThat(embeddingResponse).isNotNull(); @@ -157,16 +157,16 @@ class PostgresMlEmbeddingClientIT { @Test void embedCallWithRequestOptionsOverride() { - PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.jdbcTemplate, + PostgresMlEmbeddingModel embeddingModel = new PostgresMlEmbeddingModel(this.jdbcTemplate, PostgresMlEmbeddingOptions.builder() .withTransformer("distilbert-base-uncased") .withVectorType(VectorType.PG_VECTOR) .build()); - embeddingClient.afterPropertiesSet(); + embeddingModel.afterPropertiesSet(); var request1 = new EmbeddingRequest(List.of("Hello World!", "Spring AI!", "LLM!"), EmbeddingOptions.EMPTY); - EmbeddingResponse embeddingResponse = embeddingClient.call(request1); + EmbeddingResponse embeddingResponse = embeddingModel.call(request1); assertThat(embeddingResponse).isNotNull(); assertThat(embeddingResponse.getResults()).hasSize(3); @@ -188,7 +188,7 @@ class PostgresMlEmbeddingClientIT { .withKwargs(Map.of("device", "cpu")) .build()); - embeddingResponse = embeddingClient.call(request2); + embeddingResponse = embeddingModel.call(request2); assertThat(embeddingResponse).isNotNull(); assertThat(embeddingResponse.getResults()).hasSize(3); @@ -205,11 +205,11 @@ class PostgresMlEmbeddingClientIT { @Test void dimensions() { - PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.jdbcTemplate); - embeddingClient.afterPropertiesSet(); - Assertions.assertThat(embeddingClient.dimensions()).isEqualTo(768); + PostgresMlEmbeddingModel embeddingModel = new PostgresMlEmbeddingModel(this.jdbcTemplate); + embeddingModel.afterPropertiesSet(); + Assertions.assertThat(embeddingModel.dimensions()).isEqualTo(768); // cached - Assertions.assertThat(embeddingClient.dimensions()).isEqualTo(768); + Assertions.assertThat(embeddingModel.dimensions()).isEqualTo(768); } @SpringBootApplication diff --git a/models/spring-ai-postgresml/src/test/java/org/springframework/ai/postgresml/PostgresMlEmbeddingOptionsTests.java b/models/spring-ai-postgresml/src/test/java/org/springframework/ai/postgresml/PostgresMlEmbeddingOptionsTests.java index aeceb5c58..07ce531b7 100644 --- a/models/spring-ai-postgresml/src/test/java/org/springframework/ai/postgresml/PostgresMlEmbeddingOptionsTests.java +++ b/models/spring-ai-postgresml/src/test/java/org/springframework/ai/postgresml/PostgresMlEmbeddingOptionsTests.java @@ -34,8 +34,8 @@ public class PostgresMlEmbeddingOptionsTests { public void defaultOptions() { PostgresMlEmbeddingOptions options = PostgresMlEmbeddingOptions.builder().build(); - assertThat(options.getTransformer()).isEqualTo(PostgresMlEmbeddingClient.DEFAULT_TRANSFORMER_MODEL); - assertThat(options.getVectorType()).isEqualTo(PostgresMlEmbeddingClient.VectorType.PG_ARRAY); + assertThat(options.getTransformer()).isEqualTo(PostgresMlEmbeddingModel.DEFAULT_TRANSFORMER_MODEL); + assertThat(options.getVectorType()).isEqualTo(PostgresMlEmbeddingModel.VectorType.PG_ARRAY); assertThat(options.getKwargs()).isEqualTo(Map.of()); assertThat(options.getMetadataMode()).isEqualTo(org.springframework.ai.document.MetadataMode.EMBED); } @@ -44,13 +44,13 @@ public class PostgresMlEmbeddingOptionsTests { public void newOptions() { PostgresMlEmbeddingOptions options = PostgresMlEmbeddingOptions.builder() .withTransformer("intfloat/e5-small") - .withVectorType(PostgresMlEmbeddingClient.VectorType.PG_VECTOR) + .withVectorType(PostgresMlEmbeddingModel.VectorType.PG_VECTOR) .withMetadataMode(org.springframework.ai.document.MetadataMode.ALL) .withKwargs(Map.of("device", "cpu")) .build(); assertThat(options.getTransformer()).isEqualTo("intfloat/e5-small"); - assertThat(options.getVectorType()).isEqualTo(PostgresMlEmbeddingClient.VectorType.PG_VECTOR); + assertThat(options.getVectorType()).isEqualTo(PostgresMlEmbeddingModel.VectorType.PG_VECTOR); assertThat(options.getKwargs()).isEqualTo(Map.of("device", "cpu")); assertThat(options.getMetadataMode()).isEqualTo(org.springframework.ai.document.MetadataMode.ALL); } @@ -59,37 +59,37 @@ public class PostgresMlEmbeddingOptionsTests { public void mergeOptions() { var jdbcTemplate = Mockito.mock(JdbcTemplate.class); - PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(jdbcTemplate); + PostgresMlEmbeddingModel embeddingModel = new PostgresMlEmbeddingModel(jdbcTemplate); - PostgresMlEmbeddingOptions options = embeddingClient.mergeOptions(EmbeddingOptions.EMPTY); + PostgresMlEmbeddingOptions options = embeddingModel.mergeOptions(EmbeddingOptions.EMPTY); // Default options - assertThat(options.getTransformer()).isEqualTo(PostgresMlEmbeddingClient.DEFAULT_TRANSFORMER_MODEL); - assertThat(options.getVectorType()).isEqualTo(PostgresMlEmbeddingClient.VectorType.PG_ARRAY); + assertThat(options.getTransformer()).isEqualTo(PostgresMlEmbeddingModel.DEFAULT_TRANSFORMER_MODEL); + assertThat(options.getVectorType()).isEqualTo(PostgresMlEmbeddingModel.VectorType.PG_ARRAY); assertThat(options.getKwargs()).isEqualTo(Map.of()); assertThat(options.getMetadataMode()).isEqualTo(org.springframework.ai.document.MetadataMode.EMBED); // Partial override - options = embeddingClient.mergeOptions(PostgresMlEmbeddingOptions.builder() + options = embeddingModel.mergeOptions(PostgresMlEmbeddingOptions.builder() .withTransformer("intfloat/e5-small") .withKwargs(Map.of("device", "cpu")) .build()); assertThat(options.getTransformer()).isEqualTo("intfloat/e5-small"); - assertThat(options.getVectorType()).isEqualTo(PostgresMlEmbeddingClient.VectorType.PG_ARRAY); // Default + assertThat(options.getVectorType()).isEqualTo(PostgresMlEmbeddingModel.VectorType.PG_ARRAY); // Default assertThat(options.getKwargs()).isEqualTo(Map.of("device", "cpu")); assertThat(options.getMetadataMode()).isEqualTo(org.springframework.ai.document.MetadataMode.EMBED); // Default // Complete override - options = embeddingClient.mergeOptions(PostgresMlEmbeddingOptions.builder() + options = embeddingModel.mergeOptions(PostgresMlEmbeddingOptions.builder() .withTransformer("intfloat/e5-small") - .withVectorType(PostgresMlEmbeddingClient.VectorType.PG_VECTOR) + .withVectorType(PostgresMlEmbeddingModel.VectorType.PG_VECTOR) .withMetadataMode(org.springframework.ai.document.MetadataMode.ALL) .withKwargs(Map.of("device", "cpu")) .build()); assertThat(options.getTransformer()).isEqualTo("intfloat/e5-small"); - assertThat(options.getVectorType()).isEqualTo(PostgresMlEmbeddingClient.VectorType.PG_VECTOR); + assertThat(options.getVectorType()).isEqualTo(PostgresMlEmbeddingModel.VectorType.PG_VECTOR); assertThat(options.getKwargs()).isEqualTo(Map.of("device", "cpu")); assertThat(options.getMetadataMode()).isEqualTo(org.springframework.ai.document.MetadataMode.ALL); } diff --git a/models/spring-ai-stability-ai/src/main/java/org/springframework/ai/stabilityai/StabilityAiImageClient.java b/models/spring-ai-stability-ai/src/main/java/org/springframework/ai/stabilityai/StabilityAiImageModel.java similarity index 90% rename from models/spring-ai-stability-ai/src/main/java/org/springframework/ai/stabilityai/StabilityAiImageClient.java rename to models/spring-ai-stability-ai/src/main/java/org/springframework/ai/stabilityai/StabilityAiImageModel.java index e632e5235..abb52c9a9 100644 --- a/models/spring-ai-stability-ai/src/main/java/org/springframework/ai/stabilityai/StabilityAiImageClient.java +++ b/models/spring-ai-stability-ai/src/main/java/org/springframework/ai/stabilityai/StabilityAiImageModel.java @@ -22,7 +22,7 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.image.Image; -import org.springframework.ai.image.ImageClient; +import org.springframework.ai.image.ImageModel; import org.springframework.ai.image.ImageGeneration; import org.springframework.ai.image.ImageOptions; import org.springframework.ai.image.ImagePrompt; @@ -34,10 +34,10 @@ import org.springframework.ai.stabilityai.api.StabilityAiImageOptions; import org.springframework.util.Assert; /** - * StabilityAiImageClient is a class that implements the ImageClient interface. It - * provides a client for calling the StabilityAI image generation API. + * StabilityAiImageModel is a class that implements the ImageModel interface. It provides + * a client for calling the StabilityAI image generation API. */ -public class StabilityAiImageClient implements ImageClient { +public class StabilityAiImageModel implements ImageModel { private final Logger logger = LoggerFactory.getLogger(getClass()); @@ -45,11 +45,11 @@ public class StabilityAiImageClient implements ImageClient { private final StabilityAiApi stabilityAiApi; - public StabilityAiImageClient(StabilityAiApi stabilityAiApi) { + public StabilityAiImageModel(StabilityAiApi stabilityAiApi) { this(stabilityAiApi, StabilityAiImageOptions.builder().build()); } - public StabilityAiImageClient(StabilityAiApi stabilityAiApi, StabilityAiImageOptions options) { + public StabilityAiImageModel(StabilityAiApi stabilityAiApi, StabilityAiImageOptions options) { Assert.notNull(stabilityAiApi, "StabilityAiApi must not be null"); Assert.notNull(options, "StabilityAiImageOptions must not be null"); this.stabilityAiApi = stabilityAiApi; @@ -61,20 +61,20 @@ public class StabilityAiImageClient implements ImageClient { } /** - * Calls the StabilityAiImageClient with the given StabilityAiImagePrompt and returns + * Calls the StabilityAiImageModel with the given StabilityAiImagePrompt and returns * the ImageResponse. This overloaded call method lets you pass the full set of Prompt * instructions that StabilityAI supports. * @param imagePrompt the StabilityAiImagePrompt containing the prompt and image model * options - * @return the ImageResponse generated by the StabilityAiImageClient + * @return the ImageResponse generated by the StabilityAiImageModel */ public ImageResponse call(ImagePrompt imagePrompt) { ImageOptions runtimeOptions = imagePrompt.getOptions(); - // Merge the runtime options passed via the prompt with the StabilityAiImageClient + // Merge the runtime options passed via the prompt with the StabilityAiImageModel // options configured via Autoconfiguration. - // Runtime options overwrite StabilityAiImageClient options + // Runtime options overwrite StabilityAiImageModel options StabilityAiImageOptions optionsToUse = ModelOptionsUtils.merge(runtimeOptions, this.options, StabilityAiImageOptions.class); diff --git a/models/spring-ai-stability-ai/src/test/java/org/springframework/ai/stabilityai/StabilityAiImageClientIT.java b/models/spring-ai-stability-ai/src/test/java/org/springframework/ai/stabilityai/StabilityAiImageModelIT.java similarity index 91% rename from models/spring-ai-stability-ai/src/test/java/org/springframework/ai/stabilityai/StabilityAiImageClientIT.java rename to models/spring-ai-stability-ai/src/test/java/org/springframework/ai/stabilityai/StabilityAiImageModelIT.java index afa4b0d55..b2de03a19 100644 --- a/models/spring-ai-stability-ai/src/test/java/org/springframework/ai/stabilityai/StabilityAiImageClientIT.java +++ b/models/spring-ai-stability-ai/src/test/java/org/springframework/ai/stabilityai/StabilityAiImageModelIT.java @@ -19,7 +19,7 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.image.Image; -import org.springframework.ai.image.ImageClient; +import org.springframework.ai.image.ImageModel; import org.springframework.ai.image.ImageGeneration; import org.springframework.ai.image.ImagePrompt; import org.springframework.ai.image.ImageResponse; @@ -36,10 +36,10 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest(classes = StabilityAiImageTestConfiguration.class) @EnabledIfEnvironmentVariable(named = "STABILITYAI_API_KEY", matches = ".*") -public class StabilityAiImageClientIT { +public class StabilityAiImageModelIT { @Autowired - protected ImageClient stabilityAiImageClient; + protected ImageModel stabilityAiImageModel; @Test void imageAsBase64Test() throws IOException { @@ -54,7 +54,7 @@ public class StabilityAiImageClientIT { ImagePrompt imagePrompt = new ImagePrompt(instructions, imageOptions); - ImageResponse imageResponse = this.stabilityAiImageClient.call(imagePrompt); + ImageResponse imageResponse = this.stabilityAiImageModel.call(imagePrompt); ImageGeneration imageGeneration = imageResponse.getResult(); Image image = imageGeneration.getOutput(); diff --git a/models/spring-ai-stability-ai/src/test/java/org/springframework/ai/stabilityai/StabilityAiImageTestConfiguration.java b/models/spring-ai-stability-ai/src/test/java/org/springframework/ai/stabilityai/StabilityAiImageTestConfiguration.java index 394d384ff..c5271ff00 100644 --- a/models/spring-ai-stability-ai/src/test/java/org/springframework/ai/stabilityai/StabilityAiImageTestConfiguration.java +++ b/models/spring-ai-stability-ai/src/test/java/org/springframework/ai/stabilityai/StabilityAiImageTestConfiguration.java @@ -29,8 +29,8 @@ public class StabilityAiImageTestConfiguration { } @Bean - StabilityAiImageClient stabilityAiImageClient(StabilityAiApi stabilityAiApi) { - return new StabilityAiImageClient(stabilityAiApi); + StabilityAiImageModel stabilityAiImageModel(StabilityAiApi stabilityAiApi) { + return new StabilityAiImageModel(stabilityAiApi); } private String getApiKey() { diff --git a/models/spring-ai-transformers/src/main/java/org/springframework/ai/transformers/TransformersEmbeddingClient.java b/models/spring-ai-transformers/src/main/java/org/springframework/ai/transformers/TransformersEmbeddingModel.java similarity index 97% rename from models/spring-ai-transformers/src/main/java/org/springframework/ai/transformers/TransformersEmbeddingClient.java rename to models/spring-ai-transformers/src/main/java/org/springframework/ai/transformers/TransformersEmbeddingModel.java index d5b1d1933..20d24f8a6 100644 --- a/models/spring-ai-transformers/src/main/java/org/springframework/ai/transformers/TransformersEmbeddingClient.java +++ b/models/spring-ai-transformers/src/main/java/org/springframework/ai/transformers/TransformersEmbeddingModel.java @@ -40,7 +40,7 @@ import org.apache.commons.logging.LogFactory; import org.springframework.ai.document.Document; import org.springframework.ai.document.MetadataMode; -import org.springframework.ai.embedding.AbstractEmbeddingClient; +import org.springframework.ai.embedding.AbstractEmbeddingModel; import org.springframework.ai.embedding.Embedding; import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingRequest; @@ -56,9 +56,9 @@ import org.springframework.util.StringUtils; * * @author Christian Tzolov */ -public class TransformersEmbeddingClient extends AbstractEmbeddingClient implements InitializingBean { +public class TransformersEmbeddingModel extends AbstractEmbeddingModel implements InitializingBean { - private static final Log logger = LogFactory.getLog(TransformersEmbeddingClient.class); + private static final Log logger = LogFactory.getLog(TransformersEmbeddingModel.class); // ONNX tokenizer for the all-MiniLM-L6-v2 generative public final static String DEFAULT_ONNX_TOKENIZER_URI = "https://raw.githubusercontent.com/spring-projects/spring-ai/main/models/spring-ai-transformers/src/main/resources/onnx/all-MiniLM-L6-v2/tokenizer.json"; @@ -126,11 +126,11 @@ public class TransformersEmbeddingClient extends AbstractEmbeddingClient impleme private Set onnxModelInputs; - public TransformersEmbeddingClient() { + public TransformersEmbeddingModel() { this(MetadataMode.NONE); } - public TransformersEmbeddingClient(MetadataMode metadataMode) { + public TransformersEmbeddingModel(MetadataMode metadataMode) { Assert.notNull(metadataMode, "Metadata mode should not be null"); this.metadataMode = metadataMode; } diff --git a/models/spring-ai-transformers/src/test/java/org/springframework/ai/transformers/TransformersEmbeddingClientTests.java b/models/spring-ai-transformers/src/test/java/org/springframework/ai/transformers/TransformersEmbeddingModelTests.java similarity index 72% rename from models/spring-ai-transformers/src/test/java/org/springframework/ai/transformers/TransformersEmbeddingClientTests.java rename to models/spring-ai-transformers/src/test/java/org/springframework/ai/transformers/TransformersEmbeddingModelTests.java index d2526bbe8..023496eff 100644 --- a/models/spring-ai-transformers/src/test/java/org/springframework/ai/transformers/TransformersEmbeddingClientTests.java +++ b/models/spring-ai-transformers/src/test/java/org/springframework/ai/transformers/TransformersEmbeddingModelTests.java @@ -28,16 +28,16 @@ import static org.assertj.core.api.Assertions.assertThat; /** * @author Christian Tzolov */ -public class TransformersEmbeddingClientTests { +public class TransformersEmbeddingModelTests { private static DecimalFormat DF = new DecimalFormat("#.######"); @Test void embed() throws Exception { - TransformersEmbeddingClient embeddingClient = new TransformersEmbeddingClient(); - embeddingClient.afterPropertiesSet(); - List embed = embeddingClient.embed("Hello world"); + TransformersEmbeddingModel embeddingModel = new TransformersEmbeddingModel(); + embeddingModel.afterPropertiesSet(); + List embed = embeddingModel.embed("Hello world"); assertThat(embed).hasSize(384); assertThat(DF.format(embed.get(0))).isEqualTo(DF.format(-0.19744634628295898)); assertThat(DF.format(embed.get(383))).isEqualTo(DF.format(0.17298996448516846)); @@ -45,9 +45,9 @@ public class TransformersEmbeddingClientTests { @Test void embedDocument() throws Exception { - TransformersEmbeddingClient embeddingClient = new TransformersEmbeddingClient(); - embeddingClient.afterPropertiesSet(); - List embed = embeddingClient.embed(new Document("Hello world")); + TransformersEmbeddingModel embeddingModel = new TransformersEmbeddingModel(); + embeddingModel.afterPropertiesSet(); + List embed = embeddingModel.embed(new Document("Hello world")); assertThat(embed).hasSize(384); assertThat(DF.format(embed.get(0))).isEqualTo(DF.format(-0.19744634628295898)); assertThat(DF.format(embed.get(383))).isEqualTo(DF.format(0.17298996448516846)); @@ -55,9 +55,9 @@ public class TransformersEmbeddingClientTests { @Test void embedList() throws Exception { - TransformersEmbeddingClient embeddingClient = new TransformersEmbeddingClient(); - embeddingClient.afterPropertiesSet(); - List> embed = embeddingClient.embed(List.of("Hello world", "World is big")); + TransformersEmbeddingModel embeddingModel = new TransformersEmbeddingModel(); + embeddingModel.afterPropertiesSet(); + List> embed = embeddingModel.embed(List.of("Hello world", "World is big")); assertThat(embed).hasSize(2); assertThat(embed.get(0)).hasSize(384); assertThat(DF.format(embed.get(0).get(0))).isEqualTo(DF.format(-0.19744634628295898)); @@ -72,9 +72,9 @@ public class TransformersEmbeddingClientTests { @Test void embedForResponse() throws Exception { - TransformersEmbeddingClient embeddingClient = new TransformersEmbeddingClient(); - embeddingClient.afterPropertiesSet(); - EmbeddingResponse embed = embeddingClient.embedForResponse(List.of("Hello world", "World is big")); + TransformersEmbeddingModel embeddingModel = new TransformersEmbeddingModel(); + embeddingModel.afterPropertiesSet(); + EmbeddingResponse embed = embeddingModel.embedForResponse(List.of("Hello world", "World is big")); assertThat(embed.getResults()).hasSize(2); assertThat(embed.getMetadata()).isEmpty(); @@ -90,11 +90,11 @@ public class TransformersEmbeddingClientTests { @Test void dimensions() throws Exception { - TransformersEmbeddingClient embeddingClient = new TransformersEmbeddingClient(); - embeddingClient.afterPropertiesSet(); - assertThat(embeddingClient.dimensions()).isEqualTo(384); + TransformersEmbeddingModel embeddingModel = new TransformersEmbeddingModel(); + embeddingModel.afterPropertiesSet(); + assertThat(embeddingModel.dimensions()).isEqualTo(384); // cached - assertThat(embeddingClient.dimensions()).isEqualTo(384); + assertThat(embeddingModel.dimensions()).isEqualTo(384); } } diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClient.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java similarity index 95% rename from models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClient.java rename to models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java index e30df03ce..80d668f0c 100644 --- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClient.java +++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java @@ -32,16 +32,17 @@ import com.google.cloud.vertexai.generativeai.PartMaker; import com.google.cloud.vertexai.generativeai.ResponseStream; import com.google.protobuf.Struct; import com.google.protobuf.util.JsonFormat; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.MessageType; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.model.ModelDescription; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.model.function.AbstractFunctionCallSupport; import org.springframework.ai.model.function.FunctionCallbackContext; @@ -65,9 +66,9 @@ import java.util.stream.Collectors; * @author Grogdunn * @since 0.8.1 */ -public class VertexAiGeminiChatClient - extends AbstractFunctionCallSupport - implements ChatClient, StreamingChatClient, DisposableBean { +public class VertexAiGeminiChatModel + extends AbstractFunctionCallSupport + implements ChatModel, StreamingChatModel, DisposableBean { private final static boolean IS_RUNTIME_CALL = true; @@ -95,7 +96,7 @@ public class VertexAiGeminiChatClient } - public enum ChatModel { + public enum ChatModel implements ModelDescription { GEMINI_PRO_VISION("gemini-pro-vision"), @@ -115,9 +116,14 @@ public class VertexAiGeminiChatClient return this.value; } + @Override + public String getModelName() { + return this.value; + } + } - public VertexAiGeminiChatClient(VertexAI vertexAI) { + public VertexAiGeminiChatModel(VertexAI vertexAI) { this(vertexAI, VertexAiGeminiChatOptions.builder() .withModel(ChatModel.GEMINI_PRO_VISION) @@ -125,11 +131,11 @@ public class VertexAiGeminiChatClient .build()); } - public VertexAiGeminiChatClient(VertexAI vertexAI, VertexAiGeminiChatOptions options) { + public VertexAiGeminiChatModel(VertexAI vertexAI, VertexAiGeminiChatOptions options) { this(vertexAI, options, null); } - public VertexAiGeminiChatClient(VertexAI vertexAI, VertexAiGeminiChatOptions options, + public VertexAiGeminiChatModel(VertexAI vertexAI, VertexAiGeminiChatOptions options, FunctionCallbackContext functionCallbackContext) { super(functionCallbackContext); @@ -478,4 +484,9 @@ public class VertexAiGeminiChatClient return response.getCandidatesList().get(0).getContent().getPartsList().get(0).hasFunctionCall(); } + @Override + public ChatOptions getDefaultOptions() { + return VertexAiGeminiChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatOptions.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatOptions.java index 311ff16af..6a83f4a00 100644 --- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatOptions.java +++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatOptions.java @@ -28,7 +28,7 @@ import com.fasterxml.jackson.annotation.JsonProperty; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallingOptions; -import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatClient.ChatModel; +import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatModel.ChatModel; import org.springframework.boot.context.properties.NestedConfigurationProperty; import org.springframework.util.Assert; @@ -78,10 +78,10 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions, ChatOp private @JsonProperty("modelName") String model; /** - * Tool Function Callbacks to register with the ChatClient. + * Tool Function Callbacks to register with the ChatModel. * For Prompt Options the functionCallbacks are automatically enabled for the duration of the prompt execution. * For Default Options the functionCallbacks are registered but disabled by default. Use the enableFunctions to set the functions - * from the registry to be used by the ChatClient chat completion requests. + * from the registry to be used by the ChatModel chat completion requests. */ @NestedConfigurationProperty @JsonIgnore @@ -336,4 +336,18 @@ public class VertexAiGeminiChatOptions implements FunctionCallingOptions, ChatOp return true; } + public static VertexAiGeminiChatOptions fromOptions(VertexAiGeminiChatOptions fromOptions) { + VertexAiGeminiChatOptions options = new VertexAiGeminiChatOptions(); + options.setStopSequences(fromOptions.getStopSequences()); + options.setTemperature(fromOptions.getTemperature()); + options.setTopP(fromOptions.getTopP()); + options.setTopK(fromOptions.getTopK()); + options.setCandidateCount(fromOptions.getCandidateCount()); + options.setMaxOutputTokens(fromOptions.getMaxOutputTokens()); + options.setModel(fromOptions.getModel()); + options.setFunctionCallbacks(fromOptions.getFunctionCallbacks()); + options.setFunctions(fromOptions.getFunctions()); + return options; + } + } diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/aot/VertexAiGeminiRuntimeHints.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/aot/VertexAiGeminiRuntimeHints.java index 8911288fa..0a46b9f2f 100644 --- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/aot/VertexAiGeminiRuntimeHints.java +++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/aot/VertexAiGeminiRuntimeHints.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.vertexai.gemini.aot; -import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatClient; +import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatModel; import org.springframework.aot.hint.MemberCategory; import org.springframework.aot.hint.RuntimeHints; import org.springframework.aot.hint.RuntimeHintsRegistrar; @@ -34,7 +34,7 @@ public class VertexAiGeminiRuntimeHints implements RuntimeHintsRegistrar { @Override public void registerHints(RuntimeHints hints, ClassLoader classLoader) { var mcs = MemberCategory.values(); - for (var tr : findJsonAnnotatedClassesInPackage(VertexAiGeminiChatClient.class)) + for (var tr : findJsonAnnotatedClassesInPackage(VertexAiGeminiChatModel.class)) hints.reflection().registerType(tr, mcs); } diff --git a/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClientIT.java b/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModelIT.java similarity index 91% rename from models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClientIT.java rename to models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModelIT.java index 58f687eb4..5e81d683b 100644 --- a/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClientIT.java +++ b/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModelIT.java @@ -53,10 +53,10 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @EnabledIfEnvironmentVariable(named = "VERTEX_AI_GEMINI_PROJECT_ID", matches = ".*") @EnabledIfEnvironmentVariable(named = "VERTEX_AI_GEMINI_LOCATION", matches = ".*") -class VertexAiGeminiChatClientIT { +class VertexAiGeminiChatModelIT { @Autowired - private VertexAiGeminiChatClient client; + private VertexAiGeminiChatModel chatModel; @Value("classpath:/prompts/system-message.st") private Resource systemResource; @@ -70,7 +70,7 @@ class VertexAiGeminiChatClientIT { SystemPromptTemplate systemPromptTemplate = new SystemPromptTemplate(systemResource); Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", name, "voice", voice)); Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); - ChatResponse response = client.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); } @@ -87,7 +87,7 @@ class VertexAiGeminiChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "ice cream flavors.", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = this.client.call(prompt).getResult(); + Generation generation = this.chatModel.call(prompt).getResult(); List list = outputParser.convert(generation.getOutput().getContent()); assertThat(list).hasSize(5); @@ -106,7 +106,7 @@ class VertexAiGeminiChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); Map result = outputConverter.convert(generation.getOutput().getContent()); assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); @@ -129,7 +129,7 @@ class VertexAiGeminiChatClientIT { """; PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); ActorsFilmsRecord actorsFilms = outputConvert.convert(generation.getOutput().getContent()); assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks"); @@ -139,7 +139,7 @@ class VertexAiGeminiChatClientIT { @Test void textStream() { - String generationTextFromStream = client.stream(new Prompt("Explain Bulgaria? Answer in 10 paragraphs.")) + String generationTextFromStream = chatModel.stream(new Prompt("Explain Bulgaria? Answer in 10 paragraphs.")) .collectList() .block() .stream() @@ -167,7 +167,7 @@ class VertexAiGeminiChatClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - String generationTextFromStream = client.stream(prompt) + String generationTextFromStream = chatModel.stream(prompt) .collectList() .block() .stream() @@ -191,7 +191,7 @@ class VertexAiGeminiChatClientIT { var userMessage = new UserMessage("Explain what do you see o this picture?", List.of(new Media(MimeTypeUtils.IMAGE_PNG, data))); - var response = client.call(new Prompt(List.of(userMessage))); + var response = chatModel.call(new Prompt(List.of(userMessage))); // Response should contain something like: // I see a bunch of bananas in a golden basket. The bananas are ripe and yellow. @@ -231,10 +231,10 @@ class VertexAiGeminiChatClientIT { } @Bean - public VertexAiGeminiChatClient vertexAiEmbedding(VertexAI vertexAi) { - return new VertexAiGeminiChatClient(vertexAi, + public VertexAiGeminiChatModel vertexAiEmbedding(VertexAI vertexAi) { + return new VertexAiGeminiChatModel(vertexAi, VertexAiGeminiChatOptions.builder() - .withModel(VertexAiGeminiChatClient.ChatModel.GEMINI_PRO_VISION) + .withModel(VertexAiGeminiChatModel.ChatModel.GEMINI_PRO_VISION) .build()); } diff --git a/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/aot/VertexAiGeminiRuntimeHintsTests.java b/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/aot/VertexAiGeminiRuntimeHintsTests.java index 2e3f12f12..a4aaf3988 100644 --- a/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/aot/VertexAiGeminiRuntimeHintsTests.java +++ b/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/aot/VertexAiGeminiRuntimeHintsTests.java @@ -19,7 +19,7 @@ import java.util.Set; import org.junit.jupiter.api.Test; -import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatClient; +import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatModel; import org.springframework.aot.hint.RuntimeHints; import org.springframework.aot.hint.TypeReference; @@ -38,7 +38,7 @@ class VertexAiGeminiRuntimeHintsTests { RuntimeHints runtimeHints = new RuntimeHints(); VertexAiGeminiRuntimeHints vertexAiGeminiRuntimeHints = new VertexAiGeminiRuntimeHints(); vertexAiGeminiRuntimeHints.registerHints(runtimeHints, null); - Set jsonAnnotatedClasses = findJsonAnnotatedClassesInPackage(VertexAiGeminiChatClient.class); + Set jsonAnnotatedClasses = findJsonAnnotatedClassesInPackage(VertexAiGeminiChatModel.class); for (TypeReference jsonAnnotatedClass : jsonAnnotatedClasses) { assertThat(runtimeHints).matches(reflection().onType(jsonAnnotatedClass)); } diff --git a/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/function/VertexAiGeminiChatClientFunctionCallingIT.java b/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/function/VertexAiGeminiChatModelFunctionCallingIT.java similarity index 92% rename from models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/function/VertexAiGeminiChatClientFunctionCallingIT.java rename to models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/function/VertexAiGeminiChatModelFunctionCallingIT.java index 11b852222..e71d58bc9 100644 --- a/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/function/VertexAiGeminiChatClientFunctionCallingIT.java +++ b/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/function/VertexAiGeminiChatModelFunctionCallingIT.java @@ -28,6 +28,7 @@ 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.vertexai.gemini.VertexAiGeminiChatModel; import reactor.core.publisher.Flux; import org.springframework.ai.chat.ChatResponse; @@ -38,7 +39,6 @@ import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.function.FunctionCallbackWrapper; import org.springframework.ai.model.function.FunctionCallbackWrapper.Builder.SchemaType; -import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatClient; import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatOptions; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.SpringBootConfiguration; @@ -50,12 +50,12 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @EnabledIfEnvironmentVariable(named = "VERTEX_AI_GEMINI_PROJECT_ID", matches = ".*") @EnabledIfEnvironmentVariable(named = "VERTEX_AI_GEMINI_LOCATION", matches = ".*") -public class VertexAiGeminiChatClientFunctionCallingIT { +public class VertexAiGeminiChatModelFunctionCallingIT { private final Logger logger = LoggerFactory.getLogger(getClass()); @Autowired - private VertexAiGeminiChatClient vertexGeminiClient; + private VertexAiGeminiChatModel vertexGeminiClient; @AfterEach public void afterEach() { @@ -98,8 +98,8 @@ public class VertexAiGeminiChatClientFunctionCallingIT { """; var promptOptions = VertexAiGeminiChatOptions.builder() - .withModel(VertexAiGeminiChatClient.ChatModel.GEMINI_PRO) - // .withModel(VertexAiGeminiChatClient.ChatModel.GEMINI_PRO_1_5_PRO) + .withModel(VertexAiGeminiChatModel.ChatModel.GEMINI_PRO) + // .withModel(VertexAiGeminiModelCall.ChatModel.GEMINI_PRO_1_5_PRO) .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) .withName("get_current_weather") .withDescription("Get the current weather in a given location") @@ -126,8 +126,8 @@ public class VertexAiGeminiChatClientFunctionCallingIT { List messages = new ArrayList<>(List.of(userMessage)); var promptOptions = VertexAiGeminiChatOptions.builder() - // .withModel(VertexAiGeminiChatClient.ChatModel.GEMINI_PRO_1_5_PRO) - .withModel(VertexAiGeminiChatClient.ChatModel.GEMINI_PRO.getValue()) + // .withModel(VertexAiGeminiModelCall.ChatModel.GEMINI_PRO_1_5_PRO) + .withModel(VertexAiGeminiChatModel.ChatModel.GEMINI_PRO.getValue()) .withFunctionCallbacks(List.of( FunctionCallbackWrapper.builder(new MockWeatherService()) .withSchemaType(SchemaType.OPEN_API_SCHEMA) @@ -168,7 +168,7 @@ public class VertexAiGeminiChatClientFunctionCallingIT { List messages = new ArrayList<>(List.of(userMessage)); var promptOptions = VertexAiGeminiChatOptions.builder() - .withModel(VertexAiGeminiChatClient.ChatModel.GEMINI_PRO) + .withModel(VertexAiGeminiChatModel.ChatModel.GEMINI_PRO) .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) .withSchemaType(SchemaType.OPEN_API_SCHEMA) .withName("getCurrentWeather") @@ -224,10 +224,10 @@ public class VertexAiGeminiChatClientFunctionCallingIT { } @Bean - public VertexAiGeminiChatClient vertexAiEmbedding(VertexAI vertexAi) { - return new VertexAiGeminiChatClient(vertexAi, + public VertexAiGeminiChatModel vertexAiEmbedding(VertexAI vertexAi) { + return new VertexAiGeminiChatModel(vertexAi, VertexAiGeminiChatOptions.builder() - .withModel(VertexAiGeminiChatClient.ChatModel.GEMINI_PRO) + .withModel(VertexAiGeminiChatModel.ChatModel.GEMINI_PRO) .withTemperature(0.9f) .build()); } diff --git a/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatClient.java b/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatModel.java similarity index 90% rename from models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatClient.java rename to models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatModel.java index 671097615..7b9c00955 100644 --- a/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatClient.java +++ b/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatModel.java @@ -18,7 +18,7 @@ package org.springframework.ai.vertexai.palm2; import java.util.List; import java.util.stream.Collectors; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; @@ -35,18 +35,18 @@ import org.springframework.util.CollectionUtils; /** * @author Christian Tzolov */ -public class VertexAiPaLm2ChatClient implements ChatClient { +public class VertexAiPaLm2ChatModel implements ChatModel { private final VertexAiPaLm2Api vertexAiApi; private final VertexAiPaLm2ChatOptions defaultOptions; - public VertexAiPaLm2ChatClient(VertexAiPaLm2Api vertexAiApi) { + public VertexAiPaLm2ChatModel(VertexAiPaLm2Api vertexAiApi) { this(vertexAiApi, VertexAiPaLm2ChatOptions.builder().withTemperature(0.7f).withCandidateCount(1).withTopK(20).build()); } - public VertexAiPaLm2ChatClient(VertexAiPaLm2Api vertexAiApi, VertexAiPaLm2ChatOptions defaultOptions) { + public VertexAiPaLm2ChatModel(VertexAiPaLm2Api vertexAiApi, VertexAiPaLm2ChatOptions defaultOptions) { Assert.notNull(defaultOptions, "Default options must not be null!"); Assert.notNull(vertexAiApi, "VertexAiPaLm2Api must not be null!"); @@ -111,4 +111,9 @@ public class VertexAiPaLm2ChatClient implements ChatClient { return request; } + @Override + public ChatOptions getDefaultOptions() { + return VertexAiPaLm2ChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatOptions.java b/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatOptions.java index cae09e7bb..55742ff05 100644 --- a/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatOptions.java +++ b/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatOptions.java @@ -127,4 +127,13 @@ public class VertexAiPaLm2ChatOptions implements ChatOptions { this.topK = topK; } + public static VertexAiPaLm2ChatOptions fromOptions(VertexAiPaLm2ChatOptions fromOptions) { + return VertexAiPaLm2ChatOptions.builder() + .withTemperature(fromOptions.getTemperature()) + .withCandidateCount(fromOptions.getCandidateCount()) + .withTopP(fromOptions.getTopP()) + .withTopK(fromOptions.getTopK()) + .build(); + } + } diff --git a/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2EmbeddingClient.java b/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2EmbeddingModel.java similarity index 88% rename from models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2EmbeddingClient.java rename to models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2EmbeddingModel.java index b864c8b4c..59ee71cd1 100644 --- a/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2EmbeddingClient.java +++ b/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2EmbeddingModel.java @@ -19,7 +19,7 @@ import java.util.List; import java.util.concurrent.atomic.AtomicInteger; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.AbstractEmbeddingClient; +import org.springframework.ai.embedding.AbstractEmbeddingModel; import org.springframework.ai.embedding.Embedding; import org.springframework.ai.embedding.EmbeddingRequest; import org.springframework.ai.embedding.EmbeddingResponse; @@ -28,11 +28,11 @@ import org.springframework.ai.vertexai.palm2.api.VertexAiPaLm2Api; /** * @author Christian Tzolov */ -public class VertexAiPaLm2EmbeddingClient extends AbstractEmbeddingClient { +public class VertexAiPaLm2EmbeddingModel extends AbstractEmbeddingModel { private final VertexAiPaLm2Api vertexAiApi; - public VertexAiPaLm2EmbeddingClient(VertexAiPaLm2Api vertexAiApi) { + public VertexAiPaLm2EmbeddingModel(VertexAiPaLm2Api vertexAiApi) { this.vertexAiApi = vertexAiApi; } diff --git a/models/spring-ai-vertex-ai-palm2/src/test/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatGenerationClientIT.java b/models/spring-ai-vertex-ai-palm2/src/test/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatGenerationClientIT.java index 5fd43b0d5..4039969fd 100644 --- a/models/spring-ai-vertex-ai-palm2/src/test/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatGenerationClientIT.java +++ b/models/spring-ai-vertex-ai-palm2/src/test/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatGenerationClientIT.java @@ -48,7 +48,7 @@ import static org.assertj.core.api.Assertions.assertThat; class VertexAiPaLm2ChatGenerationClientIT { @Autowired - private VertexAiPaLm2ChatClient client; + private VertexAiPaLm2ChatModel chatModel; @Value("classpath:/prompts/system-message.st") private Resource systemResource; @@ -62,7 +62,7 @@ class VertexAiPaLm2ChatGenerationClientIT { SystemPromptTemplate systemPromptTemplate = new SystemPromptTemplate(systemResource); Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", name, "voice", voice)); Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); - ChatResponse response = client.call(prompt); + ChatResponse response = chatModel.call(prompt); assertThat(response.getResult().getOutput().getContent()).contains("Bartholomew"); } @@ -79,7 +79,7 @@ class VertexAiPaLm2ChatGenerationClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "ice cream flavors.", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = this.client.call(prompt).getResult(); + Generation generation = this.chatModel.call(prompt).getResult(); List list = outputConverter.convert(generation.getOutput().getContent()); assertThat(list).hasSize(5); @@ -98,7 +98,7 @@ class VertexAiPaLm2ChatGenerationClientIT { PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); Map result = outputConverter.convert(generation.getOutput().getContent()); assertThat(result.get("numbers")).isEqualTo(Arrays.asList(1, 2, 3, 4, 5, 6, 7, 8, 9)); @@ -120,7 +120,7 @@ class VertexAiPaLm2ChatGenerationClientIT { """; PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); - Generation generation = client.call(prompt).getResult(); + Generation generation = chatModel.call(prompt).getResult(); ActorsFilmsRecord actorsFilms = outputConverter.convert(generation.getOutput().getContent()); assertThat(actorsFilms.actor()).isEqualTo("Tom Hanks"); @@ -136,8 +136,8 @@ class VertexAiPaLm2ChatGenerationClientIT { } @Bean - public VertexAiPaLm2ChatClient vertexAiEmbedding(VertexAiPaLm2Api vertexAiApi) { - return new VertexAiPaLm2ChatClient(vertexAiApi); + public VertexAiPaLm2ChatModel vertexAiEmbedding(VertexAiPaLm2Api vertexAiApi) { + return new VertexAiPaLm2ChatModel(vertexAiApi); } } diff --git a/models/spring-ai-vertex-ai-palm2/src/test/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatRequestTests.java b/models/spring-ai-vertex-ai-palm2/src/test/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatRequestTests.java index 6c5478f67..81a07a150 100644 --- a/models/spring-ai-vertex-ai-palm2/src/test/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatRequestTests.java +++ b/models/spring-ai-vertex-ai-palm2/src/test/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatRequestTests.java @@ -29,12 +29,12 @@ import static org.assertj.core.api.Assertions.assertThat; */ public class VertexAiPaLm2ChatRequestTests { - VertexAiPaLm2ChatClient client = new VertexAiPaLm2ChatClient(new VertexAiPaLm2Api("bla")); + VertexAiPaLm2ChatModel chatModel = new VertexAiPaLm2ChatModel(new VertexAiPaLm2Api("bla")); @Test public void createRequestWithDefaultOptions() { - var request = client.createRequest(new Prompt("Test message content")); + var request = chatModel.createRequest(new Prompt("Test message content")); assertThat(request.prompt().messages()).hasSize(1); @@ -55,7 +55,7 @@ public class VertexAiPaLm2ChatRequestTests { // .withCandidateCount(2) .build(); - var request = client.createRequest(new Prompt("Test message content", promptOptions)); + var request = chatModel.createRequest(new Prompt("Test message content", promptOptions)); assertThat(request.prompt().messages()).hasSize(1); @@ -75,7 +75,7 @@ public class VertexAiPaLm2ChatRequestTests { .withTopP(0.6f) .build(); - var request = client.createRequest(new Prompt("Test message content", portablePromptOptions)); + var request = chatModel.createRequest(new Prompt("Test message content", portablePromptOptions)); assertThat(request.prompt().messages()).hasSize(1); diff --git a/models/spring-ai-vertex-ai-palm2/src/test/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2EmbeddingClientIT.java b/models/spring-ai-vertex-ai-palm2/src/test/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2EmbeddingModelIT.java similarity index 78% rename from models/spring-ai-vertex-ai-palm2/src/test/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2EmbeddingClientIT.java rename to models/spring-ai-vertex-ai-palm2/src/test/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2EmbeddingModelIT.java index 6014d5b0e..2e05dbfd1 100644 --- a/models/spring-ai-vertex-ai-palm2/src/test/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2EmbeddingClientIT.java +++ b/models/spring-ai-vertex-ai-palm2/src/test/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2EmbeddingModelIT.java @@ -31,24 +31,24 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest @EnabledIfEnvironmentVariable(named = "PALM_API_KEY", matches = ".*") -class VertexAiPaLm2EmbeddingClientIT { +class VertexAiPaLm2EmbeddingModelIT { @Autowired - private VertexAiPaLm2EmbeddingClient embeddingClient; + private VertexAiPaLm2EmbeddingModel embeddingModel; @Test void simpleEmbedding() { - assertThat(embeddingClient).isNotNull(); - EmbeddingResponse embeddingResponse = embeddingClient.embedForResponse(List.of("Hello World")); + assertThat(embeddingModel).isNotNull(); + EmbeddingResponse embeddingResponse = embeddingModel.embedForResponse(List.of("Hello World")); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); - assertThat(embeddingClient.dimensions()).isEqualTo(768); + assertThat(embeddingModel.dimensions()).isEqualTo(768); } @Test void batchEmbedding() { - assertThat(embeddingClient).isNotNull(); - EmbeddingResponse embeddingResponse = embeddingClient + assertThat(embeddingModel).isNotNull(); + EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); assertThat(embeddingResponse.getResults()).hasSize(2); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); @@ -56,7 +56,7 @@ class VertexAiPaLm2EmbeddingClientIT { assertThat(embeddingResponse.getResults().get(1).getOutput()).isNotEmpty(); assertThat(embeddingResponse.getResults().get(1).getIndex()).isEqualTo(1); - assertThat(embeddingClient.dimensions()).isEqualTo(768); + assertThat(embeddingModel.dimensions()).isEqualTo(768); } @SpringBootConfiguration @@ -68,8 +68,8 @@ class VertexAiPaLm2EmbeddingClientIT { } @Bean - public VertexAiPaLm2EmbeddingClient vertexAiEmbedding(VertexAiPaLm2Api vertexAiApi) { - return new VertexAiPaLm2EmbeddingClient(vertexAiApi); + public VertexAiPaLm2EmbeddingModel vertexAiEmbedding(VertexAiPaLm2Api vertexAiApi) { + return new VertexAiPaLm2EmbeddingModel(vertexAiApi); } } diff --git a/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatClient.java b/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatModel.java similarity index 90% rename from models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatClient.java rename to models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatModel.java index b4f5c14ff..b04cd0051 100644 --- a/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatClient.java +++ b/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatModel.java @@ -18,12 +18,12 @@ package org.springframework.ai.watsonx; import java.util.List; import java.util.Map; +import org.springframework.ai.chat.ChatModel; import reactor.core.publisher.Flux; -import org.springframework.ai.chat.ChatClient; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; @@ -35,7 +35,7 @@ import org.springframework.ai.watsonx.utils.MessageToPromptConverter; import org.springframework.util.Assert; /** - * {@link ChatClient} implementation for {@literal watsonx.ai}. + * {@link ChatModel} implementation for {@literal watsonx.ai}. * * watsonx.ai allows developers to use large language models within a SaaS service. It * supports multiple open-source models as well as IBM created models @@ -48,13 +48,13 @@ import org.springframework.util.Assert; * @author Christian Tzolov * @since 1.0.0 */ -public class WatsonxAiChatClient implements ChatClient, StreamingChatClient { +public class WatsonxAiChatModel implements ChatModel, StreamingChatModel { private final WatsonxAiApi watsonxAiApi; private final WatsonxAiChatOptions defaultOptions; - public WatsonxAiChatClient(WatsonxAiApi watsonxAiApi) { + public WatsonxAiChatModel(WatsonxAiApi watsonxAiApi) { this(watsonxAiApi, WatsonxAiChatOptions.builder() .withTemperature(0.7f) @@ -68,7 +68,7 @@ public class WatsonxAiChatClient implements ChatClient, StreamingChatClient { .build()); } - public WatsonxAiChatClient(WatsonxAiApi watsonxAiApi, WatsonxAiChatOptions defaultOptions) { + public WatsonxAiChatModel(WatsonxAiApi watsonxAiApi, WatsonxAiChatOptions defaultOptions) { Assert.notNull(watsonxAiApi, "watsonxAiApi cannot be null"); Assert.notNull(defaultOptions, "defaultOptions cannot be null"); this.watsonxAiApi = watsonxAiApi; @@ -140,4 +140,9 @@ public class WatsonxAiChatClient implements ChatClient, StreamingChatClient { return WatsonxAiRequest.builder(convertedPrompt).withParameters(parameters).build(); } + @Override + public ChatOptions getDefaultOptions() { + return WatsonxAiChatOptions.fromOptions(this.defaultOptions); + } + } \ No newline at end of file diff --git a/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatOptions.java b/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatOptions.java index 0b10febdb..28bd02159 100644 --- a/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatOptions.java +++ b/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatOptions.java @@ -324,5 +324,21 @@ public class WatsonxAiChatOptions implements ChatOptions { return input != null ? input.replaceAll("([a-z])([A-Z]+)", "$1_$2").toLowerCase() : null; } + public static WatsonxAiChatOptions fromOptions(WatsonxAiChatOptions fromOptions) { + return WatsonxAiChatOptions.builder() + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withTopK(fromOptions.getTopK()) + .withDecodingMethod(fromOptions.getDecodingMethod()) + .withMaxNewTokens(fromOptions.getMaxNewTokens()) + .withMinNewTokens(fromOptions.getMinNewTokens()) + .withStopSequences(fromOptions.getStopSequences()) + .withRepetitionPenalty(fromOptions.getRepetitionPenalty()) + .withRandomSeed(fromOptions.getRandomSeed()) + .withModel(fromOptions.getModel()) + .withAdditionalProperties(fromOptions.getAdditionalProperties()) + .build(); + } + } // @formatter:on \ No newline at end of file diff --git a/models/spring-ai-watsonx-ai/src/test/java/org/springframework/ai/watsonx/WatsonxAiChatClientTest.java b/models/spring-ai-watsonx-ai/src/test/java/org/springframework/ai/watsonx/WatsonxAiChatModelTest.java similarity index 93% rename from models/spring-ai-watsonx-ai/src/test/java/org/springframework/ai/watsonx/WatsonxAiChatClientTest.java rename to models/spring-ai-watsonx-ai/src/test/java/org/springframework/ai/watsonx/WatsonxAiChatModelTest.java index c74afa427..18b89ff09 100644 --- a/models/spring-ai-watsonx-ai/src/test/java/org/springframework/ai/watsonx/WatsonxAiChatClientTest.java +++ b/models/spring-ai-watsonx-ai/src/test/java/org/springframework/ai/watsonx/WatsonxAiChatModelTest.java @@ -46,9 +46,9 @@ import static org.mockito.Mockito.when; * @author Pablo Sanchidrian Herrera * @author John Jairo Moreno Rojas */ -public class WatsonxAiChatClientTest { +public class WatsonxAiChatModelTest { - WatsonxAiChatClient chatClient = new WatsonxAiChatClient(mock(WatsonxAiApi.class)); + WatsonxAiChatModel chatModel = new WatsonxAiChatModel(mock(WatsonxAiApi.class)); @Test public void testCreateRequestWithNoModelId() { @@ -57,7 +57,7 @@ public class WatsonxAiChatClientTest { Prompt prompt = new Prompt("Test message", options); Exception exception = Assert.assertThrows(IllegalArgumentException.class, () -> { - WatsonxAiRequest request = chatClient.request(prompt); + WatsonxAiRequest request = chatModel.request(prompt); }); } @@ -71,7 +71,7 @@ public class WatsonxAiChatClientTest { .build(); Prompt prompt = new Prompt(msg, modelOptions); - WatsonxAiRequest request = chatClient.request(prompt); + WatsonxAiRequest request = chatModel.request(prompt); Assert.assertEquals(request.getModelId(), "meta-llama/llama-2-70b-chat"); assertThat(request.getParameters().get("decoding_method")).isEqualTo("greedy"); @@ -105,7 +105,7 @@ public class WatsonxAiChatClientTest { Prompt prompt = new Prompt(msg, modelOptions); - WatsonxAiRequest request = chatClient.request(prompt); + WatsonxAiRequest request = chatModel.request(prompt); Assert.assertEquals(request.getModelId(), "meta-llama/llama-2-70b-chat"); assertThat(request.getParameters().get("decoding_method")).isEqualTo("sample"); @@ -139,7 +139,7 @@ public class WatsonxAiChatClientTest { Prompt prompt = new Prompt(msg, modelOptions); - WatsonxAiRequest request = chatClient.request(prompt); + WatsonxAiRequest request = chatModel.request(prompt); Assert.assertEquals(request.getModelId(), "meta-llama/llama-2-70b-chat"); assertThat(request.getInput()).isEqualTo(msg); @@ -157,7 +157,7 @@ public class WatsonxAiChatClientTest { @Test public void testCallMethod() { WatsonxAiApi mockChatApi = mock(WatsonxAiApi.class); - WatsonxAiChatClient client = new WatsonxAiChatClient(mockChatApi); + WatsonxAiChatModel chatModel = new WatsonxAiChatModel(mockChatApi); Prompt prompt = new Prompt(List.of(new SystemMessage("Your prompt here")), WatsonxAiChatOptions.builder().withModel("google/flan-ul2").build()); @@ -177,7 +177,7 @@ public class WatsonxAiChatClientTest { Map.of("warnings", List.of(Map.of("message", "the message", "id", "disclaimer_warning"))))); ChatResponse expectedResponse = new ChatResponse(List.of(expectedGenerator)); - ChatResponse response = client.call(prompt); + ChatResponse response = chatModel.call(prompt); Assert.assertEquals(expectedResponse.getResults().size(), response.getResults().size()); Assert.assertEquals(expectedResponse.getResult().getOutput(), response.getResult().getOutput()); @@ -186,7 +186,7 @@ public class WatsonxAiChatClientTest { @Test public void testStreamMethod() { WatsonxAiApi mockChatApi = mock(WatsonxAiApi.class); - WatsonxAiChatClient client = new WatsonxAiChatClient(mockChatApi); + WatsonxAiChatModel chatModel = new WatsonxAiChatModel(mockChatApi); Prompt prompt = new Prompt(List.of(new SystemMessage("Your prompt here")), WatsonxAiChatOptions.builder().withModel("google/flan-ul2").build()); @@ -210,7 +210,7 @@ public class WatsonxAiChatClientTest { Map.of("warnings", List.of(Map.of("message", "the message", "id", "disclaimer_warning"))))); Generation secondGen = new Generation("onse"); - Flux response = client.stream(prompt); + Flux response = chatModel.stream(prompt); StepVerifier.create(response).assertNext(current -> { diff --git a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatClient.java b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatModel.java similarity index 93% rename from models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatClient.java rename to models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatModel.java index f0c612672..2b5d804b1 100644 --- a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatClient.java +++ b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatModel.java @@ -17,10 +17,12 @@ package org.springframework.ai.zhipuai; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.chat.ChatClient; + +import org.springframework.ai.chat.ChatModel; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.metadata.ChatGenerationMetadata; import org.springframework.ai.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; @@ -55,20 +57,20 @@ import java.util.Set; import java.util.concurrent.ConcurrentHashMap; /** - * {@link ChatClient} and {@link StreamingChatClient} implementation for - * {@literal ZhiPuAI} backed by {@link ZhiPuAiApi}. + * {@link ChatModel} and {@link StreamingChatModel} implementation for {@literal ZhiPuAI} + * backed by {@link ZhiPuAiApi}. * * @author Geng Rong - * @see ChatClient - * @see StreamingChatClient + * @see ChatModel + * @see StreamingChatModel * @see ZhiPuAiApi * @since 1.0.0 M1 */ -public class ZhiPuAiChatClient extends +public class ZhiPuAiChatModel extends AbstractFunctionCallSupport> - implements ChatClient, StreamingChatClient { + implements ChatModel, StreamingChatModel { - private static final Logger logger = LoggerFactory.getLogger(ZhiPuAiChatClient.class); + private static final Logger logger = LoggerFactory.getLogger(ZhiPuAiChatModel.class); /** * The default options used for the chat completion requests. @@ -86,35 +88,35 @@ public class ZhiPuAiChatClient extends private final ZhiPuAiApi zhiPuAiApi; /** - * Creates an instance of the ZhiPuAiChatClient. + * Creates an instance of the ZhiPuAiChatModel. * @param zhiPuAiApi The ZhiPuAiApi instance to be used for interacting with the * ZhiPuAI Chat API. * @throws IllegalArgumentException if zhiPuAiApi is null */ - public ZhiPuAiChatClient(ZhiPuAiApi zhiPuAiApi) { + public ZhiPuAiChatModel(ZhiPuAiApi zhiPuAiApi) { this(zhiPuAiApi, ZhiPuAiChatOptions.builder().withModel(ZhiPuAiApi.DEFAULT_CHAT_MODEL).withTemperature(0.7f).build()); } /** - * Initializes an instance of the ZhiPuAiChatClient. + * Initializes an instance of the ZhiPuAiChatModel. * @param zhiPuAiApi The ZhiPuAiApi instance to be used for interacting with the * ZhiPuAI Chat API. - * @param options The ZhiPuAiChatOptions to configure the chat client. + * @param options The ZhiPuAiChatOptions to configure the chat model. */ - public ZhiPuAiChatClient(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options) { + public ZhiPuAiChatModel(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options) { this(zhiPuAiApi, options, null, RetryUtils.DEFAULT_RETRY_TEMPLATE); } /** - * Initializes a new instance of the ZhiPuAiChatClient. + * Initializes a new instance of the ZhiPuAiChatModel. * @param zhiPuAiApi The ZhiPuAiApi instance to be used for interacting with the * ZhiPuAI Chat API. - * @param options The ZhiPuAiChatOptions to configure the chat client. + * @param options The ZhiPuAiChatOptions to configure the chat model. * @param functionCallbackContext The function callback context. * @param retryTemplate The retry template. */ - public ZhiPuAiChatClient(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options, + public ZhiPuAiChatModel(ZhiPuAiApi zhiPuAiApi, ZhiPuAiChatOptions options, FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate) { super(functionCallbackContext); Assert.notNull(zhiPuAiApi, "ZhiPuAiApi must not be null"); @@ -382,4 +384,9 @@ public class ZhiPuAiChatClient extends && choice.finishReason() == ChatCompletionFinishReason.TOOL_CALLS; } + @Override + public ChatOptions getDefaultOptions() { + return ZhiPuAiChatOptions.fromOptions(this.defaultOptions); + } + } diff --git a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatOptions.java b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatOptions.java index 13dc8d519..6280c10ac 100644 --- a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatOptions.java +++ b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatOptions.java @@ -117,10 +117,10 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions, ChatOptions { private @JsonProperty("user_id") String user; /** - * ZhiPuAI Tool Function Callbacks to register with the ChatClient. + * ZhiPuAI Tool Function Callbacks to register with the ChatModel. * For Prompt Options the functionCallbacks are automatically enabled for the duration of the prompt execution. * For Default Options the functionCallbacks are registered but disabled by default. Use the enableFunctions to set the functions - * from the registry to be used by the ChatClient chat completion requests. + * from the registry to be used by the ChatModel chat completion requests. */ @NestedConfigurationProperty @JsonIgnore @@ -490,4 +490,24 @@ public class ZhiPuAiChatOptions implements FunctionCallingOptions, ChatOptions { throw new UnsupportedOperationException("Unimplemented method 'setTopK'"); } + public static ZhiPuAiChatOptions fromOptions(ZhiPuAiChatOptions fromOptions) { + return ZhiPuAiChatOptions.builder() + .withModel(fromOptions.getModel()) + .withFrequencyPenalty(fromOptions.getFrequencyPenalty()) + .withMaxTokens(fromOptions.getMaxTokens()) + .withN(fromOptions.getN()) + .withPresencePenalty(fromOptions.getPresencePenalty()) + .withResponseFormat(fromOptions.getResponseFormat()) + .withSeed(fromOptions.getSeed()) + .withStop(fromOptions.getStop()) + .withTemperature(fromOptions.getTemperature()) + .withTopP(fromOptions.getTopP()) + .withTools(fromOptions.getTools()) + .withToolChoice(fromOptions.getToolChoice()) + .withUser(fromOptions.getUser()) + .withFunctionCallbacks(fromOptions.getFunctionCallbacks()) + .withFunctions(fromOptions.getFunctions()) + .build(); + } + } diff --git a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiEmbeddingClient.java b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiEmbeddingModel.java similarity index 87% rename from models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiEmbeddingClient.java rename to models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiEmbeddingModel.java index 25aecfc74..8ec16a109 100644 --- a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiEmbeddingClient.java +++ b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiEmbeddingModel.java @@ -19,7 +19,7 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; import org.springframework.ai.document.MetadataMode; -import org.springframework.ai.embedding.AbstractEmbeddingClient; +import org.springframework.ai.embedding.AbstractEmbeddingModel; import org.springframework.ai.embedding.Embedding; import org.springframework.ai.embedding.EmbeddingOptions; import org.springframework.ai.embedding.EmbeddingRequest; @@ -39,9 +39,9 @@ import java.util.List; * @author Geng Rong * @since 1.0.0 M1 */ -public class ZhiPuAiEmbeddingClient extends AbstractEmbeddingClient { +public class ZhiPuAiEmbeddingModel extends AbstractEmbeddingModel { - private static final Logger logger = LoggerFactory.getLogger(ZhiPuAiEmbeddingClient.class); + private static final Logger logger = LoggerFactory.getLogger(ZhiPuAiEmbeddingModel.class); private final ZhiPuAiEmbeddingOptions defaultOptions; @@ -52,43 +52,43 @@ public class ZhiPuAiEmbeddingClient extends AbstractEmbeddingClient { private final MetadataMode metadataMode; /** - * Constructor for the ZhiPuAiEmbeddingClient class. + * Constructor for the ZhiPuAiEmbeddingModel class. * @param zhiPuAiApi The ZhiPuAiApi instance to use for making API requests. */ - public ZhiPuAiEmbeddingClient(ZhiPuAiApi zhiPuAiApi) { + public ZhiPuAiEmbeddingModel(ZhiPuAiApi zhiPuAiApi) { this(zhiPuAiApi, MetadataMode.EMBED); } /** - * Initializes a new instance of the ZhiPuAiEmbeddingClient class. + * Initializes a new instance of the ZhiPuAiEmbeddingModel class. * @param zhiPuAiApi The ZhiPuAiApi instance to use for making API requests. * @param metadataMode The mode for generating metadata. */ - public ZhiPuAiEmbeddingClient(ZhiPuAiApi zhiPuAiApi, MetadataMode metadataMode) { + public ZhiPuAiEmbeddingModel(ZhiPuAiApi zhiPuAiApi, MetadataMode metadataMode) { this(zhiPuAiApi, metadataMode, ZhiPuAiEmbeddingOptions.builder().withModel(ZhiPuAiApi.DEFAULT_EMBEDDING_MODEL).build(), RetryUtils.DEFAULT_RETRY_TEMPLATE); } /** - * Initializes a new instance of the ZhiPuAiEmbeddingClient class. + * Initializes a new instance of the ZhiPuAiEmbeddingModel class. * @param zhiPuAiApi The ZhiPuAiApi instance to use for making API requests. * @param metadataMode The mode for generating metadata. * @param zhiPuAiEmbeddingOptions The options for ZhiPuAI embedding. */ - public ZhiPuAiEmbeddingClient(ZhiPuAiApi zhiPuAiApi, MetadataMode metadataMode, + public ZhiPuAiEmbeddingModel(ZhiPuAiApi zhiPuAiApi, MetadataMode metadataMode, ZhiPuAiEmbeddingOptions zhiPuAiEmbeddingOptions) { this(zhiPuAiApi, metadataMode, zhiPuAiEmbeddingOptions, RetryUtils.DEFAULT_RETRY_TEMPLATE); } /** - * Initializes a new instance of the ZhiPuAiEmbeddingClient class. + * Initializes a new instance of the ZhiPuAiEmbeddingModel class. * @param zhiPuAiApi - The ZhiPuAiApi instance to use for making API requests. * @param metadataMode - The mode for generating metadata. * @param options - The options for ZhiPuAI embedding. * @param retryTemplate - The RetryTemplate for retrying failed API requests. */ - public ZhiPuAiEmbeddingClient(ZhiPuAiApi zhiPuAiApi, MetadataMode metadataMode, ZhiPuAiEmbeddingOptions options, + public ZhiPuAiEmbeddingModel(ZhiPuAiApi zhiPuAiApi, MetadataMode metadataMode, ZhiPuAiEmbeddingOptions options, RetryTemplate retryTemplate) { Assert.notNull(zhiPuAiApi, "ZhiPuAiApi must not be null"); Assert.notNull(metadataMode, "metadataMode must not be null"); diff --git a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiImageClient.java b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiImageModel.java similarity index 92% rename from models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiImageClient.java rename to models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiImageModel.java index 09caabb63..5b31aca1c 100644 --- a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiImageClient.java +++ b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiImageModel.java @@ -19,7 +19,7 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.image.Image; -import org.springframework.ai.image.ImageClient; +import org.springframework.ai.image.ImageModel; import org.springframework.ai.image.ImageGeneration; import org.springframework.ai.image.ImageOptions; import org.springframework.ai.image.ImagePrompt; @@ -34,15 +34,15 @@ import org.springframework.util.Assert; import java.util.List; /** - * ZhiPuAiImageClient is a class that implements the ImageClient interface. It provides a + * ZhiPuAiImageModel is a class that implements the ImageModel interface. It provides a * client for calling the ZhiPuAI image generation API. * * @author Geng Rong * @since 1.0.0 M1 */ -public class ZhiPuAiImageClient implements ImageClient { +public class ZhiPuAiImageModel implements ImageModel { - private final static Logger logger = LoggerFactory.getLogger(ZhiPuAiImageClient.class); + private final static Logger logger = LoggerFactory.getLogger(ZhiPuAiImageModel.class); private final ZhiPuAiImageOptions defaultOptions; @@ -50,11 +50,11 @@ public class ZhiPuAiImageClient implements ImageClient { public final RetryTemplate retryTemplate; - public ZhiPuAiImageClient(ZhiPuAiImageApi zhiPuAiImageApi) { + public ZhiPuAiImageModel(ZhiPuAiImageApi zhiPuAiImageApi) { this(zhiPuAiImageApi, ZhiPuAiImageOptions.builder().build(), RetryUtils.DEFAULT_RETRY_TEMPLATE); } - public ZhiPuAiImageClient(ZhiPuAiImageApi zhiPuAiImageApi, ZhiPuAiImageOptions defaultOptions, + public ZhiPuAiImageModel(ZhiPuAiImageApi zhiPuAiImageApi, ZhiPuAiImageOptions defaultOptions, RetryTemplate retryTemplate) { Assert.notNull(zhiPuAiImageApi, "ZhiPuAiImageApi must not be null"); Assert.notNull(defaultOptions, "defaultOptions must not be null"); diff --git a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/api/ZhiPuAiApi.java b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/api/ZhiPuAiApi.java index 5984e189c..45e21a500 100644 --- a/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/api/ZhiPuAiApi.java +++ b/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/api/ZhiPuAiApi.java @@ -18,6 +18,8 @@ package org.springframework.ai.zhipuai.api; import com.fasterxml.jackson.annotation.JsonInclude; import com.fasterxml.jackson.annotation.JsonInclude.Include; import com.fasterxml.jackson.annotation.JsonProperty; + +import org.springframework.ai.model.ModelDescription; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.ai.retry.RetryUtils; import org.springframework.boot.context.properties.bind.ConstructorBinding; @@ -110,7 +112,7 @@ public class ZhiPuAiApi { * ZhiPuAI Chat Completion Models: * ZhiPuAI Model. */ - public enum ChatModel { + public enum ChatModel implements ModelDescription { GLM_4("GLM-4"), GLM_3_Turbo("GLM-3-Turbo"); @@ -123,6 +125,11 @@ public class ZhiPuAiApi { public String getValue() { return value; } + + @Override + public String getModelName() { + return this.value; + } } /** diff --git a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ChatCompletionRequestTests.java b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ChatCompletionRequestTests.java index b958bc013..972280627 100644 --- a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ChatCompletionRequestTests.java +++ b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ChatCompletionRequestTests.java @@ -33,7 +33,7 @@ public class ChatCompletionRequestTests { @Test public void createRequestWithChatOptions() { - var client = new ZhiPuAiChatClient(new ZhiPuAiApi("TEST"), + var client = new ZhiPuAiChatModel(new ZhiPuAiApi("TEST"), ZhiPuAiChatOptions.builder().withModel("DEFAULT_MODEL").withTemperature(66.6f).build()); var request = client.createRequest(new Prompt("Test message content"), false); @@ -59,7 +59,7 @@ public class ChatCompletionRequestTests { final String TOOL_FUNCTION_NAME = "CurrentWeather"; - var client = new ZhiPuAiChatClient(new ZhiPuAiApi("TEST"), + var client = new ZhiPuAiChatModel(new ZhiPuAiApi("TEST"), ZhiPuAiChatOptions.builder().withModel("DEFAULT_MODEL").build()); var request = client.createRequest(new Prompt("Test message content", @@ -89,7 +89,7 @@ public class ChatCompletionRequestTests { final String TOOL_FUNCTION_NAME = "CurrentWeather"; - var client = new ZhiPuAiChatClient(new ZhiPuAiApi("TEST"), + var client = new ZhiPuAiChatModel(new ZhiPuAiApi("TEST"), ZhiPuAiChatOptions.builder() .withModel("DEFAULT_MODEL") .withFunctionCallbacks(List.of(FunctionCallbackWrapper.builder(new MockWeatherService()) diff --git a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ZhiPuAiTestConfiguration.java b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ZhiPuAiTestConfiguration.java index 92f0bbb28..d35ac839a 100644 --- a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ZhiPuAiTestConfiguration.java +++ b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/ZhiPuAiTestConfiguration.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.zhipuai; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.zhipuai.api.ZhiPuAiApi; import org.springframework.ai.zhipuai.api.ZhiPuAiImageApi; import org.springframework.boot.SpringBootConfiguration; @@ -48,18 +48,18 @@ public class ZhiPuAiTestConfiguration { } @Bean - public ZhiPuAiChatClient zhiPuAiChatClient(ZhiPuAiApi api) { - return new ZhiPuAiChatClient(api); + public ZhiPuAiChatModel zhiPuAiChatModel(ZhiPuAiApi api) { + return new ZhiPuAiChatModel(api); } @Bean - public ZhiPuAiImageClient zhiPuAiImageClient(ZhiPuAiImageApi imageApi) { - return new ZhiPuAiImageClient(imageApi); + public ZhiPuAiImageModel zhiPuAiImageModel(ZhiPuAiImageApi imageApi) { + return new ZhiPuAiImageModel(imageApi); } @Bean - public EmbeddingClient zhiPuAiEmbeddingClient(ZhiPuAiApi api) { - return new ZhiPuAiEmbeddingClient(api); + public EmbeddingModel zhiPuAiEmbeddingModel(ZhiPuAiApi api) { + return new ZhiPuAiEmbeddingModel(api); } } diff --git a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/api/ZhiPuAiRetryTests.java b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/api/ZhiPuAiRetryTests.java index 3326af098..6feb7747c 100644 --- a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/api/ZhiPuAiRetryTests.java +++ b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/api/ZhiPuAiRetryTests.java @@ -26,11 +26,11 @@ import org.springframework.ai.image.ImageMessage; import org.springframework.ai.image.ImagePrompt; import org.springframework.ai.retry.RetryUtils; import org.springframework.ai.retry.TransientAiException; -import org.springframework.ai.zhipuai.ZhiPuAiChatClient; +import org.springframework.ai.zhipuai.ZhiPuAiChatModel; import org.springframework.ai.zhipuai.ZhiPuAiChatOptions; -import org.springframework.ai.zhipuai.ZhiPuAiEmbeddingClient; +import org.springframework.ai.zhipuai.ZhiPuAiEmbeddingModel; import org.springframework.ai.zhipuai.ZhiPuAiEmbeddingOptions; -import org.springframework.ai.zhipuai.ZhiPuAiImageClient; +import org.springframework.ai.zhipuai.ZhiPuAiImageModel; import org.springframework.ai.zhipuai.ZhiPuAiImageOptions; import org.springframework.ai.zhipuai.api.ZhiPuAiApi.ChatCompletion; import org.springframework.ai.zhipuai.api.ZhiPuAiApi.ChatCompletionChunk; @@ -93,11 +93,11 @@ public class ZhiPuAiRetryTests { private @Mock ZhiPuAiImageApi zhiPuAiImageApi; - private ZhiPuAiChatClient chatClient; + private ZhiPuAiChatModel chatModel; - private ZhiPuAiEmbeddingClient embeddingClient; + private ZhiPuAiEmbeddingModel embeddingModel; - private ZhiPuAiImageClient imageClient; + private ZhiPuAiImageModel imageModel; @BeforeEach public void beforeEach() { @@ -105,10 +105,10 @@ public class ZhiPuAiRetryTests { retryListener = new TestRetryListener(); retryTemplate.registerListener(retryListener); - chatClient = new ZhiPuAiChatClient(zhiPuAiApi, ZhiPuAiChatOptions.builder().build(), null, retryTemplate); - embeddingClient = new ZhiPuAiEmbeddingClient(zhiPuAiApi, MetadataMode.EMBED, + chatModel = new ZhiPuAiChatModel(zhiPuAiApi, ZhiPuAiChatOptions.builder().build(), null, retryTemplate); + embeddingModel = new ZhiPuAiEmbeddingModel(zhiPuAiApi, MetadataMode.EMBED, ZhiPuAiEmbeddingOptions.builder().build(), retryTemplate); - imageClient = new ZhiPuAiImageClient(zhiPuAiImageApi, ZhiPuAiImageOptions.builder().build(), retryTemplate); + imageModel = new ZhiPuAiImageModel(zhiPuAiImageApi, ZhiPuAiImageOptions.builder().build(), retryTemplate); } @Test @@ -124,7 +124,7 @@ public class ZhiPuAiRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(ResponseEntity.of(Optional.of(expectedChatCompletion))); - var result = chatClient.call(new Prompt("text")); + var result = chatModel.call(new Prompt("text")); assertThat(result).isNotNull(); assertThat(result.getResult().getOutput().getContent()).isSameAs("Response"); @@ -136,7 +136,7 @@ public class ZhiPuAiRetryTests { public void zhiPuAiChatNonTransientError() { when(zhiPuAiApi.chatCompletionEntity(isA(ChatCompletionRequest.class))) .thenThrow(new RuntimeException("Non Transient Error")); - assertThrows(RuntimeException.class, () -> chatClient.call(new Prompt("text"))); + assertThrows(RuntimeException.class, () -> chatModel.call(new Prompt("text"))); } @Test @@ -152,7 +152,7 @@ public class ZhiPuAiRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(Flux.just(expectedChatCompletion)); - var result = chatClient.stream(new Prompt("text")); + var result = chatModel.stream(new Prompt("text")); assertThat(result).isNotNull(); assertThat(result.collectList().block().get(0).getResult().getOutput().getContent()).isSameAs("Response"); @@ -164,7 +164,7 @@ public class ZhiPuAiRetryTests { public void zhiPuAiChatStreamNonTransientError() { when(zhiPuAiApi.chatCompletionStream(isA(ChatCompletionRequest.class))) .thenThrow(new RuntimeException("Non Transient Error")); - assertThrows(RuntimeException.class, () -> chatClient.stream(new Prompt("text"))); + assertThrows(RuntimeException.class, () -> chatModel.stream(new Prompt("text"))); } @Test @@ -178,7 +178,7 @@ public class ZhiPuAiRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(ResponseEntity.of(Optional.of(expectedEmbeddings))); - var result = embeddingClient + var result = embeddingModel .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), null)); assertThat(result).isNotNull(); @@ -191,7 +191,7 @@ public class ZhiPuAiRetryTests { public void zhiPuAiEmbeddingNonTransientError() { when(zhiPuAiApi.embeddings(isA(EmbeddingRequest.class))) .thenThrow(new RuntimeException("Non Transient Error")); - assertThrows(RuntimeException.class, () -> embeddingClient + assertThrows(RuntimeException.class, () -> embeddingModel .call(new org.springframework.ai.embedding.EmbeddingRequest(List.of("text1", "text2"), null))); } @@ -205,7 +205,7 @@ public class ZhiPuAiRetryTests { .thenThrow(new TransientAiException("Transient Error 2")) .thenReturn(ResponseEntity.of(Optional.of(expectedResponse))); - var result = imageClient.call(new ImagePrompt(List.of(new ImageMessage("Image Message")))); + var result = imageModel.call(new ImagePrompt(List.of(new ImageMessage("Image Message")))); assertThat(result).isNotNull(); assertThat(result.getResult().getOutput().getUrl()).isEqualTo("url678"); @@ -218,7 +218,7 @@ public class ZhiPuAiRetryTests { when(zhiPuAiImageApi.createImage(isA(ZhiPuAiImageRequest.class))) .thenThrow(new RuntimeException("Transient Error 1")); assertThrows(RuntimeException.class, - () -> imageClient.call(new ImagePrompt(List.of(new ImageMessage("Image Message"))))); + () -> imageModel.call(new ImagePrompt(List.of(new ImageMessage("Image Message"))))); } } diff --git a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/image/ZhiPuAiImageClientIT.java b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/image/ZhiPuAiImageModelIT.java similarity index 92% rename from models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/image/ZhiPuAiImageClientIT.java rename to models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/image/ZhiPuAiImageModelIT.java index c150cce63..2c95a1499 100644 --- a/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/image/ZhiPuAiImageClientIT.java +++ b/models/spring-ai-zhipuai/src/test/java/org/springframework/ai/zhipuai/image/ZhiPuAiImageModelIT.java @@ -19,7 +19,7 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.image.Image; -import org.springframework.ai.image.ImageClient; +import org.springframework.ai.image.ImageModel; import org.springframework.ai.image.ImageOptionsBuilder; import org.springframework.ai.image.ImagePrompt; import org.springframework.ai.image.ImageResponse; @@ -32,10 +32,10 @@ import static org.assertj.core.api.Assertions.assertThat; @SpringBootTest(classes = ZhiPuAiTestConfiguration.class) @EnabledIfEnvironmentVariable(named = "ZHIPU_AI_API_KEY", matches = ".+") -public class ZhiPuAiImageClientIT { +public class ZhiPuAiImageModelIT { @Autowired - protected ImageClient imageClient; + protected ImageModel imageModel; @Test void imageAsUrlTest() { @@ -46,7 +46,7 @@ public class ZhiPuAiImageClientIT { ImagePrompt imagePrompt = new ImagePrompt(instructions, options); - ImageResponse imageResponse = imageClient.call(imagePrompt); + ImageResponse imageResponse = imageModel.call(imagePrompt); assertThat(imageResponse.getResults()).hasSize(1); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java index cff4f8674..bc2163765 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java @@ -15,30 +15,626 @@ */ package org.springframework.ai.chat; -import org.springframework.ai.chat.prompt.Prompt; - +import java.io.IOException; +import java.net.URL; +import java.nio.charset.Charset; +import java.util.ArrayList; import java.util.Arrays; +import java.util.Collection; +import java.util.HashMap; +import java.util.HashSet; +import java.util.List; +import java.util.Map; +import java.util.function.Consumer; +import reactor.core.publisher.Flux; + +import org.springframework.ai.chat.messages.Media; import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.SystemMessage; import org.springframework.ai.chat.messages.UserMessage; -import org.springframework.ai.model.ModelClient; +import org.springframework.ai.chat.prompt.ChatOptions; +import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.chat.prompt.PromptTemplate; +import org.springframework.ai.converter.BeanOutputConverter; +import org.springframework.ai.model.function.FunctionCallback; +import org.springframework.ai.model.function.FunctionCallbackWrapper; +import org.springframework.ai.model.function.FunctionCallingOptions; +import org.springframework.core.ParameterizedTypeReference; +import org.springframework.core.io.Resource; +import org.springframework.util.Assert; +import org.springframework.util.CollectionUtils; +import org.springframework.util.MimeType; +import org.springframework.util.StringUtils; -@FunctionalInterface -public interface ChatClient extends ModelClient { +// todo support plugging in a outputConverter at runtime +// todo figure out stream and list methods +/** + * @author Mark Pollack + * @author Christian Tzolov + * @author Josh Long + * @author Arjen Poutsma + */ +public interface ChatClient { + + static ChatClientBuilder builder(ChatModel chatModel) { + return new ChatClientBuilder(chatModel); + } + + ChatClientRequest prompt(); + + ChatClientPromptRequest prompt(Prompt prompt); + + interface PromptSpec { + + T text(String text); + + T text(Resource text, Charset charset); + + T text(Resource text); + + T params(Map p); + + T param(String k, String v); + + } + + abstract class AbstractPromptSpec> implements PromptSpec { + + private String text = ""; + + private final Map params = new HashMap<>(); + + @Override + public T text(String text) { + this.text = text; + return self(); + } + + @Override + public T text(Resource text, Charset charset) { + try { + this.text(text.getContentAsString(charset)); + } + catch (IOException e) { + throw new RuntimeException(e); + } + return self(); + } + + @Override + public T text(Resource text) { + this.text(text, Charset.defaultCharset()); + return self(); + } + + @Override + public T param(String k, String v) { + this.params.put(k, v); + return self(); + } + + @Override + public T params(Map p) { + this.params.putAll(p); + return self(); + } + + protected abstract T self(); + + protected String text() { + return this.text; + } + + protected Map params() { + return this.params; + } + + } + + class UserSpec extends AbstractPromptSpec implements PromptSpec { + + private final List media = new ArrayList<>(); + + public UserSpec media(Media... media) { + this.media.addAll(Arrays.asList(media)); + return self(); + } + + public UserSpec media(MimeType mimeType, URL url) { + this.media.add(new Media(mimeType, url)); + return self(); + } + + public UserSpec media(MimeType mimeType, Resource resource) { + this.media.add(new Media(mimeType, resource)); + return self(); + } + + protected List media() { + return this.media; + } + + @Override + protected UserSpec self() { + return this; + } + + } + + class SystemSpec extends AbstractPromptSpec implements PromptSpec { + + @Override + protected SystemSpec self() { + return this; + } + + } + + class ChatClientPromptRequest { + + private final ChatModel chatModel; + + private final Prompt prompt; + + public ChatClientPromptRequest(ChatModel chatModel, Prompt prompt) { + this.chatModel = chatModel; + this.prompt = prompt; + } + + public ChatClientRequest.CallPromptResponseSpec call() { + return new ChatClientRequest.CallPromptResponseSpec(this.chatModel, this.prompt); + } + + public ChatClientRequest.StreamPromptResponseSpec stream() { + return new ChatClientRequest.StreamPromptResponseSpec((StreamingChatModel) this.chatModel, this.prompt); + } + + } + + class ChatClientRequest { + + private final ChatModel chatModel; + + private String userText = ""; + + private String systemText = ""; + + private ChatOptions chatOptions; + + private final List media = new ArrayList<>(); + + private final List functionNames = new ArrayList<>(); + + private final List functionCallbacks = new ArrayList<>(); + + private final List messages = new ArrayList<>(); + + private final Map userParams = new HashMap<>(); + + private final Map systemParams = new HashMap<>(); + + /* copy constructor */ + ChatClientRequest(ChatClientRequest ccr) { + this(ccr.chatModel, ccr.userText, ccr.systemText, ccr.functionCallbacks, ccr.functionNames, ccr.media, + ccr.chatOptions); + } + + public ChatClientRequest(ChatModel chatModel, String userText, String systemText, + List functionCallbacks, List functionNames, List media, + ChatOptions chatOptions) { + + this.chatModel = chatModel; + this.chatOptions = chatOptions != null ? chatOptions : chatModel.getDefaultOptions(); + + this.userText = userText; + this.systemText = systemText; + + this.functionNames.addAll(functionNames); + this.functionCallbacks.addAll(functionCallbacks); + this.media.addAll(media); + } + + public ChatClientRequest messages(Message... messages) { + this.messages.addAll(List.of(messages)); + return this; + } + + public ChatClientRequest messages(List messages) { + this.messages.addAll(messages); + return this; + } + + public ChatClientRequest options(T options) { + this.chatOptions = options; + return this; + } + + public ChatClientRequest function(String name, String description, + java.util.function.Function function) { + var fcw = FunctionCallbackWrapper.builder(function) + .withDescription(description) + .withName(name) + .withResponseConverter(Object::toString) + .build(); + this.functionCallbacks.add(fcw); + return this; + } + + public ChatClientRequest functions(String... functionBeanNames) { + this.functionNames.addAll(List.of(functionBeanNames)); + return this; + } + + public ChatClientRequest chatOptions(ChatOptions chatOptions) { + this.chatOptions = chatOptions; + return this; + } + + public ChatClientRequest system(String text) { + this.systemText = text; + return this; + } + + public ChatClientRequest system(Consumer consumer) { + var ss = new SystemSpec(); + consumer.accept(ss); + this.systemText = StringUtils.hasText(ss.text()) ? ss.text() : this.systemText; + this.systemParams.putAll(ss.params()); + return this; + } + + public ChatClientRequest user(String text) { + this.userText = text; + return this; + } + + public ChatClientRequest user(Consumer consumer) { + var us = new UserSpec(); + consumer.accept(us); + this.userText = StringUtils.hasText(us.text()) ? us.text() : this.userText; + this.userParams.putAll(us.params()); + this.media.addAll(us.media()); + return this; + } + + public static class StreamPromptResponseSpec { + + private final Prompt prompt; + + private final StreamingChatModel chatModel; + + public StreamPromptResponseSpec(StreamingChatModel streamingChatModel, Prompt prompt) { + this.chatModel = streamingChatModel; + this.prompt = prompt; + } + + public Flux chatResponse() { + return doGetFluxChatResponse(this.prompt); + } + + private Flux doGetFluxChatResponse(Prompt prompt) { + return this.chatModel.stream(prompt); + } + + public Flux content() { + return doGetFluxChatResponse(this.prompt).map(r -> { + if (r.getResult() == null || r.getResult().getOutput() == null + || r.getResult().getOutput().getContent() == null) { + return ""; + } + return r.getResult().getOutput().getContent(); + }).filter(v -> StringUtils.hasText(v)); + } + + } + + public static class CallPromptResponseSpec { + + private final ChatModel chatModel; + + private final Prompt prompt; + + public CallPromptResponseSpec(ChatModel chatModel, Prompt prompt) { + this.chatModel = chatModel; + this.prompt = prompt; + } + + public String content() { + return doGetChatResponse(this.prompt).getResult().getOutput().getContent(); + } + + public List contents() { + return doGetChatResponse(this.prompt).getResults() + .stream() + .map(r -> r.getOutput().getContent()) + .toList(); + } + + public ChatResponse chatResponse() { + return doGetChatResponse(this.prompt); + } + + private ChatResponse doGetChatResponse(Prompt prompt) { + return chatModel.call(prompt); + } + + } + + public static class CallResponseSpec { + + private final ChatClientRequest request; + + private final ChatModel chatModel; + + public CallResponseSpec(ChatModel chatModel, ChatClientRequest request) { + this.chatModel = chatModel; + this.request = request; + } + + public T entity(ParameterizedTypeReference type) { + return doSingleWithBeanOutputConverter(new BeanOutputConverter(type)); + } + + private T doSingleWithBeanOutputConverter(BeanOutputConverter boc) { + var processedUserText = this.request.userText + System.lineSeparator() + System.lineSeparator() + + "{format}"; + var chatResponse = doGetChatResponse(processedUserText, boc.getFormat()); + var stringResponse = chatResponse.getResult().getOutput().getContent(); + return boc.convert(stringResponse); + } + + public T entity(Class type) { + Assert.notNull(type, "the class must be non-null"); + var boc = new BeanOutputConverter(type); + return doSingleWithBeanOutputConverter(boc); + } + + private ChatResponse doGetChatResponse(String processedUserText) { + return this.doGetChatResponse(processedUserText, ""); + } + + private ChatResponse doGetChatResponse(String processedUserText, String formatParam) { + Map userParams = new HashMap<>(this.request.userParams); + if (StringUtils.hasText(formatParam)) { + userParams.put("format", formatParam); + } + + var messages = new ArrayList(); + var textsAreValid = (StringUtils.hasText(processedUserText) + || StringUtils.hasText(this.request.systemText)); + var messagesAreValid = !this.request.messages.isEmpty(); + Assert.state(!(messagesAreValid && textsAreValid), "you must specify either " + Message.class.getName() + + " instances or user/system texts, but not both"); + if (textsAreValid) { + UserMessage userMessage = null; + if (!CollectionUtils.isEmpty(userParams)) { + userMessage = new UserMessage(new PromptTemplate(processedUserText, userParams).render(), + this.request.media); + } + else { + userMessage = new UserMessage(processedUserText, this.request.media); + } + if (StringUtils.hasText(this.request.systemText) || !this.request.systemParams.isEmpty()) { + var systemMessage = new SystemMessage( + new PromptTemplate(this.request.systemText, this.request.systemParams).render()); + messages.add(systemMessage); + } + messages.add(userMessage); + } + else { + messages.addAll(this.request.messages); + } + if (this.request.chatOptions instanceof FunctionCallingOptions functionCallingOptions) { + // if (this.request.chatOptions instanceof + // FunctionCallingOptionsBuilder.PortableFunctionCallingOptions + // functionCallingOptions) { + if (!this.request.functionNames.isEmpty()) { + functionCallingOptions.setFunctions(new HashSet<>(this.request.functionNames)); + } + if (!this.request.functionCallbacks.isEmpty()) { + functionCallingOptions.setFunctionCallbacks(this.request.functionCallbacks); + } + } + var prompt = new Prompt(messages, this.request.chatOptions); + return this.chatModel.call(prompt); + } + + public ChatResponse chatResponse() { + return doGetChatResponse(this.request.userText); + } + + public String content() { + return doGetChatResponse(this.request.userText).getResult().getOutput().getContent(); + } + + public List contents() { + return doGetChatResponse(this.request.userText).getResults() + .stream() + .map(r -> r.getOutput().getContent()) + .toList(); + } + + } + + public static class StreamResponseSpec { + + private final ChatClientRequest request; + + private final StreamingChatModel chatModel; + + public StreamResponseSpec(StreamingChatModel streamingChatModel, ChatClientRequest request) { + this.chatModel = streamingChatModel; + this.request = request; + } + + private Flux doGetFluxChatResponse(String processedUserText) { + Map userParams = new HashMap<>(this.request.userParams); + + var messages = new ArrayList(); + var textsAreValid = (StringUtils.hasText(processedUserText) + || StringUtils.hasText(this.request.systemText)); + var messagesAreValid = !this.request.messages.isEmpty(); + Assert.state(!(messagesAreValid && textsAreValid), "you must specify either " + Message.class.getName() + + " instances or user/system texts, but not both"); + if (textsAreValid) { + UserMessage userMessage = null; + if (!CollectionUtils.isEmpty(userParams)) { + userMessage = new UserMessage(new PromptTemplate(processedUserText, userParams).render(), + this.request.media); + } + else { + userMessage = new UserMessage(processedUserText, this.request.media); + } + if (StringUtils.hasText(this.request.systemText) || !this.request.systemParams.isEmpty()) { + var systemMessage = new SystemMessage( + new PromptTemplate(this.request.systemText, this.request.systemParams).render()); + messages.add(systemMessage); + } + messages.add(userMessage); + } + else { + messages.addAll(this.request.messages); + } + if (this.request.chatOptions instanceof FunctionCallingOptions functionCallingOptions) { + // if (this.request.chatOptions instanceof + // FunctionCallingOptionsBuilder.PortableFunctionCallingOptions + // functionCallingOptions) { + if (!this.request.functionNames.isEmpty()) { + functionCallingOptions.setFunctions(new HashSet<>(this.request.functionNames)); + } + if (!this.request.functionCallbacks.isEmpty()) { + functionCallingOptions.setFunctionCallbacks(this.request.functionCallbacks); + } + } + var prompt = new Prompt(messages, this.request.chatOptions); + return this.chatModel.stream(prompt); + } + + public Flux chatResponse() { + return doGetFluxChatResponse(this.request.userText); + } + + public Flux content() { + return doGetFluxChatResponse(this.request.userText).map(r -> { + if (r.getResult() == null || r.getResult().getOutput() == null + || r.getResult().getOutput().getContent() == null) { + return ""; + } + return r.getResult().getOutput().getContent(); + }).filter(v -> StringUtils.hasText(v)); + } + + } + + public CallResponseSpec call() { + return new CallResponseSpec(this.chatModel, this); + } + + public StreamResponseSpec stream() { + return new StreamResponseSpec((StreamingChatModel) this.chatModel, this); + } + + } + + class ChatClientBuilder { + + private final ChatClientRequest defaultRequest; + + private final ChatModel chatModel; + + ChatClientBuilder(ChatModel chatModel) { + Assert.notNull(chatModel, "the " + ChatModel.class.getName() + " must be non-null"); + this.chatModel = chatModel; + this.defaultRequest = new ChatClientRequest(chatModel, "", "", List.of(), List.of(), List.of(), null); + } + + public ChatClient build() { + return new DefaultChatClient(this.chatModel, this.defaultRequest); + } + + public ChatClientBuilder defaultRuntimeOptions(ChatOptions chatOptions) { + this.defaultRequest.chatOptions(chatOptions); + return this; + } + + public ChatClientBuilder defaultUser(String text) { + this.defaultRequest.user(text); + return this; + } + + public ChatClientBuilder defaultUser(Consumer userSpecConsumer) { + this.defaultRequest.user(userSpecConsumer); + return this; + } + + public ChatClientBuilder defaultSystem(String text) { + this.defaultRequest.system(text); + return this; + } + + public ChatClientBuilder defaultSystem(Consumer systemSpecConsumer) { + this.defaultRequest.system(systemSpecConsumer); + return this; + } + + public ChatClientBuilder defaultFunctionWrappers(String name, String description, + java.util.function.Function function) { + this.defaultRequest.function(name, description, function); + return this; + } + + public ChatClientBuilder defaultFunctions(String... functionNames) { + this.defaultRequest.functions(functionNames); + return this; + } + + } + + /** + * Calls the underlying chat model with a prompt message and returns the output + * content of the first generation. + * @param message The message to be used as a prompt for the chat model. + * @return The output content of the first generation. + * @deprecated This method is deprecated as of version 1.0.0 M1 and will be removed in + * a future release. Use the method + * builder(chatModel).build().prompt().user(message).call().content() instead + * + */ + @Deprecated(since = "1.0.0 M1", forRemoval = true) default String call(String message) { - Prompt prompt = new Prompt(new UserMessage(message)); - Generation generation = call(prompt).getResult(); + var prompt = new Prompt(new UserMessage(message)); + var generation = call(prompt).getResult(); return (generation != null) ? generation.getOutput().getContent() : ""; } + /** + * Calls the underlying chat model with a prompt message and returns the output + * content of the first generation. + * @param messages The messages to be used as a prompt for the chat model. + * @return The output content of the first generation. + * @deprecated This method is deprecated as of version 1.0.0 M1 and will be removed in + * a future release. Use the method + * builder(chatModel).build().prompt().messages(messages).call().content() instead. + */ + @Deprecated(since = "1.0.0 M1", forRemoval = true) default String call(Message... messages) { - Prompt prompt = new Prompt(Arrays.asList(messages)); - Generation generation = call(prompt).getResult(); + var prompt = new Prompt(Arrays.asList(messages)); + var generation = call(prompt).getResult(); return (generation != null) ? generation.getOutput().getContent() : ""; } - @Override + /** + * Calls the underlying chat model with a prompt and returns the corresponding chat + * response. + * @param prompt The prompt to be used for the chat model. + * @return The chat response containing the generated messages. + * @deprecated This method is deprecated as of version 1.0.0 M1 and will be removed in + * a future release. Use the method builder(chatModel).build().prompt(prompt).call() + * instead. + */ + @Deprecated(since = "1.0.0 M1", forRemoval = true) ChatResponse call(Prompt prompt); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatModel.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatModel.java new file mode 100644 index 000000000..4d345bd5c --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatModel.java @@ -0,0 +1,47 @@ +/* + * Copyright 2023 - 2024 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.ai.chat; + +import org.springframework.ai.chat.prompt.ChatOptions; +import org.springframework.ai.chat.prompt.Prompt; + +import java.util.Arrays; + +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.model.Model; + +// @FunctionalInterface +public interface ChatModel extends Model { + + default String call(String message) { + Prompt prompt = new Prompt(new UserMessage(message)); + Generation generation = call(prompt).getResult(); + return (generation != null) ? generation.getOutput().getContent() : ""; + } + + default String call(Message... messages) { + Prompt prompt = new Prompt(Arrays.asList(messages)); + Generation generation = call(prompt).getResult(); + return (generation != null) ? generation.getOutput().getContent() : ""; + } + + @Override + ChatResponse call(Prompt prompt); + + ChatOptions getDefaultOptions(); + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/DefaultChatClient.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/DefaultChatClient.java new file mode 100644 index 000000000..a2e4dcae6 --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/DefaultChatClient.java @@ -0,0 +1,43 @@ +package org.springframework.ai.chat; + +import org.springframework.ai.chat.prompt.Prompt; + +/** + * @author Mark Pollack + * @author Christian Tzolov + * @author Josh Long + * @author Arjen Poutsma + */ +class DefaultChatClient implements ChatClient { + + private final ChatModel chatModel; + + private final ChatClientRequest defaultChatClientRequest; + + public DefaultChatClient(ChatModel chatModel, ChatClientRequest defaultChatClientRequest) { + this.chatModel = chatModel; + this.defaultChatClientRequest = defaultChatClientRequest; + } + + @Override + public ChatClientRequest prompt() { + return new ChatClientRequest(this.defaultChatClientRequest); + } + + @Override + public ChatClientPromptRequest prompt(Prompt prompt) { + return new ChatClientPromptRequest(this.chatModel, prompt); + } + + /** + * use the new fluid DSL starting in {@link #prompt()} + * @param prompt the {@link Prompt prompt} object + * @return a {@link ChatResponse chat response} + */ + @Deprecated(forRemoval = true, since = "1.0.0 M1") + @Override + public ChatResponse call(Prompt prompt) { + return this.chatModel.call(prompt); + } + +} \ No newline at end of file diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/StreamingChatClient.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/StreamingChatModel.java similarity index 91% rename from spring-ai-core/src/main/java/org/springframework/ai/chat/StreamingChatClient.java rename to spring-ai-core/src/main/java/org/springframework/ai/chat/StreamingChatModel.java index 69634b192..e6b96f30e 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/StreamingChatClient.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/StreamingChatModel.java @@ -21,10 +21,10 @@ import reactor.core.publisher.Flux; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.model.StreamingModelClient; +import org.springframework.ai.model.StreamingModel; @FunctionalInterface -public interface StreamingChatClient extends StreamingModelClient { +public interface StreamingChatModel extends StreamingModel { default Flux stream(String message) { Prompt prompt = new Prompt(message); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/service/PromptTransformingChatService.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/PromptTransformingChatService.java index 761ae5f42..4cb686c3c 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/service/PromptTransformingChatService.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/PromptTransformingChatService.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.chat.service; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; import org.springframework.ai.chat.prompt.transformer.PromptTransformer; @@ -35,7 +35,7 @@ import java.util.Objects; */ public class PromptTransformingChatService implements ChatService { - private ChatClient chatClient; + private ChatModel chatModel; private List retrievers; @@ -45,19 +45,19 @@ public class PromptTransformingChatService implements ChatService { private List chatServiceListeners; - public PromptTransformingChatService(ChatClient chatClient, List retrievers, + public PromptTransformingChatService(ChatModel chatModel, List retrievers, List documentPostProcessors, List augmentors, List chatServiceListeners) { - Objects.requireNonNull(chatClient, "chatClient must not be null"); - this.chatClient = chatClient; + Objects.requireNonNull(chatModel, "chatModel must not be null"); + this.chatModel = chatModel; this.retrievers = retrievers; this.documentPostProcessors = documentPostProcessors; this.augmentors = augmentors; this.chatServiceListeners = chatServiceListeners; } - public static Builder builder(ChatClient chatClient) { - return new Builder().withChatClient(chatClient); + public static Builder builder(ChatModel chatModel) { + return new Builder().withChatModel(chatModel); } @Override @@ -86,7 +86,7 @@ public class PromptTransformingChatService implements ChatService { } // Perform generation - ChatResponse chatResponse = this.chatClient.call(chatServiceContext.getPrompt()); + ChatResponse chatResponse = this.chatModel.call(chatServiceContext.getPrompt()); // Invoke Listeners onComplete ChatServiceResponse chatServiceResponse = new ChatServiceResponse(chatServiceContext, chatResponse); @@ -98,7 +98,7 @@ public class PromptTransformingChatService implements ChatService { public static class Builder { - private ChatClient chatClient; + private ChatModel chatModel; private List retrievers = new ArrayList<>(); @@ -108,8 +108,8 @@ public class PromptTransformingChatService implements ChatService { private List chatServiceListeners = new ArrayList<>(); - public Builder withChatClient(ChatClient chatClient) { - this.chatClient = chatClient; + public Builder withChatModel(ChatModel chatModel) { + this.chatModel = chatModel; return this; } @@ -134,7 +134,7 @@ public class PromptTransformingChatService implements ChatService { } public PromptTransformingChatService build() { - return new PromptTransformingChatService(chatClient, retrievers, documentPostProcessors, augmentors, + return new PromptTransformingChatService(chatModel, retrievers, documentPostProcessors, augmentors, chatServiceListeners); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/service/StreamingPromptTransformingChatService.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/StreamingPromptTransformingChatService.java index c94241531..6a33a2262 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/service/StreamingPromptTransformingChatService.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/service/StreamingPromptTransformingChatService.java @@ -23,7 +23,7 @@ import org.springframework.ai.chat.prompt.transformer.ChatServiceContext; import reactor.core.publisher.Flux; import org.springframework.ai.chat.ChatResponse; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.messages.MessageAggregator; import org.springframework.ai.chat.prompt.transformer.PromptTransformer; @@ -33,7 +33,7 @@ import org.springframework.ai.chat.prompt.transformer.PromptTransformer; */ public class StreamingPromptTransformingChatService implements StreamingChatService { - private StreamingChatClient streamingChatClient; + private StreamingChatModel streamingChatModel; private List retrievers; @@ -43,19 +43,19 @@ public class StreamingPromptTransformingChatService implements StreamingChatServ private List chatServiceListeners; - public StreamingPromptTransformingChatService(StreamingChatClient chatClient, List retrievers, + public StreamingPromptTransformingChatService(StreamingChatModel chatModel, List retrievers, List documentPostProcessors, List augmentors, List chatServiceListeners) { - Objects.requireNonNull(chatClient, "chatClient must not be null"); - this.streamingChatClient = chatClient; + Objects.requireNonNull(chatModel, "chatModel must not be null"); + this.streamingChatModel = chatModel; this.retrievers = retrievers; this.documentPostProcessors = documentPostProcessors; this.augmentors = augmentors; this.chatServiceListeners = chatServiceListeners; } - public static Builder builder(StreamingChatClient chatClient) { - return new Builder().withChatClient(chatClient); + public static Builder builder(StreamingChatModel chatModel) { + return new Builder().withChatModel(chatModel); } @Override @@ -87,7 +87,7 @@ public class StreamingPromptTransformingChatService implements StreamingChatServ final var promptContext2 = chatServiceContext; Flux fluxChatResponse = new MessageAggregator() - .aggregate(this.streamingChatClient.stream(chatServiceContext.getPrompt()), chatResponse -> { + .aggregate(this.streamingChatModel.stream(chatServiceContext.getPrompt()), chatResponse -> { for (ChatServiceListener listener : this.chatServiceListeners) { listener.onComplete(new ChatServiceResponse(promptContext2, chatResponse)); } @@ -99,7 +99,7 @@ public class StreamingPromptTransformingChatService implements StreamingChatServ public static class Builder { - private StreamingChatClient chatClient; + private StreamingChatModel chatModel; private List retrievers = new ArrayList<>(); @@ -109,8 +109,8 @@ public class StreamingPromptTransformingChatService implements StreamingChatServ private List chatServiceListeners = new ArrayList<>(); - public Builder withChatClient(StreamingChatClient chatClient) { - this.chatClient = chatClient; + public Builder withChatModel(StreamingChatModel chatModel) { + this.chatModel = chatModel; return this; } @@ -135,8 +135,8 @@ public class StreamingPromptTransformingChatService implements StreamingChatServ } public StreamingPromptTransformingChatService build() { - return new StreamingPromptTransformingChatService(chatClient, retrievers, documentPostProcessors, - augmentors, chatServiceListeners); + return new StreamingPromptTransformingChatService(chatModel, retrievers, documentPostProcessors, augmentors, + chatServiceListeners); } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/converter/README.md b/spring-ai-core/src/main/java/org/springframework/ai/converter/README.md index f1a4f7004..6b4f9aa3f 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/converter/README.md +++ b/spring-ai-core/src/main/java/org/springframework/ai/converter/README.md @@ -1,7 +1,7 @@ # Structured Output * [Documentation](https://docs.spring.io/spring-ai/reference/concepts.html#_output_parsing) -* [Usage examples](https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java) +* [Usage examples](https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelIT.java) The output of AI models traditionally arrives as a text, even if you ask for the reply to be in JSON. It may be the correct JSON, but it isn’t a JSON data structure. diff --git a/spring-ai-core/src/main/java/org/springframework/ai/embedding/AbstractEmbeddingClient.java b/spring-ai-core/src/main/java/org/springframework/ai/embedding/AbstractEmbeddingModel.java similarity index 82% rename from spring-ai-core/src/main/java/org/springframework/ai/embedding/AbstractEmbeddingClient.java rename to spring-ai-core/src/main/java/org/springframework/ai/embedding/AbstractEmbeddingModel.java index b7749153f..6d857d3a3 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/embedding/AbstractEmbeddingClient.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/embedding/AbstractEmbeddingModel.java @@ -24,12 +24,12 @@ import java.util.stream.Collectors; import org.springframework.core.io.DefaultResourceLoader; /** - * Abstract implementation of the {@link EmbeddingClient} interface that provides + * Abstract implementation of the {@link EmbeddingModel} interface that provides * dimensions calculation caching. * * @author Christian Tzolov */ -public abstract class AbstractEmbeddingClient implements EmbeddingClient { +public abstract class AbstractEmbeddingModel implements EmbeddingModel { protected final AtomicInteger embeddingDimensions = new AtomicInteger(-1); @@ -37,14 +37,14 @@ public abstract class AbstractEmbeddingClient implements EmbeddingClient { /** * Return the dimension of the requested embedding generative name. If the generative - * name is unknown uses the EmbeddingClient to perform a dummy EmbeddingClient#embed - * and count the response dimensions. - * @param embeddingClient Fall-back client to determine, empirically the dimensions. + * name is unknown uses the EmbeddingModel to perform a dummy EmbeddingModel#embed and + * count the response dimensions. + * @param embeddingModel Fall-back client to determine, empirically the dimensions. * @param modelName Embedding generative name to retrieve the dimensions for. * @param dummyContent Dummy content to use for the empirical dimension calculation. * @return Returns the embedding dimensions for the modelName. */ - public static int dimensions(EmbeddingClient embeddingClient, String modelName, String dummyContent) { + public static int dimensions(EmbeddingModel embeddingModel, String modelName, String dummyContent) { if (KNOWN_EMBEDDING_DIMENSIONS.containsKey(modelName)) { // Retrieve the dimension from a pre-configured file. @@ -53,7 +53,7 @@ public abstract class AbstractEmbeddingClient implements EmbeddingClient { else { // Determine the dimensions empirically. // Generate an embedding and count the dimension size; - return embeddingClient.embed(dummyContent).size(); + return embeddingModel.embed(dummyContent).size(); } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/embedding/EmbeddingClient.java b/spring-ai-core/src/main/java/org/springframework/ai/embedding/EmbeddingModel.java similarity index 91% rename from spring-ai-core/src/main/java/org/springframework/ai/embedding/EmbeddingClient.java rename to spring-ai-core/src/main/java/org/springframework/ai/embedding/EmbeddingModel.java index d9dc2f8a3..66a972f62 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/embedding/EmbeddingClient.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/embedding/EmbeddingModel.java @@ -16,15 +16,15 @@ package org.springframework.ai.embedding; import org.springframework.ai.document.Document; -import org.springframework.ai.model.ModelClient; +import org.springframework.ai.model.Model; import org.springframework.util.Assert; import java.util.List; /** - * EmbeddingClient is a generic interface for embedding clients. + * EmbeddingModel is a generic interface for embedding models. */ -public interface EmbeddingClient extends ModelClient { +public interface EmbeddingModel extends Model { @Override EmbeddingResponse call(EmbeddingRequest request); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/evaluation/RelevancyEvaluator.java b/spring-ai-core/src/main/java/org/springframework/ai/evaluation/RelevancyEvaluator.java index 8afc9dcff..8a6b91815 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/evaluation/RelevancyEvaluator.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/evaluation/RelevancyEvaluator.java @@ -1,6 +1,6 @@ package org.springframework.ai.evaluation; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.messages.MessageType; @@ -31,14 +31,14 @@ public class RelevancyEvaluator implements Evaluator { private final ChatOptions chatOptions; - private ChatClient chatClient; + private ChatModel chatModel; - public RelevancyEvaluator(ChatClient chatClient) { - this(chatClient, ChatOptionsBuilder.builder().build()); + public RelevancyEvaluator(ChatModel chatModel) { + this(chatModel, ChatOptionsBuilder.builder().build()); } - public RelevancyEvaluator(ChatClient chatClient, ChatOptions chatOptions) { - this.chatClient = chatClient; + public RelevancyEvaluator(ChatModel chatModel, ChatOptions chatOptions) { + this.chatModel = chatModel; this.chatOptions = chatOptions; } @@ -52,7 +52,7 @@ public class RelevancyEvaluator implements Evaluator { Message message = promptTemplate .createMessage(Map.of("query", query, "response", response, "context", context)); - ChatResponse chatResponse = this.chatClient.call(new Prompt(message, this.chatOptions)); + ChatResponse chatResponse = this.chatModel.call(new Prompt(message, this.chatOptions)); var evaluationResponse = chatResponse.getResult().getOutput().getContent(); boolean passing = false; diff --git a/spring-ai-core/src/main/java/org/springframework/ai/image/ImageClient.java b/spring-ai-core/src/main/java/org/springframework/ai/image/ImageModel.java similarity index 85% rename from spring-ai-core/src/main/java/org/springframework/ai/image/ImageClient.java rename to spring-ai-core/src/main/java/org/springframework/ai/image/ImageModel.java index 7ce6a4749..493da50bf 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/image/ImageClient.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/image/ImageModel.java @@ -15,10 +15,10 @@ */ package org.springframework.ai.image; -import org.springframework.ai.model.ModelClient; +import org.springframework.ai.model.Model; @FunctionalInterface -public interface ImageClient extends ModelClient { +public interface ImageModel extends Model { ImageResponse call(ImagePrompt request); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/ModelClient.java b/spring-ai-core/src/main/java/org/springframework/ai/model/Model.java similarity index 82% rename from spring-ai-core/src/main/java/org/springframework/ai/model/ModelClient.java rename to spring-ai-core/src/main/java/org/springframework/ai/model/Model.java index 1f114a59b..1671a3dc4 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/ModelClient.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/Model.java @@ -16,8 +16,8 @@ package org.springframework.ai.model; /** - * The ModelClient interface provides a generic API for invoking AI models. It is designed - * to handle the interaction with various types of AI models by abstracting the process of + * The Model interface provides a generic API for invoking AI models. It is designed to + * handle the interaction with various types of AI models by abstracting the process of * sending requests and receiving responses. The interface uses Java generics to * accommodate different types of requests and responses, enhancing flexibility and * adaptability across different AI model implementations. @@ -27,7 +27,7 @@ package org.springframework.ai.model; * @author Mark Pollack * @since 0.8.0 */ -public interface ModelClient, TRes extends ModelResponse> { +public interface Model, TRes extends ModelResponse> { /** * Executes a method call to the AI model. diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/ModelDescription.java b/spring-ai-core/src/main/java/org/springframework/ai/model/ModelDescription.java new file mode 100644 index 000000000..57b083fab --- /dev/null +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/ModelDescription.java @@ -0,0 +1,38 @@ +/* + * Copyright 2024-2024 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.springframework.ai.model; + +/** + * @author Christian Tzolov + */ +public interface ModelDescription { + + String getModelName(); + + default String getDescription() { + return ""; + } + + default String getVersion() { + return ""; + } + + default int getContextLength() { + return -1; + } + +} diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/StreamingModelClient.java b/spring-ai-core/src/main/java/org/springframework/ai/model/StreamingModel.java similarity index 82% rename from spring-ai-core/src/main/java/org/springframework/ai/model/StreamingModelClient.java rename to spring-ai-core/src/main/java/org/springframework/ai/model/StreamingModel.java index 4dddced14..2c1de77a9 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/StreamingModelClient.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/StreamingModel.java @@ -18,8 +18,8 @@ package org.springframework.ai.model; import reactor.core.publisher.Flux; /** - * The StreamingModelClient interface provides a generic API for invoking an AI models - * with streaming response. It abstracts the process of sending requests and receiving a + * The StreamingModel interface provides a generic API for invoking an AI models with + * streaming response. It abstracts the process of sending requests and receiving a * streaming responses. The interface uses Java generics to accommodate different types of * requests and responses, enhancing flexibility and adaptability across different AI * model implementations. @@ -30,7 +30,7 @@ import reactor.core.publisher.Flux; * @author Christian Tzolov * @since 0.8.0 */ -public interface StreamingModelClient, TResChunk extends ModelResponse> { +public interface StreamingModel, TResChunk extends ModelResponse> { /** * Executes a method call to the AI model. diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractFunctionCallback.java b/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractFunctionCallback.java index f4cdd4ef8..6bd639c88 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractFunctionCallback.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractFunctionCallback.java @@ -54,7 +54,7 @@ abstract class AbstractFunctionCallback implements Function, Functio /** * Constructs a new {@link AbstractFunctionCallback} with the given name, description, * input type and default object mapper. - * @param name Function name. Should be unique within the ChatClient's function + * @param name Function name. Should be unique within the ChatModel's function * registry. * @param description Function description. Used as a "system prompt" by the model to * decide if the function should be called. diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptions.java b/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptions.java index 146b35c47..df603d2b3 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptions.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallingOptions.java @@ -24,32 +24,32 @@ import java.util.Set; public interface FunctionCallingOptions { /** - * Function Callbacks to be registered with the ChatClient. For Prompt Options the + * Function Callbacks to be registered with the ChatModel. For Prompt Options the * functionCallbacks are automatically enabled for the duration of the prompt * execution. For Default Options the FunctionCallbacks are registered but disabled by * default. You have to use "functions" property to list the function names from the - * ChatClient registry to be used in the chat completion requests. - * @return Return the Function Callbacks to be registered with the ChatClient. + * ChatModel registry to be used in the chat completion requests. + * @return Return the Function Callbacks to be registered with the ChatModel. */ List getFunctionCallbacks(); /** - * Set the Function Callbacks to be registered with the ChatClient. + * Set the Function Callbacks to be registered with the ChatModel. * @param functionCallbacks the Function Callbacks to be registered with the - * ChatClient. + * ChatModel. */ void setFunctionCallbacks(List functionCallbacks); /** - * @return List of function names from the ChatClient registry to be used in the next + * @return List of function names from the ChatModel registry to be used in the next * chat completion requests. */ Set getFunctions(); /** - * Set the list of function names from the ChatClient registry to be used in the next + * Set the list of function names from the ChatModel registry to be used in the next * chat completion requests. - * @param functions the list of function names from the ChatClient registry to be used + * @param functions the list of function names from the ChatModel registry to be used * in the next chat completion requests. */ void setFunctions(Set functions); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/transformer/KeywordMetadataEnricher.java b/spring-ai-core/src/main/java/org/springframework/ai/transformer/KeywordMetadataEnricher.java index 67807e419..6156be16d 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/transformer/KeywordMetadataEnricher.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/transformer/KeywordMetadataEnricher.java @@ -18,7 +18,7 @@ package org.springframework.ai.transformer; import java.util.List; import java.util.Map; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.document.Document; import org.springframework.ai.document.DocumentTransformer; import org.springframework.ai.chat.prompt.Prompt; @@ -43,18 +43,18 @@ public class KeywordMetadataEnricher implements DocumentTransformer { /** * Model predictor */ - private final ChatClient chatClient; + private final ChatModel chatModel; /** * The number of keywords to extract. */ private final int keywordCount; - public KeywordMetadataEnricher(ChatClient chatClient, int keywordCount) { - Assert.notNull(chatClient, "ChatClient must not be null"); + public KeywordMetadataEnricher(ChatModel chatModel, int keywordCount) { + Assert.notNull(chatModel, "ChatModel must not be null"); Assert.isTrue(keywordCount >= 1, "Document count must be >= 1"); - this.chatClient = chatClient; + this.chatModel = chatModel; this.keywordCount = keywordCount; } @@ -64,7 +64,7 @@ public class KeywordMetadataEnricher implements DocumentTransformer { var template = new PromptTemplate(String.format(KEYWORDS_TEMPLATE, keywordCount)); Prompt prompt = template.create(Map.of(CONTEXT_STR_PLACEHOLDER, document.getContent())); - String keywords = this.chatClient.call(prompt).getResult().getOutput().getContent(); + String keywords = this.chatModel.call(prompt).getResult().getOutput().getContent(); document.getMetadata().putAll(Map.of(EXCERPT_KEYWORDS_METADATA_KEY, keywords)); } return documents; diff --git a/spring-ai-core/src/main/java/org/springframework/ai/transformer/SummaryMetadataEnricher.java b/spring-ai-core/src/main/java/org/springframework/ai/transformer/SummaryMetadataEnricher.java index e882ac497..154886a61 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/transformer/SummaryMetadataEnricher.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/transformer/SummaryMetadataEnricher.java @@ -20,7 +20,7 @@ import java.util.HashMap; import java.util.List; import java.util.Map; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.document.Document; import org.springframework.ai.document.DocumentTransformer; import org.springframework.ai.document.MetadataMode; @@ -62,7 +62,7 @@ public class SummaryMetadataEnricher implements DocumentTransformer { /** * AI client. */ - private final ChatClient chatClient; + private final ChatModel chatModel; /** * Number of documents from front to use for title extraction. @@ -76,16 +76,16 @@ public class SummaryMetadataEnricher implements DocumentTransformer { */ private final String summaryTemplate; - public SummaryMetadataEnricher(ChatClient chatClient, List summaryTypes) { - this(chatClient, summaryTypes, DEFAULT_SUMMARY_EXTRACT_TEMPLATE, MetadataMode.ALL); + public SummaryMetadataEnricher(ChatModel chatModel, List summaryTypes) { + this(chatModel, summaryTypes, DEFAULT_SUMMARY_EXTRACT_TEMPLATE, MetadataMode.ALL); } - public SummaryMetadataEnricher(ChatClient chatClient, List summaryTypes, String summaryTemplate, + public SummaryMetadataEnricher(ChatModel chatModel, List summaryTypes, String summaryTemplate, MetadataMode metadataMode) { - Assert.notNull(chatClient, "ChatClient must not be null"); + Assert.notNull(chatModel, "ChatModel must not be null"); Assert.hasText(summaryTemplate, "Summary template must not be empty"); - this.chatClient = chatClient; + this.chatModel = chatModel; this.summaryTypes = CollectionUtils.isEmpty(summaryTypes) ? List.of(SummaryType.CURRENT) : summaryTypes; this.metadataMode = metadataMode; this.summaryTemplate = summaryTemplate; @@ -101,7 +101,7 @@ public class SummaryMetadataEnricher implements DocumentTransformer { Prompt prompt = new PromptTemplate(this.summaryTemplate) .create(Map.of(CONTEXT_STR_PLACEHOLDER, documentContext)); - documentSummaries.add(this.chatClient.call(prompt).getResult().getOutput().getContent()); + documentSummaries.add(this.chatModel.call(prompt).getResult().getOutput().getContent()); } for (int i = 0; i < documentSummaries.size(); i++) { diff --git a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/SimpleVectorStore.java b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/SimpleVectorStore.java index 21db81da0..7c681b678 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/SimpleVectorStore.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/vectorstore/SimpleVectorStore.java @@ -22,7 +22,7 @@ import com.fasterxml.jackson.databind.ObjectWriter; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.core.io.Resource; import java.io.File; @@ -61,18 +61,18 @@ public class SimpleVectorStore implements VectorStore { protected Map store = new ConcurrentHashMap<>(); - protected EmbeddingClient embeddingClient; + protected EmbeddingModel embeddingModel; - public SimpleVectorStore(EmbeddingClient embeddingClient) { - Objects.requireNonNull(embeddingClient, "EmbeddingClient must not be null"); - this.embeddingClient = embeddingClient; + public SimpleVectorStore(EmbeddingModel embeddingModel) { + Objects.requireNonNull(embeddingModel, "EmbeddingModel must not be null"); + this.embeddingModel = embeddingModel; } @Override public void add(List documents) { for (Document document : documents) { - logger.info("Calling EmbeddingClient for document id = {}", document.getId()); - List embedding = this.embeddingClient.embed(document); + logger.info("Calling EmbeddingModel for document id = {}", document.getId()); + List embedding = this.embeddingModel.embed(document); document.setEmbedding(embedding); this.store.put(document.getId(), document); } @@ -187,7 +187,7 @@ public class SimpleVectorStore implements VectorStore { } private List getUserQueryEmbedding(String query) { - return this.embeddingClient.embed(query); + return this.embeddingModel.embed(query); } public static class Similarity { diff --git a/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatClientTest.java b/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatClientTest.java new file mode 100644 index 000000000..44645a4e4 --- /dev/null +++ b/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatClientTest.java @@ -0,0 +1,136 @@ +/* + * Copyright 2024-2024 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.springframework.ai.chat; + +import java.net.MalformedURLException; +import java.net.URL; +import java.util.List; + +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Captor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.MessageType; +import org.springframework.ai.chat.messages.UserMessage; +import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.model.function.FunctionCallingOptions; +import org.springframework.util.MimeTypeUtils; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.when; + +/** + * @author Christian Tzolov + */ +@ExtendWith(MockitoExtension.class) +public class ChatClientTest { + + @Mock + ChatModel chatModel; + + @Captor + ArgumentCaptor promptCaptor; + + @BeforeEach + public void beforeAll() { + when(chatModel.call(promptCaptor.capture())) + .thenReturn(new ChatResponse(List.of(new Generation("response")))); + } + + @Test + public void simpleUserPrompt() { + assertThat(ChatClient.builder(chatModel).build().prompt().user("User prompt").call().content()) + .isEqualTo("response"); + + Message userMessage = promptCaptor.getValue().getInstructions().get(0); + assertThat(userMessage.getContent()).isEqualTo("User prompt"); + assertThat(userMessage.getMessageType()).isEqualTo(MessageType.USER); + } + + @Test + public void simpleUserPromptObject() throws MalformedURLException { + UserMessage message = new UserMessage("User prompt"); + Prompt prompt = new Prompt(message); + assertThat(ChatClient.builder(chatModel).build().prompt(prompt).call().content()).isEqualTo("response"); + + Message userMessage = promptCaptor.getValue().getInstructions().get(0); + assertThat(userMessage.getContent()).isEqualTo("User prompt"); + assertThat(userMessage.getMessageType()).isEqualTo(MessageType.USER); + } + + @Test + public void simpleSystemPrompt() throws MalformedURLException { + String response = ChatClient.builder(chatModel).build().prompt().system("System prompt").call().content(); + + assertThat(response).isEqualTo("response"); + + assertThat(promptCaptor.getValue().getInstructions()).hasSize(2); + + Message systemMessage = promptCaptor.getValue().getInstructions().get(0); + assertThat(systemMessage.getContent()).isEqualTo("System prompt"); + assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM); + + // Is this expected? + Message userMessage = promptCaptor.getValue().getInstructions().get(1); + assertThat(userMessage.getContent()).isEqualTo(""); + assertThat(userMessage.getMessageType()).isEqualTo(MessageType.USER); + } + + @Test + public void complexCall() throws MalformedURLException { + + var options = FunctionCallingOptions.builder().build(); + when(chatModel.getDefaultOptions()).thenReturn(options); + + var url = new URL("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png"); + + // @formatter:off + ChatClient client = ChatClient.builder(chatModel) + .defaultSystem("System text") + .defaultFunctions("function1") + .build(); + + String response = client.prompt() + .user(u -> u.text("User text {music}").param("music", "Rock").media(MimeTypeUtils.IMAGE_PNG, url)) + .call() + .content(); + // @formatter:on + + assertThat(response).isEqualTo("response"); + assertThat(promptCaptor.getValue().getInstructions()).hasSize(2); + + Message systemMessage = promptCaptor.getValue().getInstructions().get(0); + assertThat(systemMessage.getContent()).isEqualTo("System text"); + assertThat(systemMessage.getMessageType()).isEqualTo(MessageType.SYSTEM); + + Message userMessage = promptCaptor.getValue().getInstructions().get(1); + assertThat(userMessage.getContent()).isEqualTo("User text Rock"); + assertThat(userMessage.getMessageType()).isEqualTo(MessageType.USER); + assertThat(userMessage.getMedia()).hasSize(1); + assertThat(userMessage.getMedia().iterator().next().getMimeType()).isEqualTo(MimeTypeUtils.IMAGE_PNG); + assertThat(userMessage.getMedia().iterator().next().getData()) + .isEqualTo("https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png"); + + assertThat(options.getFunctions()).containsExactly("function1"); + } + +} diff --git a/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatClientTests.java b/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatModelTests.java similarity index 96% rename from spring-ai-core/src/test/java/org/springframework/ai/chat/ChatClientTests.java rename to spring-ai-core/src/test/java/org/springframework/ai/chat/ChatModelTests.java index 28e4a528c..5eb9796f1 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatClientTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatModelTests.java @@ -34,12 +34,12 @@ import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.prompt.Prompt; /** - * Unit Tests for {@link ChatClient}. + * Unit Tests for {@link ChatModel}. * * @author John Blum * @since 0.2.0 */ -class ChatClientTests { +class ChatModelTests { @Test void generateWithStringCallsGenerateWithPromptAndReturnsResponseCorrectly() { @@ -47,7 +47,7 @@ class ChatClientTests { String userMessage = "Zero Wing"; String responseMessage = "All your bases are belong to us"; - ChatClient mockClient = Mockito.mock(ChatClient.class); + ChatModel mockClient = Mockito.mock(ChatModel.class); AssistantMessage mockAssistantMessage = Mockito.mock(AssistantMessage.class); when(mockAssistantMessage.getContent()).thenReturn(responseMessage); diff --git a/spring-ai-core/src/test/java/org/springframework/ai/chat/memory/ChatMemoryTests.java b/spring-ai-core/src/test/java/org/springframework/ai/chat/memory/ChatMemoryTests.java index 741fd5341..5859d66c6 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/chat/memory/ChatMemoryTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/chat/memory/ChatMemoryTests.java @@ -25,10 +25,10 @@ import org.mockito.Captor; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; -import org.springframework.ai.chat.StreamingChatClient; +import org.springframework.ai.chat.StreamingChatModel; import org.springframework.ai.chat.service.ChatServiceResponse; import org.springframework.ai.chat.service.PromptTransformingChatService; import org.springframework.ai.chat.messages.Message; @@ -48,10 +48,10 @@ import static org.mockito.Mockito.when; public class ChatMemoryTests { @Mock - ChatClient chatClient; + ChatModel chatModel; @Mock - StreamingChatClient streamingChatClient; + StreamingChatModel streamingChatModel; @Captor ArgumentCaptor promptCaptor; @@ -61,7 +61,7 @@ public class ChatMemoryTests { ChatMemory chatHistory = new InMemoryChatMemory(); - PromptTransformingChatService chatService = PromptTransformingChatService.builder(chatClient) + PromptTransformingChatService chatService = PromptTransformingChatService.builder(chatModel) .withRetrievers(List.of(ChatMemoryRetriever.builder().withChatHistory(chatHistory).build())) .withContentPostProcessors( List.of(new LastMaxTokenSizeContentTransformer(new JTokkitTokenCountEstimator(), 10))) @@ -69,7 +69,7 @@ public class ChatMemoryTests { .withChatServiceListeners(List.of(new ChatMemoryChatServiceListener(chatHistory))) .build(); - chatClientUserMessages(chatService, chatHistory); + chatModelUserMessages(chatService, chatHistory); } @Test @@ -77,7 +77,7 @@ public class ChatMemoryTests { ChatMemory chatHistory = new InMemoryChatMemory(); - PromptTransformingChatService chatService = PromptTransformingChatService.builder(chatClient) + PromptTransformingChatService chatService = PromptTransformingChatService.builder(chatModel) .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) .withContentPostProcessors( List.of(new LastMaxTokenSizeContentTransformer(new JTokkitTokenCountEstimator(), 10))) @@ -85,12 +85,12 @@ public class ChatMemoryTests { .withChatServiceListeners(List.of(new ChatMemoryChatServiceListener(chatHistory))) .build(); - chatClientUserMessages(chatService, chatHistory); + chatModelUserMessages(chatService, chatHistory); } - public void chatClientUserMessages(PromptTransformingChatService chatService, ChatMemory chatHistory) { + public void chatModelUserMessages(PromptTransformingChatService chatService, ChatMemory chatHistory) { - when(chatClient.call(promptCaptor.capture())) + when(chatModel.call(promptCaptor.capture())) .thenReturn(new ChatResponse(List.of(new Generation("assistant:1")))) .thenReturn(new ChatResponse(List.of(new Generation("assistant:2")))) .thenReturn(new ChatResponse(List.of(new Generation("assistant:3")))); diff --git a/spring-ai-core/src/test/java/org/springframework/ai/embedding/AbstractEmbeddingClientTests.java b/spring-ai-core/src/test/java/org/springframework/ai/embedding/AbstractEmbeddingModelTests.java similarity index 82% rename from spring-ai-core/src/test/java/org/springframework/ai/embedding/AbstractEmbeddingClientTests.java rename to spring-ai-core/src/test/java/org/springframework/ai/embedding/AbstractEmbeddingModelTests.java index ad42fc271..d6d51454b 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/embedding/AbstractEmbeddingClientTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/embedding/AbstractEmbeddingModelTests.java @@ -37,15 +37,15 @@ import static org.mockito.Mockito.when; * @author Christian Tzolov */ @ExtendWith(MockitoExtension.class) -public class AbstractEmbeddingClientTests { +public class AbstractEmbeddingModelTests { @Mock - private EmbeddingClient embeddingClient; + private EmbeddingModel embeddingModel; @Test public void testDefaultMethodImplementation() { - EmbeddingClient dummy = new EmbeddingClient() { + EmbeddingModel dummy = new EmbeddingModel() { @Override public List embed(String text) { @@ -79,16 +79,16 @@ public class AbstractEmbeddingClientTests { @ParameterizedTest @CsvFileSource(resources = "/embedding/embedding-model-dimensions.properties", numLinesToSkip = 1, delimiter = '=') public void testKnownEmbeddingModelDimensions(String model, String dimension) { - assertThat(AbstractEmbeddingClient.dimensions(embeddingClient, model, "Hello world!")) + assertThat(AbstractEmbeddingModel.dimensions(embeddingModel, model, "Hello world!")) .isEqualTo(Integer.valueOf(dimension)); - verify(embeddingClient, never()).embed(any(String.class)); - verify(embeddingClient, never()).embed(any(Document.class)); + verify(embeddingModel, never()).embed(any(String.class)); + verify(embeddingModel, never()).embed(any(Document.class)); } @Test public void testUnknownModelDimension() { - when(embeddingClient.embed(eq("Hello world!"))).thenReturn(List.of(0.1, 0.1, 0.1)); - assertThat(AbstractEmbeddingClient.dimensions(embeddingClient, "unknown_model", "Hello world!")).isEqualTo(3); + when(embeddingModel.embed(eq("Hello world!"))).thenReturn(List.of(0.1, 0.1, 0.1)); + assertThat(AbstractEmbeddingModel.dimensions(embeddingModel, "unknown_model", "Hello world!")).isEqualTo(3); } } diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/images/chat-options-flow.jpg b/spring-ai-docs/src/main/antora/modules/ROOT/images/chat-options-flow.jpg index 4d259d79a..a4a4ece10 100644 Binary files a/spring-ai-docs/src/main/antora/modules/ROOT/images/chat-options-flow.jpg and b/spring-ai-docs/src/main/antora/modules/ROOT/images/chat-options-flow.jpg differ diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/images/embeddings-api.jpg b/spring-ai-docs/src/main/antora/modules/ROOT/images/embeddings-api.jpg index 41b383ff3..69d2f5335 100644 Binary files a/spring-ai-docs/src/main/antora/modules/ROOT/images/embeddings-api.jpg and b/spring-ai-docs/src/main/antora/modules/ROOT/images/embeddings-api.jpg differ diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/images/openai-chatclient-function-call.jpg b/spring-ai-docs/src/main/antora/modules/ROOT/images/openai-chatclient-function-call.jpg index b92da642d..bef40fc97 100644 Binary files a/spring-ai-docs/src/main/antora/modules/ROOT/images/openai-chatclient-function-call.jpg and b/spring-ai-docs/src/main/antora/modules/ROOT/images/openai-chatclient-function-call.jpg differ diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/images/spring-ai-chat-api.jpg b/spring-ai-docs/src/main/antora/modules/ROOT/images/spring-ai-chat-api.jpg index 371b3f4e9..96f33fd52 100644 Binary files a/spring-ai-docs/src/main/antora/modules/ROOT/images/spring-ai-chat-api.jpg and b/spring-ai-docs/src/main/antora/modules/ROOT/images/spring-ai-chat-api.jpg differ diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/images/spring-ai-chat-completions-clients.jpg b/spring-ai-docs/src/main/antora/modules/ROOT/images/spring-ai-chat-completions-clients.jpg index f20701c82..5e4ff4b7f 100644 Binary files a/spring-ai-docs/src/main/antora/modules/ROOT/images/spring-ai-chat-completions-clients.jpg and b/spring-ai-docs/src/main/antora/modules/ROOT/images/spring-ai-chat-completions-clients.jpg differ diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/images/spring-ai-generic-model-api.jpg b/spring-ai-docs/src/main/antora/modules/ROOT/images/spring-ai-generic-model-api.jpg index 4fcdadaf6..d23f3f8f5 100644 Binary files a/spring-ai-docs/src/main/antora/modules/ROOT/images/spring-ai-generic-model-api.jpg and b/spring-ai-docs/src/main/antora/modules/ROOT/images/spring-ai-generic-model-api.jpg differ diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/images/structured-output-api.jpg b/spring-ai-docs/src/main/antora/modules/ROOT/images/structured-output-api.jpg index 7a4ae6d96..de792b4b3 100644 Binary files a/spring-ai-docs/src/main/antora/modules/ROOT/images/structured-output-api.jpg and b/spring-ai-docs/src/main/antora/modules/ROOT/images/structured-output-api.jpg differ diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/nav.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/nav.adoc index 62d2a59ae..ba62c808f 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/nav.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/nav.adoc @@ -46,7 +46,7 @@ *** xref:api/image/openai-image.adoc[OpenAI] *** xref:api/image/stabilityai-image.adoc[Stability] *** xref:api/image/zhipuai-image.adoc[ZhiPuAI] -** xref:api/audio[Audio API] +** xref:api/audio[Audio Model API] *** xref:api/audio/transcriptions.adoc[] **** xref:api/audio/transcriptions/openai-transcriptions.adoc[OpenAI] *** xref:api/audio/speech.adoc[] diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/aimetadata.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/aimetadata.adoc index 6ab4d105f..6efd46764 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/aimetadata.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/aimetadata.adoc @@ -42,7 +42,7 @@ class MyService { Prompt prompt = createPrompt(request); - ChatResponse response = chatClient.call(prompt); + ChatResponse response = chatModel.call(prompt); // Process the chat response diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/speech.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/speech.adoc index a581a57da..adabcd80c 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/speech.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/speech.adoc @@ -2,4 +2,4 @@ = Text-To-Speech (TTS) API Spring AI provides support for OpenAI's Speech API. -When additional providers for Speech are implemented, a common `SpeechClient` and `StreamingSpeechClient` interface will be extracted. \ No newline at end of file +When additional providers for Speech are implemented, a common `SpeechModel` and `StreamingSpeechModel` interface will be extracted. \ No newline at end of file diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/speech/openai-speech.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/speech/openai-speech.adoc index 94a5f4c7c..6ba1841ac 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/speech/openai-speech.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/speech/openai-speech.adoc @@ -68,7 +68,7 @@ OpenAiAudioSpeechOptions speechOptions = OpenAiAudioSpeechOptions.builder() .build(); SpeechPrompt speechPrompt = new SpeechPrompt("Hello, this is a text-to-speech example.", speechOptions); -SpeechResponse response = openAiAudioSpeechClient.call(speechPrompt); +SpeechResponse response = openAiAudioSpeechModel.call(speechPrompt); ---- == Manual Configuration @@ -94,13 +94,13 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create an `OpenAiAudioSpeechClient`: +Next, create an `OpenAiAudioSpeechModel`: [source,java] ---- var openAiAudioApi = new OpenAiAudioApi(System.getenv("OPENAI_API_KEY")); -var openAiAudioSpeechClient = new OpenAiAudioSpeechClient(openAiAudioApi); +var openAiAudioSpeechModel = new OpenAiAudioSpeechModel(openAiAudioApi); var speechOptions = OpenAiAudioSpeechOptions.builder() .withResponseFormat(OpenAiAudioApi.SpeechRequest.AudioResponseFormat.MP3) @@ -109,7 +109,7 @@ var speechOptions = OpenAiAudioSpeechOptions.builder() .build(); var speechPrompt = new SpeechPrompt("Hello, this is a text-to-speech example.", speechOptions); -SpeechResponse response = openAiAudioSpeechClient.call(speechPrompt); +SpeechResponse response = openAiAudioSpeechModel.call(speechPrompt); // Accessing metadata (rate limit info) OpenAiAudioSpeechResponseMetadata metadata = response.getMetadata(); @@ -125,7 +125,7 @@ The Speech API provides support for real-time audio streaming using chunk transf ---- var openAiAudioApi = new OpenAiAudioApi(System.getenv("OPENAI_API_KEY")); -var openAiAudioSpeechClient = new OpenAiAudioSpeechClient(openAiAudioApi); +var openAiAudioSpeechModel = new OpenAiAudioSpeechModel(openAiAudioApi); OpenAiAudioSpeechOptions speechOptions = OpenAiAudioSpeechOptions.builder() .withVoice(OpenAiAudioApi.SpeechRequest.Voice.ALLOY) @@ -136,9 +136,9 @@ OpenAiAudioSpeechOptions speechOptions = OpenAiAudioSpeechOptions.builder() SpeechPrompt speechPrompt = new SpeechPrompt("Today is a wonderful day to build something people love!", speechOptions); -Flux responseStream = openAiAudioSpeechClient.stream(speechPrompt); +Flux responseStream = openAiAudioSpeechModel.stream(speechPrompt); ---- == Example Code -* The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechClientIT.java[OpenAiSpeechClientIT.java] test provides some general examples of how to use the library. +* The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechModelIT.java[OpenAiSpeechModelIT.java] test provides some general examples of how to use the library. diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/transcriptions.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/transcriptions.adoc index 703f19908..08d41a9a5 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/transcriptions.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/transcriptions.adoc @@ -2,4 +2,4 @@ = Transcription API Spring AI provides support for OpenAI's Transcription API. -When additional providers for Transcription are implemented, a common `AudioTranscriptionClient` interface will be extracted. \ No newline at end of file +When additional providers for Transcription are implemented, a common `AudioTranscriptionModel` interface will be extracted. \ No newline at end of file diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/transcriptions/openai-transcriptions.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/transcriptions/openai-transcriptions.adoc index 5592aa266..5f352f4f9 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/transcriptions/openai-transcriptions.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/audio/transcriptions/openai-transcriptions.adoc @@ -37,7 +37,7 @@ TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Man === Transcription Properties -The prefix `spring.ai.openai.audio.transcription` is used as the property prefix that lets you configure the retry mechanism for the OpenAI Image client. +The prefix `spring.ai.openai.audio.transcription` is used as the property prefix that lets you configure the retry mechanism for the OpenAI image model. [cols="3,5,2"] |==== @@ -69,7 +69,7 @@ OpenAiAudioTranscriptionOptions transcriptionOptions = OpenAiAudioTranscriptionO .withResponseFormat(responseFormat) .build(); AudioTranscriptionPrompt transcriptionRequest = new AudioTranscriptionPrompt(audioFile, transcriptionOptions); -AudioTranscriptionResponse response = openAiTranscriptionClient.call(transcriptionRequest); +AudioTranscriptionResponse response = openAiTranscriptionModel.call(transcriptionRequest); ---- == Manual Configuration @@ -95,13 +95,13 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create a `OpenAiAudioTranscriptionClient` +Next, create a `OpenAiAudioTranscriptionModel` [source,java] ---- var openAiAudioApi = new OpenAiAudioApi(System.getenv("OPENAI_API_KEY")); -var openAiAudioTranscriptionClient = new OpenAiAudioTranscriptionClient(openAiAudioApi); +var openAiAudioTranscriptionModel = new OpenAiAudioTranscriptionModel(openAiAudioApi); var transcriptionOptions = OpenAiAudioTranscriptionOptions.builder() .withResponseFormat(TranscriptResponseFormat.TEXT) @@ -111,8 +111,8 @@ var transcriptionOptions = OpenAiAudioTranscriptionOptions.builder() var audioFile = new FileSystemResource("/path/to/your/resource/speech/jfk.flac"); AudioTranscriptionPrompt transcriptionRequest = new AudioTranscriptionPrompt(audioFile, transcriptionOptions); -AudioTranscriptionResponse response = openAiTranscriptionClient.call(transcriptionRequest); +AudioTranscriptionResponse response = openAiTranscriptionModel.call(transcriptionRequest); ---- == Example Code -* The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionClientIT.java[OpenAiTranscriptionClientIT.java] test provides some general examples how to use the library. \ No newline at end of file +* The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionModelIT.java[OpenAiTranscriptionModelIT.java] test provides some general examples how to use the library. \ No newline at end of file diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/bedrock.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/bedrock.adoc index 3bdb84e10..f8b2b2062 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/bedrock.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/bedrock.adoc @@ -2,7 +2,7 @@ link:https://docs.aws.amazon.com/bedrock/latest/userguide/what-is-bedrock.html[Amazon Bedrock] is a managed service that provides foundation models from various AI providers, available through a unified API. -Spring AI supports https://docs.aws.amazon.com/bedrock/latest/userguide/model-ids-arns.html[all the Chat and Embedding AI models] available through Amazon Bedrock by implementing the Spring interfaces `ChatClient`, `StreamingChatClient`, and `EmbeddingClient`. +Spring AI supports https://docs.aws.amazon.com/bedrock/latest/userguide/model-ids-arns.html[all the Chat and Embedding AI models] available through Amazon Bedrock by implementing the Spring interfaces `ChatModel`, `StreamingChatModel`, and `EmbeddingModel`. Additionally, Spring AI provides Spring Auto-Configurations and Boot Starters for all clients, making it easy to bootstrap and configure for the Bedrock models. @@ -94,7 +94,7 @@ Here are the supported `` and `` combinations: | titan | Yes | Yes | Yes (however, no batch support) |==== -For example, to enable the Bedrock Llama Chat client, you need to set `spring.ai.bedrock.llama.chat.enabled=true`. +For example, to enable the Bedrock Llama chat model, you need to set `spring.ai.bedrock.llama.chat.enabled=true`. Next, you can use the `spring.ai.bedrock...*` properties to configure each model as provided. diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/anthropic-chat.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/anthropic-chat.adoc index def68f7c2..5858c9cef 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/anthropic-chat.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/anthropic-chat.adoc @@ -56,7 +56,7 @@ TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Man ==== Retry Properties -The prefix `spring.ai.retry` is used as the property prefix that lets you configure the retry mechanism for the Anthropic Chat client. +The prefix `spring.ai.retry` is used as the property prefix that lets you configure the retry mechanism for the Anthropic chat model. [cols="3,5,1"] |==== @@ -88,13 +88,13 @@ The prefix `spring.ai.anthropic` is used as the property prefix that lets you co ==== Configuration Properties -The prefix `spring.ai.anthropic.chat` is the property prefix that lets you configure the chat client implementation for Anthropic. +The prefix `spring.ai.anthropic.chat` is the property prefix that lets you configure the chat model implementation for Anthropic. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.anthropic.chat.enabled | Enable Anthropic chat client. | true +| spring.ai.anthropic.chat.enabled | Enable Anthropic chat model. | true | spring.ai.anthropic.chat.options.model | This is the Anthropic Chat model to use. Supports `claude-3-opus-20240229`, `claude-3-sonnet-20240229`, `claude-3-haiku-20240307` and the legacy `claude-2.1`, `claude-2.0` and `claude-instant-1.2` models. | `claude-3-opus-20240229` | spring.ai.anthropic.chat.options.temperature | The sampling temperature to use that controls the apparent creativity of generated completions. Higher values will make output more random while lower values will make results more focused and deterministic. It is not recommended to modify temperature and top_p for the same completions request as the interaction of these two settings is difficult to predict. | 0.8 | spring.ai.anthropic.chat.options.max-tokens | The maximum number of tokens to generate in the chat completion. The total length of input tokens and generated tokens is limited by the model's context length. | 500 @@ -102,7 +102,7 @@ The prefix `spring.ai.anthropic.chat` is the property prefix that lets you confi | spring.ai.anthropic.chat.options.top-p | Use nucleus sampling. In nucleus sampling, we compute the cumulative distribution over all the options for each subsequent token in decreasing probability order and cut it off once it reaches a particular probability specified by top_p. You should either alter temperature or top_p, but not both. Recommended for advanced use cases only. You usually only need to use temperature. | - | spring.ai.anthropic.chat.options.top-k | Only sample from the top K options for each subsequent token. Used to remove "long tail" low probability responses. Learn more technical details here. Recommended for advanced use cases only. You usually only need to use temperature. | - | spring.ai.mistralai.chat.options.functions | List of functions, identified by their names, to enable for function calling in a single prompt requests. Functions with those names must exist in the functionCallbacks registry. | - -| spring.ai.mistralai.chat.options.functionCallbacks | MistralAI Tool Function Callbacks to register with the ChatClient. | - +| spring.ai.mistralai.chat.options.functionCallbacks | MistralAI Tool Function Callbacks to register with the ChatModel. | - |==== TIP: All properties prefixed with `spring.ai.anthropic.chat.options` can be overridden at runtime by adding a request specific <> to the `Prompt` call. @@ -111,14 +111,14 @@ TIP: All properties prefixed with `spring.ai.anthropic.chat.options` can be over The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatOptions.java[AnthropicChatOptions.java] provides model configurations, such as the model to use, the temperature, the max token count, etc. -On start-up, the default options can be configured with the `AnthropicChatClient(api, options)` constructor or the `spring.ai.anthropic.chat.options.*` properties. +On start-up, the default options can be configured with the `AnthropicChatModel(api, options)` constructor or the `spring.ai.anthropic.chat.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default model and temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", AnthropicChatOptions.builder() @@ -132,7 +132,7 @@ TIP: In addition to the model specific https://github.com/spring-projects/spring == Function Calling -You can register custom Java functions with the `AnthropicChatClient` and have the Anthropic Claude model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. +You can register custom Java functions with the `AnthropicChatModel` and have the Anthropic Claude model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. This is a powerful technique to connect the LLM capabilities with external tools and APIs. Read more about xref:api/chat/functions/anthropic-chat-functions.adoc[Anthropic Function Calling]. @@ -146,7 +146,7 @@ Check the link:https://docs.anthropic.com/claude/docs/vision[Vision guide] for m Spring AI's `Message` interface supports multimodal AI models by introducing the Media type. This type contains data and information about media attachments in messages, using Spring's `org.springframework.util.MimeType` and a `java.lang.Object` for the raw media data. -Below is a simple code example extracted from https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatClientIT.java[AnthropicChatClientIT.java], demonstrating the combination of user text with an image. +Below is a simple code example extracted from https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatModelIT.java[AnthropicChatModelIT.java], demonstrating the combination of user text with an image. [source,java] ---- @@ -155,7 +155,7 @@ byte[] imageData = new ClassPathResource("/multimodal.test.png").getContentAsByt var userMessage = new UserMessage("Explain what do you see on this picture?", List.of(new Media(MimeTypeUtils.IMAGE_PNG, imageData))); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage))); +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage))); logger.info(response.getResult().getOutput().getContent()); ---- @@ -182,7 +182,7 @@ The composition and lighting give the image a clean, minimalist aesthetic that h https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-anthropic-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Anthropic Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Anthropic chat model: [source,application.properties] ---- @@ -194,37 +194,37 @@ spring.ai.anthropic.chat.options.max-tokens=450 TIP: replace the `api-key` with your Anthropic credentials. -This will create a `AnthropicChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `AnthropicChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class ChatController { - private final AnthropicChatClient chatClient; + private final AnthropicChatModel chatModel; @Autowired - public ChatController(AnthropicChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(AnthropicChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } @GetMapping("/ai/generateStream") public Flux generateStream(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { Prompt prompt = new Prompt(new UserMessage(message)); - return chatClient.stream(prompt); + return chatModel.stream(prompt); } } ---- == Manual Configuration -The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatClient.java[AnthropicChatClient] implements the `ChatClient` and `StreamingChatClient` and uses the <> to connect to the Anthropic service. +The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java[AnthropicChatModel] implements the `ChatModel` and `StreamingChatModel` and uses the <> to connect to the Anthropic service. Add the `spring-ai-anthropic` dependency to your project's Maven `pom.xml` file: @@ -247,24 +247,24 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create a `AnthropicChatClient` and use it for text generations: +Next, create a `AnthropicChatModel` and use it for text generations: [source,java] ---- var anthropicApi = new AnthropicApi(System.getenv("ANTHROPIC_API_KEY")); -var chatClient = new AnthropicChatClient(anthropicApi, +var chatModel = new AnthropicChatModel(anthropicApi, AnthropicChatOptions.builder() .withModel("claude-3-opus-20240229") .withTemperature(0.4) .withMaxTokens(200) .build()); -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); // Or with streaming responses -Flux response = chatClient.stream( +Flux response = chatModel.stream( new Prompt("Generate the names of 5 famous pirates.")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/azure-openai-chat.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/azure-openai-chat.adoc index 759aca404..88a7519ed 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/azure-openai-chat.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/azure-openai-chat.adoc @@ -88,13 +88,13 @@ The prefix `spring.ai.azure.openai` is the property prefix to configure the conn | spring.ai.azure.openai.endpoint | The endpoint from the Azure AI OpenAI `Keys and Endpoint` section under `Resource Management` | - |==== -The prefix `spring.ai.azure.openai.chat` is the property prefix that configures the `ChatClient` implementation for Azure OpenAI. +The prefix `spring.ai.azure.openai.chat` is the property prefix that configures the `ChatModel` implementation for Azure OpenAI. [cols="3,5,3"] |==== | Property | Description | Default -| spring.ai.azure.openai.chat.enabled | Enable Azure OpenAI chat client. | true +| spring.ai.azure.openai.chat.enabled | Enable Azure OpenAI chat model. | true | spring.ai.azure.openai.chat.options.deployment-name | * In use with Azure, this refers to the "Deployment Name" of your model, which you can find at https://oai.azure.com/portal. It's important to note that within an Azure OpenAI deployment, the "Deployment Name" is distinct from the model itself. The confusion around these terms stems from the intention to make the Azure OpenAI client library compatible with the original OpenAI endpoint. The deployment structures offered by Azure OpenAI and Sam Altman's OpenAI differ significantly. Deployments model name to provide as part of this completions request. | gpt-35-turbo @@ -115,14 +115,14 @@ TIP: All properties prefixed with `spring.ai.azure.openai.chat.options` can be o The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatOptions.java[AzureOpenAiChatOptions.java] provides model configurations, such as the model to use, the temperature, the frequency penalty, etc. -On start-up, the default options can be configured with the `AzureOpenAiChatClient(api, options)` constructor or the `spring.ai.azure.openai.chat.options.*` properties. +On start-up, the default options can be configured with the `AzureOpenAiChatModel(api, options)` constructor or the `spring.ai.azure.openai.chat.options.*` properties. At runtime you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default model and temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", AzureOpenAiChatOptions.builder() @@ -137,7 +137,7 @@ TIP: In addition to the model specific link:https://github.com/spring-projects/s == Function Calling -You can register custom Java functions with the AzureOpenAiChatClient and have the model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. +You can register custom Java functions with the AzureOpenAiChatModel and have the model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. This is a powerful technique to connect the LLM capabilities with external tools and APIs. Read more about xref:api/chat/functions/azure-open-ai-chat-functions.adoc[Azure OpenAI Function Calling]. @@ -145,7 +145,7 @@ Read more about xref:api/chat/functions/azure-open-ai-chat-functions.adoc[Azure https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-azure-openai-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the OpenAi Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the OpenAi chat model: [source,application.properties] ---- @@ -157,8 +157,8 @@ spring.ai.azure.openai.chat.options.temperature=0.7 TIP: replace the `api-key` and `endpoint` with your Azure OpenAI credentials. -This will create a `AzureOpenAiChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `AzureOpenAiChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] @@ -166,29 +166,29 @@ Here is an example of a simple `@Controller` class that uses the chat client for @RestController public class ChatController { - private final AzureOpenAiChatClient chatClient; + private final AzureOpenAiChatModel chatModel; @Autowired - public ChatController(AzureOpenAiChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(AzureOpenAiChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } @GetMapping("/ai/generateStream") public Flux generateStream(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { Prompt prompt = new Prompt(new UserMessage(message)); - return chatClient.stream(prompt); + return chatModel.stream(prompt); } } ---- == Manual Configuration -The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatClient.java[AzureOpenAiChatClient] implements the `ChatClient` and `StreamingChatClient` and uses the link:https://learn.microsoft.com/en-us/java/api/overview/azure/ai-openai-readme?view=azure-java-preview[Azure OpenAI Java Client]. +The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatModel.java[AzureOpenAiChatModel] implements the `ChatModel` and `StreamingChatModel` and uses the link:https://learn.microsoft.com/en-us/java/api/overview/azure/ai-openai-readme?view=azure-java-preview[Azure OpenAI Java Client]. To enable it, add the `spring-ai-azure-openai` dependency to your project's Maven `pom.xml` file: [source, xml] @@ -210,9 +210,9 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -TIP: The `spring-ai-azure-openai` dependency also provide the access to the `AzureOpenAiChatClient`. For more information about the `AzureOpenAiChatClient` refer to the link:../chat/azure-openai-chat.html[Azure OpenAI Chat] section. +TIP: The `spring-ai-azure-openai` dependency also provide the access to the `AzureOpenAiChatModel`. For more information about the `AzureOpenAiChatModel` refer to the link:../chat/azure-openai-chat.html[Azure OpenAI Chat] section. -Next, create an `AzureOpenAiChatClient` instance and use it to generate text responses: +Next, create an `AzureOpenAiChatModel` instance and use it to generate text responses: [source,java] ---- @@ -227,13 +227,13 @@ var openAIChatOptions = AzureOpenAiChatOptions.builder() .withMaxTokens(200) .build(); -var chatClient = new AzureOpenAiChatClient(openAIClient, openAIChatOptions); +var chatModel = new AzureOpenAiChatModel(openAIClient, openAIChatOptions); -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); // Or with streaming responses -Flux response = chatClient.stream( +Flux response = chatModel.stream( new Prompt("Generate the names of 5 famous pirates.")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-anthropic.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-anthropic.adoc index d6fc03d60..dde93f185 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-anthropic.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-anthropic.adoc @@ -74,13 +74,13 @@ The prefix `spring.ai.bedrock.aws` is the property prefix to configure the conne | spring.ai.bedrock.aws.secret-key | AWS secret key. | - |==== -The prefix `spring.ai.bedrock.anthropic.chat` is the property prefix that configures the chat client implementation for Claude. +The prefix `spring.ai.bedrock.anthropic.chat` is the property prefix that configures the chat model implementation for Claude. [cols="2,5,1"] |==== | Property | Description | Default -| spring.ai.bedrock.anthropic.chat.enable | Enable Bedrock Anthropic chat client. Disabled by default | false +| spring.ai.bedrock.anthropic.chat.enable | Enable Bedrock Anthropic chat model. Disabled by default | false | spring.ai.bedrock.anthropic.chat.model | The model id to use. See the https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/api/AnthropicChatBedrockApi.java[AnthropicChatModel] for the supported models. | anthropic.claude-v2 | spring.ai.bedrock.anthropic.chat.options.temperature | Controls the randomness of the output. Values can range over [0.0,1.0] | 0.8 | spring.ai.bedrock.anthropic.chat.options.topP | The maximum cumulative probability of tokens to consider when sampling. | AWS Bedrock default @@ -100,14 +100,14 @@ TIP: All properties prefixed with `spring.ai.bedrock.anthropic.chat.options` can The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/AnthropicChatOptions.java[AnthropicChatOptions.java] provides model configurations, such as temperature, topK, topP, etc. -On start-up, the default options can be configured with the `BedrockAnthropicChatClient(api, options)` constructor or the `spring.ai.bedrock.anthropic.chat.options.*` properties. +On start-up, the default options can be configured with the `BedrockAnthropicChatModel(api, options)` constructor or the `spring.ai.bedrock.anthropic.chat.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", AnthropicChatOptions.builder() @@ -122,7 +122,7 @@ TIP: In addition to the model specific https://github.com/spring-projects/spring https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-bedrock-ai-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Anthropic Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Anthropic chat model: [source] ---- @@ -138,37 +138,37 @@ spring.ai.bedrock.anthropic.chat.options.top-k=15 TIP: replace the `regions`, `access-key` and `secret-key` with your AWS credentials. -This will create a `BedrockAnthropicChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `BedrockAnthropicChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class ChatController { - private final BedrockAnthropicChatClient chatClient; + private final BedrockAnthropicChatModel chatModel; @Autowired - public ChatController(BedrockAnthropicChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(BedrockAnthropicChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } @GetMapping("/ai/generateStream") public Flux generateStream(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { Prompt prompt = new Prompt(new UserMessage(message)); - return chatClient.stream(prompt); + return chatModel.stream(prompt); } } ---- == Manual Configuration -The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatClient.java[BedrockAnthropicChatClient] implements the `ChatClient` and `StreamingChatClient` and uses the <> to connect to the Bedrock Anthropic service. +The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatModel.java[BedrockAnthropicChatModel] implements the `ChatModel` and `StreamingChatModel` and uses the <> to connect to the Bedrock Anthropic service. Add the `spring-ai-bedrock` dependency to your project's Maven `pom.xml` file: @@ -191,7 +191,7 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatClient.java[BedrockAnthropicChatClient] and use it for text generations: +Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatModel.java[BedrockAnthropicChatModel] and use it for text generations: [source,java] ---- @@ -202,7 +202,7 @@ AnthropicChatBedrockApi anthropicApi = new AnthropicChatBedrockApi( new ObjectMapper(), Duration.ofMillis(1000L)); -BedrockAnthropicChatClient chatClient = new BedrockAnthropicChatClient(anthropicApi, +BedrockAnthropicChatModel chatModel = new BedrockAnthropicChatModel(anthropicApi, AnthropicChatOptions.builder() .withTemperature(0.6f) .withTopK(10) @@ -211,11 +211,11 @@ BedrockAnthropicChatClient chatClient = new BedrockAnthropicChatClient(anthropic .withAnthropicVersion(AnthropicChatBedrockApi.DEFAULT_ANTHROPIC_VERSION) .build()); -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); // Or with streaming responses -Flux response = chatClient.stream( +Flux response = chatModel.stream( new Prompt("Generate the names of 5 famous pirates.")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-anthropic3.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-anthropic3.adoc index 3b4cd66af..c03b1d7da 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-anthropic3.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-anthropic3.adoc @@ -71,13 +71,13 @@ The prefix `spring.ai.bedrock.aws` is the property prefix to configure the conne | spring.ai.bedrock.aws.secret-key | AWS secret key. | - |==== -The prefix `spring.ai.bedrock.anthropic3.chat` is the property prefix that configures the chat client implementation for Claude. +The prefix `spring.ai.bedrock.anthropic3.chat` is the property prefix that configures the chat model implementation for Claude. [cols="2,5,1"] |==== | Property | Description | Default -| spring.ai.bedrock.anthropic3.chat.enable | Enable Bedrock Anthropic chat client. Disabled by default | false +| spring.ai.bedrock.anthropic3.chat.enable | Enable Bedrock Anthropic chat model. Disabled by default | false | spring.ai.bedrock.anthropic3.chat.model | The model id to use. Supports the `anthropic.claude-3-sonnet-20240229-v1:0`,`anthropic.claude-3-haiku-20240307-v1:0` and the legacy `anthropic.claude-v2`, `anthropic.claude-v2:1` and `anthropic.claude-instant-v1` models for both synchronous and streaming responses. | `anthropic.claude-3-sonnet-20240229-v1:0` | spring.ai.bedrock.anthropic3.chat.options.temperature | Controls the randomness of the output. Values can range over [0.0,1.0] | 0.8 | spring.ai.bedrock.anthropic3.chat.options.top-p | The maximum cumulative probability of tokens to consider when sampling. | AWS Bedrock default @@ -97,14 +97,14 @@ TIP: All properties prefixed with `spring.ai.bedrock.anthropic3.chat.options` ca The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/Anthropic3ChatOptions.java[Anthropic3ChatOptions.java] provides model configurations, such as temperature, topK, topP, etc. -On start-up, the default options can be configured with the `BedrockAnthropicChatClient(api, options)` constructor or the `spring.ai.bedrock.anthropic3.chat.options.*` properties. +On start-up, the default options can be configured with the `BedrockAnthropicChatModel(api, options)` constructor or the `spring.ai.bedrock.anthropic3.chat.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", Anthropic3ChatOptions.builder() @@ -126,7 +126,7 @@ Check the link:https://docs.anthropic.com/claude/docs/vision[Vision guide] for m Spring AI's `Message` interface supports multimodal AI models by introducing the Media type. This type contains data and information about media attachments in messages, using Spring's `org.springframework.util.MimeType` and a `java.lang.Object` for the raw media data. -Below is a simple code example extracted from https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic3/Anthropic3ChatClientIT.java[Anthropic3ChatClientIT.java], demonstrating the combination of user text with an image. +Below is a simple code example extracted from https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic3/Anthropic3ChatModelIT.java[Anthropic3ChatModelIT.java], demonstrating the combination of user text with an image. [source,java] ---- @@ -135,7 +135,7 @@ Below is a simple code example extracted from https://github.com/spring-projects var userMessage = new UserMessage("Explain what do you see o this picture?", List.of(new Media(MimeTypeUtils.IMAGE_PNG, imageData))); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage))); + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage))); assertThat(response.getResult().getOutput().getContent()).contains("bananas", "apple", "basket"); ---- @@ -163,7 +163,7 @@ The composition and lighting give the image a clean, minimalist aesthetic that h https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-bedrock-ai-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Anthropic Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Anthropic chat model: [source] ---- @@ -179,37 +179,37 @@ spring.ai.bedrock.anthropic3.chat.options.top-k=15 TIP: replace the `regions`, `access-key` and `secret-key` with your AWS credentials. -This will create a `BedrockAnthropicChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `BedrockAnthropicChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class ChatController { - private final BedrockAnthropic3ChatClient chatClient; + private final BedrockAnthropic3ChatModel chatModel; @Autowired - public ChatController(BedrockAnthropic3ChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(BedrockAnthropic3ChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } @GetMapping("/ai/generateStream") public Flux generateStream(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { Prompt prompt = new Prompt(new UserMessage(message)); - return chatClient.stream(prompt); + return chatModel.stream(prompt); } } ---- == Manual Configuration -The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatClient.java[BedrockAnthropic3ChatClient] implements the `ChatClient` and `StreamingChatClient` and uses the <> to connect to the Bedrock Anthropic service. +The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatModel.java[BedrockAnthropic3ChatModel] implements the `ChatModel` and `StreamingChatModel` and uses the <> to connect to the Bedrock Anthropic service. Add the `spring-ai-bedrock` dependency to your project's Maven `pom.xml` file: @@ -232,7 +232,7 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatClient.java[BedrockAnthropic3ChatClient] and use it for text generations: +Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatModel.java[BedrockAnthropic3ChatModel] and use it for text generations: [source,java] ---- @@ -243,7 +243,7 @@ Anthropic3ChatBedrockApi anthropicApi = new Anthropic3ChatBedrockApi( new ObjectMapper(), Duration.ofMillis(1000L)); -BedrockAnthropic3ChatClient chatClient = new BedrockAnthropic3ChatClient(anthropicApi, +BedrockAnthropic3ChatModel chatModel = new BedrockAnthropic3ChatModel(anthropicApi, AnthropicChatOptions.builder() .withTemperature(0.6f) .withTopK(10) @@ -252,11 +252,11 @@ BedrockAnthropic3ChatClient chatClient = new BedrockAnthropic3ChatClient(anthrop .withAnthropicVersion(AnthropicChatBedrockApi.DEFAULT_ANTHROPIC_VERSION) .build()); -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); // Or with streaming responses -Flux response = chatClient.stream( +Flux response = chatModel.stream( new Prompt("Generate the names of 5 famous pirates.")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-cohere.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-cohere.adoc index 8133f3bc2..c4345a46c 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-cohere.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-cohere.adoc @@ -1,6 +1,6 @@ = Cohere Chat -Provides Bedrock Cohere Chat client. +Provides Bedrock Cohere chat model. Integrate generative AI capabilities into essential apps and workflows that improve business outcomes. The https://aws.amazon.com/bedrock/cohere-command-embed/[AWS Bedrock Cohere Model Page] and https://docs.aws.amazon.com/bedrock/latest/userguide/what-is-bedrock.html[Amazon Bedrock User Guide] contains detailed information on how to use the AWS hosted model. @@ -64,7 +64,7 @@ The prefix `spring.ai.bedrock.aws` is the property prefix to configure the conne | spring.ai.bedrock.aws.secret-key | AWS secret key. | - |==== -The prefix `spring.ai.bedrock.cohere.chat` is the property prefix that configures the chat client implementation for Cohere. +The prefix `spring.ai.bedrock.cohere.chat` is the property prefix that configures the chat model implementation for Cohere. [cols="2,5,1"] |==== @@ -93,14 +93,14 @@ TIP: All properties prefixed with `spring.ai.bedrock.cohere.chat.options` can be The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatOptions.java[BedrockCohereChatOptions.java] provides model configurations, such as temperature, topK, topP, etc. -On start-up, the default options can be configured with the `BedrockCohereChatClient(api, options)` constructor or the `spring.ai.bedrock.cohere.chat.options.*` properties. +On start-up, the default options can be configured with the `BedrockCohereChatModel(api, options)` constructor or the `spring.ai.bedrock.cohere.chat.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", BedrockCohereChatOptions.builder() @@ -115,7 +115,7 @@ TIP: In addition to the model specific https://github.com/spring-projects/spring https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-bedrock-ai-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Cohere Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Cohere chat model: [source] ---- @@ -130,37 +130,37 @@ spring.ai.bedrock.cohere.chat.options.temperature=0.8 TIP: replace the `regions`, `access-key` and `secret-key` with your AWS credentials. -This will create a `BedrockCohereChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `BedrockCohereChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class ChatController { - private final BedrockCohereChatClient chatClient; + private final BedrockCohereChatModel chatModel; @Autowired - public ChatController(BedrockCohereChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(BedrockCohereChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } @GetMapping("/ai/generateStream") public Flux generateStream(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { Prompt prompt = new Prompt(new UserMessage(message)); - return chatClient.stream(prompt); + return chatModel.stream(prompt); } } ---- == Manual Configuration -The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatClient.java[BedrockCohereChatClient] implements the `ChatClient` and `StreamingChatClient` and uses the <> to connect to the Bedrock Cohere service. +The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatModel.java[BedrockCohereChatModel] implements the `ChatModel` and `StreamingChatModel` and uses the <> to connect to the Bedrock Cohere service. Add the `spring-ai-bedrock` dependency to your project's Maven `pom.xml` file: @@ -183,7 +183,7 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatClient.java[BedrockCohereChatClient] and use it for text generations: +Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatModel.java[BedrockCohereChatModel] and use it for text generations: [source,java] ---- @@ -193,7 +193,7 @@ CohereChatBedrockApi api = new CohereChatBedrockApi(CohereChatModel.COHERE_COMMA new ObjectMapper(), Duration.ofMillis(1000L)); -BedrockCohereChatClient chatClient = new BedrockCohereChatClient(api, +BedrockCohereChatModel chatModel = new BedrockCohereChatModel(api, BedrockCohereChatOptions.builder() .withTemperature(0.6f) .withTopK(10) @@ -201,11 +201,11 @@ BedrockCohereChatClient chatClient = new BedrockCohereChatClient(api, .withMaxTokens(678) .build() -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); // Or with streaming responses -Flux response = chatClient.stream( +Flux response = chatModel.stream( new Prompt("Generate the names of 5 famous pirates.")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-jurassic2.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-jurassic2.adoc index cf7a083de..d1d8956ce 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-jurassic2.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-jurassic2.adoc @@ -64,7 +64,7 @@ The prefix `spring.ai.bedrock.aws` is the property prefix to configure the conne |==== -The prefix `spring.ai.bedrock.jurassic2.chat` is the property prefix that configures the chat client implementation for Jurassic-2. +The prefix `spring.ai.bedrock.jurassic2.chat` is the property prefix that configures the chat model implementation for Jurassic-2. [cols="2,5,1"] |==== @@ -86,14 +86,14 @@ TIP: All properties prefixed with `spring.ai.bedrock.jurassic2.chat.options` can The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatOptions.java[BedrockAi21Jurassic2ChatOptions.java] provides model configurations, such as temperature, topP, maxTokens, etc. -On start-up, the default options can be configured with the `BedrockAi21Jurassic2ChatClient(api, options)` constructor or the `spring.ai.bedrock.jurassic2.chat.options.*` properties. +On start-up, the default options can be configured with the `BedrockAi21Jurassic2ChatModel(api, options)` constructor or the `spring.ai.bedrock.jurassic2.chat.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", BedrockAi21Jurassic2ChatOptions.builder() @@ -108,7 +108,7 @@ TIP: In addition to the model specific https://github.com/spring-projects/spring https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-bedrock-ai-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Jurassic-2 Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Jurassic-2 chat model: [source] ---- @@ -123,24 +123,24 @@ spring.ai.bedrock.jurassic2.chat.options.temperature=0.8 TIP: replace the `regions`, `access-key` and `secret-key` with your AWS credentials. -This will create a `BedrockAi21Jurassic2ChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `BedrockAi21Jurassic2ChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class ChatController { - private final BedrockAi21Jurassic2ChatClient chatClient; + private final BedrockAi21Jurassic2ChatModel chatModel; @Autowired - public ChatController(BedrockAi21Jurassic2ChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(BedrockAi21Jurassic2ChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } } @@ -148,7 +148,7 @@ public class ChatController { == Manual Configuration -The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatClient.java[BedrockAi21Jurassic2ChatClient] implements the `ChatClient` uses the <> to connect to the Bedrock Jurassic-2 service. +The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatModel.java[BedrockAi21Jurassic2ChatModel] implements the `ChatModel` uses the <> to connect to the Bedrock Jurassic-2 service. Add the `spring-ai-bedrock` dependency to your project's Maven `pom.xml` file: @@ -171,7 +171,7 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatClient.java[BedrockAi21Jurassic2ChatClient] and use it for text generations: +Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatModel.java[BedrockAi21Jurassic2ChatModel] and use it for text generations: [source,java] ---- @@ -181,13 +181,13 @@ Ai21Jurassic2ChatBedrockApi api = new Ai21Jurassic2ChatBedrockApi(Ai21Jurassic2C new ObjectMapper(), Duration.ofMillis(1000L)); -BedrockAi21Jurassic2ChatClient chatClient = new BedrockAi21Jurassic2ChatClient(api, +BedrockAi21Jurassic2ChatModel chatModel = new BedrockAi21Jurassic2ChatModel(api, BedrockAi21Jurassic2ChatOptions.builder() .withTemperature(0.5f) .withMaxTokens(100) .withTopP(0.9f).build()); -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-llama.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-llama.adoc index 435b71e47..d8a0c6347 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-llama.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-llama.adoc @@ -69,7 +69,7 @@ The prefix `spring.ai.bedrock.aws` is the property prefix to configure the conne |==== -The prefix `spring.ai.bedrock.llama.chat` is the property prefix that configures the chat client implementation for Llama. +The prefix `spring.ai.bedrock.llama.chat` is the property prefix that configures the chat model implementation for Llama. [cols="2,5,1"] |==== @@ -91,14 +91,14 @@ TIP: All properties prefixed with `spring.ai.bedrock.llama.chat.options` can be The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatOptions.java[BedrockLlChatOptions.java] provides model configurations, such as temperature, topK, topP, etc. -On start-up, the default options can be configured with the `BedrockLlamaChatClient(api, options)` constructor or the `spring.ai.bedrock.llama.chat.options.*` properties. +On start-up, the default options can be configured with the `BedrockLlamaChatModel(api, options)` constructor or the `spring.ai.bedrock.llama.chat.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", BedrockLlamaChatOptions.builder() @@ -113,7 +113,7 @@ TIP: In addition to the model specific https://github.com/spring-projects/spring https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-bedrock-ai-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Anthropic Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Anthropic chat model: [source] ---- @@ -128,37 +128,37 @@ spring.ai.bedrock.llama.chat.options.temperature=0.8 TIP: replace the `regions`, `access-key` and `secret-key` with your AWS credentials. -This will create a `BedrockLlamaChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `BedrockLlamaChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class ChatController { - private final BedrockLlamaChatClient chatClient; + private final BedrockLlamaChatModel chatModel; @Autowired - public ChatController(BedrockLlamaChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(BedrockLlamaChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } @GetMapping("/ai/generateStream") public Flux generateStream(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { Prompt prompt = new Prompt(new UserMessage(message)); - return chatClient.stream(prompt); + return chatModel.stream(prompt); } } ---- == Manual Configuration -The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatClient.java[BedrockLlamaChatClient] implements the `ChatClient` and `StreamingChatClient` and uses the <> to connect to the Bedrock Anthropic service. +The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatModel.java[BedrockLlamaChatModel] implements the `ChatModel` and `StreamingChatModel` and uses the <> to connect to the Bedrock Anthropic service. Add the `spring-ai-bedrock` dependency to your project's Maven `pom.xml` file: @@ -181,7 +181,7 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatClient.java[BedrockLlamaChatClient] and use it for text generations: +Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatModel.java[BedrockLlamaChatModel] and use it for text generations: [source,java] ---- @@ -191,17 +191,17 @@ LlamaChatBedrockApi api = new LlamaChatBedrockApi(LlamaChatModel.LLAMA2_70B_CHAT new ObjectMapper(), Duration.ofMillis(1000L)); -BedrockLlamaChatClient chatClient = new BedrockLlamaChatClient(api, +BedrockLlamaChatModel chatModel = new BedrockLlamaChatModel(api, BedrockLlamaChatOptions.builder() .withTemperature(0.5f) .withMaxGenLen(100) .withTopP(0.9f).build()); -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); // Or with streaming responses -Flux response = chatClient.stream( +Flux response = chatModel.stream( new Prompt("Generate the names of 5 famous pirates.")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-titan.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-titan.adoc index 38978fcfc..a1e57a357 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-titan.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/bedrock/bedrock-titan.adoc @@ -65,13 +65,13 @@ The prefix `spring.ai.bedrock.aws` is the property prefix to configure the conne | spring.ai.bedrock.aws.secret-key | AWS secret key. | - |==== -The prefix `spring.ai.bedrock.titan.chat` is the property prefix that configures the chat client implementation for Titan. +The prefix `spring.ai.bedrock.titan.chat` is the property prefix that configures the chat model implementation for Titan. [cols="3,4,1"] |==== | Property | Description | Default -| spring.ai.bedrock.titan.chat.enable | Enable Bedrock Titan chat client. Disabled by default | false +| spring.ai.bedrock.titan.chat.enable | Enable Bedrock Titan chat model. Disabled by default | false | spring.ai.bedrock.titan.chat.model | The model id to use. See the link:https://github.com/spring-projects/spring-ai/blob/4839a6175cd1ec89498b97d3efb6647022c3c7cb/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/api/TitanChatBedrockApi.java#L220[TitanChatBedrockApi#TitanChatModel] for the supported models. | amazon.titan-text-lite-v1 | spring.ai.bedrock.titan.chat.options.temperature | Controls the randomness of the output. Values can range over [0.0,1.0] | 0.7 | spring.ai.bedrock.titan.chat.options.topP | The maximum cumulative probability of tokens to consider when sampling. | AWS Bedrock default @@ -89,14 +89,14 @@ TIP: All properties prefixed with `spring.ai.bedrock.titan.chat.options` can be The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatOptions.java[BedrockTitanChatOptions.java] provides model configurations, such as temperature, topP, etc. -On start-up, the default options can be configured with the `BedrockTitanChatClient(api, options)` constructor or the `spring.ai.bedrock.titan.chat.options.*` properties. +On start-up, the default options can be configured with the `BedrockTitanChatModel(api, options)` constructor or the `spring.ai.bedrock.titan.chat.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", BedrockTitanChatOptions.builder() @@ -111,7 +111,7 @@ TIP: In addition to the model specific https://github.com/spring-projects/spring https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-bedrock-ai-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Titan Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Titan chat model: [source] ---- @@ -126,37 +126,37 @@ spring.ai.bedrock.titan.chat.options.temperature=0.8 TIP: replace the `regions`, `access-key` and `secret-key` with your AWS credentials. -This will create a `BedrockTitanChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `BedrockTitanChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class ChatController { - private final BedrockTitanChatClient chatClient; + private final BedrockTitanChatModel chatModel; @Autowired - public ChatController(BedrockTitanChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(BedrockTitanChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } @GetMapping("/ai/generateStream") public Flux generateStream(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { Prompt prompt = new Prompt(new UserMessage(message)); - return chatClient.stream(prompt); + return chatModel.stream(prompt); } } ---- == Manual Configuration -The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatClient.java[BedrockTitanChatClient] implements the `ChatClient` and `StreamingChatClient` and uses the <> to connect to the Bedrock Titanic service. +The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatModel.java[BedrockTitanChatModel] implements the `ChatModel` and `StreamingChatModel` and uses the <> to connect to the Bedrock Titanic service. Add the `spring-ai-bedrock` dependency to your project's Maven `pom.xml` file: @@ -179,7 +179,7 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatClient.java[BedrockTitanChatClient] and use it for text generations: +Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanChatModel.java[BedrockTitanChatModel] and use it for text generations: [source,java] ---- @@ -190,18 +190,18 @@ TitanChatBedrockApi titanApi = new TitanChatBedrockApi( new ObjectMapper(), Duration.ofMillis(1000L)); -BedrockTitanChatClient chatClient = new BedrockTitanChatClient(titanApi, +BedrockTitanChatModel chatModel = new BedrockTitanChatModel(titanApi, BedrockTitanChatOptions.builder() .withTemperature(0.6f) .withTopP(0.8f) .withMaxTokenCount(100) .build()); -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); // Or with streaming responses -Flux response = chatClient.stream( +Flux response = chatModel.stream( new Prompt("Generate the names of 5 famous pirates.")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/anthropic-chat-functions.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/anthropic-chat-functions.adoc index 3b4d9fcee..fea58dc1d 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/anthropic-chat-functions.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/anthropic-chat-functions.adoc @@ -1,6 +1,6 @@ = Anthropic Function Calling -You can register custom Java functions with the `AnthropicChatClient` and have the Anthropic models intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. +You can register custom Java functions with the `AnthropicChatModel` and have the Anthropic models intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. This allows you to connect the LLM capabilities with external tools and APIs. The `claude-3-opus`, `claude-3-sonnet` and `claude-3-haiku` link:https://docs.anthropic.com/claude/docs/tool-use#tool-use-best-practices-and-limitations[models are trained to detect when a function should be called] and to respond with JSON that adheres to the function signature. @@ -15,7 +15,7 @@ The `description` helps the model to understand when to call the function. As a developer, you need to implement a function that takes the function call arguments sent from the AI model, and respond with the result back to the model. Your function can in turn invoke other 3rd party services to provide the results. -Spring AI makes this as easy as defining a `@Bean` definition that returns a `java.util.Function` and supplying the bean name as an option when invoking the `ChatClient`. +Spring AI makes this as easy as defining a `@Bean` definition that returns a `java.util.Function` and supplying the bean name as an option when invoking the `ChatModel`. Under the hood, Spring wraps your POJO (the function) with the appropriate adapter code that enables interaction with the AI Model, saving you from writing tedious boilerplate code. The basis of the underlying infrastructure is the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallback.java[FunctionCallback.java] interface and the companion link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallbackWrapper.java[FunctionCallbackWrapper.java] utility class to simplify the implementation and registration of Java callback functions. @@ -62,7 +62,7 @@ public class MockWeatherService implements Function { === Registering Functions as Beans -With the link:../anthropic-chat.html#_auto_configuration[AnthropicChatClient Auto-Configuration] you have multiple ways to register custom functions as beans in the Spring context. +With the link:../anthropic-chat.html#_auto_configuration[AnthropicChatModel Auto-Configuration] you have multiple ways to register custom functions as beans in the Spring context. We start with describing the most POJO friendly options. @@ -70,7 +70,7 @@ We start with describing the most POJO friendly options. In this approach you define `@Beans` in your application context as you would any other Spring managed object. -Internally, Spring AI `ChatClient` will create an instance of a `FunctionCallbackWrapper` wrapper that adds the logic for it being invoked via the AI model. +Internally, Spring AI `ChatModel` will create an instance of a `FunctionCallbackWrapper` wrapper that adds the logic for it being invoked via the AI model. The name of the `@Bean` is passed as a `ChatOption`. @@ -136,7 +136,7 @@ static class Config { } ---- -It wraps the 3rd party `MockWeatherService` function and registers it as a `CurrentWeather` function with the `AnthropicChatClient`. +It wraps the 3rd party `MockWeatherService` function and registers it as a `CurrentWeather` function with the `AnthropicChatModel`. It also provides a description (2) and an optional response converter (3) to convert the response into a text as expected by the model. NOTE: By default, the response converter does a JSON serialization of the Response object. @@ -149,17 +149,17 @@ To let the model know and call your `CurrentWeather` function you need to enable [source,java] ---- -AnthropicChatClient chatClient = ... +AnthropicChatModel chatModel = ... UserMessage userMessage = new UserMessage("What's the weather like in Paris?"); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), AnthropicChatOptions.builder().withFunction("CurrentWeather").build())); // (1) Enable the function logger.info("Response: {}", response); ---- -// NOTE: You can can have multiple functions registered in your `ChatClient` but only those enabled in the prompt request will be considered for the function calling. +// NOTE: You can can have multiple functions registered in your `ChatModel` but only those enabled in the prompt request will be considered for the function calling. Above user question will trigger 3 calls to `CurrentWeather` function (one for each city) and produce the final response. @@ -169,7 +169,7 @@ In addition to the auto-configuration you can register callback functions, dynam [source,java] ---- -AnthropicChatClient chatClient = ... +AnthropicChatModel chatModel = ... UserMessage userMessage = new UserMessage("What's the weather like in Paris?"); @@ -180,12 +180,12 @@ var promptOptions = AnthropicChatOptions.builder() new MockWeatherService()))) // function code .build(); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), promptOptions)); +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), promptOptions)); ---- NOTE: The in-prompt registered functions are enabled by default for the duration of this request. This approach allows to dynamically chose different functions to be called based on the user input. -The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/tool/FunctionCallWithPromptFunctionIT.java[FunctionCallWithPromptFunctionIT.java] integration test provides a complete example of how to register a function with the `AnthropicChatClient` and use it in a prompt request. +The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/tool/FunctionCallWithPromptFunctionIT.java[FunctionCallWithPromptFunctionIT.java] integration test provides a complete example of how to register a function with the `AnthropicChatModel` and use it in a prompt request. diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/azure-open-ai-chat-functions.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/azure-open-ai-chat-functions.adoc index 755355a98..ca69472f2 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/azure-open-ai-chat-functions.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/azure-open-ai-chat-functions.adoc @@ -2,7 +2,7 @@ Function calling lets developers create a description of a function in their code, then pass that description to a language model in a request. The response from the model includes the name of a function that matches the description and the arguments to call it with. -You can register custom Java functions with the `AzureOpenAiChatClient` and have the model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. +You can register custom Java functions with the `AzureOpenAiChatModel` and have the model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. This allows you to connect the LLM capabilities with external tools and APIs. The Azure models are trained to detect when a function should be called and to respond with JSON that adheres to the function signature. @@ -16,7 +16,7 @@ In general, the custom functions need to provide a function `name`, `description As a developer, you need to implement a function that takes the function call arguments sent from the AI model, and respond with the result back to the model. Your function can in turn invoke other 3rd party services to provide the results. -Spring AI makes this as easy as defining a `@Bean` definition that returns a `java.util.Function` and supplying the bean name as an option when invoking the `ChatClient`. +Spring AI makes this as easy as defining a `@Bean` definition that returns a `java.util.Function` and supplying the bean name as an option when invoking the `ChatModel`. Under the hood, Spring wraps your POJO (the function) with the appropriate adapter code that enables interaction with the AI Model, saving you from writing tedious boilerplate code. The basis of the underlying infrastructure is the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallback.java[FunctionCallback.java] interface and the companion link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallbackWrapper.java[FunctionCallbackWrapper.java] utility class to simplify the implementation and registration of Java callback functions. @@ -62,7 +62,7 @@ public class MockWeatherService implements Function { === Registering Functions as Beans -With the link:../azure-openai-chat.html#_auto_configuration[AzureOpenAiChatClient Auto-Configuration] you have multiple ways to register custom functions as beans in the Spring context. +With the link:../azure-openai-chat.html#_auto_configuration[AzureOpenAiChatModelAuto-Configuration] you have multiple ways to register custom functions as beans in the Spring context. We start with describing the most POJO friendly options. @@ -70,7 +70,7 @@ We start with describing the most POJO friendly options. In this approach you define `@Beans` in your application context as you would any other Spring managed object. -Internally, Spring AI `ChatClient` will create an instance of a `FunctionCallbackWrapper` wrapper that adds the logic for it being invoked via the AI model. +Internally, Spring AI `ChatModel` will create an instance of a `FunctionCallbackWrapper` wrapper that adds the logic for it being invoked via the AI model. The name of the `@Bean` is passed as a `ChatOption`. @@ -134,7 +134,7 @@ static class Config { } ---- -It wraps the 3rd party `MockWeatherService` function and registers it as a `CurrentWeather` function with the `AzureAiChatClient` and provides a description (2). +It wraps the 3rd party `MockWeatherService` function and registers it as a `CurrentWeather` function with the `AzureAiChatModel` and provides a description (2). NOTE: The default response converter does a JSON serialization of the Response object. @@ -146,17 +146,17 @@ To let the model know and call your `CurrentWeather` function you need to enable [source,java] ---- -AzureOpenAiChatClient chatClient = ... +AzureOpenAiChatModel chatModel = ... UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), AzureOpenAiChatOptions.builder().withFunction("CurrentWeather").build())); // (1) Enable the function logger.info("Response: {}", response); ---- -// NOTE: You can can have multiple functions registered in your `ChatClient` but only those enabled in the prompt request will be considered for the function calling. +// NOTE: You can can have multiple functions registered in your `ChatModel` but only those enabled in the prompt request will be considered for the function calling. Above user question will trigger 3 calls to `CurrentWeather` function (one for each city) and the final response will be something like this: @@ -176,7 +176,7 @@ In addition to the auto-configuration you can register callback functions, dynam [source,java] ---- -AzureOpenAiChatClient chatClient = ... +AzureOpenAiChatModel chatModel = ... UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris? Use Multi-turn function calling."); @@ -187,12 +187,12 @@ var promptOptions = AzureOpenAiChatOptions.builder() .build())) .build(); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), promptOptions)); +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), promptOptions)); ---- NOTE: The in-prompt registered functions are enabled by default for the duration of this request. This approach allows to dynamically chose different functions to be called based on the user input. -The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithPromptFunctionIT.java[FunctionCallWithPromptFunctionIT.java] integration test provides a complete example of how to register a function with the `AzureOpenAiChatClient` and use it in a prompt request. +The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithPromptFunctionIT.java[FunctionCallWithPromptFunctionIT.java] integration test provides a complete example of how to register a function with the `AzureOpenAiChatModel` and use it in a prompt request. diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/minimax-chat-functions.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/minimax-chat-functions.adoc index 473595ca2..2d6380bd7 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/minimax-chat-functions.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/minimax-chat-functions.adoc @@ -1,6 +1,6 @@ = Function Calling -You can register custom Java functions with the `MiniMaxChatClient` and have the MiniMax model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. +You can register custom Java functions with the `MiniMaxChatModel` and have the MiniMax model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. This allows you to connect the LLM capabilities with external tools and APIs. The MiniMax models are trained to detect when a function should be called and to respond with JSON that adheres to the function signature. @@ -11,12 +11,12 @@ In general, the custom functions need to provide a function `name`, `descriptio As a developer, you need to implement a functions that takes the function call arguments sent from the AI model, and respond with the result back to the model. Your function can in turn invoke other 3rd party services to provide the results. -Spring AI makes this as easy as defining a `@Bean` definition that returns a `java.util.Function` and supplying the bean name as an option when invoking the `ChatClient`. +Spring AI makes this as easy as defining a `@Bean` definition that returns a `java.util.Function` and supplying the bean name as an option when invoking the `ChatModel`. Under the hood, Spring wraps your POJO (the function) with the appropriate adapter code that enables interaction with the AI Model, saving you from writing tedious boilerplate code. The basis of the underlying infrastructure is the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallback.java[FunctionCallback.java] interface and the companion link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallbackWrapper.java[FunctionCallbackWrapper.java] utility class to simplify the implementation and registration of Java callback functions. -// Additionally, the Auto-Configuration provides a way to auto-register any Function beans definition as function calling candidates in the `ChatClient`. +// Additionally, the Auto-Configuration provides a way to auto-register any Function beans definition as function calling candidates in the `ChatModel`. == How it works @@ -62,7 +62,7 @@ public class MockWeatherService implements Function { === Registering Functions as Beans -With the link:../minimax-chat.html#_auto_configuration[MiniMaxChatClient Auto-Configuration] you have multiple ways to register custom functions as beans in the Spring context. +With the link:../minimax-chat.html#_auto_configuration[MiniMaxChatModel Auto-Configuration] you have multiple ways to register custom functions as beans in the Spring context. We start with describing the most POJO friendly options. @@ -71,7 +71,7 @@ We start with describing the most POJO friendly options. In this approach you define `@Beans` in your application context as you would any other Spring managed object. -Internally, Spring AI `ChatClient` will create an instance of a `FunctionCallbackWrapper` wrapper that adds the logic for it being invoked via the AI model. +Internally, Spring AI `ChatModel` will create an instance of a `FunctionCallbackWrapper` wrapper that adds the logic for it being invoked via the AI model. The name of the `@Bean` is passed as a `ChatOption`. @@ -136,7 +136,7 @@ static class Config { } ---- -It wraps the 3rd party, `MockWeatherService` function and registers it as a `CurrentWeather` function with the `MiniMaxChatClient`. +It wraps the 3rd party, `MockWeatherService` function and registers it as a `CurrentWeather` function with the `MiniMaxChatModel`. It also provides a description (2) and an optional response converter (3) to convert the response into a text as expected by the model. NOTE: By default, the response converter does a JSON serialization of the Response object. @@ -149,17 +149,17 @@ To let the model know and call your `CurrentWeather` function you need to enable [source,java] ---- -MiniMaxChatClient chatClient = ... +MiniMaxChatModel chatModel = ... UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), MiniMaxChatOptions.builder().withFunction("CurrentWeather").build())); // (1) Enable the function logger.info("Response: {}", response); ---- -// NOTE: You can can have multiple functions registered in your `ChatClient` but only those enabled in the prompt request will be considered for the function calling. +// NOTE: You can can have multiple functions registered in your `ChatModel` but only those enabled in the prompt request will be considered for the function calling. Above user question will trigger 3 calls to `CurrentWeather` function (one for each city) and the final response will be something like this: @@ -179,7 +179,7 @@ In addition to the auto-configuration you can register callback functions, dynam [source,java] ---- -MiniMaxChatClient chatClient = ... +MiniMaxChatModel chatModel = ... UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -190,18 +190,18 @@ var promptOptions = MiniMaxChatOptions.builder() new MockWeatherService()))) // function code .build(); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), promptOptions)); +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), promptOptions)); ---- NOTE: The in-prompt registered functions are enabled by default for the duration of this request. This approach allows to dynamically chose different functions to be called based on the user input. -The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/tool/FunctionCallbackInPromptIT.java[FunctionCallbackInPromptIT.java] integration test provides a complete example of how to register a function with the `MiniMaxChatClient` and use it in a prompt request. +The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/tool/FunctionCallbackInPromptIT.java[FunctionCallbackInPromptIT.java] integration test provides a complete example of how to register a function with the `MiniMaxChatModel` and use it in a prompt request. // // === Register Functions with Default Options // -// You can programmatically register functions with the `MiniMaxChatClient` using the `MiniMaxChatOptions#withFunctionCallbacks`: +// You can programmatically register functions with the `MiniMaxChatModel` using the `MiniMaxChatOptions#withFunctionCallbacks`: // // [source,java] // ---- @@ -215,12 +215,12 @@ The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot // new MockWeatherService()))) // function code // .build(); // -// MiniMaxChatClient chatClient = new MiniMaxChatClient(miniMaxApi, defaultOptions); +// MiniMaxChatModel chatModel = new MiniMaxChatModel(miniMaxApi, defaultOptions); // // UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); // -// ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), +// ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), // MiniMaxChatOptions.builder().withFunction("CurrentWeather").build())); // Enable the function // ---- // -// NOTE: Functions are registered when MiniMaxChatClient is created, by you must enable in the Prompt the functions to be used in the request. \ No newline at end of file +// NOTE: Functions are registered when MiniMaxChatModel is created, by you must enable in the Prompt the functions to be used in the request. \ No newline at end of file diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/mistralai-chat-functions.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/mistralai-chat-functions.adoc index eb1184795..77fd3dd2e 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/mistralai-chat-functions.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/mistralai-chat-functions.adoc @@ -1,6 +1,6 @@ = Mistral AI Function Calling -You can register custom Java functions with the `MistralAiChatClient` and have the Mistral AI models intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. +You can register custom Java functions with the `MistralAiChatModel` and have the Mistral AI models intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. This allows you to connect the LLM capabilities with external tools and APIs. The `open-mixtral-8x22b`, `mistral_small_latest`, and `mistral_large_latest` models are trained to detect when a function should be called and to respond with JSON that adheres to the function signature. @@ -15,7 +15,7 @@ The `description` helps the model to understand when to call the function. As a developer, you need to implement a function that takes the function call arguments sent from the AI model, and respond with the result back to the model. Your function can in turn invoke other 3rd party services to provide the results. -Spring AI makes this as easy as defining a `@Bean` definition that returns a `java.util.Function` and supplying the bean name as an option when invoking the `ChatClient`. +Spring AI makes this as easy as defining a `@Bean` definition that returns a `java.util.Function` and supplying the bean name as an option when invoking the `ChatModel`. Under the hood, Spring wraps your POJO (the function) with the appropriate adapter code that enables interaction with the AI Model, saving you from writing tedious boilerplate code. The basis of the underlying infrastructure is the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallback.java[FunctionCallback.java] interface and the companion link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallbackWrapper.java[FunctionCallbackWrapper.java] utility class to simplify the implementation and registration of Java callback functions. @@ -62,7 +62,7 @@ public class MockWeatherService implements Function { === Registering Functions as Beans -With the link:../mistralai-chat.html#_auto_configuration[MistralAiChatClient Auto-Configuration] you have multiple ways to register custom functions as beans in the Spring context. +With the link:../mistralai-chat.html#_auto_configuration[MistralAiChatModel Auto-Configuration] you have multiple ways to register custom functions as beans in the Spring context. We start with describing the most POJO friendly options. @@ -70,7 +70,7 @@ We start with describing the most POJO friendly options. In this approach you define `@Beans` in your application context as you would any other Spring managed object. -Internally, Spring AI `ChatClient` will create an instance of a `FunctionCallbackWrapper` wrapper that adds the logic for it being invoked via the AI model. +Internally, Spring AI `ChatModel` will create an instance of a `FunctionCallbackWrapper` wrapper that adds the logic for it being invoked via the AI model. The name of the `@Bean` is passed as a `ChatOption`. @@ -139,7 +139,7 @@ static class Config { } ---- -It wraps the 3rd party `MockWeatherService` function and registers it as a `CurrentWeather` function with the `MistralAiChatClient`. +It wraps the 3rd party `MockWeatherService` function and registers it as a `CurrentWeather` function with the `MistralAiChatModel`. It also provides a description (2) and an optional response converter (3) to convert the response into a text as expected by the model. NOTE: By default, the response converter does a JSON serialization of the Response object. @@ -152,17 +152,17 @@ To let the model know and call your `CurrentWeather` function you need to enable [source,java] ---- -MistralAiChatClient chatClient = ... +MistralAiChatModel chatModel = ... UserMessage userMessage = new UserMessage("What's the weather like in Paris?"); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), MistralAiChatOptions.builder().withFunction("CurrentWeather").build())); // (1) Enable the function logger.info("Response: {}", response); ---- -// NOTE: You can can have multiple functions registered in your `ChatClient` but only those enabled in the prompt request will be considered for the function calling. +// NOTE: You can can have multiple functions registered in your `ChatModel` but only those enabled in the prompt request will be considered for the function calling. Above user question will trigger 3 calls to `CurrentWeather` function (one for each city) and produce the final response. @@ -172,7 +172,7 @@ In addition to the auto-configuration you can register callback functions, dynam [source,java] ---- -MistralAiChatClient chatClient = ... +MistralAiChatModel chatModel = ... UserMessage userMessage = new UserMessage("What's the weather like in Paris?"); @@ -183,14 +183,14 @@ var promptOptions = MistralAiChatOptions.builder() new MockWeatherService()))) // function code .build(); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), promptOptions)); +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), promptOptions)); ---- NOTE: The in-prompt registered functions are enabled by default for the duration of this request. This approach allows to dynamically chose different functions to be called based on the user input. -The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/PaymentStatusPromptIT.java[PaymentStatusPromptIT.java] integration test provides a complete example of how to register a function with the `MistralAiChatClient` and use it in a prompt request. +The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/PaymentStatusPromptIT.java[PaymentStatusPromptIT.java] integration test provides a complete example of how to register a function with the `MistralAiChatModel` and use it in a prompt request. == Appendices diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/openai-chat-functions.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/openai-chat-functions.adoc index 0d85e488c..699bf7163 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/openai-chat-functions.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/openai-chat-functions.adoc @@ -1,6 +1,6 @@ = Function Calling -You can register custom Java functions with the `OpenAiChatClient` and have the OpenAI model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. +You can register custom Java functions with the `OpenAiChatModel` and have the OpenAI model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. This allows you to connect the LLM capabilities with external tools and APIs. The OpenAI models are trained to detect when a function should be called and to respond with JSON that adheres to the function signature. @@ -11,12 +11,12 @@ In general, the custom functions need to provide a function `name`, `descriptio As a developer, you need to implement a function that takes the function call arguments sent from the AI model, and respond with the result back to the model. Your function can in turn invoke other 3rd party services to provide the results. -Spring AI makes this as easy as defining a `@Bean` definition that returns a `java.util.Function` and supplying the bean name as an option when invoking the `ChatClient`. +Spring AI makes this as easy as defining a `@Bean` definition that returns a `java.util.Function` and supplying the bean name as an option when invoking the `ChatModel`. Under the hood, Spring wraps your POJO (the function) with the appropriate adapter code that enables interaction with the AI Model, saving you from writing tedious boilerplate code. The basis of the underlying infrastructure is the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallback.java[FunctionCallback.java] interface and the companion link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallbackWrapper.java[FunctionCallbackWrapper.java] utility class to simplify the implementation and registration of Java callback functions. -// Additionally, the Auto-Configuration provides a way to auto-register any Function beans definition as function calling candidates in the `ChatClient`. +// Additionally, the Auto-Configuration provides a way to auto-register any Function beans definition as function calling candidates in the `ChatModel`. == How it works @@ -62,7 +62,7 @@ public class MockWeatherService implements Function { === Registering Functions as Beans -With the link:../openai-chat.html#_auto_configuration[OpenAiChatClient Auto-Configuration] you have multiple ways to register custom functions as beans in the Spring context. +With the link:../openai-chat.html#_auto_configuration[OpenAiChatModel Auto-Configuration] you have multiple ways to register custom functions as beans in the Spring context. We start with describing the most POJO friendly options. @@ -71,7 +71,7 @@ We start with describing the most POJO friendly options. In this approach you define `@Beans` in your application context as you would any other Spring managed object. -Internally, Spring AI `ChatClient` will create an instance of a `FunctionCallbackWrapper` wrapper that adds the logic for it being invoked via the AI model. +Internally, Spring AI `ChatModel` will create an instance of a `FunctionCallbackWrapper` wrapper that adds the logic for it being invoked via the AI model. The name of the `@Bean` is passed as a `ChatOption`. @@ -136,7 +136,7 @@ static class Config { } ---- -It wraps the 3rd party `MockWeatherService` function and registers it as a `CurrentWeather` function with the `OpenAiChatClient`. +It wraps the 3rd party `MockWeatherService` function and registers it as a `CurrentWeather` function with the `OpenAiChatModel`. It also provides a description (2) and an optional response converter (3) to convert the response into a text as expected by the model. NOTE: By default, the response converter does a JSON serialization of the Response object. @@ -149,17 +149,17 @@ To let the model know and call your `CurrentWeather` function you need to enable [source,java] ---- -OpenAiChatClient chatClient = ... +OpenAiChatModel chatModel = ... UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withFunction("CurrentWeather").build())); // (1) Enable the function logger.info("Response: {}", response); ---- -// NOTE: You can can have multiple functions registered in your `ChatClient` but only those enabled in the prompt request will be considered for the function calling. +// NOTE: You can can have multiple functions registered in your `ChatModel` but only those enabled in the prompt request will be considered for the function calling. Above user question will trigger 3 calls to `CurrentWeather` function (one for each city) and the final response will be something like this: @@ -179,7 +179,7 @@ In addition to the auto-configuration you can register callback functions, dynam [source,java] ---- -OpenAiChatClient chatClient = ... +OpenAiChatModel chatModel = ... UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -190,18 +190,18 @@ var promptOptions = OpenAiChatOptions.builder() new MockWeatherService()))) // function code .build(); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), promptOptions)); +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), promptOptions)); ---- NOTE: The in-prompt registered functions are enabled by default for the duration of this request. This approach allows to dynamically chose different functions to be called based on the user input. -The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackInPromptIT.java[FunctionCallbackInPromptIT.java] integration test provides a complete example of how to register a function with the `OpenAiChatClient` and use it in a prompt request. +The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackInPromptIT.java[FunctionCallbackInPromptIT.java] integration test provides a complete example of how to register a function with the `OpenAiChatModel` and use it in a prompt request. // // === Register Functions with Default Options // -// You can programmatically register functions with the `OpenAiChatClient` using the `OpenAiChatOptions#withFunctionCallbacks`: +// You can programmatically register functions with the `OpenAiChatModel` using the `OpenAiChatOptions#withFunctionCallbacks`: // // [source,java] // ---- @@ -215,24 +215,24 @@ The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot // new MockWeatherService()))) // function code // .build(); // -// OpenAiChatClient chatClient = new OpenAiChatClient(openaiApi, defaultOptions); +// OpenAiChatModel chatModel = new OpenAiChatModel(openaiApi, defaultOptions); // // UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); // -// ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), +// ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), // OpenAiChatOptions.builder().withFunction("CurrentWeather").build())); // Enable the function // ---- // -// NOTE: Functions are registered when OpenAiChatClient is created, by you must enable in the Prompt the functions to be used in the request. +// NOTE: Functions are registered when OpenAiChatModel is created, by you must enable in the Prompt the functions to be used in the request. == Appendices: === Spring AI Function Calling Flow [[spring-ai-function-calling-flow]] -The following diagram illustrates the flow of the OpenAiChatClient Function Calling: +The following diagram illustrates the flow of the OpenAiChatModel Function Calling: -image:openai-chatclient-function-call.jpg[width=800, title="OpenAiChatClient Function Calling Flow"] +image:openai-chatclient-function-call.jpg[width=800, title="OpenAiChatModel Function Calling Flow"] === OpenAI API Function Calling Flow diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/vertexai-gemini-chat-functions.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/vertexai-gemini-chat-functions.adoc index a464ccf51..d09e0a727 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/vertexai-gemini-chat-functions.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/vertexai-gemini-chat-functions.adoc @@ -6,7 +6,7 @@ The parallel function calling is gone as well. Function calling lets developers create a description of a function in their code, then pass that description to a language model in a request. The response from the model includes the name of a function that matches the description and the arguments to call it with. -You can register custom Java functions with the `VertexAiGeminiChatClient` and have the Gemini Pro model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. +You can register custom Java functions with the `VertexAiGeminiChatModel` and have the Gemini Pro model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. This allows you to connect the LLM capabilities with external tools and APIs. The VertexAI Gemini Pro model is trained to detect when a function should be called and to respond with JSON that adheres to the function signature. @@ -18,12 +18,12 @@ In general, the custom functions need to provide a function `name`, `description As a developer, you need to implement a function that takes the function call arguments sent from the AI model, and respond with the result back to the model. Your function can in turn invoke other 3rd party services to provide the results. -Spring AI makes this as easy as defining a `@Bean` definition that returns a `java.util.Function` and supplying the bean name as an option when invoking the `ChatClient`. +Spring AI makes this as easy as defining a `@Bean` definition that returns a `java.util.Function` and supplying the bean name as an option when invoking the `ChatModel`. Under the hood, Spring wraps your POJO (the function) with the appropriate adapter code that enables interaction with the AI Model, saving you from writing tedious boilerplate code. The basis of the underlying infrastructure is the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallback.java[FunctionCallback.java] interface and the companion link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallbackWrapper.java[FunctionCallbackWrapper.java] utility class to simplify the implementation and registration of Java callback functions. -// Additionally, the Auto-Configuration provides a way to auto-register any Function beans definition as function calling candidates in the `ChatClient`. +// Additionally, the Auto-Configuration provides a way to auto-register any Function beans definition as function calling candidates in the `ChatModel`. == How it works @@ -66,7 +66,7 @@ public class MockWeatherService implements Function { === Registering Functions as Beans -With the link:../vertexai-gemini-chat.html#_auto_configuration[VertexAiGeminiChatClient Auto-Configuration] you have multiple ways to register custom functions as beans in the Spring context. +With the link:../vertexai-gemini-chat.html#_auto_configuration[VertexAiGeminiChatModel Auto-Configuration] you have multiple ways to register custom functions as beans in the Spring context. We start with describing the most POJO friendly options. @@ -74,7 +74,7 @@ We start with describing the most POJO friendly options. In this approach you define `@Beans` in your application context as you would any other Spring managed object. -Internally, Spring AI `ChatClient` will create an instance of a `FunctionCallbackWrapper` wrapper that adds the logic for it being invoked via the AI model. +Internally, Spring AI `ChatModel` will create an instance of a `FunctionCallbackWrapper` wrapper that adds the logic for it being invoked via the AI model. The name of the `@Bean` is passed as a `ChatOption`. @@ -139,7 +139,7 @@ static class Config { } ---- -It wraps the 3rd party `MockWeatherService` function and registers it as a `CurrentWeather` function with the `VertexAiGeminiChatClient`. +It wraps the 3rd party `MockWeatherService` function and registers it as a `CurrentWeather` function with the `VertexAiGeminiChatModel`. It also provides a description (2) and sets the Schema type to Open API type (3). NOTE: The default response converter does a JSON serialization of the Response object. @@ -152,17 +152,17 @@ To let the model know and call your `CurrentWeather` function you need to enable [source,java] ---- -VertexAiGeminiChatClient chatClient = ... +VertexAiGeminiChatModel chatModel = ... UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), VertexAiGeminiChatOptions.builder().withFunction("CurrentWeather").build())); // (1) Enable the function logger.info("Response: {}", response); ---- -// NOTE: You can can have multiple functions registered in your `ChatClient` but only those enabled in the prompt request will be considered for the function calling. +// NOTE: You can can have multiple functions registered in your `ChatModel` but only those enabled in the prompt request will be considered for the function calling. Above user question will trigger 3 calls to `CurrentWeather` function (one for each city) and the final response will be something like this: @@ -182,7 +182,7 @@ In addition to the auto-configuration you can register callback functions, dynam [source,java] ---- -VertexAiGeminiChatClient chatClient = ... +VertexAiGeminiChatModel chatModel = ... UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris? Use Multi-turn function calling."); @@ -194,12 +194,12 @@ var promptOptions = VertexAiGeminiChatOptions.builder() .build())) .build(); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), promptOptions)); +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), promptOptions)); ---- NOTE: The in-prompt registered functions are enabled by default for the duration of this request. This approach allows to dynamically chose different functions to be called based on the user input. -The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/gemini/tool/FunctionCallWithPromptFunctionIT.java[FunctionCallWithPromptFunctionIT.java] integration test provides a complete example of how to register a function with the `VertexAiGeminiChatClient` and use it in a prompt request. +The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/gemini/tool/FunctionCallWithPromptFunctionIT.java[FunctionCallWithPromptFunctionIT.java] integration test provides a complete example of how to register a function with the `VertexAiGeminiChatModel` and use it in a prompt request. diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/zhipuai-chat-functions.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/zhipuai-chat-functions.adoc index 338505b78..de25ce8b0 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/zhipuai-chat-functions.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/functions/zhipuai-chat-functions.adoc @@ -1,6 +1,6 @@ = Function Calling -You can register custom Java functions with the `ZhiPuAiChatClient` and have the ZhiPuAI model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. +You can register custom Java functions with the `ZhiPuAiChatModel` and have the ZhiPuAI model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. This allows you to connect the LLM capabilities with external tools and APIs. The ZhiPuAI models are trained to detect when a function should be called and to respond with JSON that adheres to the function signature. @@ -11,12 +11,12 @@ In general, the custom functions need to provide a function `name`, `descriptio As a developer, you need to implement a functions that takes the function call arguments sent from the AI model, and respond with the result back to the model. Your function can in turn invoke other 3rd party services to provide the results. -Spring AI makes this as easy as defining a `@Bean` definition that returns a `java.util.Function` and supplying the bean name as an option when invoking the `ChatClient`. +Spring AI makes this as easy as defining a `@Bean` definition that returns a `java.util.Function` and supplying the bean name as an option when invoking the `ChatModel`. Under the hood, Spring wraps your POJO (the function) with the appropriate adapter code that enables interaction with the AI Model, saving you from writing tedious boilerplate code. The basis of the underlying infrastructure is the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallback.java[FunctionCallback.java] interface and the companion link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallbackWrapper.java[FunctionCallbackWrapper.java] utility class to simplify the implementation and registration of Java callback functions. -// Additionally, the Auto-Configuration provides a way to auto-register any Function beans definition as function calling candidates in the `ChatClient`. +// Additionally, the Auto-Configuration provides a way to auto-register any Function beans definition as function calling candidates in the `ChatModel`. == How it works @@ -62,7 +62,7 @@ public class MockWeatherService implements Function { === Registering Functions as Beans -With the link:../zhipuai-chat.html#_auto_configuration[ZhiPuAiChatClient Auto-Configuration] you have multiple ways to register custom functions as beans in the Spring context. +With the link:../zhipuai-chat.html#_auto_configuration[ZhiPuAiChatModel Auto-Configuration] you have multiple ways to register custom functions as beans in the Spring context. We start with describing the most POJO friendly options. @@ -71,7 +71,7 @@ We start with describing the most POJO friendly options. In this approach you define `@Beans` in your application context as you would any other Spring managed object. -Internally, Spring AI `ChatClient` will create an instance of a `FunctionCallbackWrapper` wrapper that adds the logic for it being invoked via the AI model. +Internally, Spring AI `ChatModel` will create an instance of a `FunctionCallbackWrapper` wrapper that adds the logic for it being invoked via the AI model. The name of the `@Bean` is passed as a `ChatOption`. @@ -136,7 +136,7 @@ static class Config { } ---- -It wraps the 3rd party, `MockWeatherService` function and registers it as a `CurrentWeather` function with the `ZhiPuAiChatClient`. +It wraps the 3rd party, `MockWeatherService` function and registers it as a `CurrentWeather` function with the `ZhiPuAiChatModel`. It also provides a description (2) and an optional response converter (3) to convert the response into a text as expected by the model. NOTE: By default, the response converter does a JSON serialization of the Response object. @@ -149,17 +149,17 @@ To let the model know and call your `CurrentWeather` function you need to enable [source,java] ---- -ZhiPuAiChatClient chatClient = ... +ZhiPuAiChatModel chatModel = ... UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), ZhiPuAiChatOptions.builder().withFunction("CurrentWeather").build())); // (1) Enable the function logger.info("Response: {}", response); ---- -// NOTE: You can can have multiple functions registered in your `ChatClient` but only those enabled in the prompt request will be considered for the function calling. +// NOTE: You can can have multiple functions registered in your `ChatModel` but only those enabled in the prompt request will be considered for the function calling. Above user question will trigger 3 calls to `CurrentWeather` function (one for each city) and the final response will be something like this: @@ -179,7 +179,7 @@ In addition to the auto-configuration you can register callback functions, dynam [source,java] ---- -ZhiPuAiChatClient chatClient = ... +ZhiPuAiChatModel chatModel = ... UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -190,18 +190,18 @@ var promptOptions = ZhiPuAiChatOptions.builder() new MockWeatherService()))) // function code .build(); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), promptOptions)); +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), promptOptions)); ---- NOTE: The in-prompt registered functions are enabled by default for the duration of this request. This approach allows to dynamically chose different functions to be called based on the user input. -The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackInPromptIT.java[FunctionCallbackInPromptIT.java] integration test provides a complete example of how to register a function with the `ZhiPuAiChatClient` and use it in a prompt request. +The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackInPromptIT.java[FunctionCallbackInPromptIT.java] integration test provides a complete example of how to register a function with the `ZhiPuAiChatModel` and use it in a prompt request. // // === Register Functions with Default Options // -// You can programmatically register functions with the `ZhiPuAiChatClient` using the `ZhiPuAiChatOptions#withFunctionCallbacks`: +// You can programmatically register functions with the `ZhiPuAiChatModel using the `ZhiPuAiChatOptions#withFunctionCallbacks`: // // [source,java] // ---- @@ -215,12 +215,12 @@ The https://github.com/spring-projects/spring-ai/blob/main/spring-ai-spring-boot // new MockWeatherService()))) // function code // .build(); // -// ZhiPuAiChatClient chatClient = new ZhiPuAiChatClient(zhiPuAiApi, defaultOptions); +// ZhiPuAiChatModel chatModel = new ZhiPuAiChatModel(zhiPuAiApi, defaultOptions); // // UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); // -// ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), +// ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), // ZhiPuAiChatOptions.builder().withFunction("CurrentWeather").build())); // Enable the function // ---- // -// NOTE: Functions are registered when ZhiPuAiChatClient is created, by you must enable in the Prompt the functions to be used in the request. \ No newline at end of file +// NOTE: Functions are registered when ZhiPuAiChatModel is created, by you must enable in the Prompt the functions to be used in the request. \ No newline at end of file diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/huggingface.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/huggingface.adoc index eeb815229..a356d54f2 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/huggingface.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/huggingface.adoc @@ -27,7 +27,7 @@ export HUGGINGFACE_API_KEY=your_api_key_here TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Note, there is not yet a Spring Boot Starter for this client implementation. +Note, there is not yet a Spring Boot Starter for this chat implementation. Obtain the endpoint URL of the Inference Endpoint. You can find this on the Inference Endpoint's UI link:https://ui.endpoints.huggingface.co/[here]. @@ -36,9 +36,9 @@ You can find this on the Inference Endpoint's UI link:https://ui.endpoints.huggi [source,java] ---- -HuggingfaceChatClient client = new HuggingfaceChatClient(apiKey, basePath); +HuggingfaceChatModel chatModel = new HuggingfaceChatModel(apiKey, basePath); Prompt prompt = new Prompt("Your text here..."); -ChatResponse response = client.call(prompt); +ChatResponse response = chatModel.call(prompt); System.out.println(response.getGeneration().getText()); ---- @@ -56,7 +56,7 @@ String mistral7bInstruct = """ Just generate the JSON object without explanations: [/INST]"""; Prompt prompt = new Prompt(mistral7bInstruct); -ChatResponse aiResponse = huggingfaceChatClient.call(prompt); +ChatResponse aiResponse = huggingfaceChatModel.call(prompt); System.out.println(response.getGeneration().getText()); ---- Will produce the output diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/minimax-chat.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/minimax-chat.adoc index 29c7ff35e..4827ac891 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/minimax-chat.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/minimax-chat.adoc @@ -52,7 +52,7 @@ TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Man ==== Retry Properties -The prefix `spring.ai.retry` is used as the property prefix that lets you configure the retry mechanism for the MiniMax Chat client. +The prefix `spring.ai.retry` is used as the property prefix that lets you configure the retry mechanism for the MiniMax chat model. [cols="3,5,1"] |==== @@ -81,13 +81,13 @@ The prefix `spring.ai.minimax` is used as the property prefix that lets you conn ==== Configuration Properties -The prefix `spring.ai.minimax.chat` is the property prefix that lets you configure the chat client implementation for MiniMax. +The prefix `spring.ai.minimax.chat` is the property prefix that lets you configure the chat model implementation for MiniMax. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.minimax.chat.enabled | Enable MiniMax chat client. | true +| spring.ai.minimax.chat.enabled | Enable MiniMax chat model. | true | spring.ai.minimax.chat.base-url | Optional overrides the spring.ai.minimax.base-url to provide chat specific url | https://api.minimax.chat | spring.ai.minimax.chat.api-key | Optional overrides the spring.ai.minimax.api-key to provide chat specific api-key | - | spring.ai.minimax.chat.options.model | This is the MiniMax Chat model to use | `abab5.5-chat` (the `abab5.5s-chat`, `abab5.5-chat`, and `abab6-chat` point to the latest model versions) @@ -100,7 +100,7 @@ The prefix `spring.ai.minimax.chat` is the property prefix that lets you configu | spring.ai.minimax.chat.options.stop | The model will stop generating characters specified by stop, and currently only supports a single stop word in the format of ["stop_word1"] | - |==== -NOTE: You can override the common `spring.ai.minimax.base-url` and `spring.ai.minimax.api-key` for the `ChatClient` implementations. +NOTE: You can override the common `spring.ai.minimax.base-url` and `spring.ai.minimax.api-key` for the `ChatModel` implementations. The `spring.ai.minimax.chat.base-url` and `spring.ai.minimax.chat.api-key` properties if set take precedence over the common properties. This is useful if you want to use different MiniMax accounts for different models and different model endpoints. @@ -110,14 +110,14 @@ TIP: All properties prefixed with `spring.ai.minimax.chat.options` can be overri The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatOptions.java[MiniMaxChatOptions.java] provides model configurations, such as the model to use, the temperature, the frequency penalty, etc. -On start-up, the default options can be configured with the `MiniMaxChatClient(api, options)` constructor or the `spring.ai.minimax.chat.options.*` properties. +On start-up, the default options can be configured with the `MiniMaxChatModel(api, options)` constructor or the `spring.ai.minimax.chat.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default model and temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", MiniMaxChatOptions.builder() @@ -133,7 +133,7 @@ TIP: In addition to the model specific link:https://github.com/spring-projects/s https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-minimax-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the MiniMax Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the MiniMax chat model: [source,application.properties] ---- @@ -144,37 +144,37 @@ spring.ai.minimax.chat.options.temperature=0.7 TIP: replace the `api-key` with your MiniMax credentials. -This will create a `MiniMaxChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `MiniMaxChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class ChatController { - private final MiniMaxChatClient chatClient; + private final MiniMaxChatModel chatModel; @Autowired - public ChatController(MiniMaxChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(MiniMaxChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } @GetMapping("/ai/generateStream") public Flux generateStream(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { var prompt = new Prompt(new UserMessage(message)); - return chatClient.stream(prompt); + return chatModel.stream(prompt); } } ---- == Manual Configuration -The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatClient.java[MiniMaxChatClient] implements the `ChatClient` and `StreamingChatClient` and uses the <> to connect to the MiniMax service. +The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-minimax/src/main/java/org/springframework/ai/minimax/MiniMaxChatModel.java[MiniMaxChatModel] implements the `ChatModel` and `StreamingChatModel` and uses the <> to connect to the MiniMax service. Add the `spring-ai-minimax` dependency to your project's Maven `pom.xml` file: @@ -197,23 +197,23 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create a `MiniMaxChatClient` and use it for text generations: +Next, create a `MiniMaxChatModel` and use it for text generations: [source,java] ---- var miniMaxApi = new MiniMaxApi(System.getenv("MINIMAX_API_KEY")); -var chatClient = new MiniMaxChatClient(miniMaxApi, MiniMaxChatOptions.builder() +var chatModel = new MiniMaxChatModel(miniMaxApi, MiniMaxChatOptions.builder() .withModel(MiniMaxApi.ChatModel.GLM_3_Turbo.getValue()) .withTemperature(0.4f) .withMaxTokens(200) .build()); -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); // Or with streaming responses -Flux streamResponse = chatClient.stream( +Flux streamResponse = chatModel.stream( new Prompt("Generate the names of 5 famous pirates.")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/mistralai-chat.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/mistralai-chat.adoc index 26fb6afee..ba184e63f 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/mistralai-chat.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/mistralai-chat.adoc @@ -51,7 +51,7 @@ TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Man ==== Retry Properties -The prefix `spring.ai.retry` is used as the property prefix that lets you configure the retry mechanism for the Mistral AI Chat client. +The prefix `spring.ai.retry` is used as the property prefix that lets you configure the retry mechanism for the Mistral AI chat model. [cols="3,5,1"] |==== @@ -80,13 +80,13 @@ The prefix `spring.ai.mistralai` is used as the property prefix that lets you co ==== Configuration Properties -The prefix `spring.ai.mistralai.chat` is the property prefix that lets you configure the chat client implementation for MistralAI. +The prefix `spring.ai.mistralai.chat` is the property prefix that lets you configure the chat model implementation for MistralAI. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.mistralai.chat.enabled | Enable MistralAI chat client. | true +| spring.ai.mistralai.chat.enabled | Enable MistralAI chat model. | true | spring.ai.mistralai.chat.base-url | Optional overrides the spring.ai.mistralai.base-url to provide chat specific url | - | spring.ai.mistralai.chat.api-key | Optional overrides the spring.ai.mistralai.api-key to provide chat specific api-key | - | spring.ai.mistralai.chat.options.model | This is the MistralAI Chat model to use | `open-mistral-7b`, `open-mixtral-8x7b`, `mistral-small-latest`, `mistral-medium-latest`, `mistral-large-latest` @@ -100,10 +100,10 @@ The prefix `spring.ai.mistralai.chat` is the property prefix that lets you confi | spring.ai.mistralai.chat.options.tools | A list of tools the model may call. Currently, only functions are supported as a tool. Use this to provide a list of functions the model may generate JSON inputs for. | - | spring.ai.mistralai.chat.options.toolChoice | Controls which (if any) function is called by the model. none means the model will not call a function and instead generates a message. auto means the model can pick between generating a message or calling a function. Specifying a particular function via {"type: "function", "function": {"name": "my_function"}} forces the model to call that function. none is the default when no functions are present. auto is the default if functions are present. | - | spring.ai.mistralai.chat.options.functions | List of functions, identified by their names, to enable for function calling in a single prompt requests. Functions with those names must exist in the functionCallbacks registry. | - -| spring.ai.mistralai.chat.options.functionCallbacks | MistralAI Tool Function Callbacks to register with the ChatClient. | - +| spring.ai.mistralai.chat.options.functionCallbacks | MistralAI Tool Function Callbacks to register with the ChatModel. | - |==== -NOTE: You can override the common `spring.ai.mistralai.base-url` and `spring.ai.mistralai.api-key` for the `ChatClient` and `EmbeddingClient` implementations. +NOTE: You can override the common `spring.ai.mistralai.base-url` and `spring.ai.mistralai.api-key` for the `ChatModel` and `EmbeddingModel` implementations. The `spring.ai.mistralai.chat.base-url` and `spring.ai.mistralai.chat.api-key` properties if set take precedence over the common properties. This is useful if you want to use different MistralAI accounts for different models and different model endpoints. @@ -113,14 +113,14 @@ TIP: All properties prefixed with `spring.ai.mistralai.chat.options` can be over The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatOptions.java[MistralAiChatOptions.java] provides model configurations, such as the model to use, the temperature, the frequency penalty, etc. -On start-up, the default options can be configured with the `MistralAiChatClient(api, options)` constructor or the `spring.ai.mistralai.chat.options.*` properties. +On start-up, the default options can be configured with the `MistralAiChatModel(api, options)` constructor or the `spring.ai.mistralai.chat.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default model and temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", MistralAiChatOptions.builder() @@ -134,7 +134,7 @@ TIP: In addition to the model specific link:https://github.com/spring-projects/s == Function Calling -You can register custom Java functions with the MistralAiChatClient and have the Mistral AI model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. +You can register custom Java functions with the MistralAiChatModel and have the Mistral AI model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. This is a powerful technique to connect the LLM capabilities with external tools and APIs. Read more about xref:api/chat/functions/mistralai-chat-functions.adoc[Mistral AI Function Calling]. @@ -142,7 +142,7 @@ Read more about xref:api/chat/functions/mistralai-chat-functions.adoc[Mistral AI https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-mistral-ai-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the OpenAi Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the OpenAi chat model: [source,application.properties] ---- @@ -153,37 +153,37 @@ spring.ai.mistralai.chat.options.temperature=0.7 TIP: replace the `api-key` with your OpenAI credentials. -This will create a `MistralAiChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `MistralAiChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class ChatController { - private final MistralAiChatClient chatClient; + private final MistralAiChatModel chatModel; @Autowired - public ChatController(MistralAiChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(MistralAiChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } @GetMapping("/ai/generateStream") public Flux generateStream(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { var prompt = new Prompt(new UserMessage(message)); - return chatClient.stream(prompt); + return chatModel.stream(prompt); } } ---- == Manual Configuration -The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatClient.java[MistralAiChatClient] implements the `ChatClient` and `StreamingChatClient` and uses the <> to connect to the MistralAI service. +The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java[MistralAiChatModel] implements the `ChatModel` and `StreamingChatModel` and uses the <> to connect to the MistralAI service. Add the `spring-ai-mistral-ai` dependency to your project's Maven `pom.xml` file: @@ -206,23 +206,23 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create a `MistralAiChatClient` and use it for text generations: +Next, create a `MistralAiChatModel` and use it for text generations: [source,java] ---- var mistralAiApi = new MistralAiApi(System.getenv("MISTRAL_AI_API_KEY")); -var chatClient = new MistralAiChatClient(mistralAiApi, MistralAiChatOptions.builder() +var chatModel = new MistralAiChatModel(mistralAiApi, MistralAiChatOptions.builder() .withModel(MistralAiApi.ChatModel.LARGE.getValue()) .withTemperature(0.4f) .withMaxToken(200) .build()); -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); // Or with streaming responses -Flux response = chatClient.stream( +Flux response = chatModel.stream( new Prompt("Generate the names of 5 famous pirates.")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/ollama-chat.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/ollama-chat.adoc index 0fb32b997..df8a98146 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/ollama-chat.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/ollama-chat.adoc @@ -1,7 +1,7 @@ = Ollama Chat With https://ollama.ai/[Ollama] you can run various Large Language Models (LLMs) locally and generate text from them. -Spring AI supports the Ollama text generation with `OllamaChatClient`. +Spring AI supports the Ollama text generation with `OllamaChatModel`. == Prerequisites @@ -52,16 +52,16 @@ The prefix `spring.ai.ollama` is the property prefix to configure the connection | spring.ai.ollama.base-url | Base URL where Ollama API server is running. | `http://localhost:11434` |==== -The prefix `spring.ai.ollama.chat.options` is the property prefix that configures the Ollama chat client . +The prefix `spring.ai.ollama.chat.options` is the property prefix that configures the Ollama chat model . It includes the Ollama request (advanced) parameters such as the `model`, `keep-alive`, and `format` as well as the Ollama model `options` properties. -Here are the advanced request parameter for the Ollama chat client: +Here are the advanced request parameter for the Ollama chat model: [cols="3,6,1"] |==== | Property | Description | Default -| spring.ai.ollama.chat.enabled | Enable Ollama chat client. | true +| spring.ai.ollama.chat.enabled | Enable Ollama chat model. | true | spring.ai.ollama.chat.options.model | The name of the https://github.com/ollama/ollama?tab=readme-ov-file#model-library[supported model] to use. | mistral | spring.ai.ollama.chat.options.format | The format to return a response in. Currently the only accepted value is `json` | - | spring.ai.ollama.chat.options.keep_alive | Controls how long the model will stay loaded into memory following the request | 5m @@ -110,14 +110,14 @@ TIP: All properties prefixed with `spring.ai.ollama.chat.options` can be overrid The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/api/OllamaOptions.java[OllamaOptions.java] provides model configurations, such as the model to use, the temperature, etc. -On start-up, the default options can be configured with the `OllamaChatClient(api, options)` constructor or the `spring.ai.ollama.chat.options.*` properties. +On start-up, the default options can be configured with the `OllamaChatModel(api, options)` constructor or the `spring.ai.ollama.chat.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default model and temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", OllamaOptions.create() @@ -141,7 +141,7 @@ The Ollama link:https://github.com/ollama/ollama/blob/main/docs/api.md#parameter Spring AI’s link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/Message.java[Message] interface facilitates multimodal AI models by introducing the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/Media.java[Media] type. This type encompasses data and details regarding media attachments in messages, utilizing Spring’s `org.springframework.util.MimeType` and a `java.lang.Object` for the raw media data. -Below is a straightforward code example excerpted from link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatClientMultimodalIT.java[OllamaChatClientMultimodalIT.java], illustrating the fusion of user text with an image. +Below is a straightforward code example excerpted from link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelMultimodalIT.java[OllamaChatModelMultimodalIT.java], illustrating the fusion of user text with an image. [source,java] ---- @@ -150,7 +150,7 @@ byte[] imageData = new ClassPathResource("/multimodal.test.png").getContentAsByt var userMessage = new UserMessage("Explain what do you see on this picture?", List.of(new Media(MimeTypeUtils.IMAGE_PNG, imageData))); -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt(List.of(userMessage), OllamaOptions.create().withModel("llava"))); logger.info(response.getResult().getOutput().getContent()); @@ -174,7 +174,7 @@ where fruits are being displayed, possibly for convenience or aesthetic purposes https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-ollama-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Ollama Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the Ollama chat model: [source,application.properties] ---- @@ -185,30 +185,30 @@ spring.ai.ollama.chat.options.temperature=0.7 TIP: replace the `base-url` with your Ollama server URL. -This will create a `OllamaChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `OllamaChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class ChatController { - private final OllamaChatClient chatClient; + private final OllamaChatModel chatModel; @Autowired - public ChatController(OllamaChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(OllamaChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } @GetMapping("/ai/generateStream") public Flux generateStream(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { Prompt prompt = new Prompt(new UserMessage(message)); - return chatClient.stream(prompt); + return chatModel.stream(prompt); } } @@ -216,8 +216,8 @@ public class ChatController { == Manual Configuration -If you don't want to use the Spring Boot auto-configuration, you can manually configure the `OllamaChatClient` in your application. -The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatClient.java[OllamaChatClient] implements the `ChatClient` and `StreamingChatClient` and uses the <> to connect to the Ollama service. +If you don't want to use the Spring Boot auto-configuration, you can manually configure the `OllamaChatModel` in your application. +The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-ollama/src/main/java/org/springframework/ai/ollama/OllamaChatModel.java[OllamaChatModel] implements the `ChatModel` and `StreamingChatModel` and uses the <> to connect to the Ollama service. To use it add the `spring-ai-ollama` dependency to your project's Maven `pom.xml` file: @@ -240,25 +240,25 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -TIP: The `spring-ai-ollama` dependency provides access also to the `OllamaEmbeddingClient`. -For more information about the `OllamaEmbeddingClient` refer to the link:../embeddings/ollama-embeddings.html[Ollama Embedding Client] section. +TIP: The `spring-ai-ollama` dependency provides access also to the `OllamaEmbeddingModel`. +For more information about the `OllamaEmbeddingModel` refer to the link:../embeddings/ollama-embeddings.html[Ollama Embedding Client] section. -Next, create an `OllamaChatClient` instance and use it to text generations requests: +Next, create an `OllamaChatModel` instance and use it to text generations requests: [source,java] ---- var ollamaApi = new OllamaApi(); -var chatClient = new OllamaChatClient(ollamaApi, +var chatModel = new OllamaChatModel(ollamaApi, OllamaOptions.create() .withModel(OllamaOptions.DEFAULT_MODEL) .withTemperature(0.9f)); -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); // Or with streaming responses -Flux response = chatClient.stream( +Flux response = chatModel.stream( new Prompt("Generate the names of 5 famous pirates.")); ---- @@ -274,7 +274,7 @@ image::ollama-chat-completion-api.jpg[OllamaApi Chat Completion API Diagram, 800 Here is a simple snippet showing how to use the API programmatically: -NOTE: The `OllamaApi` is low level api and is not recommended for direct use. Use the `OllamaChatClient` instead. +NOTE: The `OllamaApi` is low level api and is not recommended for direct use. Use the `OllamaChatModel` instead. [source,java] ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/openai-chat.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/openai-chat.adoc index 440a4a6b6..5155c388e 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/openai-chat.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/openai-chat.adoc @@ -51,7 +51,7 @@ TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Man ==== Retry Properties -The prefix `spring.ai.retry` is used as the property prefix that lets you configure the retry mechanism for the OpenAI Chat client. +The prefix `spring.ai.retry` is used as the property prefix that lets you configure the retry mechanism for the OpenAI chat model. [cols="3,5,1"] |==== @@ -81,13 +81,13 @@ The prefix `spring.ai.openai` is used as the property prefix that lets you conne ==== Configuration Properties -The prefix `spring.ai.openai.chat` is the property prefix that lets you configure the chat client implementation for OpenAI. +The prefix `spring.ai.openai.chat` is the property prefix that lets you configure the chat model implementation for OpenAI. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.openai.chat.enabled | Enable OpenAI chat client. | true +| spring.ai.openai.chat.enabled | Enable OpenAI chat model. | true | spring.ai.openai.chat.base-url | Optional overrides the spring.ai.openai.base-url to provide chat specific url | - | spring.ai.openai.chat.api-key | Optional overrides the spring.ai.openai.api-key to provide chat specific api-key | - | spring.ai.openai.chat.options.model | This is the OpenAI Chat model to use | `gpt-3.5-turbo` (the `gpt-3.5-turbo`, `gpt-4`, and `gpt-4-32k` point to the latest model versions) @@ -107,7 +107,7 @@ The prefix `spring.ai.openai.chat` is the property prefix that lets you configur | spring.ai.openai.chat.options.functions | List of functions, identified by their names, to enable for function calling in a single prompt requests. Functions with those names must exist in the functionCallbacks registry. | - |==== -NOTE: You can override the common `spring.ai.openai.base-url` and `spring.ai.openai.api-key` for the `ChatClient` and `EmbeddingClient` implementations. +NOTE: You can override the common `spring.ai.openai.base-url` and `spring.ai.openai.api-key` for the `ChatModel` and `EmbeddingModel` implementations. The `spring.ai.openai.chat.base-url` and `spring.ai.openai.chat.api-key` properties if set take precedence over the common properties. This is useful if you want to use different OpenAI accounts for different models and different model endpoints. @@ -117,14 +117,14 @@ TIP: All properties prefixed with `spring.ai.openai.chat.options` can be overrid The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatOptions.java[OpenAiChatOptions.java] provides model configurations, such as the model to use, the temperature, the frequency penalty, etc. -On start-up, the default options can be configured with the `OpenAiChatClient(api, options)` constructor or the `spring.ai.openai.chat.options.*` properties. +On start-up, the default options can be configured with the `OpenAiChatModel(api, options)` constructor or the `spring.ai.openai.chat.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default model and temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", OpenAiChatOptions.builder() @@ -138,7 +138,7 @@ TIP: In addition to the model specific https://github.com/spring-projects/spring == Function Calling -You can register custom Java functions with the OpenAiChatClient and have the OpenAI model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. +You can register custom Java functions with the OpenAiChatModel and have the OpenAI model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. This is a powerful technique to connect the LLM capabilities with external tools and APIs. Read more about xref:api/chat/functions/openai-chat-functions.adoc[OpenAI Function Calling]. @@ -152,7 +152,7 @@ The OpenAI link:https://platform.openai.com/docs/api-reference/chat/create#chat- Spring AI’s link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/Message.java[Message] interface facilitates multimodal AI models by introducing the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/Media.java[Media] type. This type encompasses data and details regarding media attachments in messages, utilizing Spring’s `org.springframework.util.MimeType` and a `java.lang.Object` for the raw media data. -Below is a code example excerpted from link:https://github.com/spring-projects/spring-ai/blob/b3cfa2b900ea785e055e4ff71086eeb52f6578a3/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java[OpenAiChatClientIT.java], illustrating the fusion of user text with an image using the the `GPT_4_VISION_PREVIEW` model. +Below is a code example excerpted from link:https://github.com/spring-projects/spring-ai/blob/b3cfa2b900ea785e055e4ff71086eeb52f6578a3/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelIT.java[OpenAiChatModelIT.java], illustrating the fusion of user text with an image using the the `GPT_4_VISION_PREVIEW` model. [source,java] ---- @@ -161,7 +161,7 @@ byte[] imageData = new ClassPathResource("/multimodal.test.png").getContentAsByt var userMessage = new UserMessage("Explain what do you see on this picture?", List.of(new Media(MimeTypeUtils.IMAGE_PNG, imageData))); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withModel(OpenAiApi.ChatModel.GPT_4_VISION_PREVIEW.getValue()).build())); ---- @@ -173,7 +173,7 @@ var userMessage = new UserMessage("Explain what do you see on this picture?", List.of(new Media(MimeTypeUtils.IMAGE_PNG, "https://docs.spring.io/spring-ai/reference/1.0-SNAPSHOT/_images/multimodal.test.png"))); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withModel(OpenAiApi.ChatModel.GPT_4_O.getValue()).build())); ---- @@ -198,7 +198,7 @@ view of the fruit inside. https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-openai-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the OpenAi Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the OpenAi chat model: [source,application.properties] ---- @@ -209,37 +209,37 @@ spring.ai.openai.chat.options.temperature=0.7 TIP: replace the `api-key` with your OpenAI credentials. -This will create a `OpenAiChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `OpenAiChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class ChatController { - private final OpenAiChatClient chatClient; + private final OpenAiChatModel chatModel; @Autowired - public ChatController(OpenAiChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(OpenAiChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } @GetMapping("/ai/generateStream") public Flux generateStream(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { Prompt prompt = new Prompt(new UserMessage(message)); - return chatClient.stream(prompt); + return chatModel.stream(prompt); } } ---- == Manual Configuration -The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatClient.java[OpenAiChatClient] implements the `ChatClient` and `StreamingChatClient` and uses the <> to connect to the OpenAI service. +The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java[OpenAiChatModel] implements the `ChatModel` and `StreamingChatModel` and uses the <> to connect to the OpenAI service. Add the `spring-ai-openai` dependency to your project's Maven `pom.xml` file: @@ -262,7 +262,7 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create a `OpenAiChatClient` and use it for text generations: +Next, create a `OpenAiChatModel` and use it for text generations: [source,java] ---- @@ -272,14 +272,14 @@ var openAiChatOptions = OpenAiChatOptions.builder() .withTemperature(0.4) .withMaxTokens(200) .build(); -var chatClient = new OpenAiChatClient(openAiApi, openAiChatOptions) +var chatModel = new OpenAiChatModel(openAiApi, openAiChatOptions) -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); // Or with streaming responses -Flux response = chatClient.stream( +Flux response = chatModel.stream( new Prompt("Generate the names of 5 famous pirates.")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/vertexai-gemini-chat.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/vertexai-gemini-chat.adoc index ecf11add0..2f362a6a9 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/vertexai-gemini-chat.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/vertexai-gemini-chat.adoc @@ -59,7 +59,7 @@ The prefix `spring.ai.vertex.ai.gemini` is used as the property prefix that lets | spring.ai.vertex.ai.gemini.transport | API transport. GRPC or REST. | GRPC |==== -The prefix `spring.ai.vertex.ai.gemini.chat` is the property prefix that lets you configure the chat client implementation for VertexAI Gemini Chat. +The prefix `spring.ai.vertex.ai.gemini.chat` is the property prefix that lets you configure the chat model implementation for VertexAI Gemini Chat. [cols="3,5,1"] |==== @@ -84,14 +84,14 @@ TIP: All properties prefixed with `spring.ai.vertex.ai.gemini.chat.options` can The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatOptions.java[VertexAiGeminiChatOptions.java] provides model configurations, such as the temperature, the topK, etc. -On start-up, the default options can be configured with the `VertexAiGeminiChatClient(api, options)` constructor or the `spring.ai.vertex.ai.chat.options.*` properties. +On start-up, the default options can be configured with the `VertexAiGeminiChatModel(api, options)` constructor or the `spring.ai.vertex.ai.chat.options.*` properties. At runtime you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", VertexAiPaLm2ChatOptions.builder() @@ -109,7 +109,7 @@ WARNING: As of 30th of April 2023, the Vertex AI `Gemini Pro` model has signific Apparently the Gemini Pro can not handle anymore the function name correctly. The parallel function calling is gone as well. -You can register custom Java functions with the VertexAiGeminiChatClient and have the Gemini Pro model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. +You can register custom Java functions with the VertexAiGeminiChatModel and have the Gemini Pro model intelligently choose to output a JSON object containing arguments to call one or many of the registered functions. This is a powerful technique to connect the LLM capabilities with external tools and APIs. Read more about xref:api/chat/functions/vertexai-gemini-chat-functions.adoc[Vertex AI Gemini Function Calling]. @@ -122,7 +122,7 @@ Google's Gemini AI models support this capability by comprehending and integrati Spring AI's `Message` interface supports multimodal AI models by introducing the Media type. This type contains data and information about media attachments in messages, using Spring's `org.springframework.util.MimeType` and a `java.lang.Object` for the raw media data. -Below is a simple code example extracted from https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClientIT.java[VertexAiGeminiChatClientIT.java], demonstrating the combination of user text with an image. +Below is a simple code example extracted from https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModelIT.java[VertexAiGeminiChatModelIT.java], demonstrating the combination of user text with an image. [source,java] @@ -132,14 +132,14 @@ byte[] data = new ClassPathResource("/vertex-test.png").getContentAsByteArray(); var userMessage = new UserMessage("Explain what do you see o this picture?", List.of(new Media(MimeTypeUtils.IMAGE_PNG, data))); -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage))); +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage))); ---- == Sample Controller https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-vertex-ai-palm2-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the VertexAi Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the VertexAi chat model: [source,application.properties] ---- @@ -151,37 +151,37 @@ spring.ai.vertex.ai.gemini.chat.options.temperature=0.5 TIP: replace the `api-key` with your VertexAI credentials. -This will create a `VertexAiGeminiChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `VertexAiGeminiChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class ChatController { - private final VertexAiGeminiChatClient chatClient; + private final VertexAiGeminiChatModel chatModel; @Autowired - public ChatController(VertexAiGeminiChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(VertexAiGeminiChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } @GetMapping("/ai/generateStream") public Flux generateStream(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { Prompt prompt = new Prompt(new UserMessage(message)); - return chatClient.stream(prompt); + return chatModel.stream(prompt); } } ---- == Manual Configuration -The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClient.java[VertexAiGeminiChatClient] implements the `ChatClient` and uses the `VertexAI` to connect to the Vertex AI Gemini service. +The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java[VertexAiGeminiChatModel] implements the `ChatModel` and uses the `VertexAI` to connect to the Vertex AI Gemini service. Add the `spring-ai-vertex-ai-gemini` dependency to your project's Maven `pom.xml` file: @@ -204,19 +204,19 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create a `VertexAiGeminiChatClient` and use it for text generations: +Next, create a `VertexAiGeminiChatModel` and use it for text generations: [source,java] ---- VertexAI vertexApi = new VertexAI(projectId, location); -var chatClient = new VertexAiGeminiChatClient(vertexApi, +var chatModel = new VertexAiGeminiChatModel(vertexApi, VertexAiGeminiChatOptions.builder() .withModel(ChatModel.GEMINI_PRO_1_5_PRO) .withTemperature(0.4) .build()); -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/vertexai-palm2-chat.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/vertexai-palm2-chat.adoc index 045df41e2..4b928067e 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/vertexai-palm2-chat.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/vertexai-palm2-chat.adoc @@ -61,13 +61,13 @@ The prefix `spring.ai.vertex.ai` is used as the property prefix that lets you co | spring.ai.vertex.ai.api-key | The API Key | - |==== -The prefix `spring.ai.vertex.ai.chat` is the property prefix that lets you configure the chat client implementation for VertexAI Chat. +The prefix `spring.ai.vertex.ai.chat` is the property prefix that lets you configure the chat model implementation for VertexAI Chat. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.vertex.ai.chat.enabled | Enable Vertex AI PaLM API Chat client. | true +| spring.ai.vertex.ai.chat.enabled | Enable Vertex AI PaLM API chat model. | true | spring.ai.vertex.ai.chat.model | This is the https://cloud.google.com/vertex-ai/docs/generative-ai/model-reference/text-chat[Vertex Chat model] to use | chat-bison-001 | spring.ai.vertex.ai.chat.options.temperature | Controls the randomness of the output. Values can range over [0.0,1.0], inclusive. A value closer to 1.0 will produce responses that are more varied, while a value closer to 0.0 will typically result in less surprising responses from the generative. This value specifies default to be used by the backend while making the call to the generative. | 0.7 | spring.ai.vertex.ai.chat.options.topK | The maximum number of tokens to consider when sampling. The generative uses combined Top-k and nucleus sampling. Top-k sampling considers the set of topK most probable tokens. | 20 @@ -81,14 +81,14 @@ TIP: All properties prefixed with `spring.ai.vertex.ai.chat.options` can be over The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2ChatOptions.java[VertexAiPaLm2ChatOptions.java] provides model configurations, such as the temperature, the topK, etc. -On start-up, the default options can be configured with the `VertexAiPaLm2ChatClient(api, options)` constructor or the `spring.ai.vertex.ai.chat.options.*` properties. +On start-up, the default options can be configured with the `VertexAiPaLm2ChatModel(api, options)` constructor or the `spring.ai.vertex.ai.chat.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", VertexAiPaLm2ChatOptions.builder() @@ -103,7 +103,7 @@ TIP: In addition to the model specific `VertexAiPaLm2ChatOptions` you can use a https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-vertex-ai-palm2-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the VertexAi Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the VertexAi chat model: [source,application.properties] ---- @@ -114,37 +114,37 @@ spring.ai.vertex.ai.chat.options.temperature=0.5 TIP: replace the `api-key` with your VertexAI credentials. -This will create a `VertexAiPaLm2ChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `VertexAiPaLm2ChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class ChatController { - private final VertexAiPaLm2ChatClient chatClient; + private final VertexAiPaLm2ChatModel chatModel; @Autowired - public ChatController(VertexAiPaLm2ChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(VertexAiPaLm2ChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } @GetMapping("/ai/generateStream") public Flux generateStream(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { Prompt prompt = new Prompt(new UserMessage(message)); - return chatClient.stream(prompt); + return chatModel.stream(prompt); } } ---- == Manual Configuration -The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/main/java/org/springframework/ai/vertexai/paml2/VertexAiPaLm2ChatClient.java[VertexAiPaLm2ChatClient] implements the `ChatClient` and uses the <> to connect to the VertexAI service. +The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/main/java/org/springframework/ai/vertexai/paml2/VertexAiPaLm2ChatModel.java[VertexAiPaLm2ChatModel] implements the `ChatModel` and uses the <> to connect to the VertexAI service. Add the `spring-ai-vertex-ai-palm2` dependency to your project's Maven `pom.xml` file: @@ -167,18 +167,18 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create a `VertexAiPaLm2ChatClient` and use it for text generations: +Next, create a `VertexAiPaLm2ChatModel` and use it for text generations: [source,java] ---- VertexAiPaLm2Api vertexAiApi = new VertexAiPaLm2Api(< YOUR PALM_API_KEY>); -var chatClient = new VertexAiPaLm2ChatClient(vertexAiApi, +var chatModel = new VertexAiPaLm2ChatModel(vertexAiApi, VertexAiPaLm2ChatOptions.builder() .withTemperature(0.4) .build()); -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/watsonx-ai-chat.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/watsonx-ai-chat.adoc index 79ce3d8d9..4b085b72b 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/watsonx-ai-chat.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/watsonx-ai-chat.adoc @@ -1,7 +1,7 @@ = watsonx.ai Chat With https://dataplatform.cloud.ibm.com/docs/content/wsj/getting-started/overview-wx.html?context=wx&audience=wdp[watsonx.ai] you can run various Large Language Models (LLMs) locally and generate text from them. -Spring AI supports the watsonx.ai text generation with `WatsonxAiChatClient`. +Spring AI supports the watsonx.ai text generation with `WatsonxAiChatModel`. == Prerequisites @@ -53,13 +53,13 @@ The prefix `spring.ai.watsonx.ai` is used as the property prefix that lets you c ==== Configuration Properties -The prefix `spring.ai.watsonx.ai.chat` is the property prefix that lets you configure the chat client implementation for Watsonx.AI. +The prefix `spring.ai.watsonx.ai.chat` is the property prefix that lets you configure the chat model implementation for Watsonx.AI. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.watsonx.ai.chat.enabled | Enable Watsonx.AI chat client. | true +| spring.ai.watsonx.ai.chat.enabled | Enable Watsonx.AI chat model. | true | spring.ai.watsonx.ai.chat.options.temperature | The temperature of the model. Increasing the temperature will make the model answer more creatively. | 0.7 | spring.ai.watsonx.ai.chat.options.top-p | Works together with top-k. A higher value (e.g., 0.95) will lead to more diverse text, while a lower value (e.g., 0.2) will generate more focused and conservative text. | 1.0 | spring.ai.watsonx.ai.chat.options.top-k | Reduces the probability of generating nonsense. A higher value (e.g. 100) will give more diverse answers, while a lower value (e.g. 10) will be more conservative. | 50 @@ -76,14 +76,14 @@ The prefix `spring.ai.watsonx.ai.chat` is the property prefix that lets you conf The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-watsonx-ai/src/main/java/org/springframework/ai/watsonx/WatsonxAiChatOptions.java[WatsonxAiChatOptions.java] provides model configurations, such as the model to use, the temperature, the frequency penalty, etc. -On start-up, the default options can be configured with the `WatsonxAiChatClient(api, options)` constructor or the `spring.ai.watsonxai.chat.options.*` properties. +On start-up, the default options can be configured with the `WatsonxAiChatModel(api, options)` constructor or the `spring.ai.watsonxai.chat.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default model and temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", WatsonxAiChatOptions.builder() @@ -103,11 +103,11 @@ NOTE: For more information go to https://dataplatform.cloud.ibm.com/docs/content public class MyClass { private final static String MODEL = "google/flan-ul2"; - private final WatsonxAiChatClient chat; + private final WatsonxAiChatModel chatModel; @Autowired - MyClass(WatsonxAiChatClient chat) { - this.chat = chat; + MyClass(WatsonxAiChatModel chatModel) { + this.chatModel = chatModel; } public String generate(String userInput) { @@ -119,7 +119,7 @@ public class MyClass { Prompt prompt = new Prompt(new SystemMessage(userInput), options); - var results = chat.call(prompt); + var results = chatModel.call(prompt); var generatedText = results.getResult().getOutput().getContent(); @@ -135,7 +135,7 @@ public class MyClass { Prompt prompt = new Prompt(new SystemMessage(userInput), options); - var results = chat.stream(prompt).collectList().block(); // wait till the stream is resolved (completed) + var results = chatModel.stream(prompt).collectList().block(); // wait till the stream is resolved (completed) var generatedText = results.stream() .map(generation -> generation.getResult().getOutput().getContent()) diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/zhipuai-chat.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/zhipuai-chat.adoc index a88b2526e..fd4cf5a03 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/zhipuai-chat.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chat/zhipuai-chat.adoc @@ -52,7 +52,7 @@ TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Man ==== Retry Properties -The prefix `spring.ai.retry` is used as the property prefix that lets you configure the retry mechanism for the ZhiPu AI Chat client. +The prefix `spring.ai.retry` is used as the property prefix that lets you configure the retry mechanism for the ZhiPu AI chat model. [cols="3,5,1"] |==== @@ -81,13 +81,13 @@ The prefix `spring.ai.zhiPu` is used as the property prefix that lets you connec ==== Configuration Properties -The prefix `spring.ai.zhipuai.chat` is the property prefix that lets you configure the chat client implementation for ZhiPuAI. +The prefix `spring.ai.zhipuai.chat` is the property prefix that lets you configure the chat model implementation for ZhiPuAI. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.zhipuai.chat.enabled | Enable ZhiPuAI chat client. | true +| spring.ai.zhipuai.chat.enabled | Enable ZhiPuAI chat model. | true | spring.ai.zhipuai.chat.base-url | Optional overrides the spring.ai.zhipuai.base-url to provide chat specific url | https://open.bigmodel.cn/api/paas | spring.ai.zhipuai.chat.api-key | Optional overrides the spring.ai.zhipuai.api-key to provide chat specific api-key | - | spring.ai.zhipuai.chat.options.model | This is the ZhiPuAI Chat model to use | `GLM-3-Turbo` (the `GLM-3-Turbo`, `GLM-4`, and `GLM-4V` point to the latest model versions) @@ -101,7 +101,7 @@ The prefix `spring.ai.zhipuai.chat` is the property prefix that lets you configu | spring.ai.zhipuai.chat.options.user | A unique identifier representing your end-user, which can help ZhiPuAI to monitor and detect abuse. | - |==== -NOTE: You can override the common `spring.ai.zhipuai.base-url` and `spring.ai.zhipuai.api-key` for the `ChatClient` implementations. +NOTE: You can override the common `spring.ai.zhipuai.base-url` and `spring.ai.zhipuai.api-key` for the `ChatModel` implementations. The `spring.ai.zhipuai.chat.base-url` and `spring.ai.zhipuai.chat.api-key` properties if set take precedence over the common properties. This is useful if you want to use different ZhiPuAI accounts for different models and different model endpoints. @@ -111,14 +111,14 @@ TIP: All properties prefixed with `spring.ai.zhipuai.chat.options` can be overri The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatOptions.java[ZhiPuAiChatOptions.java] provides model configurations, such as the model to use, the temperature, the frequency penalty, etc. -On start-up, the default options can be configured with the `ZhiPuAiChatClient(api, options)` constructor or the `spring.ai.zhipuai.chat.options.*` properties. +On start-up, the default options can be configured with the `ZhiPuAiChatModel(api, options)` constructor or the `spring.ai.zhipuai.chat.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `Prompt` call. For example to override the default model and temperature for a specific request: [source,java] ---- -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt( "Generate the names of 5 famous pirates.", ZhiPuAiChatOptions.builder() @@ -134,7 +134,7 @@ TIP: In addition to the model specific link:https://github.com/spring-projects/s https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-zhipuai-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the ZhiPuAi Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the ZhiPuAi chat model: [source,application.properties] ---- @@ -145,37 +145,37 @@ spring.ai.zhipuai.chat.options.temperature=0.7 TIP: replace the `api-key` with your ZhiPuAI credentials. -This will create a `ZhiPuAiChatClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `ZhiPuAiChatModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class ChatController { - private final ZhiPuAiChatClient chatClient; + private final ZhiPuAiChatModel chatModel; @Autowired - public ChatController(ZhiPuAiChatClient chatClient) { - this.chatClient = chatClient; + public ChatController(ZhiPuAiChatModel chatModel) { + this.chatModel = chatModel; } @GetMapping("/ai/generate") public Map generate(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - return Map.of("generation", chatClient.call(message)); + return Map.of("generation", chatModel.call(message)); } @GetMapping("/ai/generateStream") public Flux generateStream(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { var prompt = new Prompt(new UserMessage(message)); - return chatClient.stream(prompt); + return chatModel.stream(prompt); } } ---- == Manual Configuration -The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatClient.java[ZhiPuAiChatClient] implements the `ChatClient` and `StreamingChatClient` and uses the <> to connect to the ZhiPuAI service. +The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiChatModel.java[ZhiPuAiChatModel] implements the `ChatModel` and `StreamingChatModel` and uses the <> to connect to the ZhiPuAI service. Add the `spring-ai-zhipuai` dependency to your project's Maven `pom.xml` file: @@ -198,23 +198,23 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create a `ZhiPuAiChatClient` and use it for text generations: +Next, create a `ZhiPuAiChatModel` and use it for text generations: [source,java] ---- var zhiPuAiApi = new ZhiPuAiApi(System.getenv("ZHIPU_AI_API_KEY")); -var chatClient = new ZhiPuAiChatClient(zhiPuAiApi, ZhiPuAiChatOptions.builder() +var chatModel = new ZhiPuAiChatModel(zhiPuAiApi, ZhiPuAiChatOptions.builder() .withModel(ZhiPuAiApi.ChatModel.GLM_3_Turbo.getValue()) .withTemperature(0.4f) .withMaxTokens(200) .build()); -ChatResponse response = chatClient.call( +ChatResponse response = chatModel.call( new Prompt("Generate the names of 5 famous pirates.")); // Or with streaming responses -Flux streamResponse = chatClient.stream( +Flux streamResponse = chatModel.stream( new Prompt("Generate the names of 5 famous pirates.")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chatclient.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chatclient.adoc index 42f7857a5..1dfc16309 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chatclient.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/chatclient.adoc @@ -1,27 +1,27 @@ -[[ChatClient]] -= Chat Completion API +[[ChatModel]] += Chat Model API -The Chat Completion API offers developers the ability to integrate AI-powered chat completion capabilities into their applications. It leverages pre-trained language models, such as GPT (Generative Pre-trained Transformer), to generate human-like responses to user inputs in natural language. +The Chat Model API offers developers the ability to integrate AI-powered chat completion capabilities into their applications. It leverages pre-trained language models, such as GPT (Generative Pre-trained Transformer), to generate human-like responses to user inputs in natural language. The API typically works by sending a prompt or partial conversation to the AI model, which then generates a completion or continuation of the conversation based on its training data and understanding of natural language patterns. The completed response is then returned to the application, which can present it to the user or use it for further processing. -The `Spring AI Chat Completion API` is designed to be a simple and portable interface for interacting with various xref:concepts.adoc#_models[AI Models], allowing developers to switch between different models with minimal code changes. +The `Spring AI Chat Model API` is designed to be a simple and portable interface for interacting with various xref:concepts.adoc#_models[AI Models], allowing developers to switch between different models with minimal code changes. This design aligns with Spring's philosophy of modularity and interchangeability. -Also with the help of companion classes like `Prompt` for input encapsulation and `ChatResponse` for output handling, the Chat Completion API unifies the communication with AI Models. +Also with the help of companion classes like `Prompt` for input encapsulation and `ChatResponse` for output handling, the Chat Model API unifies the communication with AI Models. It manages the complexity of request preparation and response parsing, offering a direct and simplified API interaction. == API Overview -This section provides a guide to the Spring AI Chat Completion API interface and associated classes. +This section provides a guide to the Spring AI Chat Model API interface and associated classes. -=== ChatClient +=== ChatModel -Here is the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatClient.java[ChatClient] interface definition: +Here is the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/chat/ChatModel.java[ChatModel] interface definition: [source,java] ---- -public interface ChatClient extends ModelClient { +public interface ChatModel extends Model { default String call(String message) {// implementation omitted } @@ -35,19 +35,19 @@ public interface ChatClient extends ModelClient { The `call` method with a `String` parameter simplifies initial use, avoiding the complexities of the more sophisticated `Prompt` and `ChatResponse` classes. In real-world applications, it is more common to use the `call` method that takes a `Prompt` instance and returns an `ChatResponse`. -=== StreamingChatClient +=== StreamingChatModel -Here is the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/chat/StreamingChatClient.java[StreamingChatClient] interface definition: +Here is the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/chat/StreamingChatModel.java[StreamingChatModel] interface definition: [source,java] ---- -public interface StreamingChatClient extends StreamingModelClient { +public interface StreamingChatModel extends StreamingModel { @Override Flux stream(Prompt prompt); } ---- -The `stream` method takes a `Prompt` request similar to `ChatClient` but it streams the responses using the reactive Flux API. +The `stream` method takes a `Prompt` request similar to `ChatModel` but it streams the responses using the reactive Flux API. === Prompt @@ -130,7 +130,7 @@ public interface ChatOptions extends ModelOptions { } ---- -Additionally, every model specific ChatClient/StreamingChatClient implementation can have its own options that can be passed to the AI model. For example, the OpenAI Chat Completion model has its own options like `presencePenalty`, `frequencyPenalty`, `bestOf` etc. +Additionally, every model specific ChatModel/StreamingChatModel implementation can have its own options that can be passed to the AI model. For example, the OpenAI Chat Completion model has its own options like `presencePenalty`, `frequencyPenalty`, `bestOf` etc. This is a powerful feature that allows developers to use model specific options when starting the application and then override them with at runtime using the Prompt request: @@ -184,7 +184,7 @@ public class Generation implements ModelResult { == Available Implementations -The `ChatClient` and `StreamingChatClient` implementations are provided for the following Model providers: +The `ChatModel` and `StreamingChatModel` implementations are provided for the following Model providers: image::spring-ai-chat-completions-clients.jpg[align="center", width="800px"] @@ -205,7 +205,7 @@ image::spring-ai-chat-completions-clients.jpg[align="center", width="800px"] == Chat Model API -The Spring AI Chat Completion API is build on top of the Spring AI `Generic Model API` providing Chat specific abstractions and implementations. Following class diagram illustrates the main classes and interfaces of the Spring AI Chat Completion API. +The Spring AI Chat Model API is build on top of the Spring AI `Generic Model API` providing Chat specific abstractions and implementations. Following class diagram illustrates the main classes and interfaces of the Spring AI Chat Model API. image::spring-ai-chat-api.jpg[align="center", width="900px"] diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings.adoc index 5bd3ce6c7..d6140e23d 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings.adoc @@ -1,23 +1,23 @@ -[[EmbeddingClient]] -= Embeddings API +[[EmbeddingModel]] += Embeddings Model API -The `EmbeddingClient` interface is designed for straightforward integration with embedding models in AI and machine learning. +The `EmbeddingModel` interface is designed for straightforward integration with embedding models in AI and machine learning. Its primary function is to convert text into numerical vectors, commonly referred to as embeddings. These embeddings are crucial for various tasks such as semantic analysis and text classification. -The design of the EmbeddingClient interface centers around two primary goals: +The design of the EmbeddingModel interface centers around two primary goals: * *Portability*: This interface ensures easy adaptability across various embedding models. It allows developers to switch between different embedding techniques or models with minimal code changes. This design aligns with Spring's philosophy of modularity and interchangeability. -* *Simplicity*: EmbeddingClient simplifies the process of converting text to embeddings. +* *Simplicity*: EmbeddingModel simplifies the process of converting text to embeddings. By providing straightforward methods like `embed(String text)` and `embed(Document document)`, it takes the complexity out of dealing with raw text data and embedding algorithms. This design choice makes it easier for developers, especially those new to AI, to utilize embeddings in their applications without delving deep into the underlying mechanics. == API Overview -The Embedding API is built on top of the generic https://github.com/spring-projects/spring-ai/tree/main/spring-ai-core/src/main/java/org/springframework/ai/model[Spring AI Model API], which is a part of the Spring AI library. -As such, the EmbeddingClient interface extends the `ModelClient` interface, which provides a standard set of methods for interacting with AI models. The `EmbeddingRequest` and `EmbeddingResponse` classes extend from the `ModelRequest` and `ModelResponse` are used to encapsulate the input and output of the embedding models, respectively. +The Embedding Model API is built on top of the generic https://github.com/spring-projects/spring-ai/tree/main/spring-ai-core/src/main/java/org/springframework/ai/model[Spring AI Model API], which is a part of the Spring AI library. +As such, the EmbeddingModel interface extends the `Model` interface, which provides a standard set of methods for interacting with AI models. The `EmbeddingRequest` and `EmbeddingResponse` classes extend from the `ModelRequest` and `ModelResponse` are used to encapsulate the input and output of the embedding models, respectively. The Embedding API in turn is used by higher-level components to implement Embedding Clients for specific embedding models, such as OpenAI, Titan, Azure OpenAI, Ollie, and others. @@ -25,13 +25,13 @@ Following diagram illustrates the Embedding API and its relationship with the Sp image:embeddings-api.jpg[title=Embeddings API,align=center,width=900] -=== EmbeddingClient +=== EmbeddingModel -This section provides a guide to the `EmbeddingClient` interface and associated classes. +This section provides a guide to the `EmbeddingModel` interface and associated classes. [source,java] ---- -public interface EmbeddingClient extends ModelClient { +public interface EmbeddingModel extends Model { @Override EmbeddingResponse call(EmbeddingRequest request); @@ -148,7 +148,7 @@ public class Embedding implements ModelResult> { == Available Implementations [[available-implementations]] -Internally the various `EmbeddingClient` implementations use different low-level libraries and APIs to perform the embedding tasks. The following are some of the available implementations of the `EmbeddingClient` implementations: +Internally the various `EmbeddingModel` implementations use different low-level libraries and APIs to perform the embedding tasks. The following are some of the available implementations of the `EmbeddingModel` implementations: * xref:api/embeddings/openai-embeddings.adoc[Spring AI OpenAI Embeddings] * xref:api/embeddings/azure-openai-embeddings.adoc[Spring AI Azure OpenAI Embeddings] diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/azure-openai-embeddings.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/azure-openai-embeddings.adoc index 402e6f0c2..53ae973dc 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/azure-openai-embeddings.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/azure-openai-embeddings.adoc @@ -66,13 +66,13 @@ The prefix `spring.ai.azure.openai` is the property prefix to configure the conn |==== -The prefix `spring.ai.azure.openai.embedding` is the property prefix that configures the `EmbeddingClient` implementation for Azure OpenAI +The prefix `spring.ai.azure.openai.embedding` is the property prefix that configures the `EmbeddingModel` implementation for Azure OpenAI [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.azure.openai.embedding.enabled | Enable Azure OpenAI embedding client. | true +| spring.ai.azure.openai.embedding.enabled | Enable Azure OpenAI embedding model. | true | spring.ai.azure.openai.embedding.metadata-mode | Document content extraction mode | EMBED | spring.ai.azure.openai.embedding.options.deployment-name | This is the value of the 'Deployment Name' as presented in the Azure AI Portal | text-embedding-ada-002 | spring.ai.azure.openai.embedding.options.user | An identifier for the caller or end user of the operation. This may be used for tracking or rate-limiting purposes. | - @@ -85,14 +85,14 @@ TIP: All properties prefixed with `spring.ai.azure.openai.embedding.options` can The `AzureOpenAiEmbeddingOptions` provides the configuration information for the embedding requests. The `AzureOpenAiEmbeddingOptions` offers a builder to create the options. -At start time use the `AzureOpenAiEmbeddingClient` constructor to set the default options used for all embedding requests. +At start time use the `AzureOpenAiEmbeddingModel` constructor to set the default options used for all embedding requests. At run-time you can override the default options, by passing a `AzureOpenAiEmbeddingOptions` instance with your to the `EmbeddingRequest` request. For example to override the default model name for a specific request: [source,java] ---- -EmbeddingResponse embeddingResponse = embeddingClient.call( +EmbeddingResponse embeddingResponse = embeddingModel.call( new EmbeddingRequest(List.of("Hello World", "World is big and salvation is near"), AzureOpenAiEmbeddingOptions.builder() .withModel("Different-Embedding-Model-Deployment-Name") @@ -102,8 +102,8 @@ EmbeddingResponse embeddingResponse = embeddingClient.call( == Sample Code -This will create a `EmbeddingClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the `EmbeddingClient` implementation. +This will create a `EmbeddingModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the `EmbeddingModel` implementation. [source,application.properties] ---- @@ -117,16 +117,16 @@ spring.ai.azure.openai.embedding.options.model=text-embedding-ada-002 @RestController public class EmbeddingController { - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; @Autowired - public EmbeddingController(EmbeddingClient embeddingClient) { - this.embeddingClient = embeddingClient; + public EmbeddingController(EmbeddingModel embeddingModel) { + this.embeddingModel = embeddingModel; } @GetMapping("/ai/embedding") public Map embed(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - EmbeddingResponse embeddingResponse = this.embeddingClient.embedForResponse(List.of(message)); + EmbeddingResponse embeddingResponse = this.embeddingModel.embedForResponse(List.of(message)); return Map.of("embedding", embeddingResponse); } } @@ -134,7 +134,7 @@ public class EmbeddingController { == Manual Configuration -If you prefer not to use the Spring Boot auto-configuration, you can manually configure the `AzureOpenAiEmbeddingClient` in your application. +If you prefer not to use the Spring Boot auto-configuration, you can manually configure the `AzureOpenAiEmbeddingModel` in your application. For this add the `spring-ai-azure-openai` dependency to your project's Maven `pom.xml` file: [source, xml] ---- @@ -155,9 +155,9 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -NOTE: The `spring-ai-azure-openai` dependency also provide the access to the `AzureOpenAiEmbeddingClient`. For more information about the `AzureOpenAiChatClient` refer to the link:../embeddings/azure-openai-embeddings.html[Azure OpenAI Embeddings] section. +NOTE: The `spring-ai-azure-openai` dependency also provide the access to the `AzureOpenAiEmbeddingModel`. For more information about the `AzureOpenAiChatModel` refer to the link:../embeddings/azure-openai-embeddings.html[Azure OpenAI Embeddings] section. -Next, create an `AzureOpenAiEmbeddingClient` instance and use it to compute the similarity between two input texts: +Next, create an `AzureOpenAiEmbeddingModel` instance and use it to compute the similarity between two input texts: [source,java] ---- @@ -166,13 +166,13 @@ var openAIClient = OpenAIClientBuilder() .endpoint(System.getenv("AZURE_OPENAI_ENDPOINT")) .buildClient(); -var embeddingClient = new AzureOpenAiEmbeddingClient(openAIClient) +var embeddingModel = new AzureOpenAiEmbeddingModel(openAIClient) .withDefaultOptions(AzureOpenAiEmbeddingOptions.builder() .withModel("text-embedding-ada-002") .withUser("user-6") .build()); -EmbeddingResponse embeddingResponse = embeddingClient +EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/bedrock-cohere-embedding.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/bedrock-cohere-embedding.adoc index 2d596bb4b..e4dcee9b6 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/bedrock-cohere-embedding.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/bedrock-cohere-embedding.adoc @@ -62,7 +62,7 @@ The prefix `spring.ai.bedrock.aws` is the property prefix to configure the conne | spring.ai.bedrock.aws.secret-key | AWS secret key. | - |==== -The prefix `spring.ai.bedrock.cohere.embedding` (defined in `BedrockCohereEmbeddingProperties`) is the property prefix that configures the embedding client implementation for Cohere. +The prefix `spring.ai.bedrock.cohere.embedding` (defined in `BedrockCohereEmbeddingProperties`) is the property prefix that configures the embedding model implementation for Cohere. [cols="3,4,1"] |==== @@ -83,14 +83,14 @@ TIP: All properties prefixed with `spring.ai.bedrock.cohere.embedding.options` c The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingOptions.java[BedrockCohereEmbeddingOptions.java] provides model configurations, such as `input-type` or `truncate`. -On start-up, the default options can be configured with the `BedrockCohereEmbeddingClient(api, options)` constructor or the `spring.ai.bedrock.cohere.embedding.options.*` properties. +On start-up, the default options can be configured with the `BedrockCohereEmbeddingModel(api, options)` constructor or the `spring.ai.bedrock.cohere.embedding.options.*` properties. At run-time you can override the default options by adding new, request specific, options to the `EmbeddingRequest` call. For example to override the default temperature for a specific request: [source,java] ---- -EmbeddingResponse embeddingResponse = embeddingClient.call( +EmbeddingResponse embeddingResponse = embeddingModel.call( new EmbeddingRequest(List.of("Hello World", "World is big and salvation is near"), BedrockCohereEmbeddingOptions.builder() .withInputType(InputType.SEARCH_DOCUMENT) @@ -115,24 +115,24 @@ spring.ai.bedrock.cohere.embedding.options.input-type=search-document TIP: replace the `regions`, `access-key` and `secret-key` with your AWS credentials. -This will create a `BedrockCohereEmbeddingClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +This will create a `BedrockCohereEmbeddingModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class EmbeddingController { - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; @Autowired - public EmbeddingController(EmbeddingClient embeddingClient) { - this.embeddingClient = embeddingClient; + public EmbeddingController(EmbeddingModel embeddingModel) { + this.embeddingModel = embeddingModel; } @GetMapping("/ai/embedding") public Map embed(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - EmbeddingResponse embeddingResponse = this.embeddingClient.embedForResponse(List.of(message)); + EmbeddingResponse embeddingResponse = this.embeddingModel.embedForResponse(List.of(message)); return Map.of("embedding", embeddingResponse); } } @@ -140,7 +140,7 @@ public class EmbeddingController { == Manual Configuration -The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingClient.java[BedrockCohereEmbeddingClient] implements the `EmbeddingClient` and uses the <> to connect to the Bedrock Cohere service. +The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingModel.java[BedrockCohereEmbeddingModel] implements the `EmbeddingModel` and uses the <> to connect to the Bedrock Cohere service. Add the `spring-ai-bedrock` dependency to your project's Maven `pom.xml` file: @@ -163,7 +163,7 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingClient.java[BedrockCohereEmbeddingClient] and use it for text embeddings: +Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/cohere/BedrockCohereEmbeddingModel.java[BedrockCohereEmbeddingModel] and use it for text embeddings: [source,java] ---- @@ -172,9 +172,9 @@ var cohereEmbeddingApi =new CohereEmbeddingBedrockApi( EnvironmentVariableCredentialsProvider.create(), Region.US_EAST_1.id(), new ObjectMapper()); -var embeddingClient = new BedrockCohereEmbeddingClient(cohereEmbeddingApi); +var embeddingModel = new BedrockCohereEmbeddingModel(cohereEmbeddingApi); -EmbeddingResponse embeddingResponse = embeddingClient +EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/bedrock-titan-embedding.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/bedrock-titan-embedding.adoc index e1a5a6600..b93685458 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/bedrock-titan-embedding.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/bedrock-titan-embedding.adoc @@ -69,7 +69,7 @@ The prefix `spring.ai.bedrock.aws` is the property prefix to configure the conne | spring.ai.bedrock.aws.secret-key | AWS secret key. | - |==== -The prefix `spring.ai.bedrock.titan.embedding` (defined in `BedrockTitanEmbeddingProperties`) is the property prefix that configures the embedding client implementation for Titan. +The prefix `spring.ai.bedrock.titan.embedding` (defined in `BedrockTitanEmbeddingProperties`) is the property prefix that configures the embedding model implementation for Titan. [cols="3,4,1"] |==== @@ -84,14 +84,14 @@ Model ID values can also be found in the https://docs.aws.amazon.com/bedrock/lat == Runtime Options [[embedding-options]] The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingOptions.java[BedrockTitanEmbeddingOptions.java] provides model configurations, such as `input-type`. -On start-up, the default options can be configured with the `BedrockTitanEmbeddingClient(api).withInputType(type)` method or the `spring.ai.bedrock.titan.embedding.input-type` properties. +On start-up, the default options can be configured with the `BedrockTitanEmbeddingModel(api).withInputType(type)` method or the `spring.ai.bedrock.titan.embedding.input-type` properties. At run-time you can override the default options by adding new, request specific, options to the `EmbeddingRequest` call. For example to override the default temperature for a specific request: [source,java] ---- -EmbeddingResponse embeddingResponse = embeddingClient.call( +EmbeddingResponse embeddingResponse = embeddingModel.call( new EmbeddingRequest(List.of("Hello World", "World is big and salvation is near"), BedrockTitanEmbeddingOptions.builder() .withInputType(InputType.TEXT) @@ -116,23 +116,23 @@ spring.ai.bedrock.titan.embedding.enabled=true TIP: replace the `regions`, `access-key` and `secret-key` with your AWS credentials. This will create a `EmbeddingController` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the chat client for text generations. +Here is an example of a simple `@Controller` class that uses the chat model for text generations. [source,java] ---- @RestController public class EmbeddingController { - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; @Autowired - public EmbeddingController(EmbeddingClient embeddingClient) { - this.embeddingClient = embeddingClient; + public EmbeddingController(EmbeddingModel embeddingModel) { + this.embeddingModel = embeddingModel; } @GetMapping("/ai/embedding") public Map embed(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - EmbeddingResponse embeddingResponse = this.embeddingClient.embedForResponse(List.of(message)); + EmbeddingResponse embeddingResponse = this.embeddingModel.embedForResponse(List.of(message)); return Map.of("embedding", embeddingResponse); } } @@ -140,7 +140,7 @@ public class EmbeddingController { == Manual Configuration -The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingClient.java[BedrockTitanEmbeddingClient] implements the `EmbeddingClient` and uses the <> to connect to the Bedrock Titan service. +The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingModel.java[BedrockTitanEmbeddingModel] implements the `EmbeddingModel` and uses the <> to connect to the Bedrock Titan service. Add the `spring-ai-bedrock` dependency to your project's Maven `pom.xml` file: @@ -163,16 +163,16 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingClient.java[BedrockTitanEmbeddingClient] and use it for text embeddings: +Next, create an https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/titan/BedrockTitanEmbeddingModel.java[BedrockTitanEmbeddingModel] and use it for text embeddings: [source,java] ---- var titanEmbeddingApi = new TitanEmbeddingBedrockApi( TitanEmbeddingModel.TITAN_EMBED_IMAGE_V1.id(), Region.US_EAST_1.id()); -var embeddingClient = new BedrockTitanEmbeddingClient(titanEmbeddingApi); +var embeddingModel = new BedrockTitanEmbeddingModel(titanEmbeddingApi); -EmbeddingResponse embeddingResponse = embeddingClient +EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World")); // NOTE titan does not support batch embedding. ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/minimax-embeddings.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/minimax-embeddings.adoc index da452c69d..f554b4664 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/minimax-embeddings.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/minimax-embeddings.adoc @@ -81,19 +81,19 @@ The prefix `spring.ai.minimax` is used as the property prefix that lets you conn ==== Configuration Properties -The prefix `spring.ai.minimax.embedding` is property prefix that configures the `EmbeddingClient` implementation for MiniMax. +The prefix `spring.ai.minimax.embedding` is property prefix that configures the `EmbeddingModel` implementation for MiniMax. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.minimax.embedding.enabled | Enable MiniMax embedding client. | true +| spring.ai.minimax.embedding.enabled | Enable MiniMax embedding model. | true | spring.ai.minimax.embedding.base-url | Optional overrides the spring.ai.minimax.base-url to provide embedding specific url | - | spring.ai.minimax.embedding.api-key | Optional overrides the spring.ai.minimax.api-key to provide embedding specific api-key | - | spring.ai.minimax.embedding.options.model | The model to use | embo-01 |==== -NOTE: You can override the common `spring.ai.minimax.base-url` and `spring.ai.minimax.api-key` for the `ChatClient` and `EmbeddingClient` implementations. +NOTE: You can override the common `spring.ai.minimax.base-url` and `spring.ai.minimax.api-key` for the `ChatModel` and `EmbeddingModel` implementations. The `spring.ai.minimax.embedding.base-url` and `spring.ai.minimax.embedding.api-key` properties if set take precedence over the common properties. Similarly, the `spring.ai.minimax.embedding.base-url` and `spring.ai.minimax.embedding.api-key` properties if set take precedence over the common properties. This is useful if you want to use different MiniMax accounts for different models and different model endpoints. @@ -106,14 +106,14 @@ The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-mini The default options can be configured using the `spring.ai.minimax.embedding.options` properties as well. -At start-time use the `MiniMaxEmbeddingClient` constructor to set the default options used for all embedding requests. +At start-time use the `MiniMaxEmbeddingModel` constructor to set the default options used for all embedding requests. At run-time you can override the default options, using a `MiniMaxEmbeddingOptions` instance as part of your `EmbeddingRequest`. For example to override the default model name for a specific request: [source,java] ---- -EmbeddingResponse embeddingResponse = embeddingClient.call( +EmbeddingResponse embeddingResponse = embeddingModel.call( new EmbeddingRequest(List.of("Hello World", "World is big and salvation is near"), MiniMaxEmbeddingOptions.builder() .withModel("Different-Embedding-Model-Deployment-Name") @@ -122,8 +122,8 @@ EmbeddingResponse embeddingResponse = embeddingClient.call( == Sample Controller -This will create a `EmbeddingClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the `EmbeddingClient` implementation. +This will create a `EmbeddingModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the `EmbeddingC` implementation. [source,application.properties] ---- @@ -136,16 +136,16 @@ spring.ai.minimax.embedding.options.model=embo-01 @RestController public class EmbeddingController { - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; @Autowired - public EmbeddingController(EmbeddingClient embeddingClient) { - this.embeddingClient = embeddingClient; + public EmbeddingController(EmbeddingModel embeddingModel) { + this.embeddingModel = embeddingModel; } @GetMapping("/ai/embedding") public Map embed(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - EmbeddingResponse embeddingResponse = this.embeddingClient.embedForResponse(List.of(message)); + EmbeddingResponse embeddingResponse = this.embeddingModel.embedForResponse(List.of(message)); return Map.of("embedding", embeddingResponse); } } @@ -174,21 +174,21 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -NOTE: The `spring-ai-minimax` dependency provides access also to the `MiniMaxChatClient`. -For more information about the `MiniMaxChatClient` refer to the link:../chat/minimax-chat.html[MiniMax Chat Client] section. +NOTE: The `spring-ai-minimax` dependency provides access also to the `MiniMaxChatModel`. +For more information about the `MiniMaxChatModel refer to the link:../chat/minimax-chat.html[MiniMax Chat Client] section. -Next, create an `MiniMaxEmbeddingClient` instance and use it to compute the similarity between two input texts: +Next, create an `MiniMaxEmbeddingModel` instance and use it to compute the similarity between two input texts: [source,java] ---- var miniMaxApi = new MiniMaxApi(System.getenv("MINIMAX_API_KEY")); -var embeddingClient = new MiniMaxEmbeddingClient(miniMaxApi) +var embeddingModel = new MiniMaxEmbeddingModel(miniMaxApi) .withDefaultOptions(MiniMaxChatOptions.build() .withModel("embo-01") .build()); -EmbeddingResponse embeddingResponse = embeddingClient +EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/mistralai-embeddings.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/mistralai-embeddings.adoc index 57d8c19db..77afada94 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/mistralai-embeddings.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/mistralai-embeddings.adoc @@ -80,13 +80,13 @@ The prefix `spring.ai.mistralai` is used as the property prefix that lets you co ==== Configuration Properties -The prefix `spring.ai.mistralai.embedding` is property prefix that configures the `EmbeddingClient` implementation for MistralAI. +The prefix `spring.ai.mistralai.embedding` is property prefix that configures the `EmbeddingModel` implementation for MistralAI. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.mistralai.embedding.enabled | Enable OpenAI embedding client. | true +| spring.ai.mistralai.embedding.enabled | Enable OpenAI embedding model. | true | spring.ai.mistralai.embedding.base-url | Optional overrides the spring.ai.mistralai.base-url to provide embedding specific url | - | spring.ai.mistralai.embedding.api-key | Optional overrides the spring.ai.mistralai.api-key to provide embedding specific api-key | - | spring.ai.mistralai.embedding.metadata-mode | Document content extraction mode. | EMBED @@ -94,7 +94,7 @@ The prefix `spring.ai.mistralai.embedding` is property prefix that configures th | spring.ai.mistralai.embedding.options.encodingFormat | The format to return the embeddings in. Can be either float or base64. | - |==== -NOTE: You can override the common `spring.ai.mistralai.base-url` and `spring.ai.mistralai.api-key` for the `ChatClient` and `EmbeddingClient` implementations. +NOTE: You can override the common `spring.ai.mistralai.base-url` and `spring.ai.mistralai.api-key` for the `ChatModel` and `EmbeddingModel` implementations. The `spring.ai.mistralai.embedding.base-url` and `spring.ai.mistralai.embedding.api-key` properties if set take precedence over the common properties. Similarly, the `spring.ai.mistralai.embedding.base-url` and `spring.ai.mistralai.embedding.api-key` properties if set take precedence over the common properties. This is useful if you want to use different MistralAI accounts for different models and different model endpoints. @@ -107,14 +107,14 @@ The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-mist The default options can be configured using the `spring.ai.mistralai.embedding.options` properties as well. -At start-time use the `MistralAiEmbeddingClient` constructor to set the default options used for all embedding requests. +At start-time use the `MistralAiEmbeddingModel` constructor to set the default options used for all embedding requests. At run-time you can override the default options, using a `MistralAiEmbeddingOptions` instance as part of your `EmbeddingRequest`. For example to override the default model name for a specific request: [source,java] ---- -EmbeddingResponse embeddingResponse = embeddingClient.call( +EmbeddingResponse embeddingResponse = embeddingModel.call( new EmbeddingRequest(List.of("Hello World", "World is big and salvation is near"), MistralAiEmbeddingOptions.builder() .withModel("Different-Embedding-Model-Deployment-Name") @@ -123,8 +123,8 @@ EmbeddingResponse embeddingResponse = embeddingClient.call( == Sample Controller -This will create a `EmbeddingClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the `EmbeddingClient` implementation. +This will create a `EmbeddingModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the `EmbeddingModel` implementation. [source,application.properties] ---- @@ -137,16 +137,16 @@ spring.ai.mistralai.embedding.options.model=mistral-embed @RestController public class EmbeddingController { - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; @Autowired - public EmbeddingController(EmbeddingClient embeddingClient) { - this.embeddingClient = embeddingClient; + public EmbeddingController(EmbeddingModel embeddingModel) { + this.embeddingModel = embeddingModel; } @GetMapping("/ai/embedding") public Map embed(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - var embeddingResponse = this.embeddingClient.embedForResponse(List.of(message)); + var embeddingResponse = this.embeddingModel.embedForResponse(List.of(message)); return Map.of("embedding", embeddingResponse); } } @@ -175,22 +175,22 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -NOTE: The `spring-ai-mistral-ai` dependency provides access also to the `MistralAiChatClient`. -For more information about the `MistralAiChatClient` refer to the link:../chat/mistralai-chat.html[MistralAI Chat Client] section. +NOTE: The `spring-ai-mistral-ai` dependency provides access also to the `MistralAiChatModel`. +For more information about the `MistralAiChatModel` refer to the link:../chat/mistralai-chat.html[MistralAI Chat Client] section. -Next, create an `MistralAiEmbeddingClient` instance and use it to compute the similarity between two input texts: +Next, create an `MistralAiEmbeddingModel` instance and use it to compute the similarity between two input texts: [source,java] ---- var mistralAiApi = new MistralAiApi(System.getenv("MISTRAL_AI_API_KEY")); -var embeddingClient = new MistralAiEmbeddingClient(mistralAiApi, +var embeddingModel = new MistralAiEmbeddingModel(mistralAiApi, MistralAiEmbeddingOptions.builder() .withModel("mistral-embed") .withEncodingFormat("float") .build()); -EmbeddingResponse embeddingResponse = embeddingClient +EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/ollama-embeddings.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/ollama-embeddings.adoc index 386051e71..095692862 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/ollama-embeddings.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/ollama-embeddings.adoc @@ -1,7 +1,7 @@ = Ollama Embeddings With https://ollama.ai/[Ollama] you can run various Large Language Models (LLMs) locally and generate embeddings from them. -Spring AI supports the Ollama text embeddings with `OllamaEmbeddingClient`. +Spring AI supports the Ollama text embeddings with `OllamaEmbeddingModel`. An embedding is a vector (list) of floating point numbers. The distance between two vectors measures their relatedness. @@ -47,7 +47,7 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. The `spring.ai.ollama.embedding.options.*` properties are used to configure the default options used for all embedding requests. -(It is used as `OllamaEmbeddingClient#withDefaultOptions()` instance). +(It is used as `OllamaEmbeddingModel#withDefaultOptions()` instance). === Embedding Properties @@ -60,13 +60,13 @@ The prefix `spring.ai.ollama` is the property prefix to configure the connection | spring.ai.ollama.base-url | Base URL where Ollama API server is running. | `http://localhost:11434` |==== -The prefix `spring.ai.ollama.embedding.options` is the property prefix that configures the `EmbeddingClient` implementation for Ollama. +The prefix `spring.ai.ollama.embedding.options` is the property prefix that configures the `EmbeddingModel` implementation for Ollama. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.ollama.embedding.enabled | Enable Ollama embedding client. | true +| spring.ai.ollama.embedding.enabled | Enable Ollama embedding model. | true | spring.ai.ollama.embedding.options.model | The name of the https://github.com/ollama/ollama?tab=readme-ov-file#model-library[supported model] to use. | mistral |==== @@ -115,14 +115,14 @@ The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-olla The default options can be configured using the `spring.ai.ollama.embedding.options` properties as well. -At start-time use the `OllamaEmbeddingClient#withDefaultOptions()` to configure the default options used for all embedding requests. +At start-time use the `OllamaEmbeddingModel#withDefaultOptions()` to configure the default options used for all embedding requests. At run-time you can override the default options, using a `OllamaOptions` instance as part of your `EmbeddingRequest`. For example to override the default model name for a specific request: [source,java] ---- -EmbeddingResponse embeddingResponse = embeddingClient.call( +EmbeddingResponse embeddingResponse = embeddingModel.call( new EmbeddingRequest(List.of("Hello World", "World is big and salvation is near"), OllamaOptions.create() .withModel("Different-Embedding-Model-Deployment-Name")); @@ -130,24 +130,24 @@ EmbeddingResponse embeddingResponse = embeddingClient.call( == Sample Controller -This will create a `EmbeddingClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the `EmbeddingClient` implementation. +This will create a `EmbeddingModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the `EmbeddingModel` implementation. [source,java] ---- @RestController public class EmbeddingController { - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; @Autowired - public EmbeddingController(EmbeddingClient embeddingClient) { - this.embeddingClient = embeddingClient; + public EmbeddingController(EmbeddingModel embeddingModel) { + this.embeddingModel = embeddingModel; } @GetMapping("/ai/embedding") public Map embed(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - EmbeddingResponse embeddingResponse = this.embeddingClient.embedForResponse(List.of(message)); + EmbeddingResponse embeddingResponse = this.embeddingModel.embedForResponse(List.of(message)); return Map.of("embedding", embeddingResponse); } } @@ -155,7 +155,7 @@ public class EmbeddingController { == Manual Configuration -If you are not using Spring Boot, you can manually configure the `OllamaEmbeddingClient`. +If you are not using Spring Boot, you can manually configure the `OllamaEmbeddingModel`. For this add the spring-ai-ollama dependency to your project’s Maven pom.xml file: [source,xml] @@ -177,21 +177,21 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -NOTE: The `spring-ai-ollama` dependency provides access also to the `OllamaChatClient`. -For more information about the `OllamaChatClient` refer to the link:../chat/ollama-chat.html[Ollama Chat Client] section. +NOTE: The `spring-ai-ollama` dependency provides access also to the `OllamaChatModel`. +For more information about the `OllamaChatModel` refer to the link:../chat/ollama-chat.html[Ollama Chat Client] section. -Next, create an `OllamaEmbeddingClient` instance and use it to compute the similarity between two input texts: +Next, create an `OllamaEmbeddingModel` instance and use it to compute the similarity between two input texts: [source,java] ---- var ollamaApi = new OllamaApi(); -var embeddingClient = new OllamaEmbeddingClient(ollamaApi) +var embeddingModel = new OllamaEmbeddingModel(ollamaApi) .withDefaultOptions(OllamaOptions.create() .withModel(OllamaOptions.DEFAULT_MODEL) .toMap()); -EmbeddingResponse embeddingResponse = embeddingClient +EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/onnx.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/onnx.adoc index 565dafc22..286f2d8fa 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/onnx.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/onnx.adoc @@ -1,6 +1,6 @@ = Transformers (ONNX) Embeddings -The `TransformersEmbeddingClient` is an `EmbeddingClient` implementation that locally computes https://www.sbert.net/examples/applications/computing-embeddings/README.html#sentence-embeddings-with-transformers[sentence embeddings] using a selected https://www.sbert.net/[sentence transformer]. +The `TransformersEmbeddingModel` is an `EmbeddingModel` implementation that locally computes https://www.sbert.net/examples/applications/computing-embeddings/README.html#sentence-embeddings-with-transformers[sentence embeddings] using a selected https://www.sbert.net/[sentence transformer]. It uses https://www.sbert.net/docs/pretrained_models.html[pre-trained] transformer models, serialized into the https://onnx.ai/[Open Neural Network Exchange (ONNX)] format. @@ -25,7 +25,7 @@ source ./venv/bin/activate (venv) optimum-cli export onnx --generative sentence-transformers/all-MiniLM-L6-v2 onnx-output-folder ---- -The snippet exports the https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2[sentence-transformers/all-MiniLM-L6-v2] transformer into the `onnx-output-folder` folder. Later includes the `tokenizer.json` and `model.onnx` files used by the embedding client. +The snippet exports the https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2[sentence-transformers/all-MiniLM-L6-v2] transformer into the `onnx-output-folder` folder. Later includes the `tokenizer.json` and `model.onnx` files used by the embedding model. In place of the all-MiniLM-L6-v2 you can pick any huggingface transformer identifier or provide direct file path. @@ -43,9 +43,9 @@ Add the `spring-ai-transformers` project to your maven dependencies: TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -then create a new `TransformersEmbeddingClient` instance and use the `setTokenizerResource(tokenizerJsonUri)` and `setModelResource(modelOnnxUri)` methods to set the URIs of the exported `tokenizer.json` and `model.onnx` files. (`classpath:`, `file:` or `https:` URI schemas are supported). +then create a new `TransformersEmbeddingModel` instance and use the `setTokenizerResource(tokenizerJsonUri)` and `setModelResource(modelOnnxUri)` methods to set the URIs of the exported `tokenizer.json` and `model.onnx` files. (`classpath:`, `file:` or `https:` URI schemas are supported). -If the model is not explicitly set, `TransformersEmbeddingClient` defaults to https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2[sentence-transformers/all-MiniLM-L6-v2]: +If the model is not explicitly set, `TransformersEmbeddingModel` defaults to https://huggingface.co/sentence-transformers/all-MiniLM-L6-v2[sentence-transformers/all-MiniLM-L6-v2]: [cols="2*"] |=== @@ -55,53 +55,53 @@ If the model is not explicitly set, `TransformersEmbeddingClient` defaults to ht | Size | 80MB |=== -The following snippet illustrates how to use the `TransformersEmbeddingClient` manually: +The following snippet illustrates how to use the `TransformersEmbeddingModel` manually: [source,java] ---- -TransformersEmbeddingClient embeddingClient = new TransformersEmbeddingClient(); +TransformersEmbeddingModel embeddingModel = new TransformersEmbeddingModel(); // (optional) defaults to classpath:/onnx/all-MiniLM-L6-v2/tokenizer.json -embeddingClient.setTokenizerResource("classpath:/onnx/all-MiniLM-L6-v2/tokenizer.json"); +embeddingModel.setTokenizerResource("classpath:/onnx/all-MiniLM-L6-v2/tokenizer.json"); // (optional) defaults to classpath:/onnx/all-MiniLM-L6-v2/model.onnx -embeddingClient.setModelResource("classpath:/onnx/all-MiniLM-L6-v2/model.onnx"); +embeddingModel.setModelResource("classpath:/onnx/all-MiniLM-L6-v2/model.onnx"); // (optional) defaults to ${java.io.tmpdir}/spring-ai-onnx-model // Only the http/https resources are cached by default. -embeddingClient.setResourceCacheDirectory("/tmp/onnx-zoo"); +embeddingModel.setResourceCacheDirectory("/tmp/onnx-zoo"); // (optional) Set the tokenizer padding if you see an errors like: // "ai.onnxruntime.OrtException: Supplied array is ragged, ..." -embeddingClient.setTokenizerOptions(Map.of("padding", "true")); +embeddingModel.setTokenizerOptions(Map.of("padding", "true")); -embeddingClient.afterPropertiesSet(); +embeddingModel.afterPropertiesSet(); -List> embeddings = embeddingClient.embed(List.of("Hello world", "World is big")); +List> embeddings = embeddingModel.embed(List.of("Hello world", "World is big")); ---- -NOTE: If you create an instance of `TransformersEmbeddingClient` manually, you must call the `afterPropertiesSet()` method after setting the properties and before using the client. +NOTE: If you create an instance of `TransformersEmbeddingModel` manually, you must call the `afterPropertiesSet()` method after setting the properties and before using the client. The first `embed()` call downloads the large ONNX model and caches it on the local file system. Therefore, the first call might take longer than usual. Use the `#setResourceCacheDirectory()` method to set the local folder where the ONNX models as stored. The default cache folder is `${java.io.tmpdir}/spring-ai-onnx-model`. -It is more convenient (and preferred) to create the TransformersEmbeddingClient as a `Bean`. +It is more convenient (and preferred) to create the TransformersEmbeddingModel as a `Bean`. Then you don't have to call the `afterPropertiesSet()` manually. [source,java] ---- @Bean -public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); +public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } ---- == Transformers Embedding Spring Boot Starter -You can bootstrap and autowire the `TransformersEmbeddingClient` with the following Spring Boot starter: +You can bootstrap and autowire the `TransformersEmbeddingModel` with the following Spring Boot starter: [source,xml] ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/openai-embeddings.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/openai-embeddings.adoc index 75a0d45f3..b57ee6cda 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/openai-embeddings.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/openai-embeddings.adoc @@ -81,13 +81,13 @@ The prefix `spring.ai.openai` is used as the property prefix that lets you conne ==== Configuration Properties -The prefix `spring.ai.openai.embedding` is property prefix that configures the `EmbeddingClient` implementation for OpenAI. +The prefix `spring.ai.openai.embedding` is property prefix that configures the `EmbeddingModel` implementation for OpenAI. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.openai.embedding.enabled | Enable OpenAI embedding client. | true +| spring.ai.openai.embedding.enabled | Enable OpenAI embedding model. | true | spring.ai.openai.embedding.base-url | Optional overrides the spring.ai.openai.base-url to provide embedding specific url | - | spring.ai.openai.embedding.api-key | Optional overrides the spring.ai.openai.api-key to provide embedding specific api-key | - | spring.ai.openai.embedding.metadata-mode | Document content extraction mode. | EMBED @@ -97,7 +97,7 @@ The prefix `spring.ai.openai.embedding` is property prefix that configures the ` | spring.ai.openai.embedding.options.dimensions | The number of dimensions the resulting output embeddings should have. Only supported in `text-embedding-3` and later models. | - |==== -NOTE: You can override the common `spring.ai.openai.base-url` and `spring.ai.openai.api-key` for the `ChatClient` and `EmbeddingClient` implementations. +NOTE: You can override the common `spring.ai.openai.base-url` and `spring.ai.openai.api-key` for the `ChatModel` and `EmbeddingModel` implementations. The `spring.ai.openai.embedding.base-url` and `spring.ai.openai.embedding.api-key` properties if set take precedence over the common properties. Similarly, the `spring.ai.openai.embedding.base-url` and `spring.ai.openai.embedding.api-key` properties if set take precedence over the common properties. This is useful if you want to use different OpenAI accounts for different models and different model endpoints. @@ -110,14 +110,14 @@ The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-open The default options can be configured using the `spring.ai.openai.embedding.options` properties as well. -At start-time use the `OpenAiEmbeddingClient` constructor to set the default options used for all embedding requests. +At start-time use the `OpenAiEmbeddingModel` constructor to set the default options used for all embedding requests. At run-time you can override the default options, using a `OpenAiEmbeddingOptions` instance as part of your `EmbeddingRequest`. For example to override the default model name for a specific request: [source,java] ---- -EmbeddingResponse embeddingResponse = embeddingClient.call( +EmbeddingResponse embeddingResponse = embeddingModel.call( new EmbeddingRequest(List.of("Hello World", "World is big and salvation is near"), OpenAiEmbeddingOptions.builder() .withModel("Different-Embedding-Model-Deployment-Name") @@ -126,8 +126,8 @@ EmbeddingResponse embeddingResponse = embeddingClient.call( == Sample Controller -This will create a `EmbeddingClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the `EmbeddingClient` implementation. +This will create a `EmbeddingModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the `EmbeddingModel` implementation. [source,application.properties] ---- @@ -140,16 +140,16 @@ spring.ai.openai.embedding.options.model=text-embedding-ada-002 @RestController public class EmbeddingController { - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; @Autowired - public EmbeddingController(EmbeddingClient embeddingClient) { - this.embeddingClient = embeddingClient; + public EmbeddingController(EmbeddingModel embeddingModel) { + this.embeddingModel = embeddingModel; } @GetMapping("/ai/embedding") public Map embed(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - EmbeddingResponse embeddingResponse = this.embeddingClient.embedForResponse(List.of(message)); + EmbeddingResponse embeddingResponse = this.embeddingModel.embedForResponse(List.of(message)); return Map.of("embedding", embeddingResponse); } } @@ -178,16 +178,16 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -NOTE: The `spring-ai-openai` dependency provides access also to the `OpenAiChatClient`. -For more information about the `OpenAiChatClient` refer to the link:../chat/openai-chat.html[OpenAI Chat Client] section. +NOTE: The `spring-ai-openai` dependency provides access also to the `OpenAiChatModel`. +For more information about the `OpenAiChatModel` refer to the link:../chat/openai-chat.html[OpenAI Chat Client] section. -Next, create an `OpenAiEmbeddingClient` instance and use it to compute the similarity between two input texts: +Next, create an `OpenAiEmbeddingModel` instance and use it to compute the similarity between two input texts: [source,java] ---- var openAiApi = new OpenAiApi(System.getenv("OPENAI_API_KEY")); -var embeddingClient = new OpenAiEmbeddingClient( +var embeddingModel = new OpenAiEmbeddingModel( openAiApi, MetadataMode.EMBED, OpenAiEmbeddingOptions.builder() @@ -196,7 +196,7 @@ var embeddingClient = new OpenAiEmbeddingClient( .build(), RetryUtils.DEFAULT_RETRY_TEMPLATE); -EmbeddingResponse embeddingResponse = embeddingClient +EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/postgresml-embeddings.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/postgresml-embeddings.adoc index a8a506a3f..8f4458961 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/postgresml-embeddings.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/postgresml-embeddings.adoc @@ -40,16 +40,16 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Use the `spring.ai.postgresml.embedding.options.*` properties to configure your `PostgresMlEmbeddingClient`. links +Use the `spring.ai.postgresml.embedding.options.*` properties to configure your `PostgresMlEmbeddingModel`. links === Embedding Properties -The prefix `spring.ai.postgresml.embedding` is property prefix that configures the `EmbeddingClient` implementation for PostgresML embeddings. +The prefix `spring.ai.postgresml.embedding` is property prefix that configures the `EmbeddingModel` implementation for PostgresML embeddings. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.postgresml.embedding.enabled | Enable PostgresML embedding client. | true +| spring.ai.postgresml.embedding.enabled | Enable PostgresML embedding model. | true | spring.ai.postgresml.embedding.options.transformer | The Huggingface transformer model to use for the embedding. | distilbert-base-uncased | spring.ai.postgresml.embedding.options.kwargs | Additional transformer specific options. | empty map | spring.ai.postgresml.embedding.options.vectorType | PostgresML vector type to use for the embedding. Two options are supported: `PG_ARRAY` and `PG_VECTOR`. | PG_ARRAY @@ -61,10 +61,10 @@ TIP: All properties prefixed with `spring.ai.postgresml.embedding.options` can b == Runtime Options [[embedding-options]] -Use the https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingOptions.java[PostgresMlEmbeddingOptions.java] to configure the `PostgresMlEmbeddingClient` with options, such as the model to use and etc. +Use the https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/main/java/org/springframework/ai/postgresml/PostgresMlEmbeddingOptions.java[PostgresMlEmbeddingOptions.java] to configure the `PostgresMlEmbeddingModel` with options, such as the model to use and etc. -On start you can pass a `PostgresMlEmbeddingOptions` to the `PostgresMlEmbeddingClient` constructor to configure the default options used for all embedding requests. +On start you can pass a `PostgresMlEmbeddingOptions` to the `PostgresMlEmbeddingModel` constructor to configure the default options used for all embedding requests. At run-time you can override the default options, using a `PostgresMlEmbeddingOptions` in your `EmbeddingRequest`. @@ -73,7 +73,7 @@ For example to override the default model name for a specific request: [source,java] ---- -EmbeddingResponse embeddingResponse = embeddingClient.call( +EmbeddingResponse embeddingResponse = embeddingModel.call( new EmbeddingRequest(List.of("Hello World", "World is big and salvation is near"), PostgresMlEmbeddingOptions.builder() .withTransformer("intfloat/e5-small") @@ -84,8 +84,8 @@ EmbeddingResponse embeddingResponse = embeddingClient.call( == Sample Controller -This will create a `EmbeddingClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the `EmbeddingClient` implementation. +This will create a `EmbeddingModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the `EmbeddingModel` implementation. [source,application.properties] ---- @@ -100,16 +100,16 @@ spring.ai.postgresml.embedding.options.kwargs.device=cpu @RestController public class EmbeddingController { - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; @Autowired - public EmbeddingController(EmbeddingClient embeddingClient) { - this.embeddingClient = embeddingClient; + public EmbeddingController(EmbeddingModel embeddingModel) { + this.embeddingModel = embeddingModel; } @GetMapping("/ai/embedding") public Map embed(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - EmbeddingResponse embeddingResponse = this.embeddingClient.embedForResponse(List.of(message)); + EmbeddingResponse embeddingResponse = this.embeddingModel.embedForResponse(List.of(message)); return Map.of("embedding", embeddingResponse); } } @@ -117,7 +117,7 @@ public class EmbeddingController { == Manual configuration -Instead of using the Spring Boot auto-configuration, you can create the `PostgresMlEmbeddingClient` manually. +Instead of using the Spring Boot auto-configuration, you can create the `PostgresMlEmbeddingModel` manually. For this add the `spring-ai-postgresml` dependency to your project's Maven `pom.xml` file: [source, xml] @@ -139,13 +139,13 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create an `PostgresMlEmbeddingClient` instance and use it to compute the similarity between two input texts: +Next, create an `PostgresMlEmbeddingModel` instance and use it to compute the similarity between two input texts: [source,java] ---- var jdbcTemplate = new JdbcTemplate(dataSource); // your posgresml data source -PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.jdbcTemplate, +PostgresMlEmbeddingModel embeddingModel = new PostgresMlEmbeddingModel(this.jdbcTemplate, PostgresMlEmbeddingOptions.builder() .withTransformer("distilbert-base-uncased") // huggingface transformer model name. .withVectorType(VectorType.PG_VECTOR) //vector type in PostgreSQL. @@ -153,21 +153,21 @@ PostgresMlEmbeddingClient embeddingClient = new PostgresMlEmbeddingClient(this.j .withMetadataMode(MetadataMode.EMBED) // Document metadata mode. .build()); -embeddingClient.afterPropertiesSet(); // initialize the jdbc template and database. +embeddingModel.afterPropertiesSet(); // initialize the jdbc template and database. -EmbeddingResponse embeddingResponse = embeddingClient +EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); ---- NOTE: When created manually, you must call the `afterPropertiesSet()` after setting the properties and before using the client. -It is more convenient (and preferred) to create the PostgresMlEmbeddingClient as a `@Bean`. +It is more convenient (and preferred) to create the PostgresMlEmbeddingModel as a `@Bean`. Then you don’t have to call the `afterPropertiesSet()` manually: [source,java] ---- @Bean -public EmbeddingClient embeddingClient(JdbcTemplate jdbcTemplate) { - return new PostgresMlEmbeddingClient(jdbcTemplate, +public EmbeddingModel embeddingModel(JdbcTemplate jdbcTemplate) { + return new PostgresMlEmbeddingModel(jdbcTemplate, PostgresMlEmbeddingOptions.builder() .... .build()); diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/vertexai-embeddings.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/vertexai-embeddings.adoc index d66b94071..34e8bcd7a 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/vertexai-embeddings.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/vertexai-embeddings.adoc @@ -76,7 +76,7 @@ The prefix `spring.ai.vertex.ai.embedding` is the property prefix that lets you https://start.spring.io/[Create] a new Spring Boot project and add the `spring-ai-vertex-ai-palm2-spring-boot-starter` to your pom (or gradle) dependencies. -Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the VertexAi Chat client: +Add a `application.properties` file, under the `src/main/resources` directory, to enable and configure the VertexAi chat model: [source,application.properties] ---- @@ -86,7 +86,7 @@ spring.ai.vertex.ai.embedding.model=embedding-gecko-001 TIP: replace the `api-key` with your VertexAI credentials. -This will create a `VertexAiPaLm2EmbeddingClient` implementation that you can inject into your class. +This will create a `VertexAiPaLm2EmbeddingModel` implementation that you can inject into your class. Here is an example of a simple `@Controller` class that uses the embedding client for text generations. [source,java] @@ -94,16 +94,16 @@ Here is an example of a simple `@Controller` class that uses the embedding clien @RestController public class EmbeddingController { - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; @Autowired - public EmbeddingController(EmbeddingClient embeddingClient) { - this.embeddingClient = embeddingClient; + public EmbeddingController(EmbeddingModel embeddingModel) { + this.embeddingModel = embeddingModel; } @GetMapping("/ai/embedding") public Map embed(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - EmbeddingResponse embeddingResponse = this.embeddingClient.embedForResponse(List.of(message)); + EmbeddingResponse embeddingResponse = this.embeddingModel.embedForResponse(List.of(message)); return Map.of("embedding", embeddingResponse); } } @@ -111,7 +111,7 @@ public class EmbeddingController { == Manual Configuration -The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2EmbeddingClient.java[VertexAiPaLm2EmbeddingClient] implements the `EmbeddingClient` and uses the <> to connect to the VertexAI service. +The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-vertex-ai-palm2/src/main/java/org/springframework/ai/vertexai/palm2/VertexAiPaLm2EmbeddingModel.java[VertexAiPaLm2EmbeddingModel] implements the `EmbeddingModel` and uses the <> to connect to the VertexAI service. Add the `spring-ai-vertex-ai` dependency to your project's Maven `pom.xml` file: @@ -134,15 +134,15 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create a `VertexAiPaLm2EmbeddingClient` and use it for text generations: +Next, create a `VertexAiPaLm2EmbeddingModel` and use it for text generations: [source,java] ---- VertexAiPaLm2Api vertexAiApi = new VertexAiPaLm2Api(< YOUR PALM_API_KEY>); -var embeddingClient = new VertexAiPaLm2EmbeddingClient(vertexAiApi); +var embeddingModel = new VertexAiPaLm2EmbeddingModel(vertexAiApi); -EmbeddingResponse embeddingResponse = embeddingClient +EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/zhipuai-embeddings.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/zhipuai-embeddings.adoc index a189b448a..1ad2abdca 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/zhipuai-embeddings.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/embeddings/zhipuai-embeddings.adoc @@ -81,19 +81,19 @@ The prefix `spring.ai.zhipuai` is used as the property prefix that lets you conn ==== Configuration Properties -The prefix `spring.ai.zhipuai.embedding` is property prefix that configures the `EmbeddingClient` implementation for ZhiPuAI. +The prefix `spring.ai.zhipuai.embedding` is property prefix that configures the `EmbeddingModel` implementation for ZhiPuAI. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.zhipuai.embedding.enabled | Enable ZhiPuAI embedding client. | true +| spring.ai.zhipuai.embedding.enabled | Enable ZhiPuAI embedding model. | true | spring.ai.zhipuai.embedding.base-url | Optional overrides the spring.ai.zhipuai.base-url to provide embedding specific url | - | spring.ai.zhipuai.embedding.api-key | Optional overrides the spring.ai.zhipuai.api-key to provide embedding specific api-key | - | spring.ai.zhipuai.embedding.options.model | The model to use | embedding-2 |==== -NOTE: You can override the common `spring.ai.zhipuai.base-url` and `spring.ai.zhipuai.api-key` for the `ChatClient` and `EmbeddingClient` implementations. +NOTE: You can override the common `spring.ai.zhipuai.base-url` and `spring.ai.zhipuai.api-key` for the `ChatModel` and `EmbeddingModel` implementations. The `spring.ai.zhipuai.embedding.base-url` and `spring.ai.zhipuai.embedding.api-key` properties if set take precedence over the common properties. Similarly, the `spring.ai.zhipuai.embedding.base-url` and `spring.ai.zhipuai.embedding.api-key` properties if set take precedence over the common properties. This is useful if you want to use different ZhiPuAI accounts for different models and different model endpoints. @@ -106,14 +106,14 @@ The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-zhip The default options can be configured using the `spring.ai.zhipuai.embedding.options` properties as well. -At start-time use the `ZhiPuAiEmbeddingClient` constructor to set the default options used for all embedding requests. +At start-time use the `ZhiPuAiEmbeddingModel` constructor to set the default options used for all embedding requests. At run-time you can override the default options, using a `ZhiPuAiEmbeddingOptions` instance as part of your `EmbeddingRequest`. For example to override the default model name for a specific request: [source,java] ---- -EmbeddingResponse embeddingResponse = embeddingClient.call( +EmbeddingResponse embeddingResponse = embeddingModel.call( new EmbeddingRequest(List.of("Hello World", "World is big and salvation is near"), ZhiPuAiEmbeddingOptions.builder() .withModel("Different-Embedding-Model-Deployment-Name") @@ -122,8 +122,8 @@ EmbeddingResponse embeddingResponse = embeddingClient.call( == Sample Controller -This will create a `EmbeddingClient` implementation that you can inject into your class. -Here is an example of a simple `@Controller` class that uses the `EmbeddingClient` implementation. +This will create a `EmbeddingModel` implementation that you can inject into your class. +Here is an example of a simple `@Controller` class that uses the `EmbeddingModel` implementation. [source,application.properties] ---- @@ -136,16 +136,16 @@ spring.ai.zhipuai.embedding.options.model=embedding-2 @RestController public class EmbeddingController { - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; @Autowired - public EmbeddingController(EmbeddingClient embeddingClient) { - this.embeddingClient = embeddingClient; + public EmbeddingController(EmbeddingModel embeddingModel) { + this.embeddingModel = embeddingModel; } @GetMapping("/ai/embedding") public Map embed(@RequestParam(value = "message", defaultValue = "Tell me a joke") String message) { - EmbeddingResponse embeddingResponse = this.embeddingClient.embedForResponse(List.of(message)); + EmbeddingResponse embeddingResponse = this.embeddingModel.embedForResponse(List.of(message)); return Map.of("embedding", embeddingResponse); } } @@ -174,21 +174,21 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -NOTE: The `spring-ai-zhipuai` dependency provides access also to the `ZhiPuAiChatClient`. -For more information about the `ZhiPuAiChatClient` refer to the link:../chat/zhipuai-chat.html[ZhiPuAI Chat Client] section. +NOTE: The `spring-ai-zhipuai` dependency provides access also to the `ZhiPuAiChatModel`. +For more information about the `ZhiPuAiChatModel` refer to the link:../chat/zhipuai-chat.html[ZhiPuAI Chat Client] section. -Next, create an `ZhiPuAiEmbeddingClient` instance and use it to compute the similarity between two input texts: +Next, create an `ZhiPuAiEmbeddingModel` instance and use it to compute the similarity between two input texts: [source,java] ---- var zhiPuAiApi = new ZhiPuAiApi(System.getenv("ZHIPU_AI_API_KEY")); -var embeddingClient = new ZhiPuAiEmbeddingClient(zhiPuAiApi) +var embeddingModel = new ZhiPuAiEmbeddingModel(zhiPuAiApi) .withDefaultOptions(ZhiPuAiChatOptions.build() .withModel("embedding-2") .build()); -EmbeddingResponse embeddingResponse = embeddingClient +EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/generic-model.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/generic-model.adoc index 6b30f90bc..602bb4b95 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/generic-model.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/generic-model.adoc @@ -1,7 +1,7 @@ [[generic-model-api]] = Generic Model API -In order to provide a foundation for all AI Model clients, the Generic Model API was created. +In order to provide a foundation for all AI Models, the Generic Model API was created. This makes it easy to contribute new AI Model support to Spring AI by following a common pattern. The following sections walk through this API. @@ -9,15 +9,15 @@ The following sections walk through this API. image::spring-ai-generic-model-api.jpg[width=900, align="center"] -== ModelClient +== Model -The ModelClient interface provides a generic API for invoking AI models. It is designed to handle the interaction with various types of AI models by abstracting the process of sending requests and receiving responses. The interface uses Java generics to accommodate different types of requests and responses, enhancing flexibility and adaptability across different AI model implementations. +The Model interface provides a generic API for invoking AI models. It is designed to handle the interaction with various types of AI models by abstracting the process of sending requests and receiving responses. The interface uses Java generics to accommodate different types of requests and responses, enhancing flexibility and adaptability across different AI model implementations. The interface is defined below: [source,java] ---- -public interface ModelClient, TRes extends ModelResponse> { +public interface Model, TRes extends ModelResponse> { /** * Executes a method call to the AI model. @@ -29,13 +29,13 @@ public interface ModelClient, TRes extends ModelRes } ---- -== StreamingModelClient +== StreamingModel -The StreamingModelClient interface provides a generic API for invoking an AI model with streaming response. It abstracts the process of sending requests and receiving a streaming response. The interface uses Java generics to accommodate different types of requests and responses, enhancing flexibility and adaptability across different AI model implementations. +The StreamingModel interface provides a generic API for invoking an AI model with streaming response. It abstracts the process of sending requests and receiving a streaming response. The interface uses Java generics to accommodate different types of requests and responses, enhancing flexibility and adaptability across different AI model implementations. [source,java] ---- -public interface StreamingModelClient, TResChunk extends ModelResponse> { +public interface StreamingModel, TResChunk extends ModelResponse> { /** * Executes a method call to the AI model. diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/image/openai-image.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/image/openai-image.adoc index ff544bd85..66766c960 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/image/openai-image.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/image/openai-image.adoc @@ -42,12 +42,12 @@ TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Man === Image Generation Properties -The prefix `spring.ai.openai.image` is the property prefix that lets you configure the `ImageClient` implementation for OpenAI. +The prefix `spring.ai.openai.image` is the property prefix that lets you configure the `ImageModel` implementation for OpenAI. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.openai.image.enabled | Enable OpenAI image client. | true +| spring.ai.openai.image.enabled | Enable OpenAI image model. | true | spring.ai.openai.image.base-url | Optional overrides the spring.ai.openai.base-url to provide chat specific url | - | spring.ai.openai.image.api-key | Optional overrides the spring.ai.openai.api-key to provide chat specific api-key | - | spring.ai.openai.image.options.n | The number of images to generate. Must be between 1 and 10. For dall-e-3, only n=1 is supported. | - @@ -97,14 +97,14 @@ The prefix `spring.ai.retry` is used as the property prefix that lets you config The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiImageOptions.java[OpenAiImageOptions.java] provides model configurations, such as the model to use, the quality, the size, etc. -On start-up, the default options can be configured with the `OpenAiImageClient(OpenAiImageApi openAiImageApi)` constructor and the `withDefaultOptions(OpenAiImageOptions defaultOptions)` method. Alternatively, use the `spring.ai.openai.image.options.*` properties described previously. +On start-up, the default options can be configured with the `OpenAiImageModel(OpenAiImageApi openAiImageApi)` constructor and the `withDefaultOptions(OpenAiImageOptions defaultOptions)` method. Alternatively, use the `spring.ai.openai.image.options.*` properties described previously. At runtime you can override the default options by adding new, request specific, options to the `ImagePrompt` call. For example to override the OpenAI specific options such as quality and the number of images to create, use the following code example: [source,java] ---- -ImageResponse response = openaiImageClient.call( +ImageResponse response = openaiImageModel.call( new ImagePrompt("A light cream colored mini golden doodle", OpenAiImageOptions.builder() .withQuality("hd") diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/image/stabilityai-image.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/image/stabilityai-image.adoc index b526b5405..01f1347f7 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/image/stabilityai-image.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/image/stabilityai-image.adoc @@ -51,13 +51,13 @@ The prefix `spring.ai.stabilityai` is used as the property prefix that lets you | spring.ai.stabilityai.api-key | The API Key | - |==== -The prefix `spring.ai.stabilityai.image` is the property prefix that lets you configure the `ImageClient` implementation for Stability AI. +The prefix `spring.ai.stabilityai.image` is the property prefix that lets you configure the `ImageModel` implementation for Stability AI. [cols="2,5,1"] |==== | Property | Description | Default -| spring.ai.stabilityai.image.enabled | Enable Stability AI image client. | true +| spring.ai.stabilityai.image.enabled | Enable Stability AI image model. | true | spring.ai.stabilityai.image.base-url | Optional overrides the spring.ai.openai.base-url to provide a specific url | `https://api.stability.ai/v1` | spring.ai.stabilityai.image.api-key | Optional overrides the spring.ai.openai.api-key to provide a specific api-key | - | spring.ai.stabilityai.image.option.n | The number of images to be generated. Must be between 1 and 10. | 1 @@ -78,14 +78,14 @@ The prefix `spring.ai.stabilityai.image` is the property prefix that lets you co The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-stabilityai/src/main/java/org/springframework/ai/stabilityai/api/StabilityAiImageOptions.java[StabilityAiImageOptions.java] provides model configurations, such as the model to use, the style, the size, etc. -On start-up, the default options can be configured with the `StabilityAiImageClient(StabilityAiApi stabilityAiApi, StabilityAiImageOptions options)` constructor. Alternatively, use the `spring.ai.openai.image.options.*` properties described previously. +On start-up, the default options can be configured with the `StabilityAiImageModel(StabilityAiApi stabilityAiApi, StabilityAiImageOptions options)` constructor. Alternatively, use the `spring.ai.openai.image.options.*` properties described previously. At runtime, you can override the default options by adding new, request specific, options to the `ImagePrompt` call. For example to override the Stability AI specific options such as quality and the number of images to create, use the following code example: [source,java] ---- -ImageResponse response = openaiImageClient.call( +ImageResponse response = openaiImageModel.call( new ImagePrompt("A light cream colored mini golden doodle", StabilityAiImageOptions.builder() .withStylePreset("cinematic") diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/image/zhipuai-image.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/image/zhipuai-image.adoc index b6a02d3bd..aeedede6c 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/image/zhipuai-image.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/image/zhipuai-image.adoc @@ -48,12 +48,12 @@ TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Man === Image Generation Properties -The prefix `spring.ai.zhipuai.image` is the property prefix that lets you configure the `ImageClient` implementation for ZhiPuAI. +The prefix `spring.ai.zhipuai.image` is the property prefix that lets you configure the `ImageModel` implementation for ZhiPuAI. [cols="3,5,1"] |==== | Property | Description | Default -| spring.ai.zhipuai.image.enabled | Enable ZhiPuAI image client. | true +| spring.ai.zhipuai.image.enabled | Enable ZhiPuAI image model. | true | spring.ai.zhipuai.image.base-url | Optional overrides the spring.ai.zhipuai.base-url to provide chat specific url | - | spring.ai.zhipuai.image.api-key | Optional overrides the spring.ai.zhipuai.api-key to provide chat specific api-key | - | spring.ai.zhipuai.image.options.model | The model to use for image generation. | cogview-3 @@ -96,14 +96,14 @@ The prefix `spring.ai.retry` is used as the property prefix that lets you config The https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-zhipuai/src/main/java/org/springframework/ai/zhipuai/ZhiPuAiImageOptions.java[ZhiPuAiImageOptions.java] provides model configurations, such as the model to use, the quality, the size, etc. -On start-up, the default options can be configured with the `ZhiPuAiImageClient(ZhiPuAiImageApi zhiPuAiImageApi)` constructor and the `withDefaultOptions(ZhiPuAiImageOptions defaultOptions)` method. Alternatively, use the `spring.ai.zhipuai.image.options.*` properties described previously. +On start-up, the default options can be configured with the `ZhiPuAiImageModel(ZhiPuAiImageApi zhiPuAiImageApi)` constructor and the `withDefaultOptions(ZhiPuAiImageOptions defaultOptions)` method. Alternatively, use the `spring.ai.zhipuai.image.options.*` properties described previously. At runtime you can override the default options by adding new, request specific, options to the `ImagePrompt` call. For example to override the ZhiPuAI specific options such as quality and the number of images to create, use the following code example: [source,java] ---- -ImageResponse response = zhiPuAiImageClient.call( +ImageResponse response = zhiPuAiImageModel.call( new ImagePrompt("A light cream colored mini golden doodle", ZhiPuAiImageOptions.builder() .withQuality("hd") diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/imageclient.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/imageclient.adoc index 3d73ce061..314d5a5a7 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/imageclient.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/imageclient.adoc @@ -1,27 +1,27 @@ -[[ImageClient]] -= Image Generation API +[[ImageModel]] += Image Model API -The `Spring Image Generation API` is designed to be a simple and portable interface for interacting with various xref:concepts.adoc#_models[AI Models] specialized in image generation, allowing developers to switch between different image-related models with minimal code changes. +The `Spring Image Model API` is designed to be a simple and portable interface for interacting with various xref:concepts.adoc#_models[AI Models] specialized in image generation, allowing developers to switch between different image-related models with minimal code changes. This design aligns with Spring's philosophy of modularity and interchangeability, ensuring developers can quickly adapt their applications to different AI capabilities related to image processing. -Additionally, with the support of companion classes like `ImagePrompt` for input encapsulation and `ImageResponse` for output handling, the Image Generation API unifies the communication with AI Models dedicated to image generation. +Additionally, with the support of companion classes like `ImagePrompt` for input encapsulation and `ImageResponse` for output handling, the Image Model API unifies the communication with AI Models dedicated to image generation. It manages the complexity of request preparation and response parsing, offering a direct and simplified API interaction for image-generation functionalities. -The Spring Image Generation API is built on top of the Spring AI `Generic Model API`, providing image-specific abstractions and implementations. +The Spring Image Model API is built on top of the Spring AI `Generic Model API`, providing image-specific abstractions and implementations. == API Overview -This section provides a guide to the Spring Image Generation API interface and associated classes. +This section provides a guide to the Spring Image Model API interface and associated classes. -== Image Client +== Image Model -Here is the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/image/ImageClient.java[ImageClient] interface definition: +Here is the link:https://github.com/spring-projects/spring-ai/blob/main/spring-ai-core/src/main/java/org/springframework/ai/image/ImageModel.java[ImageModel] interface definition: [source,java] ---- @FunctionalInterface -public interface ImageClient extends ModelClient { +public interface ImageModel extends Model { ImageResponse call(ImagePrompt request); @@ -68,7 +68,7 @@ public class ImageMessage { public Float getWeight() {...} // constructors and utility methods omitted - +} ---- ==== ImageOptions @@ -94,7 +94,7 @@ public interface ImageOptions extends ModelOptions { } ---- -Additionally, every model specific ImageClient implementation can have its own options that can be passed to the AI model. For example, the OpenAI Image Generation model has its own options like `quality`, `style`, etc. +Additionally, every model specific ImageModel implementation can have its own options that can be passed to the AI model. For example, the OpenAI Image Generation model has its own options like `quality`, `style`, etc. This is a powerful feature that allows developers to use model specific options when starting the application and then override them with at runtime using the `ImagePrompt`. @@ -157,7 +157,7 @@ public class ImageGeneration implements ModelResult { == Available Implementations -`ImageClient` implementations are provided for the following Model providers: +`ImageModel` implementations are provided for the following Model providers: * xref:api/image/openai-image.adoc[OpenAI Image Generation] * xref:api/image/stabilityai-image.adoc[StabilityAI Image Generation] diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/multimodality.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/multimodality.adoc index 320dd8910..2630a2cbb 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/multimodality.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/multimodality.adoc @@ -48,7 +48,7 @@ var userMessage = new UserMessage( "Explain what do you see in this picture?", // content List.of(new Media(MimeTypeUtils.IMAGE_PNG, imageData))); // media -ChatResponse response = chatClient.call(new Prompt(List.of(userMessage))); +ChatResponse response = chatModel.call(new Prompt(List.of(userMessage))); ---- and produce a response like: diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/prompt.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/prompt.adoc index 96fab52f4..c1ea0ef91 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/prompt.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/prompt.adoc @@ -11,7 +11,7 @@ Another analogy is a SQL statement that contain placeholders for certain express As Spring AI evolves, it will introduce higher levels of abstraction for interacting with AI models. The foundational classes described in this section can be likened to JDBC in terms of their role and functionality. -The `ChatClient` class, for instance, is analogous to the core JDBC library in the JDK. +The `ChatModel` class, for instance, is analogous to the core JDBC library in the JDK. Building upon this, Spring AI can provide helper classes similar to `JdbcTemplate`, Spring Data Repositories, and eventually, more advanced constructs like ChatEngines and Agents that consider past interactions with the model. The structure of prompts has evolved over time within the AI field. @@ -24,7 +24,7 @@ OpenAI have introduced even more structure to prompts by categorizing multiple m === Prompt -It is common to use the `call` method of `ChatClient` that takes a `Prompt` instance and returns an `ChatResponse`. +It is common to use the `call` method of `ChatModel` that takes a `Prompt` instance and returns an `ChatResponse`. The Prompt class functions as a container for an organized series of Message objects, with each one forming a segment of the overall prompt. Every Message embodies a unique role within the prompt, differing in its content and intent. @@ -152,7 +152,7 @@ The interfaces implemented by this class support different aspects of prompt cre `PromptTemplateMessageActions` is tailored for prompt creation through the generation and manipulation of Message objects. -`PromptTemplateActions` is designed to return the Prompt object, which can be passed to ChatClient for generating a response. +`PromptTemplateActions` is designed to return the Prompt object, which can be passed to ChatModel for generating a response. While these interfaces might not be used extensively in many projects, they show the different approaches to prompt creation. @@ -213,7 +213,7 @@ PromptTemplate promptTemplate = new PromptTemplate("Tell me a {adjective} joke a Prompt prompt = promptTemplate.create(Map.of("adjective", adjective, "topic", topic)); -return chatClient.call(prompt).getResult(); +return chatModel.call(prompt).getResult(); ``` Another example taken from the https://github.com/Azure-Samples/spring-ai-azure-workshop/blob/main/3-README-prompt-roles.md[AI Workshop on Roles] is shown below. @@ -237,13 +237,13 @@ Message systemMessage = systemPromptTemplate.createMessage(Map.of("name", name, Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); -List response = chatClient.call(prompt).getResults(); +List response = chatModel.call(prompt).getResults(); ``` This shows how you can build up the `Prompt` instance by using the `SystemPromptTemplate` to create a `Message` with the system role passing in placeholder values. The message with the role `user` is then combined with the message of the role `system` to form the prompt. -The prompt is then passed to the ChatClient to get a generative response. +The prompt is then passed to the ChatModel to get a generative response. === Using resources instead of raw Strings diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/speech.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/speech.adoc index a581a57da..26f76603a 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/speech.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/speech.adoc @@ -1,5 +1,5 @@ [[Speech]] -= Text-To-Speech (TTS) API += Speech Model API -Spring AI provides support for OpenAI's Speech API. -When additional providers for Speech are implemented, a common `SpeechClient` and `StreamingSpeechClient` interface will be extracted. \ No newline at end of file +Spring AI provides support for OpenAI's Text-To-Speech (TTS) API. +When additional providers for Speech are implemented, a common `SpeechModel` and `StreamingSpeechModel` interface will be extracted. \ No newline at end of file diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/speech/openai-speech.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/speech/openai-speech.adoc index 94a5f4c7c..6ba1841ac 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/speech/openai-speech.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/speech/openai-speech.adoc @@ -68,7 +68,7 @@ OpenAiAudioSpeechOptions speechOptions = OpenAiAudioSpeechOptions.builder() .build(); SpeechPrompt speechPrompt = new SpeechPrompt("Hello, this is a text-to-speech example.", speechOptions); -SpeechResponse response = openAiAudioSpeechClient.call(speechPrompt); +SpeechResponse response = openAiAudioSpeechModel.call(speechPrompt); ---- == Manual Configuration @@ -94,13 +94,13 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create an `OpenAiAudioSpeechClient`: +Next, create an `OpenAiAudioSpeechModel`: [source,java] ---- var openAiAudioApi = new OpenAiAudioApi(System.getenv("OPENAI_API_KEY")); -var openAiAudioSpeechClient = new OpenAiAudioSpeechClient(openAiAudioApi); +var openAiAudioSpeechModel = new OpenAiAudioSpeechModel(openAiAudioApi); var speechOptions = OpenAiAudioSpeechOptions.builder() .withResponseFormat(OpenAiAudioApi.SpeechRequest.AudioResponseFormat.MP3) @@ -109,7 +109,7 @@ var speechOptions = OpenAiAudioSpeechOptions.builder() .build(); var speechPrompt = new SpeechPrompt("Hello, this is a text-to-speech example.", speechOptions); -SpeechResponse response = openAiAudioSpeechClient.call(speechPrompt); +SpeechResponse response = openAiAudioSpeechModel.call(speechPrompt); // Accessing metadata (rate limit info) OpenAiAudioSpeechResponseMetadata metadata = response.getMetadata(); @@ -125,7 +125,7 @@ The Speech API provides support for real-time audio streaming using chunk transf ---- var openAiAudioApi = new OpenAiAudioApi(System.getenv("OPENAI_API_KEY")); -var openAiAudioSpeechClient = new OpenAiAudioSpeechClient(openAiAudioApi); +var openAiAudioSpeechModel = new OpenAiAudioSpeechModel(openAiAudioApi); OpenAiAudioSpeechOptions speechOptions = OpenAiAudioSpeechOptions.builder() .withVoice(OpenAiAudioApi.SpeechRequest.Voice.ALLOY) @@ -136,9 +136,9 @@ OpenAiAudioSpeechOptions speechOptions = OpenAiAudioSpeechOptions.builder() SpeechPrompt speechPrompt = new SpeechPrompt("Today is a wonderful day to build something people love!", speechOptions); -Flux responseStream = openAiAudioSpeechClient.stream(speechPrompt); +Flux responseStream = openAiAudioSpeechModel.stream(speechPrompt); ---- == Example Code -* The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechClientIT.java[OpenAiSpeechClientIT.java] test provides some general examples of how to use the library. +* The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/speech/OpenAiSpeechModelIT.java[OpenAiSpeechModelIT.java] test provides some general examples of how to use the library. diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/structured-output-converter.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/structured-output-converter.adoc index 1ac816123..47b63233e 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/structured-output-converter.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/structured-output-converter.adoc @@ -123,7 +123,7 @@ String template = """ {format} """; -Generation generation = chatClient.call( +Generation generation = chatModel.call( new Prompt(new PromptTemplate(template, Map.of("actor", actor, "format", format)).createMessage())).getResult(); ActorsFilms actorsFilms = beanOutputConverter.convert(generation.getOutput().getContent()); @@ -147,7 +147,7 @@ String template = """ Prompt prompt = new Prompt(new PromptTemplate(template, Map.of("format", format)).createMessage()); -Generation generation = chatClient.call(prompt).getResult(); +Generation generation = chatModel.call(prompt).getResult(); List actorsFilms = outputConverter.convert(generation.getOutput().getContent()); ---- @@ -168,7 +168,7 @@ String template = """ PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "an array of numbers from 1 to 9 under they key name 'numbers'", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); -Generation generation = chatClient.call(prompt).getResult(); +Generation generation = chatModel.call(prompt).getResult(); Map result = mapOutputConverter.convert(generation.getOutput().getContent()); ---- @@ -189,7 +189,7 @@ String template = """ PromptTemplate promptTemplate = new PromptTemplate(template, Map.of("subject", "ice cream flavors", "format", format)); Prompt prompt = new Prompt(promptTemplate.createMessage()); -Generation generation = this.chatClient.call(prompt).getResult(); +Generation generation = this.chatModel.call(prompt).getResult(); List list = listOutputConverter.convert(generation.getOutput().getContent()); ---- @@ -201,16 +201,16 @@ The following AI Models have been tested to support List, Map and Bean structure [cols="2,5"] |==== | Model | Integration Tests / Samples -| xref:api/chat/openai-chat.adoc[OpenAI] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatClientIT.java[OpenAiChatClientIT] -| xref:api/chat/anthropic-chat.adoc[Anthropic Claude 3] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatClientIT.java[AnthropicChatClientIT.java] -| xref:api/chat/azure-openai-chat.adoc[Azure OpenAI] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatClientIT.java[AzureOpenAiChatClientIT.java] -| xref:api/chat/mistralai-chat.adoc[Mistral AI] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiChatClientIT.java[MistralAiChatClientIT.java] -| xref:api/chat/ollama-chat.adoc[Ollama] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatClientIT.java[OllamaChatClientIT.java] -| xref:api/chat/vertexai-gemini-chat.adoc[Vertex AI Gemini] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatClientIT.java[VertexAiGeminiChatClientIT.java] -| xref:api/chat/bedrock/bedrock-anthropic.adoc[Bedrock Anthropic 2] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatClientIT.java[BedrockAnthropicChatClientIT.java] -| xref:api/chat/bedrock/bedrock-anthropic3.adoc[Bedrock Anthropic 3] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatClientIT.java[BedrockAnthropic3ChatClientIT.java] -| xref:api/chat/bedrock/bedrock-cohere.adoc[Bedrock Cohere] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatClientIT.java[BedrockCohereChatClientIT.java] -| xref:api/chat/bedrock/bedrock-llama.adoc[Bedrock Llama] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatClientIT.java[BedrockLlamaChatClientIT.java.java] +| xref:api/chat/openai-chat.adoc[OpenAI] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/OpenAiChatModelIT.java[OpenAiChatModelIT] +| xref:api/chat/anthropic-chat.adoc[Anthropic Claude 3] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/AnthropicChatModelIT.java[AnthropicChatModelIT.java] +| xref:api/chat/azure-openai-chat.adoc[Azure OpenAI] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-azure-openai/src/test/java/org/springframework/ai/azure/openai/AzureOpenAiChatModelIT.java[AzureOpenAiChatModelIT.java] +| xref:api/chat/mistralai-chat.adoc[Mistral AI] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-mistral-ai/src/test/java/org/springframework/ai/mistralai/MistralAiChatModelIT.java[MistralAiChatModelIT.java] +| xref:api/chat/ollama-chat.adoc[Ollama] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-ollama/src/test/java/org/springframework/ai/ollama/OllamaChatModelIT.java[OllamaChatModelIT.java] +| xref:api/chat/vertexai-gemini-chat.adoc[Vertex AI Gemini] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-vertex-ai-gemini/src/test/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModelIT.java[VertexAiGeminiChatModelIT.java] +| xref:api/chat/bedrock/bedrock-anthropic.adoc[Bedrock Anthropic 2] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic/BedrockAnthropicChatModelIT.java[BedrockAnthropicChatModelIT.java] +| xref:api/chat/bedrock/bedrock-anthropic3.adoc[Bedrock Anthropic 3] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/anthropic3/BedrockAnthropic3ChatModelIT.java[BedrockAnthropic3ChatModelIT.java] +| xref:api/chat/bedrock/bedrock-cohere.adoc[Bedrock Cohere] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/cohere/BedrockCohereChatModelIT.java[BedrockCohereChatModelIT.java] +| xref:api/chat/bedrock/bedrock-llama.adoc[Bedrock Llama] | link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-bedrock/src/test/java/org/springframework/ai/bedrock/llama/BedrockLlamaChatModelIT.java[BedrockLlamaChatModelIT.java.java] |==== == Build-in JSON mode diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/transcriptions.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/transcriptions.adoc index 703f19908..45156a9fe 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/transcriptions.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/transcriptions.adoc @@ -1,5 +1,5 @@ [[Transcription]] -= Transcription API += Transcription Model API -Spring AI provides support for OpenAI's Transcription API. -When additional providers for Transcription are implemented, a common `AudioTranscriptionClient` interface will be extracted. \ No newline at end of file +Spring AI provides support for OpenAI's Transcription Model API. +When additional providers for Transcription are implemented, a common `AudioTranscriptionModel` interface will be extracted. \ No newline at end of file diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/transcriptions/openai-transcriptions.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/transcriptions/openai-transcriptions.adoc index 5592aa266..5f352f4f9 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/transcriptions/openai-transcriptions.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/transcriptions/openai-transcriptions.adoc @@ -37,7 +37,7 @@ TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Man === Transcription Properties -The prefix `spring.ai.openai.audio.transcription` is used as the property prefix that lets you configure the retry mechanism for the OpenAI Image client. +The prefix `spring.ai.openai.audio.transcription` is used as the property prefix that lets you configure the retry mechanism for the OpenAI image model. [cols="3,5,2"] |==== @@ -69,7 +69,7 @@ OpenAiAudioTranscriptionOptions transcriptionOptions = OpenAiAudioTranscriptionO .withResponseFormat(responseFormat) .build(); AudioTranscriptionPrompt transcriptionRequest = new AudioTranscriptionPrompt(audioFile, transcriptionOptions); -AudioTranscriptionResponse response = openAiTranscriptionClient.call(transcriptionRequest); +AudioTranscriptionResponse response = openAiTranscriptionModel.call(transcriptionRequest); ---- == Manual Configuration @@ -95,13 +95,13 @@ dependencies { TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Management] section to add the Spring AI BOM to your build file. -Next, create a `OpenAiAudioTranscriptionClient` +Next, create a `OpenAiAudioTranscriptionModel` [source,java] ---- var openAiAudioApi = new OpenAiAudioApi(System.getenv("OPENAI_API_KEY")); -var openAiAudioTranscriptionClient = new OpenAiAudioTranscriptionClient(openAiAudioApi); +var openAiAudioTranscriptionModel = new OpenAiAudioTranscriptionModel(openAiAudioApi); var transcriptionOptions = OpenAiAudioTranscriptionOptions.builder() .withResponseFormat(TranscriptResponseFormat.TEXT) @@ -111,8 +111,8 @@ var transcriptionOptions = OpenAiAudioTranscriptionOptions.builder() var audioFile = new FileSystemResource("/path/to/your/resource/speech/jfk.flac"); AudioTranscriptionPrompt transcriptionRequest = new AudioTranscriptionPrompt(audioFile, transcriptionOptions); -AudioTranscriptionResponse response = openAiTranscriptionClient.call(transcriptionRequest); +AudioTranscriptionResponse response = openAiTranscriptionModel.call(transcriptionRequest); ---- == Example Code -* The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionClientIT.java[OpenAiTranscriptionClientIT.java] test provides some general examples how to use the library. \ No newline at end of file +* The link:https://github.com/spring-projects/spring-ai/blob/main/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/audio/transcription/OpenAiTranscriptionModelIT.java[OpenAiTranscriptionModelIT.java] test provides some general examples how to use the library. \ No newline at end of file diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs.adoc index e7725d6e6..93539a19d 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs.adoc @@ -72,7 +72,7 @@ It also contains metadata in the form of key-value pairs, including details such Upon insertion into the vector database, the text content is transformed into a numerical array, or a `List`, known as vector embeddings, using an embedding model. Embedding models, such as https://en.wikipedia.org/wiki/Word2vec[Word2Vec], https://en.wikipedia.org/wiki/GloVe_(machine_learning)[GLoVE], and https://en.wikipedia.org/wiki/BERT_(language_model)[BERT], or OpenAI's `text-embedding-ada-002`, are used to convert words, sentences, or paragraphs into these vector embeddings. -The vector database's role is to store and facilitate similarity searches for these embeddings. It does not generate the embeddings itself. For creating vector embeddings, the `EmbeddingClient` should be utilized. +The vector database's role is to store and facilitate similarity searches for these embeddings. It does not generate the embeddings itself. For creating vector embeddings, the `EmbeddingModel` should be utilized. The `similaritySearch` methods in the interface allow for retrieving documents similar to a given query string. These methods can be fine-tuned by using the following parameters: @@ -113,9 +113,9 @@ Information on each of the `VectorStore` implementations can be found in the sub To compute the embeddings for a vector database, you need to pick an embedding model that matches the higher-level AI model being used. -For example, with OpenAI's ChatGPT, we use the `OpenAiEmbeddingClient` and a model named `text-embedding-ada-002`. +For example, with OpenAI's ChatGPT, we use the `OpenAiEmbeddingModel` and a model named `text-embedding-ada-002`. -The Spring Boot starter's auto-configuration for OpenAI makes an implementation of `EmbeddingClient` available in the Spring application context for dependency injection. +The Spring Boot starter's auto-configuration for OpenAI makes an implementation of `EmbeddingModel` available in the Spring application context for dependency injection. The general usage of loading data into a vector store is something you would do in a batch-like job, by first loading data into Spring AI's `Document` class and then calling the `save` method. diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/apache-cassandra.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/apache-cassandra.adoc index b1c9d9f0d..71776e4cf 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/apache-cassandra.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/apache-cassandra.adoc @@ -37,7 +37,7 @@ It stands out from other HNSW Vector Similarity Search implementations by being == Prerequisites -1. A `EmbeddingClient` instance to compute the document embeddings. This is usually configured as a Spring Bean. Several options are available: +1. A `Embedding` instance to compute the document embeddings. This is usually configured as a Spring Bean. Several options are available: - `Transformers Embedding` - computes the embedding in your local environment. The default is via ONNX and the all-MiniLM-L6-v2 Sentence Transformers. This just works. - If you want to use OpenAI's Embeddings` - uses the OpenAI embedding endpoint. You need to create an account at link:https://platform.openai.com/signup[OpenAI Signup] and generate the api-key token at link:https://platform.openai.com/account/api-keys[API Keys]. @@ -82,11 +82,11 @@ Create a CassandraVectorStore instance connected to your Apache Cassandra databa [source,java] ---- @Bean -public VectorStore vectorStore(EmbeddingClient embeddingClient) { +public VectorStore vectorStore(EmbeddingModel embeddingModel) { CassandraVectorStoreConfig config = CassandraVectorStoreConfig.builder().build(); - return new CassandraVectorStore(config, embeddingClient); + return new CassandraVectorStore(config, embeddingModel); } ---- @@ -189,7 +189,7 @@ Then configure the store like: [source,java] ---- @Bean -public CassandraVectorStore store(EmbeddingClient embeddingClient) { +public CassandraVectorStore store(EmbeddingModel embeddingModel) { List partitionColumns = List.of(new SchemaColumn("wiki", DataTypes.TEXT), new SchemaColumn("language", DataTypes.TEXT), new SchemaColumn("title", DataTypes.TEXT)); @@ -226,13 +226,13 @@ public CassandraVectorStore store(EmbeddingClient embeddingClient) { }) .build(); - return new CassandraVectorStore(conf, embeddingClient()); + return new CassandraVectorStore(conf, embeddingModel()); } @Bean -public EmbeddingClient embeddingClient() { +public EmbeddingModel embeddingModel() { // default is ONNX all-MiniLM-L6-v2 which is what we want - return new TransformersEmbeddingClient(); + return new TransformersEmbeddingModel(); } ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/azure.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/azure.adoc index b9f303c07..e935e4e64 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/azure.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/azure.adoc @@ -97,13 +97,13 @@ public SearchIndexClient searchIndexClient() { } ---- -To create a vector store, you can use the following code by injecting the `SearchIndexClient` bean created in the above sample along with an `EmbeddingClient` provided by the Spring AI library that implements the desired Embeddings interface. +To create a vector store, you can use the following code by injecting the `SearchIndexClient` bean created in the above sample along with an `EmbeddingModel` provided by the Spring AI library that implements the desired Embeddings interface. [source,java] ---- @Bean -public VectorStore vectorStore(SearchIndexClient searchIndexClient, EmbeddingClient embeddingClient) { - return new AzureVectorStore(searchIndexClient, embeddingClient, +public VectorStore vectorStore(SearchIndexClient searchIndexClient, EmbeddingModel embeddingModel) { + return new AzureVectorStore(searchIndexClient, embeddingModel, // Define the metadata fields to be used // in the similarity search filters. List.of(MetadataField.text("country"), diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/chroma.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/chroma.adoc index cde40fc10..19d63685a 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/chroma.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/chroma.adoc @@ -8,8 +8,8 @@ link:https://docs.trychroma.com/[Chroma] is the open-source embedding database. 1. Access to ChromeDB. The <> appendix shows how to set up a DB locally with a Docker container. -2. `EmbeddingClient` instance to compute the document embeddings. Several options are available: -- If required, an API key for the xref:api/embeddings.adoc#available-implementations[EmbeddingClient] to generate the embeddings stored by the `ChromaVectorStore`. +2. `EmbeddingModel` instance to compute the document embeddings. Several options are available: +- If required, an API key for the xref:api/embeddings.adoc#available-implementations[EmbeddingModel] to generate the embeddings stored by the `ChromaVectorStore`. On startup, the `ChromaVectorStore` creates the required collection if one is not provisioned already. @@ -39,16 +39,16 @@ TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Man TIP: Refer to the xref:getting-started.adoc#repositories[Repositories] section to add Milestone and/or Snapshot Repositories to your build file. -Additionally, you will need a configured `EmbeddingClient` bean. Refer to the xref:api/embeddings.adoc#available-implementations[EmbeddingClient] section for more information. +Additionally, you will need a configured `EmbeddingModel` bean. Refer to the xref:api/embeddings.adoc#available-implementations[EmbeddingModel] section for more information. Here is an example of the needed bean: [source,java] ---- @Bean -public EmbeddingClient embeddingClient() { - // Can be any other EmbeddingClient implementation. - return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("SPRING_AI_OPENAI_API_KEY"))); +public EmbeddingModel embeddingModel() { + // Can be any other EmbeddingModel implementation. + return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("SPRING_AI_OPENAI_API_KEY"))); } ---- @@ -181,7 +181,7 @@ Add these dependencies to your project: ---- -* OpenAI: Required for calculating embeddings. You can use any other embedding client implementation. +* OpenAI: Required for calculating embeddings. You can use any other embedding model implementation. [source,xml] ---- @@ -218,8 +218,8 @@ Integrate with OpenAI's embeddings by adding the Spring Boot OpenAI starter to y [source,java] ---- @Bean -public VectorStore chromaVectorStore(EmbeddingClient embeddingClient, ChromaApi chromaApi) { - return new ChromaVectorStore(embeddingClient, chromaApi, "TestCollection"); +public VectorStore chromaVectorStore(EmbeddingModel embeddingModel, ChromaApi chromaApi) { + return new ChromaVectorStore(embeddingModel, chromaApi, "TestCollection"); } ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/elasticsearch.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/elasticsearch.adoc index eaff6cf4b..3c0f4da8d 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/elasticsearch.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/elasticsearch.adoc @@ -41,7 +41,7 @@ TIP: Refer to the xref:getting-started.adoc#repositories[Repositories] section t Please have a look at the list of <> for the vector store to learn about the default values and configuration options. -Additionally, you will need a configured `EmbeddingClient` bean. Refer to the xref:api/embeddings.adoc#available-implementations[EmbeddingClient] section for more information. +Additionally, you will need a configured `EmbeddingModel` bean. Refer to the xref:api/embeddings.adoc#available-implementations[EmbeddingModel] section for more information. Now you can auto-wire the `ElasticsearchVectorStore` as a vector store in your application. @@ -217,14 +217,14 @@ and then create the `ElasticsearchVectorStore` bean: [source,java] ---- @Bean -public ElasticsearchVectorStore vectorStore(EmbeddingClient embeddingClient, RestClient restClient) { - return new ElasticsearchVectorStore( restClient, embeddingClient); +public ElasticsearchVectorStore vectorStore(EmbeddingModel embeddingModel, RestClient restClient) { + return new ElasticsearchVectorStore( restClient, embeddingModel); } -// This can be any EmbeddingClient implementation. +// This can be any EmbeddingModel implementation. @Bean -public EmbeddingClient embeddingClient() { - return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); +public EmbeddingModel embeddingModel() { + return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); } ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/gemfire.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/gemfire.adoc index c09a45ef7..1ccb7add3 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/gemfire.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/gemfire.adoc @@ -65,8 +65,8 @@ public GemFireVectorStoreConfig gemFireVectorStoreConfig() { [source,java] ---- @Bean -public VectorStore vectorStore(GemFireVectorStoreConfig config, EmbeddingClient embeddingClient) { - return new GemFireVectorStore(config, embeddingClient); +public VectorStore vectorStore(GemFireVectorStoreConfig config, EmbeddingModel embeddingModel) { + return new GemFireVectorStore(config, embeddingModel); } ---- - Create a Vector Index which will configure GemFire region. diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/hana.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/hana.adoc index 182356e7c..2538c5d90 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/hana.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/hana.adoc @@ -3,7 +3,7 @@ == Prerequisites * You need a SAP HANA Cloud vector engine account - Refer xref:api/vectordbs/hanadb-provision-a-trial-account.adoc[SAP HANA Cloud vector engine - provision a trial account] guide to create a trial account. -* If required, an API key for the xref:api/embeddings.adoc#available-implementations[EmbeddingClient] to generate the embeddings stored by the vector store. +* If required, an API key for the xref:api/embeddings.adoc#available-implementations[EmbeddingModel] to generate the embeddings stored by the vector store. == Auto-configuration @@ -34,7 +34,7 @@ Please have a look at the list of xref:#_hanacloudvectorstore_properties[configu TIP: Refer to the xref:getting-started.adoc#repositories[Repositories] section to add Milestone and/or Snapshot Repositories to your build file. -Additionally, you will need a configured `EmbeddingClient` bean. Refer to the xref:api/embeddings.adoc#available-implementations[EmbeddingClient] section for more information. +Additionally, you will need a configured `EmbeddingModel` bean. Refer to the xref:api/embeddings.adoc#available-implementations[EmbeddingModel] section for more information. == HanaCloudVectorStore properties @@ -226,7 +226,7 @@ public class CricketWorldCupRepository implements HanaVectorRepository", embeddingClient); +public QdrantVectorStore vectorStore(EmbeddingModel embeddingModel, QdrantClient qdrantClient) { + return new QdrantVectorStore(qdrantClient, "", embeddingModel); } ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/redis.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/redis.adoc index 436fd59ce..2c19ee544 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/redis.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/redis.adoc @@ -16,8 +16,8 @@ link:https://redis.io/docs/interact/search-and-query/[Redis Search and Query] ex - https://app.redislabs.com/#/[Redis Cloud] (recommended) - link:https://hub.docker.com/r/redis/redis-stack[Docker] image _redis/redis-stack:latest_ -2. `EmbeddingClient` instance to compute the document embeddings. Several options are available: -- If required, an API key for the xref:api/embeddings.adoc#available-implementations[EmbeddingClient] to generate the embeddings stored by the `RedisVectorStore`. +2. `EmbeddingModel` instance to compute the document embeddings. Several options are available: +- If required, an API key for the xref:api/embeddings.adoc#available-implementations[EmbeddingModel] to generate the embeddings stored by the `RedisVectorStore`. == Auto-configuration @@ -45,16 +45,16 @@ TIP: Refer to the xref:getting-started.adoc#dependency-management[Dependency Man TIP: Refer to the xref:getting-started.adoc#repositories[Repositories] section to add Milestone and/or Snapshot Repositories to your build file. -Additionally, you will need a configured `EmbeddingClient` bean. Refer to the xref:api/embeddings.adoc#available-implementations[EmbeddingClient] section for more information. +Additionally, you will need a configured `EmbeddingModel` bean. Refer to the xref:api/embeddings.adoc#available-implementations[EmbeddingModel] section for more information. Here is an example of the needed bean: [source,java] ---- @Bean -public EmbeddingClient embeddingClient() { - // Can be any other EmbeddingClient implementation. - return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("SPRING_AI_OPENAI_API_KEY"))); +public EmbeddingModel embeddingModel() { + // Can be any other EmbeddingModel implementation. + return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("SPRING_AI_OPENAI_API_KEY"))); } ---- @@ -179,7 +179,7 @@ Then, create a `RedisVectorStore` bean in your Spring configuration: [source,java] ---- @Bean -public VectorStore vectorStore(EmbeddingClient embeddingClient) { +public VectorStore vectorStore(EmbeddingModel embeddingModel) { RedisVectorStoreConfig config = RedisVectorStoreConfig.builder() .withURI("redis://localhost:6379") // Define the metadata fields to be used @@ -189,7 +189,7 @@ public VectorStore vectorStore(EmbeddingClient embeddingClient) { MetadataField.numeric("year")) .build(); - return new RedisVectorStore(config, embeddingClient); + return new RedisVectorStore(config, embeddingModel); } ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/weaviate.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/weaviate.adoc index 060bb322e..6c8280796 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/weaviate.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/api/vectordbs/weaviate.adoc @@ -10,7 +10,7 @@ It provides tools to store document embeddings, content, and metadata and to sea == Prerequisites -1. `EmbeddingClient` instance to compute the document embeddings. Several options are available: +1. `EmbeddingModel` instance to compute the document embeddings. Several options are available: - `Transformers Embedding` - computes the embedding in your local environment. Follow the ONNX Transformers Embedding instructions. - `OpenAI Embedding` - uses the OpenAI embedding endpoint. You need to create an account at link:https://platform.openai.com/signup[OpenAI Signup] and generate the api-key token at link:https://platform.openai.com/account/api-keys[API Keys]. @@ -40,10 +40,10 @@ dependencies { } ---- -The Vector Store, also requires an `EmbeddingClient` instance to calculate embeddings for the documents. -You can pick one of the available xref:api/embeddings.adoc#available-implementations[EmbeddingClient Implementations]. +The Vector Store, also requires an `EmbeddingModel` instance to calculate embeddings for the documents. +You can pick one of the available xref:api/embeddings.adoc#available-implementations[EmbeddingModel Implementations]. -For example to use the xref:api/embeddings/openai-embeddings.adoc[OpenAI EmbeddingClient] add the following dependency to your project: +For example to use the xref:api/embeddings/openai-embeddings.adoc[OpenAI EmbeddingModel] add the following dependency to your project: [source,xml] ---- @@ -230,13 +230,13 @@ This provides you with an implementation of the Embeddings client: [source,java] ---- @Bean -public WeaviateVectorStore vectorStore(EmbeddingClient embeddingClient, WeaviateClient weaviateClient) { +public WeaviateVectorStore vectorStore(EmbeddingModel embeddingModel, WeaviateClient weaviateClient) { WeaviateVectorStoreConfig.Builder configBuilder = WeaviateVectorStore.WeaviateVectorStoreConfig.builder() .withObjectClass() .withConsistencyLevel(); - return new WeaviateVectorStore(configBuilder.build(), embeddingClient, weaviateClient); + return new WeaviateVectorStore(configBuilder.build(), embeddingModel, weaviateClient); } ---- diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/contribution-guidelines.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/contribution-guidelines.adoc index b56717152..aca1a3c7b 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/contribution-guidelines.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/contribution-guidelines.adoc @@ -16,7 +16,7 @@ To contribute a new model, adhere to the following steps: you'll need to develop a low-level client API class. This often involves utilizing the `RestClient` class from the Spring Framework, similar to the `OpenAiApi` class. -. *Create a ModelClient implementation* +. *Create a Model implementation* Ensure your client conforms to the link:https://docs.spring.io/spring-ai/reference/api/generic-model.html[Generic Model API]. Use existing request and response classes if your model's inputs and outputs are supported. If not, create new classes for the Generic Model API and establish a new Java package. diff --git a/spring-ai-docs/src/main/antora/modules/ROOT/pages/upgrade-notes.adoc b/spring-ai-docs/src/main/antora/modules/ROOT/pages/upgrade-notes.adoc index 5072f31be..7195e6209 100644 --- a/spring-ai-docs/src/main/antora/modules/ROOT/pages/upgrade-notes.adoc +++ b/spring-ai-docs/src/main/antora/modules/ROOT/pages/upgrade-notes.adoc @@ -57,9 +57,9 @@ to === January 24, 2024 Update * Moving the `prompt` and `messages` and `metadata` packages to subpackages of `org.sf.ai.chat` -* New functionality is *text to image* clients. Classes are `OpenAiImageClient` and `StabilityAiImageClient`. See the integration tests for usage, docs are coming soon. +* New functionality is *text to image* clients. Classes are `OpenAiImageModel` and `StabilityAiImageModel`. See the integration tests for usage, docs are coming soon. * A new package `model` that contains interfaces and base classes to support creating AI Model Clients for any input/output data type combination. At the moment the chat and image model packages implement this. We will be updating the embedding package to this new model soon. -* A new "portable options" design pattern. We wanted to provide as much portability in the `ChatClient` as possible across different chat based AI Models. There is a common set of generation options and then those that are specific to a model provider. A sort of "duck typing" approach is used. `ModelOptions` in the model package is a marker interface indicating implementations of this class will provide the options for a model. See `ImageOptions`, a subinterface that defines portable options across all text->image `ImageClient` implementations. Then `StabilityAiImageOptions` and `OpenAiImageOptions` provide the options specific to each model provider. All options classes are created via a fluent API builder all can be passed into the portable `ImageClient` API. These option data types are using in autoconfiguration/configuration properties for the `ImageClient` implementations. +* A new "portable options" design pattern. We wanted to provide as much portability in the `ModelCall` as possible across different chat based AI Models. There is a common set of generation options and then those that are specific to a model provider. A sort of "duck typing" approach is used. `ModelOptions` in the model package is a marker interface indicating implementations of this class will provide the options for a model. See `ImageOptions`, a subinterface that defines portable options across all text->image `ImageModel` implementations. Then `StabilityAiImageOptions` and `OpenAiImageOptions` provide the options specific to each model provider. All options classes are created via a fluent API builder all can be passed into the portable `ImageModel` API. These option data types are using in autoconfiguration/configuration properties for the `ImageModel` implementations. === January 13, 2024 Update @@ -79,7 +79,7 @@ Merge SimplePersistentVectorStore and InMemoryVectorStore into SimpleVectorStore Refactor the Ollama client and related classes and package names -* Replace the org.springframework.ai.ollama.client.OllamaClient by org.springframework.ai.ollama.OllamaChatClient. +* Replace the org.springframework.ai.ollama.client.OllamaClient by org.springframework.ai.ollama.OllamaModelCall. * The OllamaChatClient method signatures have changed. * Rename the org.springframework.ai.autoconfigure.ollama.OllamaProperties into org.springframework.ai.autoconfigure.ollama.OllamaChatProperties and change the suffix to: `spring.ai.ollama.chat`. Some of the properties have changed as well. diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/anthropic/AnthropicAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/anthropic/AnthropicAutoConfiguration.java index e94598a8d..6179cb793 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/anthropic/AnthropicAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/anthropic/AnthropicAutoConfiguration.java @@ -17,7 +17,7 @@ package org.springframework.ai.autoconfigure.anthropic; import java.util.List; -import org.springframework.ai.anthropic.AnthropicChatClient; +import org.springframework.ai.anthropic.AnthropicChatModel; import org.springframework.ai.anthropic.api.AnthropicApi; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; import org.springframework.ai.model.function.FunctionCallback; @@ -57,7 +57,7 @@ public class AnthropicAutoConfiguration { @Bean @ConditionalOnMissingBean - public AnthropicChatClient anthropicChatClient(AnthropicApi anthropicApi, AnthropicChatProperties chatProperties, + public AnthropicChatModel anthropicChatModel(AnthropicApi anthropicApi, AnthropicChatProperties chatProperties, RetryTemplate retryTemplate, FunctionCallbackContext functionCallbackContext, List toolFunctionCallbacks) { @@ -65,7 +65,7 @@ public class AnthropicAutoConfiguration { chatProperties.getOptions().getFunctionCallbacks().addAll(toolFunctionCallbacks); } - return new AnthropicChatClient(anthropicApi, chatProperties.getOptions(), retryTemplate, + return new AnthropicChatModel(anthropicApi, chatProperties.getOptions(), retryTemplate, functionCallbackContext); } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/anthropic/AnthropicChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/anthropic/AnthropicChatProperties.java index d9cdc0565..b83ba4540 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/anthropic/AnthropicChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/anthropic/AnthropicChatProperties.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.autoconfigure.anthropic; -import org.springframework.ai.anthropic.AnthropicChatClient; +import org.springframework.ai.anthropic.AnthropicChatModel; import org.springframework.ai.anthropic.AnthropicChatOptions; import org.springframework.boot.context.properties.ConfigurationProperties; import org.springframework.boot.context.properties.NestedConfigurationProperty; @@ -32,7 +32,7 @@ public class AnthropicChatProperties { public static final String CONFIG_PREFIX = "spring.ai.anthropic.chat"; /** - * Enable Anthropic chat client. + * Enable Anthropic chat model. */ private boolean enabled = true; @@ -43,9 +43,9 @@ public class AnthropicChatProperties { */ @NestedConfigurationProperty private AnthropicChatOptions options = AnthropicChatOptions.builder() - .withModel(AnthropicChatClient.DEFAULT_MODEL_NAME) - .withMaxTokens(AnthropicChatClient.DEFAULT_MAX_TOKENS) - .withTemperature(AnthropicChatClient.DEFAULT_TEMPERATURE) + .withModel(AnthropicChatModel.DEFAULT_MODEL_NAME) + .withMaxTokens(AnthropicChatModel.DEFAULT_MAX_TOKENS) + .withTemperature(AnthropicChatModel.DEFAULT_TEMPERATURE) .build(); public AnthropicChatOptions getOptions() { diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/azure/openai/AzureOpenAiAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/azure/openai/AzureOpenAiAutoConfiguration.java index a9b13e304..b4fb7ed4f 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/azure/openai/AzureOpenAiAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/azure/openai/AzureOpenAiAutoConfiguration.java @@ -22,8 +22,8 @@ import com.azure.ai.openai.OpenAIClientBuilder; import com.azure.core.credential.AzureKeyCredential; import com.azure.core.util.ClientOptions; -import org.springframework.ai.azure.openai.AzureOpenAiChatClient; -import org.springframework.ai.azure.openai.AzureOpenAiEmbeddingClient; +import org.springframework.ai.azure.openai.AzureOpenAiChatModel; +import org.springframework.ai.azure.openai.AzureOpenAiEmbeddingModel; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; import org.springframework.boot.autoconfigure.AutoConfiguration; @@ -37,7 +37,7 @@ import org.springframework.util.Assert; import org.springframework.util.CollectionUtils; @AutoConfiguration -@ConditionalOnClass({ OpenAIClientBuilder.class, AzureOpenAiChatClient.class }) +@ConditionalOnClass({ OpenAIClientBuilder.class, AzureOpenAiChatModel.class }) @EnableConfigurationProperties({ AzureOpenAiChatProperties.class, AzureOpenAiEmbeddingProperties.class, AzureOpenAiConnectionProperties.class }) public class AzureOpenAiAutoConfiguration { @@ -58,7 +58,7 @@ public class AzureOpenAiAutoConfiguration { @Bean @ConditionalOnProperty(prefix = AzureOpenAiChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public AzureOpenAiChatClient azureOpenAiChatClient(OpenAIClient openAIClient, + public AzureOpenAiChatModel azureOpenAiChatModel(OpenAIClient openAIClient, AzureOpenAiChatProperties chatProperties, List toolFunctionCallbacks, FunctionCallbackContext functionCallbackContext) { @@ -66,18 +66,18 @@ public class AzureOpenAiAutoConfiguration { chatProperties.getOptions().getFunctionCallbacks().addAll(toolFunctionCallbacks); } - AzureOpenAiChatClient azureOpenAiChatClient = new AzureOpenAiChatClient(openAIClient, - chatProperties.getOptions(), functionCallbackContext); + AzureOpenAiChatModel azureOpenAiChatModel = new AzureOpenAiChatModel(openAIClient, chatProperties.getOptions(), + functionCallbackContext); - return azureOpenAiChatClient; + return azureOpenAiChatModel; } @Bean @ConditionalOnProperty(prefix = AzureOpenAiEmbeddingProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public AzureOpenAiEmbeddingClient azureOpenAiEmbeddingClient(OpenAIClient openAIClient, + public AzureOpenAiEmbeddingModel azureOpenAiEmbeddingModel(OpenAIClient openAIClient, AzureOpenAiEmbeddingProperties embeddingProperties) { - return new AzureOpenAiEmbeddingClient(openAIClient, embeddingProperties.getMetadataMode(), + return new AzureOpenAiEmbeddingModel(openAIClient, embeddingProperties.getMetadataMode(), embeddingProperties.getOptions()); } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/azure/openai/AzureOpenAiChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/azure/openai/AzureOpenAiChatProperties.java index 5ce9c2669..7e974c1a9 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/azure/openai/AzureOpenAiChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/azure/openai/AzureOpenAiChatProperties.java @@ -29,7 +29,7 @@ public class AzureOpenAiChatProperties { private static final Double DEFAULT_TEMPERATURE = 0.7; /** - * Enable Azure OpenAI chat client. + * Enable Azure OpenAI chat model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/azure/openai/AzureOpenAiEmbeddingProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/azure/openai/AzureOpenAiEmbeddingProperties.java index e320eec2f..d7eb357ba 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/azure/openai/AzureOpenAiEmbeddingProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/azure/openai/AzureOpenAiEmbeddingProperties.java @@ -27,7 +27,7 @@ public class AzureOpenAiEmbeddingProperties { public static final String CONFIG_PREFIX = "spring.ai.azure.openai.embedding"; /** - * Enable Azure OpenAI embedding client. + * Enable Azure OpenAI embedding model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic/BedrockAnthropicChatAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic/BedrockAnthropicChatAutoConfiguration.java index 7b3d4d2d5..3e3032454 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic/BedrockAnthropicChatAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic/BedrockAnthropicChatAutoConfiguration.java @@ -18,7 +18,7 @@ package org.springframework.ai.autoconfigure.bedrock.anthropic; import com.fasterxml.jackson.databind.ObjectMapper; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionConfiguration; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.anthropic.BedrockAnthropicChatClient; +import org.springframework.ai.bedrock.anthropic.BedrockAnthropicChatModel; import org.springframework.ai.bedrock.anthropic.api.AnthropicChatBedrockApi; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnBean; @@ -59,9 +59,9 @@ public class BedrockAnthropicChatAutoConfiguration { @Bean @ConditionalOnBean(AnthropicChatBedrockApi.class) - public BedrockAnthropicChatClient anthropicChatClient(AnthropicChatBedrockApi anthropicApi, + public BedrockAnthropicChatModel anthropicChatModel(AnthropicChatBedrockApi anthropicApi, BedrockAnthropicChatProperties properties) { - return new BedrockAnthropicChatClient(anthropicApi, properties.getOptions()); + return new BedrockAnthropicChatModel(anthropicApi, properties.getOptions()); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic/BedrockAnthropicChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic/BedrockAnthropicChatProperties.java index ed1800b5c..e9b263677 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic/BedrockAnthropicChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic/BedrockAnthropicChatProperties.java @@ -35,7 +35,7 @@ public class BedrockAnthropicChatProperties { public static final String CONFIG_PREFIX = "spring.ai.bedrock.anthropic.chat"; /** - * Enable Bedrock Anthropic chat client. Disabled by default. + * Enable Bedrock Anthropic chat model. Disabled by default. */ private boolean enabled = false; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic3/BedrockAnthropic3ChatAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic3/BedrockAnthropic3ChatAutoConfiguration.java index 60e5cdce6..3e53f026b 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic3/BedrockAnthropic3ChatAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic3/BedrockAnthropic3ChatAutoConfiguration.java @@ -18,7 +18,7 @@ package org.springframework.ai.autoconfigure.bedrock.anthropic3; import com.fasterxml.jackson.databind.ObjectMapper; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionConfiguration; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.anthropic3.BedrockAnthropic3ChatClient; +import org.springframework.ai.bedrock.anthropic3.BedrockAnthropic3ChatModel; import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnBean; @@ -59,9 +59,9 @@ public class BedrockAnthropic3ChatAutoConfiguration { @Bean @ConditionalOnBean(Anthropic3ChatBedrockApi.class) - public BedrockAnthropic3ChatClient anthropic3ChatClient(Anthropic3ChatBedrockApi anthropicApi, + public BedrockAnthropic3ChatModel anthropic3ChatModel(Anthropic3ChatBedrockApi anthropicApi, BedrockAnthropic3ChatProperties properties) { - return new BedrockAnthropic3ChatClient(anthropicApi, properties.getOptions()); + return new BedrockAnthropic3ChatModel(anthropicApi, properties.getOptions()); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic3/BedrockAnthropic3ChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic3/BedrockAnthropic3ChatProperties.java index 54d516239..71086b0d6 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic3/BedrockAnthropic3ChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/anthropic3/BedrockAnthropic3ChatProperties.java @@ -34,7 +34,7 @@ public class BedrockAnthropic3ChatProperties { public static final String CONFIG_PREFIX = "spring.ai.bedrock.anthropic3.chat"; /** - * Enable Bedrock Anthropic chat client. Disabled by default. + * Enable Bedrock Anthropic chat model. Disabled by default. */ private boolean enabled = false; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereChatAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereChatAutoConfiguration.java index 66b6d5e5e..896078e5b 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereChatAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereChatAutoConfiguration.java @@ -18,7 +18,7 @@ package org.springframework.ai.autoconfigure.bedrock.cohere; import com.fasterxml.jackson.databind.ObjectMapper; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionConfiguration; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.cohere.BedrockCohereChatClient; +import org.springframework.ai.bedrock.cohere.BedrockCohereChatModel; import org.springframework.ai.bedrock.cohere.api.CohereChatBedrockApi; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnBean; @@ -57,10 +57,10 @@ public class BedrockCohereChatAutoConfiguration { @Bean @ConditionalOnBean(CohereChatBedrockApi.class) - public BedrockCohereChatClient cohereChatClient(CohereChatBedrockApi cohereChatApi, + public BedrockCohereChatModel cohereChatModel(CohereChatBedrockApi cohereChatApi, BedrockCohereChatProperties properties) { - return new BedrockCohereChatClient(cohereChatApi, properties.getOptions()); + return new BedrockCohereChatModel(cohereChatApi, properties.getOptions()); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereEmbeddingAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereEmbeddingAutoConfiguration.java index 76e3ea603..82b6292de 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereEmbeddingAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereEmbeddingAutoConfiguration.java @@ -21,7 +21,7 @@ import software.amazon.awssdk.regions.providers.AwsRegionProvider; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionConfiguration; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.cohere.BedrockCohereEmbeddingClient; +import org.springframework.ai.bedrock.cohere.BedrockCohereEmbeddingModel; import org.springframework.ai.bedrock.cohere.api.CohereEmbeddingBedrockApi; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnBean; @@ -59,10 +59,10 @@ public class BedrockCohereEmbeddingAutoConfiguration { @Bean @ConditionalOnMissingBean @ConditionalOnBean(CohereEmbeddingBedrockApi.class) - public BedrockCohereEmbeddingClient cohereEmbeddingClient(CohereEmbeddingBedrockApi cohereEmbeddingApi, + public BedrockCohereEmbeddingModel cohereEmbeddingModel(CohereEmbeddingBedrockApi cohereEmbeddingApi, BedrockCohereEmbeddingProperties properties) { - return new BedrockCohereEmbeddingClient(cohereEmbeddingApi, properties.getOptions()); + return new BedrockCohereEmbeddingModel(cohereEmbeddingApi, properties.getOptions()); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/jurrasic2/BedrockAi21Jurassic2ChatAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/jurrasic2/BedrockAi21Jurassic2ChatAutoConfiguration.java index e8266a0e4..8ad3c0bb1 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/jurrasic2/BedrockAi21Jurassic2ChatAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/jurrasic2/BedrockAi21Jurassic2ChatAutoConfiguration.java @@ -19,7 +19,7 @@ package org.springframework.ai.autoconfigure.bedrock.jurrasic2; import com.fasterxml.jackson.databind.ObjectMapper; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionConfiguration; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.jurassic2.BedrockAi21Jurassic2ChatClient; +import org.springframework.ai.bedrock.jurassic2.BedrockAi21Jurassic2ChatModel; import org.springframework.ai.bedrock.jurassic2.api.Ai21Jurassic2ChatBedrockApi; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnBean; @@ -59,10 +59,10 @@ public class BedrockAi21Jurassic2ChatAutoConfiguration { @Bean @ConditionalOnBean(Ai21Jurassic2ChatBedrockApi.class) - public BedrockAi21Jurassic2ChatClient jurassic2ChatClient(Ai21Jurassic2ChatBedrockApi ai21Jurassic2ChatBedrockApi, + public BedrockAi21Jurassic2ChatModel jurassic2ChatModel(Ai21Jurassic2ChatBedrockApi ai21Jurassic2ChatBedrockApi, BedrockAi21Jurassic2ChatProperties properties) { - return BedrockAi21Jurassic2ChatClient.builder(ai21Jurassic2ChatBedrockApi) + return BedrockAi21Jurassic2ChatModel.builder(ai21Jurassic2ChatBedrockApi) .withOptions(properties.getOptions()) .build(); } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/jurrasic2/BedrockAi21Jurassic2ChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/jurrasic2/BedrockAi21Jurassic2ChatProperties.java index 217c7c7eb..eccd7e0c9 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/jurrasic2/BedrockAi21Jurassic2ChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/jurrasic2/BedrockAi21Jurassic2ChatProperties.java @@ -33,7 +33,7 @@ public class BedrockAi21Jurassic2ChatProperties { public static final String CONFIG_PREFIX = "spring.ai.bedrock.jurassic2.chat"; /** - * Enable Bedrock Ai21Jurassic2 chat client. Disabled by default. + * Enable Bedrock Ai21Jurassic2 chat model. Disabled by default. */ private boolean enabled = false; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/llama/BedrockLlamaChatAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/llama/BedrockLlamaChatAutoConfiguration.java index 9293acc84..6e105b8f2 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/llama/BedrockLlamaChatAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/llama/BedrockLlamaChatAutoConfiguration.java @@ -16,12 +16,12 @@ package org.springframework.ai.autoconfigure.bedrock.llama; import com.fasterxml.jackson.databind.ObjectMapper; +import org.springframework.ai.bedrock.llama.BedrockLlamaChatModel; import software.amazon.awssdk.auth.credentials.AwsCredentialsProvider; import software.amazon.awssdk.regions.providers.AwsRegionProvider; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionConfiguration; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.llama.BedrockLlamaChatClient; import org.springframework.ai.bedrock.llama.api.LlamaChatBedrockApi; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnBean; @@ -59,9 +59,9 @@ public class BedrockLlamaChatAutoConfiguration { @Bean @ConditionalOnBean(LlamaChatBedrockApi.class) - public BedrockLlamaChatClient llamaChatClient(LlamaChatBedrockApi llamaApi, BedrockLlamaChatProperties properties) { + public BedrockLlamaChatModel llamaChatModel(LlamaChatBedrockApi llamaApi, BedrockLlamaChatProperties properties) { - return new BedrockLlamaChatClient(llamaApi, properties.getOptions()); + return new BedrockLlamaChatModel(llamaApi, properties.getOptions()); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/llama/BedrockLlamaChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/llama/BedrockLlamaChatProperties.java index f93b65aca..048b7dde2 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/llama/BedrockLlamaChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/llama/BedrockLlamaChatProperties.java @@ -32,7 +32,7 @@ public class BedrockLlamaChatProperties { public static final String CONFIG_PREFIX = "spring.ai.bedrock.llama.chat"; /** - * Enable Bedrock Llama chat client. Disabled by default. + * Enable Bedrock Llama chat model. Disabled by default. */ private boolean enabled = false; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanChatAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanChatAutoConfiguration.java index 67995b9e3..0115967fe 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanChatAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanChatAutoConfiguration.java @@ -18,7 +18,7 @@ package org.springframework.ai.autoconfigure.bedrock.titan; import com.fasterxml.jackson.databind.ObjectMapper; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionConfiguration; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.titan.BedrockTitanChatClient; +import org.springframework.ai.bedrock.titan.BedrockTitanChatModel; import org.springframework.ai.bedrock.titan.api.TitanChatBedrockApi; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnBean; @@ -57,10 +57,10 @@ public class BedrockTitanChatAutoConfiguration { @Bean @ConditionalOnBean(TitanChatBedrockApi.class) - public BedrockTitanChatClient titanChatClient(TitanChatBedrockApi titanChatApi, + public BedrockTitanChatModel titanChatModel(TitanChatBedrockApi titanChatApi, BedrockTitanChatProperties properties) { - return new BedrockTitanChatClient(titanChatApi, properties.getOptions()); + return new BedrockTitanChatModel(titanChatApi, properties.getOptions()); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanEmbeddingAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanEmbeddingAutoConfiguration.java index 5ea79d451..b019dc1c6 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanEmbeddingAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanEmbeddingAutoConfiguration.java @@ -21,7 +21,7 @@ import software.amazon.awssdk.regions.providers.AwsRegionProvider; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionConfiguration; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.titan.BedrockTitanEmbeddingClient; +import org.springframework.ai.bedrock.titan.BedrockTitanEmbeddingModel; import org.springframework.ai.bedrock.titan.api.TitanEmbeddingBedrockApi; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnBean; @@ -59,10 +59,10 @@ public class BedrockTitanEmbeddingAutoConfiguration { @Bean @ConditionalOnMissingBean @ConditionalOnBean(TitanEmbeddingBedrockApi.class) - public BedrockTitanEmbeddingClient titanEmbeddingClient(TitanEmbeddingBedrockApi titanEmbeddingApi, + public BedrockTitanEmbeddingModel titanEmbeddingModel(TitanEmbeddingBedrockApi titanEmbeddingApi, BedrockTitanEmbeddingProperties properties) { - return new BedrockTitanEmbeddingClient(titanEmbeddingApi).withInputType(properties.getInputType()); + return new BedrockTitanEmbeddingModel(titanEmbeddingApi).withInputType(properties.getInputType()); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanEmbeddingProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanEmbeddingProperties.java index 3e80cfed8..9d9b1bd6e 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanEmbeddingProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanEmbeddingProperties.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.autoconfigure.bedrock.titan; -import org.springframework.ai.bedrock.titan.BedrockTitanEmbeddingClient.InputType; +import org.springframework.ai.bedrock.titan.BedrockTitanEmbeddingModel.InputType; import org.springframework.ai.bedrock.titan.api.TitanEmbeddingBedrockApi.TitanEmbeddingModel; import org.springframework.boot.context.properties.ConfigurationProperties; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/huggingface/HuggingfaceChatAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/huggingface/HuggingfaceChatAutoConfiguration.java index 42c3b1e73..ef28fe2b3 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/huggingface/HuggingfaceChatAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/huggingface/HuggingfaceChatAutoConfiguration.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.autoconfigure.huggingface; -import org.springframework.ai.huggingface.HuggingfaceChatClient; +import org.springframework.ai.huggingface.HuggingfaceChatModel; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean; @@ -24,7 +24,7 @@ import org.springframework.boot.context.properties.EnableConfigurationProperties import org.springframework.context.annotation.Bean; @AutoConfiguration -@ConditionalOnClass(HuggingfaceChatClient.class) +@ConditionalOnClass(HuggingfaceChatModel.class) @EnableConfigurationProperties(HuggingfaceChatProperties.class) public class HuggingfaceChatAutoConfiguration { @@ -32,8 +32,8 @@ public class HuggingfaceChatAutoConfiguration { @ConditionalOnMissingBean @ConditionalOnProperty(prefix = HuggingfaceChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public HuggingfaceChatClient huggingfaceChatClient(HuggingfaceChatProperties huggingfaceChatProperties) { - return new HuggingfaceChatClient(huggingfaceChatProperties.getApiKey(), huggingfaceChatProperties.getUrl()); + public HuggingfaceChatModel huggingfaceChatModel(HuggingfaceChatProperties huggingfaceChatProperties) { + return new HuggingfaceChatModel(huggingfaceChatProperties.getApiKey(), huggingfaceChatProperties.getUrl()); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/huggingface/HuggingfaceChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/huggingface/HuggingfaceChatProperties.java index 4495a5123..d31bef3fb 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/huggingface/HuggingfaceChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/huggingface/HuggingfaceChatProperties.java @@ -27,7 +27,7 @@ public class HuggingfaceChatProperties { private String url; /** - * Enable Huggingface chat client. + * Enable Huggingface chat model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfiguration.java index f3c805db2..80cde6dfe 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfiguration.java @@ -16,8 +16,8 @@ package org.springframework.ai.autoconfigure.minimax; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; -import org.springframework.ai.minimax.MiniMaxChatClient; -import org.springframework.ai.minimax.MiniMaxEmbeddingClient; +import org.springframework.ai.minimax.MiniMaxChatModel; +import org.springframework.ai.minimax.MiniMaxEmbeddingModel; import org.springframework.ai.minimax.api.MiniMaxApi; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; @@ -51,7 +51,7 @@ public class MiniMaxAutoConfiguration { @ConditionalOnMissingBean @ConditionalOnProperty(prefix = MiniMaxChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public MiniMaxChatClient miniMaxChatClient(MiniMaxConnectionProperties commonProperties, + public MiniMaxChatModel miniMaxChatModel(MiniMaxConnectionProperties commonProperties, MiniMaxChatProperties chatProperties, RestClient.Builder restClientBuilder, List toolFunctionCallbacks, FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { @@ -63,21 +63,21 @@ public class MiniMaxAutoConfiguration { chatProperties.getOptions().getFunctionCallbacks().addAll(toolFunctionCallbacks); } - return new MiniMaxChatClient(miniMaxApi, chatProperties.getOptions(), functionCallbackContext, retryTemplate); + return new MiniMaxChatModel(miniMaxApi, chatProperties.getOptions(), functionCallbackContext, retryTemplate); } @Bean @ConditionalOnMissingBean @ConditionalOnProperty(prefix = MiniMaxEmbeddingProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public MiniMaxEmbeddingClient miniMaxEmbeddingClient(MiniMaxConnectionProperties commonProperties, + public MiniMaxEmbeddingModel miniMaxEmbeddingModel(MiniMaxConnectionProperties commonProperties, MiniMaxEmbeddingProperties embeddingProperties, RestClient.Builder restClientBuilder, RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { var miniMaxApi = miniMaxApi(embeddingProperties.getBaseUrl(), commonProperties.getBaseUrl(), embeddingProperties.getApiKey(), commonProperties.getApiKey(), restClientBuilder, responseErrorHandler); - return new MiniMaxEmbeddingClient(miniMaxApi, embeddingProperties.getMetadataMode(), + return new MiniMaxEmbeddingModel(miniMaxApi, embeddingProperties.getMetadataMode(), embeddingProperties.getOptions(), retryTemplate); } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxChatProperties.java index c7f3716f3..e32748b30 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxChatProperties.java @@ -33,7 +33,7 @@ public class MiniMaxChatProperties extends MiniMaxParentProperties { private static final Double DEFAULT_TEMPERATURE = 0.7; /** - * Enable MiniMax chat client. + * Enable MiniMax chat model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxEmbeddingProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxEmbeddingProperties.java index 21fbb752e..bfdb49174 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxEmbeddingProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/minimax/MiniMaxEmbeddingProperties.java @@ -32,7 +32,7 @@ public class MiniMaxEmbeddingProperties extends MiniMaxParentProperties { public static final String DEFAULT_EMBEDDING_MODEL = MiniMaxApi.EmbeddingModel.Embo_01.value; /** - * Enable MiniMax embedding client. + * Enable MiniMax embedding model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/mistralai/MistralAiAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/mistralai/MistralAiAutoConfiguration.java index 1328f3488..abc3104b2 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/mistralai/MistralAiAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/mistralai/MistralAiAutoConfiguration.java @@ -18,8 +18,8 @@ package org.springframework.ai.autoconfigure.mistralai; import java.util.List; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; -import org.springframework.ai.mistralai.MistralAiChatClient; -import org.springframework.ai.mistralai.MistralAiEmbeddingClient; +import org.springframework.ai.mistralai.MistralAiChatModel; +import org.springframework.ai.mistralai.MistralAiEmbeddingModel; import org.springframework.ai.mistralai.api.MistralAiApi; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; @@ -53,7 +53,7 @@ public class MistralAiAutoConfiguration { @ConditionalOnMissingBean @ConditionalOnProperty(prefix = MistralAiEmbeddingProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public MistralAiEmbeddingClient mistralAiEmbeddingClient(MistralAiCommonProperties commonProperties, + public MistralAiEmbeddingModel mistralAiEmbeddingModel(MistralAiCommonProperties commonProperties, MistralAiEmbeddingProperties embeddingProperties, RestClient.Builder restClientBuilder, RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { @@ -61,7 +61,7 @@ public class MistralAiAutoConfiguration { embeddingProperties.getBaseUrl(), commonProperties.getBaseUrl(), restClientBuilder, responseErrorHandler); - return new MistralAiEmbeddingClient(mistralAiApi, embeddingProperties.getMetadataMode(), + return new MistralAiEmbeddingModel(mistralAiApi, embeddingProperties.getMetadataMode(), embeddingProperties.getOptions(), retryTemplate); } @@ -69,7 +69,7 @@ public class MistralAiAutoConfiguration { @ConditionalOnMissingBean @ConditionalOnProperty(prefix = MistralAiChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public MistralAiChatClient mistralAiChatClient(MistralAiCommonProperties commonProperties, + public MistralAiChatModel mistralAiChatModel(MistralAiCommonProperties commonProperties, MistralAiChatProperties chatProperties, RestClient.Builder restClientBuilder, List toolFunctionCallbacks, FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { @@ -81,7 +81,7 @@ public class MistralAiAutoConfiguration { chatProperties.getOptions().getFunctionCallbacks().addAll(toolFunctionCallbacks); } - return new MistralAiChatClient(mistralAiApi, chatProperties.getOptions(), functionCallbackContext, + return new MistralAiChatModel(mistralAiApi, chatProperties.getOptions(), functionCallbackContext, retryTemplate); } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/mistralai/MistralAiChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/mistralai/MistralAiChatProperties.java index 4f9e00401..79b05a81f 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/mistralai/MistralAiChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/mistralai/MistralAiChatProperties.java @@ -43,7 +43,7 @@ public class MistralAiChatProperties extends MistralAiParentProperties { } /** - * Enable OpenAI chat client. + * Enable OpenAI chat model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/mistralai/MistralAiEmbeddingProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/mistralai/MistralAiEmbeddingProperties.java index dbb5d8f3c..450ac479a 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/mistralai/MistralAiEmbeddingProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/mistralai/MistralAiEmbeddingProperties.java @@ -35,7 +35,7 @@ public class MistralAiEmbeddingProperties extends MistralAiParentProperties { public static final String DEFAULT_ENCODING_FORMAT = "float"; /** - * Enable MistralAI embedding client. + * Enable MistralAI embedding model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/ollama/OllamaAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/ollama/OllamaAutoConfiguration.java index 8b788e73b..0839b271e 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/ollama/OllamaAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/ollama/OllamaAutoConfiguration.java @@ -15,8 +15,8 @@ */ package org.springframework.ai.autoconfigure.ollama; -import org.springframework.ai.ollama.OllamaChatClient; -import org.springframework.ai.ollama.OllamaEmbeddingClient; +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.condition.ConditionalOnClass; @@ -56,17 +56,17 @@ public class OllamaAutoConfiguration { @ConditionalOnMissingBean @ConditionalOnProperty(prefix = OllamaChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public OllamaChatClient ollamaChatClient(OllamaApi ollamaApi, OllamaChatProperties properties) { - return new OllamaChatClient(ollamaApi, properties.getOptions()); + public OllamaChatModel ollamaChatModel(OllamaApi ollamaApi, OllamaChatProperties properties) { + return new OllamaChatModel(ollamaApi, properties.getOptions()); } @Bean @ConditionalOnMissingBean @ConditionalOnProperty(prefix = OllamaEmbeddingProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public OllamaEmbeddingClient ollamaEmbeddingClient(OllamaApi ollamaApi, OllamaEmbeddingProperties properties) { + public OllamaEmbeddingModel ollamaEmbeddingModel(OllamaApi ollamaApi, OllamaEmbeddingProperties properties) { - return new OllamaEmbeddingClient(ollamaApi, properties.getOptions()); + return new OllamaEmbeddingModel(ollamaApi, properties.getOptions()); } private static class PropertiesOllamaConnectionDetails implements OllamaConnectionDetails { diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/ollama/OllamaChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/ollama/OllamaChatProperties.java index f3bc8ad78..2439c83e8 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/ollama/OllamaChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/ollama/OllamaChatProperties.java @@ -31,7 +31,7 @@ public class OllamaChatProperties { public static final String CONFIG_PREFIX = "spring.ai.ollama.chat"; /** - * Enable Ollama chat client. + * Enable Ollama chat model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/ollama/OllamaEmbeddingProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/ollama/OllamaEmbeddingProperties.java index 69f809ea1..a2368cd2e 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/ollama/OllamaEmbeddingProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/ollama/OllamaEmbeddingProperties.java @@ -31,7 +31,7 @@ public class OllamaEmbeddingProperties { public static final String CONFIG_PREFIX = "spring.ai.ollama.embedding"; /** - * Enable Ollama embedding client. + * Enable Ollama embedding model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfiguration.java index 848456468..3a601abb2 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfiguration.java @@ -20,11 +20,8 @@ import java.util.List; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; -import org.springframework.ai.openai.OpenAiAudioTranscriptionClient; -import org.springframework.ai.openai.OpenAiChatClient; -import org.springframework.ai.openai.OpenAiEmbeddingClient; -import org.springframework.ai.openai.OpenAiImageClient; -import org.springframework.ai.openai.OpenAiAudioSpeechClient; +import org.springframework.ai.openai.*; +import org.springframework.ai.openai.OpenAiChatModel; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.openai.api.OpenAiAudioApi; import org.springframework.ai.openai.api.OpenAiImageApi; @@ -57,7 +54,7 @@ public class OpenAiAutoConfiguration { @ConditionalOnMissingBean @ConditionalOnProperty(prefix = OpenAiChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public OpenAiChatClient openAiChatClient(OpenAiConnectionProperties commonProperties, + public OpenAiChatModel openAiChatModel(OpenAiConnectionProperties commonProperties, OpenAiChatProperties chatProperties, RestClient.Builder restClientBuilder, List toolFunctionCallbacks, FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { @@ -69,21 +66,21 @@ public class OpenAiAutoConfiguration { chatProperties.getOptions().getFunctionCallbacks().addAll(toolFunctionCallbacks); } - return new OpenAiChatClient(openAiApi, chatProperties.getOptions(), functionCallbackContext, retryTemplate); + return new OpenAiChatModel(openAiApi, chatProperties.getOptions(), functionCallbackContext, retryTemplate); } @Bean @ConditionalOnMissingBean @ConditionalOnProperty(prefix = OpenAiEmbeddingProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public OpenAiEmbeddingClient openAiEmbeddingClient(OpenAiConnectionProperties commonProperties, + public OpenAiEmbeddingModel openAiEmbeddingModel(OpenAiConnectionProperties commonProperties, OpenAiEmbeddingProperties embeddingProperties, RestClient.Builder restClientBuilder, RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { var openAiApi = openAiApi(embeddingProperties.getBaseUrl(), commonProperties.getBaseUrl(), embeddingProperties.getApiKey(), commonProperties.getApiKey(), restClientBuilder, responseErrorHandler); - return new OpenAiEmbeddingClient(openAiApi, embeddingProperties.getMetadataMode(), + return new OpenAiEmbeddingModel(openAiApi, embeddingProperties.getMetadataMode(), embeddingProperties.getOptions(), retryTemplate); } @@ -103,7 +100,7 @@ public class OpenAiAutoConfiguration { @ConditionalOnMissingBean @ConditionalOnProperty(prefix = OpenAiImageProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public OpenAiImageClient openAiImageClient(OpenAiConnectionProperties commonProperties, + public OpenAiImageModel openAiImageModel(OpenAiConnectionProperties commonProperties, OpenAiImageProperties imageProperties, RestClient.Builder restClientBuilder, RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { @@ -118,12 +115,12 @@ public class OpenAiAutoConfiguration { var openAiImageApi = new OpenAiImageApi(baseUrl, apiKey, restClientBuilder, responseErrorHandler); - return new OpenAiImageClient(openAiImageApi, imageProperties.getOptions(), retryTemplate); + return new OpenAiImageModel(openAiImageApi, imageProperties.getOptions(), retryTemplate); } @Bean @ConditionalOnMissingBean - public OpenAiAudioTranscriptionClient openAiAudioTranscriptionClient(OpenAiConnectionProperties commonProperties, + public OpenAiAudioTranscriptionModel openAiAudioTranscriptionModel(OpenAiConnectionProperties commonProperties, OpenAiAudioTranscriptionProperties transcriptionProperties, RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { @@ -138,15 +135,15 @@ public class OpenAiAutoConfiguration { var openAiAudioApi = new OpenAiAudioApi(baseUrl, apiKey, RestClient.builder(), responseErrorHandler); - OpenAiAudioTranscriptionClient openAiChatClient = new OpenAiAudioTranscriptionClient(openAiAudioApi, + OpenAiAudioTranscriptionModel openAiChatModel = new OpenAiAudioTranscriptionModel(openAiAudioApi, transcriptionProperties.getOptions(), retryTemplate); - return openAiChatClient; + return openAiChatModel; } @Bean @ConditionalOnMissingBean - public OpenAiAudioSpeechClient openAiAudioSpeechClient(OpenAiConnectionProperties commonProperties, + public OpenAiAudioSpeechModel openAiAudioSpeechModel(OpenAiConnectionProperties commonProperties, OpenAiAudioSpeechProperties speechProperties, ResponseErrorHandler responseErrorHandler) { String apiKey = StringUtils.hasText(speechProperties.getApiKey()) ? speechProperties.getApiKey() @@ -160,10 +157,10 @@ public class OpenAiAutoConfiguration { var openAiAudioApi = new OpenAiAudioApi(baseUrl, apiKey, RestClient.builder(), responseErrorHandler); - OpenAiAudioSpeechClient openAiSpeechClient = new OpenAiAudioSpeechClient(openAiAudioApi, + OpenAiAudioSpeechModel openAiSpeechModel = new OpenAiAudioSpeechModel(openAiAudioApi, speechProperties.getOptions()); - return openAiSpeechClient; + return openAiSpeechModel; } @Bean diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiChatProperties.java index 5e9940b42..41524ed39 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiChatProperties.java @@ -29,7 +29,7 @@ public class OpenAiChatProperties extends OpenAiParentProperties { private static final Double DEFAULT_TEMPERATURE = 0.7; /** - * Enable OpenAI chat client. + * Enable OpenAI chat model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiEmbeddingProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiEmbeddingProperties.java index fa796d92f..5901d4013 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiEmbeddingProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiEmbeddingProperties.java @@ -28,7 +28,7 @@ public class OpenAiEmbeddingProperties extends OpenAiParentProperties { public static final String DEFAULT_EMBEDDING_MODEL = "text-embedding-ada-002"; /** - * Enable OpenAI embedding client. + * Enable OpenAI embedding model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiImageProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiImageProperties.java index aa16503fa..06fb24bf6 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiImageProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/openai/OpenAiImageProperties.java @@ -34,7 +34,7 @@ public class OpenAiImageProperties extends OpenAiParentProperties { public static final String DEFAULT_IMAGE_MODEL = OpenAiImageApi.ImageModel.DALL_E_3.getValue(); /** - * Enable OpenAI Image client. + * Enable OpenAI image model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlAutoConfiguration.java index 78ac62dc4..ca30501b5 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlAutoConfiguration.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.autoconfigure.postgresml; -import org.springframework.ai.postgresml.PostgresMlEmbeddingClient; +import org.springframework.ai.postgresml.PostgresMlEmbeddingModel; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean; @@ -26,13 +26,13 @@ import org.springframework.context.annotation.Bean; import org.springframework.jdbc.core.JdbcTemplate; /** - * Auto-configuration class for PostgresMlEmbeddingClient. + * Auto-configuration class for PostgresMlEmbeddingModel. * * @author Utkarsh Srivastava * @author Christian Tzolov */ @AutoConfiguration(after = JdbcTemplateAutoConfiguration.class) -@ConditionalOnClass(PostgresMlEmbeddingClient.class) +@ConditionalOnClass(PostgresMlEmbeddingModel.class) @EnableConfigurationProperties(PostgresMlEmbeddingProperties.class) public class PostgresMlAutoConfiguration { @@ -40,10 +40,10 @@ public class PostgresMlAutoConfiguration { @ConditionalOnMissingBean @ConditionalOnProperty(prefix = PostgresMlEmbeddingProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public PostgresMlEmbeddingClient postgresMlEmbeddingClient(JdbcTemplate jdbcTemplate, + public PostgresMlEmbeddingModel postgresMlEmbeddingModel(JdbcTemplate jdbcTemplate, PostgresMlEmbeddingProperties embeddingProperties) { - return new PostgresMlEmbeddingClient(jdbcTemplate, embeddingProperties.getOptions()); + return new PostgresMlEmbeddingModel(jdbcTemplate, embeddingProperties.getOptions()); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlEmbeddingProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlEmbeddingProperties.java index c0b13b540..53dba7f93 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlEmbeddingProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlEmbeddingProperties.java @@ -18,7 +18,7 @@ package org.springframework.ai.autoconfigure.postgresml; import java.util.Map; import org.springframework.ai.document.MetadataMode; -import org.springframework.ai.postgresml.PostgresMlEmbeddingClient; +import org.springframework.ai.postgresml.PostgresMlEmbeddingModel; import org.springframework.ai.postgresml.PostgresMlEmbeddingOptions; import org.springframework.boot.context.properties.ConfigurationProperties; import org.springframework.boot.context.properties.NestedConfigurationProperty; @@ -36,14 +36,14 @@ public class PostgresMlEmbeddingProperties { public static final String CONFIG_PREFIX = "spring.ai.postgresml.embedding"; /** - * Enable Postgres ML embedding client. + * Enable Postgres ML embedding model. */ private boolean enabled = true; @NestedConfigurationProperty private PostgresMlEmbeddingOptions options = PostgresMlEmbeddingOptions.builder() - .withTransformer(PostgresMlEmbeddingClient.DEFAULT_TRANSFORMER_MODEL) - .withVectorType(PostgresMlEmbeddingClient.VectorType.PG_ARRAY) + .withTransformer(PostgresMlEmbeddingModel.DEFAULT_TRANSFORMER_MODEL) + .withVectorType(PostgresMlEmbeddingModel.VectorType.PG_ARRAY) .withKwargs(Map.of()) .withMetadataMode(MetadataMode.EMBED) .build(); diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageAutoConfiguration.java index fa1660983..5983499f2 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageAutoConfiguration.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.autoconfigure.stabilityai; -import org.springframework.ai.stabilityai.StabilityAiImageClient; +import org.springframework.ai.stabilityai.StabilityAiImageModel; import org.springframework.ai.stabilityai.api.StabilityAiApi; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; @@ -58,9 +58,9 @@ public class StabilityAiImageAutoConfiguration { @ConditionalOnMissingBean @ConditionalOnProperty(prefix = StabilityAiImageProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public StabilityAiImageClient stabilityAiImageClient(StabilityAiApi stabilityAiApi, + public StabilityAiImageModel stabilityAiImageModel(StabilityAiApi stabilityAiApi, StabilityAiImageProperties stabilityAiImageProperties) { - return new StabilityAiImageClient(stabilityAiApi, stabilityAiImageProperties.getOptions()); + return new StabilityAiImageModel(stabilityAiApi, stabilityAiImageProperties.getOptions()); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageProperties.java index 4b81fe9e7..d307750df 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImageProperties.java @@ -30,7 +30,7 @@ public class StabilityAiImageProperties extends StabilityAiParentProperties { public static final String CONFIG_PREFIX = "spring.ai.stabilityai.image"; /** - * Enable Stability Image client. + * Enable Stability image model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingClientAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingModelAutoConfiguration.java similarity index 60% rename from spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingClientAutoConfiguration.java rename to spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingModelAutoConfiguration.java index adbacb01d..71ec54d5f 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingClientAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingModelAutoConfiguration.java @@ -18,7 +18,7 @@ package org.springframework.ai.autoconfigure.transformers; import ai.djl.huggingface.tokenizers.HuggingFaceTokenizer; import ai.onnxruntime.OrtSession; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean; @@ -30,29 +30,29 @@ import org.springframework.context.annotation.Bean; * @author Christian Tzolov */ @AutoConfiguration -@EnableConfigurationProperties({ TransformersEmbeddingClientProperties.class }) -@ConditionalOnClass({ OrtSession.class, HuggingFaceTokenizer.class, TransformersEmbeddingClient.class }) -public class TransformersEmbeddingClientAutoConfiguration { +@EnableConfigurationProperties({ TransformersEmbeddingModelProperties.class }) +@ConditionalOnClass({ OrtSession.class, HuggingFaceTokenizer.class, TransformersEmbeddingModel.class }) +public class TransformersEmbeddingModelAutoConfiguration { @Bean @ConditionalOnMissingBean - @ConditionalOnProperty(prefix = TransformersEmbeddingClientProperties.CONFIG_PREFIX, name = "enabled", + @ConditionalOnProperty(prefix = TransformersEmbeddingModelProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public TransformersEmbeddingClient embeddingClient(TransformersEmbeddingClientProperties properties) { + public TransformersEmbeddingModel embeddingModel(TransformersEmbeddingModelProperties properties) { - TransformersEmbeddingClient embeddingClient = new TransformersEmbeddingClient(properties.getMetadataMode()); + TransformersEmbeddingModel embeddingModel = new TransformersEmbeddingModel(properties.getMetadataMode()); - embeddingClient.setDisableCaching(!properties.getCache().isEnabled()); - embeddingClient.setResourceCacheDirectory(properties.getCache().getDirectory()); + embeddingModel.setDisableCaching(!properties.getCache().isEnabled()); + embeddingModel.setResourceCacheDirectory(properties.getCache().getDirectory()); - embeddingClient.setTokenizerResource(properties.getTokenizer().getUri()); - embeddingClient.setTokenizerOptions(properties.getTokenizer().getOptions()); + embeddingModel.setTokenizerResource(properties.getTokenizer().getUri()); + embeddingModel.setTokenizerOptions(properties.getTokenizer().getOptions()); - embeddingClient.setModelResource(properties.getOnnx().getModelUri()); + embeddingModel.setModelResource(properties.getOnnx().getModelUri()); - embeddingClient.setGpuDeviceId(properties.getOnnx().getGpuDeviceId()); + embeddingModel.setGpuDeviceId(properties.getOnnx().getGpuDeviceId()); - return embeddingClient; + return embeddingModel; } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingClientProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingModelProperties.java similarity index 90% rename from spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingClientProperties.java rename to spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingModelProperties.java index 230f86b5a..0c0344eaf 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingClientProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingModelProperties.java @@ -24,17 +24,17 @@ import ai.djl.huggingface.tokenizers.HuggingFaceTokenizer; import org.springframework.ai.document.Document; import org.springframework.ai.document.MetadataMode; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.boot.context.properties.ConfigurationProperties; import org.springframework.boot.context.properties.NestedConfigurationProperty; -import static org.springframework.ai.autoconfigure.transformers.TransformersEmbeddingClientProperties.CONFIG_PREFIX; +import static org.springframework.ai.autoconfigure.transformers.TransformersEmbeddingModelProperties.CONFIG_PREFIX; /** * @author Christian Tzolov */ @ConfigurationProperties(CONFIG_PREFIX) -public class TransformersEmbeddingClientProperties { +public class TransformersEmbeddingModelProperties { public static final String CONFIG_PREFIX = "spring.ai.embedding.transformer"; @@ -65,7 +65,7 @@ public class TransformersEmbeddingClientProperties { * URI of a pre-trained HuggingFaceTokenizer created by the ONNX engine (e.g. * tokenizer.json). */ - private String uri = TransformersEmbeddingClient.DEFAULT_ONNX_TOKENIZER_URI; + private String uri = TransformersEmbeddingModel.DEFAULT_ONNX_TOKENIZER_URI; /** * HuggingFaceTokenizer options such as 'addSpecialTokens', 'modelMaxLength', @@ -145,12 +145,12 @@ public class TransformersEmbeddingClientProperties { * https://sbert.net/docs/pretrained_models.html. Defaults to * sentence-transformers/all-MiniLM-L6-v2. */ - private String modelUri = TransformersEmbeddingClient.DEFAULT_ONNX_MODEL_URI; + private String modelUri = TransformersEmbeddingModel.DEFAULT_ONNX_MODEL_URI; /** * Defaults to: 'last_hidden_state'. */ - private String modelOutputName = TransformersEmbeddingClient.DEFAULT_MODEL_OUTPUT_NAME; + private String modelOutputName = TransformersEmbeddingModel.DEFAULT_MODEL_OUTPUT_NAME; /** * Run on a GPU or with another provider (optional). @@ -196,9 +196,9 @@ public class TransformersEmbeddingClientProperties { /** * Specifies what parts of the {@link Document}'s content and metadata will be used * for computing the embeddings. Applicable for the - * {@link TransformersEmbeddingClient#embed(Document)} method only. Has no effect on - * the {@link TransformersEmbeddingClient#embed(String)} or - * {@link TransformersEmbeddingClient#embed(List)}. Defaults to + * {@link TransformersEmbeddingModel#embed(Document)} method only. Has no effect on + * the {@link TransformersEmbeddingModel#embed(String)} or + * {@link TransformersEmbeddingModel#embed(List)}. Defaults to * {@link MetadataMode#NONE}. */ private MetadataMode metadataMode = MetadataMode.NONE; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/azure/AzureVectorStoreAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/azure/AzureVectorStoreAutoConfiguration.java index b36626fc7..d18527235 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/azure/AzureVectorStoreAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/azure/AzureVectorStoreAutoConfiguration.java @@ -19,7 +19,7 @@ import com.azure.core.credential.AzureKeyCredential; import com.azure.search.documents.indexes.SearchIndexClient; import com.azure.search.documents.indexes.SearchIndexClientBuilder; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.azure.AzureVectorStore; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; @@ -32,7 +32,7 @@ import org.springframework.context.annotation.Bean; * @author Christian Tzolov */ @AutoConfiguration -@ConditionalOnClass({ EmbeddingClient.class, SearchIndexClient.class, AzureVectorStore.class }) +@ConditionalOnClass({ EmbeddingModel.class, SearchIndexClient.class, AzureVectorStore.class }) @EnableConfigurationProperties({ AzureVectorStoreProperties.class }) @ConditionalOnProperty(prefix = "spring.ai.vectorstore.azure", value = { "url", "api-key", "index-name" }) public class AzureVectorStoreAutoConfiguration { @@ -47,10 +47,10 @@ public class AzureVectorStoreAutoConfiguration { @Bean @ConditionalOnMissingBean - public AzureVectorStore vectorStore(SearchIndexClient searchIndexClient, EmbeddingClient embeddingClient, + public AzureVectorStore vectorStore(SearchIndexClient searchIndexClient, EmbeddingModel embeddingModel, AzureVectorStoreProperties properties) { - var vectorStore = new AzureVectorStore(searchIndexClient, embeddingClient); + var vectorStore = new AzureVectorStore(searchIndexClient, embeddingModel); vectorStore.setIndexName(properties.getIndexName()); diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStoreAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStoreAutoConfiguration.java index f9e760169..eb053338f 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStoreAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStoreAutoConfiguration.java @@ -20,7 +20,7 @@ import java.time.Duration; import com.datastax.oss.driver.api.core.CqlSession; import com.datastax.oss.driver.api.core.config.DefaultDriverOption; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.CassandraVectorStore; import org.springframework.ai.vectorstore.CassandraVectorStoreConfig; import org.springframework.boot.autoconfigure.AutoConfiguration; @@ -42,7 +42,7 @@ public class CassandraVectorStoreAutoConfiguration { @Bean @ConditionalOnMissingBean - public CassandraVectorStore vectorStore(EmbeddingClient embeddingClient, CassandraVectorStoreProperties properties, + public CassandraVectorStore vectorStore(EmbeddingModel embeddingModel, CassandraVectorStoreProperties properties, CqlSession cqlSession) { var builder = CassandraVectorStoreConfig.builder().withCqlSession(cqlSession); @@ -61,7 +61,7 @@ public class CassandraVectorStoreAutoConfiguration { builder = builder.returnEmbeddings(); } - return new CassandraVectorStore(builder.build(), embeddingClient); + return new CassandraVectorStore(builder.build(), embeddingModel); } @Bean diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfiguration.java index 74a2d3200..cd3cf9313 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfiguration.java @@ -18,7 +18,7 @@ package org.springframework.ai.autoconfigure.vectorstore.chroma; import com.fasterxml.jackson.databind.ObjectMapper; import org.springframework.ai.chroma.ChromaApi; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.ChromaVectorStore; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; @@ -33,7 +33,7 @@ import org.springframework.web.client.RestTemplate; * @author Eddú Meléndez */ @AutoConfiguration -@ConditionalOnClass({ EmbeddingClient.class, RestTemplate.class, ChromaVectorStore.class, ObjectMapper.class }) +@ConditionalOnClass({ EmbeddingModel.class, RestTemplate.class, ChromaVectorStore.class, ObjectMapper.class }) @EnableConfigurationProperties({ ChromaApiProperties.class, ChromaVectorStoreProperties.class }) public class ChromaVectorStoreAutoConfiguration { @@ -70,9 +70,9 @@ public class ChromaVectorStoreAutoConfiguration { @Bean @ConditionalOnMissingBean - public ChromaVectorStore vectorStore(EmbeddingClient embeddingClient, ChromaApi chromaApi, + public ChromaVectorStore vectorStore(EmbeddingModel embeddingModel, ChromaApi chromaApi, ChromaVectorStoreProperties storeProperties) { - return new ChromaVectorStore(embeddingClient, chromaApi, storeProperties.getCollectionName()); + return new ChromaVectorStore(embeddingModel, chromaApi, storeProperties.getCollectionName()); } private static class PropertiesChromaConnectionDetails implements ChromaConnectionDetails { diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/elasticsearch/ElasticsearchVectorStoreAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/elasticsearch/ElasticsearchVectorStoreAutoConfiguration.java index df614dbdb..460ec67f5 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/elasticsearch/ElasticsearchVectorStoreAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/elasticsearch/ElasticsearchVectorStoreAutoConfiguration.java @@ -17,7 +17,7 @@ package org.springframework.ai.autoconfigure.vectorstore.elasticsearch; import org.elasticsearch.client.RestClient; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.ElasticsearchVectorStore; import org.springframework.ai.vectorstore.ElasticsearchVectorStoreOptions; import org.springframework.boot.autoconfigure.AutoConfiguration; @@ -35,14 +35,14 @@ import org.springframework.util.StringUtils; */ @AutoConfiguration(after = ElasticsearchRestClientAutoConfiguration.class) -@ConditionalOnClass({ ElasticsearchVectorStore.class, EmbeddingClient.class, RestClient.class }) +@ConditionalOnClass({ ElasticsearchVectorStore.class, EmbeddingModel.class, RestClient.class }) @EnableConfigurationProperties(ElasticsearchVectorStoreProperties.class) class ElasticsearchVectorStoreAutoConfiguration { @Bean @ConditionalOnMissingBean ElasticsearchVectorStore vectorStore(ElasticsearchVectorStoreProperties properties, RestClient restClient, - EmbeddingClient embeddingClient) { + EmbeddingModel embeddingModel) { ElasticsearchVectorStoreOptions elasticsearchVectorStoreOptions = new ElasticsearchVectorStoreOptions(); if (StringUtils.hasText(properties.getIndexName())) { @@ -58,7 +58,7 @@ class ElasticsearchVectorStoreAutoConfiguration { elasticsearchVectorStoreOptions.setSimilarity(properties.getSimilarity()); } - return new ElasticsearchVectorStore(elasticsearchVectorStoreOptions, restClient, embeddingClient); + return new ElasticsearchVectorStore(elasticsearchVectorStoreOptions, restClient, embeddingModel); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/hanadb/HanaCloudVectorStoreAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/hanadb/HanaCloudVectorStoreAutoConfiguration.java index edb69a13f..10076958a 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/hanadb/HanaCloudVectorStoreAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/hanadb/HanaCloudVectorStoreAutoConfiguration.java @@ -17,7 +17,7 @@ package org.springframework.ai.autoconfigure.vectorstore.hanadb; import javax.sql.DataSource; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.HanaCloudVectorStore; import org.springframework.ai.vectorstore.HanaCloudVectorStoreConfig; import org.springframework.ai.vectorstore.HanaVectorEntity; @@ -41,9 +41,9 @@ public class HanaCloudVectorStoreAutoConfiguration { @Bean @ConditionalOnMissingBean public HanaCloudVectorStore vectorStore(HanaVectorRepository repository, - EmbeddingClient embeddingClient, HanaCloudVectorStoreProperties properties) { + EmbeddingModel embeddingModel, HanaCloudVectorStoreProperties properties) { - return new HanaCloudVectorStore(repository, embeddingClient, + return new HanaCloudVectorStore(repository, embeddingModel, HanaCloudVectorStoreConfig.builder() .tableName(properties.getTableName()) .topK(properties.getTopK()) diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/milvus/MilvusVectorStoreAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/milvus/MilvusVectorStoreAutoConfiguration.java index c1b2ab70a..db325d31a 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/milvus/MilvusVectorStoreAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/milvus/MilvusVectorStoreAutoConfiguration.java @@ -22,7 +22,7 @@ import io.milvus.param.ConnectParam; import io.milvus.param.IndexType; import io.milvus.param.MetricType; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.MilvusVectorStore; import org.springframework.ai.vectorstore.MilvusVectorStore.MilvusVectorStoreConfig; import org.springframework.boot.autoconfigure.AutoConfiguration; @@ -37,7 +37,7 @@ import org.springframework.util.StringUtils; * @author Eddú Meléndez */ @AutoConfiguration -@ConditionalOnClass({ MilvusVectorStore.class, EmbeddingClient.class }) +@ConditionalOnClass({ MilvusVectorStore.class, EmbeddingModel.class }) @EnableConfigurationProperties({ MilvusServiceClientProperties.class, MilvusVectorStoreProperties.class }) public class MilvusVectorStoreAutoConfiguration { @@ -50,7 +50,7 @@ public class MilvusVectorStoreAutoConfiguration { @Bean @ConditionalOnMissingBean - public MilvusVectorStore vectorStore(MilvusServiceClient milvusClient, EmbeddingClient embeddingClient, + public MilvusVectorStore vectorStore(MilvusServiceClient milvusClient, EmbeddingModel embeddingModel, MilvusVectorStoreProperties properties) { MilvusVectorStoreConfig config = MilvusVectorStoreConfig.builder() @@ -62,7 +62,7 @@ public class MilvusVectorStoreAutoConfiguration { .withEmbeddingDimension(properties.getEmbeddingDimension()) .build(); - return new MilvusVectorStore(milvusClient, embeddingClient, config); + return new MilvusVectorStore(milvusClient, embeddingModel, config); } @Bean diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/mongo/MongoDBAtlasVectorStoreAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/mongo/MongoDBAtlasVectorStoreAutoConfiguration.java index 2088a6fca..8baabd87f 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/mongo/MongoDBAtlasVectorStoreAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/mongo/MongoDBAtlasVectorStoreAutoConfiguration.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.autoconfigure.vectorstore.mongo; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.MongoDBAtlasVectorStore; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; @@ -32,13 +32,13 @@ import org.springframework.util.StringUtils; * @since 1.0.0 */ @AutoConfiguration(after = MongoDataAutoConfiguration.class) -@ConditionalOnClass({ MongoDBAtlasVectorStore.class, EmbeddingClient.class, MongoTemplate.class }) +@ConditionalOnClass({ MongoDBAtlasVectorStore.class, EmbeddingModel.class, MongoTemplate.class }) @EnableConfigurationProperties(MongoDBAtlasVectorStoreProperties.class) public class MongoDBAtlasVectorStoreAutoConfiguration { @Bean @ConditionalOnMissingBean - MongoDBAtlasVectorStore vectorStore(MongoTemplate mongoTemplate, EmbeddingClient embeddingClient, + MongoDBAtlasVectorStore vectorStore(MongoTemplate mongoTemplate, EmbeddingModel embeddingModel, MongoDBAtlasVectorStoreProperties properties) { var builder = MongoDBAtlasVectorStore.MongoDBVectorStoreConfig.builder(); @@ -54,7 +54,7 @@ public class MongoDBAtlasVectorStoreAutoConfiguration { } MongoDBAtlasVectorStore.MongoDBVectorStoreConfig config = builder.build(); - return new MongoDBAtlasVectorStore(mongoTemplate, embeddingClient, config); + return new MongoDBAtlasVectorStore(mongoTemplate, embeddingModel, config); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/neo4j/Neo4jVectorStoreAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/neo4j/Neo4jVectorStoreAutoConfiguration.java index 28be1a7d1..e5abcf5fb 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/neo4j/Neo4jVectorStoreAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/neo4j/Neo4jVectorStoreAutoConfiguration.java @@ -17,7 +17,7 @@ package org.springframework.ai.autoconfigure.vectorstore.neo4j; import org.neo4j.driver.Driver; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.Neo4jVectorStore; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; @@ -30,13 +30,13 @@ import org.springframework.context.annotation.Bean; * @author Jingzhou Ou */ @AutoConfiguration(after = Neo4jAutoConfiguration.class) -@ConditionalOnClass({ Neo4jVectorStore.class, EmbeddingClient.class, Driver.class }) +@ConditionalOnClass({ Neo4jVectorStore.class, EmbeddingModel.class, Driver.class }) @EnableConfigurationProperties({ Neo4jVectorStoreProperties.class }) public class Neo4jVectorStoreAutoConfiguration { @Bean @ConditionalOnMissingBean - public Neo4jVectorStore vectorStore(Driver driver, EmbeddingClient embeddingClient, + public Neo4jVectorStore vectorStore(Driver driver, EmbeddingModel embeddingModel, Neo4jVectorStoreProperties properties) { Neo4jVectorStore.Neo4jVectorStoreConfig config = Neo4jVectorStore.Neo4jVectorStoreConfig.builder() .withDatabaseName(properties.getDatabaseName()) @@ -49,7 +49,7 @@ public class Neo4jVectorStoreAutoConfiguration { .withConstraintName(properties.getConstraintName()) .build(); - return new Neo4jVectorStore(driver, embeddingClient, config); + return new Neo4jVectorStore(driver, embeddingModel, config); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/pgvector/PgVectorStoreAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/pgvector/PgVectorStoreAutoConfiguration.java index 0ad62f417..58ff912cf 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/pgvector/PgVectorStoreAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/pgvector/PgVectorStoreAutoConfiguration.java @@ -17,7 +17,7 @@ package org.springframework.ai.autoconfigure.vectorstore.pgvector; import javax.sql.DataSource; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.PgVectorStore; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; @@ -37,11 +37,11 @@ public class PgVectorStoreAutoConfiguration { @Bean @ConditionalOnMissingBean - public PgVectorStore vectorStore(JdbcTemplate jdbcTemplate, EmbeddingClient embeddingClient, + public PgVectorStore vectorStore(JdbcTemplate jdbcTemplate, EmbeddingModel embeddingModel, PgVectorStoreProperties properties) { - return new PgVectorStore(jdbcTemplate, embeddingClient, properties.getDimensions(), - properties.getDistanceType(), properties.isRemoveExistingVectorStoreTable(), properties.getIndexType()); + return new PgVectorStore(jdbcTemplate, embeddingModel, properties.getDimensions(), properties.getDistanceType(), + properties.isRemoveExistingVectorStoreTable(), properties.getIndexType()); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/pinecone/PineconeVectorStoreAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/pinecone/PineconeVectorStoreAutoConfiguration.java index 871b7dc55..2da22bb6f 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/pinecone/PineconeVectorStoreAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/pinecone/PineconeVectorStoreAutoConfiguration.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.autoconfigure.vectorstore.pinecone; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.PineconeVectorStore; import org.springframework.ai.vectorstore.PineconeVectorStore.PineconeVectorStoreConfig; import org.springframework.boot.autoconfigure.AutoConfiguration; @@ -28,13 +28,13 @@ import org.springframework.context.annotation.Bean; * @author Christian Tzolov */ @AutoConfiguration -@ConditionalOnClass({ PineconeVectorStore.class, EmbeddingClient.class }) +@ConditionalOnClass({ PineconeVectorStore.class, EmbeddingModel.class }) @EnableConfigurationProperties(PineconeVectorStoreProperties.class) public class PineconeVectorStoreAutoConfiguration { @Bean @ConditionalOnMissingBean - public PineconeVectorStore vectorStore(EmbeddingClient embeddingClient, PineconeVectorStoreProperties properties) { + public PineconeVectorStore vectorStore(EmbeddingModel embeddingModel, PineconeVectorStoreProperties properties) { var config = PineconeVectorStoreConfig.builder() .withApiKey(properties.getApiKey()) @@ -45,7 +45,7 @@ public class PineconeVectorStoreAutoConfiguration { .withServerSideTimeout(properties.getServerSideTimeout()) .build(); - return new PineconeVectorStore(config, embeddingClient); + return new PineconeVectorStore(config, embeddingModel); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/qdrant/QdrantVectorStoreAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/qdrant/QdrantVectorStoreAutoConfiguration.java index c0dd23640..5de119d99 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/qdrant/QdrantVectorStoreAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/qdrant/QdrantVectorStoreAutoConfiguration.java @@ -17,7 +17,7 @@ package org.springframework.ai.autoconfigure.vectorstore.qdrant; import io.qdrant.client.QdrantClient; import io.qdrant.client.QdrantGrpcClient; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.qdrant.QdrantVectorStore; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; @@ -31,7 +31,7 @@ import org.springframework.context.annotation.Bean; * @since 0.8.1 */ @AutoConfiguration -@ConditionalOnClass({ QdrantVectorStore.class, EmbeddingClient.class }) +@ConditionalOnClass({ QdrantVectorStore.class, EmbeddingModel.class }) @EnableConfigurationProperties(QdrantVectorStoreProperties.class) public class QdrantVectorStoreAutoConfiguration { @@ -56,9 +56,9 @@ public class QdrantVectorStoreAutoConfiguration { @Bean @ConditionalOnMissingBean - public QdrantVectorStore vectorStore(EmbeddingClient embeddingClient, QdrantVectorStoreProperties properties, + public QdrantVectorStore vectorStore(EmbeddingModel embeddingModel, QdrantVectorStoreProperties properties, QdrantClient qdrantClient) { - return new QdrantVectorStore(qdrantClient, properties.getCollectionName(), embeddingClient); + return new QdrantVectorStore(qdrantClient, properties.getCollectionName(), embeddingModel); } static class PropertiesQdrantConnectionDetails implements QdrantConnectionDetails { diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/redis/RedisVectorStoreAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/redis/RedisVectorStoreAutoConfiguration.java index e17d76e7d..22df3a02d 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/redis/RedisVectorStoreAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/redis/RedisVectorStoreAutoConfiguration.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.autoconfigure.vectorstore.redis; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.RedisVectorStore; import org.springframework.ai.vectorstore.RedisVectorStore.RedisVectorStoreConfig; import org.springframework.boot.autoconfigure.AutoConfiguration; @@ -29,7 +29,7 @@ import org.springframework.context.annotation.Bean; * @author Eddú Meléndez */ @AutoConfiguration -@ConditionalOnClass({ RedisVectorStore.class, EmbeddingClient.class }) +@ConditionalOnClass({ RedisVectorStore.class, EmbeddingModel.class }) @EnableConfigurationProperties(RedisVectorStoreProperties.class) public class RedisVectorStoreAutoConfiguration { @@ -41,7 +41,7 @@ public class RedisVectorStoreAutoConfiguration { @Bean @ConditionalOnMissingBean - public RedisVectorStore vectorStore(EmbeddingClient embeddingClient, RedisVectorStoreProperties properties, + public RedisVectorStore vectorStore(EmbeddingModel embeddingModel, RedisVectorStoreProperties properties, RedisConnectionDetails redisConnectionDetails) { var config = RedisVectorStoreConfig.builder() @@ -50,7 +50,7 @@ public class RedisVectorStoreAutoConfiguration { .withPrefix(properties.getPrefix()) .build(); - return new RedisVectorStore(config, embeddingClient); + return new RedisVectorStore(config, embeddingModel); } private static class PropertiesRedisConnectionDetails implements RedisConnectionDetails { diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/weaviate/WeaviateVectorStoreAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/weaviate/WeaviateVectorStoreAutoConfiguration.java index 431b70723..a358ce3dd 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/weaviate/WeaviateVectorStoreAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vectorstore/weaviate/WeaviateVectorStoreAutoConfiguration.java @@ -19,7 +19,7 @@ import io.weaviate.client.Config; import io.weaviate.client.WeaviateAuthClient; import io.weaviate.client.WeaviateClient; import io.weaviate.client.v1.auth.exception.AuthException; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.WeaviateVectorStore; import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig; import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig.MetadataField; @@ -34,7 +34,7 @@ import org.springframework.context.annotation.Bean; * @author Eddú Meléndez */ @AutoConfiguration -@ConditionalOnClass({ EmbeddingClient.class, WeaviateVectorStore.class }) +@ConditionalOnClass({ EmbeddingModel.class, WeaviateVectorStore.class }) @EnableConfigurationProperties({ WeaviateVectorStoreProperties.class }) public class WeaviateVectorStoreAutoConfiguration { @@ -60,7 +60,7 @@ public class WeaviateVectorStoreAutoConfiguration { @Bean @ConditionalOnMissingBean - public WeaviateVectorStore vectorStore(EmbeddingClient embeddingClient, WeaviateClient weaviateClient, + public WeaviateVectorStore vectorStore(EmbeddingModel embeddingModel, WeaviateClient weaviateClient, WeaviateVectorStoreProperties properties) { WeaviateVectorStoreConfig.Builder configBuilder = WeaviateVectorStore.WeaviateVectorStoreConfig.builder() @@ -72,7 +72,7 @@ public class WeaviateVectorStoreAutoConfiguration { .toList()) .withConsistencyLevel(properties.getConsistencyLevel()); - return new WeaviateVectorStore(configBuilder.build(), embeddingClient, weaviateClient); + return new WeaviateVectorStore(configBuilder.build(), embeddingModel, weaviateClient); } static class PropertiesWeaviateConnectionDetails implements WeaviateConnectionDetails { diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/gemini/VertexAiGeminiAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/gemini/VertexAiGeminiAutoConfiguration.java index 12dfb8134..813103603 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/gemini/VertexAiGeminiAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/gemini/VertexAiGeminiAutoConfiguration.java @@ -24,7 +24,7 @@ import com.google.cloud.vertexai.VertexAI; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; import org.springframework.ai.model.function.FunctionCallbackWrapper.Builder.SchemaType; -import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatClient; +import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatModel; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean; import org.springframework.boot.context.properties.EnableConfigurationProperties; @@ -40,7 +40,7 @@ import org.springframework.util.StringUtils; * @author Christian Tzolov * @since 0.8.0 */ -@ConditionalOnClass({ VertexAI.class, VertexAiGeminiChatClient.class }) +@ConditionalOnClass({ VertexAI.class, VertexAiGeminiChatModel.class }) @EnableConfigurationProperties({ VertexAiGeminiChatProperties.class, VertexAiGeminiConnectionProperties.class }) public class VertexAiGeminiAutoConfiguration { @@ -74,7 +74,7 @@ public class VertexAiGeminiAutoConfiguration { @Bean @ConditionalOnMissingBean - public VertexAiGeminiChatClient vertexAiGeminiChat(VertexAI vertexAi, VertexAiGeminiChatProperties chatProperties, + public VertexAiGeminiChatModel vertexAiGeminiChat(VertexAI vertexAi, VertexAiGeminiChatProperties chatProperties, List toolFunctionCallbacks, ApplicationContext context) { FunctionCallbackContext functionCallbackContext = springAiFunctionManager(context); @@ -83,7 +83,7 @@ public class VertexAiGeminiAutoConfiguration { chatProperties.getOptions().getFunctionCallbacks().addAll(toolFunctionCallbacks); } - return new VertexAiGeminiChatClient(vertexAi, chatProperties.getOptions(), functionCallbackContext); + return new VertexAiGeminiChatModel(vertexAi, chatProperties.getOptions(), functionCallbackContext); } /** diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/gemini/VertexAiGeminiChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/gemini/VertexAiGeminiChatProperties.java index 866f4f352..141e4ce98 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/gemini/VertexAiGeminiChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/gemini/VertexAiGeminiChatProperties.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.autoconfigure.vertexai.gemini; -import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatClient; +import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatModel; import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatOptions; import org.springframework.boot.context.properties.ConfigurationProperties; @@ -30,7 +30,7 @@ public class VertexAiGeminiChatProperties { public static final String CONFIG_PREFIX = "spring.ai.vertex.ai.gemini.chat"; - public static final String DEFAULT_MODEL = VertexAiGeminiChatClient.ChatModel.GEMINI_PRO_VISION.getValue(); + public static final String DEFAULT_MODEL = VertexAiGeminiChatModel.ChatModel.GEMINI_PRO_VISION.getValue(); /** * Vertex AI Gemini API generative options. diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPalm2AutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPalm2AutoConfiguration.java index c1ac9fa75..96708399f 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPalm2AutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPalm2AutoConfiguration.java @@ -15,8 +15,8 @@ */ package org.springframework.ai.autoconfigure.vertexai.palm2; -import org.springframework.ai.vertexai.palm2.VertexAiPaLm2ChatClient; -import org.springframework.ai.vertexai.palm2.VertexAiPaLm2EmbeddingClient; +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.condition.ConditionalOnClass; @@ -47,17 +47,17 @@ public class VertexAiPalm2AutoConfiguration { @ConditionalOnMissingBean @ConditionalOnProperty(prefix = VertexAiPlam2ChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public VertexAiPaLm2ChatClient vertexAiChatClient(VertexAiPaLm2Api vertexAiApi, + public VertexAiPaLm2ChatModel vertexAiChatModel(VertexAiPaLm2Api vertexAiApi, VertexAiPlam2ChatProperties chatProperties) { - return new VertexAiPaLm2ChatClient(vertexAiApi, chatProperties.getOptions()); + return new VertexAiPaLm2ChatModel(vertexAiApi, chatProperties.getOptions()); } @Bean @ConditionalOnMissingBean @ConditionalOnProperty(prefix = VertexAiPalm2EmbeddingProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public VertexAiPaLm2EmbeddingClient vertexAiEmbeddingClient(VertexAiPaLm2Api vertexAiApi) { - return new VertexAiPaLm2EmbeddingClient(vertexAiApi); + public VertexAiPaLm2EmbeddingModel vertexAiEmbeddingModel(VertexAiPaLm2Api vertexAiApi) { + return new VertexAiPaLm2EmbeddingModel(vertexAiApi); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPalm2EmbeddingProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPalm2EmbeddingProperties.java index 2e01dbeab..0dc079b03 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPalm2EmbeddingProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPalm2EmbeddingProperties.java @@ -24,7 +24,7 @@ public class VertexAiPalm2EmbeddingProperties { public static final String CONFIG_PREFIX = "spring.ai.vertex.ai.embedding"; /** - * Enable Vertex AI PaLM API embedding client. + * Enable Vertex AI PaLM API embedding model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPlam2ChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPlam2ChatProperties.java index 881f8e8dc..d966f9d7e 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPlam2ChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPlam2ChatProperties.java @@ -25,7 +25,7 @@ public class VertexAiPlam2ChatProperties { public static final String CONFIG_PREFIX = "spring.ai.vertex.ai.chat"; /** - * Enable Vertex AI PaLM API chat client. + * Enable Vertex AI PaLM API chat model */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/watsonxai/WatsonxAiAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/watsonxai/WatsonxAiAutoConfiguration.java index ce651ee17..fd62de3ce 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/watsonxai/WatsonxAiAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/watsonxai/WatsonxAiAutoConfiguration.java @@ -15,7 +15,7 @@ */ package org.springframework.ai.autoconfigure.watsonxai; -import org.springframework.ai.watsonx.WatsonxAiChatClient; +import org.springframework.ai.watsonx.WatsonxAiChatModel; import org.springframework.ai.watsonx.api.WatsonxAiApi; import org.springframework.boot.autoconfigure.AutoConfiguration; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; @@ -50,8 +50,8 @@ public class WatsonxAiAutoConfiguration { @Bean @ConditionalOnMissingBean - public WatsonxAiChatClient watsonxChatClient(WatsonxAiApi watsonxApi, WatsonxAiChatProperties chatProperties) { - return new WatsonxAiChatClient(watsonxApi, chatProperties.getOptions()); + public WatsonxAiChatModel watsonxChatModel(WatsonxAiApi watsonxApi, WatsonxAiChatProperties chatProperties) { + return new WatsonxAiChatModel(watsonxApi, chatProperties.getOptions()); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/watsonxai/WatsonxAiChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/watsonxai/WatsonxAiChatProperties.java index 3da222e93..c0db014dd 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/watsonxai/WatsonxAiChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/watsonxai/WatsonxAiChatProperties.java @@ -33,7 +33,7 @@ public class WatsonxAiChatProperties { public static final String CONFIG_PREFIX = "spring.ai.watsonx.ai.chat"; /** - * Enable Watsonx.AI chat client. + * Enable Watsonx.AI chat model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfiguration.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfiguration.java index e2c8b1126..894b533b0 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfiguration.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfiguration.java @@ -18,9 +18,9 @@ package org.springframework.ai.autoconfigure.zhipuai; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackContext; -import org.springframework.ai.zhipuai.ZhiPuAiChatClient; -import org.springframework.ai.zhipuai.ZhiPuAiEmbeddingClient; -import org.springframework.ai.zhipuai.ZhiPuAiImageClient; +import org.springframework.ai.zhipuai.ZhiPuAiChatModel; +import org.springframework.ai.zhipuai.ZhiPuAiEmbeddingModel; +import org.springframework.ai.zhipuai.ZhiPuAiImageModel; import org.springframework.ai.zhipuai.api.ZhiPuAiApi; import org.springframework.ai.zhipuai.api.ZhiPuAiImageApi; import org.springframework.boot.autoconfigure.AutoConfiguration; @@ -53,7 +53,7 @@ public class ZhiPuAiAutoConfiguration { @ConditionalOnMissingBean @ConditionalOnProperty(prefix = ZhiPuAiChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public ZhiPuAiChatClient zhiPuAiChatClient(ZhiPuAiConnectionProperties commonProperties, + public ZhiPuAiChatModel zhiPuAiChatModel(ZhiPuAiConnectionProperties commonProperties, ZhiPuAiChatProperties chatProperties, RestClient.Builder restClientBuilder, List toolFunctionCallbacks, FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { @@ -65,21 +65,21 @@ public class ZhiPuAiAutoConfiguration { chatProperties.getOptions().getFunctionCallbacks().addAll(toolFunctionCallbacks); } - return new ZhiPuAiChatClient(zhiPuAiApi, chatProperties.getOptions(), functionCallbackContext, retryTemplate); + return new ZhiPuAiChatModel(zhiPuAiApi, chatProperties.getOptions(), functionCallbackContext, retryTemplate); } @Bean @ConditionalOnMissingBean @ConditionalOnProperty(prefix = ZhiPuAiEmbeddingProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public ZhiPuAiEmbeddingClient zhiPuAiEmbeddingClient(ZhiPuAiConnectionProperties commonProperties, + public ZhiPuAiEmbeddingModel zhiPuAiEmbeddingModel(ZhiPuAiConnectionProperties commonProperties, ZhiPuAiEmbeddingProperties embeddingProperties, RestClient.Builder restClientBuilder, RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { var zhiPuAiApi = zhiPuAiApi(embeddingProperties.getBaseUrl(), commonProperties.getBaseUrl(), embeddingProperties.getApiKey(), commonProperties.getApiKey(), restClientBuilder, responseErrorHandler); - return new ZhiPuAiEmbeddingClient(zhiPuAiApi, embeddingProperties.getMetadataMode(), + return new ZhiPuAiEmbeddingModel(zhiPuAiApi, embeddingProperties.getMetadataMode(), embeddingProperties.getOptions(), retryTemplate); } @@ -99,7 +99,7 @@ public class ZhiPuAiAutoConfiguration { @ConditionalOnMissingBean @ConditionalOnProperty(prefix = ZhiPuAiImageProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) - public ZhiPuAiImageClient zhiPuAiImageClient(ZhiPuAiConnectionProperties commonProperties, + public ZhiPuAiImageModel zhiPuAiImageModel(ZhiPuAiConnectionProperties commonProperties, ZhiPuAiImageProperties imageProperties, RestClient.Builder restClientBuilder, RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { @@ -114,7 +114,7 @@ public class ZhiPuAiAutoConfiguration { var zhiPuAiImageApi = new ZhiPuAiImageApi(baseUrl, apiKey, restClientBuilder, responseErrorHandler); - return new ZhiPuAiImageClient(zhiPuAiImageApi, imageProperties.getOptions(), retryTemplate); + return new ZhiPuAiImageModel(zhiPuAiImageApi, imageProperties.getOptions(), retryTemplate); } @Bean diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiChatProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiChatProperties.java index 1ec554b44..8b54e7e81 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiChatProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiChatProperties.java @@ -33,7 +33,7 @@ public class ZhiPuAiChatProperties extends ZhiPuAiParentProperties { private static final Double DEFAULT_TEMPERATURE = 0.7; /** - * Enable ZhiPuAI chat client. + * Enable ZhiPuAI chat model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiEmbeddingProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiEmbeddingProperties.java index 78357ceff..4e1c6ef80 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiEmbeddingProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiEmbeddingProperties.java @@ -32,7 +32,7 @@ public class ZhiPuAiEmbeddingProperties extends ZhiPuAiParentProperties { public static final String DEFAULT_EMBEDDING_MODEL = ZhiPuAiApi.EmbeddingModel.Embedding_2.value; /** - * Enable ZhiPuAI embedding client. + * Enable ZhiPuAI embedding model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiImageProperties.java b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiImageProperties.java index 19ce4f7bb..7463d4573 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiImageProperties.java +++ b/spring-ai-spring-boot-autoconfigure/src/main/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiImageProperties.java @@ -28,7 +28,7 @@ public class ZhiPuAiImageProperties extends ZhiPuAiParentProperties { public static final String CONFIG_PREFIX = "spring.ai.zhipuai.image"; /** - * Enable ZhiPuAI Image client. + * Enable ZhiPuAI image model. */ private boolean enabled = true; diff --git a/spring-ai-spring-boot-autoconfigure/src/main/resources/META-INF/spring/org.springframework.boot.autoconfigure.AutoConfiguration.imports b/spring-ai-spring-boot-autoconfigure/src/main/resources/META-INF/spring/org.springframework.boot.autoconfigure.AutoConfiguration.imports index 4cbab5ad6..3ce1ae92c 100644 --- a/spring-ai-spring-boot-autoconfigure/src/main/resources/META-INF/spring/org.springframework.boot.autoconfigure.AutoConfiguration.imports +++ b/spring-ai-spring-boot-autoconfigure/src/main/resources/META-INF/spring/org.springframework.boot.autoconfigure.AutoConfiguration.imports @@ -1,7 +1,7 @@ org.springframework.ai.autoconfigure.openai.OpenAiAutoConfiguration org.springframework.ai.autoconfigure.azure.openai.AzureOpenAiAutoConfiguration org.springframework.ai.autoconfigure.stabilityai.StabilityAiImageAutoConfiguration -org.springframework.ai.autoconfigure.transformers.TransformersEmbeddingClientAutoConfiguration +org.springframework.ai.autoconfigure.transformers.TransformersEmbeddingModelAutoConfiguration org.springframework.ai.autoconfigure.huggingface.HuggingfaceChatAutoConfiguration org.springframework.ai.autoconfigure.vertexai.palm2.VertexAiPalm2AutoConfiguration org.springframework.ai.autoconfigure.vertexai.gemini.VertexAiGeminiAutoConfiguration diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/AnthropicAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/AnthropicAutoConfigurationIT.java index 69e33a304..587bcb715 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/AnthropicAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/AnthropicAutoConfigurationIT.java @@ -22,9 +22,9 @@ 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.anthropic.AnthropicChatModel; import reactor.core.publisher.Flux; -import org.springframework.ai.anthropic.AnthropicChatClient; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; @@ -50,8 +50,8 @@ public class AnthropicAutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - AnthropicChatClient chatClient = context.getBean(AnthropicChatClient.class); - String response = chatClient.call("Hello"); + AnthropicChatModel chatModel = context.getBean(AnthropicChatModel.class); + String response = chatModel.call("Hello"); assertThat(response).isNotEmpty(); logger.info("Response: " + response); }); @@ -60,8 +60,8 @@ public class AnthropicAutoConfigurationIT { @Test void generateStreaming() { contextRunner.run(context -> { - AnthropicChatClient chatClient = context.getBean(AnthropicChatClient.class); - Flux responseFlux = chatClient.stream(new Prompt(new UserMessage("Hello"))); + AnthropicChatModel chatModel = context.getBean(AnthropicChatModel.class); + Flux responseFlux = chatModel.stream(new Prompt(new UserMessage("Hello"))); String response = responseFlux.collectList() .block() diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/AnthropicPropertiesTests.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/AnthropicPropertiesTests.java index 0c2c6aae3..086e7e5b5 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/AnthropicPropertiesTests.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/AnthropicPropertiesTests.java @@ -17,7 +17,7 @@ package org.springframework.ai.autoconfigure.anthropic; import org.junit.jupiter.api.Test; -import org.springframework.ai.anthropic.AnthropicChatClient; +import org.springframework.ai.anthropic.AnthropicChatModel; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; @@ -102,7 +102,7 @@ public class AnthropicPropertiesTests { RestClientAutoConfiguration.class, AnthropicAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(AnthropicChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(AnthropicChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(AnthropicChatModel.class)).isNotEmpty(); }); // Explicitly enable the chat auto-configuration. @@ -111,7 +111,7 @@ public class AnthropicPropertiesTests { RestClientAutoConfiguration.class, AnthropicAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(AnthropicChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(AnthropicChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(AnthropicChatModel.class)).isNotEmpty(); }); // Explicitly disable the chat auto-configuration. @@ -120,7 +120,7 @@ public class AnthropicPropertiesTests { RestClientAutoConfiguration.class, AnthropicAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(AnthropicChatProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(AnthropicChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(AnthropicChatModel.class)).isEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/tool/FunctionCallWithFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/tool/FunctionCallWithFunctionBeanIT.java index d170add72..6e7366e2e 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/tool/FunctionCallWithFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/tool/FunctionCallWithFunctionBeanIT.java @@ -23,7 +23,7 @@ import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.anthropic.AnthropicChatClient; +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; @@ -61,19 +61,19 @@ class FunctionCallWithFunctionBeanIT { "spring.ai.anthropic.chat.options.model=" + AnthropicApi.ChatModel.CLAUDE_3_OPUS.getValue()) .run(context -> { - AnthropicChatClient chatClient = context.getBean(AnthropicChatClient.class); + AnthropicChatModel chatModel = context.getBean(AnthropicChatModel.class); var userMessage = new UserMessage( "What's the weather like in San Francisco, in Paris, France and in Tokyo, Japan? Return the temperature in Celsius."); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), AnthropicChatOptions.builder().withFunction("weatherFunction").build())); logger.info("Response: {}", response); assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); - response = chatClient.call(new Prompt(List.of(userMessage), + response = chatModel.call(new Prompt(List.of(userMessage), AnthropicChatOptions.builder().withFunction("weatherFunction3").build())); logger.info("Response: {}", response); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/tool/FunctionCallWithPromptFunctionIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/tool/FunctionCallWithPromptFunctionIT.java index a9592cae7..31e67231f 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/tool/FunctionCallWithPromptFunctionIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/anthropic/tool/FunctionCallWithPromptFunctionIT.java @@ -22,7 +22,7 @@ import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.anthropic.AnthropicChatClient; +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; @@ -54,7 +54,7 @@ public class FunctionCallWithPromptFunctionIT { "spring.ai.anthropic.chat.options.model=" + AnthropicApi.ChatModel.CLAUDE_3_OPUS.getValue()) .run(context -> { - AnthropicChatClient chatClient = context.getBean(AnthropicChatClient.class); + AnthropicChatModel chatModel = context.getBean(AnthropicChatModel.class); UserMessage userMessage = new UserMessage( "What's the weather like in San Francisco, in Paris and in Tokyo? Return the temperature in Celsius."); @@ -66,7 +66,7 @@ public class FunctionCallWithPromptFunctionIT { .build())) .build(); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), promptOptions)); + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), promptOptions)); logger.info("Response: {}", response); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/AzureOpenAiAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/AzureOpenAiAutoConfigurationIT.java index ce7075250..2b3c3be0e 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/AzureOpenAiAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/AzureOpenAiAutoConfigurationIT.java @@ -21,12 +21,12 @@ import java.util.stream.Collectors; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.springframework.ai.azure.openai.AzureOpenAiChatModel; import org.springframework.ai.chat.messages.AssistantMessage; import reactor.core.publisher.Flux; import org.springframework.ai.autoconfigure.azure.openai.AzureOpenAiAutoConfiguration; -import org.springframework.ai.azure.openai.AzureOpenAiChatClient; -import org.springframework.ai.azure.openai.AzureOpenAiEmbeddingClient; +import org.springframework.ai.azure.openai.AzureOpenAiEmbeddingModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; import org.springframework.ai.embedding.EmbeddingResponse; @@ -77,8 +77,8 @@ public class AzureOpenAiAutoConfigurationIT { @Test public void chatCompletion() { contextRunner.run(context -> { - AzureOpenAiChatClient chatClient = context.getBean(AzureOpenAiChatClient.class); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage, systemMessage))); + AzureOpenAiChatModel chatModel = context.getBean(AzureOpenAiChatModel.class); + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage, systemMessage))); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); }); } @@ -87,9 +87,9 @@ public class AzureOpenAiAutoConfigurationIT { public void chatCompletionStreaming() { contextRunner.run(context -> { - AzureOpenAiChatClient chatClient = context.getBean(AzureOpenAiChatClient.class); + AzureOpenAiChatModel chatModel = context.getBean(AzureOpenAiChatModel.class); - Flux response = chatClient.stream(new Prompt(List.of(userMessage, systemMessage))); + Flux response = chatModel.stream(new Prompt(List.of(userMessage, systemMessage))); List responses = response.collectList().block(); assertThat(responses.size()).isGreaterThan(1); @@ -108,9 +108,9 @@ public class AzureOpenAiAutoConfigurationIT { @Test void embedding() { contextRunner.run(context -> { - AzureOpenAiEmbeddingClient embeddingClient = context.getBean(AzureOpenAiEmbeddingClient.class); + AzureOpenAiEmbeddingModel embeddingModel = context.getBean(AzureOpenAiEmbeddingModel.class); - EmbeddingResponse embeddingResponse = embeddingClient + EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); assertThat(embeddingResponse.getResults()).hasSize(2); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); @@ -118,7 +118,7 @@ public class AzureOpenAiAutoConfigurationIT { assertThat(embeddingResponse.getResults().get(1).getOutput()).isNotEmpty(); assertThat(embeddingResponse.getResults().get(1).getIndex()).isEqualTo(1); - assertThat(embeddingClient.dimensions()).isEqualTo(1536); + assertThat(embeddingModel.dimensions()).isEqualTo(1536); }); } @@ -127,17 +127,17 @@ public class AzureOpenAiAutoConfigurationIT { // Disable the chat auto-configuration. contextRunner.withPropertyValues("spring.ai.azure.openai.chat.enabled=false").run(context -> { - assertThat(context.getBeansOfType(AzureOpenAiChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(AzureOpenAiChatModel.class)).isEmpty(); }); // The chat auto-configuration is enabled by default. contextRunner.run(context -> { - assertThat(context.getBeansOfType(AzureOpenAiChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(AzureOpenAiChatModel.class)).isNotEmpty(); }); // Explicitly enable the chat auto-configuration. contextRunner.withPropertyValues("spring.ai.azure.openai.chat.enabled=true").run(context -> { - assertThat(context.getBeansOfType(AzureOpenAiChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(AzureOpenAiChatModel.class)).isNotEmpty(); }); } @@ -146,17 +146,17 @@ public class AzureOpenAiAutoConfigurationIT { // Disable the embedding auto-configuration. contextRunner.withPropertyValues("spring.ai.azure.openai.embedding.enabled=false").run(context -> { - assertThat(context.getBeansOfType(AzureOpenAiEmbeddingClient.class)).isEmpty(); + assertThat(context.getBeansOfType(AzureOpenAiEmbeddingModel.class)).isEmpty(); }); // The embedding auto-configuration is enabled by default. contextRunner.run(context -> { - assertThat(context.getBeansOfType(AzureOpenAiEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(AzureOpenAiEmbeddingModel.class)).isNotEmpty(); }); // Explicitly enable the embedding auto-configuration. contextRunner.withPropertyValues("spring.ai.azure.openai.embedding.enabled=true").run(context -> { - assertThat(context.getBeansOfType(AzureOpenAiEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(AzureOpenAiEmbeddingModel.class)).isNotEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionBeanIT.java index e8a7e2126..30ec5df8b 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionBeanIT.java @@ -24,9 +24,9 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.autoconfigure.azure.openai.AzureOpenAiAutoConfiguration; -import org.springframework.ai.azure.openai.AzureOpenAiChatClient; +import org.springframework.ai.azure.openai.AzureOpenAiChatModel; import org.springframework.ai.azure.openai.AzureOpenAiChatOptions; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; @@ -57,19 +57,19 @@ class FunctionCallWithFunctionBeanIT { contextRunner.withPropertyValues("spring.ai.azure.openai.chat.options..deployment-name=gpt-4-0125-preview") .run(context -> { - ChatClient chatClient = context.getBean(AzureOpenAiChatClient.class); + ChatModel chatModel = context.getBean(AzureOpenAiChatModel.class); UserMessage userMessage = new UserMessage( "What's the weather like in San Francisco, Paris and in Tokyo? Use Multi-turn function calling."); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), AzureOpenAiChatOptions.builder().withFunction("weatherFunction").build())); logger.info("Response: {}", response); assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); - response = chatClient.call(new Prompt(List.of(userMessage), + response = chatModel.call(new Prompt(List.of(userMessage), AzureOpenAiChatOptions.builder().withFunction("weatherFunction3").build())); logger.info("Response: {}", response); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionWrapperIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionWrapperIT.java index 934b84cf1..06fb413ca 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionWrapperIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithFunctionWrapperIT.java @@ -23,7 +23,7 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.autoconfigure.azure.openai.AzureOpenAiAutoConfiguration; -import org.springframework.ai.azure.openai.AzureOpenAiChatClient; +import org.springframework.ai.azure.openai.AzureOpenAiChatModel; import org.springframework.ai.azure.openai.AzureOpenAiChatOptions; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.UserMessage; @@ -56,12 +56,12 @@ public class FunctionCallWithFunctionWrapperIT { contextRunner.withPropertyValues("spring.ai.azure.openai.chat.options.deployment-name=gpt-4-0125-preview") .run(context -> { - AzureOpenAiChatClient chatClient = context.getBean(AzureOpenAiChatClient.class); + AzureOpenAiChatModel chatModel = context.getBean(AzureOpenAiChatModel.class); UserMessage userMessage = new UserMessage( "What's the weather like in San Francisco, Paris and in Tokyo?"); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), AzureOpenAiChatOptions.builder().withFunction("WeatherInfo").build())); logger.info("Response: {}", response); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithPromptFunctionIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithPromptFunctionIT.java index 9bef922d2..048d4fa6c 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithPromptFunctionIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/azure/tool/FunctionCallWithPromptFunctionIT.java @@ -23,7 +23,7 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.autoconfigure.azure.openai.AzureOpenAiAutoConfiguration; -import org.springframework.ai.azure.openai.AzureOpenAiChatClient; +import org.springframework.ai.azure.openai.AzureOpenAiChatModel; import org.springframework.ai.azure.openai.AzureOpenAiChatOptions; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.UserMessage; @@ -52,7 +52,7 @@ public class FunctionCallWithPromptFunctionIT { contextRunner.withPropertyValues("spring.ai.azure.openai.chat.options.deployment-name=gpt-4-0125-preview") .run(context -> { - AzureOpenAiChatClient chatClient = context.getBean(AzureOpenAiChatClient.class); + AzureOpenAiChatModel chatModel = context.getBean(AzureOpenAiChatModel.class); UserMessage userMessage = new UserMessage( "What's the weather like in San Francisco, in Paris and in Tokyo? Use Multi-turn function calling."); @@ -64,7 +64,7 @@ public class FunctionCallWithPromptFunctionIT { .build())) .build(); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), promptOptions)); + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), promptOptions)); logger.info("Response: {}", response); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/anthropic/BedrockAnthropicChatAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/anthropic/BedrockAnthropicChatAutoConfigurationIT.java index f8aca8b9a..014ba672b 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/anthropic/BedrockAnthropicChatAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/anthropic/BedrockAnthropicChatAutoConfigurationIT.java @@ -21,12 +21,12 @@ import java.util.stream.Collectors; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.springframework.ai.bedrock.anthropic.BedrockAnthropicChatModel; import org.springframework.ai.chat.messages.AssistantMessage; import reactor.core.publisher.Flux; import software.amazon.awssdk.regions.Region; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.anthropic.BedrockAnthropicChatClient; import org.springframework.ai.bedrock.anthropic.api.AnthropicChatBedrockApi.AnthropicChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; @@ -69,8 +69,8 @@ public class BedrockAnthropicChatAutoConfigurationIT { @Test public void chatCompletion() { contextRunner.run(context -> { - BedrockAnthropicChatClient anthropicChatClient = context.getBean(BedrockAnthropicChatClient.class); - ChatResponse response = anthropicChatClient.call(new Prompt(List.of(userMessage, systemMessage))); + BedrockAnthropicChatModel anthropicChatModel = context.getBean(BedrockAnthropicChatModel.class); + ChatResponse response = anthropicChatModel.call(new Prompt(List.of(userMessage, systemMessage))); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); }); } @@ -79,9 +79,9 @@ public class BedrockAnthropicChatAutoConfigurationIT { public void chatCompletionStreaming() { contextRunner.run(context -> { - BedrockAnthropicChatClient anthropicChatClient = context.getBean(BedrockAnthropicChatClient.class); + BedrockAnthropicChatModel anthropicChatModel = context.getBean(BedrockAnthropicChatModel.class); - Flux response = anthropicChatClient.stream(new Prompt(List.of(userMessage, systemMessage))); + Flux response = anthropicChatModel.stream(new Prompt(List.of(userMessage, systemMessage))); List responses = response.collectList().block(); assertThat(responses.size()).isGreaterThan(2); @@ -130,7 +130,7 @@ public class BedrockAnthropicChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockAnthropicChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockAnthropicChatProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockAnthropicChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockAnthropicChatModel.class)).isEmpty(); }); // Explicitly enable the chat auto-configuration. @@ -138,7 +138,7 @@ public class BedrockAnthropicChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockAnthropicChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockAnthropicChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(BedrockAnthropicChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(BedrockAnthropicChatModel.class)).isNotEmpty(); }); // Explicitly disable the chat auto-configuration. @@ -146,7 +146,7 @@ public class BedrockAnthropicChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockAnthropicChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockAnthropicChatProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockAnthropicChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockAnthropicChatModel.class)).isEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/anthropic3/BedrockAnthropic3ChatAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/anthropic3/BedrockAnthropic3ChatAutoConfigurationIT.java index 467fedc7e..8cf6d91b0 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/anthropic3/BedrockAnthropic3ChatAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/anthropic3/BedrockAnthropic3ChatAutoConfigurationIT.java @@ -21,12 +21,12 @@ import java.util.stream.Collectors; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.springframework.ai.bedrock.anthropic3.BedrockAnthropic3ChatModel; import org.springframework.ai.chat.messages.AssistantMessage; import reactor.core.publisher.Flux; import software.amazon.awssdk.regions.Region; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.anthropic3.BedrockAnthropic3ChatClient; import org.springframework.ai.bedrock.anthropic3.api.Anthropic3ChatBedrockApi.AnthropicChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; @@ -69,8 +69,8 @@ public class BedrockAnthropic3ChatAutoConfigurationIT { @Test public void chatCompletion() { contextRunner.run(context -> { - BedrockAnthropic3ChatClient anthropicChatClient = context.getBean(BedrockAnthropic3ChatClient.class); - ChatResponse response = anthropicChatClient.call(new Prompt(List.of(userMessage, systemMessage))); + BedrockAnthropic3ChatModel anthropicChatModel = context.getBean(BedrockAnthropic3ChatModel.class); + ChatResponse response = anthropicChatModel.call(new Prompt(List.of(userMessage, systemMessage))); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); }); } @@ -79,9 +79,9 @@ public class BedrockAnthropic3ChatAutoConfigurationIT { public void chatCompletionStreaming() { contextRunner.run(context -> { - BedrockAnthropic3ChatClient anthropicChatClient = context.getBean(BedrockAnthropic3ChatClient.class); + BedrockAnthropic3ChatModel anthropicChatModel = context.getBean(BedrockAnthropic3ChatModel.class); - Flux response = anthropicChatClient.stream(new Prompt(List.of(userMessage, systemMessage))); + Flux response = anthropicChatModel.stream(new Prompt(List.of(userMessage, systemMessage))); List responses = response.collectList().block(); assertThat(responses.size()).isGreaterThan(2); @@ -130,7 +130,7 @@ public class BedrockAnthropic3ChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockAnthropic3ChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockAnthropic3ChatProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockAnthropic3ChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockAnthropic3ChatModel.class)).isEmpty(); }); // Explicitly enable the chat auto-configuration. @@ -138,7 +138,7 @@ public class BedrockAnthropic3ChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockAnthropic3ChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockAnthropic3ChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(BedrockAnthropic3ChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(BedrockAnthropic3ChatModel.class)).isNotEmpty(); }); // Explicitly disable the chat auto-configuration. @@ -146,7 +146,7 @@ public class BedrockAnthropic3ChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockAnthropic3ChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockAnthropic3ChatProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockAnthropic3ChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockAnthropic3ChatModel.class)).isEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereChatAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereChatAutoConfigurationIT.java index e515f6afd..815030555 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereChatAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereChatAutoConfigurationIT.java @@ -21,13 +21,13 @@ import java.util.stream.Collectors; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.springframework.ai.bedrock.cohere.BedrockCohereChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.AssistantMessage; import reactor.core.publisher.Flux; import software.amazon.awssdk.regions.Region; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.cohere.BedrockCohereChatClient; import org.springframework.ai.bedrock.cohere.api.CohereChatBedrockApi.CohereChatModel; import org.springframework.ai.bedrock.cohere.api.CohereChatBedrockApi.CohereChatRequest.ReturnLikelihoods; import org.springframework.ai.bedrock.cohere.api.CohereChatBedrockApi.CohereChatRequest.Truncate; @@ -72,8 +72,8 @@ public class BedrockCohereChatAutoConfigurationIT { @Test public void chatCompletion() { contextRunner.run(context -> { - BedrockCohereChatClient cohereChatClient = context.getBean(BedrockCohereChatClient.class); - ChatResponse response = cohereChatClient.call(new Prompt(List.of(userMessage, systemMessage))); + BedrockCohereChatModel cohereChatModel = context.getBean(BedrockCohereChatModel.class); + ChatResponse response = cohereChatModel.call(new Prompt(List.of(userMessage, systemMessage))); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); }); } @@ -82,9 +82,9 @@ public class BedrockCohereChatAutoConfigurationIT { public void chatCompletionStreaming() { contextRunner.run(context -> { - BedrockCohereChatClient cohereChatClient = context.getBean(BedrockCohereChatClient.class); + BedrockCohereChatModel cohereChatModel = context.getBean(BedrockCohereChatModel.class); - Flux response = cohereChatClient.stream(new Prompt(List.of(userMessage, systemMessage))); + Flux response = cohereChatModel.stream(new Prompt(List.of(userMessage, systemMessage))); List responses = response.collectList().block(); assertThat(responses.size()).isGreaterThan(2); @@ -146,7 +146,7 @@ public class BedrockCohereChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockCohereChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockCohereChatProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockCohereChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockCohereChatModel.class)).isEmpty(); }); // Explicitly enable the chat auto-configuration. @@ -154,7 +154,7 @@ public class BedrockCohereChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockCohereChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockCohereChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(BedrockCohereChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(BedrockCohereChatModel.class)).isNotEmpty(); }); // Explicitly disable the chat auto-configuration. @@ -162,7 +162,7 @@ public class BedrockCohereChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockCohereChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockCohereChatProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockCohereChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockCohereChatModel.class)).isEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereEmbeddingAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereEmbeddingAutoConfigurationIT.java index 49498f7c5..14d388955 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereEmbeddingAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/cohere/BedrockCohereEmbeddingAutoConfigurationIT.java @@ -18,7 +18,7 @@ package org.springframework.ai.autoconfigure.bedrock.cohere; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.cohere.BedrockCohereEmbeddingClient; +import org.springframework.ai.bedrock.cohere.BedrockCohereEmbeddingModel; import org.springframework.ai.bedrock.cohere.api.CohereEmbeddingBedrockApi.CohereEmbeddingModel; import org.springframework.ai.bedrock.cohere.api.CohereEmbeddingBedrockApi.CohereEmbeddingRequest; import org.springframework.ai.bedrock.cohere.api.CohereEmbeddingBedrockApi.CohereEmbeddingRequest.InputType; @@ -52,12 +52,12 @@ public class BedrockCohereEmbeddingAutoConfigurationIT { @Test public void singleEmbedding() { contextRunner.run(context -> { - BedrockCohereEmbeddingClient embeddingClient = context.getBean(BedrockCohereEmbeddingClient.class); - assertThat(embeddingClient).isNotNull(); - EmbeddingResponse embeddingResponse = embeddingClient.embedForResponse(List.of("Hello World")); + BedrockCohereEmbeddingModel embeddingModel = context.getBean(BedrockCohereEmbeddingModel.class); + assertThat(embeddingModel).isNotNull(); + EmbeddingResponse embeddingResponse = embeddingModel.embedForResponse(List.of("Hello World")); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); - assertThat(embeddingClient.dimensions()).isEqualTo(1024); + assertThat(embeddingModel.dimensions()).isEqualTo(1024); }); } @@ -65,10 +65,10 @@ public class BedrockCohereEmbeddingAutoConfigurationIT { public void batchEmbedding() { contextRunner.run(context -> { - BedrockCohereEmbeddingClient embeddingClient = context.getBean(BedrockCohereEmbeddingClient.class); + BedrockCohereEmbeddingModel embeddingModel = context.getBean(BedrockCohereEmbeddingModel.class); - assertThat(embeddingClient).isNotNull(); - EmbeddingResponse embeddingResponse = embeddingClient + assertThat(embeddingModel).isNotNull(); + EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); assertThat(embeddingResponse.getResults()).hasSize(2); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); @@ -76,7 +76,7 @@ public class BedrockCohereEmbeddingAutoConfigurationIT { assertThat(embeddingResponse.getResults().get(1).getOutput()).isNotEmpty(); assertThat(embeddingResponse.getResults().get(1).getIndex()).isEqualTo(1); - assertThat(embeddingClient.dimensions()).isEqualTo(1024); + assertThat(embeddingModel.dimensions()).isEqualTo(1024); }); } @@ -116,7 +116,7 @@ public class BedrockCohereEmbeddingAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockCohereEmbeddingAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockCohereEmbeddingProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockCohereEmbeddingClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockCohereEmbeddingModel.class)).isEmpty(); }); // Explicitly enable the embedding auto-configuration. @@ -124,7 +124,7 @@ public class BedrockCohereEmbeddingAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockCohereEmbeddingAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockCohereEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(BedrockCohereEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(BedrockCohereEmbeddingModel.class)).isNotEmpty(); }); // Explicitly disable the embedding auto-configuration. @@ -132,7 +132,7 @@ public class BedrockCohereEmbeddingAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockCohereEmbeddingAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockCohereEmbeddingProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockCohereEmbeddingClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockCohereEmbeddingModel.class)).isEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/jurassic2/BedrockAi21Jurassic2ChatAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/jurassic2/BedrockAi21Jurassic2ChatAutoConfigurationIT.java index 057e46483..a2cbf3f0d 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/jurassic2/BedrockAi21Jurassic2ChatAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/jurassic2/BedrockAi21Jurassic2ChatAutoConfigurationIT.java @@ -21,7 +21,7 @@ import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; import org.springframework.ai.autoconfigure.bedrock.jurrasic2.BedrockAi21Jurassic2ChatAutoConfiguration; import org.springframework.ai.autoconfigure.bedrock.jurrasic2.BedrockAi21Jurassic2ChatProperties; -import org.springframework.ai.bedrock.jurassic2.BedrockAi21Jurassic2ChatClient; +import org.springframework.ai.bedrock.jurassic2.BedrockAi21Jurassic2ChatModel; import org.springframework.ai.bedrock.jurassic2.api.Ai21Jurassic2ChatBedrockApi; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.Message; @@ -69,9 +69,8 @@ public class BedrockAi21Jurassic2ChatAutoConfigurationIT { @Test public void chatCompletion() { contextRunner.run(context -> { - BedrockAi21Jurassic2ChatClient ai21Jurassic2ChatClient = context - .getBean(BedrockAi21Jurassic2ChatClient.class); - ChatResponse response = ai21Jurassic2ChatClient.call(new Prompt(List.of(userMessage, systemMessage))); + BedrockAi21Jurassic2ChatModel ai21Jurassic2ChatModel = context.getBean(BedrockAi21Jurassic2ChatModel.class); + ChatResponse response = ai21Jurassic2ChatModel.call(new Prompt(List.of(userMessage, systemMessage))); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); }); } @@ -111,7 +110,7 @@ public class BedrockAi21Jurassic2ChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockAi21Jurassic2ChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockAi21Jurassic2ChatProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockAi21Jurassic2ChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockAi21Jurassic2ChatModel.class)).isEmpty(); }); // Explicitly enable the chat auto-configuration. @@ -119,7 +118,7 @@ public class BedrockAi21Jurassic2ChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockAi21Jurassic2ChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockAi21Jurassic2ChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(BedrockAi21Jurassic2ChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(BedrockAi21Jurassic2ChatModel.class)).isNotEmpty(); }); // Explicitly disable the chat auto-configuration. @@ -127,7 +126,7 @@ public class BedrockAi21Jurassic2ChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockAi21Jurassic2ChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockAi21Jurassic2ChatProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockAi21Jurassic2ChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockAi21Jurassic2ChatModel.class)).isEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/llama/BedrockLlamaChatAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/llama/BedrockLlamaChatAutoConfigurationIT.java index 6ca69b655..5e9241181 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/llama/BedrockLlamaChatAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/llama/BedrockLlamaChatAutoConfigurationIT.java @@ -21,13 +21,13 @@ import java.util.stream.Collectors; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.springframework.ai.bedrock.llama.BedrockLlamaChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.AssistantMessage; import reactor.core.publisher.Flux; import software.amazon.awssdk.regions.Region; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.llama.BedrockLlamaChatClient; import org.springframework.ai.bedrock.llama.api.LlamaChatBedrockApi.LlamaChatModel; import org.springframework.ai.chat.Generation; import org.springframework.ai.chat.prompt.Prompt; @@ -71,8 +71,8 @@ public class BedrockLlamaChatAutoConfigurationIT { @Test public void chatCompletion() { contextRunner.run(context -> { - BedrockLlamaChatClient llamaChatClient = context.getBean(BedrockLlamaChatClient.class); - ChatResponse response = llamaChatClient.call(new Prompt(List.of(userMessage, systemMessage))); + BedrockLlamaChatModel llamaChatModel = context.getBean(BedrockLlamaChatModel.class); + ChatResponse response = llamaChatModel.call(new Prompt(List.of(userMessage, systemMessage))); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); }); } @@ -81,9 +81,9 @@ public class BedrockLlamaChatAutoConfigurationIT { public void chatCompletionStreaming() { contextRunner.run(context -> { - BedrockLlamaChatClient llamaChatClient = context.getBean(BedrockLlamaChatClient.class); + BedrockLlamaChatModel llamaChatModel = context.getBean(BedrockLlamaChatModel.class); - Flux response = llamaChatClient.stream(new Prompt(List.of(userMessage, systemMessage))); + Flux response = llamaChatModel.stream(new Prompt(List.of(userMessage, systemMessage))); List responses = response.collectList().block(); assertThat(responses.size()).isGreaterThan(2); @@ -133,7 +133,7 @@ public class BedrockLlamaChatAutoConfigurationIT { new ApplicationContextRunner().withConfiguration(AutoConfigurations.of(BedrockLlamaChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockLlamaChatProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockLlamaChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockLlamaChatModel.class)).isEmpty(); }); // Explicitly enable the chat auto-configuration. @@ -141,7 +141,7 @@ public class BedrockLlamaChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockLlamaChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockLlamaChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(BedrockLlamaChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(BedrockLlamaChatModel.class)).isNotEmpty(); }); // Explicitly disable the chat auto-configuration. @@ -149,7 +149,7 @@ public class BedrockLlamaChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockLlamaChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockLlamaChatProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockLlamaChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockLlamaChatModel.class)).isEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanChatAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanChatAutoConfigurationIT.java index 93c654178..0783d8149 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanChatAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanChatAutoConfigurationIT.java @@ -27,7 +27,7 @@ import reactor.core.publisher.Flux; import software.amazon.awssdk.regions.Region; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.titan.BedrockTitanChatClient; +import org.springframework.ai.bedrock.titan.BedrockTitanChatModel; import org.springframework.ai.bedrock.titan.api.TitanChatBedrockApi.TitanChatModel; import org.springframework.ai.chat.Generation; import org.springframework.ai.chat.prompt.Prompt; @@ -70,8 +70,8 @@ public class BedrockTitanChatAutoConfigurationIT { @Test public void chatCompletion() { contextRunner.run(context -> { - BedrockTitanChatClient chatClient = context.getBean(BedrockTitanChatClient.class); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage, systemMessage))); + BedrockTitanChatModel chatModel = context.getBean(BedrockTitanChatModel.class); + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage, systemMessage))); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); }); } @@ -80,9 +80,9 @@ public class BedrockTitanChatAutoConfigurationIT { public void chatCompletionStreaming() { contextRunner.run(context -> { - BedrockTitanChatClient chatClient = context.getBean(BedrockTitanChatClient.class); + BedrockTitanChatModel chatModel = context.getBean(BedrockTitanChatModel.class); - Flux response = chatClient.stream(new Prompt(List.of(userMessage, systemMessage))); + Flux response = chatModel.stream(new Prompt(List.of(userMessage, systemMessage))); List responses = response.collectList().block(); assertThat(responses.size()).isGreaterThan(1); @@ -137,7 +137,7 @@ public class BedrockTitanChatAutoConfigurationIT { new ApplicationContextRunner().withConfiguration(AutoConfigurations.of(BedrockTitanChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockTitanChatProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockTitanChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockTitanChatModel.class)).isEmpty(); }); // Explicitly enable the chat auto-configuration. @@ -145,7 +145,7 @@ public class BedrockTitanChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockTitanChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockTitanChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(BedrockTitanChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(BedrockTitanChatModel.class)).isNotEmpty(); }); // Explicitly disable the chat auto-configuration. @@ -153,7 +153,7 @@ public class BedrockTitanChatAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockTitanChatAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockTitanChatProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockTitanChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockTitanChatModel.class)).isEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanEmbeddingAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanEmbeddingAutoConfigurationIT.java index 637598fcd..5a5a2ad4c 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanEmbeddingAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/bedrock/titan/BedrockTitanEmbeddingAutoConfigurationIT.java @@ -23,8 +23,8 @@ import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import software.amazon.awssdk.regions.Region; import org.springframework.ai.autoconfigure.bedrock.BedrockAwsConnectionProperties; -import org.springframework.ai.bedrock.titan.BedrockTitanEmbeddingClient; -import org.springframework.ai.bedrock.titan.BedrockTitanEmbeddingClient.InputType; +import org.springframework.ai.bedrock.titan.BedrockTitanEmbeddingModel; +import org.springframework.ai.bedrock.titan.BedrockTitanEmbeddingModel.InputType; import org.springframework.ai.bedrock.titan.api.TitanEmbeddingBedrockApi.TitanEmbeddingModel; import org.springframework.ai.embedding.EmbeddingResponse; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -52,31 +52,31 @@ public class BedrockTitanEmbeddingAutoConfigurationIT { @Test public void singleTextEmbedding() { contextRunner.withPropertyValues("spring.ai.bedrock.titan.embedding.inputType=TEXT").run(context -> { - BedrockTitanEmbeddingClient embeddingClient = context.getBean(BedrockTitanEmbeddingClient.class); - assertThat(embeddingClient).isNotNull(); - EmbeddingResponse embeddingResponse = embeddingClient.embedForResponse(List.of("Hello World")); + BedrockTitanEmbeddingModel embeddingModel = context.getBean(BedrockTitanEmbeddingModel.class); + assertThat(embeddingModel).isNotNull(); + EmbeddingResponse embeddingResponse = embeddingModel.embedForResponse(List.of("Hello World")); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); - assertThat(embeddingClient.dimensions()).isEqualTo(1024); + assertThat(embeddingModel.dimensions()).isEqualTo(1024); }); } @Test public void singleImageEmbedding() { contextRunner.withPropertyValues("spring.ai.bedrock.titan.embedding.inputType=IMAGE").run(context -> { - BedrockTitanEmbeddingClient embeddingClient = context.getBean(BedrockTitanEmbeddingClient.class); - assertThat(embeddingClient).isNotNull(); + BedrockTitanEmbeddingModel embeddingModel = context.getBean(BedrockTitanEmbeddingModel.class); + assertThat(embeddingModel).isNotNull(); byte[] image = new DefaultResourceLoader().getResource("classpath:/spring_framework.png") .getContentAsByteArray(); var base64Image = Base64.getEncoder().encodeToString(image); - EmbeddingResponse embeddingResponse = embeddingClient.embedForResponse(List.of(base64Image)); + EmbeddingResponse embeddingResponse = embeddingModel.embedForResponse(List.of(base64Image)); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); - assertThat(embeddingClient.dimensions()).isEqualTo(1024); + assertThat(embeddingModel.dimensions()).isEqualTo(1024); }); } @@ -111,7 +111,7 @@ public class BedrockTitanEmbeddingAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockTitanEmbeddingAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockTitanEmbeddingProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockTitanEmbeddingClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockTitanEmbeddingModel.class)).isEmpty(); }); // Explicitly enable the embedding auto-configuration. @@ -119,7 +119,7 @@ public class BedrockTitanEmbeddingAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockTitanEmbeddingAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockTitanEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(BedrockTitanEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(BedrockTitanEmbeddingModel.class)).isNotEmpty(); }); // Explicitly disable the embedding auto-configuration. @@ -127,7 +127,7 @@ public class BedrockTitanEmbeddingAutoConfigurationIT { .withConfiguration(AutoConfigurations.of(BedrockTitanEmbeddingAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(BedrockTitanEmbeddingProperties.class)).isEmpty(); - assertThat(context.getBeansOfType(BedrockTitanEmbeddingClient.class)).isEmpty(); + assertThat(context.getBeansOfType(BedrockTitanEmbeddingModel.class)).isEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackInPromptIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackInPromptIT.java index 7026812a6..e59895a89 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackInPromptIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackInPromptIT.java @@ -25,7 +25,7 @@ import org.springframework.ai.chat.Generation; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.minimax.MiniMaxChatClient; +import org.springframework.ai.minimax.MiniMaxChatModel; import org.springframework.ai.minimax.MiniMaxChatOptions; import org.springframework.ai.model.function.FunctionCallbackWrapper; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -55,7 +55,7 @@ public class FunctionCallbackInPromptIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.minimax.chat.options.model=abab6-chat").run(context -> { - MiniMaxChatClient chatClient = context.getBean(MiniMaxChatClient.class); + MiniMaxChatModel chatModel = context.getBean(MiniMaxChatModel.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -67,7 +67,7 @@ public class FunctionCallbackInPromptIT { .build())) .build(); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), promptOptions)); + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), promptOptions)); logger.info("Response: {}", response); @@ -80,7 +80,7 @@ public class FunctionCallbackInPromptIT { contextRunner.withPropertyValues("spring.ai.minimax.chat.options.model=abab6-chat").run(context -> { - MiniMaxChatClient chatClient = context.getBean(MiniMaxChatClient.class); + MiniMaxChatModel chatModel = context.getBean(MiniMaxChatModel.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -92,7 +92,7 @@ public class FunctionCallbackInPromptIT { .build())) .build(); - Flux response = chatClient.stream(new Prompt(List.of(userMessage), promptOptions)); + Flux response = chatModel.stream(new Prompt(List.of(userMessage), promptOptions)); String content = response.collectList() .block() diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWithPlainFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWithPlainFunctionBeanIT.java index 1c1492e65..b16d358cd 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWithPlainFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWithPlainFunctionBeanIT.java @@ -25,7 +25,7 @@ import org.springframework.ai.chat.Generation; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.minimax.MiniMaxChatClient; +import org.springframework.ai.minimax.MiniMaxChatModel; import org.springframework.ai.minimax.MiniMaxChatOptions; import org.springframework.ai.model.function.FunctionCallingOptions; import org.springframework.ai.model.function.FunctionCallingOptionsBuilder.PortableFunctionCallingOptions; @@ -61,12 +61,12 @@ class FunctionCallbackWithPlainFunctionBeanIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.minimax.chat.options.model=abab6-chat").run(context -> { - MiniMaxChatClient chatClient = context.getBean(MiniMaxChatClient.class); + MiniMaxChatModel chatModel = context.getBean(MiniMaxChatModel.class); // Test weatherFunction UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), MiniMaxChatOptions.builder().withFunction("weatherFunction").build())); logger.info("Response: {}", response); @@ -74,7 +74,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); // Test weatherFunctionTwo - response = chatClient.call(new Prompt(List.of(userMessage), + response = chatModel.call(new Prompt(List.of(userMessage), MiniMaxChatOptions.builder().withFunction("weatherFunctionTwo").build())); logger.info("Response: {}", response); @@ -88,7 +88,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { void functionCallWithPortableFunctionCallingOptions() { contextRunner.withPropertyValues("spring.ai.minimax.chat.options.model=abab6-chat").run(context -> { - MiniMaxChatClient chatClient = context.getBean(MiniMaxChatClient.class); + MiniMaxChatModel chatModel = context.getBean(MiniMaxChatModel.class); // Test weatherFunction UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -97,7 +97,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { .withFunction("weatherFunction") .build(); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), functionOptions)); + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), functionOptions)); logger.info("Response: {}", response); }); @@ -107,12 +107,12 @@ class FunctionCallbackWithPlainFunctionBeanIT { void streamFunctionCallTest() { contextRunner.withPropertyValues("spring.ai.minimax.chat.options.model=abab6-chat").run(context -> { - MiniMaxChatClient chatClient = context.getBean(MiniMaxChatClient.class); + MiniMaxChatModel chatModel = context.getBean(MiniMaxChatModel.class); // Test weatherFunction UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - Flux response = chatClient.stream(new Prompt(List.of(userMessage), + Flux response = chatModel.stream(new Prompt(List.of(userMessage), MiniMaxChatOptions.builder().withFunction("weatherFunction").build())); String content = response.collectList() @@ -130,7 +130,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { assertThat(content).containsAnyOf("15.0", "15"); // Test weatherFunctionTwo - response = chatClient.stream(new Prompt(List.of(userMessage), + response = chatModel.stream(new Prompt(List.of(userMessage), MiniMaxChatOptions.builder().withFunction("weatherFunctionTwo").build())); content = response.collectList() diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWrapperIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWrapperIT.java index 75530376f..47ddb1500 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWrapperIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/FunctionCallbackWrapperIT.java @@ -25,7 +25,7 @@ import org.springframework.ai.chat.Generation; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.minimax.MiniMaxChatClient; +import org.springframework.ai.minimax.MiniMaxChatModel; import org.springframework.ai.minimax.MiniMaxChatOptions; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackWrapper; @@ -59,11 +59,11 @@ public class FunctionCallbackWrapperIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.minimax.chat.options.model=abab6-chat").run(context -> { - MiniMaxChatClient chatClient = context.getBean(MiniMaxChatClient.class); + MiniMaxChatModel chatModel = context.getBean(MiniMaxChatModel.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - ChatResponse response = chatClient.call( + ChatResponse response = chatModel.call( new Prompt(List.of(userMessage), MiniMaxChatOptions.builder().withFunction("WeatherInfo").build())); logger.info("Response: {}", response); @@ -77,11 +77,11 @@ public class FunctionCallbackWrapperIT { void streamFunctionCallTest() { contextRunner.withPropertyValues("spring.ai.minimax.chat.options.model=abab6-chat").run(context -> { - MiniMaxChatClient chatClient = context.getBean(MiniMaxChatClient.class); + MiniMaxChatModel chatModel = context.getBean(MiniMaxChatModel.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - Flux response = chatClient.stream( + Flux response = chatModel.stream( new Prompt(List.of(userMessage), MiniMaxChatOptions.builder().withFunction("WeatherInfo").build())); String content = response.collectList() diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfigurationIT.java index d400b2c47..ac5e6ac05 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxAutoConfigurationIT.java @@ -24,8 +24,8 @@ import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.embedding.EmbeddingResponse; -import org.springframework.ai.minimax.MiniMaxChatClient; -import org.springframework.ai.minimax.MiniMaxEmbeddingClient; +import org.springframework.ai.minimax.MiniMaxChatModel; +import org.springframework.ai.minimax.MiniMaxEmbeddingModel; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -52,8 +52,8 @@ public class MiniMaxAutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - MiniMaxChatClient client = context.getBean(MiniMaxChatClient.class); - String response = client.call("Hello"); + MiniMaxChatModel chatModel = context.getBean(MiniMaxChatModel.class); + String response = chatModel.call("Hello"); assertThat(response).isNotEmpty(); logger.info("Response: " + response); }); @@ -62,8 +62,8 @@ public class MiniMaxAutoConfigurationIT { @Test void generateStreaming() { contextRunner.run(context -> { - MiniMaxChatClient client = context.getBean(MiniMaxChatClient.class); - Flux responseFlux = client.stream(new Prompt(new UserMessage("Hello"))); + MiniMaxChatModel chatModel = context.getBean(MiniMaxChatModel.class); + Flux responseFlux = chatModel.stream(new Prompt(new UserMessage("Hello"))); String response = responseFlux.collectList().block().stream().map(chatResponse -> { return chatResponse.getResults().get(0).getOutput().getContent(); }).collect(Collectors.joining()); @@ -76,9 +76,9 @@ public class MiniMaxAutoConfigurationIT { @Test void embedding() { contextRunner.run(context -> { - MiniMaxEmbeddingClient embeddingClient = context.getBean(MiniMaxEmbeddingClient.class); + MiniMaxEmbeddingModel embeddingModel = context.getBean(MiniMaxEmbeddingModel.class); - EmbeddingResponse embeddingResponse = embeddingClient + EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); assertThat(embeddingResponse.getResults()).hasSize(2); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); @@ -86,7 +86,7 @@ public class MiniMaxAutoConfigurationIT { assertThat(embeddingResponse.getResults().get(1).getOutput()).isNotEmpty(); assertThat(embeddingResponse.getResults().get(1).getIndex()).isEqualTo(1); - assertThat(embeddingClient.dimensions()).isEqualTo(1536); + assertThat(embeddingModel.dimensions()).isEqualTo(1536); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxPropertiesTests.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxPropertiesTests.java index b2d834b49..07b2d4e6a 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxPropertiesTests.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/minimax/MiniMaxPropertiesTests.java @@ -19,8 +19,8 @@ import org.junit.jupiter.api.Test; import org.skyscreamer.jsonassert.JSONAssert; import org.skyscreamer.jsonassert.JSONCompareMode; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; -import org.springframework.ai.minimax.MiniMaxChatClient; -import org.springframework.ai.minimax.MiniMaxEmbeddingClient; +import org.springframework.ai.minimax.MiniMaxChatModel; +import org.springframework.ai.minimax.MiniMaxEmbeddingModel; import org.springframework.ai.minimax.api.MiniMaxApi; import org.springframework.ai.model.ModelOptionsUtils; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -270,7 +270,7 @@ public class MiniMaxPropertiesTests { RestClientAutoConfiguration.class, MiniMaxAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(MiniMaxEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(MiniMaxEmbeddingClient.class)).isEmpty(); + assertThat(context.getBeansOfType(MiniMaxEmbeddingModel.class)).isEmpty(); }); new ApplicationContextRunner() @@ -279,7 +279,7 @@ public class MiniMaxPropertiesTests { RestClientAutoConfiguration.class, MiniMaxAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(MiniMaxEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(MiniMaxEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(MiniMaxEmbeddingModel.class)).isNotEmpty(); }); new ApplicationContextRunner() @@ -289,7 +289,7 @@ public class MiniMaxPropertiesTests { RestClientAutoConfiguration.class, MiniMaxAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(MiniMaxEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(MiniMaxEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(MiniMaxEmbeddingModel.class)).isNotEmpty(); }); } @@ -302,7 +302,7 @@ public class MiniMaxPropertiesTests { RestClientAutoConfiguration.class, MiniMaxAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(MiniMaxChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(MiniMaxChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(MiniMaxChatModel.class)).isEmpty(); }); new ApplicationContextRunner() @@ -311,7 +311,7 @@ public class MiniMaxPropertiesTests { RestClientAutoConfiguration.class, MiniMaxAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(MiniMaxChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(MiniMaxChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(MiniMaxChatModel.class)).isNotEmpty(); }); new ApplicationContextRunner() @@ -321,7 +321,7 @@ public class MiniMaxPropertiesTests { RestClientAutoConfiguration.class, MiniMaxAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(MiniMaxChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(MiniMaxChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(MiniMaxChatModel.class)).isNotEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/MistralAiAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/MistralAiAutoConfigurationIT.java index 25816e8c0..a629fe52e 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/MistralAiAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/MistralAiAutoConfigurationIT.java @@ -22,6 +22,7 @@ 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; @@ -29,8 +30,7 @@ import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.embedding.EmbeddingResponse; -import org.springframework.ai.mistralai.MistralAiChatClient; -import org.springframework.ai.mistralai.MistralAiEmbeddingClient; +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; @@ -54,8 +54,8 @@ public class MistralAiAutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - MistralAiChatClient client = context.getBean(MistralAiChatClient.class); - String response = client.call("Hello"); + MistralAiChatModel chatModel = context.getBean(MistralAiChatModel.class); + String response = chatModel.call("Hello"); assertThat(response).isNotEmpty(); logger.info("Response: " + response); }); @@ -64,8 +64,8 @@ public class MistralAiAutoConfigurationIT { @Test void generateStreaming() { contextRunner.run(context -> { - MistralAiChatClient client = context.getBean(MistralAiChatClient.class); - Flux responseFlux = client.stream(new Prompt(new UserMessage("Hello"))); + MistralAiChatModel chatModel = context.getBean(MistralAiChatModel.class); + Flux responseFlux = chatModel.stream(new Prompt(new UserMessage("Hello"))); String response = responseFlux.collectList().block().stream().map(chatResponse -> { return chatResponse.getResults().get(0).getOutput().getContent(); }).collect(Collectors.joining()); @@ -78,9 +78,9 @@ public class MistralAiAutoConfigurationIT { @Test void embedding() { contextRunner.run(context -> { - MistralAiEmbeddingClient embeddingClient = context.getBean(MistralAiEmbeddingClient.class); + MistralAiEmbeddingModel embeddingModel = context.getBean(MistralAiEmbeddingModel.class); - EmbeddingResponse embeddingResponse = embeddingClient + EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); assertThat(embeddingResponse.getResults()).hasSize(2); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); @@ -88,7 +88,7 @@ public class MistralAiAutoConfigurationIT { assertThat(embeddingResponse.getResults().get(1).getOutput()).isNotEmpty(); assertThat(embeddingResponse.getResults().get(1).getIndex()).isEqualTo(1); - assertThat(embeddingClient.dimensions()).isEqualTo(1024); + assertThat(embeddingModel.dimensions()).isEqualTo(1024); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/PaymentStatusBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/PaymentStatusBeanIT.java index d4259f3ee..b4a581982 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/PaymentStatusBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/PaymentStatusBeanIT.java @@ -30,7 +30,7 @@ import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.mistralai.MistralAiChatClient; +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; @@ -60,9 +60,9 @@ class PaymentStatusBeanIT { .withPropertyValues("spring.ai.mistralai.chat.options.model=" + MistralAiApi.ChatModel.LARGE.getValue()) .run(context -> { - MistralAiChatClient chatClient = context.getBean(MistralAiChatClient.class); + MistralAiChatModel chatModel = context.getBean(MistralAiChatModel.class); - ChatResponse response = chatClient + ChatResponse response = chatModel .call(new Prompt(List.of(new UserMessage("What's the status of my transaction with id T1001?")), MistralAiChatOptions.builder() .withFunction("retrievePaymentStatus") diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/PaymentStatusBeanOpenAiIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/PaymentStatusBeanOpenAiIT.java index 15796f921..069fe5b23 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/PaymentStatusBeanOpenAiIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/PaymentStatusBeanOpenAiIT.java @@ -31,7 +31,7 @@ import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.mistralai.api.MistralAiApi; -import org.springframework.ai.openai.OpenAiChatClient; +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; @@ -43,7 +43,7 @@ import org.springframework.context.annotation.Description; import static org.assertj.core.api.Assertions.assertThat; /** - * Same test as {@link PaymentStatusBeanIT.java} but using {@link OpenAiChatClient} for + * Same test as {@link PaymentStatusBeanIT.java} but using {@link OpenAiChatModel} for * Mistral AI Function Calling implementation. * * @author Christian Tzolov @@ -67,9 +67,9 @@ class PaymentStatusBeanOpenAiIT { .withPropertyValues("spring.ai.openai.chat.options.model=" + MistralAiApi.ChatModel.SMALL.getValue()) .run(context -> { - OpenAiChatClient chatClient = context.getBean(OpenAiChatClient.class); + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); - ChatResponse response = chatClient + ChatResponse response = chatModel .call(new Prompt(List.of(new UserMessage("What's the status of my transaction with id T1001?")), OpenAiChatOptions.builder() .withFunction("retrievePaymentStatus") diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/PaymentStatusPromptIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/PaymentStatusPromptIT.java index 58bd9e84c..086bc6683 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/PaymentStatusPromptIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/PaymentStatusPromptIT.java @@ -30,7 +30,7 @@ import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.mistralai.MistralAiChatClient; +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; @@ -71,7 +71,7 @@ public class PaymentStatusPromptIT { .withPropertyValues("spring.ai.mistralai.chat.options.model=" + MistralAiApi.ChatModel.SMALL.getValue()) .run(context -> { - MistralAiChatClient chatClient = context.getBean(MistralAiChatClient.class); + MistralAiChatModel chatModel = context.getBean(MistralAiChatModel.class); UserMessage userMessage = new UserMessage("What's the status of my transaction with id T1001?"); @@ -86,7 +86,7 @@ public class PaymentStatusPromptIT { .build())) .build(); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), promptOptions)); + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), promptOptions)); logger.info("Response: {}", response); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/WeatherServicePromptIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/WeatherServicePromptIT.java index 6bf052f56..1cfc7f245 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/WeatherServicePromptIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/mistralai/tool/WeatherServicePromptIT.java @@ -33,7 +33,7 @@ import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.mistralai.MistralAiChatClient; +import org.springframework.ai.mistralai.MistralAiChatModel; import org.springframework.ai.mistralai.MistralAiChatOptions; import org.springframework.ai.mistralai.api.MistralAiApi; import org.springframework.ai.mistralai.api.MistralAiApi.ChatCompletionRequest.ToolChoice; @@ -64,7 +64,7 @@ public class WeatherServicePromptIT { .withPropertyValues("spring.ai.mistralai.chat.options.model=" + MistralAiApi.ChatModel.LARGE.getValue()) .run(context -> { - MistralAiChatClient chatClient = context.getBean(MistralAiChatClient.class); + MistralAiChatModel chatModel = context.getBean(MistralAiChatModel.class); UserMessage userMessage = new UserMessage("What's the weather like in Paris?"); // UserMessage userMessage = new UserMessage("What's the weather like in @@ -79,7 +79,7 @@ public class WeatherServicePromptIT { .build())) .build(); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), promptOptions)); + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), promptOptions)); logger.info("Response: {}", response); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaChatAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaChatAutoConfigurationIT.java index affa26197..4ed3448c7 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaChatAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaChatAutoConfigurationIT.java @@ -26,7 +26,7 @@ import org.junit.jupiter.api.Test; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; import org.springframework.ai.chat.messages.AssistantMessage; -import org.springframework.ai.ollama.OllamaChatClient; +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; @@ -104,8 +104,8 @@ public class OllamaChatAutoConfigurationIT { @Test public void chatCompletion() { contextRunner.run(context -> { - OllamaChatClient chatClient = context.getBean(OllamaChatClient.class); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage, systemMessage))); + OllamaChatModel chatModel = context.getBean(OllamaChatModel.class); + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage, systemMessage))); assertThat(response.getResult().getOutput().getContent()).contains("Blackbeard"); }); } @@ -114,9 +114,9 @@ public class OllamaChatAutoConfigurationIT { public void chatCompletionStreaming() { contextRunner.run(context -> { - OllamaChatClient chatClient = context.getBean(OllamaChatClient.class); + OllamaChatModel chatModel = context.getBean(OllamaChatModel.class); - Flux response = chatClient.stream(new Prompt(List.of(userMessage, systemMessage))); + Flux response = chatModel.stream(new Prompt(List.of(userMessage, systemMessage))); List responses = response.collectList().block(); assertThat(responses.size()).isGreaterThan(1); @@ -136,17 +136,17 @@ public class OllamaChatAutoConfigurationIT { void chatActivation() { contextRunner.withPropertyValues("spring.ai.ollama.chat.enabled=false").run(context -> { assertThat(context.getBeansOfType(OllamaChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OllamaChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(OllamaChatModel.class)).isEmpty(); }); contextRunner.run(context -> { assertThat(context.getBeansOfType(OllamaChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OllamaChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(OllamaChatModel.class)).isNotEmpty(); }); contextRunner.withPropertyValues("spring.ai.ollama.chat.enabled=true").run(context -> { assertThat(context.getBeansOfType(OllamaChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OllamaChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(OllamaChatModel.class)).isNotEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaEmbeddingAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaEmbeddingAutoConfigurationIT.java index b9afde511..0599e0da6 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaEmbeddingAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/ollama/OllamaEmbeddingAutoConfigurationIT.java @@ -27,7 +27,7 @@ import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.embedding.EmbeddingResponse; -import org.springframework.ai.ollama.OllamaEmbeddingClient; +import org.springframework.ai.ollama.OllamaEmbeddingModel; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -69,12 +69,12 @@ public class OllamaEmbeddingAutoConfigurationIT { @Test public void singleTextEmbedding() { contextRunner.run(context -> { - OllamaEmbeddingClient embeddingClient = context.getBean(OllamaEmbeddingClient.class); - assertThat(embeddingClient).isNotNull(); - EmbeddingResponse embeddingResponse = embeddingClient.embedForResponse(List.of("Hello World")); + OllamaEmbeddingModel embeddingModel = context.getBean(OllamaEmbeddingModel.class); + assertThat(embeddingModel).isNotNull(); + EmbeddingResponse embeddingResponse = embeddingModel.embedForResponse(List.of("Hello World")); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); - assertThat(embeddingClient.dimensions()).isEqualTo(3200); + assertThat(embeddingModel.dimensions()).isEqualTo(3200); }); } @@ -82,17 +82,17 @@ public class OllamaEmbeddingAutoConfigurationIT { void embeddingActivation() { contextRunner.withPropertyValues("spring.ai.ollama.embedding.enabled=false").run(context -> { assertThat(context.getBeansOfType(OllamaEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OllamaEmbeddingClient.class)).isEmpty(); + assertThat(context.getBeansOfType(OllamaEmbeddingModel.class)).isEmpty(); }); contextRunner.run(context -> { assertThat(context.getBeansOfType(OllamaEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OllamaEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(OllamaEmbeddingModel.class)).isNotEmpty(); }); contextRunner.withPropertyValues("spring.ai.ollama.embedding.enabled=true").run(context -> { assertThat(context.getBeansOfType(OllamaEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OllamaEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(OllamaEmbeddingModel.class)).isNotEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfigurationIT.java index 44c2d7375..d9e7a381a 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/OpenAiAutoConfigurationIT.java @@ -27,18 +27,15 @@ import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.image.ImagePrompt; import org.springframework.ai.image.ImageResponse; -import org.springframework.ai.openai.OpenAiImageClient; +import org.springframework.ai.openai.*; import org.springframework.core.io.ClassPathResource; import org.springframework.core.io.Resource; import reactor.core.publisher.Flux; -import org.springframework.ai.openai.OpenAiAudioSpeechClient; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.embedding.EmbeddingResponse; -import org.springframework.ai.openai.OpenAiAudioTranscriptionClient; -import org.springframework.ai.openai.OpenAiChatClient; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +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; @@ -58,8 +55,8 @@ public class OpenAiAutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - OpenAiChatClient client = context.getBean(OpenAiChatClient.class); - String response = client.call("Hello"); + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); + String response = chatModel.call("Hello"); assertThat(response).isNotEmpty(); logger.info("Response: " + response); }); @@ -68,9 +65,9 @@ public class OpenAiAutoConfigurationIT { @Test void transcribe() { contextRunner.run(context -> { - OpenAiAudioTranscriptionClient client = context.getBean(OpenAiAudioTranscriptionClient.class); + OpenAiAudioTranscriptionModel transcriptionModel = context.getBean(OpenAiAudioTranscriptionModel.class); Resource audioFile = new ClassPathResource("/speech/jfk.flac"); - String response = client.call(audioFile); + String response = transcriptionModel.call(audioFile); assertThat(response).isNotEmpty(); logger.info("Response: " + response); }); @@ -79,8 +76,8 @@ public class OpenAiAutoConfigurationIT { @Test void speech() { contextRunner.run(context -> { - OpenAiAudioSpeechClient client = context.getBean(OpenAiAudioSpeechClient.class); - byte[] response = client.call("H"); + OpenAiAudioSpeechModel speechModel = context.getBean(OpenAiAudioSpeechModel.class); + byte[] response = speechModel.call("H"); assertThat(response).isNotNull(); assertThat(verifyMp3FrameHeader(response)) .withFailMessage("Expected MP3 frame header to be present in the response, but it was not found.") @@ -105,8 +102,8 @@ public class OpenAiAutoConfigurationIT { @Test void generateStreaming() { contextRunner.run(context -> { - OpenAiChatClient client = context.getBean(OpenAiChatClient.class); - Flux responseFlux = client.stream(new Prompt(new UserMessage("Hello"))); + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); + Flux responseFlux = chatModel.stream(new Prompt(new UserMessage("Hello"))); String response = responseFlux.collectList().block().stream().map(chatResponse -> { return chatResponse.getResults().get(0).getOutput().getContent(); }).collect(Collectors.joining()); @@ -119,9 +116,9 @@ public class OpenAiAutoConfigurationIT { @Test void embedding() { contextRunner.run(context -> { - OpenAiEmbeddingClient embeddingClient = context.getBean(OpenAiEmbeddingClient.class); + OpenAiEmbeddingModel embeddingModel = context.getBean(OpenAiEmbeddingModel.class); - EmbeddingResponse embeddingResponse = embeddingClient + EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); assertThat(embeddingResponse.getResults()).hasSize(2); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); @@ -129,15 +126,15 @@ public class OpenAiAutoConfigurationIT { assertThat(embeddingResponse.getResults().get(1).getOutput()).isNotEmpty(); assertThat(embeddingResponse.getResults().get(1).getIndex()).isEqualTo(1); - assertThat(embeddingClient.dimensions()).isEqualTo(1536); + assertThat(embeddingModel.dimensions()).isEqualTo(1536); }); } @Test void generateImage() { contextRunner.withPropertyValues("spring.ai.openai.image.options.size=1024x1024").run(context -> { - OpenAiImageClient client = context.getBean(OpenAiImageClient.class); - ImageResponse imageResponse = client.call(new ImagePrompt("forest")); + OpenAiImageModel imageModel = context.getBean(OpenAiImageModel.class); + ImageResponse imageResponse = imageModel.call(new ImagePrompt("forest")); assertThat(imageResponse.getResults()).hasSize(1); assertThat(imageResponse.getResult().getOutput().getUrl()).isNotEmpty(); logger.info("Generated image: " + imageResponse.getResult().getOutput().getUrl()); @@ -151,8 +148,8 @@ public class OpenAiAutoConfigurationIT { .withPropertyValues("spring.ai.openai.image.options.model=dall-e-2", "spring.ai.openai.image.options.size=256x256") .run(context -> { - OpenAiImageClient client = context.getBean(OpenAiImageClient.class); - ImageResponse imageResponse = client.call(new ImagePrompt("forest")); + OpenAiImageModel imageModel = context.getBean(OpenAiImageModel.class); + ImageResponse imageResponse = imageModel.call(new ImagePrompt("forest")); assertThat(imageResponse.getResults()).hasSize(1); assertThat(imageResponse.getResult().getOutput().getUrl()).isNotEmpty(); logger.info("Generated image: " + imageResponse.getResult().getOutput().getUrl()); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/OpenAiPropertiesTests.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/OpenAiPropertiesTests.java index 10583ecd4..0956020dc 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/OpenAiPropertiesTests.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/OpenAiPropertiesTests.java @@ -21,9 +21,9 @@ import org.skyscreamer.jsonassert.JSONCompareMode; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; import org.springframework.ai.model.ModelOptionsUtils; -import org.springframework.ai.openai.OpenAiChatClient; -import org.springframework.ai.openai.OpenAiEmbeddingClient; -import org.springframework.ai.openai.OpenAiImageClient; +import org.springframework.ai.openai.OpenAiChatModel; +import org.springframework.ai.openai.OpenAiEmbeddingModel; +import org.springframework.ai.openai.OpenAiImageModel; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionRequest.ResponseFormat; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionRequest.ToolChoiceBuilder; import org.springframework.ai.openai.api.OpenAiApi.FunctionTool.Type; @@ -561,7 +561,7 @@ public class OpenAiPropertiesTests { RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(OpenAiEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OpenAiEmbeddingClient.class)).isEmpty(); + assertThat(context.getBeansOfType(OpenAiEmbeddingModel.class)).isEmpty(); }); new ApplicationContextRunner() @@ -570,7 +570,7 @@ public class OpenAiPropertiesTests { RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(OpenAiEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OpenAiEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(OpenAiEmbeddingModel.class)).isNotEmpty(); }); new ApplicationContextRunner() @@ -580,7 +580,7 @@ public class OpenAiPropertiesTests { RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(OpenAiEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OpenAiEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(OpenAiEmbeddingModel.class)).isNotEmpty(); }); } @@ -593,7 +593,7 @@ public class OpenAiPropertiesTests { RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(OpenAiChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OpenAiChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(OpenAiChatModel.class)).isEmpty(); }); new ApplicationContextRunner() @@ -602,7 +602,7 @@ public class OpenAiPropertiesTests { RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(OpenAiChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OpenAiChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(OpenAiChatModel.class)).isNotEmpty(); }); new ApplicationContextRunner() @@ -612,7 +612,7 @@ public class OpenAiPropertiesTests { RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(OpenAiChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OpenAiChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(OpenAiChatModel.class)).isNotEmpty(); }); } @@ -626,7 +626,7 @@ public class OpenAiPropertiesTests { RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(OpenAiImageProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OpenAiImageClient.class)).isEmpty(); + assertThat(context.getBeansOfType(OpenAiImageModel.class)).isEmpty(); }); new ApplicationContextRunner() @@ -635,7 +635,7 @@ public class OpenAiPropertiesTests { RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(OpenAiImageProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OpenAiImageClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(OpenAiImageModel.class)).isNotEmpty(); }); new ApplicationContextRunner() @@ -645,7 +645,7 @@ public class OpenAiPropertiesTests { RestClientAutoConfiguration.class, OpenAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(OpenAiImageProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(OpenAiImageClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(OpenAiImageModel.class)).isNotEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackInPrompt2IT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackInPrompt2IT.java new file mode 100644 index 000000000..7d956c736 --- /dev/null +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackInPrompt2IT.java @@ -0,0 +1,120 @@ +/* + * Copyright 2023 - 2024 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.ai.autoconfigure.openai.tool; + +import java.util.function.Function; +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.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 { + + private final Logger logger = LoggerFactory.getLogger(FunctionCallbackInPromptIT.class); + + 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)); + + @Test + void functionCallTest() { + contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { + + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); + + ChatClient chatClient = ChatClient.builder(chatModel).build(); + + // @formatter:off + chatClient.prompt() + .user("Tell me a joke?") + .call().content(); + + String content = ChatClient.builder(chatModel).build().prompt() + .user("What's the weather like in San Francisco, Tokyo, and Paris?") + .function("CurrentWeatherService", "Get the weather in location", new MockWeatherService()) + .call().content(); + // @formatter:on + + logger.info("Response: {}", content); + + assertThat(content).containsAnyOf("30.0", "30"); + assertThat(content).containsAnyOf("10.0", "10"); + assertThat(content).containsAnyOf("15.0", "15"); + }); + } + + @Test + void functionCallTest2() { + contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { + + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); + + // @formatter:off + String content = ChatClient.builder(chatModel).build().prompt() + .user("What's the weather like in Amsterdam?") + .function("CurrentWeatherService", "Get the weather in location", + new Function() { + @Override + public String apply(MockWeatherService.Request request) { + return "18 degrees Celsius"; + } + }) + .call().content(); + // @formatter:on + logger.info("Response: {}", content); + + assertThat(content).contains("18"); + }); + } + + @Test + void streamingFunctionCallTest() { + + contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { + + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); + + // @formatter:off + String content = ChatClient.builder(chatModel).build().prompt() + .user("What's the weather like in San Francisco, Tokyo, and Paris?") + .function("CurrentWeatherService", "Get the weather in location", new MockWeatherService()) + .stream().content() + .collectList().block().stream().collect(Collectors.joining()); + // @formatter:on + + logger.info("Response: {}", content); + + assertThat(content).containsAnyOf("30.0", "30"); + assertThat(content).containsAnyOf("10.0", "10"); + assertThat(content).containsAnyOf("15.0", "15"); + }); + } + +} \ No newline at end of file diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackInPromptIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackInPromptIT.java index 67770f6af..5f14c9789 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackInPromptIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackInPromptIT.java @@ -32,7 +32,7 @@ import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.function.FunctionCallbackWrapper; -import org.springframework.ai.openai.OpenAiChatClient; +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; @@ -54,7 +54,7 @@ public class FunctionCallbackInPromptIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { - OpenAiChatClient chatClient = context.getBean(OpenAiChatClient.class); + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -66,7 +66,7 @@ public class FunctionCallbackInPromptIT { .build())) .build(); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), promptOptions)); + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), promptOptions)); logger.info("Response: {}", response); @@ -79,7 +79,7 @@ public class FunctionCallbackInPromptIT { contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { - OpenAiChatClient chatClient = context.getBean(OpenAiChatClient.class); + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -91,7 +91,7 @@ public class FunctionCallbackInPromptIT { .build())) .build(); - Flux response = chatClient.stream(new Prompt(List.of(userMessage), promptOptions)); + Flux response = chatModel.stream(new Prompt(List.of(userMessage), promptOptions)); String content = response.collectList() .block() diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWithPlainFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWithPlainFunctionBeanIT.java index d5437a62f..50a57b00f 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWithPlainFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWithPlainFunctionBeanIT.java @@ -27,15 +27,14 @@ import reactor.core.publisher.Flux; import org.springframework.ai.autoconfigure.openai.OpenAiAutoConfiguration; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; +import org.springframework.ai.chat.ChatClient; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.Generation; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.model.function.FunctionCallingOptions; -import org.springframework.ai.model.function.FunctionCallingOptionsBuilder.PortableFunctionCallingOptions; -import org.springframework.ai.openai.OpenAiChatClient; import org.springframework.ai.openai.OpenAiChatOptions; +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; @@ -60,12 +59,12 @@ class FunctionCallbackWithPlainFunctionBeanIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { - OpenAiChatClient chatClient = context.getBean(OpenAiChatClient.class); + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); // Test weatherFunction UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withFunction("weatherFunction").build())); logger.info("Response: {}", response); @@ -73,7 +72,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); // Test weatherFunctionTwo - response = chatClient.call(new Prompt(List.of(userMessage), + response = chatModel.call(new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withFunction("weatherFunctionTwo").build())); logger.info("Response: {}", response); @@ -87,18 +86,17 @@ class FunctionCallbackWithPlainFunctionBeanIT { void functionCallWithPortableFunctionCallingOptions() { contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { - OpenAiChatClient chatClient = context.getBean(OpenAiChatClient.class); + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); - // Test weatherFunction - UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); + // @formatter:off + String content = ChatClient.builder(chatModel).build().prompt() + .functions("weatherFunction") + .user("What's the weather like in San Francisco, Tokyo, and Paris?") + .stream().content() + .collectList().block().stream().collect(Collectors.joining()); + // @formatter:on - PortableFunctionCallingOptions functionOptions = FunctionCallingOptions.builder() - .withFunction("weatherFunction") - .build(); - - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), functionOptions)); - - logger.info("Response: {}", response); + logger.info("Response: {}", content); }); } @@ -106,12 +104,12 @@ class FunctionCallbackWithPlainFunctionBeanIT { void streamFunctionCallTest() { contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { - OpenAiChatClient chatClient = context.getBean(OpenAiChatClient.class); + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); // Test weatherFunction UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - Flux response = chatClient.stream(new Prompt(List.of(userMessage), + Flux response = chatModel.stream(new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withFunction("weatherFunction").build())); String content = response.collectList() @@ -129,7 +127,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { assertThat(content).containsAnyOf("15.0", "15"); // Test weatherFunctionTwo - response = chatClient.stream(new Prompt(List.of(userMessage), + response = chatModel.stream(new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withFunction("weatherFunctionTwo").build())); content = response.collectList() diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapper2IT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapper2IT.java new file mode 100644 index 000000000..d9d3cf01e --- /dev/null +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapper2IT.java @@ -0,0 +1,112 @@ +/* + * Copyright 2023 - 2024 the original author or authors. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * https://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.springframework.ai.autoconfigure.openai.tool; + +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.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 { + + private final Logger logger = LoggerFactory.getLogger(FunctionCallbackWrapperIT.class); + + 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)) + .withUserConfiguration(Config.class); + + @Test + void functionCallTest() { + contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { + + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); + + // @formatter:off + ChatClient chatClient = ChatClient.builder(chatModel) + .defaultFunctions("WeatherInfo") + .defaultUser(u -> u.text("What's the weather like in {cities}?")) + .build(); + + String content = chatClient.prompt() + .user(u -> u.param("cities", "San Francisco, Tokyo, Paris")) + .call().content(); + // @formatter:on + + logger.info("Response: {}", content); + + 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-preview").run(context -> { + + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); + + // @formatter:off + String content = ChatClient.builder(chatModel).build().prompt() + .functions("WeatherInfo") + .user("What's the weather like in San Francisco, Tokyo, and Paris?") + .stream().content() + .collectList().block().stream().collect(Collectors.joining()); + // @formatter:on + + logger.info("Response: {}", content); + + assertThat(content).containsAnyOf("30.0", "30"); + assertThat(content).containsAnyOf("10.0", "10"); + assertThat(content).containsAnyOf("15.0", "15"); + }); + } + + @Configuration + static class Config { + + @Bean + public FunctionCallback weatherFunctionInfo() { + + return FunctionCallbackWrapper.builder(new MockWeatherService()) + .withName("WeatherInfo") + .withDescription("Get the weather in location") + .withResponseConverter((response) -> "" + response.temp() + response.unit()) + .build(); + } + + } + +} \ No newline at end of file diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapperIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapperIT.java index 44e2c3af7..e4dcaf355 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapperIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/openai/tool/FunctionCallbackWrapperIT.java @@ -22,6 +22,7 @@ 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; @@ -33,7 +34,6 @@ import org.springframework.ai.chat.messages.UserMessage; 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.openai.OpenAiChatClient; import org.springframework.ai.openai.OpenAiChatOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; @@ -58,11 +58,11 @@ public class FunctionCallbackWrapperIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { - OpenAiChatClient chatClient = context.getBean(OpenAiChatClient.class); + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - ChatResponse response = chatClient.call( + ChatResponse response = chatModel.call( new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withFunction("WeatherInfo").build())); logger.info("Response: {}", response); @@ -76,11 +76,11 @@ public class FunctionCallbackWrapperIT { void streamFunctionCallTest() { contextRunner.withPropertyValues("spring.ai.openai.chat.options.model=gpt-4-turbo-preview").run(context -> { - OpenAiChatClient chatClient = context.getBean(OpenAiChatClient.class); + OpenAiChatModel chatModel = context.getBean(OpenAiChatModel.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - Flux response = chatClient.stream( + Flux response = chatModel.stream( new Prompt(List.of(userMessage), OpenAiChatOptions.builder().withFunction("WeatherInfo").build())); String content = response.collectList() diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlAutoConfigurationIT.java index c08bbde28..491db782c 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlAutoConfigurationIT.java @@ -28,7 +28,7 @@ import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.utility.DockerImageName; import org.springframework.ai.embedding.EmbeddingResponse; -import org.springframework.ai.postgresml.PostgresMlEmbeddingClient; +import org.springframework.ai.postgresml.PostgresMlEmbeddingModel; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.test.autoconfigure.jdbc.AutoConfigureTestDatabase; @@ -69,9 +69,9 @@ public class PostgresMlAutoConfigurationIT { .withBean(JdbcTemplate.class, () -> jdbcTemplate) .withConfiguration(AutoConfigurations.of(PostgresMlAutoConfiguration.class)); contextRunner.run(context -> { - PostgresMlEmbeddingClient embeddingClient = context.getBean(PostgresMlEmbeddingClient.class); + PostgresMlEmbeddingModel embeddingModel = context.getBean(PostgresMlEmbeddingModel.class); - EmbeddingResponse embeddingResponse = embeddingClient + EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); assertThat(embeddingResponse.getResults()).hasSize(2); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); @@ -79,7 +79,7 @@ public class PostgresMlAutoConfigurationIT { assertThat(embeddingResponse.getResults().get(1).getOutput()).isNotEmpty(); assertThat(embeddingResponse.getResults().get(1).getIndex()).isEqualTo(1); - assertThat(embeddingClient.dimensions()).isEqualTo(768); + assertThat(embeddingModel.dimensions()).isEqualTo(768); }); } @@ -90,7 +90,7 @@ public class PostgresMlAutoConfigurationIT { .withPropertyValues("spring.ai.postgresml.embedding.enabled=false") .run(context -> { assertThat(context.getBeansOfType(PostgresMlEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(PostgresMlEmbeddingClient.class)).isEmpty(); + assertThat(context.getBeansOfType(PostgresMlEmbeddingModel.class)).isEmpty(); }); new ApplicationContextRunner().withBean(JdbcTemplate.class, () -> jdbcTemplate) @@ -98,14 +98,14 @@ public class PostgresMlAutoConfigurationIT { .withPropertyValues("spring.ai.postgresml.embedding.enabled=true") .run(context -> { assertThat(context.getBeansOfType(PostgresMlEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(PostgresMlEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(PostgresMlEmbeddingModel.class)).isNotEmpty(); }); new ApplicationContextRunner().withBean(JdbcTemplate.class, () -> jdbcTemplate) .withConfiguration(AutoConfigurations.of(PostgresMlAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(PostgresMlEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(PostgresMlEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(PostgresMlEmbeddingModel.class)).isNotEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlEmbeddingPropertiesTests.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlEmbeddingPropertiesTests.java index 6f5a84d2f..0576e7b5b 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlEmbeddingPropertiesTests.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/postgresml/PostgresMlEmbeddingPropertiesTests.java @@ -20,7 +20,7 @@ import java.util.Map; import org.junit.jupiter.api.Test; import org.springframework.ai.document.MetadataMode; -import org.springframework.ai.postgresml.PostgresMlEmbeddingClient; +import org.springframework.ai.postgresml.PostgresMlEmbeddingModel; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.context.properties.EnableConfigurationProperties; @@ -48,7 +48,7 @@ class PostgresMlEmbeddingPropertiesTests { assertThat(this.postgresMlProperties).isNotNull(); assertThat(this.postgresMlProperties.getOptions().getTransformer()).isEqualTo("abc123"); assertThat(this.postgresMlProperties.getOptions().getVectorType()) - .isEqualTo(PostgresMlEmbeddingClient.VectorType.PG_ARRAY); + .isEqualTo(PostgresMlEmbeddingModel.VectorType.PG_ARRAY); assertThat(this.postgresMlProperties.getOptions().getKwargs()) .isEqualTo(Map.of("key1", "value1", "key2", "value2")); assertThat(this.postgresMlProperties.getOptions().getMetadataMode()).isEqualTo(MetadataMode.ALL); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiAutoConfigurationIT.java index 2597e7e20..0423170af 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiAutoConfigurationIT.java @@ -18,7 +18,7 @@ package org.springframework.ai.autoconfigure.stabilityai; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.image.Image; -import org.springframework.ai.image.ImageClient; +import org.springframework.ai.image.ImageModel; import org.springframework.ai.image.ImageGeneration; import org.springframework.ai.image.ImagePrompt; import org.springframework.ai.image.ImageResponse; @@ -39,7 +39,7 @@ public class StabilityAiAutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - ImageClient imageClient = context.getBean(ImageClient.class); + ImageModel imageModel = context.getBean(ImageModel.class); StabilityAiImageOptions imageOptions = StabilityAiImageOptions.builder() .withStylePreset(StyleEnum.PHOTOGRAPHIC) .build(); @@ -49,7 +49,7 @@ public class StabilityAiAutoConfigurationIT { """; ImagePrompt imagePrompt = new ImagePrompt(instructions, imageOptions); - ImageResponse imageResponse = imageClient.call(imagePrompt); + ImageResponse imageResponse = imageModel.call(imagePrompt); ImageGeneration imageGeneration = imageResponse.getResult(); Image image = imageGeneration.getOutput(); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImagePropertiesTests.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImagePropertiesTests.java index 63802bec5..6fbb85094 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImagePropertiesTests.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/stabilityai/StabilityAiImagePropertiesTests.java @@ -17,7 +17,7 @@ package org.springframework.ai.autoconfigure.stabilityai; import org.junit.jupiter.api.Test; -import org.springframework.ai.stabilityai.StabilityAiImageClient; +import org.springframework.ai.stabilityai.StabilityAiImageModel; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -80,7 +80,7 @@ public class StabilityAiImagePropertiesTests { .withConfiguration(AutoConfigurations.of(StabilityAiImageAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(StabilityAiImageProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(StabilityAiImageClient.class)).isEmpty(); + assertThat(context.getBeansOfType(StabilityAiImageModel.class)).isEmpty(); }); @@ -90,7 +90,7 @@ public class StabilityAiImagePropertiesTests { .withConfiguration(AutoConfigurations.of(StabilityAiImageAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(StabilityAiImageProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(StabilityAiImageClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(StabilityAiImageModel.class)).isNotEmpty(); }); @@ -100,7 +100,7 @@ public class StabilityAiImagePropertiesTests { .withConfiguration(AutoConfigurations.of(StabilityAiImageAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(StabilityAiImageProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(StabilityAiImageClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(StabilityAiImageModel.class)).isNotEmpty(); }); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingClientAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingModelAutoConfigurationIT.java similarity index 66% rename from spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingClientAutoConfigurationIT.java rename to spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingModelAutoConfigurationIT.java index 257becedd..2cb96123e 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingClientAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/transformers/TransformersEmbeddingModelAutoConfigurationIT.java @@ -21,8 +21,8 @@ import java.util.List; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.io.TempDir; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -31,29 +31,29 @@ import static org.assertj.core.api.Assertions.assertThat; /** * @author Christian Tzolov */ -public class TransformersEmbeddingClientAutoConfigurationIT { +public class TransformersEmbeddingModelAutoConfigurationIT { @TempDir File tempDir; private final ApplicationContextRunner contextRunner = new ApplicationContextRunner() - .withConfiguration(AutoConfigurations.of(TransformersEmbeddingClientAutoConfiguration.class)); + .withConfiguration(AutoConfigurations.of(TransformersEmbeddingModelAutoConfiguration.class)); @Test public void embedding() { contextRunner.run(context -> { - var properties = context.getBean(TransformersEmbeddingClientProperties.class); + var properties = context.getBean(TransformersEmbeddingModelProperties.class); assertThat(properties.getCache().isEnabled()).isTrue(); assertThat(properties.getCache().getDirectory()).isEqualTo( new File(System.getProperty("java.io.tmpdir"), "spring-ai-onnx-generative").getAbsolutePath()); - EmbeddingClient embeddingClient = context.getBean(EmbeddingClient.class); - assertThat(embeddingClient).isInstanceOf(TransformersEmbeddingClient.class); + EmbeddingModel embeddingModel = context.getBean(EmbeddingModel.class); + assertThat(embeddingModel).isInstanceOf(TransformersEmbeddingModel.class); - List> embeddings = embeddingClient.embed(List.of("Spring Framework", "Spring AI")); + List> embeddings = embeddingModel.embed(List.of("Spring Framework", "Spring AI")); assertThat(embeddings.size()).isEqualTo(2); // batch size - assertThat(embeddings.get(0).size()).isEqualTo(embeddingClient.dimensions()); // dimensions + assertThat(embeddings.get(0).size()).isEqualTo(embeddingModel.dimensions()); // dimensions // size }); } @@ -65,7 +65,7 @@ public class TransformersEmbeddingClientAutoConfigurationIT { "spring.ai.embedding.transformer.onnx.modelUri=https://huggingface.co/intfloat/e5-small-v2/resolve/main/model.onnx", "spring.ai.embedding.transformer.tokenizer.uri=https://huggingface.co/intfloat/e5-small-v2/raw/main/tokenizer.json") .run(context -> { - var properties = context.getBean(TransformersEmbeddingClientProperties.class); + var properties = context.getBean(TransformersEmbeddingModelProperties.class); assertThat(properties.getOnnx().getModelUri()) .isEqualTo("https://huggingface.co/intfloat/e5-small-v2/resolve/main/model.onnx"); assertThat(properties.getTokenizer().getUri()) @@ -75,15 +75,15 @@ public class TransformersEmbeddingClientAutoConfigurationIT { assertThat(properties.getCache().getDirectory()).isEqualTo(tempDir.getAbsolutePath()); assertThat(tempDir.listFiles()).hasSize(2); - EmbeddingClient embeddingClient = context.getBean(EmbeddingClient.class); - assertThat(embeddingClient).isInstanceOf(TransformersEmbeddingClient.class); + EmbeddingModel embeddingModel = context.getBean(EmbeddingModel.class); + assertThat(embeddingModel).isInstanceOf(TransformersEmbeddingModel.class); - assertThat(embeddingClient.dimensions()).isEqualTo(384); + assertThat(embeddingModel.dimensions()).isEqualTo(384); - List> embeddings = embeddingClient.embed(List.of("Spring Framework", "Spring AI")); + List> embeddings = embeddingModel.embed(List.of("Spring Framework", "Spring AI")); assertThat(embeddings.size()).isEqualTo(2); // batch size - assertThat(embeddings.get(0).size()).isEqualTo(embeddingClient.dimensions()); // dimensions + assertThat(embeddings.get(0).size()).isEqualTo(embeddingModel.dimensions()); // dimensions // size }); } @@ -91,18 +91,18 @@ public class TransformersEmbeddingClientAutoConfigurationIT { @Test void embeddingActivation() { contextRunner.withPropertyValues("spring.ai.embedding.transformer.enabled=false").run(context -> { - assertThat(context.getBeansOfType(TransformersEmbeddingClientProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(TransformersEmbeddingClient.class)).isEmpty(); + assertThat(context.getBeansOfType(TransformersEmbeddingModelProperties.class)).isNotEmpty(); + assertThat(context.getBeansOfType(TransformersEmbeddingModel.class)).isEmpty(); }); contextRunner.withPropertyValues("spring.ai.embedding.transformer.enabled=true").run(context -> { - assertThat(context.getBeansOfType(TransformersEmbeddingClientProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(TransformersEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(TransformersEmbeddingModelProperties.class)).isNotEmpty(); + assertThat(context.getBeansOfType(TransformersEmbeddingModel.class)).isNotEmpty(); }); contextRunner.run(context -> { - assertThat(context.getBeansOfType(TransformersEmbeddingClientProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(TransformersEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(TransformersEmbeddingModelProperties.class)).isNotEmpty(); + assertThat(context.getBeansOfType(TransformersEmbeddingModel.class)).isNotEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/azure/AzureVectorStoreAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/azure/AzureVectorStoreAutoConfigurationIT.java index 09d2b4127..6f44bed48 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/azure/AzureVectorStoreAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/azure/AzureVectorStoreAutoConfigurationIT.java @@ -28,8 +28,8 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.azure.AzureVectorStore; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; @@ -127,8 +127,8 @@ public class AzureVectorStoreAutoConfigurationIT { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStoreAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStoreAutoConfigurationIT.java index 119bde944..2d767a14c 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStoreAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/cassandra/CassandraVectorStoreAutoConfigurationIT.java @@ -26,8 +26,8 @@ import org.testcontainers.utility.DockerImageName; import org.springframework.ai.ResourceUtils; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -93,8 +93,8 @@ class CassandraVectorStoreAutoConfigurationIT { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfigurationIT.java index def4dd970..9caba58e1 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/chroma/ChromaVectorStoreAutoConfigurationIT.java @@ -24,8 +24,8 @@ import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -90,8 +90,8 @@ public class ChromaVectorStoreAutoConfigurationIT { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/milvus/MilvusVectorStoreAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/milvus/MilvusVectorStoreAutoConfigurationIT.java index a8500ed39..ac6c10e43 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/milvus/MilvusVectorStoreAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/milvus/MilvusVectorStoreAutoConfigurationIT.java @@ -30,8 +30,8 @@ import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.ResourceUtils; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -118,8 +118,8 @@ public class MilvusVectorStoreAutoConfigurationIT { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/neo4j/Neo4jVectorStoreAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/neo4j/Neo4jVectorStoreAutoConfigurationIT.java index fcd3281be..39a9812ad 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/neo4j/Neo4jVectorStoreAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/neo4j/Neo4jVectorStoreAutoConfigurationIT.java @@ -27,8 +27,8 @@ import org.testcontainers.utility.DockerImageName; import org.springframework.ai.ResourceUtils; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -97,8 +97,8 @@ public class Neo4jVectorStoreAutoConfigurationIT { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/pgvector/PgVectorStoreAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/pgvector/PgVectorStoreAutoConfigurationIT.java index b068685a1..d3cbf9575 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/pgvector/PgVectorStoreAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/pgvector/PgVectorStoreAutoConfigurationIT.java @@ -26,8 +26,8 @@ import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -104,8 +104,8 @@ public class PgVectorStoreAutoConfigurationIT { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/pinecone/PineconeVectorStoreAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/pinecone/PineconeVectorStoreAutoConfigurationIT.java index a137636ba..e7cf1c3d9 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/pinecone/PineconeVectorStoreAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/pinecone/PineconeVectorStoreAutoConfigurationIT.java @@ -28,8 +28,8 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -113,8 +113,8 @@ public class PineconeVectorStoreAutoConfigurationIT { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/qdrant/QdrantVectorStoreAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/qdrant/QdrantVectorStoreAutoConfigurationIT.java index 52dda7e96..26448bafd 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/qdrant/QdrantVectorStoreAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/qdrant/QdrantVectorStoreAutoConfigurationIT.java @@ -26,8 +26,8 @@ import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.qdrant.QdrantContainer; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -99,8 +99,8 @@ public class QdrantVectorStoreAutoConfigurationIT { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/qdrant/QdrantVectorStoreCloudAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/qdrant/QdrantVectorStoreCloudAutoConfigurationIT.java index 40261466e..6884cef10 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/qdrant/QdrantVectorStoreCloudAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/qdrant/QdrantVectorStoreCloudAutoConfigurationIT.java @@ -30,8 +30,8 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -137,8 +137,8 @@ public class QdrantVectorStoreCloudAutoConfigurationIT { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/redis/RedisVectorStoreAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/redis/RedisVectorStoreAutoConfigurationIT.java index c27725f9c..b5721db77 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/redis/RedisVectorStoreAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/redis/RedisVectorStoreAutoConfigurationIT.java @@ -23,8 +23,8 @@ import java.util.Map; import org.junit.jupiter.api.Test; import org.springframework.ai.ResourceUtils; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.boot.autoconfigure.AutoConfigurations; @@ -84,8 +84,8 @@ class RedisVectorStoreAutoConfigurationIT { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/weaviate/WeaviateVectorStoreAutoConfigurationTests.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/weaviate/WeaviateVectorStoreAutoConfigurationTests.java index 87b69c461..3946d75e6 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/weaviate/WeaviateVectorStoreAutoConfigurationTests.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vectorstore/weaviate/WeaviateVectorStoreAutoConfigurationTests.java @@ -24,8 +24,8 @@ import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig.MetadataField; @@ -122,8 +122,8 @@ public class WeaviateVectorStoreAutoConfigurationTests { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/VertexAiGeminiAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/VertexAiGeminiAutoConfigurationIT.java index af0e77de9..9ae96d6f7 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/VertexAiGeminiAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/VertexAiGeminiAutoConfigurationIT.java @@ -21,12 +21,12 @@ 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.vertexai.gemini.VertexAiGeminiChatModel; import reactor.core.publisher.Flux; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatClient; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -46,8 +46,8 @@ public class VertexAiGeminiAutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - VertexAiGeminiChatClient client = context.getBean(VertexAiGeminiChatClient.class); - String response = client.call("Hello"); + VertexAiGeminiChatModel chatModel = context.getBean(VertexAiGeminiChatModel.class); + String response = chatModel.call("Hello"); assertThat(response).isNotEmpty(); logger.info("Response: " + response); }); @@ -56,8 +56,8 @@ public class VertexAiGeminiAutoConfigurationIT { @Test void generateStreaming() { contextRunner.run(context -> { - VertexAiGeminiChatClient client = context.getBean(VertexAiGeminiChatClient.class); - Flux responseFlux = client.stream(new Prompt(new UserMessage("Hello"))); + VertexAiGeminiChatModel chatModel = context.getBean(VertexAiGeminiChatModel.class); + Flux responseFlux = chatModel.stream(new Prompt(new UserMessage("Hello"))); String response = responseFlux.collectList().block().stream().map(chatResponse -> { return chatResponse.getResults().get(0).getOutput().getContent(); }).collect(Collectors.joining()); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionBeanIT.java index 908445675..cfdaaea48 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionBeanIT.java @@ -28,7 +28,7 @@ import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.messages.SystemMessage; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatClient; +import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatModel; import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -55,12 +55,12 @@ class FunctionCallWithFunctionBeanIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.vertex.ai.gemini.chat.options.model=" - // + VertexAiGeminiChatClient.ChatModel.GEMINI_PRO.getValue()) - + VertexAiGeminiChatClient.ChatModel.GEMINI_PRO_1_5_PRO.getValue()) - // + VertexAiGeminiChatClient.ChatModel.GEMINI_PRO_1_5_FLASH.getValue()) + // + VertexAiGeminiChatModel.ChatModel.GEMINI_PRO.getValue()) + + VertexAiGeminiChatModel.ChatModel.GEMINI_PRO_1_5_PRO.getValue()) + // + VertexAiGeminiChatModel.ChatModel.GEMINI_PRO_1_5_FLASH.getValue()) .run(context -> { - VertexAiGeminiChatClient chatClient = context.getBean(VertexAiGeminiChatClient.class); + VertexAiGeminiChatModel chatModel = context.getBean(VertexAiGeminiChatModel.class); var systemMessage = new SystemMessage(""" Use Multi-turn function calling. @@ -72,9 +72,9 @@ class FunctionCallWithFunctionBeanIT { // Please let me know how many function calls you've preformed."); "What's the weather like in San Francisco, Paris and in Tokyo?"); - ChatResponse response = chatClient.call(new Prompt(List.of(systemMessage, userMessage), + ChatResponse response = chatModel.call(new Prompt(List.of(systemMessage, userMessage), VertexAiGeminiChatOptions.builder().withFunction("weatherFunction").build())); - // ChatResponse response = chatClient.call(new + // ChatResponse response = chatModel.call(new // Prompt(List.of(userMessage), // VertexAiGeminiChatOptions.builder().withFunction("weatherFunction").build())); @@ -84,14 +84,14 @@ class FunctionCallWithFunctionBeanIT { Thread.sleep(10000); - response = chatClient.call(new Prompt(List.of(systemMessage, userMessage), + response = chatModel.call(new Prompt(List.of(systemMessage, userMessage), VertexAiGeminiChatOptions.builder().withFunction("weatherFunction3").build())); logger.info("Response: {}", response); assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); - response = chatClient + response = chatModel .call(new Prompt(List.of(systemMessage, userMessage), VertexAiGeminiChatOptions.builder().build())); logger.info("Response: {}", response); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionWrapperIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionWrapperIT.java index cbad663c9..58a041b19 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionWrapperIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithFunctionWrapperIT.java @@ -30,7 +30,7 @@ import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackWrapper; import org.springframework.ai.model.function.FunctionCallbackWrapper.Builder.SchemaType; -import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatClient; +import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatModel; import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -55,10 +55,10 @@ public class FunctionCallWithFunctionWrapperIT { void functionCallTest() { contextRunner .withPropertyValues("spring.ai.vertex.ai.gemini.chat.options.model=" - + VertexAiGeminiChatClient.ChatModel.GEMINI_PRO.getValue()) + + VertexAiGeminiChatModel.ChatModel.GEMINI_PRO.getValue()) .run(context -> { - VertexAiGeminiChatClient chatClient = context.getBean(VertexAiGeminiChatClient.class); + VertexAiGeminiChatModel chatModel = context.getBean(VertexAiGeminiChatModel.class); var systemMessage = new SystemMessage(""" Use Multi-turn function calling. @@ -67,7 +67,7 @@ public class FunctionCallWithFunctionWrapperIT { """); var userMessage = new UserMessage("What's the weather like in San Francisco, Paris and in Tokyo?"); - ChatResponse response = chatClient.call(new Prompt(List.of(systemMessage, userMessage), + ChatResponse response = chatModel.call(new Prompt(List.of(systemMessage, userMessage), VertexAiGeminiChatOptions.builder().withFunction("WeatherInfo").build())); logger.info("Response: {}", response); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithPromptFunctionIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithPromptFunctionIT.java index b654fb124..dd03ee40d 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithPromptFunctionIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/gemini/tool/FunctionCallWithPromptFunctionIT.java @@ -29,7 +29,7 @@ import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.function.FunctionCallbackWrapper; import org.springframework.ai.model.function.FunctionCallbackWrapper.Builder.SchemaType; -import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatClient; +import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatModel; import org.springframework.ai.vertexai.gemini.VertexAiGeminiChatOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -51,10 +51,10 @@ public class FunctionCallWithPromptFunctionIT { void functionCallTest() { contextRunner .withPropertyValues("spring.ai.vertex.ai.gemini.chat.options.model=" - + VertexAiGeminiChatClient.ChatModel.GEMINI_PRO.getValue()) + + VertexAiGeminiChatModel.ChatModel.GEMINI_PRO.getValue()) .run(context -> { - VertexAiGeminiChatClient chatClient = context.getBean(VertexAiGeminiChatClient.class); + VertexAiGeminiChatModel chatModel = context.getBean(VertexAiGeminiChatModel.class); var systemMessage = new SystemMessage(""" Use Multi-turn function calling. @@ -72,14 +72,14 @@ public class FunctionCallWithPromptFunctionIT { .build())) .build(); - ChatResponse response = chatClient.call(new Prompt(List.of(systemMessage, userMessage), promptOptions)); + ChatResponse response = chatModel.call(new Prompt(List.of(systemMessage, userMessage), promptOptions)); logger.info("Response: {}", response); assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); // Verify that no function call is made. - response = chatClient + response = chatModel .call(new Prompt(List.of(systemMessage, userMessage), VertexAiGeminiChatOptions.builder().build())); logger.info("Response: {}", response); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPaLm2AutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPaLm2AutoConfigurationIT.java index 0ea8f663e..1f4b2f7ba 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPaLm2AutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/vertexai/palm2/VertexAiPaLm2AutoConfigurationIT.java @@ -23,8 +23,8 @@ 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.VertexAiPaLm2ChatClient; -import org.springframework.ai.vertexai.palm2.VertexAiPaLm2EmbeddingClient; +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; @@ -48,9 +48,9 @@ public class VertexAiPaLm2AutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - VertexAiPaLm2ChatClient client = context.getBean(VertexAiPaLm2ChatClient.class); + VertexAiPaLm2ChatModel chatModel = context.getBean(VertexAiPaLm2ChatModel.class); - String response = client.call("Hello"); + String response = chatModel.call("Hello"); assertThat(response).isNotEmpty(); logger.info("Response: " + response); @@ -60,9 +60,9 @@ public class VertexAiPaLm2AutoConfigurationIT { @Test void embedding() { contextRunner.run(context -> { - VertexAiPaLm2EmbeddingClient embeddingClient = context.getBean(VertexAiPaLm2EmbeddingClient.class); + VertexAiPaLm2EmbeddingModel embeddingModel = context.getBean(VertexAiPaLm2EmbeddingModel.class); - EmbeddingResponse embeddingResponse = embeddingClient + EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); assertThat(embeddingResponse.getResults()).hasSize(2); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); @@ -70,7 +70,7 @@ public class VertexAiPaLm2AutoConfigurationIT { assertThat(embeddingResponse.getResults().get(1).getOutput()).isNotEmpty(); assertThat(embeddingResponse.getResults().get(1).getIndex()).isEqualTo(1); - assertThat(embeddingClient.dimensions()).isEqualTo(768); + assertThat(embeddingModel.dimensions()).isEqualTo(768); }); } @@ -80,19 +80,19 @@ public class VertexAiPaLm2AutoConfigurationIT { // Disable the embedding auto-configuration. contextRunner.withPropertyValues("spring.ai.vertex.ai.embedding.enabled=false").run(context -> { assertThat(context.getBeansOfType(VertexAiPalm2EmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(VertexAiPaLm2EmbeddingClient.class)).isEmpty(); + assertThat(context.getBeansOfType(VertexAiPaLm2EmbeddingModel.class)).isEmpty(); }); // The embedding auto-configuration is enabled by default. contextRunner.run(context -> { assertThat(context.getBeansOfType(VertexAiPalm2EmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(VertexAiPaLm2EmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(VertexAiPaLm2EmbeddingModel.class)).isNotEmpty(); }); // Explicitly enable the embedding auto-configuration. contextRunner.withPropertyValues("spring.ai.vertex.ai.embedding.enabled=true").run(context -> { assertThat(context.getBeansOfType(VertexAiPalm2EmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(VertexAiPaLm2EmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(VertexAiPaLm2EmbeddingModel.class)).isNotEmpty(); }); } @@ -102,19 +102,19 @@ public class VertexAiPaLm2AutoConfigurationIT { // Disable the chat auto-configuration. contextRunner.withPropertyValues("spring.ai.vertex.ai.chat.enabled=false").run(context -> { assertThat(context.getBeansOfType(VertexAiPlam2ChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(VertexAiPaLm2ChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(VertexAiPaLm2ChatModel.class)).isEmpty(); }); // The chat auto-configuration is enabled by default. contextRunner.run(context -> { assertThat(context.getBeansOfType(VertexAiPlam2ChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(VertexAiPaLm2ChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(VertexAiPaLm2ChatModel.class)).isNotEmpty(); }); // Explicitly enable the chat auto-configuration. contextRunner.withPropertyValues("spring.ai.vertex.ai.chat.enabled=true").run(context -> { assertThat(context.getBeansOfType(VertexAiPlam2ChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(VertexAiPaLm2ChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(VertexAiPaLm2ChatModel.class)).isNotEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfigurationIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfigurationIT.java index 073c81ffc..9004412ee 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfigurationIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiAutoConfigurationIT.java @@ -26,9 +26,9 @@ 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.zhipuai.ZhiPuAiChatClient; -import org.springframework.ai.zhipuai.ZhiPuAiEmbeddingClient; -import org.springframework.ai.zhipuai.ZhiPuAiImageClient; +import org.springframework.ai.zhipuai.ZhiPuAiChatModel; +import org.springframework.ai.zhipuai.ZhiPuAiEmbeddingModel; +import org.springframework.ai.zhipuai.ZhiPuAiImageModel; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -55,8 +55,8 @@ public class ZhiPuAiAutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - ZhiPuAiChatClient client = context.getBean(ZhiPuAiChatClient.class); - String response = client.call("Hello"); + ZhiPuAiChatModel chatModel = context.getBean(ZhiPuAiChatModel.class); + String response = chatModel.call("Hello"); assertThat(response).isNotEmpty(); logger.info("Response: " + response); }); @@ -65,8 +65,8 @@ public class ZhiPuAiAutoConfigurationIT { @Test void generateStreaming() { contextRunner.run(context -> { - ZhiPuAiChatClient client = context.getBean(ZhiPuAiChatClient.class); - Flux responseFlux = client.stream(new Prompt(new UserMessage("Hello"))); + ZhiPuAiChatModel chatModel = context.getBean(ZhiPuAiChatModel.class); + Flux responseFlux = chatModel.stream(new Prompt(new UserMessage("Hello"))); String response = responseFlux.collectList().block().stream().map(chatResponse -> { return chatResponse.getResults().get(0).getOutput().getContent(); }).collect(Collectors.joining()); @@ -79,9 +79,9 @@ public class ZhiPuAiAutoConfigurationIT { @Test void embedding() { contextRunner.run(context -> { - ZhiPuAiEmbeddingClient embeddingClient = context.getBean(ZhiPuAiEmbeddingClient.class); + ZhiPuAiEmbeddingModel embeddingModel = context.getBean(ZhiPuAiEmbeddingModel.class); - EmbeddingResponse embeddingResponse = embeddingClient + EmbeddingResponse embeddingResponse = embeddingModel .embedForResponse(List.of("Hello World", "World is big and salvation is near")); assertThat(embeddingResponse.getResults()).hasSize(2); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); @@ -89,15 +89,15 @@ public class ZhiPuAiAutoConfigurationIT { assertThat(embeddingResponse.getResults().get(1).getOutput()).isNotEmpty(); assertThat(embeddingResponse.getResults().get(1).getIndex()).isEqualTo(1); - assertThat(embeddingClient.dimensions()).isEqualTo(1536); + assertThat(embeddingModel.dimensions()).isEqualTo(1536); }); } @Test void generateImage() { contextRunner.withPropertyValues("spring.ai.zhipuai.image.options.size=1024x1024").run(context -> { - ZhiPuAiImageClient client = context.getBean(ZhiPuAiImageClient.class); - ImageResponse imageResponse = client.call(new ImagePrompt("forest")); + ZhiPuAiImageModel ImageModel = context.getBean(ZhiPuAiImageModel.class); + ImageResponse imageResponse = ImageModel.call(new ImagePrompt("forest")); assertThat(imageResponse.getResults()).hasSize(1); assertThat(imageResponse.getResult().getOutput().getUrl()).isNotEmpty(); logger.info("Generated image: " + imageResponse.getResult().getOutput().getUrl()); diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiPropertiesTests.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiPropertiesTests.java index 121473dfd..154840460 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiPropertiesTests.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/ZhiPuAiPropertiesTests.java @@ -20,9 +20,9 @@ import org.skyscreamer.jsonassert.JSONAssert; import org.skyscreamer.jsonassert.JSONCompareMode; import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; import org.springframework.ai.model.ModelOptionsUtils; -import org.springframework.ai.zhipuai.ZhiPuAiChatClient; -import org.springframework.ai.zhipuai.ZhiPuAiEmbeddingClient; -import org.springframework.ai.zhipuai.ZhiPuAiImageClient; +import org.springframework.ai.zhipuai.ZhiPuAiChatModel; +import org.springframework.ai.zhipuai.ZhiPuAiEmbeddingModel; +import org.springframework.ai.zhipuai.ZhiPuAiImageModel; import org.springframework.ai.zhipuai.api.ZhiPuAiApi; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; @@ -346,7 +346,7 @@ public class ZhiPuAiPropertiesTests { RestClientAutoConfiguration.class, ZhiPuAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(ZhiPuAiEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(ZhiPuAiEmbeddingClient.class)).isEmpty(); + assertThat(context.getBeansOfType(ZhiPuAiEmbeddingModel.class)).isEmpty(); }); new ApplicationContextRunner() @@ -355,7 +355,7 @@ public class ZhiPuAiPropertiesTests { RestClientAutoConfiguration.class, ZhiPuAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(ZhiPuAiEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(ZhiPuAiEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(ZhiPuAiEmbeddingModel.class)).isNotEmpty(); }); new ApplicationContextRunner() @@ -365,7 +365,7 @@ public class ZhiPuAiPropertiesTests { RestClientAutoConfiguration.class, ZhiPuAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(ZhiPuAiEmbeddingProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(ZhiPuAiEmbeddingClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(ZhiPuAiEmbeddingModel.class)).isNotEmpty(); }); } @@ -378,7 +378,7 @@ public class ZhiPuAiPropertiesTests { RestClientAutoConfiguration.class, ZhiPuAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(ZhiPuAiChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(ZhiPuAiChatClient.class)).isEmpty(); + assertThat(context.getBeansOfType(ZhiPuAiChatModel.class)).isEmpty(); }); new ApplicationContextRunner() @@ -387,7 +387,7 @@ public class ZhiPuAiPropertiesTests { RestClientAutoConfiguration.class, ZhiPuAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(ZhiPuAiChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(ZhiPuAiChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(ZhiPuAiChatModel.class)).isNotEmpty(); }); new ApplicationContextRunner() @@ -397,7 +397,7 @@ public class ZhiPuAiPropertiesTests { RestClientAutoConfiguration.class, ZhiPuAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(ZhiPuAiChatProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(ZhiPuAiChatClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(ZhiPuAiChatModel.class)).isNotEmpty(); }); } @@ -411,7 +411,7 @@ public class ZhiPuAiPropertiesTests { RestClientAutoConfiguration.class, ZhiPuAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(ZhiPuAiImageProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(ZhiPuAiImageClient.class)).isEmpty(); + assertThat(context.getBeansOfType(ZhiPuAiImageModel.class)).isEmpty(); }); new ApplicationContextRunner() @@ -420,7 +420,7 @@ public class ZhiPuAiPropertiesTests { RestClientAutoConfiguration.class, ZhiPuAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(ZhiPuAiImageProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(ZhiPuAiImageClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(ZhiPuAiImageModel.class)).isNotEmpty(); }); new ApplicationContextRunner() @@ -430,7 +430,7 @@ public class ZhiPuAiPropertiesTests { RestClientAutoConfiguration.class, ZhiPuAiAutoConfiguration.class)) .run(context -> { assertThat(context.getBeansOfType(ZhiPuAiImageProperties.class)).isNotEmpty(); - assertThat(context.getBeansOfType(ZhiPuAiImageClient.class)).isNotEmpty(); + assertThat(context.getBeansOfType(ZhiPuAiImageModel.class)).isNotEmpty(); }); } diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackInPromptIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackInPromptIT.java index 29a71041f..c47471749 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackInPromptIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackInPromptIT.java @@ -27,7 +27,7 @@ import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.function.FunctionCallbackWrapper; -import org.springframework.ai.zhipuai.ZhiPuAiChatClient; +import org.springframework.ai.zhipuai.ZhiPuAiChatModel; import org.springframework.ai.zhipuai.ZhiPuAiChatOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; @@ -56,7 +56,7 @@ public class FunctionCallbackInPromptIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.zhipuai.chat.options.model=glm-4").run(context -> { - ZhiPuAiChatClient chatClient = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatModel chatModel = context.getBean(ZhiPuAiChatModel.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -68,7 +68,7 @@ public class FunctionCallbackInPromptIT { .build())) .build(); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), promptOptions)); + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), promptOptions)); logger.info("Response: {}", response); @@ -81,7 +81,7 @@ public class FunctionCallbackInPromptIT { contextRunner.withPropertyValues("spring.ai.zhipuai.chat.options.model=glm-4").run(context -> { - ZhiPuAiChatClient chatClient = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatModel chatModel = context.getBean(ZhiPuAiChatModel.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -93,7 +93,7 @@ public class FunctionCallbackInPromptIT { .build())) .build(); - Flux response = chatClient.stream(new Prompt(List.of(userMessage), promptOptions)); + Flux response = chatModel.stream(new Prompt(List.of(userMessage), promptOptions)); String content = response.collectList() .block() diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWithPlainFunctionBeanIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWithPlainFunctionBeanIT.java index 2d7c121d1..1d76347e1 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWithPlainFunctionBeanIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWithPlainFunctionBeanIT.java @@ -28,7 +28,7 @@ import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.function.FunctionCallingOptions; import org.springframework.ai.model.function.FunctionCallingOptionsBuilder.PortableFunctionCallingOptions; -import org.springframework.ai.zhipuai.ZhiPuAiChatClient; +import org.springframework.ai.zhipuai.ZhiPuAiChatModel; import org.springframework.ai.zhipuai.ZhiPuAiChatOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; @@ -62,12 +62,12 @@ class FunctionCallbackWithPlainFunctionBeanIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.zhipuai.chat.options.model=glm-4").run(context -> { - ZhiPuAiChatClient chatClient = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatModel chatModel = context.getBean(ZhiPuAiChatModel.class); // Test weatherFunction UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), ZhiPuAiChatOptions.builder().withFunction("weatherFunction").build())); logger.info("Response: {}", response); @@ -75,7 +75,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { assertThat(response.getResult().getOutput().getContent()).contains("30", "10", "15"); // Test weatherFunctionTwo - response = chatClient.call(new Prompt(List.of(userMessage), + response = chatModel.call(new Prompt(List.of(userMessage), ZhiPuAiChatOptions.builder().withFunction("weatherFunctionTwo").build())); logger.info("Response: {}", response); @@ -89,7 +89,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { void functionCallWithPortableFunctionCallingOptions() { contextRunner.withPropertyValues("spring.ai.zhipuai.chat.options.model=glm-4").run(context -> { - ZhiPuAiChatClient chatClient = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatModel chatModel = context.getBean(ZhiPuAiChatModel.class); // Test weatherFunction UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); @@ -98,7 +98,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { .withFunction("weatherFunction") .build(); - ChatResponse response = chatClient.call(new Prompt(List.of(userMessage), functionOptions)); + ChatResponse response = chatModel.call(new Prompt(List.of(userMessage), functionOptions)); logger.info("Response: {}", response); }); @@ -108,12 +108,12 @@ class FunctionCallbackWithPlainFunctionBeanIT { void streamFunctionCallTest() { contextRunner.withPropertyValues("spring.ai.zhipuai.chat.options.model=glm-4").run(context -> { - ZhiPuAiChatClient chatClient = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatModel chatModel = context.getBean(ZhiPuAiChatModel.class); // Test weatherFunction UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - Flux response = chatClient.stream(new Prompt(List.of(userMessage), + Flux response = chatModel.stream(new Prompt(List.of(userMessage), ZhiPuAiChatOptions.builder().withFunction("weatherFunction").build())); String content = response.collectList() @@ -131,7 +131,7 @@ class FunctionCallbackWithPlainFunctionBeanIT { assertThat(content).containsAnyOf("15.0", "15"); // Test weatherFunctionTwo - response = chatClient.stream(new Prompt(List.of(userMessage), + response = chatModel.stream(new Prompt(List.of(userMessage), ZhiPuAiChatOptions.builder().withFunction("weatherFunctionTwo").build())); content = response.collectList() diff --git a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWrapperIT.java b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWrapperIT.java index fefd81326..5cb4cefed 100644 --- a/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWrapperIT.java +++ b/spring-ai-spring-boot-autoconfigure/src/test/java/org/springframework/ai/autoconfigure/zhipuai/tool/FunctionCallbackWrapperIT.java @@ -28,7 +28,7 @@ import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.model.function.FunctionCallback; import org.springframework.ai.model.function.FunctionCallbackWrapper; -import org.springframework.ai.zhipuai.ZhiPuAiChatClient; +import org.springframework.ai.zhipuai.ZhiPuAiChatModel; import org.springframework.ai.zhipuai.ZhiPuAiChatOptions; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; @@ -60,11 +60,11 @@ public class FunctionCallbackWrapperIT { void functionCallTest() { contextRunner.withPropertyValues("spring.ai.zhipuai.chat.options.model=glm-4").run(context -> { - ZhiPuAiChatClient chatClient = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatModel chatModel = context.getBean(ZhiPuAiChatModel.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - ChatResponse response = chatClient.call( + ChatResponse response = chatModel.call( new Prompt(List.of(userMessage), ZhiPuAiChatOptions.builder().withFunction("WeatherInfo").build())); logger.info("Response: {}", response); @@ -78,11 +78,11 @@ public class FunctionCallbackWrapperIT { void streamFunctionCallTest() { contextRunner.withPropertyValues("spring.ai.zhipuai.chat.options.model=glm-4").run(context -> { - ZhiPuAiChatClient chatClient = context.getBean(ZhiPuAiChatClient.class); + ZhiPuAiChatModel chatModel = context.getBean(ZhiPuAiChatModel.class); UserMessage userMessage = new UserMessage("What's the weather like in San Francisco, Tokyo, and Paris?"); - Flux response = chatClient.stream( + Flux response = chatModel.stream( new Prompt(List.of(userMessage), ZhiPuAiChatOptions.builder().withFunction("WeatherInfo").build())); String content = response.collectList() diff --git a/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/chroma/ChromaContainerConnectionDetailsFactory.java b/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/chroma/ChromaContainerConnectionDetailsFactory.java index cb859f400..dffe9603f 100644 --- a/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/chroma/ChromaContainerConnectionDetailsFactory.java +++ b/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/chroma/ChromaContainerConnectionDetailsFactory.java @@ -23,7 +23,7 @@ import org.testcontainers.chromadb.ChromaDBContainer; /** * @author Eddú Meléndez */ -class ChromaContainerConnectionDetailsFactory +public class ChromaContainerConnectionDetailsFactory extends ContainerConnectionDetailsFactory { @Override diff --git a/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/milvus/MilvusContainerConnectionDetailsFactory.java b/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/milvus/MilvusContainerConnectionDetailsFactory.java index a44643c38..8d4f2d7cd 100644 --- a/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/milvus/MilvusContainerConnectionDetailsFactory.java +++ b/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/milvus/MilvusContainerConnectionDetailsFactory.java @@ -23,7 +23,7 @@ import org.testcontainers.milvus.MilvusContainer; /** * @author Eddú Meléndez */ -class MilvusContainerConnectionDetailsFactory +public class MilvusContainerConnectionDetailsFactory extends ContainerConnectionDetailsFactory { @Override diff --git a/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/ollama/OllamaContainerConnectionDetailsFactory.java b/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/ollama/OllamaContainerConnectionDetailsFactory.java index 46174bc36..46a410f86 100644 --- a/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/ollama/OllamaContainerConnectionDetailsFactory.java +++ b/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/ollama/OllamaContainerConnectionDetailsFactory.java @@ -23,7 +23,7 @@ import org.testcontainers.ollama.OllamaContainer; /** * @author Eddú Meléndez */ -class OllamaContainerConnectionDetailsFactory +public class OllamaContainerConnectionDetailsFactory extends ContainerConnectionDetailsFactory { @Override diff --git a/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/qdrant/QdrantContainerConnectionDetailsFactory.java b/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/qdrant/QdrantContainerConnectionDetailsFactory.java index 2b410c516..615610bca 100644 --- a/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/qdrant/QdrantContainerConnectionDetailsFactory.java +++ b/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/qdrant/QdrantContainerConnectionDetailsFactory.java @@ -23,7 +23,7 @@ import org.testcontainers.qdrant.QdrantContainer; /** * @author Eddú Meléndez */ -class QdrantContainerConnectionDetailsFactory +public class QdrantContainerConnectionDetailsFactory extends ContainerConnectionDetailsFactory { @Override diff --git a/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/redis/RedisContainerConnectionDetailsFactory.java b/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/redis/RedisContainerConnectionDetailsFactory.java index 72007c4e6..d1e6d7e66 100644 --- a/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/redis/RedisContainerConnectionDetailsFactory.java +++ b/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/redis/RedisContainerConnectionDetailsFactory.java @@ -23,7 +23,7 @@ import org.springframework.boot.testcontainers.service.connection.ContainerConne /** * @author Eddú Meléndez */ -class RedisContainerConnectionDetailsFactory +public class RedisContainerConnectionDetailsFactory extends ContainerConnectionDetailsFactory { @Override diff --git a/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/weaviate/WeaviateContainerConnectionDetailsFactory.java b/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/weaviate/WeaviateContainerConnectionDetailsFactory.java index 601fb6244..ff427ee37 100644 --- a/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/weaviate/WeaviateContainerConnectionDetailsFactory.java +++ b/spring-ai-spring-boot-testcontainers/src/main/java/org/springframework/ai/testcontainers/service/connection/weaviate/WeaviateContainerConnectionDetailsFactory.java @@ -23,7 +23,7 @@ import org.testcontainers.weaviate.WeaviateContainer; /** * @author Eddú Meléndez */ -class WeaviateContainerConnectionDetailsFactory +public class WeaviateContainerConnectionDetailsFactory extends ContainerConnectionDetailsFactory { @Override diff --git a/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/chroma/ChromaContainerConnectionDetailsFactoryTest.java b/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/chroma/ChromaContainerConnectionDetailsFactoryTest.java index 1a9db7bc8..2977ef7bd 100644 --- a/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/chroma/ChromaContainerConnectionDetailsFactoryTest.java +++ b/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/chroma/ChromaContainerConnectionDetailsFactoryTest.java @@ -18,8 +18,8 @@ package org.springframework.ai.testcontainers.service.connection.chroma; import org.junit.jupiter.api.Test; import org.springframework.ai.autoconfigure.vectorstore.chroma.ChromaVectorStoreAutoConfiguration; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.beans.factory.annotation.Autowired; @@ -83,8 +83,8 @@ class ChromaContainerConnectionDetailsFactoryTest { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/milvus/MilvusContainerConnectionDetailsFactoryTest.java b/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/milvus/MilvusContainerConnectionDetailsFactoryTest.java index ee43cf001..8320dad60 100644 --- a/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/milvus/MilvusContainerConnectionDetailsFactoryTest.java +++ b/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/milvus/MilvusContainerConnectionDetailsFactoryTest.java @@ -19,8 +19,8 @@ import org.junit.jupiter.api.Test; import org.springframework.ai.ResourceUtils; import org.springframework.ai.autoconfigure.vectorstore.milvus.MilvusVectorStoreAutoConfiguration; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.beans.factory.annotation.Autowired; @@ -90,8 +90,8 @@ class MilvusContainerConnectionDetailsFactoryTest { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/ollama/OllamaContainerConnectionDetailsFactoryTest.java b/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/ollama/OllamaContainerConnectionDetailsFactoryTest.java index aa470135a..10c45c24b 100644 --- a/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/ollama/OllamaContainerConnectionDetailsFactoryTest.java +++ b/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/ollama/OllamaContainerConnectionDetailsFactoryTest.java @@ -21,7 +21,7 @@ import org.junit.jupiter.api.BeforeAll; import org.junit.jupiter.api.Test; import org.springframework.ai.autoconfigure.ollama.OllamaAutoConfiguration; import org.springframework.ai.embedding.EmbeddingResponse; -import org.springframework.ai.ollama.OllamaEmbeddingClient; +import org.springframework.ai.ollama.OllamaEmbeddingModel; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.boot.autoconfigure.ImportAutoConfiguration; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; @@ -56,7 +56,7 @@ class OllamaContainerConnectionDetailsFactoryTest { static OllamaContainer ollama = new OllamaContainer("ollama/ollama:0.1.29"); @Autowired - private OllamaEmbeddingClient embeddingClient; + private OllamaEmbeddingModel embeddingModel; @BeforeAll public static void beforeAll() throws IOException, InterruptedException { @@ -67,10 +67,10 @@ class OllamaContainerConnectionDetailsFactoryTest { @Test public void singleTextEmbedding() { - EmbeddingResponse embeddingResponse = this.embeddingClient.embedForResponse(List.of("Hello World")); + EmbeddingResponse embeddingResponse = this.embeddingModel.embedForResponse(List.of("Hello World")); assertThat(embeddingResponse.getResults()).hasSize(1); assertThat(embeddingResponse.getResults().get(0).getOutput()).isNotEmpty(); - assertThat(this.embeddingClient.dimensions()).isEqualTo(3200); + assertThat(this.embeddingModel.dimensions()).isEqualTo(3200); } @Configuration(proxyBeanMethods = false) diff --git a/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/qdrant/QdrantContainerConnectionDetailsFactoryTest.java b/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/qdrant/QdrantContainerConnectionDetailsFactoryTest.java index c1c461a3e..c0cb3f036 100644 --- a/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/qdrant/QdrantContainerConnectionDetailsFactoryTest.java +++ b/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/qdrant/QdrantContainerConnectionDetailsFactoryTest.java @@ -18,8 +18,8 @@ package org.springframework.ai.testcontainers.service.connection.qdrant; import org.junit.jupiter.api.Test; import org.springframework.ai.autoconfigure.vectorstore.qdrant.QdrantVectorStoreAutoConfiguration; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.beans.factory.annotation.Autowired; @@ -91,8 +91,8 @@ public class QdrantContainerConnectionDetailsFactoryTest { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/redis/RedisContainerConnectionDetailsFactoryTest.java b/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/redis/RedisContainerConnectionDetailsFactoryTest.java index b0888689e..aa53e64fd 100644 --- a/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/redis/RedisContainerConnectionDetailsFactoryTest.java +++ b/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/redis/RedisContainerConnectionDetailsFactoryTest.java @@ -20,8 +20,8 @@ import org.junit.jupiter.api.Test; import org.springframework.ai.ResourceUtils; import org.springframework.ai.autoconfigure.vectorstore.redis.RedisVectorStoreAutoConfiguration; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.beans.factory.annotation.Autowired; @@ -82,8 +82,8 @@ class RedisContainerConnectionDetailsFactoryTest { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/weaviate/WeaviateContainerConnectionDetailsFactoryTest.java b/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/weaviate/WeaviateContainerConnectionDetailsFactoryTest.java index 685bc1b79..3ab141406 100644 --- a/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/weaviate/WeaviateContainerConnectionDetailsFactoryTest.java +++ b/spring-ai-spring-boot-testcontainers/src/test/java/org/springframework/ai/testcontainers/service/connection/weaviate/WeaviateContainerConnectionDetailsFactoryTest.java @@ -19,8 +19,8 @@ import org.junit.jupiter.api.Test; import org.springframework.ai.autoconfigure.vectorstore.weaviate.WeaviateVectorStoreAutoConfiguration; import org.springframework.ai.autoconfigure.vectorstore.weaviate.WeaviateVectorStoreProperties; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.ai.vectorstore.WeaviateVectorStore; @@ -116,8 +116,8 @@ class WeaviateContainerConnectionDetailsFactoryTest { static class Config { @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/spring-ai-test/src/main/java/org/springframework/ai/evaluation/BaseMemoryTest.java b/spring-ai-test/src/main/java/org/springframework/ai/evaluation/BaseMemoryTest.java index 5fa667481..a5d65866e 100644 --- a/spring-ai-test/src/main/java/org/springframework/ai/evaluation/BaseMemoryTest.java +++ b/spring-ai-test/src/main/java/org/springframework/ai/evaluation/BaseMemoryTest.java @@ -44,10 +44,10 @@ public class BaseMemoryTest { protected StreamingChatService streamingChatService; public BaseMemoryTest(RelevancyEvaluator relevancyEvaluator, ChatService chatService, - StreamingChatService streamingChatClient) { + StreamingChatService streamingChatModel) { this.relevancyEvaluator = relevancyEvaluator; this.chatService = chatService; - this.streamingChatService = streamingChatClient; + this.streamingChatService = streamingChatModel; } @Test diff --git a/spring-ai-test/src/main/java/org/springframework/ai/evaluation/BasicEvaluationTest.java b/spring-ai-test/src/main/java/org/springframework/ai/evaluation/BasicEvaluationTest.java index bdc0750bd..231d51631 100644 --- a/spring-ai-test/src/main/java/org/springframework/ai/evaluation/BasicEvaluationTest.java +++ b/spring-ai-test/src/main/java/org/springframework/ai/evaluation/BasicEvaluationTest.java @@ -17,7 +17,7 @@ package org.springframework.ai.evaluation; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.ChatResponse; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.chat.prompt.PromptTemplate; @@ -38,7 +38,7 @@ public class BasicEvaluationTest { private static final Logger logger = LoggerFactory.getLogger(BasicEvaluationTest.class); @Autowired - protected ChatClient openAiChatClient; + protected ChatModel openAiChatModel; @Value("classpath:/prompts/spring/test/evaluation/qa-evaluator-accurate-answer.st") protected Resource qaEvaluatorAccurateAnswerResource; @@ -68,12 +68,12 @@ public class BasicEvaluationTest { } Message userMessage = userPromptTemplate.createMessage(); Prompt prompt = new Prompt(List.of(userMessage, systemMessage)); - String yesOrNo = openAiChatClient.call(prompt).getResult().getOutput().getContent(); + String yesOrNo = openAiChatModel.call(prompt).getResult().getOutput().getContent(); logger.info("Is Answer related to question: " + yesOrNo); if (yesOrNo.equalsIgnoreCase("no")) { SystemMessage notRelatedSystemMessage = new SystemMessage(qaEvaluatorNotRelatedResource); prompt = new Prompt(List.of(userMessage, notRelatedSystemMessage)); - String reasonForFailure = openAiChatClient.call(prompt).getResult().getOutput().getContent(); + String reasonForFailure = openAiChatModel.call(prompt).getResult().getOutput().getContent(); fail(reasonForFailure); } else { diff --git a/spring-ai-test/src/main/resources/test/data/spring.ai.txt b/spring-ai-test/src/main/resources/test/data/spring.ai.txt index f41b20ec3..3fd3513cb 100644 --- a/spring-ai-test/src/main/resources/test/data/spring.ai.txt +++ b/spring-ai-test/src/main/resources/test/data/spring.ai.txt @@ -1,6 +1,6 @@ The Spring AI project aims to streamline the development of applications that incorporate artificial intelligence functionality without unnecessary complexity. The project draws inspiration from notable Python projects, such as LangChain and LlamaIndex, but Spring AI is not a direct port of those projects. The project was founded with the belief that the next wave of Generative AI applications will not be only for Python developers but will be ubiquitous across many programming languages. -At its core, Spring AI provides abstractions that serve as the foundation for developing AI applications. These abstractions have multiple implementations, enabling easy component swapping with minimal code changes. For example, Spring AI introduces the ChatClient interface with implementations for OpenAI and Azure OpenAI. +At its core, Spring AI provides abstractions that serve as the foundation for developing AI applications. These abstractions have multiple implementations, enabling easy component swapping with minimal code changes. For example, Spring AI introduces the ChatModel interface with implementations for OpenAI and Azure OpenAI. In addition to these core abstractions, Spring AI aims to provide higher-level functionalities to address common use cases such as “Q&A over your documentation” or “Chat with your documentation.” As the complexity of the use cases increases, the Spring AI project will integrate with other projects in the Spring Ecosystem, such as Spring Integration, Spring Batch, and Spring Data. To simplify setup, Spring Boot starters are available to help set up essential dependencies and classes. There is also a collection of sample applications to help you explore the project’s features. Lastly, the new Spring CLI project also enables you to get started quickly by using the spring boot new AI command for new projects or spring boot add AI for adding AI capabilities to your existing application. diff --git a/vector-stores/spring-ai-azure-store/src/main/java/org/springframework/ai/vectorstore/azure/AzureVectorStore.java b/vector-stores/spring-ai-azure-store/src/main/java/org/springframework/ai/vectorstore/azure/AzureVectorStore.java index e0d8d9188..547bcb8eb 100644 --- a/vector-stores/spring-ai-azure-store/src/main/java/org/springframework/ai/vectorstore/azure/AzureVectorStore.java +++ b/vector-stores/spring-ai-azure-store/src/main/java/org/springframework/ai/vectorstore/azure/AzureVectorStore.java @@ -45,7 +45,7 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.ai.vectorstore.filter.FilterExpressionConverter; @@ -92,7 +92,7 @@ public class AzureVectorStore implements VectorStore, InitializingBean { private final SearchIndexClient searchIndexClient; - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; private SearchClient searchClient; @@ -146,29 +146,29 @@ public class AzureVectorStore implements VectorStore, InitializingBean { * Constructs a new AzureCognitiveSearchVectorStore. * @param searchIndexClient A pre-configured Azure {@link SearchIndexClient} that CRUD * for Azure search indexes and factory for {@link SearchClient}. - * @param embeddingClient The client for embedding operations. + * @param embeddingModel The client for embedding operations. */ - public AzureVectorStore(SearchIndexClient searchIndexClient, EmbeddingClient embeddingClient) { - this(searchIndexClient, embeddingClient, List.of()); + public AzureVectorStore(SearchIndexClient searchIndexClient, EmbeddingModel embeddingModel) { + this(searchIndexClient, embeddingModel, List.of()); } /** * Constructs a new AzureCognitiveSearchVectorStore. * @param searchIndexClient A pre-configured Azure {@link SearchIndexClient} that CRUD * for Azure search indexes and factory for {@link SearchClient}. - * @param embeddingClient The client for embedding operations. + * @param embeddingModel The client for embedding operations. * @param filterMetadataFields List of metadata fields (as field name and type) that * can be used in similarity search query filter expressions. */ - public AzureVectorStore(SearchIndexClient searchIndexClient, EmbeddingClient embeddingClient, + public AzureVectorStore(SearchIndexClient searchIndexClient, EmbeddingModel embeddingModel, List filterMetadataFields) { - Assert.notNull(embeddingClient, "The embedding client can not be null."); + Assert.notNull(embeddingModel, "The embedding model can not be null."); Assert.notNull(searchIndexClient, "The search index client can not be null."); Assert.notNull(filterMetadataFields, "The filterMetadataFields can not be null."); this.searchIndexClient = searchIndexClient; - this.embeddingClient = embeddingClient; + this.embeddingModel = embeddingModel; this.filterMetadataFields = filterMetadataFields; this.filterExpressionConverter = new AzureAiSearchFilterExpressionConverter(filterMetadataFields); } @@ -211,7 +211,7 @@ public class AzureVectorStore implements VectorStore, InitializingBean { } final var searchDocuments = documents.stream().map(document -> { - final var embeddings = this.embeddingClient.embed(document); + final var embeddings = this.embeddingModel.embed(document); SearchDocument searchDocument = new SearchDocument(); searchDocument.put(ID_FIELD_NAME, document.getId()); searchDocument.put(EMBEDDING_FIELD_NAME, embeddings); @@ -277,7 +277,7 @@ public class AzureVectorStore implements VectorStore, InitializingBean { Assert.notNull(request, "The search request must not be null."); - var searchEmbedding = toFloatList(embeddingClient.embed(request.getQuery())); + var searchEmbedding = toFloatList(embeddingModel.embed(request.getQuery())); final var vectorQuery = new VectorizedQuery(searchEmbedding).setKNearestNeighborsCount(request.getTopK()) // Set the fields to compare the vector against. This is a comma-delimited @@ -328,7 +328,7 @@ public class AzureVectorStore implements VectorStore, InitializingBean { @Override public void afterPropertiesSet() throws Exception { - int dimensions = this.embeddingClient.dimensions(); + int dimensions = this.embeddingModel.dimensions(); List fields = new ArrayList<>(); diff --git a/vector-stores/spring-ai-azure-store/src/test/java/org/springframework/ai/vectorstore/azure/AzureVectorStoreIT.java b/vector-stores/spring-ai-azure-store/src/test/java/org/springframework/ai/vectorstore/azure/AzureVectorStoreIT.java index 7e4cbbbd7..9c09b5879 100644 --- a/vector-stores/spring-ai-azure-store/src/test/java/org/springframework/ai/vectorstore/azure/AzureVectorStoreIT.java +++ b/vector-stores/spring-ai-azure-store/src/test/java/org/springframework/ai/vectorstore/azure/AzureVectorStoreIT.java @@ -34,8 +34,8 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.ai.vectorstore.azure.AzureVectorStore.MetadataField; @@ -302,15 +302,15 @@ public class AzureVectorStoreIT { } @Bean - public VectorStore vectorStore(SearchIndexClient searchIndexClient, EmbeddingClient embeddingClient) { + public VectorStore vectorStore(SearchIndexClient searchIndexClient, EmbeddingModel embeddingModel) { var filterableMetaFields = List.of(MetadataField.text("country"), MetadataField.int64("year"), MetadataField.date("activationDate")); - return new AzureVectorStore(searchIndexClient, embeddingClient, filterableMetaFields); + return new AzureVectorStore(searchIndexClient, embeddingModel, filterableMetaFields); } @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/vector-stores/spring-ai-cassandra-store/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStore.java b/vector-stores/spring-ai-cassandra-store/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStore.java index c08f7d09d..3cc2fa220 100644 --- a/vector-stores/spring-ai-cassandra-store/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStore.java +++ b/vector-stores/spring-ai-cassandra-store/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStore.java @@ -31,7 +31,7 @@ import com.datastax.oss.driver.shaded.guava.common.base.Preconditions; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.CassandraVectorStoreConfig.SchemaColumn; import org.springframework.ai.vectorstore.filter.FilterExpressionConverter; import org.springframework.beans.factory.InitializingBean; @@ -57,7 +57,7 @@ import java.util.concurrent.ConcurrentMap; * * This class requires a CassandraVectorStoreConfig configuration object for * initialization, which includes settings like connection details, index name, column - * names, etc. It also requires an EmbeddingClient to convert documents into embeddings + * names, etc. It also requires an EmbeddingModel to convert documents into embeddings * before storing them. * * A schema matching the configuration is automatically created if it doesn't exist. @@ -73,7 +73,7 @@ import java.util.concurrent.ConcurrentMap; * change the schema server-side you need a new CassandraVectorStore instance. * * When adding documents with the method {@link #add(List)} it first calls - * embeddingClient to create the embeddings. This is slow. Configure + * embeddingModel to create the embeddings. This is slow. Configure * {@link CassandraVectorStoreConfig.Builder#withFixedThreadPoolExecutorSize(int)} * accordingly to improve performance so embeddings are created and the documents are * added concurrently. The default concurrency is 16 @@ -86,7 +86,7 @@ import java.util.concurrent.ConcurrentMap; * @author Mick Semb Wever * @see VectorStore * @see org.springframework.ai.vectorstore.CassandraVectorStoreConfig - * @see EmbeddingClient + * @see EmbeddingModel * @since 1.0.0 */ public class CassandraVectorStore implements VectorStore, InitializingBean, AutoCloseable { @@ -113,7 +113,7 @@ public class CassandraVectorStore implements VectorStore, InitializingBean, Auto private final CassandraVectorStoreConfig conf; - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; private final FilterExpressionConverter filterExpressionConverter; @@ -125,14 +125,14 @@ public class CassandraVectorStore implements VectorStore, InitializingBean, Auto private final Similarity similarity; - public CassandraVectorStore(CassandraVectorStoreConfig conf, EmbeddingClient embeddingClient) { + public CassandraVectorStore(CassandraVectorStoreConfig conf, EmbeddingModel embeddingModel) { Preconditions.checkArgument(null != conf, "Config must not be null"); - Preconditions.checkArgument(null != embeddingClient, "Embedding client must not be null"); + Preconditions.checkArgument(null != embeddingModel, "Embedding client must not be null"); this.conf = conf; - this.embeddingClient = embeddingClient; - conf.ensureSchemaExists(embeddingClient.dimensions()); + this.embeddingModel = embeddingModel; + conf.ensureSchemaExists(embeddingModel.dimensions()); prepareAddStatement(Set.of()); this.deleteStmt = prepareDeleteStatement(); @@ -159,7 +159,7 @@ public class CassandraVectorStore implements VectorStore, InitializingBean, Auto List primaryKeyValues = this.conf.documentIdTranslator.apply(d.getId()); if (null == d.getEmbedding() || d.getEmbedding().isEmpty()) { - d.setEmbedding(this.embeddingClient.embed(d)); + d.setEmbedding(this.embeddingModel.embed(d)); } BoundStatementBuilder builder = prepareAddStatement(d.getMetadata().keySet()).boundStatementBuilder(); @@ -204,7 +204,7 @@ public class CassandraVectorStore implements VectorStore, InitializingBean, Auto @Override public List similaritySearch(SearchRequest request) { Preconditions.checkArgument(request.getTopK() <= 1000); - var embedding = toFloatArray(this.embeddingClient.embed(request.getQuery())); + var embedding = toFloatArray(this.embeddingModel.embed(request.getQuery())); CqlVector cqlVector = CqlVector.newInstance(embedding); String whereClause = ""; @@ -256,7 +256,7 @@ public class CassandraVectorStore implements VectorStore, InitializingBean, Auto } void checkSchemaValid() { - this.conf.checkSchemaValid(embeddingClient.dimensions()); + this.conf.checkSchemaValid(embeddingModel.dimensions()); } private Similarity getIndexSimilarity(TableMetadata metadata) { diff --git a/vector-stores/spring-ai-cassandra-store/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStoreConfig.java b/vector-stores/spring-ai-cassandra-store/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStoreConfig.java index 7218f2697..e102c3091 100644 --- a/vector-stores/spring-ai-cassandra-store/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStoreConfig.java +++ b/vector-stores/spring-ai-cassandra-store/src/main/java/org/springframework/ai/vectorstore/CassandraVectorStoreConfig.java @@ -318,7 +318,7 @@ public class CassandraVectorStoreConfig implements AutoCloseable { /** * Executor to use when adding documents. The hotspot is the call to the - * embeddingClient. For remote transformers you probably want a higher value to + * embeddingModel. For remote transformers you probably want a higher value to * utilize network. For local transformers you probably want a lower value to * avoid saturation. **/ diff --git a/vector-stores/spring-ai-cassandra-store/src/test/java/org/springframework/ai/vectorstore/CassandraRichSchemaVectorStoreIT.java b/vector-stores/spring-ai-cassandra-store/src/test/java/org/springframework/ai/vectorstore/CassandraRichSchemaVectorStoreIT.java index 868c79fbb..26b336f36 100644 --- a/vector-stores/spring-ai-cassandra-store/src/test/java/org/springframework/ai/vectorstore/CassandraRichSchemaVectorStoreIT.java +++ b/vector-stores/spring-ai-cassandra-store/src/test/java/org/springframework/ai/vectorstore/CassandraRichSchemaVectorStoreIT.java @@ -43,8 +43,8 @@ import org.testcontainers.shaded.org.apache.commons.lang3.RandomStringUtils; import org.testcontainers.utility.DockerImageName; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.CassandraVectorStoreConfig.SchemaColumn; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.autoconfigure.EnableAutoConfiguration; @@ -197,7 +197,7 @@ class CassandraRichSchemaVectorStoreIT { try (CassandraVectorStore store = new CassandraVectorStore( storeBuilder(context, List.of()).withFixedThreadPoolExecutorSize(nThreads).build(), - context.getBean(EmbeddingClient.class))) { + context.getBean(EmbeddingModel.class))) { var executor = Executors.newFixedThreadPool((int) (nThreads * 1.2)); for (int k = 0; k < rounds; ++k) { @@ -489,9 +489,9 @@ class CassandraRichSchemaVectorStoreIT { public static class TestApplication { @Bean - public EmbeddingClient embeddingClient() { + public EmbeddingModel embeddingModel() { // default is ONNX all-MiniLM-L6-v2 - return new TransformersEmbeddingClient(); + return new TransformersEmbeddingModel(); } @Bean @@ -524,7 +524,7 @@ class CassandraRichSchemaVectorStoreIT { if (dropKeyspaceFirst) { conf.dropKeyspace(); } - return new StoreWrapper(new CassandraVectorStore(conf, context.getBean(EmbeddingClient.class)), conf); + return new StoreWrapper(new CassandraVectorStore(conf, context.getBean(EmbeddingModel.class)), conf); } static CassandraVectorStoreConfig.Builder storeBuilder(ApplicationContext context, diff --git a/vector-stores/spring-ai-cassandra-store/src/test/java/org/springframework/ai/vectorstore/CassandraVectorStoreIT.java b/vector-stores/spring-ai-cassandra-store/src/test/java/org/springframework/ai/vectorstore/CassandraVectorStoreIT.java index 27ed4246d..b362cc2c4 100644 --- a/vector-stores/spring-ai-cassandra-store/src/test/java/org/springframework/ai/vectorstore/CassandraVectorStoreIT.java +++ b/vector-stores/spring-ai-cassandra-store/src/test/java/org/springframework/ai/vectorstore/CassandraVectorStoreIT.java @@ -35,8 +35,8 @@ import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.utility.DockerImageName; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.CassandraVectorStoreConfig.SchemaColumn; import org.springframework.ai.vectorstore.CassandraVectorStoreConfig.SchemaColumnTags; import org.springframework.boot.SpringBootConfiguration; @@ -387,7 +387,7 @@ class CassandraVectorStoreIT { public static class TestApplication { @Bean - public CassandraVectorStore store(CqlSession cqlSession, EmbeddingClient embeddingClient) { + public CassandraVectorStore store(CqlSession cqlSession, EmbeddingModel embeddingModel) { CassandraVectorStoreConfig conf = storeBuilder(cqlSession) .addMetadataColumns(new SchemaColumn("meta1", DataTypes.TEXT), @@ -396,12 +396,12 @@ class CassandraVectorStoreIT { .build(); conf.dropKeyspace(); - return new CassandraVectorStore(conf, embeddingClient); + return new CassandraVectorStore(conf, embeddingModel); } @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } @Bean @@ -432,7 +432,7 @@ class CassandraVectorStoreIT { CassandraVectorStoreConfig.Builder builder) { CassandraVectorStoreConfig conf = builder.build(); conf.dropKeyspace(); - return new CassandraVectorStore(conf, context.getBean(EmbeddingClient.class)); + return new CassandraVectorStore(conf, context.getBean(EmbeddingModel.class)); } } diff --git a/vector-stores/spring-ai-cassandra-store/src/test/java/org/springframework/ai/vectorstore/WikiVectorStoreExample.java b/vector-stores/spring-ai-cassandra-store/src/test/java/org/springframework/ai/vectorstore/WikiVectorStoreExample.java index 910dac85a..301b61c49 100644 --- a/vector-stores/spring-ai-cassandra-store/src/test/java/org/springframework/ai/vectorstore/WikiVectorStoreExample.java +++ b/vector-stores/spring-ai-cassandra-store/src/test/java/org/springframework/ai/vectorstore/WikiVectorStoreExample.java @@ -24,8 +24,8 @@ import org.junit.jupiter.api.Assertions; import org.junit.jupiter.api.Test; import org.testcontainers.junit.jupiter.Testcontainers; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.CassandraVectorStoreConfig.SchemaColumn; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.autoconfigure.EnableAutoConfiguration; @@ -79,7 +79,7 @@ class WikiVectorStoreExample { public static class TestApplication { @Bean - public CassandraVectorStore store(CqlSession cqlSession, EmbeddingClient embeddingClient) { + public CassandraVectorStore store(CqlSession cqlSession, EmbeddingModel embeddingModel) { List partitionColumns = List.of(new SchemaColumn("wiki", DataTypes.TEXT), new SchemaColumn("language", DataTypes.TEXT), new SchemaColumn("title", DataTypes.TEXT)); @@ -119,13 +119,13 @@ class WikiVectorStoreExample { }) .build(); - return new CassandraVectorStore(conf, embeddingClient()); + return new CassandraVectorStore(conf, embeddingModel()); } @Bean - public EmbeddingClient embeddingClient() { + public EmbeddingModel embeddingModel() { // default is ONNX all-MiniLM-L6-v2 which is what we want - return new TransformersEmbeddingClient(); + return new TransformersEmbeddingModel(); } @Bean diff --git a/vector-stores/spring-ai-chroma-store/src/main/java/org/springframework/ai/vectorstore/ChromaVectorStore.java b/vector-stores/spring-ai-chroma-store/src/main/java/org/springframework/ai/vectorstore/ChromaVectorStore.java index 40ede3a67..d14d1f8b3 100644 --- a/vector-stores/spring-ai-chroma-store/src/main/java/org/springframework/ai/vectorstore/ChromaVectorStore.java +++ b/vector-stores/spring-ai-chroma-store/src/main/java/org/springframework/ai/vectorstore/ChromaVectorStore.java @@ -26,7 +26,7 @@ import org.springframework.ai.chroma.ChromaApi.AddEmbeddingsRequest; import org.springframework.ai.chroma.ChromaApi.DeleteEmbeddingsRequest; import org.springframework.ai.chroma.ChromaApi.Embedding; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.filter.FilterExpressionConverter; import org.springframework.ai.vectorstore.filter.converter.ChromaFilterExpressionConverter; import org.springframework.beans.factory.InitializingBean; @@ -37,9 +37,9 @@ import org.springframework.util.StringUtils; /** * {@link ChromaVectorStore} is a concrete implementation of the {@link VectorStore} * interface. It is responsible for adding, deleting, and searching documents based on - * their similarity to a query, using the {@link ChromaApi} and {@link EmbeddingClient} - * for embedding calculations. For more information about how it does this, see the - * official Chroma website. + * their similarity to a query, using the {@link ChromaApi} and {@link EmbeddingModel} for + * embedding calculations. For more information about how it does this, see the official + * Chroma website. */ public class ChromaVectorStore implements VectorStore, InitializingBean { @@ -51,7 +51,7 @@ public class ChromaVectorStore implements VectorStore, InitializingBean { public static final int DEFAULT_TOP_K = 4; - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; private final ChromaApi chromaApi; @@ -61,12 +61,12 @@ public class ChromaVectorStore implements VectorStore, InitializingBean { private String collectionId; - public ChromaVectorStore(EmbeddingClient embeddingClient, ChromaApi chromaApi) { - this(embeddingClient, chromaApi, DEFAULT_COLLECTION_NAME); + public ChromaVectorStore(EmbeddingModel embeddingModel, ChromaApi chromaApi) { + this(embeddingModel, chromaApi, DEFAULT_COLLECTION_NAME); } - public ChromaVectorStore(EmbeddingClient embeddingClient, ChromaApi chromaApi, String collectionName) { - this.embeddingClient = embeddingClient; + public ChromaVectorStore(EmbeddingModel embeddingModel, ChromaApi chromaApi, String collectionName) { + this.embeddingModel = embeddingModel; this.chromaApi = chromaApi; this.collectionName = collectionName; this.filterExpressionConverter = new ChromaFilterExpressionConverter(); @@ -93,7 +93,7 @@ public class ChromaVectorStore implements VectorStore, InitializingBean { ids.add(document.getId()); metadatas.add(document.getMetadata()); contents.add(document.getContent()); - document.setEmbedding(this.embeddingClient.embed(document)); + document.setEmbedding(this.embeddingModel.embed(document)); embeddings.add(JsonUtils.toFloatArray(document.getEmbedding())); } @@ -118,7 +118,7 @@ public class ChromaVectorStore implements VectorStore, InitializingBean { String query = request.getQuery(); Assert.notNull(query, "Query string must not be null"); - List embedding = this.embeddingClient.embed(query); + List embedding = this.embeddingModel.embed(query); Map where = (StringUtils.hasText(nativeFilterExpression)) ? JsonUtils.jsonToMap(nativeFilterExpression) : Map.of(); var queryRequest = new ChromaApi.QueryRequest(JsonUtils.toFloatList(embedding), request.getTopK(), where); diff --git a/vector-stores/spring-ai-chroma-store/src/test/java/org/springframework/ai/vectorstore/BasicAuthChromaWhereIT.java b/vector-stores/spring-ai-chroma-store/src/test/java/org/springframework/ai/vectorstore/BasicAuthChromaWhereIT.java index 8be4b8ac9..2d6c13101 100644 --- a/vector-stores/spring-ai-chroma-store/src/test/java/org/springframework/ai/vectorstore/BasicAuthChromaWhereIT.java +++ b/vector-stores/spring-ai-chroma-store/src/test/java/org/springframework/ai/vectorstore/BasicAuthChromaWhereIT.java @@ -25,9 +25,9 @@ import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.chroma.ChromaApi; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -107,13 +107,13 @@ public class BasicAuthChromaWhereIT { } @Bean - public VectorStore chromaVectorStore(EmbeddingClient embeddingClient, ChromaApi chromaApi) { - return new ChromaVectorStore(embeddingClient, chromaApi, "TestCollection"); + public VectorStore chromaVectorStore(EmbeddingModel embeddingModel, ChromaApi chromaApi) { + return new ChromaVectorStore(embeddingModel, chromaApi, "TestCollection"); } @Bean - public EmbeddingClient embeddingClient() { - return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); + public EmbeddingModel embeddingModel() { + return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); } } diff --git a/vector-stores/spring-ai-chroma-store/src/test/java/org/springframework/ai/vectorstore/ChromaVectorStoreIT.java b/vector-stores/spring-ai-chroma-store/src/test/java/org/springframework/ai/vectorstore/ChromaVectorStoreIT.java index 7ed5eeea8..918a90bdc 100644 --- a/vector-stores/spring-ai-chroma-store/src/test/java/org/springframework/ai/vectorstore/ChromaVectorStoreIT.java +++ b/vector-stores/spring-ai-chroma-store/src/test/java/org/springframework/ai/vectorstore/ChromaVectorStoreIT.java @@ -27,8 +27,8 @@ import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.chroma.ChromaApi; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -212,13 +212,13 @@ public class ChromaVectorStoreIT { } @Bean - public VectorStore chromaVectorStore(EmbeddingClient embeddingClient, ChromaApi chromaApi) { - return new ChromaVectorStore(embeddingClient, chromaApi, "TestCollection"); + public VectorStore chromaVectorStore(EmbeddingModel embeddingModel, ChromaApi chromaApi) { + return new ChromaVectorStore(embeddingModel, chromaApi, "TestCollection"); } @Bean - public EmbeddingClient embeddingClient() { - return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); + public EmbeddingModel embeddingModel() { + return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); } } diff --git a/vector-stores/spring-ai-chroma-store/src/test/java/org/springframework/ai/vectorstore/TokenSecuredChromaWhereIT.java b/vector-stores/spring-ai-chroma-store/src/test/java/org/springframework/ai/vectorstore/TokenSecuredChromaWhereIT.java index dac6e9bc0..22d08e57b 100644 --- a/vector-stores/spring-ai-chroma-store/src/test/java/org/springframework/ai/vectorstore/TokenSecuredChromaWhereIT.java +++ b/vector-stores/spring-ai-chroma-store/src/test/java/org/springframework/ai/vectorstore/TokenSecuredChromaWhereIT.java @@ -25,9 +25,9 @@ import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.chroma.ChromaApi; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.test.context.runner.ApplicationContextRunner; @@ -139,13 +139,13 @@ public class TokenSecuredChromaWhereIT { } @Bean - public VectorStore chromaVectorStore(EmbeddingClient embeddingClient, ChromaApi chromaApi) { - return new ChromaVectorStore(embeddingClient, chromaApi, "TestCollection"); + public VectorStore chromaVectorStore(EmbeddingModel embeddingModel, ChromaApi chromaApi) { + return new ChromaVectorStore(embeddingModel, chromaApi, "TestCollection"); } @Bean - public EmbeddingClient embeddingClient() { - return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); + public EmbeddingModel embeddingModel() { + return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); } } diff --git a/vector-stores/spring-ai-elasticsearch-store/src/main/java/org/springframework/ai/vectorstore/ElasticsearchVectorStore.java b/vector-stores/spring-ai-elasticsearch-store/src/main/java/org/springframework/ai/vectorstore/ElasticsearchVectorStore.java index ebc01eae5..9954f1828 100644 --- a/vector-stores/spring-ai-elasticsearch-store/src/main/java/org/springframework/ai/vectorstore/ElasticsearchVectorStore.java +++ b/vector-stores/spring-ai-elasticsearch-store/src/main/java/org/springframework/ai/vectorstore/ElasticsearchVectorStore.java @@ -34,7 +34,7 @@ import org.elasticsearch.client.RestClient; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.filter.Filter; import org.springframework.ai.vectorstore.filter.FilterExpressionConverter; import org.springframework.beans.factory.InitializingBean; @@ -58,7 +58,7 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean { private static final Logger logger = LoggerFactory.getLogger(ElasticsearchVectorStore.class); - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; private final ElasticsearchClient elasticsearchClient; @@ -68,17 +68,17 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean { private String similarityFunction; - public ElasticsearchVectorStore(RestClient restClient, EmbeddingClient embeddingClient) { - this(new ElasticsearchVectorStoreOptions(), restClient, embeddingClient); + public ElasticsearchVectorStore(RestClient restClient, EmbeddingModel embeddingModel) { + this(new ElasticsearchVectorStoreOptions(), restClient, embeddingModel); } public ElasticsearchVectorStore(ElasticsearchVectorStoreOptions options, RestClient restClient, - EmbeddingClient embeddingClient) { - Objects.requireNonNull(embeddingClient, "RestClient must not be null"); - Objects.requireNonNull(embeddingClient, "EmbeddingClient must not be null"); + EmbeddingModel embeddingModel) { + Objects.requireNonNull(embeddingModel, "RestClient must not be null"); + Objects.requireNonNull(embeddingModel, "EmbeddingModel must not be null"); this.elasticsearchClient = new ElasticsearchClient(new RestClientTransport(restClient, new JacksonJsonpMapper( new ObjectMapper().configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, false)))); - this.embeddingClient = embeddingClient; + this.embeddingModel = embeddingModel; this.options = options; this.filterExpressionConverter = new ElasticsearchAiSearchFilterExpressionConverter(); // the potential functions for vector fields at @@ -97,8 +97,8 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean { for (Document document : documents) { if (Objects.isNull(document.getEmbedding()) || document.getEmbedding().isEmpty()) { - logger.debug("Calling EmbeddingClient for document id = " + document.getId()); - document.setEmbedding(this.embeddingClient.embed(document)); + logger.debug("Calling EmbeddingModel for document id = " + document.getId()); + document.setEmbedding(this.embeddingModel.embed(document)); } builkRequestBuilder.operations(op -> op .index(idx -> idx.index(this.options.getIndexName()).id(document.getId()).document(document))); @@ -136,7 +136,7 @@ public class ElasticsearchVectorStore implements VectorStore, InitializingBean { @Override public List similaritySearch(SearchRequest searchRequest) { Assert.notNull(searchRequest, "The search request must not be null."); - return similaritySearch(this.embeddingClient.embed(searchRequest.getQuery()), searchRequest.getTopK(), + return similaritySearch(this.embeddingModel.embed(searchRequest.getQuery()), searchRequest.getTopK(), Double.valueOf(searchRequest.getSimilarityThreshold()).floatValue(), searchRequest.getFilterExpression()); } diff --git a/vector-stores/spring-ai-elasticsearch-store/src/test/java/org/springframework/ai/vectorstore/ElasticsearchVectorStoreIT.java b/vector-stores/spring-ai-elasticsearch-store/src/test/java/org/springframework/ai/vectorstore/ElasticsearchVectorStoreIT.java index c277ff1cb..393cfa7e6 100644 --- a/vector-stores/spring-ai-elasticsearch-store/src/test/java/org/springframework/ai/vectorstore/ElasticsearchVectorStoreIT.java +++ b/vector-stores/spring-ai-elasticsearch-store/src/test/java/org/springframework/ai/vectorstore/ElasticsearchVectorStoreIT.java @@ -39,8 +39,8 @@ import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.shaded.com.fasterxml.jackson.databind.ObjectMapper; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.autoconfigure.EnableAutoConfiguration; @@ -360,15 +360,15 @@ class ElasticsearchVectorStoreIT { public static class TestApplication { @Bean - public ElasticsearchVectorStore vectorStore(EmbeddingClient embeddingClient) { + public ElasticsearchVectorStore vectorStore(EmbeddingModel embeddingModel) { return new ElasticsearchVectorStore( RestClient.builder(HttpHost.create(elasticsearchContainer.getHttpHostAddress())).build(), - embeddingClient); + embeddingModel); } @Bean - public EmbeddingClient embeddingClient() { - return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); + public EmbeddingModel embeddingModel() { + return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); } } diff --git a/vector-stores/spring-ai-gemfire-store/src/main/java/org/springframework/ai/vectorstore/GemFireVectorStore.java b/vector-stores/spring-ai-gemfire-store/src/main/java/org/springframework/ai/vectorstore/GemFireVectorStore.java index df7e08ec5..c91fff95d 100644 --- a/vector-stores/spring-ai-gemfire-store/src/main/java/org/springframework/ai/vectorstore/GemFireVectorStore.java +++ b/vector-stores/spring-ai-gemfire-store/src/main/java/org/springframework/ai/vectorstore/GemFireVectorStore.java @@ -31,7 +31,7 @@ import com.fasterxml.jackson.databind.ObjectMapper; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.http.HttpMethod; import org.springframework.http.MediaType; import org.springframework.util.Assert; @@ -60,7 +60,7 @@ public class GemFireVectorStore implements VectorStore { private final WebClient client; - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; private final int topKPerBucket; @@ -192,11 +192,11 @@ public class GemFireVectorStore implements VectorStore { this.indexName = indexName; } - public GemFireVectorStore(GemFireVectorStoreConfig config, EmbeddingClient embedding) { + public GemFireVectorStore(GemFireVectorStoreConfig config, EmbeddingModel embedding) { Assert.notNull(config, "GemFireVectorStoreConfig must not be null"); - Assert.notNull(embedding, "EmbeddingClient must not be null"); + Assert.notNull(embedding, "EmbeddingModel must not be null"); this.client = config.client; - this.embeddingClient = embedding; + this.embeddingModel = embedding; this.topKPerBucket = config.topKPerBucket; this.topK = config.topK; this.documentField = config.documentField; @@ -417,7 +417,7 @@ public class GemFireVectorStore implements VectorStore { public void add(List documents) { UploadRequest upload = new UploadRequest(documents.stream().map(document -> { // Compute and assign an embedding to the document. - document.setEmbedding(this.embeddingClient.embed(document)); + document.setEmbedding(this.embeddingModel.embed(document)); List floatVector = document.getEmbedding().stream().map(Double::floatValue).toList(); return new UploadRequest.Embedding(document.getId(), floatVector, documentField, document.getContent(), document.getMetadata()); @@ -465,7 +465,7 @@ public class GemFireVectorStore implements VectorStore { if (request.hasFilterExpression()) { throw new UnsupportedOperationException("Gemfire does not support metadata filter expressions yet."); } - List vector = this.embeddingClient.embed(request.getQuery()); + List vector = this.embeddingModel.embed(request.getQuery()); List floatVector = vector.stream().map(Double::floatValue).toList(); return client.post() diff --git a/vector-stores/spring-ai-gemfire-store/src/test/java/org/springframework/ai/vectorstore/GemFireVectorStoreIT.java b/vector-stores/spring-ai-gemfire-store/src/test/java/org/springframework/ai/vectorstore/GemFireVectorStoreIT.java index 49d032544..21de25e50 100644 --- a/vector-stores/spring-ai-gemfire-store/src/test/java/org/springframework/ai/vectorstore/GemFireVectorStoreIT.java +++ b/vector-stores/spring-ai-gemfire-store/src/test/java/org/springframework/ai/vectorstore/GemFireVectorStoreIT.java @@ -33,8 +33,8 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.GemFireVectorStore.GemFireVectorStoreConfig; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.autoconfigure.EnableAutoConfiguration; @@ -189,15 +189,15 @@ public class GemFireVectorStoreIT { } @Bean - public GemFireVectorStore vectorStore(GemFireVectorStoreConfig config, EmbeddingClient embeddingClient) { - GemFireVectorStore gemFireVectorStore = new GemFireVectorStore(config, embeddingClient); + public GemFireVectorStore vectorStore(GemFireVectorStoreConfig config, EmbeddingModel embeddingModel) { + GemFireVectorStore gemFireVectorStore = new GemFireVectorStore(config, embeddingModel); gemFireVectorStore.setIndexName(INDEX_NAME); return gemFireVectorStore; } @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/vector-stores/spring-ai-hanadb-store/src/main/java/org/springframework/ai/vectorstore/HanaCloudVectorStore.java b/vector-stores/spring-ai-hanadb-store/src/main/java/org/springframework/ai/vectorstore/HanaCloudVectorStore.java index 4f0415b7b..db8d44e01 100644 --- a/vector-stores/spring-ai-hanadb-store/src/main/java/org/springframework/ai/vectorstore/HanaCloudVectorStore.java +++ b/vector-stores/spring-ai-hanadb-store/src/main/java/org/springframework/ai/vectorstore/HanaCloudVectorStore.java @@ -19,7 +19,7 @@ import com.fasterxml.jackson.core.JsonProcessingException; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import java.util.Collections; import java.util.List; @@ -50,7 +50,7 @@ import java.util.stream.Collectors; * 2024 * * Hana DB introduced a new datatype REAL_VECTOR that can store embeddings - * generated by org.springframework.ai.embedding.EmbeddingClient + * generated by org.springframework.ai.embedding.EmbeddingModel * * @author Rahul Mittal * @see repository; - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; private final HanaCloudVectorStoreConfig config; public HanaCloudVectorStore(HanaVectorRepository repository, - EmbeddingClient embeddingClient, HanaCloudVectorStoreConfig config) { + EmbeddingModel embeddingModel, HanaCloudVectorStoreConfig config) { this.repository = repository; - this.embeddingClient = embeddingClient; + this.embeddingModel = embeddingModel; this.config = config; } @@ -79,7 +79,7 @@ public class HanaCloudVectorStore implements VectorStore { public void add(List documents) { int count = 1; for (Document document : documents) { - logger.info("[{}/{}] Calling EmbeddingClient for document id = {}", count++, documents.size(), + logger.info("[{}/{}] Calling EmbeddingModel for document id = {}", count++, documents.size(), document.getId()); String content = document.getContent().replaceAll("\\s+", " "); String embedding = getEmbedding(document); @@ -130,15 +130,14 @@ public class HanaCloudVectorStore implements VectorStore { } private String getEmbedding(SearchRequest searchRequest) { - return "[" + this.embeddingClient.embed(searchRequest.getQuery()) + return "[" + this.embeddingModel.embed(searchRequest.getQuery()) .stream() .map(String::valueOf) .collect(Collectors.joining(", ")) + "]"; } private String getEmbedding(Document document) { - return "[" - + this.embeddingClient.embed(document).stream().map(String::valueOf).collect(Collectors.joining(", ")) + return "[" + this.embeddingModel.embed(document).stream().map(String::valueOf).collect(Collectors.joining(", ")) + "]"; } diff --git a/vector-stores/spring-ai-hanadb-store/src/test/java/org/springframework/ai/vectorstore/CricketWorldCupHanaController.java b/vector-stores/spring-ai-hanadb-store/src/test/java/org/springframework/ai/vectorstore/CricketWorldCupHanaController.java index ad3f954a7..9b24fa4db 100644 --- a/vector-stores/spring-ai-hanadb-store/src/test/java/org/springframework/ai/vectorstore/CricketWorldCupHanaController.java +++ b/vector-stores/spring-ai-hanadb-store/src/test/java/org/springframework/ai/vectorstore/CricketWorldCupHanaController.java @@ -17,7 +17,7 @@ package org.springframework.ai.vectorstore; import org.slf4j.Logger; import org.slf4j.LoggerFactory; -import org.springframework.ai.chat.ChatClient; +import org.springframework.ai.chat.ChatModel; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.chat.prompt.SystemPromptTemplate; @@ -51,11 +51,11 @@ public class CricketWorldCupHanaController { private final VectorStore hanaCloudVectorStore; - private final ChatClient chatClient; + private final ChatModel chatModel; @Autowired - public CricketWorldCupHanaController(ChatClient chatClient, VectorStore hanaCloudVectorStore) { - this.chatClient = chatClient; + public CricketWorldCupHanaController(ChatModel chatModel, VectorStore hanaCloudVectorStore) { + this.chatModel = chatModel; this.hanaCloudVectorStore = hanaCloudVectorStore; } @@ -88,7 +88,7 @@ public class CricketWorldCupHanaController { var userMessage = new UserMessage(message); Prompt prompt = new Prompt(List.of(similarDocsMessage, userMessage)); - String generation = chatClient.call(prompt).getResult().getOutput().getContent(); + String generation = chatModel.call(prompt).getResult().getOutput().getContent(); logger.info("Generation: {}", generation); return Map.of("generation", generation); } diff --git a/vector-stores/spring-ai-hanadb-store/src/test/java/org/springframework/ai/vectorstore/HanaCloudVectorStoreIT.java b/vector-stores/spring-ai-hanadb-store/src/test/java/org/springframework/ai/vectorstore/HanaCloudVectorStoreIT.java index e93b0c95a..313bd69be 100644 --- a/vector-stores/spring-ai-hanadb-store/src/test/java/org/springframework/ai/vectorstore/HanaCloudVectorStoreIT.java +++ b/vector-stores/spring-ai-hanadb-store/src/test/java/org/springframework/ai/vectorstore/HanaCloudVectorStoreIT.java @@ -28,8 +28,8 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.reader.pdf.PagePdfDocumentReader; import org.springframework.ai.transformer.splitter.TokenTextSplitter; @@ -87,8 +87,8 @@ public class HanaCloudVectorStoreIT { @Bean public VectorStore hanaCloudVectorStore(CricketWorldCupRepository cricketWorldCupRepository, - EmbeddingClient embeddingClient) { - return new HanaCloudVectorStore(cricketWorldCupRepository, embeddingClient, + EmbeddingModel embeddingModel) { + return new HanaCloudVectorStore(cricketWorldCupRepository, embeddingModel, HanaCloudVectorStoreConfig.builder().tableName("CRICKET_WORLD_CUP").topK(1).build()); } @@ -122,8 +122,8 @@ public class HanaCloudVectorStoreIT { } @Bean - public EmbeddingClient embeddingClient() { - return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); + public EmbeddingModel embeddingModel() { + return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); } } diff --git a/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java b/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java index 0a6004a1c..28627048f 100644 --- a/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java +++ b/vector-stores/spring-ai-milvus-store/src/main/java/org/springframework/ai/vectorstore/MilvusVectorStore.java @@ -52,7 +52,7 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.filter.FilterExpressionConverter; import org.springframework.ai.vectorstore.filter.converter.MilvusFilterExpressionConverter; import org.springframework.beans.factory.InitializingBean; @@ -92,7 +92,7 @@ public class MilvusVectorStore implements VectorStore, InitializingBean { private final MilvusServiceClient milvusClient; - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; private final MilvusVectorStoreConfig config; @@ -242,18 +242,18 @@ public class MilvusVectorStore implements VectorStore, InitializingBean { } - public MilvusVectorStore(MilvusServiceClient milvusClient, EmbeddingClient embeddingClient) { - this(milvusClient, embeddingClient, MilvusVectorStoreConfig.defaultConfig()); + public MilvusVectorStore(MilvusServiceClient milvusClient, EmbeddingModel embeddingModel) { + this(milvusClient, embeddingModel, MilvusVectorStoreConfig.defaultConfig()); } - public MilvusVectorStore(MilvusServiceClient milvusClient, EmbeddingClient embeddingClient, + public MilvusVectorStore(MilvusServiceClient milvusClient, EmbeddingModel embeddingModel, MilvusVectorStoreConfig config) { Assert.notNull(milvusClient, "MilvusServiceClient must not be null"); - Assert.notNull(milvusClient, "EmbeddingClient must not be null"); + Assert.notNull(milvusClient, "EmbeddingModel must not be null"); this.milvusClient = milvusClient; - this.embeddingClient = embeddingClient; + this.embeddingModel = embeddingModel; this.config = config; } @@ -268,7 +268,7 @@ public class MilvusVectorStore implements VectorStore, InitializingBean { List> embeddingArray = new ArrayList<>(); for (Document document : documents) { - List embedding = this.embeddingClient.embed(document); + List embedding = this.embeddingModel.embed(document); docIdArray.add(document.getId()); // Use a (future) DocumentTextLayoutFormatter instance to extract @@ -328,7 +328,7 @@ public class MilvusVectorStore implements VectorStore, InitializingBean { Assert.notNull(request.getQuery(), "Query string must not be null"); - List embedding = this.embeddingClient.embed(request.getQuery()); + List embedding = this.embeddingModel.embed(request.getQuery()); var searchParamBuilder = SearchParam.newBuilder() .withCollectionName(this.config.collectionName) @@ -481,13 +481,13 @@ public class MilvusVectorStore implements VectorStore, InitializingBean { return this.config.embeddingDimension; } try { - int embeddingDimensions = this.embeddingClient.dimensions(); + int embeddingDimensions = this.embeddingModel.dimensions(); if (embeddingDimensions > 0) { return embeddingDimensions; } } catch (Exception e) { - logger.warn("Failed to obtain the embedding dimensions from the embedding client and fall backs to default:" + logger.warn("Failed to obtain the embedding dimensions from the embedding model and fall backs to default:" + this.config.embeddingDimension, e); } return OPENAI_EMBEDDING_DIMENSION_SIZE; diff --git a/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusEmbeddingDimensionsTests.java b/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusEmbeddingDimensionsTests.java index 2e1e94d5b..a477cc2c8 100644 --- a/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusEmbeddingDimensionsTests.java +++ b/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusEmbeddingDimensionsTests.java @@ -21,7 +21,7 @@ import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.MilvusVectorStore.MilvusVectorStoreConfig; import static org.assertj.core.api.Assertions.assertThat; @@ -37,7 +37,7 @@ import static org.mockito.Mockito.when; public class MilvusEmbeddingDimensionsTests { @Mock - private EmbeddingClient embeddingClient; + private EmbeddingModel embeddingModel; @Mock private MilvusServiceClient milvusClient; @@ -51,36 +51,36 @@ public class MilvusEmbeddingDimensionsTests { .withEmbeddingDimension(explicitDimensions) .build(); - var dim = new MilvusVectorStore(milvusClient, embeddingClient, config).embeddingDimensions(); + var dim = new MilvusVectorStore(milvusClient, embeddingModel, config).embeddingDimensions(); assertThat(dim).isEqualTo(explicitDimensions); - verify(embeddingClient, never()).dimensions(); + verify(embeddingModel, never()).dimensions(); } @Test - public void embeddingClientDimensions() { - when(embeddingClient.dimensions()).thenReturn(969); + public void embeddingModelDimensions() { + when(embeddingModel.dimensions()).thenReturn(969); MilvusVectorStoreConfig config = MilvusVectorStoreConfig.builder().build(); - var dim = new MilvusVectorStore(milvusClient, embeddingClient, config).embeddingDimensions(); + var dim = new MilvusVectorStore(milvusClient, embeddingModel, config).embeddingDimensions(); assertThat(dim).isEqualTo(969); - verify(embeddingClient, only()).dimensions(); + verify(embeddingModel, only()).dimensions(); } @Test public void fallBackToDefaultDimensions() { - when(embeddingClient.dimensions()).thenThrow(new RuntimeException()); + when(embeddingModel.dimensions()).thenThrow(new RuntimeException()); - var dim = new MilvusVectorStore(milvusClient, embeddingClient, + var dim = new MilvusVectorStore(milvusClient, embeddingModel, MilvusVectorStoreConfig.builder().build()) .embeddingDimensions(); assertThat(dim).isEqualTo(MilvusVectorStore.OPENAI_EMBEDDING_DIMENSION_SIZE); - verify(embeddingClient, only()).dimensions(); + verify(embeddingModel, only()).dimensions(); } } diff --git a/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusVectorStoreIT.java b/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusVectorStoreIT.java index b59653888..4a1c00bab 100644 --- a/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusVectorStoreIT.java +++ b/vector-stores/spring-ai-milvus-store/src/test/java/org/springframework/ai/vectorstore/MilvusVectorStoreIT.java @@ -33,8 +33,8 @@ import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.vectorstore.MilvusVectorStore.MilvusVectorStoreConfig; import org.springframework.beans.factory.annotation.Value; @@ -258,14 +258,14 @@ public class MilvusVectorStoreIT { private MetricType metricType; @Bean - public VectorStore vectorStore(MilvusServiceClient milvusClient, EmbeddingClient embeddingClient) { + public VectorStore vectorStore(MilvusServiceClient milvusClient, EmbeddingModel embeddingModel) { MilvusVectorStoreConfig config = MilvusVectorStoreConfig.builder() .withCollectionName("test_vector_store") .withDatabaseName("default") .withIndexType(IndexType.IVF_FLAT) .withMetricType(metricType) .build(); - return new MilvusVectorStore(milvusClient, embeddingClient, config); + return new MilvusVectorStore(milvusClient, embeddingModel, config); } @Bean @@ -277,9 +277,9 @@ public class MilvusVectorStoreIT { } @Bean - public EmbeddingClient embeddingClient() { - return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); - // return new OpenAiEmbeddingClient(new + public EmbeddingModel embeddingModel() { + return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); + // return new OpenAiEmbeddingModel(new // OpenAiApi(System.getenv("OPENAI_API_KEY")), MetadataMode.EMBED, // OpenAiEmbeddingOptions.builder().withModel("text-embedding-ada-002").build()); } diff --git a/vector-stores/spring-ai-mongodb-atlas-store/src/main/java/org/springframework/ai/vectorstore/MongoDBAtlasVectorStore.java b/vector-stores/spring-ai-mongodb-atlas-store/src/main/java/org/springframework/ai/vectorstore/MongoDBAtlasVectorStore.java index 0fbb099ce..de7d8cc06 100644 --- a/vector-stores/spring-ai-mongodb-atlas-store/src/main/java/org/springframework/ai/vectorstore/MongoDBAtlasVectorStore.java +++ b/vector-stores/spring-ai-mongodb-atlas-store/src/main/java/org/springframework/ai/vectorstore/MongoDBAtlasVectorStore.java @@ -24,7 +24,7 @@ import java.util.Optional; import com.mongodb.BasicDBObject; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.beans.factory.InitializingBean; import org.springframework.data.mongodb.core.MongoTemplate; import org.springframework.data.mongodb.core.aggregation.Aggregation; @@ -58,20 +58,20 @@ public class MongoDBAtlasVectorStore implements VectorStore, InitializingBean { private final MongoTemplate mongoTemplate; - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; private final MongoDBVectorStoreConfig config; private final MongoDBAtlasFilterExpressionConverter filterExpressionConverter = new MongoDBAtlasFilterExpressionConverter(); - public MongoDBAtlasVectorStore(MongoTemplate mongoTemplate, EmbeddingClient embeddingClient) { - this(mongoTemplate, embeddingClient, MongoDBVectorStoreConfig.defaultConfig()); + public MongoDBAtlasVectorStore(MongoTemplate mongoTemplate, EmbeddingModel embeddingModel) { + this(mongoTemplate, embeddingModel, MongoDBVectorStoreConfig.defaultConfig()); } - public MongoDBAtlasVectorStore(MongoTemplate mongoTemplate, EmbeddingClient embeddingClient, + public MongoDBAtlasVectorStore(MongoTemplate mongoTemplate, EmbeddingModel embeddingModel, MongoDBVectorStoreConfig config) { this.mongoTemplate = mongoTemplate; - this.embeddingClient = embeddingClient; + this.embeddingModel = embeddingModel; this.config = config; } @@ -94,7 +94,7 @@ public class MongoDBAtlasVectorStore implements VectorStore, InitializingBean { vectorFields.add(new org.bson.Document().append("type", "vector") .append("path", this.config.pathName) - .append("numDimensions", this.embeddingClient.dimensions()) + .append("numDimensions", this.embeddingModel.dimensions()) .append("similarity", "cosine")); vectorFields.addAll(this.config.metadataFieldsToFilter.stream() @@ -129,7 +129,7 @@ public class MongoDBAtlasVectorStore implements VectorStore, InitializingBean { @Override public void add(List documents) { for (Document document : documents) { - List embedding = this.embeddingClient.embed(document); + List embedding = this.embeddingModel.embed(document); document.setEmbedding(embedding); this.mongoTemplate.save(document, this.config.collectionName); } @@ -156,7 +156,7 @@ public class MongoDBAtlasVectorStore implements VectorStore, InitializingBean { String nativeFilterExpressions = (request.getFilterExpression() != null) ? this.filterExpressionConverter.convertExpression(request.getFilterExpression()) : ""; - List queryEmbedding = this.embeddingClient.embed(request.getQuery()); + List queryEmbedding = this.embeddingModel.embed(request.getQuery()); var vectorSearch = new VectorSearchAggregation(queryEmbedding, this.config.pathName, this.config.numCandidates, this.config.vectorIndexName, request.getTopK(), nativeFilterExpressions); diff --git a/vector-stores/spring-ai-mongodb-atlas-store/src/test/java/org/springframework/ai/vectorstore/MongoDBAtlasVectorStoreIT.java b/vector-stores/spring-ai-mongodb-atlas-store/src/test/java/org/springframework/ai/vectorstore/MongoDBAtlasVectorStoreIT.java index 1e555dd22..c19963cdd 100644 --- a/vector-stores/spring-ai-mongodb-atlas-store/src/test/java/org/springframework/ai/vectorstore/MongoDBAtlasVectorStoreIT.java +++ b/vector-stores/spring-ai-mongodb-atlas-store/src/test/java/org/springframework/ai/vectorstore/MongoDBAtlasVectorStoreIT.java @@ -21,8 +21,8 @@ import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Test; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.autoconfigure.EnableAutoConfiguration; @@ -192,8 +192,8 @@ class MongoDBAtlasVectorStoreIT { public static class TestApplication { @Bean - public VectorStore vectorStore(MongoTemplate mongoTemplate, EmbeddingClient embeddingClient) { - return new MongoDBAtlasVectorStore(mongoTemplate, embeddingClient, + public VectorStore vectorStore(MongoTemplate mongoTemplate, EmbeddingModel embeddingModel) { + return new MongoDBAtlasVectorStore(mongoTemplate, embeddingModel, MongoDBAtlasVectorStore.MongoDBVectorStoreConfig.builder() .withMetadataFieldsToFilter(List.of("country", "year")) .build()); @@ -205,8 +205,8 @@ class MongoDBAtlasVectorStoreIT { } @Bean - public EmbeddingClient embeddingClient() { - return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); + public EmbeddingModel embeddingModel() { + return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); } } diff --git a/vector-stores/spring-ai-neo4j-store/src/main/java/org/springframework/ai/vectorstore/Neo4jVectorStore.java b/vector-stores/spring-ai-neo4j-store/src/main/java/org/springframework/ai/vectorstore/Neo4jVectorStore.java index 8ac8fd204..b84854b43 100644 --- a/vector-stores/spring-ai-neo4j-store/src/main/java/org/springframework/ai/vectorstore/Neo4jVectorStore.java +++ b/vector-stores/spring-ai-neo4j-store/src/main/java/org/springframework/ai/vectorstore/Neo4jVectorStore.java @@ -20,7 +20,7 @@ import org.neo4j.driver.Driver; import org.neo4j.driver.SessionConfig; import org.neo4j.driver.Values; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.filter.Neo4jVectorFilterExpressionConverter; import org.springframework.beans.factory.InitializingBean; import org.springframework.util.Assert; @@ -269,17 +269,17 @@ public class Neo4jVectorStore implements VectorStore, InitializingBean { private final Driver driver; - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; private final Neo4jVectorStoreConfig config; - public Neo4jVectorStore(Driver driver, EmbeddingClient embeddingClient, Neo4jVectorStoreConfig config) { + public Neo4jVectorStore(Driver driver, EmbeddingModel embeddingModel, Neo4jVectorStoreConfig config) { Assert.notNull(driver, "Neo4j driver must not be null"); - Assert.notNull(embeddingClient, "Embedding client must not be null"); + Assert.notNull(embeddingModel, "Embedding client must not be null"); this.driver = driver; - this.embeddingClient = embeddingClient; + this.embeddingModel = embeddingModel; this.config = config; } @@ -328,7 +328,7 @@ public class Neo4jVectorStore implements VectorStore, InitializingBean { Assert.isTrue(request.getSimilarityThreshold() >= 0 && request.getSimilarityThreshold() <= 1, "The similarity score is bounded between 0 and 1; least to most similar respectively."); - var embedding = Values.value(toFloatArray(this.embeddingClient.embed(request.getQuery()))); + var embedding = Values.value(toFloatArray(this.embeddingModel.embed(request.getQuery()))); try (var session = this.driver.session(this.config.sessionConfig)) { StringBuilder condition = new StringBuilder("score >= $threshold"); if (request.hasFilterExpression()) { @@ -372,7 +372,7 @@ public class Neo4jVectorStore implements VectorStore, InitializingBean { } private Map documentToRecord(Document document) { - var embedding = this.embeddingClient.embed(document); + var embedding = this.embeddingModel.embed(document); document.setEmbedding(embedding); var row = new HashMap(); diff --git a/vector-stores/spring-ai-neo4j-store/src/test/java/org/springframework/ai/vectorstore/Neo4jVectorStoreIT.java b/vector-stores/spring-ai-neo4j-store/src/test/java/org/springframework/ai/vectorstore/Neo4jVectorStoreIT.java index a258d8153..34433ef3b 100644 --- a/vector-stores/spring-ai-neo4j-store/src/test/java/org/springframework/ai/vectorstore/Neo4jVectorStoreIT.java +++ b/vector-stores/spring-ai-neo4j-store/src/test/java/org/springframework/ai/vectorstore/Neo4jVectorStoreIT.java @@ -34,9 +34,9 @@ import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.utility.DockerImageName; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.autoconfigure.EnableAutoConfiguration; import org.springframework.boot.autoconfigure.jdbc.DataSourceAutoConfiguration; @@ -293,9 +293,9 @@ class Neo4jVectorStoreIT { public static class TestApplication { @Bean - public VectorStore vectorStore(Driver driver, EmbeddingClient embeddingClient) { + public VectorStore vectorStore(Driver driver, EmbeddingModel embeddingModel) { - return new Neo4jVectorStore(driver, embeddingClient, + return new Neo4jVectorStore(driver, embeddingModel, Neo4jVectorStore.Neo4jVectorStoreConfig.defaultConfig()); } @@ -306,8 +306,8 @@ class Neo4jVectorStoreIT { } @Bean - public EmbeddingClient embeddingClient() { - return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); + public EmbeddingModel embeddingModel() { + return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); } } diff --git a/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java b/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java index c2b7954d9..e95840e71 100644 --- a/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java +++ b/vector-stores/spring-ai-pgvector-store/src/main/java/org/springframework/ai/vectorstore/PgVectorStore.java @@ -32,7 +32,7 @@ import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.filter.FilterExpressionConverter; import org.springframework.ai.vectorstore.filter.converter.PgVectorFilterExpressionConverter; import org.springframework.beans.factory.InitializingBean; @@ -66,7 +66,7 @@ public class PgVectorStore implements VectorStore, InitializingBean { private final JdbcTemplate jdbcTemplate; - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; private int dimensions; @@ -197,21 +197,21 @@ public class PgVectorStore implements VectorStore, InitializingBean { } - public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingClient embeddingClient) { - this(jdbcTemplate, embeddingClient, INVALID_EMBEDDING_DIMENSION, PgVectorStore.PgDistanceType.COSINE_DISTANCE, + public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingModel embeddingModel) { + this(jdbcTemplate, embeddingModel, INVALID_EMBEDDING_DIMENSION, PgVectorStore.PgDistanceType.COSINE_DISTANCE, false, PgIndexType.NONE); } - public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingClient embeddingClient, int dimensions) { - this(jdbcTemplate, embeddingClient, dimensions, PgVectorStore.PgDistanceType.COSINE_DISTANCE, false, + public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingModel embeddingModel, int dimensions) { + this(jdbcTemplate, embeddingModel, dimensions, PgVectorStore.PgDistanceType.COSINE_DISTANCE, false, PgIndexType.NONE); } - public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingClient embeddingClient, int dimensions, + public PgVectorStore(JdbcTemplate jdbcTemplate, EmbeddingModel embeddingModel, int dimensions, PgDistanceType distanceType, boolean removeExistingVectorStoreTable, PgIndexType createIndexMethod) { this.jdbcTemplate = jdbcTemplate; - this.embeddingClient = embeddingClient; + this.embeddingModel = embeddingModel; this.dimensions = dimensions; this.distanceType = distanceType; this.removeExistingVectorStoreTable = removeExistingVectorStoreTable; @@ -237,7 +237,7 @@ public class PgVectorStore implements VectorStore, InitializingBean { var document = documents.get(i); var content = document.getContent(); var json = toJson(document.getMetadata()); - var pGvector = new PGvector(toFloatArray(embeddingClient.embed(document))); + var pGvector = new PGvector(toFloatArray(embeddingModel.embed(document))); StatementCreatorUtils.setParameterValue(ps, 1, SqlTypeValue.TYPE_UNKNOWN, UUID.fromString(document.getId())); @@ -320,7 +320,7 @@ public class PgVectorStore implements VectorStore, InitializingBean { } private PGvector getQueryEmbedding(String query) { - List embedding = this.embeddingClient.embed(query); + List embedding = this.embeddingModel.embed(query); return new PGvector(toFloatArray(embedding)); } @@ -366,13 +366,13 @@ public class PgVectorStore implements VectorStore, InitializingBean { } try { - int embeddingDimensions = this.embeddingClient.dimensions(); + int embeddingDimensions = this.embeddingModel.dimensions(); if (embeddingDimensions > 0) { return embeddingDimensions; } } catch (Exception e) { - logger.warn("Failed to obtain the embedding dimensions from the embedding client and fall backs to default:" + logger.warn("Failed to obtain the embedding dimensions from the embedding model and fall backs to default:" + OPENAI_EMBEDDING_DIMENSION_SIZE, e); } return OPENAI_EMBEDDING_DIMENSION_SIZE; diff --git a/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorEmbeddingDimensionsTests.java b/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorEmbeddingDimensionsTests.java index cff66185b..4fa2c56a5 100644 --- a/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorEmbeddingDimensionsTests.java +++ b/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorEmbeddingDimensionsTests.java @@ -20,7 +20,7 @@ import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.jdbc.core.JdbcTemplate; import static org.assertj.core.api.Assertions.assertThat; @@ -36,7 +36,7 @@ import static org.mockito.Mockito.when; public class PgVectorEmbeddingDimensionsTests { @Mock - private EmbeddingClient embeddingClient; + private EmbeddingModel embeddingModel; @Mock private JdbcTemplate jdbcTemplate; @@ -46,32 +46,32 @@ public class PgVectorEmbeddingDimensionsTests { final int explicitDimensions = 696; - var dim = new PgVectorStore(jdbcTemplate, embeddingClient, explicitDimensions).embeddingDimensions(); + var dim = new PgVectorStore(jdbcTemplate, embeddingModel, explicitDimensions).embeddingDimensions(); assertThat(dim).isEqualTo(explicitDimensions); - verify(embeddingClient, never()).dimensions(); + verify(embeddingModel, never()).dimensions(); } @Test - public void embeddingClientDimensions() { - when(embeddingClient.dimensions()).thenReturn(969); + public void embeddingModelDimensions() { + when(embeddingModel.dimensions()).thenReturn(969); - var dim = new PgVectorStore(jdbcTemplate, embeddingClient).embeddingDimensions(); + var dim = new PgVectorStore(jdbcTemplate, embeddingModel).embeddingDimensions(); assertThat(dim).isEqualTo(969); - verify(embeddingClient, only()).dimensions(); + verify(embeddingModel, only()).dimensions(); } @Test public void fallBackToDefaultDimensions() { - when(embeddingClient.dimensions()).thenThrow(new RuntimeException()); + when(embeddingModel.dimensions()).thenThrow(new RuntimeException()); - var dim = new PgVectorStore(jdbcTemplate, embeddingClient).embeddingDimensions(); + var dim = new PgVectorStore(jdbcTemplate, embeddingModel).embeddingDimensions(); assertThat(dim).isEqualTo(PgVectorStore.OPENAI_EMBEDDING_DIMENSION_SIZE); - verify(embeddingClient, only()).dimensions(); + verify(embeddingModel, only()).dimensions(); } } diff --git a/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java b/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java index 9bfb63e56..ca53406ce 100644 --- a/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java +++ b/vector-stores/spring-ai-pgvector-store/src/test/java/org/springframework/ai/vectorstore/PgVectorStoreIT.java @@ -35,9 +35,9 @@ import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.openai.api.OpenAiApi; -import org.springframework.ai.openai.OpenAiEmbeddingClient; +import org.springframework.ai.openai.OpenAiEmbeddingModel; import org.springframework.ai.vectorstore.PgVectorStore.PgIndexType; import org.springframework.ai.vectorstore.filter.FilterExpressionTextParser.FilterExpressionParseException; import org.springframework.beans.factory.annotation.Value; @@ -306,8 +306,8 @@ public class PgVectorStoreIT { PgVectorStore.PgDistanceType distanceType; @Bean - public VectorStore vectorStore(JdbcTemplate jdbcTemplate, EmbeddingClient embeddingClient) { - return new PgVectorStore(jdbcTemplate, embeddingClient, PgVectorStore.INVALID_EMBEDDING_DIMENSION, + public VectorStore vectorStore(JdbcTemplate jdbcTemplate, EmbeddingModel embeddingModel) { + return new PgVectorStore(jdbcTemplate, embeddingModel, PgVectorStore.INVALID_EMBEDDING_DIMENSION, distanceType, true, PgIndexType.HNSW); } @@ -329,8 +329,8 @@ public class PgVectorStoreIT { } @Bean - public EmbeddingClient embeddingClient() { - return new OpenAiEmbeddingClient(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); + public EmbeddingModel embeddingModel() { + return new OpenAiEmbeddingModel(new OpenAiApi(System.getenv("OPENAI_API_KEY"))); } } diff --git a/vector-stores/spring-ai-pinecone-store/src/main/java/org/springframework/ai/vectorstore/PineconeVectorStore.java b/vector-stores/spring-ai-pinecone-store/src/main/java/org/springframework/ai/vectorstore/PineconeVectorStore.java index a2e083f8a..992e8c859 100644 --- a/vector-stores/spring-ai-pinecone-store/src/main/java/org/springframework/ai/vectorstore/PineconeVectorStore.java +++ b/vector-stores/spring-ai-pinecone-store/src/main/java/org/springframework/ai/vectorstore/PineconeVectorStore.java @@ -35,7 +35,7 @@ import io.pinecone.proto.UpsertRequest; import io.pinecone.proto.Vector; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.filter.FilterExpressionConverter; import org.springframework.ai.vectorstore.filter.converter.PineconeFilterExpressionConverter; import org.springframework.util.Assert; @@ -57,7 +57,7 @@ public class PineconeVectorStore implements VectorStore { public final FilterExpressionConverter filterExpressionConverter = new PineconeFilterExpressionConverter(); - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; private final PineconeConnection pineconeConnection; @@ -211,13 +211,13 @@ public class PineconeVectorStore implements VectorStore { /** * Constructs a new PineconeVectorStore. * @param config The configuration for the store. - * @param embeddingClient The client for embedding operations. + * @param embeddingModel The client for embedding operations. */ - public PineconeVectorStore(PineconeVectorStoreConfig config, EmbeddingClient embeddingClient) { + public PineconeVectorStore(PineconeVectorStoreConfig config, EmbeddingModel embeddingModel) { Assert.notNull(config, "PineconeVectorStoreConfig must not be null"); - Assert.notNull(embeddingClient, "EmbeddingClient must not be null"); + Assert.notNull(embeddingModel, "EmbeddingModel must not be null"); - this.embeddingClient = embeddingClient; + this.embeddingModel = embeddingModel; this.pineconeNamespace = config.namespace; this.pineconeConnection = new PineconeClient(config.clientConfig).connect(config.connectionConfig); this.objectMapper = new ObjectMapper(); @@ -232,7 +232,7 @@ public class PineconeVectorStore implements VectorStore { List upsertVectors = documents.stream().map(document -> { // Compute and assign an embedding to the document. - document.setEmbedding(this.embeddingClient.embed(document)); + document.setEmbedding(this.embeddingModel.embed(document)); return Vector.newBuilder() .setId(document.getId()) @@ -321,7 +321,7 @@ public class PineconeVectorStore implements VectorStore { String nativeExpressionFilters = (request.getFilterExpression() != null) ? this.filterExpressionConverter.convertExpression(request.getFilterExpression()) : ""; - List queryEmbedding = this.embeddingClient.embed(request.getQuery()); + List queryEmbedding = this.embeddingModel.embed(request.getQuery()); var queryRequestBuilder = QueryRequest.newBuilder() .addAllVector(toFloatList(queryEmbedding)) diff --git a/vector-stores/spring-ai-pinecone-store/src/test/java/org/springframework/ai/vectorstore/PineconeVectorStoreIT.java b/vector-stores/spring-ai-pinecone-store/src/test/java/org/springframework/ai/vectorstore/PineconeVectorStoreIT.java index 9e8ef5f45..ca7336f12 100644 --- a/vector-stores/spring-ai-pinecone-store/src/test/java/org/springframework/ai/vectorstore/PineconeVectorStoreIT.java +++ b/vector-stores/spring-ai-pinecone-store/src/test/java/org/springframework/ai/vectorstore/PineconeVectorStoreIT.java @@ -30,8 +30,8 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.PineconeVectorStore.PineconeVectorStoreConfig; import org.springframework.boot.SpringBootConfiguration; import org.springframework.boot.autoconfigure.EnableAutoConfiguration; @@ -287,13 +287,13 @@ public class PineconeVectorStoreIT { } @Bean - public VectorStore vectorStore(PineconeVectorStoreConfig config, EmbeddingClient embeddingClient) { - return new PineconeVectorStore(config, embeddingClient); + public VectorStore vectorStore(PineconeVectorStoreConfig config, EmbeddingModel embeddingModel) { + return new PineconeVectorStore(config, embeddingModel); } @Bean - public TransformersEmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public TransformersEmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/vector-stores/spring-ai-qdrant-store/src/main/java/org/springframework/ai/vectorstore/qdrant/QdrantVectorStore.java b/vector-stores/spring-ai-qdrant-store/src/main/java/org/springframework/ai/vectorstore/qdrant/QdrantVectorStore.java index f9d3a250c..c8001044c 100644 --- a/vector-stores/spring-ai-qdrant-store/src/main/java/org/springframework/ai/vectorstore/qdrant/QdrantVectorStore.java +++ b/vector-stores/spring-ai-qdrant-store/src/main/java/org/springframework/ai/vectorstore/qdrant/QdrantVectorStore.java @@ -27,7 +27,7 @@ import java.util.UUID; import java.util.concurrent.ExecutionException; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.beans.factory.InitializingBean; @@ -61,7 +61,7 @@ public class QdrantVectorStore implements VectorStore, InitializingBean { public static final String DEFAULT_COLLECTION_NAME = "vector_store"; - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; private final QdrantClient qdrantClient; @@ -133,27 +133,26 @@ public class QdrantVectorStore implements VectorStore, InitializingBean { /** * Constructs a new QdrantVectorStore. * @param config The configuration for the store. - * @param embeddingClient The client for embedding operations. + * @param embeddingModel The client for embedding operations. * @deprecated since 1.0.0 in favor of {@link QdrantVectorStore}. */ @Deprecated(since = "1.0.0", forRemoval = true) - public QdrantVectorStore(QdrantClient qdrantClient, QdrantVectorStoreConfig config, - EmbeddingClient embeddingClient) { - this(qdrantClient, config.collectionName, embeddingClient); + public QdrantVectorStore(QdrantClient qdrantClient, QdrantVectorStoreConfig config, EmbeddingModel embeddingModel) { + this(qdrantClient, config.collectionName, embeddingModel); } /** * Constructs a new QdrantVectorStore. * @param qdrantClient A {@link QdrantClient} instance for interfacing with Qdrant. * @param collectionName The name of the collection to use in Qdrant. - * @param embeddingClient The client for embedding operations. + * @param embeddingModel The client for embedding operations. */ - public QdrantVectorStore(QdrantClient qdrantClient, String collectionName, EmbeddingClient embeddingClient) { + public QdrantVectorStore(QdrantClient qdrantClient, String collectionName, EmbeddingModel embeddingModel) { Assert.notNull(qdrantClient, "QdrantClient must not be null"); Assert.notNull(collectionName, "collectionName must not be null"); - Assert.notNull(embeddingClient, "EmbeddingClient must not be null"); + Assert.notNull(embeddingModel, "EmbeddingModel must not be null"); - this.embeddingClient = embeddingClient; + this.embeddingModel = embeddingModel; this.collectionName = collectionName; this.qdrantClient = qdrantClient; } @@ -167,7 +166,7 @@ public class QdrantVectorStore implements VectorStore, InitializingBean { try { List points = documents.stream().map(document -> { // Compute and assign an embedding to the document. - document.setEmbedding(this.embeddingClient.embed(document)); + document.setEmbedding(this.embeddingModel.embed(document)); return PointStruct.newBuilder() .setId(id(UUID.fromString(document.getId()))) @@ -215,7 +214,7 @@ public class QdrantVectorStore implements VectorStore, InitializingBean { ? this.filterExpressionConverter.convertExpression(request.getFilterExpression()) : Filter.getDefaultInstance(); - List queryEmbedding = this.embeddingClient.embed(request.getQuery()); + List queryEmbedding = this.embeddingModel.embed(request.getQuery()); var searchPoints = SearchPoints.newBuilder() .setCollectionName(this.collectionName) @@ -290,7 +289,7 @@ public class QdrantVectorStore implements VectorStore, InitializingBean { if (!isCollectionExists()) { var vectorParams = VectorParams.newBuilder() .setDistance(Distance.Cosine) - .setSize(this.embeddingClient.dimensions()) + .setSize(this.embeddingModel.dimensions()) .build(); this.qdrantClient.createCollectionAsync(this.collectionName, vectorParams).get(); } diff --git a/vector-stores/spring-ai-qdrant-store/src/test/java/org/springframework/ai/vectorstore/qdrant/QdrantVectorStoreIT.java b/vector-stores/spring-ai-qdrant-store/src/test/java/org/springframework/ai/vectorstore/qdrant/QdrantVectorStoreIT.java index 73e6e9327..bdad3ff84 100644 --- a/vector-stores/spring-ai-qdrant-store/src/test/java/org/springframework/ai/vectorstore/qdrant/QdrantVectorStoreIT.java +++ b/vector-stores/spring-ai-qdrant-store/src/test/java/org/springframework/ai/vectorstore/qdrant/QdrantVectorStoreIT.java @@ -35,9 +35,9 @@ import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.testcontainers.qdrant.QdrantContainer; -import org.springframework.ai.azure.openai.AzureOpenAiEmbeddingClient; +import org.springframework.ai.azure.openai.AzureOpenAiEmbeddingModel; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.SearchRequest; import org.springframework.ai.vectorstore.VectorStore; import org.springframework.boot.SpringBootConfiguration; @@ -250,8 +250,8 @@ public class QdrantVectorStoreIT { } @Bean - public VectorStore qdrantVectorStore(EmbeddingClient embeddingClient, QdrantClient qdrantClient) { - return new QdrantVectorStore(qdrantClient, COLLECTION_NAME, embeddingClient); + public VectorStore qdrantVectorStore(EmbeddingModel embeddingModel, QdrantClient qdrantClient) { + return new QdrantVectorStore(qdrantClient, COLLECTION_NAME, embeddingModel); } @Bean @@ -262,8 +262,8 @@ public class QdrantVectorStoreIT { } @Bean - public AzureOpenAiEmbeddingClient azureEmbeddingClient(OpenAIClient openAIClient) { - return new AzureOpenAiEmbeddingClient(openAIClient); + public AzureOpenAiEmbeddingModel azureEmbeddingModel(OpenAIClient openAIClient) { + return new AzureOpenAiEmbeddingModel(openAIClient); } } diff --git a/vector-stores/spring-ai-redis-store/src/main/java/org/springframework/ai/vectorstore/RedisVectorStore.java b/vector-stores/spring-ai-redis-store/src/main/java/org/springframework/ai/vectorstore/RedisVectorStore.java index cb34e3b8b..6680070ec 100644 --- a/vector-stores/spring-ai-redis-store/src/main/java/org/springframework/ai/vectorstore/RedisVectorStore.java +++ b/vector-stores/spring-ai-redis-store/src/main/java/org/springframework/ai/vectorstore/RedisVectorStore.java @@ -29,7 +29,7 @@ import java.util.stream.Collectors; import org.slf4j.Logger; import org.slf4j.LoggerFactory; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.filter.FilterExpressionConverter; import org.springframework.beans.factory.InitializingBean; import org.springframework.util.Assert; @@ -63,14 +63,14 @@ import redis.clients.jedis.search.schemafields.VectorField.VectorAlgorithm; * * This class requires a RedisVectorStoreConfig configuration object for initialization, * which includes settings like Redis URI, index name, field names, and vector algorithms. - * It also requires an EmbeddingClient to convert documents into embeddings before storing + * It also requires an EmbeddingModel to convert documents into embeddings before storing * them. * * @author Julien Ruaux * @author Christian Tzolov * @see VectorStore * @see RedisVectorStoreConfig - * @see EmbeddingClient + * @see EmbeddingModel */ public class RedisVectorStore implements VectorStore, InitializingBean { @@ -280,19 +280,19 @@ public class RedisVectorStore implements VectorStore, InitializingBean { private final JedisPooled jedis; - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; private final RedisVectorStoreConfig config; private FilterExpressionConverter filterExpressionConverter; - public RedisVectorStore(RedisVectorStoreConfig config, EmbeddingClient embeddingClient) { + public RedisVectorStore(RedisVectorStoreConfig config, EmbeddingModel embeddingModel) { Assert.notNull(config, "Config must not be null"); - Assert.notNull(embeddingClient, "Embedding client must not be null"); + Assert.notNull(embeddingModel, "Embedding client must not be null"); this.jedis = new JedisPooled(config.uri); - this.embeddingClient = embeddingClient; + this.embeddingModel = embeddingModel; this.config = config; this.filterExpressionConverter = new RedisFilterExpressionConverter(this.config.metadataFields); } @@ -305,7 +305,7 @@ public class RedisVectorStore implements VectorStore, InitializingBean { public void add(List documents) { try (Pipeline pipeline = this.jedis.pipelined()) { for (Document document : documents) { - var embedding = this.embeddingClient.embed(document); + var embedding = this.embeddingModel.embed(document); document.setEmbedding(embedding); var fields = new HashMap(); @@ -365,7 +365,7 @@ public class RedisVectorStore implements VectorStore, InitializingBean { returnFields.add(this.config.embeddingFieldName); returnFields.add(this.config.contentFieldName); returnFields.add(DISTANCE_FIELD_NAME); - var embedding = toFloatArray(this.embeddingClient.embed(request.getQuery())); + var embedding = toFloatArray(this.embeddingModel.embed(request.getQuery())); Query query = new Query(queryString).addParam(EMBEDDING_PARAM_NAME, RediSearchUtil.toByteArray(embedding)) .returnFields(returnFields.toArray(new String[0])) .setSortBy(DISTANCE_FIELD_NAME, true) @@ -420,7 +420,7 @@ public class RedisVectorStore implements VectorStore, InitializingBean { private Iterable schemaFields() { Map vectorAttrs = new HashMap<>(); - vectorAttrs.put("DIM", this.embeddingClient.dimensions()); + vectorAttrs.put("DIM", this.embeddingModel.dimensions()); vectorAttrs.put("DISTANCE_METRIC", DEFAULT_DISTANCE_METRIC); vectorAttrs.put("TYPE", VECTOR_TYPE_FLOAT32); List fields = new ArrayList<>(); diff --git a/vector-stores/spring-ai-redis-store/src/test/java/org/springframework/ai/vectorstore/RedisVectorStoreIT.java b/vector-stores/spring-ai-redis-store/src/test/java/org/springframework/ai/vectorstore/RedisVectorStoreIT.java index 66d58fdf3..cce1808a0 100644 --- a/vector-stores/spring-ai-redis-store/src/test/java/org/springframework/ai/vectorstore/RedisVectorStoreIT.java +++ b/vector-stores/spring-ai-redis-store/src/test/java/org/springframework/ai/vectorstore/RedisVectorStoreIT.java @@ -27,8 +27,8 @@ import java.util.UUID; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.RedisVectorStore.MetadataField; import org.springframework.ai.vectorstore.RedisVectorStore.RedisVectorStoreConfig; import org.springframework.boot.SpringBootConfiguration; @@ -245,17 +245,17 @@ class RedisVectorStoreIT { public static class TestApplication { @Bean - public RedisVectorStore vectorStore(EmbeddingClient embeddingClient) { + public RedisVectorStore vectorStore(EmbeddingModel embeddingModel) { return new RedisVectorStore(RedisVectorStoreConfig.builder() .withURI(redisContainer.getRedisURI()) .withMetadataFields(MetadataField.tag("meta1"), MetadataField.tag("meta2"), MetadataField.tag("country"), MetadataField.numeric("year")) - .build(), embeddingClient); + .build(), embeddingModel); } @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } } diff --git a/vector-stores/spring-ai-weaviate-store/src/main/java/org/springframework/ai/vectorstore/WeaviateVectorStore.java b/vector-stores/spring-ai-weaviate-store/src/main/java/org/springframework/ai/vectorstore/WeaviateVectorStore.java index db676aadd..0f73fad8f 100644 --- a/vector-stores/spring-ai-weaviate-store/src/main/java/org/springframework/ai/vectorstore/WeaviateVectorStore.java +++ b/vector-stores/spring-ai-weaviate-store/src/main/java/org/springframework/ai/vectorstore/WeaviateVectorStore.java @@ -43,7 +43,7 @@ import io.weaviate.client.v1.graphql.query.fields.Field; import io.weaviate.client.v1.graphql.query.fields.Fields; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig.ConsistentLevel; import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig.MetadataField; import org.springframework.beans.factory.InitializingBean; @@ -80,7 +80,7 @@ public class WeaviateVectorStore implements VectorStore, InitializingBean { private static final String ADDITIONAL_VECTOR_FIELD_NAME = "vector"; - private final EmbeddingClient embeddingClient; + private final EmbeddingModel embeddingModel; private final WeaviateClient weaviateClient; @@ -277,14 +277,14 @@ public class WeaviateVectorStore implements VectorStore, InitializingBean { /** * Constructs a new WeaviateVectorStore. * @param vectorStoreConfig The configuration for the store. - * @param embeddingClient The client for embedding operations. + * @param embeddingModel The client for embedding operations. */ - public WeaviateVectorStore(WeaviateVectorStoreConfig vectorStoreConfig, EmbeddingClient embeddingClient, + public WeaviateVectorStore(WeaviateVectorStoreConfig vectorStoreConfig, EmbeddingModel embeddingModel, WeaviateClient weaviateClient) { Assert.notNull(vectorStoreConfig, "WeaviateVectorStoreConfig must not be null"); - Assert.notNull(embeddingClient, "EmbeddingClient must not be null"); + Assert.notNull(embeddingModel, "EmbeddingModel must not be null"); - this.embeddingClient = embeddingClient; + this.embeddingModel = embeddingModel; this.consistencyLevel = vectorStoreConfig.consistencyLevel; this.weaviateObjectClass = vectorStoreConfig.weaviateObjectClass; this.filterMetadataFields = vectorStoreConfig.filterMetadataFields; @@ -360,7 +360,7 @@ public class WeaviateVectorStore implements VectorStore, InitializingBean { private WeaviateObject toWeaviateObject(Document document) { if (CollectionUtils.isEmpty(document.getEmbedding())) { - List embedding = this.embeddingClient.embed(document); + List embedding = this.embeddingModel.embed(document); document.setEmbedding(embedding); } @@ -420,7 +420,7 @@ public class WeaviateVectorStore implements VectorStore, InitializingBean { @Override public List similaritySearch(SearchRequest request) { - Float[] embedding = toFloatArray(this.embeddingClient.embed(request.getQuery())); + Float[] embedding = toFloatArray(this.embeddingModel.embed(request.getQuery())); GetBuilder.GetBuilderBuilder builder = GetBuilder.builder(); diff --git a/vector-stores/spring-ai-weaviate-store/src/test/java/org/springframework/ai/vectorstore/WeaviateVectorStoreIT.java b/vector-stores/spring-ai-weaviate-store/src/test/java/org/springframework/ai/vectorstore/WeaviateVectorStoreIT.java index 013ebfaa4..1cd18ae4c 100644 --- a/vector-stores/spring-ai-weaviate-store/src/test/java/org/springframework/ai/vectorstore/WeaviateVectorStoreIT.java +++ b/vector-stores/spring-ai-weaviate-store/src/test/java/org/springframework/ai/vectorstore/WeaviateVectorStoreIT.java @@ -29,8 +29,8 @@ import org.testcontainers.junit.jupiter.Container; import org.testcontainers.junit.jupiter.Testcontainers; import org.springframework.ai.document.Document; -import org.springframework.ai.embedding.EmbeddingClient; -import org.springframework.ai.transformers.TransformersEmbeddingClient; +import org.springframework.ai.embedding.EmbeddingModel; +import org.springframework.ai.transformers.TransformersEmbeddingModel; import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig; import org.springframework.ai.vectorstore.WeaviateVectorStore.WeaviateVectorStoreConfig.MetadataField; import org.springframework.boot.SpringBootConfiguration; @@ -243,7 +243,7 @@ public class WeaviateVectorStoreIT { public static class TestApplication { @Bean - public VectorStore vectorStore(EmbeddingClient embeddingClient) { + public VectorStore vectorStore(EmbeddingModel embeddingModel) { WeaviateClient weaviateClient = new WeaviateClient( new Config("http", weaviateContainer.getHttpHostAddress())); @@ -252,14 +252,14 @@ public class WeaviateVectorStoreIT { .withConsistencyLevel(WeaviateVectorStoreConfig.ConsistentLevel.ONE) .build(); - WeaviateVectorStore vectorStore = new WeaviateVectorStore(config, embeddingClient, weaviateClient); + WeaviateVectorStore vectorStore = new WeaviateVectorStore(config, embeddingModel, weaviateClient); return vectorStore; } @Bean - public EmbeddingClient embeddingClient() { - return new TransformersEmbeddingClient(); + public EmbeddingModel embeddingModel() { + return new TransformersEmbeddingModel(); } }