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/bedrock.rs | 7 +++---- src/client/claude.rs | 7 +++---- src/client/cohere.rs | 4 ++-- src/client/ernie.rs | 6 ++++-- src/client/message.rs | 40 ++++++++++++++++++++++++++++++---------- src/client/mod.rs | 2 +- src/client/model.rs | 9 ++++----- src/client/openai.rs | 5 ++--- src/client/vertexai.rs | 3 +-- 9 files changed, 50 insertions(+), 33 deletions(-) (limited to 'src/client') diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs index e1d6c65..78f9d2b 100644 --- a/src/client/bedrock.rs +++ b/src/client/bedrock.rs @@ -363,10 +363,9 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu "content": content, })] } - MessageContent::ToolResults(results) => { - let ToolResults { - tool_results, text, .. - } = results; + MessageContent::ToolCalls(MessageContentToolCalls { + tool_results, text, .. + }) => { let mut assistant_parts = vec![]; let mut user_parts = vec![]; if !text.is_empty() { diff --git a/src/client/claude.rs b/src/client/claude.rs index 0c905b1..c7ef88a 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -202,10 +202,9 @@ pub fn claude_build_chat_completions_body( "content": content, })] } - MessageContent::ToolResults(results) => { - let ToolResults { - tool_results, text, .. - } = results; + MessageContent::ToolCalls(MessageContentToolCalls { + tool_results, text, .. + }) => { let mut assistant_parts = vec![]; let mut user_parts = vec![]; if !text.is_empty() { diff --git a/src/client/cohere.rs b/src/client/cohere.rs index c5a31f5..9726300 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -213,8 +213,8 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu .collect(); Some(json!({ "role": role, "message": list.join("\n\n") })) } - MessageContent::ToolResults(results) => { - tool_results = Some(results.tool_results); + MessageContent::ToolCalls(tool_calls) => { + tool_results = Some(tool_calls.tool_results); None } } diff --git a/src/client/ernie.rs b/src/client/ernie.rs index 0c325e5..7e1f8e7 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -231,9 +231,11 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Valu .flat_map(|message| { let Message { role, content } = message; match content { - MessageContent::ToolResults(results) => { + MessageContent::ToolCalls(MessageContentToolCalls { + tool_results, .. + }) => { let mut list = vec![]; - for tool_result in results.tool_results { + for tool_result in tool_results { list.push(json!({ "role": "assistant", "content": format!("Action: {}\nAction Input: {}", tool_result.call.name, tool_result.call.arguments) 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); diff --git a/src/client/mod.rs b/src/client/mod.rs index b22508b..5189f9f 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -6,7 +6,7 @@ mod macros; mod model; mod stream; -pub use crate::function::{ToolCall, ToolResults}; +pub use crate::function::ToolCall; pub use crate::utils::PromptKind; pub use common::*; pub use message::*; diff --git a/src/client/model.rs b/src/client/model.rs index af864e2..5d496d2 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -1,7 +1,7 @@ use super::{ list_chat_models, list_embedding_models, list_reranker_models, message::{Message, MessageContent, MessageContentPart}, - ToolResults, + MessageContentToolCalls, }; use crate::config::Config; @@ -237,10 +237,9 @@ impl Model { MessageContentPart::ImageUrl { .. } => 0, }) .sum(), - MessageContent::ToolResults(results) => { - let ToolResults { - tool_results, text, .. - } = results; + MessageContent::ToolCalls(MessageContentToolCalls { + tool_results, text, .. + }) => { estimate_token_length(text) + tool_results .iter() diff --git a/src/client/openai.rs b/src/client/openai.rs index 0f01e36..7f449fb 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -205,12 +205,11 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod .flat_map(|message| { let Message { role, content } = message; match content { - MessageContent::ToolResults(results) => { - let ToolResults { + MessageContent::ToolCalls(MessageContentToolCalls { tool_results, text, sequence, - } = results; + }) => { if !sequence { let tool_calls: Vec<_> = tool_results.iter().map(|tool_result| { json!({ diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index 452eb09..07ff2a4 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -342,8 +342,7 @@ pub fn gemini_build_chat_completions_body( .collect(); vec![json!({ "role": role, "parts": parts })] }, - MessageContent::ToolResults(results) => { - let tool_results = results.tool_results; + MessageContent::ToolCalls(MessageContentToolCalls { tool_results, .. }) => { let model_parts: Vec = tool_results.iter().map(|tool_result| { json!({ "functionCall": { -- cgit v1.2.3