From b4a40e3fedb438570770a224b890ea24f6e660a9 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 18 May 2024 19:06:21 +0800 Subject: feat: support function calling (#514) * feat: support function calling * fix on Windows OS * implement multi-steps function calling * fix on Windows OS * add error for client not support function calling * refactor message data structure and make claude client supporting function calling * support reuse previous call results * improve error handling for function calling * use prefix `may_` as indicator for `execute` type fucntions --- src/client/message.rs | 23 +++++++++++++++-------- 1 file changed, 15 insertions(+), 8 deletions(-) (limited to 'src/client/message.rs') diff --git a/src/client/message.rs b/src/client/message.rs index 9621811..d7ba698 100644 --- a/src/client/message.rs +++ b/src/client/message.rs @@ -1,4 +1,4 @@ -use crate::config::Input; +use super::ToolResults; use serde::{Deserialize, Serialize}; @@ -8,15 +8,21 @@ pub struct Message { pub content: MessageContent, } -impl Message { - pub fn new(input: &Input) -> Self { +impl Default for Message { + fn default() -> Self { Self { role: MessageRole::User, - content: input.to_message_content(), + content: MessageContent::Text(String::new()), } } } +impl Message { + pub fn new(role: MessageRole, content: MessageContent) -> Self { + Self { role, content } + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize)] #[serde(rename_all = "snake_case")] pub enum MessageRole { @@ -34,10 +40,6 @@ impl MessageRole { pub fn is_user(&self) -> bool { matches!(self, MessageRole::User) } - - pub fn is_assistant(&self) -> bool { - matches!(self, MessageRole::Assistant) - } } #[derive(Debug, Clone, Deserialize, Serialize)] @@ -45,6 +47,8 @@ impl MessageRole { pub enum MessageContent { Text(String), Array(Vec), + // Note: This type is primarily for convenience and does not exist in OpenAI's API. + ToolResults(ToolResults), } impl MessageContent { @@ -68,6 +72,7 @@ impl MessageContent { } format!(".file {}{}", files.join(" "), concated_text) } + MessageContent::ToolResults(_) => String::new(), } } @@ -83,6 +88,7 @@ impl MessageContent { *text = replace_fn(text) } } + MessageContent::ToolResults(_) => {} } } @@ -98,6 +104,7 @@ impl MessageContent { } parts.join("\n\n") } + MessageContent::ToolResults(_) => String::new(), } } } -- cgit v1.2.3