From 7e03a15cf5ce4fbdfcf219fbfed84514201d4d67 Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Fri, 26 Apr 2024 21:04:48 +0200 Subject: [PATCH] Mistral AI streaming function API change fix --- .../ai/mistralai/MistralAiChatClient.java | 8 +++++++- .../springframework/ai/mistralai/api/MistralAiApi.java | 10 ++++------ .../api/MistralAiStreamFunctionCallingHelper.java | 3 +-- 3 files changed, 12 insertions(+), 9 deletions(-) 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/MistralAiChatClient.java index 91aad8f54..ad7c347b0 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/MistralAiChatClient.java @@ -285,7 +285,13 @@ public class MistralAiChatClient extends @SuppressWarnings("null") @Override protected ChatCompletionMessage doGetToolResponseMessage(ResponseEntity chatCompletion) { - return chatCompletion.getBody().choices().iterator().next().message(); + ChatCompletionMessage msg = chatCompletion.getBody().choices().iterator().next().message(); + if (msg.role() == null) { + // add missing role + msg = new ChatCompletionMessage(msg.content(), ChatCompletionMessage.Role.ASSISTANT, msg.name(), + msg.toolCalls()); + } + return msg; } @Override 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 19ba2ac52..5373fd01d 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 @@ -535,14 +535,12 @@ public class MistralAiApi { */ @JsonProperty("model_length") MODEL_LENGTH, /** - * The model called a tool. + * */ - @JsonProperty("tool_call") TOOL_CALL, - - // anticipation of future changes. Based on: - // https://github.com/mistralai/client-python/blob/main/src/mistralai/models/chat_completion.py @JsonProperty("error") ERROR, - + /** + * The model requested a tool call. + */ @JsonProperty("tool_calls") TOOL_CALLS // @formatter:on 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 50cd22353..ed730ec48 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 @@ -190,8 +190,7 @@ public class MistralAiStreamFunctionCallingHelper { } var choice = choices.get(0); - return choice.finishReason() == ChatCompletionFinishReason.TOOL_CALL - || choice.finishReason() == ChatCompletionFinishReason.TOOL_CALLS; + return choice.finishReason() == ChatCompletionFinishReason.TOOL_CALLS; } }