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);
+ }
+
}