diff --git a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java index 8b8317b7d..abae7e749 100644 --- a/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java +++ b/models/spring-ai-anthropic/src/main/java/org/springframework/ai/anthropic/AnthropicChatModel.java @@ -64,7 +64,7 @@ import java.util.stream.Collectors; * @author Mariusz Bernacki * @since 1.0.0 */ -public class AnthropicChatModel extends AbstractToolCallSupport implements ChatModel { +public class AnthropicChatModel extends AbstractToolCallSupport implements ChatModel { private static final Logger logger = LoggerFactory.getLogger(AnthropicChatModel.class); @@ -353,7 +353,6 @@ public class AnthropicChatModel extends AbstractToolCallSupport implements ChatModel { +public class AzureOpenAiChatModel extends AbstractToolCallSupport implements ChatModel { private static final String DEFAULT_DEPLOYMENT_NAME = "gpt-35-turbo"; @@ -557,7 +557,6 @@ public class AzureOpenAiChatModel extends AbstractToolCallSupport implements ChatModel { +public class MistralAiChatModel extends AbstractToolCallSupport implements ChatModel { private final Logger log = LoggerFactory.getLogger(getClass()); @@ -358,7 +358,6 @@ public class MistralAiChatModel extends AbstractToolCallSupport }).toList(); } - @Override protected boolean isToolFunctionCall(ChatCompletion chatCompletion) { var body = chatCompletion; 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 9d313cc2e..cb0292785 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 @@ -36,7 +36,6 @@ import org.springframework.ai.model.function.FunctionCallbackContext; import org.springframework.ai.openai.api.OpenAiApi; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletion; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletion.Choice; -import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionFinishReason; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage.ChatCompletionFunction; import org.springframework.ai.openai.api.OpenAiApi.ChatCompletionMessage.MediaContent; @@ -80,7 +79,7 @@ import java.util.concurrent.ConcurrentHashMap; * @see StreamingChatModel * @see OpenAiApi */ -public class OpenAiChatModel extends AbstractToolCallSupport implements ChatModel { +public class OpenAiChatModel extends AbstractToolCallSupport implements ChatModel { private static final Logger logger = LoggerFactory.getLogger(OpenAiChatModel.class); @@ -144,50 +143,45 @@ public class OpenAiChatModel extends AbstractToolCallSupport imp ChatCompletionRequest request = createRequest(prompt, false); - return this.retryTemplate.execute(ctx -> { + ResponseEntity completionEntity = this.retryTemplate + .execute(ctx -> this.openAiApi.chatCompletionEntity(request)); - ResponseEntity completionEntity = this.openAiApi.chatCompletionEntity(request); + var chatCompletion = completionEntity.getBody(); - var chatCompletion = completionEntity.getBody(); + if (chatCompletion == null) { + logger.warn("No chat completion returned for prompt: {}", prompt); + return new ChatResponse(List.of()); + } - if (chatCompletion == null) { - logger.warn("No chat completion returned for prompt: {}", prompt); - return new ChatResponse(List.of()); - } + List choices = chatCompletion.choices(); + if (choices == null) { + logger.warn("No choices returned for prompt: {}", prompt); + return new ChatResponse(List.of()); + } - if (isToolFunctionCall(chatCompletion)) { - List toolCallMessageConversation = this.handleToolCallRequests(prompt.getInstructions(), - chatCompletion); - // Recursively call the call method with the tool call message - // conversation that contains the call responses. + List generations = choices.stream().map(choice -> { + // @formatter:off + Map metadata = Map.of( + "id", chatCompletion.id(), + "role", choice.message().role() != null ? choice.message().role().name() : "", + "finishReason", choice.finishReason() != null ? choice.finishReason().name() : ""); + // @formatter:on + return buildGeneration(choice, metadata); + }).toList(); - return this.call(new Prompt(toolCallMessageConversation, prompt.getOptions())); - } + // Non function calling. + RateLimit rateLimit = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(completionEntity); - // Non function calling. - RateLimit rateLimit = OpenAiResponseHeaderExtractor.extractAiResponseHeaders(completionEntity); + ChatResponse chatResponse = new ChatResponse(generations, from(completionEntity.getBody(), rateLimit)); - List choices = chatCompletion.choices(); - if (choices == null) { - logger.warn("No choices returned for prompt: {}", prompt); - return new ChatResponse(List.of()); - } + if (isToolCall(chatResponse, OpenAiApi.ChatCompletionFinishReason.TOOL_CALLS.name())) { + var toolCallConversation = handleToolCalls(prompt, chatResponse); + // Recursively call the call method with the tool call message + // conversation that contains the call responses. + return this.call(new Prompt(toolCallConversation, prompt.getOptions())); + } - List generations = choices.stream().map(choice -> { - Map metadata = Map.of("id", chatCompletion.id(), "role", - choice.message().role() != null ? choice.message().role().name() : "", "finishReason", - choice.finishReason() != null ? choice.finishReason().name() : ""); - var generation = new Generation(choice.message().content(), metadata); - if (choice.finishReason() != null) { - generation = generation - .withGenerationMetadata(ChatGenerationMetadata.from(choice.finishReason().name(), null)); - } - return generation; - - }).toList(); - - return new ChatResponse(generations, from(completionEntity.getBody(), rateLimit)); - }); + return chatResponse; } public static ChatResponseMetadata from(OpenAiApi.ChatCompletion result, RateLimit rateLimit) { @@ -207,60 +201,58 @@ public class OpenAiChatModel extends AbstractToolCallSupport imp ChatCompletionRequest request = createRequest(prompt, true); - return this.retryTemplate.execute(ctx -> { + Flux completionChunks = this.retryTemplate + .execute(ctx -> this.openAiApi.chatCompletionStream(request)); - Flux completionChunks = this.openAiApi.chatCompletionStream(request); + // For chunked responses, only the first chunk contains the choice role. + // The rest of the chunks with same ID share the same role. + ConcurrentHashMap roleMap = new ConcurrentHashMap<>(); - // For chunked responses, only the first chunk contains the choice role. - // The rest of the chunks with same ID share the same role. - ConcurrentHashMap roleMap = new ConcurrentHashMap<>(); + // Convert the ChatCompletionChunk into a ChatCompletion to be able to reuse + // the function call handling logic. + Flux chatResponse = completionChunks.map(this::chunkToChatCompletion) + .switchMap(chatCompletion -> Mono.just(chatCompletion).map(chatCompletion2 -> { + try { + @SuppressWarnings("null") + String id = chatCompletion2.id(); - // Convert the ChatCompletionChunk into a ChatCompletion to be able to reuse - // the function call handling logic. - return completionChunks.map(this::chunkToChatCompletion).switchMap(chatCompletion -> { + // @formatter:off + List generations = chatCompletion2.choices().stream().map(choice -> { + if (choice.message().role() != null) { + roleMap.putIfAbsent(id, choice.message().role().name()); + } + Map metadata = Map.of( + "id", chatCompletion2.id(), + "role", roleMap.getOrDefault(id, ""), + "finishReason", choice.finishReason() != null ? choice.finishReason().name() : ""); + return buildGeneration(choice, metadata); + }).toList(); + // @formatter:on - if (this.isToolFunctionCall(chatCompletion)) { - var toolCallMessageConversation = this.handleToolCallRequests(prompt.getInstructions(), - chatCompletion); - // Recursively call the stream method with the tool call message - // conversation that contains the call responses. - return this.stream(new Prompt(toolCallMessageConversation, prompt.getOptions())); + if (chatCompletion2.usage() != null) { + return new ChatResponse(generations, from(chatCompletion2)); + } + else { + return new ChatResponse(generations); + } + } + catch (Exception e) { + logger.error("Error processing chat completion", e); + return new ChatResponse(List.of()); } - // Non function calling. - return Mono.just(chatCompletion).map(chatCompletion2 -> { - try { - @SuppressWarnings("null") - String id = chatCompletion2.id(); + })); - List generations = chatCompletion2.choices().stream().map(choice -> { - if (choice.message().role() != null) { - roleMap.putIfAbsent(id, choice.message().role().name()); - } - String finish = (choice.finishReason() != null ? choice.finishReason().name() : ""); - var generation = new Generation(choice.message().content(), - Map.of("id", id, "role", roleMap.getOrDefault(id, ""), "finishReason", finish)); - if (choice.finishReason() != null) { - generation = generation.withGenerationMetadata( - ChatGenerationMetadata.from(choice.finishReason().name(), null)); - } - return generation; - }).toList(); - - if (chatCompletion2.usage() != null) { - return new ChatResponse(generations, from(chatCompletion2)); - } - else { - return new ChatResponse(generations); - } - } - catch (Exception e) { - logger.error("Error processing chat completion", e); - return new ChatResponse(List.of()); - } - - }); - }); + return chatResponse.flatMap(response -> { + if (isToolCall(response, OpenAiApi.ChatCompletionFinishReason.TOOL_CALLS.name())) { + var toolCallConversation = handleToolCalls(prompt, response); + // Recursively call the stream method with the tool call message + // conversation that contains the call responses. + return this.stream(new Prompt(toolCallConversation, prompt.getOptions())); + } + else { + return Flux.just(response); + } }); } @@ -275,26 +267,33 @@ public class OpenAiChatModel extends AbstractToolCallSupport imp .build(); } - private List handleToolCallRequests(List previousMessages, ChatCompletion chatCompletion) { + private static Generation buildGeneration(Choice choice, Map metadata) { + List toolCalls = choice.message().toolCalls() == null ? List.of() + : choice.message() + .toolCalls() + .stream() + .map(toolCall -> new AssistantMessage.ToolCall(toolCall.id(), "function", + toolCall.function().name(), toolCall.function().arguments())) + .toList(); - ChatCompletionMessage nativeAssistantMessage = this.extractAssistantMessage(chatCompletion); + var assistantMessage = new AssistantMessage(choice.message().content(), metadata, toolCalls); + var generationMetadata = ChatGenerationMetadata.from(choice.finishReason().name(), null); + var generation = new Generation(assistantMessage, generationMetadata); - List assistantToolCalls = nativeAssistantMessage.toolCalls() - .stream() - .map(toolCall -> new AssistantMessage.ToolCall(toolCall.id(), "function", toolCall.function().name(), - toolCall.function().arguments())) - .toList(); + return generation; + } - AssistantMessage assistantMessage = new AssistantMessage(nativeAssistantMessage.content(), Map.of(), - assistantToolCalls); + private List handleToolCalls(Prompt prompt, ChatResponse response) { + AssistantMessage assistantMessage = response.getResult().getOutput(); + ToolResponseMessage toolMessageResponse = this.executeFuncitons(assistantMessage); + return this.buildToolCallConversation(prompt.getInstructions(), assistantMessage, toolMessageResponse); + } - ToolResponseMessage toolResponseMessage = this.executeFuncitons(assistantMessage); - - // History + private List buildToolCallConversation(List previousMessages, AssistantMessage assistantMessage, + ToolResponseMessage toolResponseMessage) { List messages = new ArrayList<>(previousMessages); messages.add(assistantMessage); messages.add(toolResponseMessage); - return messages; } @@ -314,10 +313,6 @@ public class OpenAiChatModel extends AbstractToolCallSupport imp chunk.systemFingerprint(), "chat.completion", chunk.usage()); } - private ChatCompletionMessage extractAssistantMessage(ChatCompletion chatCompletion) { - return chatCompletion.choices().iterator().next().message(); - } - /** * Accessible for testing. */ @@ -441,22 +436,6 @@ public class OpenAiChatModel extends AbstractToolCallSupport imp }).toList(); } - @Override - protected boolean isToolFunctionCall(ChatCompletion chatCompletion) { - if (chatCompletion == null) { - return false; - } - - var choices = chatCompletion.choices(); - if (CollectionUtils.isEmpty(choices)) { - return false; - } - - var choice = choices.get(0); - return !CollectionUtils.isEmpty(choice.message().toolCalls()) - && choice.finishReason() == ChatCompletionFinishReason.TOOL_CALLS; - } - @Override public ChatOptions getDefaultOptions() { return OpenAiChatOptions.fromOptions(this.defaultOptions); diff --git a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java index 750700175..d2d928896 100644 --- a/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java +++ b/models/spring-ai-vertex-ai-gemini/src/main/java/org/springframework/ai/vertexai/gemini/VertexAiGeminiChatModel.java @@ -72,8 +72,7 @@ import java.util.Set; * @author Chris Turchin * @since 0.8.1 */ -public class VertexAiGeminiChatModel extends AbstractToolCallSupport - implements ChatModel, DisposableBean { +public class VertexAiGeminiChatModel extends AbstractToolCallSupport implements ChatModel, DisposableBean { private final static boolean IS_RUNTIME_CALL = true; @@ -503,7 +502,6 @@ public class VertexAiGeminiChatModel extends AbstractToolCallSupport { - private AssistantMessage assistantMessage; + private final AssistantMessage assistantMessage; private ChatGenerationMetadata chatGenerationMetadata; + /** + * @deprecated Use {@link #Generation(AssitantMessage)} constructor instead. + */ + @Deprecated public Generation(String text) { - this.assistantMessage = new AssistantMessage(text); + this(text, Map.of()); } + /** + * @deprecated Use {@link #Generation(AssitantMessage)} constructor instead. + */ + @Deprecated public Generation(String text, Map properties) { - this.assistantMessage = new AssistantMessage(text, properties); + this(new AssistantMessage(text, properties)); + } + + public Generation(AssistantMessage assistantMessage) { + this(assistantMessage, ChatGenerationMetadata.NULL); + } + + public Generation(AssistantMessage assistantMessage, ChatGenerationMetadata chatGenerationMetadata) { + this.assistantMessage = assistantMessage; + this.chatGenerationMetadata = chatGenerationMetadata; } @Override @@ -51,6 +68,13 @@ public class Generation implements ModelResult { return chatGenerationMetadata != null ? chatGenerationMetadata : ChatGenerationMetadata.NULL; } + /** + * @deprecated Use {@link #Generation(AssitantMessage, ChatGenerationMetadata)} + * constructor instead. + * @param chatGenerationMetadata + * @return + */ + @Deprecated public Generation withGenerationMetadata(@Nullable ChatGenerationMetadata chatGenerationMetadata) { this.chatGenerationMetadata = chatGenerationMetadata; return this; @@ -62,19 +86,19 @@ public class Generation implements ModelResult { return true; if (!(o instanceof Generation that)) return false; - return Objects.equals(assistantMessage, that.assistantMessage) - && Objects.equals(chatGenerationMetadata, that.chatGenerationMetadata); + return Objects.equals(this.assistantMessage, that.assistantMessage) + && Objects.equals(this.chatGenerationMetadata, that.chatGenerationMetadata); } @Override public int hashCode() { - return Objects.hash(assistantMessage, chatGenerationMetadata); + return Objects.hash(this.assistantMessage, this.chatGenerationMetadata); } @Override public String toString() { - return "Generation{" + "assistantMessage=" + assistantMessage + ", chatGenerationMetadata=" - + chatGenerationMetadata + '}'; + return "Generation[" + "assistantMessage=" + this.assistantMessage + ", chatGenerationMetadata=" + + this.chatGenerationMetadata + ']'; } } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractToolCallSupport.java b/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractToolCallSupport.java index db684c65d..b9872cbbb 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractToolCallSupport.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/model/function/AbstractToolCallSupport.java @@ -24,18 +24,20 @@ import java.util.concurrent.ConcurrentHashMap; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.messages.ToolResponseMessage; +import org.springframework.ai.chat.model.ChatResponse; +import org.springframework.util.Assert; import org.springframework.util.CollectionUtils; /** * Abstract base class for tool call support. Provides functionality for handling function * callbacks and executing functions. * - * @param The response type of the tool call. * @author Christian Tzolov * @author Grogdunn + * @author Thomas Vitale * @since 1.0.0 */ -public abstract class AbstractToolCallSupport { +public abstract class AbstractToolCallSupport { protected final static boolean IS_RUNTIME_CALL = true; @@ -148,6 +150,21 @@ public abstract class AbstractToolCallSupport { return new ToolResponseMessage(toolResponses, Map.of()); } - abstract protected boolean isToolFunctionCall(TRes response); + protected boolean isToolCall(ChatResponse chatResponse, String toolCallFinishReason) { + Assert.hasText(toolCallFinishReason, "toolCallFinishReason cannot be null or empty"); + + if (chatResponse == null) { + return false; + } + + var generations = chatResponse.getResults(); + if (CollectionUtils.isEmpty(generations)) { + return false; + } + + var generation = generations.get(0); + return !CollectionUtils.isEmpty(generation.getOutput().getToolCalls()) + && toolCallFinishReason.equalsIgnoreCase(generation.getMetadata().getFinishReason()); + } }