diff options
| -rw-r--r-- | src/config/bot.rs | 2 | ||||
| -rw-r--r-- | src/config/input.rs | 2 | ||||
| -rw-r--r-- | src/config/mod.rs | 582 | ||||
| -rw-r--r-- | src/config/session.rs | 16 | ||||
| -rw-r--r-- | src/main.rs | 4 | ||||
| -rw-r--r-- | src/render/mod.rs | 2 |
6 files changed, 302 insertions, 306 deletions
diff --git a/src/config/bot.rs b/src/config/bot.rs index 4d748bb..19fd5cd 100644 --- a/src/config/bot.rs +++ b/src/config/bot.rs @@ -54,7 +54,7 @@ impl Bot { } }; - let render_options = config.read().get_render_options()?; + let render_options = config.read().render_options()?; let mut markdown_render = MarkdownRender::init(render_options)?; println!("{}", markdown_render.render(&definition.banner())); diff --git a/src/config/input.rs b/src/config/input.rs index 0e2aa7a..1686245 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -185,7 +185,7 @@ impl Input { self.config.read().model.guard_max_input_tokens(&messages)?; let temperature = self.role().temperature(); let top_p = self.role().top_p(); - let functions = self.config.read().retrieve_functions(model, self.role()); + let functions = self.config.read().select_functions(model, self.role()); Ok(ChatCompletionsData { messages, temperature, diff --git a/src/config/mod.rs b/src/config/mod.rs index b874b16..7c99acb 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -196,60 +196,6 @@ impl Config { Ok(config) } - pub fn apply_prelude(&mut self) -> Result<()> { - let prelude = self.prelude.clone().unwrap_or_default(); - if prelude.is_empty() { - return Ok(()); - } - let err_msg = || format!("Invalid prelude '{}", prelude); - match prelude.split_once(':') { - Some(("role", name)) => { - if self.role.is_none() && self.session.is_none() { - self.use_role(name).with_context(err_msg)?; - } - } - Some(("session", name)) => { - if self.session.is_none() { - self.use_session(Some(name)).with_context(err_msg)?; - } - } - _ => { - bail!("{}", err_msg()) - } - } - Ok(()) - } - - pub fn buffer_editor(&self) -> Option<String> { - self.buffer_editor - .clone() - .or_else(|| env::var("VISUAL").ok().or_else(|| env::var("EDITOR").ok())) - } - - 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}`"))?; - - 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> { let env_name = get_env_name("config_dir"); let path = if let Some(v) = env::var_os(env_name) { @@ -428,39 +374,6 @@ impl Config { Ok(Self::bot_functions_dir(name)?.join(BOT_EMBEDDINGS_DIR)) } - pub fn use_prompt(&mut self, prompt: &str) -> Result<()> { - let role = Role::new(TEMP_ROLE_NAME, prompt); - self.use_role_obj(role) - } - - pub fn use_role(&mut self, name: &str) -> Result<()> { - let role = self.retrieve_role(name)?; - self.use_role_obj(role) - } - - 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(role); - } else { - self.role = Some(role); - } - Ok(()) - } - - pub fn exit_role(&mut self) -> Result<()> { - if self.role.is_some() { - if let Some(session) = self.session.as_mut() { - session.clear_role(); - } - self.role = None; - } - Ok(()) - } - pub fn state(&self) -> StateFlags { let mut flags = StateFlags::empty(); if let Some(session) = &self.session { @@ -494,6 +407,18 @@ impl Config { } } + pub fn role_like_mut(&mut self) -> Option<&mut dyn RoleLike> { + if let Some(session) = self.session.as_mut() { + Some(session) + } else if let Some(bot) = self.bot.as_mut() { + Some(bot) + } else if let Some(role) = self.role.as_mut() { + Some(role) + } else { + None + } + } + pub fn extract_role(&self) -> Role { let mut role = if let Some(session) = self.session.as_ref() { session.to_role() @@ -515,71 +440,16 @@ impl Config { role } - pub fn role_like_mut(&mut self) -> Option<&mut dyn RoleLike> { - if let Some(session) = self.session.as_mut() { - Some(session) - } else if let Some(bot) = self.bot.as_mut() { - Some(bot) - } else if let Some(role) = self.role.as_mut() { - Some(role) - } else { - 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, - } - } - - pub fn set_save_session(&mut self, value: Option<bool>) { - if let Some(session) = self.session.as_mut() { - session.set_save_session(value); - } else { - self.save_session = value; - } - } - - pub fn set_compress_threshold(&mut self, value: Option<usize>) { - if let Some(session) = self.session.as_mut() { - session.set_compress_threshold(value); - } else { - self.compress_threshold = value.unwrap_or_default(); - } - } - - pub fn set_wrap(&mut self, value: &str) -> Result<()> { - if value == "no" { - self.wrap = None; - } else if value == "auto" { - self.wrap = Some(value.into()); + pub fn info(&self) -> Result<String> { + if let Some(session) = &self.session { + session.export() + } else if let Some(role) = &self.role { + role.export() + } else if let Some(rag) = &self.rag { + rag.export() } else { - value - .parse::<u16>() - .map_err(|_| anyhow!("Invalid wrap value"))?; - self.wrap = Some(value.into()) - } - Ok(()) - } - - 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; - } + self.sysinfo() } - Ok(()) } pub fn sysinfo(&self) -> Result<String> { @@ -629,128 +499,6 @@ impl Config { Ok(output) } - pub fn role_info(&self) -> Result<String> { - if let Some(role) = &self.role { - role.export() - } else { - bail!("No role") - } - } - - pub fn session_info(&self) -> Result<String> { - if let Some(session) = &self.session { - let render_options = self.get_render_options()?; - let mut markdown_render = MarkdownRender::init(render_options)?; - session.render(&mut markdown_render) - } else { - bail!("No session") - } - } - - pub fn rag_info(&self) -> Result<String> { - if let Some(rag) = &self.rag { - rag.export() - } else { - bail!("No rag") - } - } - - 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() - } else if let Some(role) = &self.role { - role.export() - } else if let Some(rag) = &self.rag { - rag.export() - } else { - self.sysinfo() - } - } - - pub fn last_reply(&self) -> &str { - self.last_message - .as_ref() - .map(|(_, reply)| reply.as_str()) - .unwrap_or_default() - } - - 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().to_string(), None)) - .collect(), - ".model" => list_chat_models(self) - .into_iter() - .map(|v| (v.id(), Some(v.description()))) - .collect(), - ".session" => self - .list_sessions() - .into_iter() - .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", - "top_p", - "rag_top_k", - "function_calling", - "compress_threshold", - "save", - "save_session", - "highlight", - "dry_run", - "auto_copy", - ] - .into_iter() - .map(|v| (format!("{v} "), None)) - .collect(), - _ => vec![], - }; - (values, args[0]) - } else if args.len() == 2 { - let values = match args[0] { - "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 { - session.save_session() - } else { - self.save_session - }; - complete_option_bool(save_session) - } - "highlight" => complete_bool(self.highlight), - "dry_run" => complete_bool(self.dry_run), - "auto_copy" => complete_bool(self.auto_copy), - _ => vec![], - }; - (values.into_iter().map(|v| (v, None)).collect(), args[1]) - } else { - return vec![]; - }; - values - .into_iter() - .filter(|(value, _)| fuzzy_match(value, filter)) - .collect() - } - pub fn update(&mut self, data: &str) -> Result<()> { let parts: Vec<&str> = data.split_whitespace().collect(); if parts.len() != 2 { @@ -809,28 +557,124 @@ 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.")); - } + 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, + } + } + + pub fn set_save_session(&mut self, value: Option<bool>) { + if let Some(session) = self.session.as_mut() { + session.set_save_session(value); + } else { + self.save_session = value; + } + } + + pub fn set_compress_threshold(&mut self, value: Option<usize>) { + if let Some(session) = self.session.as_mut() { + session.set_compress_threshold(value); + } else { + self.compress_threshold = value.unwrap_or_default(); + } + } + + pub fn set_wrap(&mut self, value: &str) -> Result<()> { + if value == "no" { + self.wrap = None; + } else if value == "auto" { + self.wrap = Some(value.into()); + } else { + value + .parse::<u16>() + .map_err(|_| anyhow!("Invalid wrap value"))?; + self.wrap = Some(value.into()) + } + Ok(()) + } + + 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(()) + } + + pub fn use_prompt(&mut self, prompt: &str) -> Result<()> { + let role = Role::new(TEMP_ROLE_NAME, prompt); + self.use_role_obj(role) + } + + pub fn use_role(&mut self, name: &str) -> Result<()> { + let role = self.retrieve_role(name)?; + self.use_role_obj(role) + } + + 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(role); + } else { + self.role = Some(role); + } + Ok(()) + } + + pub fn role_info(&self) -> Result<String> { + if let Some(role) = &self.role { + role.export() + } else { + bail!("No role") + } + } + + pub fn exit_role(&mut self) -> Result<()> { + if self.role.is_some() { + if let Some(session) = self.session.as_mut() { + session.clear_role(); + } + self.role = None; + } + Ok(()) + } + + 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}`"))?; + + match role.model_id() { + Some(model_id) => { + if self.model.id() != model_id { + let model = Model::retrieve(self, model_id)?; + role.set_model(&model); } } - }; - functions + None => role.set_model(&self.model), + } + Ok(role) } pub fn use_session(&mut self, session: Option<&str>) -> Result<()> { @@ -877,6 +721,16 @@ impl Config { Ok(()) } + pub fn session_info(&self) -> Result<String> { + if let Some(session) = &self.session { + let render_options = self.render_options()?; + let mut markdown_render = MarkdownRender::init(render_options)?; + session.render(&mut markdown_render) + } else { + bail!("No session") + } + } + pub fn exit_session(&mut self) -> Result<()> { if let Some(mut session) = self.session.take() { let is_repl = self.working_mode == WorkingMode::Repl; @@ -991,6 +845,14 @@ impl Config { Ok(()) } + pub fn rag_info(&self) -> Result<String> { + if let Some(rag) = &self.rag { + rag.export() + } else { + bail!("No rag") + } + } + pub fn exit_rag(&mut self) -> Result<()> { self.rag.take(); Ok(()) @@ -1045,13 +907,147 @@ impl Config { Ok(()) } + pub fn bot_info(&self) -> Result<String> { + if let Some(bot) = &self.bot { + bot.export() + } else { + bail!("No rag") + } + } + pub fn exit_bot(&mut self) -> Result<()> { self.rag.take(); self.bot.take(); Ok(()) } - pub fn get_render_options(&self) -> Result<RenderOptions> { + pub fn apply_prelude(&mut self) -> Result<()> { + let prelude = self.prelude.clone().unwrap_or_default(); + if prelude.is_empty() { + return Ok(()); + } + let err_msg = || format!("Invalid prelude '{}", prelude); + match prelude.split_once(':') { + Some(("role", name)) => { + if self.role.is_none() && self.session.is_none() { + self.use_role(name).with_context(err_msg)?; + } + } + Some(("session", name)) => { + if self.session.is_none() { + self.use_session(Some(name)).with_context(err_msg)?; + } + } + _ => { + bail!("{}", err_msg()) + } + } + Ok(()) + } + + pub fn select_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 buffer_editor(&self) -> Option<String> { + self.buffer_editor + .clone() + .or_else(|| env::var("VISUAL").ok().or_else(|| env::var("EDITOR").ok())) + } + + 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().to_string(), None)) + .collect(), + ".model" => list_chat_models(self) + .into_iter() + .map(|v| (v.id(), Some(v.description()))) + .collect(), + ".session" => self + .list_sessions() + .into_iter() + .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", + "top_p", + "rag_top_k", + "function_calling", + "compress_threshold", + "save", + "save_session", + "highlight", + "dry_run", + "auto_copy", + ] + .into_iter() + .map(|v| (format!("{v} "), None)) + .collect(), + _ => vec![], + }; + (values, args[0]) + } else if args.len() == 2 { + let values = match args[0] { + "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 { + session.save_session() + } else { + self.save_session + }; + complete_option_bool(save_session) + } + "highlight" => complete_bool(self.highlight), + "dry_run" => complete_bool(self.dry_run), + "auto_copy" => complete_bool(self.auto_copy), + _ => vec![], + }; + (values.into_iter().map(|v| (v, None)).collect(), args[1]) + } else { + return vec![]; + }; + values + .into_iter() + .filter(|(value, _)| fuzzy_match(value, filter)) + .collect() + } + + pub fn last_reply(&self) -> &str { + self.last_message + .as_ref() + .map(|(_, reply)| reply.as_str()) + .unwrap_or_default() + } + + pub fn render_options(&self) -> Result<RenderOptions> { let theme = if self.highlight { let theme_mode = if self.light_theme { "light" } else { "dark" }; let theme_filename = format!("{theme_mode}.tmTheme"); diff --git a/src/config/session.rs b/src/config/session.rs index 97843bd..a525af3 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -86,6 +86,14 @@ 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() + } + pub fn name(&self) -> &str { &self.name } @@ -342,14 +350,6 @@ impl Session { Ok(()) } - 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() - } - pub fn add_message(&mut self, input: &Input, output: &str) -> Result<()> { let mut need_add_msg = true; if self.messages.is_empty() { diff --git a/src/main.rs b/src/main.rs index 459b50f..72be91c 100644 --- a/src/main.rs +++ b/src/main.rs @@ -158,7 +158,7 @@ async fn start_directive( text.clone() }; if *IS_STDOUT_TERMINAL { - let render_options = config.read().get_render_options()?; + let render_options = config.read().render_options()?; let mut markdown_render = MarkdownRender::init(render_options)?; println!("{}", markdown_render.render(&text).trim()); } else { @@ -209,7 +209,7 @@ async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) - } config.write().save_message(&mut input, &eval_str, &[])?; config.read().maybe_copy(&eval_str); - let render_options = config.read().get_render_options()?; + let render_options = config.read().render_options()?; let mut markdown_render = MarkdownRender::init(render_options)?; if config.read().dry_run { println!("{}", markdown_render.render(&eval_str).trim()); diff --git a/src/render/mod.rs b/src/render/mod.rs index 9fa6787..ce7191f 100644 --- a/src/render/mod.rs +++ b/src/render/mod.rs @@ -16,7 +16,7 @@ pub async fn render_stream( abort: AbortSignal, ) -> Result<()> { if *IS_STDOUT_TERMINAL { - let render_options = config.read().get_render_options()?; + let render_options = config.read().render_options()?; let mut render = MarkdownRender::init(render_options)?; markdown_stream(rx, &mut render, &abort).await } else { |
