diff options
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/mod.rs | 196 | ||||
| -rw-r--r-- | src/config/role.rs | 197 | ||||
| -rw-r--r-- | src/config/session.rs | 30 |
3 files changed, 272 insertions, 151 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs index 508b768..72f6d5a 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -5,7 +5,7 @@ mod session; pub use self::agent::{list_agents, Agent}; pub use self::input::Input; -pub use self::role::{Role, RoleLike, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE}; +pub use self::role::{Role, RoleLike, BUILTIN_ROLES, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE}; use self::session::Session; use crate::client::{ @@ -19,7 +19,7 @@ use crate::utils::*; use anyhow::{anyhow, bail, Context, Result}; use indexmap::IndexMap; -use inquire::{Confirm, Select}; +use inquire::{validator::Validation, Confirm, Select, Text}; use parking_lot::RwLock; use serde::Deserialize; use serde_json::json; @@ -40,7 +40,7 @@ const DARK_THEME: &[u8] = include_bytes!("../../assets/monokai-extended.theme.bi const LIGHT_THEME: &[u8] = include_bytes!("../../assets/monokai-extended-light.theme.bin"); const CONFIG_FILE_NAME: &str = "config.yaml"; -const ROLES_FILE_NAME: &str = "roles.yaml"; +const ROLES_DIR_NAME: &str = "roles"; const ENV_FILE_NAME: &str = ".env"; const MESSAGES_FILE_NAME: &str = "messages.md"; const SESSIONS_DIR_NAME: &str = "sessions"; @@ -129,8 +129,6 @@ pub struct Config { pub clients: Vec<ClientConfig>, #[serde(skip)] - pub roles: Vec<Role>, - #[serde(skip)] pub role: Option<Role>, #[serde(skip)] pub session: Option<Session>, @@ -195,7 +193,6 @@ impl Default for Config { clients: vec![], - roles: vec![], role: None, session: None, rag: None, @@ -236,7 +233,6 @@ impl Config { } config.load_functions()?; - config.load_roles()?; config.setup_model()?; config.setup_document_loaders(); @@ -269,13 +265,17 @@ impl Config { } } - pub fn roles_file() -> Result<PathBuf> { - match env::var(get_env_name("roles_file")) { + pub fn roles_dir() -> Result<PathBuf> { + match env::var(get_env_name("roles_dir")) { Ok(value) => Ok(PathBuf::from(value)), - Err(_) => Self::local_path(ROLES_FILE_NAME), + Err(_) => Self::local_path(ROLES_DIR_NAME), } } + pub fn role_file(name: &str) -> Result<PathBuf> { + Ok(Self::roles_dir()?.join(format!("{name}.md"))) + } + pub fn env_file() -> Result<PathBuf> { match env::var(get_env_name("env_file")) { Ok(value) => Ok(PathBuf::from(value)), @@ -487,7 +487,7 @@ impl Config { } else if let Some(session) = &self.session { session.export() } else if let Some(role) = &self.role { - role.export() + Ok(role.export()) } else if let Some(rag) = &self.rag { rag.export() } else { @@ -531,7 +531,7 @@ impl Config { ("highlight", self.highlight.to_string()), ("light_theme", self.light_theme.to_string()), ("config_file", display_path(&Self::config_file()?)), - ("roles_file", display_path(&Self::roles_file()?)), + ("roles_dir", display_path(&Self::roles_dir()?)), ("env_file", display_path(&Self::env_file()?)), ("functions_dir", display_path(&Self::functions_dir()?)), ("rags_dir", display_path(&Self::rags_dir()?)), @@ -543,9 +543,9 @@ impl Config { } let output = items .iter() - .map(|(name, value)| format!("{name:<24}{value}")) + .map(|(name, value)| format!("{name:<24}{value}\n")) .collect::<Vec<String>>() - .join("\n"); + .join(""); Ok(output) } @@ -716,7 +716,7 @@ impl Config { pub fn role_info(&self) -> Result<String> { if let Some(role) = &self.role { - role.export() + Ok(role.export()) } else { bail!("No role") } @@ -733,17 +733,17 @@ impl Config { } pub fn retrieve_role(&self, name: &str) -> Result<Role> { - let mut role = self - .roles - .iter() - .find(|v| v.match_name(name)) - .map(|v| { - let mut role = v.clone(); - role.complete_prompt_args(name); - role - }) - .ok_or_else(|| anyhow!("Unknown role `{name}`"))?; - + let mut role = if Self::list_roles(false).contains(&name.to_string()) { + let path = Self::role_file(name)?; + let content = read_to_string(path)?; + Role::new(name, &content) + } else { + BUILTIN_ROLES + .iter() + .find(|v| v.name() == name) + .cloned() + .ok_or_else(|| anyhow!("Unknown role `{name}`"))? + }; match role.model_id() { Some(model_id) => { if self.model.id() != model_id { @@ -758,6 +758,110 @@ impl Config { Ok(role) } + pub fn new_role(&mut self, name: &str) -> Result<()> { + let ans = Confirm::new("Create a new role?") + .with_default(true) + .prompt()?; + if ans { + self.upsert_role(name)?; + } + Ok(()) + } + + pub fn edit_role(&mut self) -> Result<()> { + if let Some(name) = self.role.as_ref().map(|v| v.name().to_string()) { + self.upsert_role(&name) + } else { + bail!("No role") + } + } + + pub fn upsert_role(&mut self, name: &str) -> Result<()> { + let role_path = Self::role_file(name)?; + ensure_parent_exists(&role_path)?; + let editor = self.editor()?; + edit_file(&editor, &role_path)?; + self.use_role(name)?; + Ok(()) + } + + pub fn save_role(&mut self, name: Option<&str>) -> Result<()> { + let mut role_name = match &self.role { + Some(role) => match name { + Some(v) => v.to_string(), + None => role.name().to_string(), + }, + None => bail!("No role"), + }; + if role_name == TEMP_ROLE_NAME { + role_name = Text::new("Role name:") + .with_validator(|input: &str| { + if input.trim().is_empty() { + Ok(Validation::Invalid("This field is required".into())) + } else { + Ok(Validation::Valid) + } + }) + .prompt()?; + } + let role_path = Self::role_file(&role_name)?; + if let Some(role) = self.role.as_mut() { + let old_name = role.name().to_string(); + role.save(&role_name, &role_path, self.working_mode.is_repl())?; + if old_name != role_name { + if let Ok(path) = Self::role_file(&old_name) { + let _ = remove_file(&path); + } + } + } + + Ok(()) + } + + pub fn all_roles() -> Vec<Role> { + let mut roles: HashMap<String, Role> = BUILTIN_ROLES + .iter() + .map(|v| (v.name().to_string(), v.clone())) + .collect(); + let names = Self::list_roles(false); + for name in names { + if let Ok(path) = Self::role_file(&name) { + if let Ok(content) = read_to_string(&path) { + let role = Role::new(&name, &content); + roles.insert(name, role); + } + } + } + let mut roles: Vec<_> = roles.into_values().collect(); + roles.sort_unstable_by(|a, b| a.name().cmp(b.name())); + roles + } + + pub fn list_roles(with_builtin: bool) -> Vec<String> { + let mut names = HashSet::new(); + if let Some(rd) = Self::roles_dir().ok().and_then(|dir| read_dir(dir).ok()) { + for entry in rd.flatten() { + if let Some(name) = entry + .file_name() + .to_str() + .and_then(|v| v.strip_suffix(".md")) + { + names.insert(name.to_string()); + } + } + } + if with_builtin { + names.extend(BUILTIN_ROLES.iter().map(|v| v.name().to_string())); + } + let mut names: Vec<_> = names.into_iter().collect(); + names.sort_unstable(); + names + } + + pub fn has_role(name: &str) -> bool { + Self::list_roles(true).iter().any(|v| v == name) + } + pub fn use_session(&mut self, session_name: Option<&str>) -> Result<()> { if self.session.is_some() { bail!( @@ -824,16 +928,22 @@ impl Config { } pub fn save_session(&mut self, name: Option<&str>) -> Result<()> { - let name = match &self.session { + let session_name = match &self.session { Some(session) => match name { Some(v) => v.to_string(), None => session.name().to_string(), }, None => bail!("No session"), }; - let session_path = self.session_file(&name)?; + let session_path = self.session_file(&session_name)?; if let Some(session) = self.session.as_mut() { - session.save(&session_path, self.working_mode.is_repl())?; + let old_name = session.name().to_string(); + session.save(&session_name, &session_path, self.working_mode.is_repl())?; + if old_name != session_name { + if let Ok(path) = self.session_file(&old_name) { + let _ = remove_file(&path); + } + } } Ok(()) } @@ -843,9 +953,9 @@ impl Config { Some(session) => session.name().to_string(), None => bail!("No session"), }; - let editor = self.editor()?; let session_path = self.session_file(&name)?; self.save_session(Some(&name))?; + let editor = self.editor()?; edit_file(&editor, &session_path).with_context(|| { format!( "Failed to edit '{}' with '{editor}'", @@ -1201,10 +1311,9 @@ impl Config { let mut filter = ""; if args.len() == 1 { values = match cmd { - ".role" => self - .roles - .iter() - .map(|v| (v.name().to_string(), None)) + ".role" => Self::list_roles(true) + .into_iter() + .map(|v| (v, None)) .collect(), ".model" => list_chat_models(self) .into_iter() @@ -1702,25 +1811,6 @@ impl Config { Ok(()) } - fn load_roles(&mut self) -> Result<()> { - let path = Self::roles_file()?; - self.roles = if !path.exists() { - vec![] - } else { - let content = read_to_string(&path) - .with_context(|| format!("Failed to load roles at {}", path.display()))?; - serde_yaml::from_str(&content).with_context(|| "Invalid roles config")? - }; - let exist_roles: HashSet<_> = self.roles.iter().map(|v| v.name().to_string()).collect(); - let builtin_roles = Role::builtin(); - for role in builtin_roles { - if !exist_roles.contains(role.name()) { - self.roles.push(role); - } - } - Ok(()) - } - fn setup_model(&mut self) -> Result<()> { let mut model_id = self.model_id.clone(); if model_id.is_empty() { diff --git a/src/config/role.rs b/src/config/role.rs index 31e894b..c61aa53 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -2,15 +2,57 @@ use super::*; use crate::client::{Message, MessageContent, MessageRole, Model}; -use anyhow::{Context, Result}; +use anyhow::Result; +use fancy_regex::Regex; use serde::{Deserialize, Serialize}; +use serde_json::Value; pub const SHELL_ROLE: &str = "%shell%"; pub const EXPLAIN_SHELL_ROLE: &str = "%explain-shell%"; pub const CODE_ROLE: &str = "%code%"; +pub const FUNCTIONS_ROLE: &str = "%functions%"; pub const INPUT_PLACEHOLDER: &str = "__INPUT__"; +lazy_static::lazy_static! { + pub static ref BUILTIN_ROLES: Vec<Role> = { + [ + (SHELL_ROLE, shell_prompt()), + ( + EXPLAIN_SHELL_ROLE, + r#"Provide a terse, single sentence description of the given shell command. +Describe each argument and option of the command. +Provide short responses in about 80 words. +APPLY MARKDOWN formatting when possible."# + .into(), + ), + ( + CODE_ROLE, + r#"Provide only code without comments or explanations. +### INPUT: +async sleep in js +### OUTPUT: +```javascript +async function timeout(ms) { + return new Promise(resolve => setTimeout(resolve, ms)); +} +``` +"# + .into(), + ), + (FUNCTIONS_ROLE, r#"--- +use_tools: all +--- + "#.into()), + ] + .into_iter() + .map(|(name, content)| Role::new(name, &content)) + .collect() + }; + + static ref RE_METADATA: Regex = Regex::new(r"(?s)-{3,}\s*(.*?)\s*-{3,}\s*(.*)").unwrap(); +} + pub trait RoleLike { fn to_role(&self) -> Role; fn model(&self) -> &Model; @@ -46,57 +88,82 @@ pub struct Role { } impl Role { - pub fn new(name: &str, prompt: &str) -> Self { - Self { - name: name.into(), - prompt: prompt.into(), + pub fn new(name: &str, content: &str) -> Self { + let mut metadata = ""; + let mut prompt = content.trim(); + if let Ok(Some(caps)) = RE_METADATA.captures(content) { + if let (Some(metadata_value), Some(prompt_value)) = (caps.get(1), caps.get(2)) { + metadata = metadata_value.as_str().trim(); + prompt = prompt_value.as_str().trim(); + } + } + let mut role = Self { + name: name.to_string(), + prompt: prompt.to_string(), ..Default::default() + }; + if !metadata.is_empty() { + if let Ok(value) = serde_yaml::from_str::<Value>(metadata) { + if let Some(value) = value.as_object() { + for (key, value) in value { + match key.as_str() { + "model" => role.model_id = value.as_str().map(|v| v.to_string()), + "temperature" => role.temperature = value.as_f64(), + "top_p" => role.top_p = value.as_f64(), + "use_tools" => role.use_tools = value.as_str().map(|v| v.to_string()), + _ => (), + } + } + } + } } + role } - pub fn builtin() -> Vec<Role> { - [ - (SHELL_ROLE, shell_prompt(), None), - ( - EXPLAIN_SHELL_ROLE, - r#"Provide a terse, single sentence description of the given shell command. -Describe each argument and option of the command. -Provide short responses in about 80 words. -APPLY MARKDOWN formatting when possible."# - .into(), - None, - ), - ( - CODE_ROLE, - r#"Provide only code without comments or explanations. -### INPUT: -async sleep in js -### OUTPUT: -```javascript -async function timeout(ms) { - return new Promise(resolve => setTimeout(resolve, ms)); -} -``` -"# - .into(), - None, - ), - ("%functions%", String::new(), Some("all".into())), - ] - .into_iter() - .map(|(name, prompt, use_tools)| Self { - name: name.into(), - prompt, - use_tools, - ..Default::default() - }) - .collect() + pub fn export(&self) -> String { + let mut metadata = vec![]; + if let Some(model) = self.model_id() { + metadata.push(format!("model: {}", model)); + } + if let Some(temperature) = self.temperature() { + metadata.push(format!("temperature: {}", temperature)); + } + if let Some(top_p) = self.top_p() { + metadata.push(format!("top_p: {}", top_p)); + } + if let Some(use_tools) = self.use_tools() { + metadata.push(format!("use_tools: {}", use_tools)); + } + if metadata.is_empty() { + format!("{}\n", self.prompt) + } else if self.prompt.is_empty() { + format!("---\n{}\n---\n", metadata.join("\n")) + } else { + format!("---\n{}\n---\n\n{}\n", metadata.join("\n"), self.prompt) + } } - pub fn export(&self) -> Result<String> { - let output = serde_yaml::to_string(&self) - .with_context(|| format!("Unable to show info about role {}", &self.name))?; - Ok(output.trim_end().to_string()) + pub fn save(&mut self, role_name: &str, role_path: &Path, is_repl: bool) -> Result<()> { + ensure_parent_exists(role_path)?; + + let content = self.export(); + std::fs::write(role_path, content).with_context(|| { + format!( + "Failed to write role {} to {}", + self.name, + role_path.display() + ) + })?; + + if is_repl { + println!("✨ Saved role to '{}'", role_path.display()); + } + + if role_name != self.name { + self.name = role_name.to_string(); + } + + Ok(()) } pub fn sync<T: RoleLike>(&mut self, role_like: &T) { @@ -150,21 +217,6 @@ async function timeout(ms) { self.prompt.contains(INPUT_PLACEHOLDER) } - pub fn complete_prompt_args(&mut self, name: &str) { - self.name = name.to_string(); - self.prompt = complete_prompt_args(&self.prompt, &self.name); - } - - pub fn match_name(&self, name: &str) -> bool { - if self.name.contains(':') { - let role_name_parts: Vec<&str> = self.name.split(':').collect(); - let name_parts: Vec<&str> = name.split(':').collect(); - role_name_parts[0] == name_parts[0] && role_name_parts.len() == name_parts.len() - } else { - self.name == name - } - } - pub fn echo_messages(&self, input: &Input) -> String { let input_markdown = input.render(); if self.is_empty_prompt() { @@ -239,6 +291,9 @@ impl RoleLike for Role { } fn set_model(&mut self, model: &Model) { + if !self.model().id().is_empty() { + self.model_id = Some(model.id().to_string()); + } self.model = model.clone(); } @@ -255,14 +310,6 @@ impl RoleLike for Role { } } -fn complete_prompt_args(prompt: &str, name: &str) -> String { - let mut prompt = prompt.trim().to_string(); - for (i, arg) in name.split(':').skip(1).enumerate() { - prompt = prompt.replace(&format!("__ARG{}__", i + 1), arg); - } - prompt -} - fn parse_structure_prompt(prompt: &str) -> (&str, Vec<(&str, &str)>) { let mut text = prompt; let mut search_input = true; @@ -333,18 +380,6 @@ mod tests { use super::*; #[test] - fn test_merge_prompt_name() { - assert_eq!( - complete_prompt_args("convert __ARG1__", "convert:foo"), - "convert foo" - ); - assert_eq!( - complete_prompt_args("convert __ARG1__ to __ARG2__", "convert:foo:bar"), - "convert foo to bar" - ); - } - - #[test] fn test_parse_structure_prompt1() { let prompt = r#" System message diff --git a/src/config/session.rs b/src/config/session.rs index eccffa9..c36c9db 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -9,7 +9,7 @@ use inquire::{validator::Validation, Confirm, Text}; use serde::{Deserialize, Serialize}; use serde_json::json; use std::collections::HashMap; -use std::fs::{self, read_to_string}; +use std::fs::{read_to_string, write}; use std::path::Path; #[derive(Debug, Clone, Default, Deserialize, Serialize)] @@ -86,10 +86,6 @@ impl Session { Ok(session) } - pub fn is_temp(&self) -> bool { - self.name == TEMP_SESSION_NAME - } - pub fn is_empty(&self) -> bool { self.messages.is_empty() && self.compressed_messages.is_empty() } @@ -218,12 +214,7 @@ impl Session { } } - if lines.last() == Some(&String::new()) { - lines.pop(); - } - - let output = lines.join("\n"); - Ok(output) + Ok(lines.join("\n")) } pub fn tokens_usage(&self) -> (usize, f32) { @@ -302,6 +293,7 @@ impl Session { pub fn exit(&mut self, session_dir: &Path, is_repl: bool) -> Result<()> { let save_session = self.save_session(); if self.dirty && save_session != Some(false) { + let mut session_name = self.name().to_string(); if save_session.is_none() { if !is_repl { return Ok(()); @@ -310,8 +302,8 @@ impl Session { if !ans { return Ok(()); } - if self.is_temp() { - self.name = Text::new("Session name:") + while session_name == TEMP_SESSION_NAME { + session_name = Text::new("Session name:") .with_validator(|input: &str| { if input.trim().is_empty() { Ok(Validation::Invalid("This field is required".into())) @@ -322,20 +314,20 @@ impl Session { .prompt()?; } } - let session_path = session_dir.join(format!("{}.yaml", self.name())); - self.save(&session_path, is_repl)?; + let session_path = session_dir.join(format!("{session_name}.yaml")); + self.save(&session_name, &session_path, is_repl)?; } Ok(()) } - pub fn save(&mut self, session_path: &Path, is_repl: bool) -> Result<()> { + pub fn save(&mut self, session_name: &str, session_path: &Path, is_repl: bool) -> Result<()> { ensure_parent_exists(session_path)?; self.path = Some(session_path.display().to_string()); let content = serde_yaml::to_string(&self) .with_context(|| format!("Failed to serde session {}", self.name))?; - fs::write(session_path, content).with_context(|| { + write(session_path, content).with_context(|| { format!( "Failed to write session {} to {}", self.name, @@ -347,6 +339,10 @@ impl Session { println!("✨ Saved session to '{}'", session_path.display()); } + if self.name() != session_name { + self.name = session_name.to_string() + } + self.dirty = false; Ok(()) |
