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/client/message.rs | 40 ++++++++++++++++++++++++++++++---------- 1 file changed, 30 insertions(+), 10 deletions(-) (limited to 'src/client/message.rs') diff --git a/src/client/message.rs b/src/client/message.rs index 061c0e2..f6b7f6d 100644 --- a/src/client/message.rs +++ b/src/client/message.rs @@ -1,6 +1,4 @@ -use super::ToolResults; - -use crate::utils::dimmed_text; +use crate::{function::ToolResult, utils::dimmed_text}; use serde::{Deserialize, Serialize}; @@ -75,7 +73,7 @@ pub enum MessageContent { Text(String), Array(Vec), // Note: This type is primarily for convenience and does not exist in OpenAI's API. - ToolResults(ToolResults), + ToolCalls(MessageContentToolCalls), } impl MessageContent { @@ -103,10 +101,9 @@ impl MessageContent { } format!(".file {}{}", files.join(" "), concated_text) } - MessageContent::ToolResults(results) => { - let ToolResults { - tool_results, text, .. - } = results; + MessageContent::ToolCalls(MessageContentToolCalls { + tool_results, text, .. + }) => { let mut lines = vec![]; if !text.is_empty() { lines.push(text.clone()) @@ -139,7 +136,7 @@ impl MessageContent { *text = replace_fn(text) } } - MessageContent::ToolResults(_) => {} + MessageContent::ToolCalls(_) => {} } } @@ -155,7 +152,7 @@ impl MessageContent { } parts.join("\n\n") } - MessageContent::ToolResults(_) => String::new(), + MessageContent::ToolCalls(_) => String::new(), } } } @@ -172,6 +169,29 @@ pub struct ImageUrl { pub url: String, } +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct MessageContentToolCalls { + pub tool_results: Vec, + pub text: String, + pub sequence: bool, +} + +impl MessageContentToolCalls { + pub fn new(tool_results: Vec, text: String) -> Self { + Self { + tool_results, + text, + sequence: false, + } + } + + pub fn merge(&mut self, tool_results: Vec, _text: String) { + self.tool_results.extend(tool_results); + self.text.clear(); + self.sequence = true; + } +} + pub fn patch_system_message(messages: &mut Vec) { if messages[0].role.is_system() { let system_message = messages.remove(0); -- cgit v1.2.3