From e92616b10e36f438bfc4def8c54945ba97f931f8 Mon Sep 17 00:00:00 2001 From: Ricken Bazolo Date: Sun, 19 Jan 2025 02:41:52 +0100 Subject: [PATCH] fix(mistral) Added index of tool call in the list of tool calls Additiona fixes: API key validation and tool calling backward compatibility - Fix API key validation in OpenAiApi builder - Standardize API key validation using Assert.notNull - Add backward compatibility support for FunctionCallback in tool calling - Update integration tests to use LegacyToolCallingManager Co-authored-by:Christian Tzolov Signed-off-by: Ricken Bazolo Signed-off-by: Christian Tzolov --- .../ai/mistralai/MistralAiChatModel.java | 2 +- .../ai/mistralai/api/MistralAiApi.java | 3 ++- .../MistralAiStreamFunctionCallingHelper.java | 12 +++++---- .../ai/openai/api/OpenAiApi.java | 6 +++++ .../proxy/MistralWithOpenAiChatModelIT.java | 19 +++++++++----- .../proxy/OllamaWithOpenAiChatModelIT.java | 13 +++++++--- .../ai/model/SimpleApiKey.java | 26 +++---------------- .../model/tool/DefaultToolCallingManager.java | 6 +++++ 8 files changed, 48 insertions(+), 39 deletions(-) diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java index 804380086..222f6f556 100644 --- a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java +++ b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/MistralAiChatModel.java @@ -384,7 +384,7 @@ public class MistralAiChatModel extends AbstractToolCallSupport implements ChatM if (!CollectionUtils.isEmpty(assistantMessage.getToolCalls())) { toolCalls = assistantMessage.getToolCalls().stream().map(toolCall -> { var function = new ChatCompletionFunction(toolCall.name(), toolCall.arguments()); - return new ToolCall(toolCall.id(), toolCall.type(), function); + return new ToolCall(toolCall.id(), toolCall.type(), function, null); }).toList(); } 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 6b9e23d22..0be9617d5 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 @@ -833,10 +833,11 @@ public class MistralAiApi { * @param type The type of tool call the output is required for. For now, this is * always function. * @param function The function definition. + * @param index The index of the tool call in the list of tool calls. */ @JsonInclude(Include.NON_NULL) public record ToolCall(@JsonProperty("id") String id, @JsonProperty("type") String type, - @JsonProperty("function") ChatCompletionFunction function) { + @JsonProperty("function") ChatCompletionFunction function, @JsonProperty("index") Integer index) { } diff --git a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/api/MistralAiStreamFunctionCallingHelper.java b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/api/MistralAiStreamFunctionCallingHelper.java index 608e8d6ff..bfcb2225a 100644 --- a/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/api/MistralAiStreamFunctionCallingHelper.java +++ b/models/spring-ai-mistral-ai/src/main/java/org/springframework/ai/mistralai/api/MistralAiStreamFunctionCallingHelper.java @@ -18,6 +18,7 @@ package org.springframework.ai.mistralai.api; import java.util.ArrayList; import java.util.List; +import java.util.Objects; import java.util.Optional; import java.util.UUID; @@ -74,16 +75,16 @@ public class MistralAiStreamFunctionCallingHelper { Optional id = current.delta() .toolCalls() .stream() - .filter(tool -> tool.id() != null) - .map(tool -> tool.id()) + .map(ToolCall::id) + .filter(Objects::nonNull) .findFirst(); - if (!id.isPresent()) { + if (id.isEmpty()) { var newId = UUID.randomUUID().toString(); var toolCallsWithID = current.delta() .toolCalls() .stream() - .map(toolCall -> new ToolCall(newId, "function", toolCall.function())) + .map(toolCall -> new ToolCall(newId, "function", toolCall.function(), toolCall.index())) .toList(); var role = current.delta().role() != null ? current.delta().role() : Role.ASSISTANT; @@ -151,7 +152,8 @@ public class MistralAiStreamFunctionCallingHelper { String id = (current.id() != null ? current.id() : previous.id()); String type = (current.type() != null ? current.type() : previous.type()); ChatCompletionFunction function = merge(previous.function(), current.function()); - return new ToolCall(id, type, function); + Integer index = (current.index() != null ? current.index() : previous.index()); + return new ToolCall(id, type, function, index); } private ChatCompletionFunction merge(ChatCompletionFunction previous, ChatCompletionFunction current) { 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 7827f034d..0183fa65a 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 @@ -1675,6 +1675,12 @@ public class OpenAiApi { return this; } + public Builder apiKey(String simpleApiKey) { + Assert.notNull(simpleApiKey, "apiKey cannot be null"); + this.apiKey = new SimpleApiKey(simpleApiKey); + return this; + } + public Builder headers(MultiValueMap headers) { Assert.notNull(headers, "headers cannot be null"); this.headers = headers; diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/proxy/MistralWithOpenAiChatModelIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/proxy/MistralWithOpenAiChatModelIT.java index 597c61b67..fcae688e2 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/proxy/MistralWithOpenAiChatModelIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/proxy/MistralWithOpenAiChatModelIT.java @@ -46,12 +46,16 @@ import org.springframework.ai.converter.BeanOutputConverter; import org.springframework.ai.converter.ListOutputConverter; import org.springframework.ai.converter.MapOutputConverter; import org.springframework.ai.model.Media; +import org.springframework.ai.model.SimpleApiKey; import org.springframework.ai.model.function.FunctionCallback; +import org.springframework.ai.model.tool.LegacyToolCallingManager; 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.tool.MockWeatherService; import org.springframework.ai.openai.chat.ActorsFilms; +import org.springframework.ai.tool.function.FunctionToolCallback; +import org.springframework.ai.tool.method.MethodToolCallback; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.beans.factory.annotation.Value; import org.springframework.boot.SpringBootConfiguration; @@ -61,6 +65,7 @@ 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 org.springframework.util.ReflectionUtils; import static org.assertj.core.api.Assertions.assertThat; @@ -251,8 +256,7 @@ class MistralWithOpenAiChatModelIT { var promptOptions = OpenAiChatOptions.builder() .model(modelName) - .functionCallbacks(List.of(FunctionCallback.builder() - .function("getCurrentWeather", new MockWeatherService()) + .toolCallbacks(List.of(FunctionToolCallback.builder("getCurrentWeather", new MockWeatherService()) .description("Get the weather in location") .inputType(MockWeatherService.Request.class) .build())) @@ -276,8 +280,7 @@ class MistralWithOpenAiChatModelIT { var promptOptions = OpenAiChatOptions.builder() .model(modelName) - .functionCallbacks(List.of(FunctionCallback.builder() - .function("getCurrentWeather", new MockWeatherService()) + .toolCallbacks(List.of(FunctionToolCallback.builder("getCurrentWeather", new MockWeatherService()) .description("Get the weather in location") .inputType(MockWeatherService.Request.class) .build())) @@ -388,12 +391,16 @@ class MistralWithOpenAiChatModelIT { @Bean public OpenAiApi chatCompletionApi() { - return new OpenAiApi(MISTRAL_BASE_URL, System.getenv("MISTRAL_AI_API_KEY")); + return OpenAiApi.builder().baseUrl(MISTRAL_BASE_URL).apiKey(System.getenv("MISTRAL_AI_API_KEY")).build(); } @Bean public OpenAiChatModel openAiClient(OpenAiApi openAiApi) { - return new OpenAiChatModel(openAiApi, OpenAiChatOptions.builder().model(MISTRAL_DEFAULT_MODEL).build()); + return OpenAiChatModel.builder() + .openAiApi(openAiApi) + .toolCallingManager(LegacyToolCallingManager.builder().build()) + .defaultOptions(OpenAiChatOptions.builder().model(MISTRAL_DEFAULT_MODEL).build()) + .build(); } } diff --git a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/proxy/OllamaWithOpenAiChatModelIT.java b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/proxy/OllamaWithOpenAiChatModelIT.java index a61bb1fe9..960aa8f46 100644 --- a/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/proxy/OllamaWithOpenAiChatModelIT.java +++ b/models/spring-ai-openai/src/test/java/org/springframework/ai/openai/chat/proxy/OllamaWithOpenAiChatModelIT.java @@ -50,11 +50,13 @@ import org.springframework.ai.converter.ListOutputConverter; import org.springframework.ai.converter.MapOutputConverter; import org.springframework.ai.model.Media; import org.springframework.ai.model.function.FunctionCallback; +import org.springframework.ai.model.tool.LegacyToolCallingManager; 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.tool.MockWeatherService; import org.springframework.ai.openai.chat.ActorsFilms; +import org.springframework.ai.tool.function.FunctionToolCallback; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.beans.factory.annotation.Value; import org.springframework.boot.SpringBootConfiguration; @@ -272,8 +274,7 @@ class OllamaWithOpenAiChatModelIT { // Note for Ollama you must set the tool choice to explicitly. Unlike OpenAI // (which defaults to "auto") Ollama defaults to "nono" .toolChoice("auto") - .functionCallbacks(List.of(FunctionCallback.builder() - .function("getCurrentWeather", new MockWeatherService()) + .toolCallbacks(List.of(FunctionToolCallback.builder("getCurrentWeather", new MockWeatherService()) .description("Get the weather in location") .inputType(MockWeatherService.Request.class) .build())) @@ -412,12 +413,16 @@ class OllamaWithOpenAiChatModelIT { @Bean public OpenAiApi chatCompletionApi() { - return new OpenAiApi(baseUrl, ""); + return OpenAiApi.builder().baseUrl(baseUrl).apiKey("").build(); } @Bean public OpenAiChatModel openAiClient(OpenAiApi openAiApi) { - return new OpenAiChatModel(openAiApi, OpenAiChatOptions.builder().model(DEFAULT_OLLAMA_MODEL).build()); + return OpenAiChatModel.builder() + .openAiApi(openAiApi) + .toolCallingManager(LegacyToolCallingManager.builder().build()) + .defaultOptions(OpenAiChatOptions.builder().model(DEFAULT_OLLAMA_MODEL).build()) + .build(); } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/SimpleApiKey.java b/spring-ai-core/src/main/java/org/springframework/ai/model/SimpleApiKey.java index 4887ff12b..ce655cf3a 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/SimpleApiKey.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/SimpleApiKey.java @@ -23,11 +23,10 @@ import org.springframework.util.Assert; * be refreshed or rotated. * * @author Adib Saikali + * @author Christian Tzolov * @since 1.0.0 */ -public final class SimpleApiKey implements ApiKey { - - private final String value; +public record SimpleApiKey(String value) implements ApiKey { /** * Create a new SimpleApiKey. @@ -35,34 +34,17 @@ public final class SimpleApiKey implements ApiKey { * @throws IllegalArgumentException if value is null or empty */ public SimpleApiKey(String value) { - Assert.hasText(value, "API key value must not be null or empty"); + Assert.notNull(value, "API key value must not be null or empty"); this.value = value; } @Override public String getValue() { - return this.value; + return this.value(); } @Override public String toString() { return "SimpleApiKey{value='***'}"; } - - @Override - public boolean equals(Object o) { - if (this == o) { - return true; - } - if (!(o instanceof SimpleApiKey that)) { - return false; - } - return this.value.equals(that.value); - } - - @Override - public int hashCode() { - return this.value.hashCode(); - } - } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/tool/DefaultToolCallingManager.java b/spring-ai-core/src/main/java/org/springframework/ai/model/tool/DefaultToolCallingManager.java index ed9b739d7..3e38ea6ea 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/tool/DefaultToolCallingManager.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/tool/DefaultToolCallingManager.java @@ -214,6 +214,12 @@ public class DefaultToolCallingManager implements ToolCallingManager { else if (toolCallback instanceof ToolCallback callback) { returnDirect = returnDirect && callback.getToolMetadata().returnDirect(); } + else if (returnDirect == null) { + // This is a temporary solution to ensure backward compatibility with + // FunctionCallback. + // TODO: remove this block when FunctionCallback is removed. + returnDirect = false; + } String toolResult; try {