diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-11 11:00:12 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-11 11:00:12 +0800 |
| commit | bb867c4fcbc6f42770471f0a0117cd8909f660d7 (patch) | |
| tree | 6293f7f1108309160d1951f53f6429e9b004870d /src/config/mod.rs | |
| parent | 5635ca6a58fb4a590419335b098b7317285bfb82 (diff) | |
| download | aichat-bb867c4fcbc6f42770471f0a0117cd8909f660d7.tar.gz | |
feat: support bot (#579)
* feat: support bots
* refactor with RoleLike
* improve exiting session
* make bot works with rag
* refactor repl assert state
* add bot banner
* repl complete bots according bots.txt
* fix on windows
* remove threadpool executing function callings
* adjust repl left_prompt
* move bot config to global config.yaml
* `.bot` throw err if funciton callings is not configured
Diffstat (limited to 'src/config/mod.rs')
| -rw-r--r-- | src/config/mod.rs | 441 |
1 files changed, 304 insertions, 137 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs index fbee667..b874b16 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -1,21 +1,23 @@ +mod bot; mod input; mod role; mod session; -pub use self::input::{Input, InputContext}; -pub use self::role::{Role, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE}; -use self::session::{Session, TEMP_SESSION_NAME}; +pub use self::bot::{list_bots, Bot, BotConfig}; +pub use self::input::Input; +pub use self::role::{Role, RoleLike, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE}; +use self::session::Session; use crate::client::{ create_client_config, list_chat_models, list_client_types, ClientConfig, Model, OPENAI_COMPATIBLE_PLATFORMS, }; -use crate::function::{Function, ToolCallResult}; -use crate::rag::{Rag, TEMP_RAG_NAME}; +use crate::function::{FunctionDeclaration, Functions, ToolCallResult}; +use crate::rag::Rag; use crate::render::{MarkdownRender, RenderOptions}; use crate::utils::{ format_option_value, fuzzy_match, get_env_name, light_theme_from_colorfgbg, now, render_prompt, - set_text, AbortSignal, IS_STDOUT_TERMINAL, + set_text, warning_text, AbortSignal, IS_STDOUT_TERMINAL, }; use anyhow::{anyhow, bail, Context, Result}; @@ -44,6 +46,15 @@ const MESSAGES_FILE_NAME: &str = "messages.md"; const SESSIONS_DIR_NAME: &str = "sessions"; const RAGS_DIR_NAME: &str = "rags"; const FUNCTIONS_DIR_NAME: &str = "functions"; +const FUNCTIONS_FILE_NAME: &str = "functions.json"; +const BOTS_DIR_NAME: &str = "bots"; +const BOT_DEFINITION_FILE_NAME: &str = "index.yaml"; +const BOT_EMBEDDINGS_DIR: &str = "embeddings"; +const BOT_RAG_FILE_NAME: &str = "rag.bin"; + +pub const TEMP_ROLE_NAME: &str = "%%"; +pub const TEMP_RAG_NAME: &str = "temp"; +pub const TEMP_SESSION_NAME: &str = "temp"; const CLIENTS_FIELD: &str = "clients"; @@ -51,7 +62,7 @@ const SUMMARIZE_PROMPT: &str = "Summarize the discussion briefly in 200 words or less to use as a prompt for future context."; const SUMMARY_PROMPT: &str = "This is a summary of the chat history as a recap: "; -const RAG_TEMPLATE: &str = r#"Answer the following question based only on the provided context: +const RAG_TEMPLATE: &str = r#"Answer the question based only on the provided context: <context> __CONTEXT__ </context> @@ -59,7 +70,7 @@ __CONTEXT__ Question: __INPUT__ "#; -const LEFT_PROMPT: &str = "{color.green}{?session {session}{?role /}}{role}{?rag #{rag}}{color.cyan}{?session )}{!session >}{color.reset} "; +const LEFT_PROMPT: &str = "{color.green}{?session {?bot {bot}#}{session}{?role /}}{!session {?bot {bot}}}{role}{?rag @{rag}}{color.cyan}{?session )}{!session >}{color.reset} "; const RIGHT_PROMPT: &str = "{color.purple}{?session {?consume_tokens {consume_tokens}({consume_percent}%)}{!consume_tokens {consume_tokens}}}{color.reset}"; #[derive(Debug, Clone, Deserialize)] @@ -91,6 +102,7 @@ pub struct Config { pub left_prompt: Option<String>, pub right_prompt: Option<String>, pub clients: Vec<ClientConfig>, + pub bots: Vec<BotConfig>, #[serde(skip)] pub roles: Vec<Role>, #[serde(skip)] @@ -100,9 +112,11 @@ pub struct Config { #[serde(skip)] pub rag: Option<Arc<Rag>>, #[serde(skip)] + pub bot: Option<Bot>, + #[serde(skip)] pub model: Model, #[serde(skip)] - pub function: Function, + pub functions: Functions, #[serde(skip)] pub working_mode: WorkingMode, #[serde(skip)] @@ -136,12 +150,14 @@ impl Default for Config { left_prompt: None, right_prompt: None, clients: vec![], + bots: vec![], roles: vec![], role: None, session: None, rag: None, + bot: None, model: Default::default(), - function: Default::default(), + functions: Default::default(), working_mode: WorkingMode::Command, last_message: None, } @@ -168,7 +184,7 @@ impl Config { config.set_wrap(&wrap)?; } - config.function = Function::init(&Self::functions_dir()?)?; + config.functions = Functions::init(&Self::functions_file()?)?; config.working_mode = working_mode; config.load_roles()?; @@ -211,7 +227,8 @@ impl Config { } pub fn retrieve_role(&self, name: &str) -> Result<Role> { - self.roles + let mut role = self + .roles .iter() .find(|v| v.match_name(name)) .map(|v| { @@ -219,7 +236,18 @@ impl Config { role.complete_prompt_args(name); role }) - .ok_or_else(|| anyhow!("Unknown role `{name}`")) + .ok_or_else(|| anyhow!("Unknown role `{name}`"))?; + + match role.model_id() { + Some(model_id) => { + if self.model.id() != model_id { + let model = Model::retrieve(self, model_id)?; + role.set_model(&model); + } + } + None => role.set_model(&self.model), + } + Ok(role) } pub fn config_dir() -> Result<PathBuf> { @@ -268,11 +296,20 @@ impl Config { let timestamp = now(); let summary = input.summary(); let input_markdown = input.render(); - let scope = match (input.role().map(|v| v.name.as_str()), input.rag()) { - (Some(role), Some(rag)) => format!(" ({role}#{rag})"), - (Some(role), _) => format!(" ({role})"), - (None, Some(rag)) => format!(" (#{rag})"), - _ => String::new(), + let scope = if self.bot.is_none() { + let role_name = if input.role().is_derived() { + None + } else { + Some(input.role().name()) + }; + match (role_name, input.rag_name()) { + (Some(role), Some(rag_name)) => format!(" ({role}#{rag_name})"), + (Some(role), _) => format!(" ({role})"), + (None, Some(rag_name)) => format!(" (#{rag_name})"), + _ => String::new(), + } + } else { + String::new() }; let output = format!("# CHAT: {summary} [{timestamp}]{scope}\n{input_markdown}\n--------\n{output}\n--------\n\n",); file.write_all(output.as_bytes()) @@ -299,17 +336,23 @@ impl Config { } } - pub fn messages_file() -> Result<PathBuf> { - match env::var(get_env_name("messages_file")) { - Ok(value) => Ok(PathBuf::from(value)), - Err(_) => Self::local_path(MESSAGES_FILE_NAME), + pub fn messages_file(&self) -> Result<PathBuf> { + match &self.bot { + None => match env::var(get_env_name("messages_file")) { + Ok(value) => Ok(PathBuf::from(value)), + Err(_) => Self::local_path(MESSAGES_FILE_NAME), + }, + Some(bot) => Ok(Self::bot_config_dir(bot.name())?.join(MESSAGES_FILE_NAME)), } } - pub fn sessions_dir() -> Result<PathBuf> { - match env::var(get_env_name("sessions_dir")) { - Ok(value) => Ok(PathBuf::from(value)), - Err(_) => Self::local_path(SESSIONS_DIR_NAME), + pub fn sessions_dir(&self) -> Result<PathBuf> { + match &self.bot { + None => match env::var(get_env_name("sessions_dir")) { + Ok(value) => Ok(PathBuf::from(value)), + Err(_) => Self::local_path(SESSIONS_DIR_NAME), + }, + Some(bot) => Ok(Self::bot_config_dir(bot.name())?.join(SESSIONS_DIR_NAME)), } } @@ -327,20 +370,66 @@ impl Config { } } - pub fn session_file(name: &str) -> Result<PathBuf> { - let mut path = Self::sessions_dir()?; - path.push(&format!("{name}.yaml")); - Ok(path) + pub fn functions_file() -> Result<PathBuf> { + Ok(Self::functions_dir()?.join(FUNCTIONS_FILE_NAME)) } - pub fn rag_file(name: &str) -> Result<PathBuf> { - let mut path = Self::rags_dir()?; - path.push(&format!("{name}.bin")); + pub fn functions_bin_dir() -> Result<PathBuf> { + Ok(Self::functions_dir()?.join("bin")) + } + + pub fn session_file(&self, name: &str) -> Result<PathBuf> { + Ok(self.sessions_dir()?.join(format!("{name}.yaml"))) + } + + pub fn rag_file(&self, name: &str) -> Result<PathBuf> { + let path = if self.bot.is_none() { + Self::rags_dir()?.join(format!("{name}.bin")) + } else { + Self::rags_dir()? + .join(BOTS_DIR_NAME) + .join(format!("{name}.bin")) + }; Ok(path) } + pub fn bots_dir() -> Result<PathBuf> { + match env::var(get_env_name("bots_config_dir")) { + Ok(value) => Ok(PathBuf::from(value)), + Err(_) => Self::local_path(BOTS_DIR_NAME), + } + } + + pub fn bot_config_dir(name: &str) -> Result<PathBuf> { + Ok(Self::bots_dir()?.join(name)) + } + + pub fn bot_rag_file(name: &str) -> Result<PathBuf> { + Ok(Self::bot_config_dir(name)?.join(BOT_RAG_FILE_NAME)) + } + + pub fn bots_functions_dir() -> Result<PathBuf> { + Ok(Self::functions_dir()?.join(BOTS_DIR_NAME)) + } + + pub fn bot_functions_dir(name: &str) -> Result<PathBuf> { + Ok(Self::bots_functions_dir()?.join(name)) + } + + pub fn bot_functions_file(name: &str) -> Result<PathBuf> { + Ok(Self::bot_functions_dir(name)?.join(FUNCTIONS_FILE_NAME)) + } + + pub fn bot_definition_file(name: &str) -> Result<PathBuf> { + Ok(Self::bot_functions_dir(name)?.join(BOT_DEFINITION_FILE_NAME)) + } + + pub fn bot_embeddings_dir(name: &str) -> Result<PathBuf> { + Ok(Self::bot_functions_dir(name)?.join(BOT_EMBEDDINGS_DIR)) + } + pub fn use_prompt(&mut self, prompt: &str) -> Result<()> { - let role = Role::temp(prompt); + let role = Role::new(TEMP_ROLE_NAME, prompt); self.use_role_obj(role) } @@ -350,22 +439,25 @@ impl Config { } pub fn use_role_obj(&mut self, role: Role) -> Result<()> { + if self.bot.is_some() { + bail!("Cannot perform this action because you are using a bot") + } if let Some(session) = self.session.as_mut() { session.guard_empty()?; - session.set_role_properties(&role); - } - if let Some(model_id) = &role.model_id { - self.set_model(model_id)?; + session.set_role(role); + } else { + self.role = Some(role); } - self.role = Some(role); Ok(()) } pub fn exit_role(&mut self) -> Result<()> { - if self.session.is_none() { - self.restore_model()?; + if self.role.is_some() { + if let Some(session) = self.session.as_mut() { + session.clear_role(); + } + self.role = None; } - self.role = None; Ok(()) } @@ -378,6 +470,9 @@ impl Config { flags |= StateFlags::SESSION; } } + if self.bot.is_some() { + flags |= StateFlags::BOT; + } if self.role.is_some() { flags |= StateFlags::ROLE; } @@ -387,27 +482,62 @@ impl Config { flags } - pub fn has_role_or_session(&self) -> bool { - self.role.is_some() || self.session.is_some() + pub fn current_model(&self) -> &Model { + if let Some(session) = self.session.as_ref() { + session.model() + } else if let Some(bot) = self.bot.as_ref() { + bot.model() + } else if let Some(role) = self.role.as_ref() { + role.model() + } else { + &self.model + } } - pub fn set_temperature(&mut self, value: Option<f64>) { - if let Some(session) = self.session.as_mut() { - session.set_temperature(value); - } else if let Some(role) = self.role.as_mut() { - role.set_temperature(value); + pub fn extract_role(&self) -> Role { + let mut role = if let Some(session) = self.session.as_ref() { + session.to_role() + } else if let Some(bot) = self.bot.as_ref() { + bot.to_role() + } else if let Some(role) = self.role.as_ref() { + role.clone() } else { - self.temperature = value; + let mut role = Role::default(); + role.batch_set(&self.model, self.temperature, self.top_p, None); + role + }; + if role.temperature().is_none() && self.temperature.is_some() { + role.set_temperature(self.temperature); } + if role.top_p().is_none() && self.top_p.is_some() { + role.set_top_p(self.top_p); + } + role } - pub fn set_top_p(&mut self, value: Option<f64>) { + pub fn role_like_mut(&mut self) -> Option<&mut dyn RoleLike> { if let Some(session) = self.session.as_mut() { - session.set_top_p(value); + Some(session) + } else if let Some(bot) = self.bot.as_mut() { + Some(bot) } else if let Some(role) = self.role.as_mut() { - role.set_top_p(value); + Some(role) } else { - self.top_p = value; + None + } + } + + pub fn set_temperature(&mut self, value: Option<f64>) { + match self.role_like_mut() { + Some(role_like) => role_like.set_temperature(value), + None => self.temperature = value, + } + } + + pub fn set_top_p(&mut self, value: Option<f64>) { + match self.role_like_mut() { + Some(role_like) => role_like.set_top_p(value), + None => self.top_p = value, } } @@ -441,46 +571,26 @@ impl Config { Ok(()) } - pub fn set_model(&mut self, value: &str) -> Result<()> { - let model = Model::find(&list_chat_models(self), value); - match model { - None => bail!("No model '{}'", value), - Some(model) => { - if let Some(session) = self.session.as_mut() { - session.set_model(&model); - } else if let Some(role) = self.role.as_mut() { - role.set_model(&model); - } + pub fn set_model(&mut self, model_id: &str) -> Result<()> { + let model = Model::retrieve(self, model_id)?; + match self.role_like_mut() { + Some(role_like) => role_like.set_model(&model), + None => { self.model = model; - Ok(()) } } + Ok(()) } - pub fn set_model_id(&mut self) { - self.model_id = self.model.id() - } - - pub fn restore_model(&mut self) -> Result<()> { - let origin_model_id = self.model_id.clone(); - self.set_model(&origin_model_id) - } - - pub fn system_info(&self) -> Result<String> { + pub fn sysinfo(&self) -> Result<String> { let display_path = |path: &Path| path.display().to_string(); let wrap = self .wrap .clone() .map_or_else(|| String::from("no"), |v| v.to_string()); - let (temperature, top_p) = if let Some(session) = &self.session { - (session.temperature(), session.top_p()) - } else if let Some(role) = &self.role { - (role.temperature, role.top_p) - } else { - (self.temperature, self.top_p) - }; + let role = self.extract_role(); let items = vec![ - ("model", self.model.id()), + ("model", role.model().id()), ( "max_output_tokens", self.model @@ -488,8 +598,8 @@ impl Config { .map(|v| format!("{v} (current model)")) .unwrap_or_else(|| "-".into()), ), - ("temperature", format_option_value(&temperature)), - ("top_p", format_option_value(&top_p)), + ("temperature", format_option_value(&role.temperature())), + ("top_p", format_option_value(&role.top_p())), ("rag_top_k", self.rag_top_k.to_string()), ("function_calling", self.function_calling.to_string()), ("compress_threshold", self.compress_threshold.to_string()), @@ -505,10 +615,11 @@ impl Config { ("prelude", format_option_value(&self.prelude)), ("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()?)), - ("rags_dir", display_path(&Self::rags_dir()?)), ("functions_dir", display_path(&Self::functions_dir()?)), + ("rags_dir", display_path(&Self::rags_dir()?)), + ("bots_dir", display_path(&Self::bots_dir()?)), + ("sessions_dir", display_path(&self.sessions_dir()?)), + ("messages_file", display_path(&self.messages_file()?)), ]; let output = items .iter() @@ -530,7 +641,7 @@ impl Config { if let Some(session) = &self.session { let render_options = self.get_render_options()?; let mut markdown_render = MarkdownRender::init(render_options)?; - session.info(&mut markdown_render) + session.render(&mut markdown_render) } else { bail!("No session") } @@ -544,6 +655,14 @@ impl Config { } } + pub fn bot_info(&self) -> Result<String> { + if let Some(bot) = &self.bot { + bot.export() + } else { + bail!("No rag") + } + } + pub fn info(&self) -> Result<String> { if let Some(session) = &self.session { session.export() @@ -552,7 +671,7 @@ impl Config { } else if let Some(rag) = &self.rag { rag.export() } else { - self.system_info() + self.sysinfo() } } @@ -563,28 +682,25 @@ impl Config { .unwrap_or_default() } - pub fn repl_complete(&self, cmd: &str, args: &[&str]) -> Vec<(String, String)> { + pub fn repl_complete(&self, cmd: &str, args: &[&str]) -> Vec<(String, Option<String>)> { let (values, filter) = if args.len() == 1 { let values = match cmd { ".role" => self .roles .iter() - .map(|v| (v.name.clone(), String::new())) + .map(|v| (v.name().to_string(), None)) .collect(), ".model" => list_chat_models(self) .into_iter() - .map(|v| (v.id(), v.description())) + .map(|v| (v.id(), Some(v.description()))) .collect(), ".session" => self .list_sessions() .into_iter() - .map(|v| (v.clone(), String::new())) - .collect(), - ".rag" => self - .list_rags() - .into_iter() - .map(|v| (v.clone(), String::new())) + .map(|v| (v, None)) .collect(), + ".rag" => self.list_rags().into_iter().map(|v| (v, None)).collect(), + ".bot" => list_bots().into_iter().map(|v| (v, None)).collect(), ".set" => vec![ "max_output_tokens", "temperature", @@ -599,7 +715,7 @@ impl Config { "auto_copy", ] .into_iter() - .map(|v| (format!("{v} "), String::new())) + .map(|v| (format!("{v} "), None)) .collect(), _ => vec![], }; @@ -625,10 +741,7 @@ impl Config { "auto_copy" => complete_bool(self.auto_copy), _ => vec![], }; - ( - values.into_iter().map(|v| (v, String::new())).collect(), - args[1], - ) + (values.into_iter().map(|v| (v, None)).collect(), args[1]) } else { return vec![]; }; @@ -696,6 +809,30 @@ impl Config { Ok(()) } + pub fn retrieve_functions( + &self, + model: &Model, + role: &Role, + ) -> Option<Vec<FunctionDeclaration>> { + let mut functions = None; + if self.function_calling { + let function_matcher = role.function_matcher(); + if let Some(matcher) = function_matcher { + functions = match &self.bot { + Some(bot) => bot.functions().select(&matcher), + None => self.functions.select(&matcher), + }; + if !model.supports_function_calling() { + functions = None; + if *IS_STDOUT_TERMINAL { + eprintln!("{}", warning_text("WARNING: the role or session includes functions, but the model or client does not support function calling.")); + } + } + } + }; + functions + } + pub fn use_session(&mut self, session: Option<&str>) -> Result<()> { if self.session.is_some() { bail!( @@ -704,7 +841,7 @@ impl Config { } match session { None => { - let session_file = Self::session_file(TEMP_SESSION_NAME)?; + let session_file = self.session_file(TEMP_SESSION_NAME)?; if session_file.exists() { remove_file(session_file).with_context(|| { format!("Failed to cleanup previous '{TEMP_SESSION_NAME}' session") @@ -714,14 +851,12 @@ impl Config { self.session = Some(session); } Some(name) => { - let session_path = Self::session_file(name)?; + let session_path = self.session_file(name)?; if !session_path.exists() { self.session = Some(Session::new(self, name)); } else { - let session = Session::load(name, &session_path)?; - let model_id = session.model_id().to_string(); + let session = Session::load(self, name, &session_path)?; self.session = Some(session); - self.set_model(&model_id)?; } } } @@ -745,20 +880,19 @@ impl Config { pub fn exit_session(&mut self) -> Result<()> { if let Some(mut session) = self.session.take() { let is_repl = self.working_mode == WorkingMode::Repl; - let sessions_dir = Self::sessions_dir()?; + let sessions_dir = self.sessions_dir()?; session.exit(&sessions_dir, is_repl)?; self.last_message = None; - self.restore_model()?; } Ok(()) } pub fn save_session(&mut self, name: &str) -> Result<()> { + let sessions_dir = self.sessions_dir()?; if let Some(session) = self.session.as_mut() { if !name.is_empty() { - session.name = name.to_string(); + session.set_name(name); } - let sessions_dir = Self::sessions_dir()?; session.save(&sessions_dir)?; } Ok(()) @@ -772,7 +906,7 @@ impl Config { } pub fn list_sessions(&self) -> Vec<String> { - let sessions_dir = match Self::sessions_dir() { + let sessions_dir = match self.sessions_dir() { Ok(dir) => dir, Err(_) => return vec![], }; @@ -795,7 +929,7 @@ impl Config { pub fn should_compress_session(&mut self) -> bool { if let Some(session) = self.session.as_mut() { if session.need_compress(self.compress_threshold) { - session.compressing = true; + session.set_compressing(true); return true; } } @@ -816,13 +950,13 @@ impl Config { pub fn is_compressing_session(&self) -> bool { self.session .as_ref() - .map(|v| v.compressing) + .map(|v| v.compressing()) .unwrap_or_default() } pub fn end_compressing_session(&mut self) { if let Some(session) = self.session.as_mut() { - session.compressing = false; + session.set_compressing(false); } } @@ -831,23 +965,23 @@ impl Config { rag: Option<&str>, abort_signal: AbortSignal, ) -> Result<()> { - if config.read().rag.is_some() { - bail!("Already in a rag, please run '.exit rag' first to exit the current rag."); + if config.read().bot.is_some() { + bail!("Cannot perform this action because you are using a bot") } let rag = match rag { None => { - let rag_path = Self::rag_file(TEMP_RAG_NAME)?; + let rag_path = config.read().rag_file(TEMP_RAG_NAME)?; if rag_path.exists() { remove_file(&rag_path).with_context(|| { format!("Failed to cleanup previous '{TEMP_RAG_NAME}' rag") })?; } - Rag::init(config, TEMP_RAG_NAME, &rag_path, abort_signal).await? + Rag::init(config, TEMP_RAG_NAME, &rag_path, &[], abort_signal).await? } Some(name) => { - let rag_path = Self::rag_file(name)?; + let rag_path = config.read().rag_file(name)?; if !rag_path.exists() { - Rag::init(config, name, &rag_path, abort_signal).await? + Rag::init(config, name, &rag_path, &[], abort_signal).await? } else { Rag::load(config, name, &rag_path)? } @@ -894,6 +1028,29 @@ impl Config { .replace("__INPUT__", text) } + pub async fn use_bot( + config: &GlobalConfig, + name: &str, + abort_signal: AbortSignal, + ) -> Result<()> { + if !config.read().function_calling { + bail!("Before using the bot, please configure function calling first."); + } + if config.read().bot.is_some() { + bail!("Already in a bot, please run '.exit bot' first to exit the current bot."); + } + let bot = Bot::init(config, name, abort_signal).await?; + config.write().rag = bot.rag(); + config.write().bot = Some(bot); + Ok(()) + } + + pub fn exit_bot(&mut self) -> Result<()> { + self.rag.take(); + self.bot.take(); + Ok(()) + } + pub fn get_render_options(&self) -> Result<RenderOptions> { let theme = if self.highlight { let theme_mode = if self.light_theme { "light" } else { "dark" }; @@ -940,22 +1097,23 @@ 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().to_string()); - output.insert("model_name", self.model.name().to_string()); + let role = self.extract_role(); + output.insert("model", role.model().id()); + output.insert("client_name", role.model().client_name().to_string()); + output.insert("model_name", role.model().name().to_string()); output.insert( "max_input_tokens", - self.model + role.model() .max_input_tokens() .unwrap_or_default() .to_string(), ); - if let Some(temperature) = self.temperature { + if let Some(temperature) = role.temperature() { if temperature != 0.0 { output.insert("temperature", temperature.to_string()); } } - if let Some(top_p) = self.top_p { + if let Some(top_p) = role.top_p() { if top_p != 0.0 { output.insert("top_p", top_p.to_string()); } @@ -974,13 +1132,13 @@ impl Config { if self.auto_copy { output.insert("auto_copy", "true".to_string()); } - if let Some(role) = &self.role { - output.insert("role", role.name.clone()); + if !role.is_derived() { + output.insert("role", role.name().to_string()); } if let Some(session) = &self.session { output.insert("session", session.name().to_string()); - output.insert("dirty", session.dirty.to_string()); - let (tokens, percent) = session.tokens_and_percent(); + output.insert("dirty", session.dirty().to_string()); + let (tokens, percent) = session.tokens_usage(); output.insert("consume_tokens", tokens.to_string()); output.insert("consume_percent", percent.to_string()); output.insert("user_messages_len", session.user_messages_len().to_string()); @@ -988,6 +1146,9 @@ impl Config { if let Some(rag) = &self.rag { output.insert("rag", rag.name().to_string()); } + if let Some(bot) = &self.bot { + output.insert("bot", bot.name().to_string()); + } if self.highlight { output.insert("color.reset", "\u{1b}[0m".to_string()); @@ -1015,7 +1176,7 @@ impl Config { } fn open_message_file(&self) -> Result<File> { - let path = Self::messages_file()?; + let path = self.messages_file()?; ensure_parent_exists(&path)?; OpenOptions::new() .create(true) @@ -1078,10 +1239,10 @@ impl Config { .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.clone()).collect(); + 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) { + if !exist_roles.contains(role.name()) { self.roles.push(role); } } @@ -1165,6 +1326,7 @@ bitflags::bitflags! { const SESSION_EMPTY = 1 << 1; const SESSION = 1 << 2; const RAG = 1 << 3; + const BOT = 1 << 4; } } @@ -1172,12 +1334,17 @@ bitflags::bitflags! { pub enum AssertState { True(StateFlags), False(StateFlags), + TrueFalse(StateFlags, StateFlags), + Equal(StateFlags), } impl AssertState { - pub fn any() -> Self { + pub fn pass() -> Self { AssertState::False(StateFlags::empty()) } + pub fn bare() -> Self { + AssertState::Equal(StateFlags::empty()) + } } fn create_config_file(config_path: &Path) -> Result<()> { |
