diff options
| -rw-r--r-- | assets/roles/%create-title%.md | 11 | ||||
| -rw-r--r-- | src/config/input.rs | 14 | ||||
| -rw-r--r-- | src/config/mod.rs | 142 | ||||
| -rw-r--r-- | src/config/role.rs | 1 | ||||
| -rw-r--r-- | src/config/session.rs | 133 | ||||
| -rw-r--r-- | src/main.rs | 2 | ||||
| -rw-r--r-- | src/repl/mod.rs | 20 |
7 files changed, 251 insertions, 72 deletions
diff --git a/assets/roles/%create-title%.md b/assets/roles/%create-title%.md new file mode 100644 index 0000000..e6f97ff --- /dev/null +++ b/assets/roles/%create-title%.md @@ -0,0 +1,11 @@ +Create a concise, 3-6 word title. + +**Notes**: +- Avoid quotation marks or emojis +- RESPOND ONLY WITH TITLE SLUG TEXT + +**Examples**: +stock-market-trends +perfect-chocolate-chip-recipe +remote-work-productivity-tips +video-game-development-insights diff --git a/src/config/input.rs b/src/config/input.rs index 18b192b..5db172f 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -14,6 +14,7 @@ use std::{collections::HashMap, fs::File, io::Read, path::Path}; use unicode_width::{UnicodeWidthChar, UnicodeWidthStr}; const IMAGE_EXTS: [&str; 5] = ["png", "jpeg", "jpg", "webp", "gif"]; +const SUMMARY_MAX_WIDTH: usize = 80; lazy_static::lazy_static! { static ref URL_RE: Regex = Regex::new(r"^[A-Za-z0-9_-]{2,}:/").unwrap(); @@ -58,12 +59,11 @@ impl Input { pub async fn from_files( config: &GlobalConfig, - text: &str, + raw_text: &str, paths: Vec<String>, role: Option<Role>, ) -> Result<Self> { let spinner = create_spinner("Loading files").await; - let raw_text = text.to_string(); let mut raw_paths = vec![]; let mut local_paths = vec![]; let mut remote_urls = vec![]; @@ -85,8 +85,8 @@ impl Input { spinner.stop(); let (files, medias, data_urls) = ret.context("Failed to load files")?; let mut texts = vec![]; - if !text.is_empty() { - texts.push(text.to_string()); + if !raw_text.is_empty() { + texts.push(raw_text.to_string()); }; if !files.is_empty() { texts.push(String::new()); @@ -98,7 +98,7 @@ impl Input { Ok(Self { config: config.clone(), text: texts.join("\n"), - raw: (raw_text, raw_paths), + raw: (raw_text.to_string(), raw_paths), patched_text: None, continue_output: None, regenerate: false, @@ -280,12 +280,12 @@ impl Input { .chars() .map(|c| if c.is_control() { ' ' } else { c }) .collect(); - if text.width_cjk() > 70 { + if text.width_cjk() > SUMMARY_MAX_WIDTH { let mut sum_width = 0; let mut chars = vec![]; for c in text.chars() { sum_width += c.width_cjk().unwrap_or(1); - if sum_width > 67 { + if sum_width > SUMMARY_MAX_WIDTH - 3 { chars.extend(['.', '.', '.']); break; } diff --git a/src/config/mod.rs b/src/config/mod.rs index 136fee1..e578620 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -5,7 +5,9 @@ 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, CODE_ROLE, CREATE_TITLE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE, +}; use self::session::Session; use crate::client::{ @@ -37,6 +39,10 @@ use std::{ }; use syntect::highlighting::ThemeSet; +pub const TEMP_ROLE_NAME: &str = "%%"; +pub const TEMP_RAG_NAME: &str = "temp"; +pub const TEMP_SESSION_NAME: &str = "temp"; + /// Monokai Extended const DARK_THEME: &[u8] = include_bytes!("../../assets/monokai-extended.theme.bin"); const LIGHT_THEME: &[u8] = include_bytes!("../../assets/monokai-extended-light.theme.bin"); @@ -52,10 +58,6 @@ const FUNCTIONS_FILE_NAME: &str = "functions.json"; const FUNCTIONS_BIN_DIR_NAME: &str = "bin"; const AGENTS_DIR_NAME: &str = "agents"; -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"; const SERVE_ADDR: &str = "127.0.0.1:8000"; @@ -339,7 +341,10 @@ impl Config { } pub fn session_file(&self, name: &str) -> PathBuf { - self.sessions_dir().join(format!("{name}.yaml")) + match name.split_once("/") { + Some((dir, name)) => self.sessions_dir().join(dir).join(format!("{name}.yaml")), + None => self.sessions_dir().join(format!("{name}.yaml")), + } } pub fn rag_file(&self, name: &str) -> PathBuf { @@ -1081,7 +1086,10 @@ impl Config { let session_name = match &self.session { Some(session) => match name { Some(v) => v.to_string(), - None => session.name().to_string(), + None => session + .autoname() + .unwrap_or_else(|| session.name()) + .to_string(), }, None => bail!("No session"), }; @@ -1124,9 +1132,9 @@ impl Config { Ok(()) } - pub fn set_append_conversation(&mut self) -> Result<()> { + pub fn set_save_session_this_time(&mut self) -> Result<()> { if let Some(session) = self.session.as_mut() { - session.set_append_conversation(); + session.set_save_session_this_time(); } else { bail!("No session") } @@ -1137,14 +1145,42 @@ impl Config { list_file_names(self.sessions_dir(), ".yaml") } - pub fn should_compress_session(&mut self) -> bool { - if let Some(session) = self.session.as_mut() { - if session.need_compress(self.compress_threshold) { - session.set_compressing(true); - return true; + pub fn list_autoname_sessions(&self) -> Vec<String> { + list_file_names(self.sessions_dir().join("_"), ".yaml") + } + + pub fn maybe_compress_session(config: GlobalConfig) { + let mut need_compress = false; + { + let mut config = config.write(); + let compress_threshold = config.compress_threshold; + if let Some(session) = config.session.as_mut() { + if session.need_compress(compress_threshold) { + session.set_compressing(true); + need_compress = true; + } } + }; + if !need_compress { + return; } - false + let color = if config.read().light_theme { + nu_ansi_term::Color::LightGray + } else { + nu_ansi_term::Color::DarkGray + }; + print!( + "\n📢 {}\n", + color.italic().paint("Compressing the session."), + ); + tokio::spawn(async move { + if let Err(err) = Config::compress_session(&config).await { + warn!("Failed to compress the session: {err}"); + } + if let Some(session) = config.write().session.as_mut() { + session.set_compressing(false); + } + }); } pub async fn compress_session(config: &GlobalConfig) -> Result<()> { @@ -1156,7 +1192,13 @@ impl Config { } None => bail!("No session"), } - let input = Input::from_str(config, config.read().summarize_prompt(), None); + + let prompt = config + .read() + .summarize_prompt + .clone() + .unwrap_or_else(|| SUMMARIZE_PROMPT.into()); + let input = Input::from_str(config, &prompt, None); let client = input.create_client()?; let summary = client.chat_completions(input).await?.text; let summary_prompt = config @@ -1171,10 +1213,6 @@ impl Config { Ok(()) } - pub fn summarize_prompt(&self) -> &str { - self.summarize_prompt.as_deref().unwrap_or(SUMMARIZE_PROMPT) - } - pub fn is_compressing_session(&self) -> bool { self.session .as_ref() @@ -1182,10 +1220,51 @@ impl Config { .unwrap_or_default() } - pub fn end_compressing_session(&mut self) { - if let Some(session) = self.session.as_mut() { - session.set_compressing(false); + pub fn maybe_autoname_session(config: GlobalConfig) { + let mut need_autoname = false; + if let Some(session) = config.write().session.as_mut() { + if session.need_autoname() { + session.set_autonaming(true); + need_autoname = true; + } + } + if !need_autoname { + return; + } + let color = if config.read().light_theme { + nu_ansi_term::Color::LightGray + } else { + nu_ansi_term::Color::DarkGray + }; + print!("\n📢 {}\n", color.italic().paint("Autonaming the session."),); + tokio::spawn(async move { + if let Err(err) = Config::autoname_session(&config).await { + warn!("Failed to autonaming the session: {err}"); + } + if let Some(session) = config.write().session.as_mut() { + session.set_autonaming(false); + } + }); + } + + pub async fn autoname_session(config: &GlobalConfig) -> Result<()> { + let text = match config + .read() + .session + .as_ref() + .and_then(|v| v.chat_history_for_autonaming()) + { + Some(v) => v, + None => bail!("No chat history"), + }; + let role = config.read().retrieve_role(CREATE_TITLE_ROLE)?; + let input = Input::from_str(config, &text, Some(role)); + let client = input.create_client()?; + let text = client.chat_completions(input).await?.text; + if let Some(session) = config.write().session.as_mut() { + session.set_autoname(&text); } + Ok(()) } pub async fn use_rag( @@ -1562,7 +1641,19 @@ impl Config { .into_iter() .map(|v| (v.id(), Some(v.description()))) .collect(), - ".session" => map_completion_values(self.list_sessions()), + ".session" => { + if args[0].starts_with("_/") { + map_completion_values( + self.list_autoname_sessions() + .iter() + .rev() + .map(|v| format!("_/{}", v)) + .collect::<Vec<String>>(), + ) + } else { + map_completion_values(self.list_sessions()) + } + } ".rag" => map_completion_values(Self::list_rags()), ".agent" => map_completion_values(list_agents()), ".starter" => match &self.agent { @@ -1772,6 +1863,9 @@ impl Config { } if let Some(session) = &self.session { output.insert("session", session.name().to_string()); + if let Some(autoname) = session.autoname() { + output.insert("session_autoname", autoname.to_string()); + } output.insert("dirty", session.dirty().to_string()); let (tokens, percent) = session.tokens_usage(); output.insert("consume_tokens", tokens.to_string()); diff --git a/src/config/role.rs b/src/config/role.rs index ecd7eb0..6612e10 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -11,6 +11,7 @@ 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 CREATE_TITLE_ROLE: &str = "%create-title%"; pub const INPUT_PLACEHOLDER: &str = "__INPUT__"; diff --git a/src/config/session.rs b/src/config/session.rs index 633f621..23658fd 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -5,6 +5,7 @@ use crate::client::{Message, MessageContent, MessageRole}; use crate::render::MarkdownRender; use anyhow::{bail, Context, Result}; +use fancy_regex::Regex; use inquire::{validator::Validation, Confirm, Text}; use serde::{Deserialize, Serialize}; use serde_json::json; @@ -12,6 +13,10 @@ use std::collections::HashMap; use std::fs::{read_to_string, write}; use std::path::Path; +lazy_static::lazy_static! { + static ref RE_AUTONAME_PREFIX: Regex = Regex::new(r"\d{8}T\d{6}-").unwrap(); +} + #[derive(Debug, Clone, Default, Deserialize, Serialize)] pub struct Session { #[serde(rename(serialize = "model", deserialize = "model"))] @@ -32,12 +37,11 @@ pub struct Session { #[serde(default, skip_serializing_if = "IndexMap::is_empty")] agent_variables: IndexMap<String, String>, - #[serde(default, skip_serializing_if = "HashMap::is_empty")] - data_urls: HashMap<String, String>, #[serde(default, skip_serializing_if = "Vec::is_empty")] compressed_messages: Vec<Message>, - messages: Vec<Message>, + #[serde(default, skip_serializing_if = "HashMap::is_empty")] + data_urls: HashMap<String, String>, #[serde(skip)] model: Model, @@ -50,9 +54,11 @@ pub struct Session { #[serde(skip)] dirty: bool, #[serde(skip)] - append_conversation: bool, + save_session_this_time: bool, #[serde(skip)] compressing: bool, + #[serde(skip)] + autoname: Option<AutoName>, } impl Session { @@ -75,8 +81,17 @@ impl Session { serde_yaml::from_str(&content).with_context(|| format!("Invalid session {}", name))?; session.model = Model::retrieve_chat(config, &session.model_id)?; - session.name = name.to_string(); - session.path = Some(path.display().to_string()); + + if let Some(autoname) = name.strip_prefix("_/") { + session.name = TEMP_SESSION_NAME.to_string(); + session.path = None; + if let Ok(true) = RE_AUTONAME_PREFIX.is_match(autoname) { + session.autoname = Some(AutoName::new(autoname[16..].to_string())); + } + } else { + session.name = name.to_string(); + session.path = Some(path.display().to_string()); + } if let Some(role_name) = &session.role_name { if let Ok(role) = config.retrieve_role(role_name) { @@ -103,19 +118,10 @@ impl Session { self.dirty } - pub fn compressing(&self) -> bool { - self.compressing - } - pub fn save_session(&self) -> Option<bool> { self.save_session } - pub fn need_compress(&self, global_compress_threshold: usize) -> bool { - let threshold = self.compress_threshold.unwrap_or(global_compress_threshold); - threshold > 0 && self.tokens() > threshold - } - pub fn tokens(&self) -> usize { self.model().total_tokens(&self.messages) } @@ -171,6 +177,10 @@ impl Session { items.push(("path", path.to_string())); } + if let Some(autoname) = self.autoname() { + items.push(("autoname", autoname.to_string())); + } + items.push(("model", self.model().id())); if let Some(temperature) = self.temperature() { @@ -279,17 +289,14 @@ impl Session { } pub fn set_save_session(&mut self, value: Option<bool>) { - if self.name == TEMP_SESSION_NAME { - return; - } if self.save_session != value { self.save_session = value; self.dirty = true; } } - pub fn set_append_conversation(&mut self) { - self.append_conversation = true; + pub fn set_save_session_this_time(&mut self) { + self.save_session_this_time = true; } pub fn set_compress_threshold(&mut self, value: Option<usize>) { @@ -299,6 +306,21 @@ impl Session { } } + pub fn need_compress(&self, global_compress_threshold: usize) -> bool { + if self.compressing { + return false; + } + let threshold = self.compress_threshold.unwrap_or(global_compress_threshold); + if threshold < 1 { + return false; + } + self.tokens() > threshold + } + + pub fn compressing(&self) -> bool { + self.compressing + } + pub fn set_compressing(&mut self, compressing: bool) { self.compressing = compressing; } @@ -323,12 +345,39 @@ impl Session { self.dirty = true; } + pub fn need_autoname(&self) -> bool { + self.autoname.as_ref().map(|v| v.need()).unwrap_or_default() + } + + pub fn set_autonaming(&mut self, naming: bool) { + if let Some(v) = self.autoname.as_mut() { + v.naming = naming; + } + } + + pub fn chat_history_for_autonaming(&self) -> Option<String> { + self.autoname.as_ref().and_then(|v| v.chat_history.clone()) + } + + pub fn autoname(&self) -> Option<&str> { + self.autoname.as_ref().and_then(|v| v.name.as_deref()) + } + + pub fn set_autoname(&mut self, value: &str) { + let name = value + .chars() + .map(|v| if v.is_alphanumeric() { v } else { '-' }) + .collect(); + self.autoname = Some(AutoName::new(name)); + } + pub fn exit(&mut self, session_dir: &Path, is_repl: bool) -> Result<()> { let mut save_session = self.save_session(); - if self.append_conversation { + if self.save_session_this_time { save_session = Some(true); } if self.dirty && save_session != Some(false) { + let mut session_dir = session_dir.to_path_buf(); let mut session_name = self.name().to_string(); if save_session.is_none() { if !is_repl { @@ -353,9 +402,16 @@ impl Session { .prompt()?; } } else if save_session == Some(true) && session_name == TEMP_SESSION_NAME { + session_dir = session_dir.join("_"); + ensure_parent_exists(&session_dir).with_context(|| { + format!("Failed to create directory '{}'", session_dir.display()) + })?; + let now = chrono::Local::now(); - let formatted_time = now.format("%Y%m%dT%H:%M:%S").to_string(); - session_name = format!("{TEMP_SESSION_NAME}-{formatted_time}"); + session_name = now.format("%Y%m%dT%H%M%S").to_string(); + if let Some(autoname) = self.autoname() { + session_name = format!("{session_name}-{autoname}") + } } let session_path = session_dir.join(format!("{session_name}.yaml")); self.save(&session_name, &session_path, is_repl)?; @@ -413,6 +469,11 @@ impl Session { } } else { if self.messages.is_empty() { + if self.name == TEMP_SESSION_NAME && self.save_session == Some(true) { + let raw_input = input.raw(); + let chat_history = format!("USER: {raw_input}\nASSISTANT: {output}\n"); + self.autoname = Some(AutoName::new_from_chat_history(chat_history)); + } self.messages.extend(input.role().build_messages(input)); } else { self.messages @@ -438,6 +499,7 @@ impl Session { self.messages.clear(); self.compressed_messages.clear(); self.data_urls.clear(); + self.autoname = None; self.dirty = true; } @@ -532,3 +594,28 @@ impl RoleLike for Session { } } } + +#[derive(Debug, Clone, Default)] +struct AutoName { + naming: bool, + chat_history: Option<String>, + name: Option<String>, +} + +impl AutoName { + pub fn new(name: String) -> Self { + Self { + name: Some(name), + ..Default::default() + } + } + pub fn new_from_chat_history(chat_history: String) -> Self { + Self { + chat_history: Some(chat_history), + ..Default::default() + } + } + pub fn need(&self) -> bool { + !self.naming && self.chat_history.is_some() && self.name.is_none() + } +} diff --git a/src/main.rs b/src/main.rs index d06b5ab..8ca1e31 100644 --- a/src/main.rs +++ b/src/main.rs @@ -134,7 +134,7 @@ async fn run(config: GlobalConfig, cli: Cli, text: Option<String>) -> Result<()> config.write().empty_session()?; } if cli.save_session { - config.write().set_append_conversation()?; + config.write().set_save_session_this_time()?; } if cli.info { let info = config.read().info()?; diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 1e84146..98d166a 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -14,7 +14,6 @@ use crate::utils::{create_abort_signal, create_spinner, set_text, temp_file, Abo use anyhow::{bail, Context, Result}; use fancy_regex::Regex; -use nu_ansi_term::Color; use reedline::{ default_emacs_keybindings, default_vi_insert_keybindings, default_vi_normal_keybindings, ColumnarMenu, EditCommand, EditMode, Emacs, KeyCode, KeyModifiers, Keybindings, Reedline, @@ -301,6 +300,7 @@ impl Repl { }, ".session" => { self.config.write().use_session(args)?; + Config::maybe_autoname_session(self.config.clone()); } ".rag" => { Config::use_rag(&self.config, args, self.abort_signal.clone()).await?; @@ -654,22 +654,8 @@ async fn ask( ) .await } else { - if config.write().should_compress_session() { - let config = config.clone(); - let color = if config.read().light_theme { - Color::LightGray - } else { - Color::DarkGray - }; - print!( - "\n📢 {}\n", - color.italic().paint("Compressing the session."), - ); - tokio::spawn(async move { - let _ = Config::compress_session(&config).await; - config.write().end_compressing_session(); - }); - } + Config::maybe_autoname_session(config.clone()); + Config::maybe_compress_session(config.clone()); Ok(()) } } |
