diff --git a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatConnector.java b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatConnector.java index c05ee47a2..4ce4bb674 100644 --- a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatConnector.java +++ b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatConnector.java @@ -117,7 +117,7 @@ public class AnthropicChatConnector extends * @param retryTemplate the retry template used to retry the Anthropic API calls. */ public AnthropicChatConnector(AnthropicApi anthropicApi, AnthropicChatOptions defaultOptions, - RetryTemplate retryTemplate) { + RetryTemplate retryTemplate) { this(anthropicApi, defaultOptions, retryTemplate, null); } @@ -130,7 +130,7 @@ public class AnthropicChatConnector extends * state of the function calls. */ public AnthropicChatConnector(AnthropicApi anthropicApi, AnthropicChatOptions defaultOptions, - RetryTemplate retryTemplate, FunctionCallbackContext functionCallbackContext) { + RetryTemplate retryTemplate, FunctionCallbackContext functionCallbackContext) { super(functionCallbackContext); diff --git a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatConnector.java b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatConnector.java index 70fb062e1..5708a5963 100644 --- a/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatConnector.java +++ b/models/spring-ai-azure-openai/src/main/java/org/springframework/ai/azure/openai/AzureOpenAiChatConnector.java @@ -107,7 +107,7 @@ public class AzureOpenAiChatConnector } public AzureOpenAiChatConnector(OpenAIClient microsoftOpenAiClient, AzureOpenAiChatOptions options, - FunctionCallbackContext functionCallbackContext) { + FunctionCallbackContext functionCallbackContext) { super(functionCallbackContext); Assert.notNull(microsoftOpenAiClient, "com.azure.ai.openai.OpenAIClient must not be null"); Assert.notNull(options, "AzureOpenAiChatOptions must not be null"); 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 3ec66d35b..f96328b4b 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,7 +59,7 @@ public class MockAzureOpenAiTestConfiguration { } @Bean - AzureOpenAiChatConnector azureOpenAiChatClient(OpenAIClient microsoftAzureOpenAiClient) { + AzureOpenAiChatConnector azureOpenAiChatClient(OpenAIClient microsoftAzureOpenAiClient) { return new AzureOpenAiChatConnector(microsoftAzureOpenAiClient); } diff --git a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatConnector.java b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatConnector.java index 079629ce2..d7a0268b0 100644 --- a/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatConnector.java +++ b/models/spring-ai-bedrock/src/main/java/org/springframework/ai/bedrock/jurassic2/BedrockAi21Jurassic2ChatConnector.java @@ -41,7 +41,7 @@ public class BedrockAi21Jurassic2ChatConnector implements ChatConnector { private final BedrockAi21Jurassic2ChatOptions defaultOptions; public BedrockAi21Jurassic2ChatConnector(Ai21Jurassic2ChatBedrockApi chatApi, - BedrockAi21Jurassic2ChatOptions options) { + BedrockAi21Jurassic2ChatOptions options) { Assert.notNull(chatApi, "Ai21Jurassic2ChatBedrockApi must not be null"); Assert.notNull(options, "BedrockAi21Jurassic2ChatOptions must not be null"); diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatConnector.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatConnector.java index 082eee1e5..54f058896 100644 --- a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatConnector.java +++ b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatConnector.java @@ -83,7 +83,7 @@ public class MistralAiChatConnector extends } public MistralAiChatConnector(MistralAiApi mistralAiApi, MistralAiChatOptions options, - FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate) { + FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate) { super(functionCallbackContext); Assert.notNull(mistralAiApi, "MistralAiApi must not be null"); Assert.notNull(options, "Options must not be null"); diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatConnector.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatConnector.java index ba76c1d52..b71f8e801 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatConnector.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatConnector.java @@ -123,7 +123,7 @@ public class OpenAiChatConnector extends * @param retryTemplate The retry template. */ public OpenAiChatConnector(OpenAiApi openAiApi, OpenAiChatOptions options, - FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate) { + FunctionCallbackContext functionCallbackContext, RetryTemplate retryTemplate) { super(functionCallbackContext); Assert.notNull(openAiApi, "OpenAiApi must not be null"); Assert.notNull(options, "Options must not be null"); diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/ChatConnectorTest.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/ChatConnectorTest.java deleted file mode 100644 index 2193c18d3..000000000 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/ChatConnectorTest.java +++ /dev/null @@ -1,56 +0,0 @@ -package org.springframework.ai.openai; - -import org.junit.jupiter.api.Test; -import org.springframework.ai.chat.ChatClient; -import org.springframework.beans.factory.annotation.Autowired; -import org.springframework.context.annotation.Bean; -import org.springframework.context.annotation.Configuration; - -import java.util.Map; - -class ChatConnectorTest { - - @Configuration - static class ChatClientTestConfiguration { - - @Bean - ChatClient client(OpenAiChatConnector openAiChatConnector) { - return ChatClient.builder(openAiChatConnector).defaultSystem(""" - you are customer service agent designed to answer questions - about a the user, {userName}'s, orders. Here are their outstanding orders. - - {orders} - - """).defaultFunctions("cancelOrder", "refundOrder").build(); - } - - } - - private final ChatClient singularity; - - ChatConnectorTest(@Autowired ChatClient singularity) { - this.singularity = singularity; - } - - @Test - void products() throws Exception { - var product0 = this.client.userPrompt("tell me about this product from the merchant {merchant}") - .userPromptParams(Map.of("merchant", "24u92")) - .execute(Product.class); - - /* - * var product1 = this.client .build() .userPromptParam("a", "b") - * .functions("cancelOrder", "refundOrder") .execute(new - * ParameterizedTypeReference() { }); - * - * var product2 = this.client - * .userPrompt("tell me about this product from the merchant {merchant}", - * Map.of("merchant", "232")) .execute(Product.class); - */ - - } - - record Product(String sku) { - } - -} 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 b2e3e3a30..cdac0c599 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 @@ -99,7 +99,7 @@ public class ChatMemoryLongTermSystemPromptIT extends BaseMemoryTest { @Bean public ChatService memoryChatService(OpenAiChatConnector chatClient, VectorStore vectorStore, - TokenCountEstimator tokenCountEstimator) { + TokenCountEstimator tokenCountEstimator) { return PromptTransformingChatService.builder(chatClient) .withRetrievers(List.of(new VectorStoreChatMemoryRetriever(vectorStore, 10))) @@ -111,7 +111,7 @@ public class ChatMemoryLongTermSystemPromptIT extends BaseMemoryTest { @Bean public StreamingChatService memoryStreamingChatService(OpenAiChatConnector streamingChatClient, - VectorStore vectorStore, TokenCountEstimator tokenCountEstimator) { + VectorStore vectorStore, TokenCountEstimator tokenCountEstimator) { return StreamingPromptTransformingChatService.builder(streamingChatClient) .withRetrievers(List.of(new VectorStoreChatMemoryRetriever(vectorStore, 10))) 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 e07516e11..2bf642da2 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 @@ -75,7 +75,7 @@ public class ChatMemoryShortTermMessageListIT extends BaseMemoryTest { @Bean public ChatService memoryChatService(OpenAiChatConnector chatClient, ChatMemory chatHistory, - TokenCountEstimator tokenCountEstimator) { + TokenCountEstimator tokenCountEstimator) { return PromptTransformingChatService.builder(chatClient) .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) @@ -87,7 +87,7 @@ public class ChatMemoryShortTermMessageListIT extends BaseMemoryTest { @Bean public StreamingChatService memoryStreamingChatService(OpenAiChatConnector streamingChatClient, - ChatMemory chatHistory, TokenCountEstimator tokenCountEstimator) { + ChatMemory chatHistory, TokenCountEstimator tokenCountEstimator) { return StreamingPromptTransformingChatService.builder(streamingChatClient) .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) 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 30a2e299e..ea9707848 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 @@ -76,7 +76,7 @@ public class ChatMemoryShortTermSystemPromptIT extends BaseMemoryTest { @Bean public ChatService memoryChatService(OpenAiChatConnector chatClient, ChatMemory chatHistory, - TokenCountEstimator tokenCountEstimator) { + TokenCountEstimator tokenCountEstimator) { return PromptTransformingChatService.builder(chatClient) .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) @@ -88,7 +88,7 @@ public class ChatMemoryShortTermSystemPromptIT extends BaseMemoryTest { @Bean public StreamingChatService memoryStreamingChatService(OpenAiChatConnector streamingChatClient, - ChatMemory chatHistory, TokenCountEstimator tokenCountEstimator) { + ChatMemory chatHistory, TokenCountEstimator tokenCountEstimator) { return StreamingPromptTransformingChatService.builder(streamingChatClient) .withRetrievers(List.of(new ChatMemoryRetriever(chatHistory))) 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 d74ef885b..9bb027c90 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 @@ -187,7 +187,7 @@ public class LongShortTermChatMemoryWithRagIT { @Bean public ChatService memoryChatService(OpenAiChatConnector chatClient, VectorStore vectorStore, - TokenCountEstimator tokenCountEstimator, ChatMemory chatHistory) { + TokenCountEstimator tokenCountEstimator, ChatMemory chatHistory) { return PromptTransformingChatService.builder(chatClient) .withRetrievers(List.of(new VectorStoreRetriever(vectorStore, SearchRequest.defaults()), 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 24db3e336..18c00ad32 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 @@ -82,7 +82,7 @@ public class OpenAiPromptTransformingChatServiceIT { @Autowired public OpenAiPromptTransformingChatServiceIT(ChatConnector chatConnector, ChatService chatService, - VectorStore vectorStore) { + VectorStore vectorStore) { this.chatConnector = chatConnector; this.chatService = chatService; this.vectorStore = vectorStore; diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatConnector.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatConnector.java index 1fa8d9df3..8e9affea0 100644 --- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatConnector.java +++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatConnector.java @@ -130,7 +130,7 @@ public class VertexAiGeminiChatConnector } public VertexAiGeminiChatConnector(VertexAI vertexAI, VertexAiGeminiChatOptions options, - FunctionCallbackContext functionCallbackContext) { + FunctionCallbackContext functionCallbackContext) { super(functionCallbackContext); 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 ae0254179..4b5850e8d 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 @@ -3,14 +3,23 @@ package org.springframework.ai.chat; import org.springframework.ai.chat.connector.ChatConnector; 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.chat.prompt.ChatOptions; import org.springframework.ai.chat.prompt.Prompt; +import org.springframework.ai.chat.prompt.PromptTemplate; +import org.springframework.ai.model.function.FunctionCallback; +import org.springframework.ai.model.function.FunctionCallbackWrapper; +import org.springframework.ai.model.function.FunctionCallingOptionsBuilder; import org.springframework.core.ParameterizedTypeReference; import org.springframework.core.io.Resource; +import org.springframework.util.Assert; import org.springframework.util.MimeType; import reactor.core.publisher.Flux; +import java.io.IOException; import java.net.URL; +import java.nio.charset.Charset; import java.util.*; import java.util.function.Consumer; @@ -93,229 +102,312 @@ public class DemoApplication { */ public interface ChatClient { - static ChatClientBuilder builder(ChatConnector connector) { - return new ChatClientBuilder(connector); - } - ChatClientRequest build(); + static ChatClientBuilder builder(ChatConnector connector) { + return new ChatClientBuilder(connector); + } - ChatResponse call(Prompt prompt); + ChatResponse call(Prompt prompt); - ChatClientRequest user(Consumer consumer); + ChatClientRequest call(); - public static class UserSpec { + interface PromptSpec { + T text(String text); - public UserSpec media(List media) { - return this; - } + T text(Resource text, Charset charset); - public UserSpec media(URL url, MimeType mimeType) { - return this; - } + T text(Resource text); - public UserSpec media(Resource resource, MimeType type) { - return this; - } + T params(Map p); - public UserSpec media(Media... m) { - return this; - } + T param(String k, String v); - public UserSpec params(Map p) { - return this; - } + } - public UserSpec param(String k, String v) { - return this; - } - } + abstract class AbstractPromptSpec> implements PromptSpec { - public static class ChatClientRequest { + private String text = ""; - private String userPrompt = ""; + private final Map params = new HashMap<>(); - private String systemPrompt = ""; + @Override + public T text(String text) { + this.text = (text); + return self(); + } - private final List media = new ArrayList<>(); + @Override + public T text(Resource text, Charset charset) { + try { + this.text(text.getContentAsString(charset)); + } + catch (IOException e) { + throw new RuntimeException(e); + } + return self(); + } - private final List functions = new ArrayList<>(); + @Override + public T text(Resource text) { + this.text(text, Charset.defaultCharset()); + return self(); + } - private final Map userPromptParams = new HashMap<>(); + @Override + public T param(String k, String v) { + this.params.put(k, v); + return self(); + } - private final Map systemPromptParams = new HashMap<>(); + @Override + public T params(Map p) { + this.params.putAll(p); + return self(); + } - List userMedia() { - return this.media; - } + protected abstract T self(); - String systemText() { - return this.systemPrompt; - } + protected String text() { + return this.text; + } - String userText() { - return this.userPrompt; - } + protected Map params() { + return this.params; + } - List functions() { - return this.functions; - } + } - public ChatClientRequest(String userPrompt, String systemPrompt, List functions, List media) { - this.userPrompt = userPrompt; - this.systemPrompt = systemPrompt; - this.functions.addAll(functions); - this.media.addAll(media); - } + class UserSpec extends AbstractPromptSpec implements PromptSpec { - public ChatClientRequest messages(Message... messages) { - return null; - } -// -// public ChatClientRequest userParam(String key, String value) { -// this.userPromptParams.put(key, value); -// return this; -// } -// -// public ChatClientRequest systemParam(String key, String value) { -// this.systemPromptParams.put(key, value); -// return this; -// } + private final List media = new ArrayList<>(); - public ChatClientRequest options(T options) { - return this; - } -// -// public ChatClientRequest systemParams(Map systemPromptParams) { -// this.systemPromptParams.putAll(systemPromptParams); -// return this; -// } -// -// public ChatClientRequest userParams(Map userPromptParams) { -// this.userPromptParams.putAll(userPromptParams); -// return this; -// } -// -// public ChatClientRequest userText(Resource resource) { -// return userText(resource, Charset.defaultCharset()); -// } + public UserSpec media(Media... media) { + this.media.addAll(Arrays.asList(media)); + return self(); + } -// public ChatClientRequest userText(Resource resource, Charset charset) { -// try { -// this.userText(resource.getContentAsString(charset)); -// } catch (IOException e) { -// throw new RuntimeException(e); -// } -// return this; -// } -// -// -// public ChatClientRequest userText(String userPrompt) { -// this.userPrompt = userPrompt; -// return this; -// } -// -// public ChatClientRequest systemText(Resource systemPrompt) { -// return systemText(systemPrompt, Charset.defaultCharset()); -// } -// -// public ChatClientRequest systemText(Resource systemPrompt, Charset charset) { -// try { -// this.systemText(systemPrompt.getContentAsString(charset)); -// } catch (IOException e) { -// throw new RuntimeException(e); -// } -// return this; -// } -// -// public ChatClientRequest systemText(String systemPrompt) { -// this.systemPrompt = systemPrompt; -// return this; -// } -// -// public ChatClientRequest userMedia(Media... media) { -// this.media.addAll(Arrays.asList(media)); -// return this; -// } + public UserSpec media(MimeType mimeType, URL url) { + this.media.add(new Media(mimeType, url)); + return self(); + } - public ChatClientRequest functions(String... functions) { - this.functions.addAll(Arrays.asList(functions)); - return this; - } + public UserSpec media(MimeType mimeType, Resource resource) { + this.media.add(new Media(mimeType, resource)); + return self(); + } + protected List media() { + return this.media; + } - public static class ChatResponseSpec { + @Override + protected UserSpec self() { + return this; + } - public T single(ParameterizedTypeReference t) { - return null; - } + } + class SystemSpec extends AbstractPromptSpec implements PromptSpec { - public T single(Class clzz) { - return null; - } + @Override + protected SystemSpec self() { + return this; + } - public ChatResponse chatResponse() { - return null; - } + } - public Flux stream(Class t) { - return null; - } + class ChatClientRequest { - public Flux stream(ParameterizedTypeReference t) { - return Flux.empty(); - } + private final ChatConnector connector; - public Collection list(Class clzz) { - return null; - } + private String userText = ""; - public Collection list(ParameterizedTypeReference> ptr) { - return List.of(); - } + private String systemText = ""; - } + private ChatOptions chatOptions; - public ChatResponseSpec chat() { - return null; - } + private final List media = new ArrayList<>(); + private final Set functionNames = new HashSet<>(); - } + private final List functionCallbacks = new ArrayList<>(); - public static class ChatClientBuilder { + private final Map userParams = new HashMap<>(); - private final ChatConnector connector; + private final List messages = new ArrayList<>(); - private final List defaultMedia = new ArrayList<>(); + private final Map systemParams = new HashMap<>(); - private final List defaultFunctions = new ArrayList<>(); + public ChatClientRequest(ChatConnector connector, String userText, String systemText, + List functionNames, List media, ChatOptions chatOptions) { + this.userText = userText; + this.systemText = systemText; + this.connector = connector; + this.functionNames.addAll(functionNames); + this.media.addAll(media); + this.chatOptions = chatOptions; + } - private String defaultSystemPrompt; + public ChatClientRequest messages(Message... messages) { + this.messages.addAll(List.of(messages)); + return this; + } - private String defaultUserPrompt; + public ChatClientRequest options(T options) { + this.chatOptions = options; + return this; + } - ChatClientBuilder(ChatConnector connector) { - this.connector = connector; - } + 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 ChatClient build() { - return new DefaultChatClient(this.connector, this.defaultSystemPrompt, this.defaultUserPrompt, - this.defaultFunctions, this.defaultMedia); - } + public ChatClientRequest functions(String... functions) { + this.functionNames.addAll(List.of(functions)); + return this; + } - public ChatClientBuilder defaultSystem(String systemPrompt) { - return this; - } + public ChatClientRequest system(Consumer consumer) { + var ss = new SystemSpec(); + consumer.accept(ss); + this.systemText = ss.text(); + this.systemParams.putAll(ss.params()); + return this; + } - public ChatClientBuilder defaultFunctions(String... functionNames) { - return this; - } + public ChatClientRequest user(Consumer consumer) { + var us = new UserSpec(); + consumer.accept(us); + this.userText = us.text(); + this.userParams.putAll(us.params()); + this.media.addAll(us.media()); + return this; + } - public ChatClientBuilder defaultUserPrompt(String userPrompt) { - return this; - } + public static class ChatResponseSpec { + + private final ChatClientRequest request; + + private final ChatConnector chatConnector; + + public ChatResponseSpec(ChatConnector chatConnector, ChatClientRequest request) { + this.chatConnector = chatConnector; + this.request = request; + } + + public T single(ParameterizedTypeReference t) { + return null; + } + + public T single(Class clzz) { + return null; + } + + public ChatResponse chatResponse() { + + var userMessage = new UserMessage( + new PromptTemplate(this.request.userText, this.request.userParams).render(), + this.request.media); + + var systemMessage = new SystemMessage( + new PromptTemplate(this.request.systemText, this.request.systemParams).render()); + + if (request.chatOptions instanceof FunctionCallingOptionsBuilder.PortableFunctionCallingOptions functionCallingOptions) { + if (!request.functionNames.isEmpty()) { + functionCallingOptions.setFunctions(request.functionNames); + } + if (!request.functionCallbacks.isEmpty()) { + functionCallingOptions.setFunctionCallbacks(request.functionCallbacks); + } + } + var prompt = new Prompt(List.of(systemMessage, userMessage), request.chatOptions); + + return this.chatConnector.call(prompt); + } + + public Flux stream(Class t) { + return null; + } + + public Flux stream(ParameterizedTypeReference t) { + return Flux.empty(); + } + + public Collection list(Class clzz) { + return null; + } + + public Collection list(ParameterizedTypeReference> ptr) { + return List.of(); + } + + } + + public ChatResponseSpec chat() { + return new ChatResponseSpec(this.connector, this); + } + + } + + class ChatClientBuilder { + + private final ChatConnector connector; + + private final List defaultMedia = new ArrayList<>(); + + private final List defaultFunctions = new ArrayList<>(); + + private String defaultSystem; + + private String defaultUser; + + ChatClientBuilder(ChatConnector connector) { + Assert.notNull(connector, "the " + ChatConnector.class.getName() + " must be non-null!"); + this.connector = connector; + } + + public ChatClient build() { + return new DefaultChatClient(this.connector, this.defaultSystem, this.defaultUser, this.defaultFunctions, + this.defaultMedia); + } + + public ChatClientBuilder defaultSystem(String systemPrompt) { + this.defaultSystem = systemPrompt; + return this; + } + + public ChatClientBuilder defaultFunctions(String... functionNames) { + this.defaultFunctions.addAll(List.of(functionNames)); + return this; + } + + public ChatClientBuilder defaultUser(String userPrompt) { + this.defaultUser = userPrompt; + return this; + } + + } + + @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(); + return (generation != null) ? generation.getOutput().getContent() : ""; + } + + @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(); + return (generation != null) ? generation.getOutput().getContent() : ""; + } - } } 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 index d265ec662..fefd88907 100644 --- 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 @@ -5,52 +5,48 @@ import org.springframework.ai.chat.messages.Media; import org.springframework.ai.chat.prompt.Prompt; import java.util.List; -import java.util.function.Consumer; - /** - * todo follow WebClient -> DefaultWebClient - * todo make sure ChatConnector also supports call(Prompt) and then mark as deprecated - * * @author Mark Pollack * @author Christian Tzolov * @author Josh Long * @author Arjen Poutsma */ -public class DefaultChatClient implements ChatClient { +class DefaultChatClient implements ChatClient { private final ChatConnector connector; - private final String userPrompt, systemPrompt; + private final String userText, systemText; - private final List functions; + private final List functionNames; private final List media; public DefaultChatClient(ChatConnector connector, String defaultSystemPrompt, String defaultUserPrompt, - List defaultFunctions, List defaultMedia) { + List defaultFunctions, List defaultMedia) { this.connector = connector; - this.userPrompt = defaultUserPrompt; - this.systemPrompt = defaultSystemPrompt; - this.functions = defaultFunctions; + this.userText = defaultUserPrompt; + this.systemText = defaultSystemPrompt; + this.functionNames = defaultFunctions; this.media = defaultMedia; } @Override - public ChatClientRequest build() { - return new ChatClientRequest(this.userPrompt, this.systemPrompt, this.functions, this.media); + public ChatClientRequest call() { + return new ChatClientRequest(this.connector, this.userText, this.systemText, this.functionNames, this.media, + null); } + /** + * use the new fluid DSL starting in {@link #call()} + * @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 null; - } - - - @Override - public ChatClientRequest user(Consumer consumer) { - return null; + return this.connector.call(prompt); } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/connector/ChatConnector.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/connector/ChatConnector.java index b83ae1b95..846e29413 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/connector/ChatConnector.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/connector/ChatConnector.java @@ -16,19 +16,14 @@ package org.springframework.ai.chat.connector; import org.springframework.ai.chat.ChatResponse; +import org.springframework.ai.chat.Generation; +import org.springframework.ai.chat.messages.Message; +import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; -public interface ChatConnector { +import java.util.Arrays; - /* - * default String call(String message) { Prompt prompt = new Prompt(new - * UserMessage(message)); Generation generation = call(prompt).getResult(); return - * (generation != null) ? generation.getOutput().getContent() : ""; } - * - * public String call(Message... messages) { Prompt prompt = new - * Prompt(Arrays.asList(messages)); Generation generation = call(prompt).getResult(); - * return (generation != null) ? generation.getOutput().getContent() : ""; } - */ +public interface ChatConnector { ChatResponse call(Prompt prompt); diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AbstractMessage.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AbstractMessage.java index b58cb3c21..3b1d51cd4 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AbstractMessage.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/messages/AbstractMessage.java @@ -66,9 +66,9 @@ public abstract class AbstractMessage implements Message { protected AbstractMessage(MessageType messageType, String textContent, Collection media, Map metadata) { - Assert.notNull(messageType, "Message type must not be null"); + Assert.notNull(messageType, "MessageType must not be null"); Assert.notNull(textContent, "Content must not be null"); - Assert.notNull(media, "media data must not be null"); + Assert.notNull(media, "Media must not be null"); this.messageType = messageType; this.textContent = textContent; 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 4b324afa4..63c6b5a65 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 @@ -46,8 +46,8 @@ public class PromptTransformingChatService implements ChatService { private List chatServiceListeners; public PromptTransformingChatService(ChatConnector chatConnector, List retrievers, - List documentPostProcessors, List augmentors, - List chatServiceListeners) { + List documentPostProcessors, List augmentors, + List chatServiceListeners) { Objects.requireNonNull(chatConnector, "chatConnector must not be null"); this.chatConnector = chatConnector; this.retrievers = retrievers; diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallback.java b/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallback.java index 6f4e8bac4..c8aa7f50a 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallback.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallback.java @@ -26,18 +26,18 @@ public interface FunctionCallback { /** * @return Returns the Function name. Unique within the model. */ - public String getName(); + String getName(); /** * @return Returns the function description. This description is used by the model do * decide if the function should be called or not. */ - public String getDescription(); + String getDescription(); /** * @return Returns the JSON schema of the function input type. */ - public String getInputTypeSchema(); + String getInputTypeSchema(); /** * Called when a model detects and triggers a function call. The model is responsible @@ -47,6 +47,6 @@ public interface FunctionCallback { * model. * @return String containing the function call response. */ - public String call(String functionInput); + String call(String functionInput); } \ No newline at end of file diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallbackContext.java b/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallbackContext.java index 59f430987..297b25113 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallbackContext.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/function/FunctionCallbackContext.java @@ -15,11 +15,7 @@ */ package org.springframework.ai.model.function; -import java.lang.reflect.Type; -import java.util.function.Function; - import com.fasterxml.jackson.annotation.JsonClassDescription; - import org.springframework.ai.model.function.FunctionCallbackWrapper.Builder.SchemaType; import org.springframework.beans.BeansException; import org.springframework.cloud.function.context.catalog.FunctionTypeUtils; @@ -32,6 +28,9 @@ import org.springframework.lang.NonNull; import org.springframework.lang.Nullable; import org.springframework.util.StringUtils; +import java.lang.reflect.Type; +import java.util.function.Function; + /** * A Spring {@link ApplicationContextAware} implementation that provides a way to retrieve * a {@link Function} from the Spring context and wrap it into a {@link FunctionCallback}. @@ -47,6 +46,7 @@ import org.springframework.util.StringUtils; * * @author Christian Tzolov * @author Christopher Smith + * @author Josh Long */ public class FunctionCallbackContext implements ApplicationContextAware { @@ -63,6 +63,19 @@ public class FunctionCallbackContext implements ApplicationContextAware { this.applicationContext = (GenericApplicationContext) applicationContext; } + public FunctionCallback getFunctionCallback(String beanName, String defaultDescription, + Function function, SchemaType schemaType) { + var beanType = FunctionTypeUtils.discoverFunctionTypeFromClass(function.getClass()); + var functionInputType = TypeResolverHelper.getFunctionArgumentType(beanType, 0); + var functionInputClass = FunctionTypeUtils.getRawType(functionInputType); + return FunctionCallbackWrapper.builder(function) + .withName(beanName) + .withSchemaType(schemaType) + .withDescription(defaultDescription) + .withInputType(functionInputClass) + .build(); + } + @SuppressWarnings({ "rawtypes", "unchecked" }) public FunctionCallback getFunctionCallback(@NonNull String beanName, @Nullable String defaultDescription) { 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 020252d0c..ee5504e7c 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 @@ -81,7 +81,7 @@ public class SummaryMetadataEnricher implements DocumentTransformer { } public SummaryMetadataEnricher(ChatConnector chatConnector, List summaryTypes, String summaryTemplate, - MetadataMode metadataMode) { + MetadataMode metadataMode) { Assert.notNull(chatConnector, "ChatConnector must not be null"); Assert.hasText(summaryTemplate, "Summary template must not be empty"); diff --git a/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatConnectorTests.java b/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatConnectorTests.java index b50390aad..60e5d6cd4 100644 --- a/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatConnectorTests.java +++ b/spring-ai-core/src/test/java/org/springframework/ai/chat/ChatConnectorTests.java @@ -48,7 +48,7 @@ class ChatConnectorTests { String userMessage = "Zero Wing"; String responseMessage = "All your bases are belong to us"; - ChatConnector mockClient = Mockito.mock(ChatConnector.class); + ChatClient mockClient = Mockito.mock(ChatClient.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 1715121cc..345449906 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 @@ -48,7 +48,7 @@ import static org.mockito.Mockito.when; public class ChatMemoryTests { @Mock - ChatConnector chatConnector; + ChatConnector chatConnector; @Mock StreamingChatClient streamingChatClient; 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 ba4bee195..5d1a6d0d1 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 @@ -58,8 +58,8 @@ public class AnthropicAutoConfiguration { @Bean @ConditionalOnMissingBean public AnthropicChatConnector anthropicChatClient(AnthropicApi anthropicApi, AnthropicChatProperties chatProperties, - RetryTemplate retryTemplate, FunctionCallbackContext functionCallbackContext, - List toolFunctionCallbacks) { + RetryTemplate retryTemplate, FunctionCallbackContext functionCallbackContext, + List toolFunctionCallbacks) { if (!CollectionUtils.isEmpty(toolFunctionCallbacks)) { chatProperties.getOptions().getFunctionCallbacks().addAll(toolFunctionCallbacks); 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 2c78dc45a..f994191c1 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 @@ -59,8 +59,8 @@ public class AzureOpenAiAutoConfiguration { @ConditionalOnProperty(prefix = AzureOpenAiChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) public AzureOpenAiChatConnector azureOpenAiChatClient(OpenAIClient openAIClient, - AzureOpenAiChatProperties chatProperties, List toolFunctionCallbacks, - FunctionCallbackContext functionCallbackContext) { + AzureOpenAiChatProperties chatProperties, List toolFunctionCallbacks, + FunctionCallbackContext functionCallbackContext) { if (!CollectionUtils.isEmpty(toolFunctionCallbacks)) { chatProperties.getOptions().getFunctionCallbacks().addAll(toolFunctionCallbacks); 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 42a985ddb..54cfbbd21 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 @@ -60,7 +60,7 @@ public class BedrockAnthropicChatAutoConfiguration { @Bean @ConditionalOnBean(AnthropicChatBedrockApi.class) public BedrockAnthropicChatConnector anthropicChatClient(AnthropicChatBedrockApi anthropicApi, - BedrockAnthropicChatProperties properties) { + BedrockAnthropicChatProperties properties) { return new BedrockAnthropicChatConnector(anthropicApi, properties.getOptions()); } 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 ad0eae187..70d985107 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 @@ -60,7 +60,7 @@ public class BedrockAnthropic3ChatAutoConfiguration { @Bean @ConditionalOnBean(Anthropic3ChatBedrockApi.class) public BedrockAnthropic3ChatConnector anthropic3ChatClient(Anthropic3ChatBedrockApi anthropicApi, - BedrockAnthropic3ChatProperties properties) { + BedrockAnthropic3ChatProperties properties) { return new BedrockAnthropic3ChatConnector(anthropicApi, properties.getOptions()); } 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 911119a8d..b9d392c68 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 @@ -58,7 +58,7 @@ public class BedrockCohereChatAutoConfiguration { @Bean @ConditionalOnBean(CohereChatBedrockApi.class) public BedrockCohereChatConnector cohereChatClient(CohereChatBedrockApi cohereChatApi, - BedrockCohereChatProperties properties) { + BedrockCohereChatProperties properties) { return new BedrockCohereChatConnector(cohereChatApi, properties.getOptions()); } 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 9ab781038..aa13ec0de 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 @@ -60,7 +60,7 @@ public class BedrockLlamaChatAutoConfiguration { @Bean @ConditionalOnBean(LlamaChatBedrockApi.class) public BedrockLlamaChatConnector llamaChatClient(LlamaChatBedrockApi llamaApi, - BedrockLlamaChatProperties properties) { + BedrockLlamaChatProperties properties) { return new BedrockLlamaChatConnector(llamaApi, properties.getOptions()); } 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 842897d73..3fdd4b8d1 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 @@ -58,7 +58,7 @@ public class BedrockTitanChatAutoConfiguration { @Bean @ConditionalOnBean(TitanChatBedrockApi.class) public BedrockTitanChatConnector titanChatClient(TitanChatBedrockApi titanChatApi, - BedrockTitanChatProperties properties) { + BedrockTitanChatProperties properties) { return new BedrockTitanChatConnector(titanChatApi, properties.getOptions()); } 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 fb622f713..a6480ec50 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 @@ -70,9 +70,9 @@ public class MistralAiAutoConfiguration { @ConditionalOnProperty(prefix = MistralAiChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) public MistralAiChatConnector mistralAiChatClient(MistralAiCommonProperties commonProperties, - MistralAiChatProperties chatProperties, RestClient.Builder restClientBuilder, - List toolFunctionCallbacks, FunctionCallbackContext functionCallbackContext, - RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { + MistralAiChatProperties chatProperties, RestClient.Builder restClientBuilder, + List toolFunctionCallbacks, FunctionCallbackContext functionCallbackContext, + RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { var mistralAiApi = mistralAiApi(chatProperties.getApiKey(), commonProperties.getApiKey(), chatProperties.getBaseUrl(), commonProperties.getBaseUrl(), restClientBuilder, responseErrorHandler); 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 93dd8e71e..644f182bc 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 @@ -55,9 +55,9 @@ public class OpenAiAutoConfiguration { @ConditionalOnProperty(prefix = OpenAiChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) public OpenAiChatConnector openAiChatClient(OpenAiConnectionProperties commonProperties, - OpenAiChatProperties chatProperties, RestClient.Builder restClientBuilder, - List toolFunctionCallbacks, FunctionCallbackContext functionCallbackContext, - RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { + OpenAiChatProperties chatProperties, RestClient.Builder restClientBuilder, + List toolFunctionCallbacks, FunctionCallbackContext functionCallbackContext, + RetryTemplate retryTemplate, ResponseErrorHandler responseErrorHandler) { var openAiApi = openAiApi(chatProperties.getBaseUrl(), commonProperties.getBaseUrl(), chatProperties.getApiKey(), commonProperties.getApiKey(), restClientBuilder, responseErrorHandler); 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 6874b04bf..0ac36e124 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 @@ -75,8 +75,8 @@ public class VertexAiGeminiAutoConfiguration { @Bean @ConditionalOnMissingBean public VertexAiGeminiChatConnector vertexAiGeminiChat(VertexAI vertexAi, - VertexAiGeminiChatProperties chatProperties, List toolFunctionCallbacks, - ApplicationContext context) { + VertexAiGeminiChatProperties chatProperties, List toolFunctionCallbacks, + ApplicationContext context) { FunctionCallbackContext functionCallbackContext = springAiFunctionManager(context); 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 f80e846d4..6005087b9 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 @@ -48,7 +48,7 @@ public class VertexAiPalm2AutoConfiguration { @ConditionalOnProperty(prefix = VertexAiPlam2ChatProperties.CONFIG_PREFIX, name = "enabled", havingValue = "true", matchIfMissing = true) public VertexAiPaLm2ChatConnector vertexAiChatClient(VertexAiPaLm2Api vertexAiApi, - VertexAiPlam2ChatProperties chatProperties) { + VertexAiPlam2ChatProperties chatProperties) { return new VertexAiPaLm2ChatConnector(vertexAiApi, chatProperties.getOptions()); } 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 911d16fca..31568c58c 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 @@ -15,17 +15,13 @@ */ package org.springframework.ai.autoconfigure.anthropic; -import java.util.List; -import java.util.stream.Collectors; - 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.AnthropicChatConnector; -import reactor.core.publisher.Flux; - 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; @@ -34,6 +30,10 @@ import org.springframework.ai.chat.prompt.Prompt; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; import org.springframework.boot.test.context.runner.ApplicationContextRunner; +import reactor.core.publisher.Flux; + +import java.util.List; +import java.util.stream.Collectors; import static org.assertj.core.api.Assertions.assertThat; @@ -50,7 +50,7 @@ public class AnthropicAutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - AnthropicChatConnector chatClient = context.getBean(AnthropicChatConnector.class); + ChatClient chatClient = ChatClient.builder(context.getBean(AnthropicChatConnector.class)).build(); String response = chatClient.call("Hello"); assertThat(response).isNotEmpty(); logger.info("Response: " + response); 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 2e5c10e52..e3b62a124 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 @@ -15,14 +15,10 @@ */ package org.springframework.ai.autoconfigure.anthropic.tool; -import java.util.List; -import java.util.function.Function; - import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; import org.slf4j.Logger; import org.slf4j.LoggerFactory; - import org.springframework.ai.anthropic.AnthropicChatConnector; import org.springframework.ai.anthropic.AnthropicChatOptions; import org.springframework.ai.anthropic.api.AnthropicApi; @@ -40,6 +36,9 @@ import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; import org.springframework.context.annotation.Description; +import java.util.List; +import java.util.function.Function; + import static org.assertj.core.api.Assertions.assertThat; @EnabledIfEnvironmentVariable(named = "ANTHROPIC_API_KEY", matches = ".*") 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 282cd8e23..0478fd087 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 @@ -15,25 +15,25 @@ */ package org.springframework.ai.autoconfigure.mistralai; -import java.util.List; -import java.util.stream.Collectors; - 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.MistralAiChatConnector; -import reactor.core.publisher.Flux; - import org.springframework.ai.autoconfigure.retry.SpringAiRetryAutoConfiguration; +import org.springframework.ai.chat.ChatClient; 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.MistralAiChatConnector; import org.springframework.ai.mistralai.MistralAiEmbeddingClient; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.autoconfigure.web.client.RestClientAutoConfiguration; import org.springframework.boot.test.context.runner.ApplicationContextRunner; +import reactor.core.publisher.Flux; + +import java.util.List; +import java.util.stream.Collectors; import static org.assertj.core.api.Assertions.assertThat; @@ -54,7 +54,7 @@ public class MistralAiAutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - MistralAiChatConnector client = context.getBean(MistralAiChatConnector.class); + ChatClient client = ChatClient.builder(context.getBean(MistralAiChatConnector.class)).build(); String response = client.call("Hello"); assertThat(response).isNotEmpty(); logger.info("Response: " + response); 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 741d98bb2..30ff67f74 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 @@ -23,6 +23,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.chat.ChatClient; import org.springframework.ai.chat.messages.UserMessage; import org.springframework.ai.chat.prompt.Prompt; import org.springframework.ai.image.ImagePrompt; @@ -55,7 +56,7 @@ public class OpenAiAutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - OpenAiChatConnector client = context.getBean(OpenAiChatConnector.class); + ChatClient client = ChatClient.builder(context.getBean(OpenAiChatConnector.class)).build(); String response = client.call("Hello"); assertThat(response).isNotEmpty(); logger.info("Response: " + response); 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 16f073bb8..021852b35 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 @@ -15,20 +15,20 @@ */ package org.springframework.ai.autoconfigure.vertexai.gemini; -import java.util.stream.Collectors; - 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 reactor.core.publisher.Flux; - +import org.springframework.ai.chat.ChatClient; 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.VertexAiGeminiChatConnector; import org.springframework.boot.autoconfigure.AutoConfigurations; import org.springframework.boot.test.context.runner.ApplicationContextRunner; +import reactor.core.publisher.Flux; + +import java.util.stream.Collectors; import static org.assertj.core.api.Assertions.assertThat; @@ -46,7 +46,7 @@ public class VertexAiGeminiAutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - VertexAiGeminiChatConnector client = context.getBean(VertexAiGeminiChatConnector.class); + ChatClient client = ChatClient.builder(context.getBean(VertexAiGeminiChatConnector.class)).build(); String response = client.call("Hello"); assertThat(response).isNotEmpty(); 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 c9fdef90e..64d534ac4 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 @@ -22,6 +22,7 @@ import org.apache.commons.logging.LogFactory; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable; +import org.springframework.ai.chat.ChatClient; import org.springframework.ai.embedding.EmbeddingResponse; import org.springframework.ai.vertexai.palm2.VertexAiPaLm2ChatConnector; import org.springframework.ai.vertexai.palm2.VertexAiPaLm2EmbeddingClient; @@ -48,8 +49,8 @@ public class VertexAiPaLm2AutoConfigurationIT { @Test void generate() { contextRunner.run(context -> { - VertexAiPaLm2ChatConnector client = context.getBean(VertexAiPaLm2ChatConnector.class); - + VertexAiPaLm2ChatConnector connector = context.getBean(VertexAiPaLm2ChatConnector.class); + ChatClient client = ChatClient.builder(connector).build(); String response = client.call("Hello"); assertThat(response).isNotEmpty();