summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--src/config/mod.rs53
-rw-r--r--src/config/session.rs10
-rw-r--r--src/main.rs21
-rw-r--r--src/repl/handler.rs2
-rw-r--r--src/repl/mod.rs2
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() {