From 80684ec84047bda0299f8fd845b1c97f6fac5c99 Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 14 Nov 2024 08:14:32 +0800 Subject: refactor: improve tool calls (#995) - rename MessageContent:ToolResults to MessageContent:ToolCalls - rename ToolResults to MessageContentToolCalls - persist tool_calls to messages.md --- src/config/input.rs | 26 +++++++++++++------------- src/config/mod.rs | 18 ++++++++++++++++-- src/config/session.rs | 4 ++-- 3 files changed, 31 insertions(+), 17 deletions(-) (limited to 'src/config') diff --git a/src/config/input.rs b/src/config/input.rs index 55a4d45..18b192b 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -2,9 +2,9 @@ use super::*; use crate::client::{ init_client, patch_system_message, ChatCompletionsData, Client, ImageUrl, Message, - MessageContent, MessageContentPart, MessageRole, Model, + MessageContent, MessageContentPart, MessageContentToolCalls, MessageRole, Model, }; -use crate::function::{ToolResult, ToolResults}; +use crate::function::ToolResult; use crate::utils::{base64_encode, sha256, AbortSignal}; use anyhow::{bail, Context, Result}; @@ -29,7 +29,7 @@ pub struct Input { regenerate: bool, medias: Vec, data_urls: HashMap, - tool_results: Option, + tool_calls: Option, rag_name: Option, role: Role, with_session: bool, @@ -48,7 +48,7 @@ impl Input { regenerate: false, medias: Default::default(), data_urls: Default::default(), - tool_results: None, + tool_calls: None, rag_name: None, role, with_session, @@ -104,7 +104,7 @@ impl Input { regenerate: false, medias, data_urls, - tool_results: Default::default(), + tool_calls: Default::default(), rag_name: None, role, with_session, @@ -120,8 +120,8 @@ impl Input { self.data_urls.clone() } - pub fn tool_results(&self) -> &Option { - &self.tool_results + pub fn tool_calls(&self) -> &Option { + &self.tool_calls } pub fn text(&self) -> String { @@ -187,12 +187,12 @@ impl Input { self.rag_name.as_deref() } - pub fn merge_tool_call(mut self, output: String, tool_results: Vec) -> Self { - match self.tool_results.as_mut() { + pub fn merge_tool_results(mut self, output: String, tool_results: Vec) -> Self { + match self.tool_calls.as_mut() { Some(exist_tool_results) => { - exist_tool_results.extend(tool_results, output); + exist_tool_results.merge(tool_results, output); } - None => self.tool_results = Some(ToolResults::new(tool_results, output)), + None => self.tool_calls = Some(MessageContentToolCalls::new(tool_results, output)), } self } @@ -232,10 +232,10 @@ impl Input { } else { self.role().build_messages(self) }; - if let Some(tool_results) = &self.tool_results { + if let Some(tool_calls) = &self.tool_calls { messages.push(Message::new( MessageRole::Assistant, - MessageContent::ToolResults(tool_results.clone()), + MessageContent::ToolCalls(tool_calls.clone()), )) } Ok(messages) diff --git a/src/config/mod.rs b/src/config/mod.rs index 9da8dae..136fee1 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -10,7 +10,7 @@ use self::session::Session; use crate::client::{ create_client_config, list_chat_models, list_client_types, list_reranker_models, ClientConfig, - Model, OPENAI_COMPATIBLE_PLATFORMS, + MessageContentToolCalls, Model, OPENAI_COMPATIBLE_PLATFORMS, }; use crate::function::{FunctionDeclaration, Functions, ToolResult}; use crate::rag::Rag; @@ -1863,8 +1863,22 @@ impl Config { } else { String::new() }; + let tool_calls = match input.tool_calls() { + Some(MessageContentToolCalls { + tool_results, text, .. + }) => { + let mut lines = vec!["".to_string()]; + if !text.is_empty() { + lines.push(text.clone()); + } + lines.push(serde_json::to_string(&tool_results).unwrap_or_default()); + lines.push("\n".to_string()); + lines.join("\n") + } + None => String::new(), + }; let output = format!( - "# CHAT: {summary} [{timestamp}]{scope}\n{raw_input}\n--------\n{output}\n--------\n\n", + "# CHAT: {summary} [{timestamp}]{scope}\n{raw_input}\n--------\n{tool_calls}{output}\n--------\n\n", ); file.write_all(output.as_bytes()) .with_context(|| "Failed to save message") diff --git a/src/config/session.rs b/src/config/session.rs index 34d6393..633f621 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -419,10 +419,10 @@ impl Session { .push(Message::new(MessageRole::User, input.message_content())); } self.data_urls.extend(input.data_urls()); - if let Some(tool_results) = input.tool_results() { + if let Some(tool_calls) = input.tool_calls() { self.messages.push(Message::new( MessageRole::Tool, - MessageContent::ToolResults(tool_results.clone()), + MessageContent::ToolCalls(tool_calls.clone()), )) } self.messages.push(Message::new( -- cgit v1.2.3