diff options
| author | sigoden <sigoden@gmail.com> | 2024-03-27 07:33:21 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-03-27 07:33:21 +0800 |
| commit | 8da9fa5f4c24ace21fbf0d8d406db44fbc6537ab (patch) | |
| tree | 39b541356c70c022a73193983884389a6c3d0f06 /src | |
| parent | 582f56e915a7c88eb0372496e298c9944f5832ae (diff) | |
| download | aichat-8da9fa5f4c24ace21fbf0d8d406db44fbc6537ab.tar.gz | |
feat: add sepereate `save_session` config item to session (#377)
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/mod.rs | 35 | ||||
| -rw-r--r-- | src/config/session.rs | 24 | ||||
| -rw-r--r-- | src/main.rs | 6 |
3 files changed, 48 insertions, 17 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs index e2fe864..d97b4b2 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -53,6 +53,8 @@ pub struct Config { pub dry_run: bool, /// Whether to save the message pub save: bool, + /// Whether to save the session automatically + pub save_session: bool, /// Whether to disable highlight pub highlight: bool, /// Whether to use a light theme @@ -67,8 +69,6 @@ pub struct Config { pub keybindings: Keybindings, /// Set a default role or session (role:<name>, session:<name>) pub prelude: String, - /// Whether to save the session automatically - pub save_session: bool, /// Compress session if tokens exceed this value (>=1000) pub compress_threshold: usize, /// The prompt for summarizing session messages @@ -362,6 +362,14 @@ impl Config { } } + pub fn set_save_session(&mut self, value: 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: usize) { self.compress_threshold = value; if let Some(session) = self.session.as_mut() { @@ -439,6 +447,7 @@ impl Config { ("temperature", temperature), ("dry_run", self.dry_run.to_string()), ("save", self.save.to_string()), + ("save_session", self.save_session.to_string()), ("highlight", self.highlight.to_string()), ("light_theme", self.light_theme.to_string()), ("wrap", wrap), @@ -505,6 +514,7 @@ impl Config { "temperature ", "compress_threshold", "save ", + "save_session ", "highlight ", "dry_run ", "auto_copy ", @@ -519,6 +529,13 @@ impl Config { let to_vec = |v: bool| vec![v.to_string()]; let values = match args[0] { "save" => to_vec(!self.save), + "save_session" => { + if let Some(session) = &self.session { + to_vec(!session.save_session()) + } else { + to_vec(!self.save_session) + } + } "highlight" => to_vec(!self.highlight), "dry_run" => to_vec(!self.dry_run), "auto_copy" => to_vec(!self.auto_copy), @@ -560,6 +577,10 @@ impl Config { let value = value.parse().with_context(|| "Invalid value")?; self.save = value; } + "save_session" => { + let value = value.parse().with_context(|| "Invalid value")?; + self.set_save_session(value); + } "highlight" => { let value = value.parse().with_context(|| "Invalid value")?; self.highlight = value; @@ -591,16 +612,12 @@ impl Config { format!("Failed to cleanup previous '{TEMP_SESSION_NAME}' session") })?; } - self.session = Some(Session::new( - TEMP_SESSION_NAME, - self.model.clone(), - self.temperature, - )); + self.session = Some(Session::new(self, TEMP_SESSION_NAME)); } Some(name) => { let session_path = Self::session_file(name)?; if !session_path.exists() { - self.session = Some(Session::new(name, self.model.clone(), self.temperature)); + self.session = Some(Session::new(self, name)); } else { let session = Session::load(name, &session_path)?; let model = session.model().to_string(); @@ -636,7 +653,7 @@ impl Config { if session.dirty { // If it's a temporary session, we'll always prompt to save on exit // If it's named, we'll save automatically if they've set the save flag and prompt if they haven't - if !self.save_session || session.is_temp() { + if !session.save_session() || session.is_temp() { if !interactive { // If we're not interactive, we will not prompt and will not save return Ok(()); diff --git a/src/config/session.rs b/src/config/session.rs index 2e0d69f..a52a86f 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -1,5 +1,5 @@ use super::input::resolve_data_url; -use super::{Input, Model}; +use super::{Config, Input, Model}; use crate::client::{Message, MessageContent, MessageRole}; use crate::render::MarkdownRender; @@ -18,6 +18,8 @@ pub struct Session { #[serde(rename(serialize = "model", deserialize = "model"))] model_id: String, temperature: Option<f64>, + #[serde(default)] + save_session: bool, messages: Vec<Message>, #[serde(default)] data_urls: HashMap<String, String>, @@ -37,10 +39,11 @@ pub struct Session { } impl Session { - pub fn new(name: &str, model: Model, temperature: Option<f64>) -> Self { + pub fn new(config: &Config, name: &str) -> Self { Self { - model_id: model.id(), - temperature, + model_id: config.model.id(), + temperature: config.temperature, + save_session: config.save_session, messages: vec![], compressed_messages: vec![], compress_threshold: None, @@ -49,7 +52,7 @@ impl Session { path: None, dirty: false, compressing: false, - model, + model: config.model.clone(), } } @@ -77,6 +80,10 @@ impl Session { self.temperature } + pub fn save_session(&self) -> bool { + self.save_session + } + pub fn need_compress(&self, current_compress_threshold: usize) -> bool { let threshold = self .compress_threshold @@ -102,6 +109,7 @@ impl Session { if let Some(temperature) = self.temperature() { data["temperature"] = temperature.into(); } + data["save_session"] = self.save_session.into(); data["total_tokens"] = tokens.into(); if let Some(conext_window) = self.model.max_input_tokens { data["max_input_tokens"] = conext_window.into(); @@ -129,6 +137,8 @@ impl Session { items.push(("temperature", temperature.to_string())); } + items.push(("save_session", self.save_session.to_string())); + if let Some(compress_threshold) = self.compress_threshold { items.push(("compress_threshold", compress_threshold.to_string())); } @@ -188,6 +198,10 @@ impl Session { self.temperature = value; } + pub fn set_save_session(&mut self, value: bool) { + self.save_session = value; + } + pub fn set_compress_threshold(&mut self, value: usize) { self.compress_threshold = Some(value); } diff --git a/src/main.rs b/src/main.rs index 502f6d3..8abca98 100644 --- a/src/main.rs +++ b/src/main.rs @@ -54,9 +54,6 @@ fn main() -> Result<()> { if let Some(wrap) = &cli.wrap { config.write().set_wrap(wrap)?; } - if cli.save_session { - config.write().save_session = true; - } if cli.light_theme { config.write().light_theme = true; } @@ -78,6 +75,9 @@ fn main() -> Result<()> { if let Some(model) = &cli.model { config.write().set_model(model)?; } + if cli.save_session { + config.write().set_save_session(true); + } if cli.no_highlight { config.write().highlight = false; } |
