From 0eacc9193b4d25f88fb2745a197933d36e6b54cd Mon Sep 17 00:00:00 2001 From: Christian Tzolov Date: Mon, 18 Nov 2024 11:02:36 +0100 Subject: [PATCH] feat(tools): Add tool call history to ToolContext - Add TOOL_CALL_HISTORY constant to store tool call history - Extend ToolContext with getToolConversationHistory method - Include tool call history in tool context during function execution - Add test coverage for tool call history verification Resolves #1202 --- ...lientMethodInvokingFunctionCallbackIT.java | 6 ++++++ .../chat/model/AbstractToolCallSupport.java | 16 ++++++++++++-- .../ai/chat/model/ToolContext.java | 21 +++++++++++++++++++ 3 files changed, 41 insertions(+), 2 deletions(-) diff --git a/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/client/AnthropicChatClientMethodInvokingFunctionCallbackIT.java b/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/client/AnthropicChatClientMethodInvokingFunctionCallbackIT.java index 097eaf153..6ed41d3d0 100644 --- a/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/client/AnthropicChatClientMethodInvokingFunctionCallbackIT.java +++ b/models/spring-ai-anthropic/src/test/java/org/springframework/ai/anthropic/client/AnthropicChatClientMethodInvokingFunctionCallbackIT.java @@ -16,6 +16,7 @@ package org.springframework.ai.anthropic.client; +import java.util.List; import java.util.Map; import java.util.concurrent.ConcurrentHashMap; @@ -27,6 +28,7 @@ import org.slf4j.LoggerFactory; import org.springframework.ai.anthropic.AnthropicTestConfiguration; import org.springframework.ai.chat.client.ChatClient; +import org.springframework.ai.chat.messages.Message; import org.springframework.ai.chat.model.ChatModel; import org.springframework.ai.chat.model.ToolContext; import org.springframework.ai.model.function.FunctionCallback; @@ -158,6 +160,9 @@ class AnthropicChatClientMethodInvokingFunctionCallbackIT { assertThat(response).contains("30", "10", "15"); assertThat(arguments).containsEntry("tool", "value"); + assertThat(arguments).containsKey(ToolContext.TOOL_CALL_HISTORY); + List tootConversationMessages = (List) arguments.get(ToolContext.TOOL_CALL_HISTORY); + assertThat(tootConversationMessages).hasSize(6); } @Test @@ -252,6 +257,7 @@ class AnthropicChatClientMethodInvokingFunctionCallbackIT { public String getWeatherWithContext(String city, Unit unit, ToolContext context) { arguments.put("tool", context.getContext().get("tool")); + arguments.put(ToolContext.TOOL_CALL_HISTORY, context.getToolCallHistory()); return getWeatherStatic(city, unit); } diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/model/AbstractToolCallSupport.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/model/AbstractToolCallSupport.java index b78a8ba7b..255444944 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/model/AbstractToolCallSupport.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/model/AbstractToolCallSupport.java @@ -17,6 +17,7 @@ package org.springframework.ai.chat.model; import java.util.ArrayList; +import java.util.HashMap; import java.util.HashSet; import java.util.List; import java.util.Map; @@ -142,12 +143,23 @@ public abstract class AbstractToolCallSupport { Map toolContextMap = Map.of(); if (prompt.getOptions() instanceof FunctionCallingOptions functionCallOptions && !CollectionUtils.isEmpty(functionCallOptions.getToolContext())) { - toolContextMap = functionCallOptions.getToolContext(); + + toolContextMap = new HashMap<>(functionCallOptions.getToolContext()); + + List toolCallHistory = new ArrayList<>(prompt.copy().getInstructions()); + toolCallHistory.add(new AssistantMessage(assistantMessage.getContent(), assistantMessage.getMetadata(), + assistantMessage.getToolCalls())); + + toolContextMap.put(ToolContext.TOOL_CALL_HISTORY, toolCallHistory); } + ToolResponseMessage toolMessageResponse = this.executeFunctions(assistantMessage, new ToolContext(toolContextMap)); - return this.buildToolCallConversation(prompt.getInstructions(), assistantMessage, toolMessageResponse); + List toolConversationHistory = this.buildToolCallConversation(prompt.getInstructions(), + assistantMessage, toolMessageResponse); + + return toolConversationHistory; } protected List buildToolCallConversation(List previousMessages, AssistantMessage assistantMessage, diff --git a/spring-ai-core/src/main/java/org/springframework/ai/chat/model/ToolContext.java b/spring-ai-core/src/main/java/org/springframework/ai/chat/model/ToolContext.java index 5ba3a60eb..69b349383 100644 --- a/spring-ai-core/src/main/java/org/springframework/ai/chat/model/ToolContext.java +++ b/spring-ai-core/src/main/java/org/springframework/ai/chat/model/ToolContext.java @@ -17,8 +17,11 @@ package org.springframework.ai.chat.model; import java.util.Collections; +import java.util.List; import java.util.Map; +import org.springframework.ai.chat.messages.Message; + /** * Represents the context for tool execution in a function calling scenario. * @@ -33,11 +36,20 @@ import java.util.Map; * {@code FunctionCallingOptions} and is used in the function execution process. *

* + *

+ * The context map can contain any information that is relevant to the tool execution. + *

+ * * @author Christian Tzolov * @since 1.0.0 */ public class ToolContext { + /** + * The key for the running, tool call history stored in the context map. + */ + public static final String TOOL_CALL_HISTORY = "TOOL_CALL_HISTORY"; + private final Map context; /** @@ -57,4 +69,13 @@ public class ToolContext { return this.context; } + /** + * Returns the tool call history from the context map. + * @return The tool call history. + */ + @SuppressWarnings("unchecked") + public List getToolCallHistory() { + return (List) this.context.get(TOOL_CALL_HISTORY); + } + }