From 588082285a7efa1c6169a16e2cbe85d7978eae98 Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Thu, 18 Jul 2024 15:04:21 +0200 Subject: [PATCH] OpenAI ChatModel: handle null finish reasons responses --- .../java/org/springframework/ai/openai/OpenAiChatModel.java | 3 ++- .../springframework/ai/openai/chat/OpenAiChatModelIT.java | 6 +++--- 2 files changed, 5 insertions(+), 4 deletions(-) diff --git a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java index cb0292785..b34328fcb 100644 --- a/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java +++ b/models/spring-ai-openai/src/main/java/org/springframework/ai/openai/OpenAiChatModel.java @@ -277,7 +277,8 @@ public class OpenAiChatModel extends AbstractToolCallSupport implements ChatMode .toList(); var assistantMessage = new AssistantMessage(choice.message().content(), metadata, toolCalls); - var generationMetadata = ChatGenerationMetadata.from(choice.finishReason().name(), null); + String finishReason = (choice.finishReason() != null ? choice.finishReason().name() : ""); + var generationMetadata = ChatGenerationMetadata.from(finishReason, null); var generation = new Generation(assistantMessage, generationMetadata); return generation; 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 index 5cedfc432..f1198d326 100644 --- 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 @@ -302,7 +302,7 @@ class OpenAiChatModelIT extends AbstractIT { logger.info(response.getResult().getOutput().getContent()); assertThat(response.getResult().getOutput().getContent()).contains("bananas", "apple"); - assertThat(response.getResult().getOutput().getContent()).containsAnyOf("bowl", "basket"); + assertThat(response.getResult().getOutput().getContent()).containsAnyOf("bowl", "basket", "fruit stand"); } @ParameterizedTest(name = "{0} : {displayName} ") @@ -318,7 +318,7 @@ class OpenAiChatModelIT extends AbstractIT { logger.info(response.getResult().getOutput().getContent()); assertThat(response.getResult().getOutput().getContent()).contains("bananas", "apple"); - assertThat(response.getResult().getOutput().getContent()).containsAnyOf("bowl", "basket"); + assertThat(response.getResult().getOutput().getContent()).containsAnyOf("bowl", "basket", "fruit stand"); } @Test @@ -341,7 +341,7 @@ class OpenAiChatModelIT extends AbstractIT { .collect(Collectors.joining()); logger.info("Response: {}", content); assertThat(content).contains("bananas", "apple"); - assertThat(content).containsAnyOf("bowl", "basket"); + assertThat(content).containsAnyOf("bowl", "basket", "fruit stand"); } @Test