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/model.rs | 196 ++++++++++++++++++++++++---------------------------- 1 file changed, 90 insertions(+), 106 deletions(-) (limited to 'src/client/model.rs') diff --git a/src/client/model.rs b/src/client/model.rs index 1af3e59..4b16ffd 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -10,15 +10,8 @@ const BASIS_TOKENS: usize = 2; #[derive(Debug, Clone)] pub struct Model { - pub client_name: String, - pub name: String, - pub max_input_tokens: Option, - pub max_output_tokens: Option, - pub pass_max_tokens: bool, - pub input_price: Option, - pub output_price: Option, - pub capabilities: ModelCapabilities, - pub extra_fields: Option>, + client_name: String, + data: ModelData, } impl Default for Model { @@ -31,30 +24,16 @@ impl Model { pub fn new(client_name: &str, name: &str) -> Self { Self { client_name: client_name.into(), - name: name.into(), - max_input_tokens: None, - max_output_tokens: None, - pass_max_tokens: false, - input_price: None, - output_price: None, - capabilities: ModelCapabilities::Text, - extra_fields: None, + data: ModelData::new(name), } } - pub fn from_config(client_name: &str, models: &[ModelConfig]) -> Vec { + pub fn from_config(client_name: &str, models: &[ModelData]) -> Vec { models .iter() - .map(|v| { - let mut model = Model::new(client_name, &v.name); - model - .set_max_input_tokens(v.max_input_tokens) - .set_max_tokens(v.max_output_tokens, v.pass_max_tokens) - .set_input_price(v.input_price) - .set_output_price(v.output_price) - .set_supports_vision(v.supports_vision) - .set_extra_fields(&v.extra_fields); - model + .map(|v| Model { + client_name: client_name.to_string(), + data: v.clone(), }) .collect() } @@ -77,7 +56,7 @@ impl Model { model = Some((*found).clone()); } else if let Some(found) = models.iter().find(|v| v.client_name == client_name) { let mut found = (*found).clone(); - found.name = model_name.to_string(); + found.data.name = model_name.to_string(); model = Some(found) } } @@ -91,99 +70,101 @@ impl Model { } pub fn id(&self) -> String { - format!("{}:{}", self.client_name, self.name) + format!("{}:{}", self.client_name, self.data.name) + } + + pub fn client_name(&self) -> &str { + &self.client_name + } + + pub fn name(&self) -> &str { + &self.data.name + } + + pub fn data(&self) -> &ModelData { + &self.data + } + + pub fn data_mut(&mut self) -> &mut ModelData { + &mut self.data } pub fn description(&self) -> String { - let max_input_tokens = format_option_value(&self.max_input_tokens); - let max_output_tokens = format_option_value(&self.max_output_tokens); - let input_price = format_option_value(&self.input_price); - let output_price = format_option_value(&self.output_price); - let vision = if self.capabilities.contains(ModelCapabilities::Vision) { - "👁" - } else { - "" + let ModelData { + max_input_tokens, + max_output_tokens, + input_price, + output_price, + supports_vision, + supports_function_calling, + .. + } = &self.data; + let max_input_tokens = format_option_value(max_input_tokens); + let max_output_tokens = format_option_value(max_output_tokens); + let input_price = format_option_value(input_price); + let output_price = format_option_value(output_price); + let mut capabilities = vec![]; + if *supports_vision { + capabilities.push('👁'); + }; + if *supports_function_calling { + capabilities.push('⚒'); }; + let capabilities: String = capabilities + .into_iter() + .map(|v| format!("{v} ")) + .collect::>() + .join(""); format!( - "{:>8} / {:>8} | {:>6} / {:>6} {}", - max_input_tokens, max_output_tokens, input_price, output_price, vision + "{:>8} / {:>8} | {:>6} / {:>6} {:>6}", + max_input_tokens, max_output_tokens, input_price, output_price, capabilities ) } + pub fn max_input_tokens(&self) -> Option { + self.data.max_input_tokens + } + + pub fn max_output_tokens(&self) -> Option { + self.data.max_output_tokens + } + pub fn supports_vision(&self) -> bool { - self.capabilities.contains(ModelCapabilities::Vision) + self.data.supports_vision + } + + pub fn supports_function_calling(&self) -> bool { + self.data.supports_function_calling } pub fn max_tokens_param(&self) -> Option { - if self.pass_max_tokens { - self.max_output_tokens + if self.data.pass_max_tokens { + self.data.max_output_tokens } else { None } } - pub fn set_max_input_tokens(&mut self, max_input_tokens: Option) -> &mut Self { - match max_input_tokens { - None | Some(0) => self.max_input_tokens = None, - _ => self.max_input_tokens = max_input_tokens, - } - self - } - pub fn set_max_tokens( &mut self, max_output_tokens: Option, pass_max_tokens: bool, ) -> &mut Self { match max_output_tokens { - None | Some(0) => self.max_output_tokens = None, - _ => self.max_output_tokens = max_output_tokens, - } - self.pass_max_tokens = pass_max_tokens; - self - } - - pub fn set_input_price(&mut self, input_price: Option) -> &mut Self { - match input_price { - None => self.input_price = None, - _ => self.input_price = input_price, - } - self - } - - pub fn set_output_price(&mut self, output_price: Option) -> &mut Self { - match output_price { - None => self.output_price = None, - _ => self.output_price = output_price, + None | Some(0) => self.data.max_output_tokens = None, + _ => self.data.max_output_tokens = max_output_tokens, } - self - } - - pub fn set_supports_vision(&mut self, supports_vision: bool) -> &mut Self { - if supports_vision { - self.capabilities |= ModelCapabilities::Vision; - } else { - self.capabilities &= !ModelCapabilities::Vision; - } - self - } - - pub fn set_extra_fields( - &mut self, - extra_fields: &Option>, - ) -> &mut Self { - self.extra_fields.clone_from(extra_fields); + self.data.pass_max_tokens = pass_max_tokens; self } pub fn messages_tokens(&self, messages: &[Message]) -> usize { messages .iter() - .map(|v| { - match &v.content { - MessageContent::Text(text) => estimate_token_length(text), - MessageContent::Array(_) => 0, // TODO - } + .map(|v| match &v.content { + MessageContent::Text(text) => estimate_token_length(text), + MessageContent::Array(_) => 0, + MessageContent::ToolResults(_) => 0, }) .sum() } @@ -203,7 +184,7 @@ impl Model { pub fn max_input_tokens_limit(&self, messages: &[Message]) -> Result<()> { let total_tokens = self.total_tokens(messages) + BASIS_TOKENS; - if let Some(max_input_tokens) = self.max_input_tokens { + if let Some(max_input_tokens) = self.data.max_input_tokens { if total_tokens >= max_input_tokens { bail!("Exceed max input tokens limit") } @@ -212,7 +193,7 @@ impl Model { } pub fn merge_extra_fields(&self, body: &mut serde_json::Value) { - if let (Some(body), Some(extra_fields)) = (body.as_object_mut(), &self.extra_fields) { + if let (Some(body), Some(extra_fields)) = (body.as_object_mut(), &self.data.extra_fields) { for (key, extra_field) in extra_fields { if body.contains_key(key) { if let (Some(sub_body), Some(extra_field)) = @@ -232,30 +213,33 @@ impl Model { } } -#[derive(Debug, Clone, Deserialize)] -pub struct ModelConfig { +#[derive(Debug, Clone, Default, Deserialize)] +pub struct ModelData { pub name: String, pub max_input_tokens: Option, pub max_output_tokens: Option, + #[serde(default)] + pub pass_max_tokens: bool, pub input_price: Option, pub output_price: Option, #[serde(default)] pub supports_vision: bool, #[serde(default)] - pub pass_max_tokens: bool, + pub supports_function_calling: bool, pub extra_fields: Option>, } +impl ModelData { + pub fn new(name: &str) -> Self { + Self { + name: name.to_string(), + ..Default::default() + } + } +} + #[derive(Debug, Clone, Deserialize)] pub struct BuiltinModels { pub platform: String, - pub models: Vec, -} - -bitflags::bitflags! { - #[derive(Debug, Clone, Copy, PartialEq)] - pub struct ModelCapabilities: u32 { - const Text = 0b00000001; - const Vision = 0b00000010; - } + pub models: Vec, } -- cgit v1.2.3