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/message.rs | |
| 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/message.rs')
| -rw-r--r-- | src/client/message.rs | 40 |
1 files changed, 30 insertions, 10 deletions
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); |
