diff options
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/input.rs | 62 | ||||
| -rw-r--r-- | src/config/mod.rs | 52 | ||||
| -rw-r--r-- | src/config/role.rs | 64 | ||||
| -rw-r--r-- | src/config/session.rs | 80 |
4 files changed, 176 insertions, 82 deletions
diff --git a/src/config/input.rs b/src/config/input.rs index 0a3201a..78da9f2 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -1,9 +1,10 @@ use super::{role::Role, session::Session, GlobalConfig}; use crate::client::{ - init_client, list_models, Client, ImageUrl, Message, MessageContent, MessageContentPart, Model, - ModelCapabilities, SendData, + init_client, list_models, Client, ImageUrl, Message, MessageContent, MessageContentPart, + MessageRole, Model, SendData, }; +use crate::function::{ToolCallResult, ToolResults}; use crate::utils::{base64_encode, sha256}; use anyhow::{bail, Context, Result}; @@ -30,6 +31,7 @@ pub struct Input { text: String, medias: Vec<String>, data_urls: HashMap<String, String>, + tool_call: Option<ToolResults>, context: InputContext, } @@ -40,6 +42,7 @@ impl Input { text: text.to_string(), medias: Default::default(), data_urls: Default::default(), + tool_call: None, context: context.unwrap_or_else(|| InputContext::from_config(config)), } } @@ -91,6 +94,7 @@ impl Input { text: texts.join("\n"), medias, data_urls, + tool_call: Default::default(), context: context.unwrap_or_else(|| InputContext::from_config(config)), }) } @@ -111,6 +115,21 @@ impl Input { self.text = text; } + pub fn merge_tool_call( + mut self, + output: String, + tool_call_results: Vec<ToolCallResult>, + ) -> Self { + match self.tool_call.as_mut() { + Some(exist_tool_call_results) => { + exist_tool_call_results.0.extend(tool_call_results); + exist_tool_call_results.1 = output; + } + None => self.tool_call = Some((tool_call_results, output)), + } + self + } + pub fn model(&self) -> Model { let model = self.config.read().model.clone(); if let Some(model_id) = self.role().and_then(|v| v.model_id.clone()) { @@ -130,7 +149,10 @@ impl Input { init_client(&self.config, Some(self.model())) } - pub fn prepare_send_data(&self, stream: bool) -> Result<SendData> { + pub fn prepare_send_data(&self, model: &Model, stream: bool) -> Result<SendData> { + if !self.medias.is_empty() && !model.supports_vision() { + bail!("The current model does not support vision."); + } let messages = self.build_messages()?; self.config.read().model.max_input_tokens_limit(&messages)?; let (temperature, top_p) = if let Some(session) = self.session(&self.config.read().session) @@ -142,23 +164,41 @@ impl Input { let config = self.config.read(); (config.temperature, config.top_p) }; + let mut functions = None; + if self.config.read().function_calling && model.supports_function_calling() { + let config = self.config.read(); + let function_filter = if let Some(session) = self.session(&config.session) { + session.function_filter() + } else if let Some(role) = self.role() { + role.function_filter.as_deref() + } else { + None + }; + functions = config.function.filtered_declarations(function_filter); + }; Ok(SendData { messages, temperature, top_p, + functions, stream, }) } pub fn build_messages(&self) -> Result<Vec<Message>> { - let messages = if let Some(session) = self.session(&self.config.read().session) { + let mut messages = if let Some(session) = self.session(&self.config.read().session) { session.build_messages(self) } else if let Some(role) = self.role() { role.build_messages(self) } else { - let message = Message::new(self); - vec![message] + vec![Message::new(MessageRole::User, self.message_content())] }; + if let Some(tool_results) = &self.tool_call { + messages.push(Message::new( + MessageRole::Assistant, + MessageContent::ToolResults(tool_results.clone()), + )) + } Ok(messages) } @@ -234,7 +274,7 @@ impl Input { format!(".file {}{}", files.join(" "), text) } - pub fn to_message_content(&self) -> MessageContent { + pub fn message_content(&self) -> MessageContent { if self.medias.is_empty() { MessageContent::Text(self.text.clone()) } else { @@ -257,14 +297,6 @@ impl Input { MessageContent::Array(list) } } - - pub fn required_capabilities(&self) -> ModelCapabilities { - if !self.medias.is_empty() { - ModelCapabilities::Vision - } else { - ModelCapabilities::Text - } - } } #[derive(Debug, Clone, Default)] 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)) { diff --git a/src/config/role.rs b/src/config/role.rs index 2a4b30c..27858a3 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -1,4 +1,5 @@ use super::Input; + use crate::{ client::{Message, MessageContent, MessageRole, Model}, utils::{detect_os, detect_shell}, @@ -18,10 +19,17 @@ pub const INPUT_PLACEHOLDER: &str = "__INPUT__"; pub struct Role { pub name: String, pub prompt: String, - #[serde(rename(serialize = "model", deserialize = "model"))] + #[serde( + rename(serialize = "model", deserialize = "model"), + skip_serializing_if = "Option::is_none" + )] pub model_id: Option<String>, + #[serde(skip_serializing_if = "Option::is_none")] pub temperature: Option<f64>, + #[serde(skip_serializing_if = "Option::is_none")] pub top_p: Option<f64>, + #[serde(skip_serializing_if = "Option::is_none")] + pub function_filter: Option<String>, } impl Role { @@ -32,12 +40,13 @@ impl Role { temperature: None, model_id: None, top_p: None, + function_filter: None, } } pub fn builtin() -> Vec<Role> { [ - (SHELL_ROLE, shell_prompt()), + (SHELL_ROLE, shell_prompt(), None), ( EXPLAIN_SHELL_ROLE, r#"Provide a terse, single sentence description of the given shell command. @@ -45,6 +54,7 @@ Describe each argument and option of the command. Provide short responses in about 80 words. APPLY MARKDOWN formatting when possible."# .into(), + None, ), ( CODE_ROLE, @@ -59,15 +69,18 @@ async function timeout(ms) { ``` "# .into(), + None, ), + ("%functions%", String::new(), Some(".*".into())), ] .into_iter() - .map(|(name, prompt)| Self { + .map(|(name, prompt, function_filter)| Self { name: name.into(), prompt, model_id: None, temperature: None, top_p: None, + function_filter, }) .collect() } @@ -78,7 +91,11 @@ async function timeout(ms) { Ok(output.trim_end().to_string()) } - pub fn embedded(&self) -> bool { + pub fn empty_prompt(&self) -> bool { + self.prompt.is_empty() + } + + pub fn embedded_prompt(&self) -> bool { self.prompt.contains(INPUT_PLACEHOLDER) } @@ -111,7 +128,9 @@ async function timeout(ms) { pub fn echo_messages(&self, input: &Input) -> String { let input_markdown = input.render(); - if self.embedded() { + if self.empty_prompt() { + input_markdown + } else if self.embedded_prompt() { self.prompt.replace(INPUT_PLACEHOLDER, &input_markdown) } else { format!("{}\n\n{}", self.prompt, input.render()) @@ -119,41 +138,30 @@ async function timeout(ms) { } pub fn build_messages(&self, input: &Input) -> Vec<Message> { - let mut content = input.to_message_content(); - - if self.embedded() { + let mut content = input.message_content(); + if self.empty_prompt() { + vec![Message::new(MessageRole::User, content)] + } else if self.embedded_prompt() { content.merge_prompt(|v: &str| self.prompt.replace(INPUT_PLACEHOLDER, v)); - vec![Message { - role: MessageRole::User, - content, - }] + vec![Message::new(MessageRole::User, content)] } else { let mut messages = vec![]; let (system, cases) = parse_structure_prompt(&self.prompt); if !system.is_empty() { - messages.push(Message { - role: MessageRole::System, - content: MessageContent::Text(system.to_string()), - }) + messages.push(Message::new( + MessageRole::System, + MessageContent::Text(system.to_string()), + )); } if !cases.is_empty() { messages.extend(cases.into_iter().flat_map(|(i, o)| { vec![ - Message { - role: MessageRole::User, - content: MessageContent::Text(i.to_string()), - }, - Message { - role: MessageRole::Assistant, - content: MessageContent::Text(o.to_string()), - }, + Message::new(MessageRole::User, MessageContent::Text(i.to_string())), + Message::new(MessageRole::Assistant, MessageContent::Text(o.to_string())), ] })); } - messages.push(Message { - role: MessageRole::User, - content, - }); + messages.push(Message::new(MessageRole::User, content)); messages } } diff --git a/src/config/session.rs b/src/config/session.rs index a458d4e..ef1eab1 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -1,5 +1,5 @@ use super::input::resolve_data_url; -use super::{Config, Input, Model}; +use super::{Config, Input, Model, Role}; use crate::client::{Message, MessageContent, MessageRole}; use crate::render::MarkdownRender; @@ -17,15 +17,20 @@ pub const TEMP_SESSION_NAME: &str = "temp"; pub struct Session { #[serde(rename(serialize = "model", deserialize = "model"))] model_id: String, + #[serde(skip_serializing_if = "Option::is_none")] temperature: Option<f64>, + #[serde(skip_serializing_if = "Option::is_none")] top_p: Option<f64>, - #[serde(default)] + #[serde(skip_serializing_if = "Option::is_none")] + function_filter: Option<String>, + #[serde(skip_serializing_if = "Option::is_none")] save_session: Option<bool>, messages: Vec<Message>, - #[serde(default)] + #[serde(default, skip_serializing_if = "HashMap::is_empty")] data_urls: HashMap<String, String>, - #[serde(default)] + #[serde(default, skip_serializing_if = "Vec::is_empty")] compressed_messages: Vec<Message>, + #[serde(skip_serializing_if = "Option::is_none")] compress_threshold: Option<usize>, #[serde(skip)] pub name: String, @@ -41,13 +46,14 @@ pub struct Session { impl Session { pub fn new(config: &Config, name: &str) -> Self { - Self { + let mut session = Self { model_id: config.model.id(), temperature: config.temperature, top_p: config.top_p, + function_filter: None, save_session: config.save_session, - messages: vec![], - compressed_messages: vec![], + messages: Default::default(), + compressed_messages: Default::default(), compress_threshold: None, data_urls: Default::default(), name: name.to_string(), @@ -55,7 +61,11 @@ impl Session { dirty: false, compressing: false, model: config.model.clone(), + }; + if let Some(role) = &config.role { + session.set_role_properties(role); } + session } pub fn load(name: &str, path: &Path) -> Result<Self> { @@ -86,6 +96,10 @@ impl Session { self.top_p } + pub fn function_filter(&self) -> Option<&str> { + self.function_filter.as_deref() + } + pub fn save_session(&self) -> Option<bool> { self.save_session } @@ -120,12 +134,15 @@ impl Session { if let Some(top_p) = self.top_p() { data["top_p"] = top_p.into(); } + if let Some(function_filter) = self.function_filter() { + data["function_filter"] = function_filter.into(); + } if let Some(save_session) = self.save_session() { data["save_session"] = save_session.into(); } data["total_tokens"] = tokens.into(); - if let Some(context_window) = self.model.max_input_tokens { - data["max_input_tokens"] = context_window.into(); + if let Some(max_input_tokens) = self.model.max_input_tokens() { + data["max_input_tokens"] = max_input_tokens.into(); } if percent != 0.0 { data["total/max"] = format!("{}%", percent).into(); @@ -153,6 +170,10 @@ impl Session { items.push(("top_p", top_p.to_string())); } + if let Some(function_filter) = self.function_filter() { + items.push(("function_filter", function_filter.into())); + } + if let Some(save_session) = self.save_session() { items.push(("save_session", save_session.to_string())); } @@ -161,7 +182,7 @@ impl Session { items.push(("compress_threshold", compress_threshold.to_string())); } - if let Some(max_input_tokens) = self.model.max_input_tokens { + if let Some(max_input_tokens) = self.model.max_input_tokens() { items.push(("max_input_tokens", max_input_tokens.to_string())); } @@ -202,7 +223,7 @@ impl Session { pub fn tokens_and_percent(&self) -> (usize, f32) { let tokens = self.tokens(); - let max_input_tokens = self.model.max_input_tokens.unwrap_or_default(); + let max_input_tokens = self.model.max_input_tokens().unwrap_or_default(); let percent = if max_input_tokens == 0 { 0.0 } else { @@ -226,6 +247,16 @@ impl Session { } } + pub fn set_functions(&mut self, function_filter: Option<&str>) { + self.function_filter = function_filter.map(|v| v.to_string()); + } + + pub fn set_role_properties(&mut self, role: &Role) { + self.set_temperature(role.temperature); + self.set_top_p(role.top_p); + self.set_functions(role.function_filter.as_deref()); + } + pub fn set_save_session(&mut self, value: Option<bool>) { if self.save_session != value { self.save_session = value; @@ -251,10 +282,10 @@ impl Session { pub fn compress(&mut self, prompt: String) { self.compressed_messages.append(&mut self.messages); - self.messages.push(Message { - role: MessageRole::System, - content: MessageContent::Text(prompt), - }); + self.messages.push(Message::new( + MessageRole::System, + MessageContent::Text(prompt), + )); self.dirty = true; } @@ -300,16 +331,14 @@ impl Session { } } if need_add_msg { - self.messages.push(Message { - role: MessageRole::User, - content: input.to_message_content(), - }); + self.messages + .push(Message::new(MessageRole::User, input.message_content())); } self.data_urls.extend(input.data_urls()); - self.messages.push(Message { - role: MessageRole::Assistant, - content: MessageContent::Text(output.to_string()), - }); + self.messages.push(Message::new( + MessageRole::Assistant, + MessageContent::Text(output.to_string()), + )); self.dirty = true; Ok(()) } @@ -340,10 +369,7 @@ impl Session { .extend(self.compressed_messages[self.compressed_messages.len() - 2..].to_vec()); } if need_add_msg { - messages.push(Message { - role: MessageRole::User, - content: input.to_message_content(), - }); + messages.push(Message::new(MessageRole::User, input.message_content())); } messages } |
