diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/mod.rs | 53 | ||||
| -rw-r--r-- | src/config/session.rs | 10 | ||||
| -rw-r--r-- | src/main.rs | 21 | ||||
| -rw-r--r-- | src/repl/handler.rs | 2 | ||||
| -rw-r--r-- | src/repl/mod.rs | 2 |
5 files changed, 52 insertions, 36 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs index ed256e6..4437833 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -47,7 +47,8 @@ pub struct Config { /// LLM model pub model: Option<String>, /// What sampling temperature to use, between 0 and 2 - pub temperature: Option<f64>, + #[serde(rename(serialize = "temperature", deserialize = "temperature"))] + pub default_temperature: Option<f64>, /// Whether to persistently save non-session chat messages pub save: bool, /// Whether to disable highlight @@ -79,13 +80,15 @@ pub struct Config { pub model_info: ModelInfo, #[serde(skip)] pub last_message: Option<(String, String)>, + #[serde(skip)] + pub temperature: Option<f64>, } impl Default for Config { fn default() -> Self { Self { model: None, - temperature: None, + default_temperature: None, save: false, highlight: true, dry_run: false, @@ -100,6 +103,7 @@ impl Default for Config { session: None, model_info: Default::default(), last_message: None, + temperature: None, } } } @@ -134,6 +138,8 @@ impl Config { config.set_wrap(&wrap)?; } + config.temperature = config.default_temperature; + config.merge_env_vars(); config.load_roles()?; config.ensure_sessions_dir()?; @@ -233,7 +239,7 @@ impl Config { Ok(path) } - pub fn change_role(&mut self, name: &str) -> Result<String> { + pub fn set_role(&mut self, name: &str) -> Result<String> { match self.get_role(name) { Some(role) => { if let Some(session) = self.session.as_mut() { @@ -241,10 +247,11 @@ impl Config { } let output = serde_yaml::to_string(&role) .unwrap_or_else(|_| "Unable to echo role details".into()); + self.temperature = role.temperature; self.role = Some(role); Ok(output) } - None => bail!("Error: Unknown role"), + None => bail!("Unknown role `{name}`"), } } @@ -252,15 +259,21 @@ impl Config { if let Some(session) = self.session.as_mut() { session.update_role(None)?; } + self.temperature = self.default_temperature; self.role = None; Ok(()) } pub fn get_temperature(&self) -> Option<f64> { - self.role - .as_ref() - .and_then(|v| v.temperature) - .or(self.temperature) + self.temperature + } + + pub fn set_temperature(&mut self, value: Option<f64>) -> Result<()> { + self.temperature = value; + if let Some(session) = self.session.as_mut() { + session.temperature = value; + } + Ok(()) } pub fn echo_messages(&self, content: &str) -> String { @@ -318,7 +331,7 @@ impl Config { None => bail!("Invalid model"), Some(model_info) => { if let Some(session) = self.session.as_mut() { - session.model = model_info.stringify(); + session.set_model(&model_info.stringify())?; } self.model_info = model_info; Ok(()) @@ -401,12 +414,13 @@ impl Config { let unset = value == "null"; match key { "temperature" => { - if unset { - self.temperature = None; + let value = if unset { + None } else { let value = value.parse().with_context(|| "Invalid value")?; - self.temperature = Some(value); - } + Some(value) + }; + self.set_temperature(value)?; } "save" => { let value = value.parse().with_context(|| "Invalid value")?; @@ -420,7 +434,7 @@ impl Config { let value = value.parse().with_context(|| "Invalid value")?; self.dry_run = value; } - _ => bail!("Error: Unknown key `{key}`"), + _ => bail!("Unknown key `{key}`"), } Ok(()) } @@ -451,13 +465,11 @@ impl Config { 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(); + let session = Session::load(name, &session_path)?; + let model = session.model.clone(); + self.temperature = session.temperature; self.session = Some(session); + self.set_model(&model)?; } } } @@ -481,6 +493,7 @@ impl Config { pub fn end_session(&mut self) -> Result<()> { if let Some(mut session) = self.session.take() { self.last_message = None; + self.temperature = self.default_temperature; if session.should_save() { let ans = Confirm::new("Save session?").with_default(true).prompt()?; if !ans { diff --git a/src/config/session.rs b/src/config/session.rs index 5ecb1c2..b2d00bb 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -13,6 +13,7 @@ pub struct Session { pub path: Option<String>, pub model: String, pub tokens: usize, + pub temperature: Option<f64>, pub messages: Vec<Message>, #[serde(skip)] pub dirty: bool, @@ -24,9 +25,11 @@ pub struct Session { impl Session { pub fn new(name: &str, model: &str, role: Option<Role>) -> Self { + let temperature = role.as_ref().and_then(|v| v.temperature); let mut value = Self { path: None, model: model.to_string(), + temperature, tokens: 0, messages: vec![], dirty: false, @@ -58,11 +61,18 @@ impl Session { pub fn update_role(&mut self, role: Option<Role>) -> Result<()> { self.guard_empty()?; + self.temperature = role.as_ref().and_then(|v| v.temperature); self.role = role; self.update_tokens(); Ok(()) } + pub fn set_model(&mut self, model: &str) -> Result<()> { + self.model = model.to_string(); + self.update_tokens(); + Ok(()) + } + pub fn save(&mut self, session_path: &Path) -> Result<()> { if !self.should_save() { return Ok(()); diff --git a/src/main.rs b/src/main.rs index 0d12731..a71a2b5 100644 --- a/src/main.rs +++ b/src/main.rs @@ -11,7 +11,7 @@ use crate::cli::Cli; use crate::client::Client; use crate::config::{Config, SharedConfig}; -use anyhow::{anyhow, Result}; +use anyhow::Result; use clap::Parser; use client::{init_client, list_models}; use crossbeam::sync::WaitGroup; @@ -56,22 +56,15 @@ fn main() -> Result<()> { if cli.dry_run { config.write().dry_run = true; } - if let Some(session) = &cli.session { - config.write().start_session(session)?; - } if let Some(model) = &cli.model { config.write().set_model(model)?; } - let role = match &cli.role { - Some(name) => Some( - config - .read() - .get_role(name) - .ok_or_else(|| anyhow!("Unknown role '{name}'"))?, - ), - None => None, - }; - config.write().role = role; + if let Some(name) = &cli.role { + config.write().set_role(name)?; + } + if let Some(session) = &cli.session { + config.write().start_session(session)?; + } if cli.no_highlight { config.write().highlight = false; } diff --git a/src/repl/handler.rs b/src/repl/handler.rs index ff6cf56..e726661 100644 --- a/src/repl/handler.rs +++ b/src/repl/handler.rs @@ -71,7 +71,7 @@ impl ReplCmdHandler { print_now!("\n"); } ReplCmd::SetRole(name) => { - let output = self.config.write().change_role(&name)?; + let output = self.config.write().set_role(&name)?; print_now!("{}\n\n", output.trim_end()); } ReplCmd::ClearRole => { diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 46d0cca..a6f173a 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -160,7 +160,7 @@ impl Repl { } fn dump_unknown_command() { - print_now!("Error: Unknown command. Type \".help\" for more information.\n\n"); + print_now!("Unknown command. Type \".help\" for more information.\n\n"); } fn dump_repl_help() { |
