diff options
| author | sigoden <sigoden@gmail.com> | 2024-11-14 08:14:32 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-11-14 08:14:32 +0800 |
| commit | 80684ec84047bda0299f8fd845b1c97f6fac5c99 (patch) | |
| tree | f37e377accd664d31a2bfe977521cfc2157a9cd1 /src/client | |
| parent | cfa9217422dfa6cf10bed6e6e3fab0722f58588d (diff) | |
| download | aichat-80684ec84047bda0299f8fd845b1c97f6fac5c99.tar.gz | |
refactor: improve tool calls (#995)
- rename MessageContent:ToolResults to MessageContent:ToolCalls
- rename ToolResults to MessageContentToolCalls
- persist tool_calls to messages.md
Diffstat (limited to 'src/client')
| -rw-r--r-- | src/client/bedrock.rs | 7 | ||||
| -rw-r--r-- | src/client/claude.rs | 7 | ||||
| -rw-r--r-- | src/client/cohere.rs | 4 | ||||
| -rw-r--r-- | src/client/ernie.rs | 6 | ||||
| -rw-r--r-- | src/client/message.rs | 40 | ||||
| -rw-r--r-- | src/client/mod.rs | 2 | ||||
| -rw-r--r-- | src/client/model.rs | 9 | ||||
| -rw-r--r-- | src/client/openai.rs | 5 | ||||
| -rw-r--r-- | src/client/vertexai.rs | 3 |
9 files changed, 50 insertions, 33 deletions
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<MessageContentPart>), // 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<ToolResult>, + pub text: String, + pub sequence: bool, +} + +impl MessageContentToolCalls { + pub fn new(tool_results: Vec<ToolResult>, text: String) -> Self { + Self { + tool_results, + text, + sequence: false, + } + } + + pub fn merge(&mut self, tool_results: Vec<ToolResult>, _text: String) { + self.tool_results.extend(tool_results); + self.text.clear(); + self.sequence = true; + } +} + pub fn patch_system_message(messages: &mut Vec<Message>) { 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<Value> = tool_results.iter().map(|tool_result| { json!({ "functionCall": { |
