diff options
| author | sigoden <sigoden@gmail.com> | 2023-10-28 21:39:17 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-10-28 21:39:17 +0800 |
| commit | bc44026ff8828f532e9eba138e2ba35f7c247271 (patch) | |
| tree | 98c55ad31a543ff3cd56e297751d32ee9ca960a7 /src/config/mod.rs | |
| parent | 1575d441724a27854c8f9fe1510a47a3050ae28d (diff) | |
| download | aichat-bc44026ff8828f532e9eba138e2ba35f7c247271.tar.gz | |
feat: enhance session/conversation (#162)
* feat: enhance session/conversation
* updates
* updates
* cut version v0.9.0-rc2
* add .session name completion
Diffstat (limited to 'src/config/mod.rs')
| -rw-r--r-- | src/config/mod.rs | 229 |
1 files changed, 167 insertions, 62 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs index 8e5157f..73c9444 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -1,10 +1,10 @@ -mod conversation; mod message; mod role; +mod session; -use self::conversation::Conversation; use self::message::Message; use self::role::Role; +use self::session::{Session, TEMP_SESSION_NAME}; use crate::client::openai::{OpenAIClient, OpenAIConfig}; use crate::client::{all_clients, create_client_config, list_models, ClientConfig, ModelInfo}; @@ -12,12 +12,12 @@ use crate::config::message::num_tokens_from_messages; use crate::utils::{get_env_name, now}; use anyhow::{anyhow, bail, Context, Result}; -use inquire::{Confirm, Select}; +use inquire::{Confirm, Select, Text}; use parking_lot::RwLock; use serde::Deserialize; use std::{ env, - fs::{create_dir_all, read_to_string, File, OpenOptions}, + fs::{create_dir_all, read_dir, read_to_string, remove_file, File, OpenOptions}, io::Write, path::{Path, PathBuf}, process::exit, @@ -27,7 +27,9 @@ use std::{ const CONFIG_FILE_NAME: &str = "config.yaml"; const ROLES_FILE_NAME: &str = "roles.yaml"; const HISTORY_FILE_NAME: &str = "history.txt"; -const MESSAGE_FILE_NAME: &str = "messages.md"; +const MESSAGES_FILE_NAME: &str = "messages.md"; +const SESSIONS_DIR_NAME: &str = "sessions"; + const SET_COMPLETIONS: [&str; 7] = [ ".set temperature", ".set save true", @@ -46,14 +48,12 @@ pub struct Config { pub model: Option<String>, /// What sampling temperature to use, between 0 and 2 pub temperature: Option<f64>, - /// Whether to persistently save chat messages + /// Whether to persistently save non-session chat messages pub save: bool, /// Whether to disable highlight pub highlight: bool, /// Used only for debugging pub dry_run: bool, - /// If set ture, start a conversation immediately upon repl - pub conversation_first: bool, /// If set true, use light theme pub light_theme: bool, /// Automatically copy the last output to the clipboard @@ -68,11 +68,13 @@ pub struct Config { /// Current selected role #[serde(skip)] pub role: Option<Role>, - /// Current conversation + /// Current session #[serde(skip)] - pub conversation: Option<Conversation>, + pub session: Option<Session>, #[serde(skip)] pub model_info: ModelInfo, + #[serde(skip)] + pub last_message: Option<(String, String)>, } impl Default for Config { @@ -83,15 +85,15 @@ impl Default for Config { save: false, highlight: true, dry_run: false, - conversation_first: false, light_theme: false, auto_copy: false, keybindings: Default::default(), - roles: vec![], clients: vec![ClientConfig::OpenAI(OpenAIConfig::default())], + roles: vec![], role: None, - conversation: None, + session: None, model_info: Default::default(), + last_message: None, } } } @@ -123,19 +125,14 @@ impl Config { if let Some(name) = config.model.clone() { config.set_model(&name)?; } + config.merge_env_vars(); config.load_roles()?; + config.ensure_sessions_dir()?; Ok(config) } - pub fn on_repl(&mut self) -> Result<()> { - if self.conversation_first { - self.start_conversation()?; - } - Ok(()) - } - pub fn get_role(&self, name: &str) -> Option<Role> { self.roles.iter().find(|v| v.match_name(name)).map(|v| { let mut role = v.clone(); @@ -156,13 +153,20 @@ impl Config { Ok(path) } - pub fn local_file(name: &str) -> Result<PathBuf> { + pub fn local_path(name: &str) -> Result<PathBuf> { let mut path = Self::config_dir()?; path.push(name); Ok(path) } - pub fn save_message(&self, input: &str, output: &str) -> Result<()> { + pub fn save_message(&mut self, input: &str, output: &str) -> Result<()> { + self.last_message = Some((input.to_string(), output.to_string())); + + if let Some(session) = self.session.as_mut() { + session.add_message(input, output)?; + return Ok(()); + } + if !self.save { return Ok(()); } @@ -193,30 +197,40 @@ impl Config { } pub fn config_file() -> Result<PathBuf> { - Self::local_file(CONFIG_FILE_NAME) + Self::local_path(CONFIG_FILE_NAME) } pub fn roles_file() -> Result<PathBuf> { let env_name = get_env_name("roles_file"); env::var(env_name).map_or_else( - |_| Self::local_file(ROLES_FILE_NAME), + |_| Self::local_path(ROLES_FILE_NAME), |value| Ok(PathBuf::from(value)), ) } pub fn history_file() -> Result<PathBuf> { - Self::local_file(HISTORY_FILE_NAME) + Self::local_path(HISTORY_FILE_NAME) } pub fn messages_file() -> Result<PathBuf> { - Self::local_file(MESSAGE_FILE_NAME) + Self::local_path(MESSAGES_FILE_NAME) + } + + pub fn sessions_dir() -> Result<PathBuf> { + Self::local_path(SESSIONS_DIR_NAME) + } + + pub fn session_file(name: &str) -> Result<PathBuf> { + let mut path = Self::sessions_dir()?; + path.push(&format!("{name}.yaml")); + Ok(path) } pub fn change_role(&mut self, name: &str) -> Result<String> { match self.get_role(name) { Some(role) => { - if let Some(conversation) = self.conversation.as_mut() { - conversation.update_role(&role)?; + if let Some(session) = self.session.as_mut() { + session.update_role(Some(role.clone()))?; } let output = serde_yaml::to_string(&role) .unwrap_or_else(|_| "Unable to echo role details".into()); @@ -228,8 +242,8 @@ impl Config { } pub fn clear_role(&mut self) -> Result<()> { - if let Some(conversation) = self.conversation.as_ref() { - conversation.can_clear_role()?; + if let Some(session) = self.session.as_mut() { + session.update_role(None)?; } self.role = None; Ok(()) @@ -237,8 +251,8 @@ impl Config { pub fn add_prompt(&mut self, prompt: &str) -> Result<()> { let role = Role::new(prompt, self.temperature); - if let Some(conversation) = self.conversation.as_mut() { - conversation.update_role(&role)?; + if let Some(session) = self.session.as_mut() { + session.update_role(Some(role.clone()))?; } self.role = Some(role); Ok(()) @@ -253,8 +267,8 @@ impl Config { pub fn echo_messages(&self, content: &str) -> String { #[allow(clippy::option_if_let_else)] - if let Some(conversation) = self.conversation.as_ref() { - conversation.echo_messages(content) + if let Some(session) = self.session.as_ref() { + session.echo_messages(content) } else if let Some(role) = self.role.as_ref() { role.echo_messages(content) } else { @@ -264,8 +278,8 @@ impl Config { pub fn build_messages(&self, content: &str) -> Result<Vec<Message>> { #[allow(clippy::option_if_let_else)] - let messages = if let Some(conversation) = self.conversation.as_ref() { - conversation.build_emssages(content) + let messages = if let Some(session) = self.session.as_ref() { + session.build_emssages(content) } else if let Some(role) = self.role.as_ref() { role.build_messages(content) } else { @@ -282,28 +296,36 @@ impl Config { pub fn set_model(&mut self, value: &str) -> Result<()> { let models = list_models(self); + let mut model_info = None; if value.contains(':') { if let Some(model) = models.iter().find(|v| v.stringify() == value) { - self.model_info = model.clone(); - return Ok(()); + model_info = Some(model.clone()); } } else if let Some(model) = models.iter().find(|v| v.client == value) { - self.model_info = model.clone(); - return Ok(()); + model_info = Some(model.clone()); + } + match model_info { + None => bail!("Invalid model"), + Some(model_info) => { + if let Some(session) = self.session.as_mut() { + session.model = model_info.stringify(); + } + self.model_info = model_info; + Ok(()) + } } - bail!("Invalid model") } pub const fn get_reamind_tokens(&self) -> usize { let mut tokens = self.model_info.max_tokens; - if let Some(conversation) = self.conversation.as_ref() { - tokens = tokens.saturating_sub(conversation.tokens); + if let Some(session) = self.session.as_ref() { + tokens = tokens.saturating_sub(session.tokens); } tokens } pub fn info(&self) -> Result<String> { - let file_info = |path: &Path| { + let path_info = |path: &Path| { let state = if path.exists() { "" } else { " ⚠️" }; format!("{}{state}", path.display()) }; @@ -311,14 +333,14 @@ impl Config { .temperature .map_or_else(|| String::from("-"), |v| v.to_string()); let items = vec![ - ("config_file", file_info(&Self::config_file()?)), - ("roles_file", file_info(&Self::roles_file()?)), - ("messages_file", file_info(&Self::messages_file()?)), + ("config_file", path_info(&Self::config_file()?)), + ("roles_file", path_info(&Self::roles_file()?)), + ("messages_file", path_info(&Self::messages_file()?)), + ("sessions_dir", path_info(&Self::sessions_dir()?)), ("model", self.model_info.stringify()), ("temperature", temperature), ("save", self.save.to_string()), ("highlight", self.highlight.to_string()), - ("conversation_first", self.conversation_first.to_string()), ("light_theme", self.light_theme.to_string()), ("dry_run", self.dry_run.to_string()), ("keybindings", self.keybindings.stringify().into()), @@ -343,6 +365,13 @@ impl Config { .iter() .map(|v| format!(".model {}", v.stringify())), ); + completion.extend( + list_models(self) + .iter() + .map(|v| format!(".model {}", v.stringify())), + ); + let sessions = self.list_sessions().unwrap_or_default(); + completion.extend(sessions.iter().map(|v| format!(".session {}", v))); completion } @@ -380,28 +409,94 @@ impl Config { Ok(()) } - pub fn start_conversation(&mut self) -> Result<()> { - if self.conversation.is_some() && self.get_reamind_tokens() > 0 { - let ans = Confirm::new("Already in a conversation, start a new one?") - .with_default(true) - .prompt()?; - if !ans { - return Ok(()); + pub fn start_session(&mut self, session: &Option<String>) -> Result<()> { + if self.session.is_some() { + bail!("Already in a session, please use '.clear session' to exit the session first?"); + } + match session { + None => { + let session_file = Self::session_file(TEMP_SESSION_NAME)?; + if session_file.exists() { + remove_file(session_file) + .with_context(|| "Failed to clean previous session")?; + } + self.session = Some(Session::new( + TEMP_SESSION_NAME, + &self.model_info.stringify(), + self.role.clone(), + )); + } + Some(name) => { + let session_path = Self::session_file(name)?; + if !session_path.exists() { + self.session = Some(Session::new( + name, + &self.model_info.stringify(), + self.role.clone(), + )); + } else { + let mut session = Session::load(name, &session_path)?; + if let Some(role) = &session.role { + self.change_role(&role.name)?; + } + self.set_model(&session.model)?; + session.update_tokens(); + self.session = Some(session); + } + } + } + if let Some(session) = self.session.as_mut() { + if session.is_empty() { + if let Some((input, output)) = &self.last_message { + let ans = Confirm::new( + "Start a session that incorporates the last question and answer?", + ) + .with_default(false) + .prompt()?; + if ans { + session.add_message(input, output)?; + } + } } } - self.conversation = Some(Conversation::new(self.role.clone())); Ok(()) } - pub fn end_conversation(&mut self) { - self.conversation = None; + pub fn end_session(&mut self) -> Result<()> { + if let Some(mut session) = self.session.take() { + self.last_message = None; + if session.should_save() { + let ans = Confirm::new("Save session?").with_default(true).prompt()?; + if !ans { + return Ok(()); + } + let mut name = session.name.clone(); + if session.is_temp() { + name = Text::new("Session name:").with_default(&name).prompt()?; + } + let session_path = Self::session_file(&name)?; + session.save(&session_path)?; + } + } + Ok(()) } - pub fn save_conversation(&mut self, input: &str, output: &str) -> Result<()> { - if let Some(conversation) = self.conversation.as_mut() { - conversation.add_message(input, output)?; + pub fn list_sessions(&self) -> Result<Vec<String>> { + let sessions_dir = Self::sessions_dir()?; + match read_dir(&sessions_dir) { + Ok(rd) => { + let mut names = vec![]; + for entry in rd { + let entry = entry?; + let name = entry.file_name(); + if let Some(name) = name.to_string_lossy().strip_suffix(".yaml") { + names.push(name.to_string()); + } + } + Ok(names) + } + Err(_) => Ok(vec![]), } - Ok(()) } pub const fn get_render_options(&self) -> (bool, bool) { @@ -463,6 +558,16 @@ impl Config { } } + fn ensure_sessions_dir(&self) -> Result<()> { + let sessions_dir = Self::sessions_dir()?; + if !sessions_dir.exists() { + create_dir_all(&sessions_dir).with_context(|| { + format!("Failed to create session_dir '{}'", sessions_dir.display()) + })?; + } + Ok(()) + } + fn compat_old_config(&mut self, config_path: &PathBuf) -> Result<()> { let content = read_to_string(config_path)?; let value: serde_json::Value = serde_yaml::from_str(&content)?; |
