diff options
| author | sigoden <sigoden@gmail.com> | 2024-05-18 19:06:21 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-18 19:06:21 +0800 |
| commit | b4a40e3fedb438570770a224b890ea24f6e660a9 (patch) | |
| tree | 344b96102da7cbedf1034d023aa82599940388b1 /src/config/mod.rs | |
| parent | 1348a62e5f8bc140a7218fbfe1b73f990ab16101 (diff) | |
| download | aichat-b4a40e3fedb438570770a224b890ea24f6e660a9.tar.gz | |
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
Diffstat (limited to 'src/config/mod.rs')
| -rw-r--r-- | src/config/mod.rs | 52 |
1 files changed, 40 insertions, 12 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs index 66abac2..d59a81e 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -10,6 +10,7 @@ use crate::client::{ create_client_config, list_client_types, list_models, ClientConfig, Model, OPENAI_COMPATIBLE_PLATFORMS, }; +use crate::function::{Function, ToolCallResult}; use crate::render::{MarkdownRender, RenderOptions}; use crate::utils::{ format_option_value, fuzzy_match, get_env_name, light_theme_from_colorfgbg, now, render_prompt, @@ -41,6 +42,7 @@ const CONFIG_FILE_NAME: &str = "config.yaml"; const ROLES_FILE_NAME: &str = "roles.yaml"; const MESSAGES_FILE_NAME: &str = "messages.md"; const SESSIONS_DIR_NAME: &str = "sessions"; +const FUNCTIONS_DIR_NAME: &str = "functions"; const CLIENTS_FIELD: &str = "clients"; @@ -69,6 +71,7 @@ pub struct Config { pub keybindings: Keybindings, pub prelude: Option<String>, pub buffer_editor: Option<String>, + pub function_calling: bool, pub compress_threshold: usize, pub summarize_prompt: Option<String>, pub summary_prompt: Option<String>, @@ -84,6 +87,8 @@ pub struct Config { #[serde(skip)] pub model: Model, #[serde(skip)] + pub function: Function, + #[serde(skip)] pub working_mode: WorkingMode, #[serde(skip)] pub last_message: Option<(Input, String)>, @@ -106,6 +111,7 @@ impl Default for Config { keybindings: Default::default(), prelude: None, buffer_editor: None, + function_calling: false, compress_threshold: 2000, summarize_prompt: None, summary_prompt: None, @@ -116,6 +122,7 @@ impl Default for Config { role: None, session: None, model: Default::default(), + function: Default::default(), working_mode: WorkingMode::Command, last_message: None, } @@ -142,6 +149,8 @@ impl Config { config.set_wrap(&wrap)?; } + config.function = Function::init(&Self::functions_dir()?)?; + config.working_mode = working_mode; config.load_roles()?; @@ -212,15 +221,20 @@ impl Config { Ok(path) } - pub fn save_message(&mut self, input: Input, output: &str) -> Result<()> { + pub fn save_message( + &mut self, + input: &Input, + output: &str, + tool_call_results: &[ToolCallResult], + ) -> Result<()> { self.last_message = Some((input.clone(), output.to_string())); - if self.dry_run { + if self.dry_run || output.is_empty() || !tool_call_results.is_empty() { return Ok(()); } if let Some(session) = input.session_mut(&mut self.session) { - session.add_message(&input, output)?; + session.add_message(input, output)?; return Ok(()); } @@ -275,6 +289,10 @@ impl Config { Self::local_path(SESSIONS_DIR_NAME) } + pub fn functions_dir() -> Result<PathBuf> { + Self::local_path(FUNCTIONS_DIR_NAME) + } + pub fn session_file(name: &str) -> Result<PathBuf> { let mut path = Self::sessions_dir()?; path.push(&format!("{name}.yaml")); @@ -294,8 +312,7 @@ impl Config { pub fn set_role_obj(&mut self, role: Role) -> Result<()> { if let Some(session) = self.session.as_mut() { session.guard_empty()?; - session.set_temperature(role.temperature); - session.set_top_p(role.top_p); + session.set_role_properties(&role); } if let Some(model_id) = &role.model_id { self.set_model(model_id)?; @@ -428,6 +445,8 @@ impl Config { ), ("temperature", format_option_value(&temperature)), ("top_p", format_option_value(&top_p)), + ("function_calling", self.function_calling.to_string()), + ("compress_threshold", self.compress_threshold.to_string()), ("dry_run", self.dry_run.to_string()), ("save", self.save.to_string()), ("save_session", format_option_value(&self.save_session)), @@ -438,11 +457,11 @@ impl Config { ("auto_copy", self.auto_copy.to_string()), ("keybindings", self.keybindings.stringify().into()), ("prelude", format_option_value(&self.prelude)), - ("compress_threshold", self.compress_threshold.to_string()), ("config_file", display_path(&Self::config_file()?)), ("roles_file", display_path(&Self::roles_file()?)), ("messages_file", display_path(&Self::messages_file()?)), ("sessions_dir", display_path(&Self::sessions_dir()?)), + ("functions_dir", display_path(&Self::functions_dir()?)), ]; let output = items .iter() @@ -508,6 +527,7 @@ impl Config { "max_output_tokens", "temperature", "top_p", + "function_calling", "compress_threshold", "save", "save_session", @@ -523,10 +543,11 @@ impl Config { (values, args[0]) } else if args.len() == 2 { let values = match args[0] { - "max_output_tokens" => match self.model.max_output_tokens { + "max_output_tokens" => match self.model.max_output_tokens() { Some(v) => vec![v.to_string()], None => vec![], }, + "function_calling" => complete_bool(self.function_calling), "save" => complete_bool(self.save), "save_session" => { let save_session = if let Some(session) = &self.session { @@ -574,6 +595,10 @@ impl Config { let value = parse_value(value)?; self.set_top_p(value); } + "function_calling" => { + let value = value.parse().with_context(|| "Invalid value")?; + self.function_calling = value; + } "compress_threshold" => { let value = parse_value(value)?; self.set_compress_threshold(value); @@ -792,11 +817,14 @@ impl Config { fn generate_prompt_context(&self) -> HashMap<&str, String> { let mut output = HashMap::new(); output.insert("model", self.model.id()); - output.insert("client_name", self.model.client_name.clone()); - output.insert("model_name", self.model.name.clone()); + output.insert("client_name", self.model.client_name().to_string()); + output.insert("model_name", self.model.name().to_string()); output.insert( "max_input_tokens", - self.model.max_input_tokens.unwrap_or_default().to_string(), + self.model + .max_input_tokens() + .unwrap_or_default() + .to_string(), ); if let Some(temperature) = self.temperature { if temperature != 0.0 { @@ -884,8 +912,8 @@ impl Config { } fn load_config_file(config_path: &Path) -> Result<Self> { - let ctx = || format!("Failed to load config at {}", config_path.display()); - let content = read_to_string(config_path).with_context(ctx)?; + let content = read_to_string(config_path) + .with_context(|| format!("Failed to load config at {}", config_path.display()))?; let config: Self = serde_yaml::from_str(&content).map_err(|err| { let err_msg = err.to_string(); let err_msg = if err_msg.starts_with(&format!("{}: ", CLIENTS_FIELD)) { |
