From 4209d5a65d16e7587c76babc726c8e173687c161 Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Tue, 12 Mar 2024 11:34:56 +0100 Subject: [PATCH] Make consistent sync/stream AssistantMessage properties for OpenAI and Mistral AI --- .../ai/mistralai/MistralAiChatClient.java | 18 +++++++++++-- .../ai/openai/OpenAiChatClient.java | 25 +++++++++++-------- 2 files changed, 31 insertions(+), 12 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 9f506e02d..91aad8f54 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 @@ -15,6 +15,7 @@ */ package org.springframework.ai.mistralai; +import java.util.HashMap; import java.util.HashSet; import java.util.List; import java.util.Map; @@ -113,8 +114,7 @@ public class MistralAiChatClient extends List generations = chatCompletion.choices() .stream() - .map(choice -> new Generation(choice.message().content(), - Map.of("role", choice.message().role().name())) + .map(choice -> new Generation(choice.message().content(), toMap(chatCompletion.id(), choice)) .withGenerationMetadata(ChatGenerationMetadata.from(choice.finishReason().name(), null))) .toList(); @@ -122,6 +122,20 @@ public class MistralAiChatClient extends }); } + private Map toMap(String id, ChatCompletion.Choice choice) { + Map map = new HashMap<>(); + + var message = choice.message(); + if (message.role() != null) { + map.put("role", message.role().name()); + } + if (choice.finishReason() != null) { + map.put("finishReason", choice.finishReason().name()); + } + map.put("id", id); + return map; + } + @Override public Flux stream(Prompt prompt) { var request = createRequest(prompt, true); 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/OpenAiChatClient.java index 65410d40f..79c4b6498 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/OpenAiChatClient.java @@ -147,7 +147,7 @@ public class OpenAiChatClient extends RateLimit rateLimits = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(completionEntity); List generations = chatCompletion.choices().stream().map(choice -> { - return new Generation(choice.message().content(), toMap(choice.message())) + return new Generation(choice.message().content(), toMap(chatCompletion.id(), choice)) .withGenerationMetadata(ChatGenerationMetadata.from(choice.finishReason().name(), null)); }).toList(); @@ -156,6 +156,20 @@ public class OpenAiChatClient extends }); } + private Map toMap(String id, ChatCompletion.Choice choice) { + Map map = new HashMap<>(); + + var message = choice.message(); + if (message.role() != null) { + map.put("role", message.role().name()); + } + if (choice.finishReason() != null) { + map.put("finishReason", choice.finishReason().name()); + } + map.put("id", id); + return map; + } + @Override public Flux stream(Prompt prompt) { @@ -280,15 +294,6 @@ public class OpenAiChatClient extends }).toList(); } - private Map toMap(ChatCompletionMessage message) { - Map map = new HashMap<>(); - - if (message.role() != null) { - map.put("role", message.role().name()); - } - return map; - } - @Override protected ChatCompletionRequest doCreateToolResponseRequest(ChatCompletionRequest previousRequest, ChatCompletionMessage responseMessage, List conversationHistory) {