diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/config/input.rs | 23 | ||||
| -rw-r--r-- | src/config/mod.rs | 54 | ||||
| -rw-r--r-- | src/config/role.rs | 10 | ||||
| -rw-r--r-- | src/config/session.rs | 35 | ||||
| -rw-r--r-- | src/rag/mod.rs | 5 | ||||
| -rw-r--r-- | src/repl/mod.rs | 2 |
6 files changed, 69 insertions, 60 deletions
diff --git a/src/config/input.rs b/src/config/input.rs index 0406a5e..4d0ed7f 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -1,7 +1,7 @@ use super::{role::Role, session::Session, GlobalConfig}; use crate::client::{ - init_client, list_chat_models, ChatCompletionsData, Client, ImageUrl, Message, MessageContent, + init_client, ChatCompletionsData, Client, ImageUrl, Message, MessageContent, MessageContentPart, MessageRole, Model, }; use crate::function::{ToolCallResult, ToolResults}; @@ -165,18 +165,15 @@ impl Input { } pub fn model(&self) -> Model { - let model = self.config.read().model.clone(); - if let Some(model_id) = self.role().and_then(|v| v.model_id.clone()) { - if model.id() != model_id { - if let Some(model) = list_chat_models(&self.config.read()) - .into_iter() - .find(|v| v.id() == model_id) - { - return model.clone(); - } - } - }; - model + if let Some(session) = self.session(&self.config.read().session) { + return session.model.clone(); + } else if let Some(model) = self + .role() + .and_then(|v| v.retrieve_model(&self.config.read())) + { + return model; + } + self.config.read().model.clone() } pub fn create_client(&self) -> Result<Box<dyn Client>> { diff --git a/src/config/mod.rs b/src/config/mod.rs index 4003a63..0240a80 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -19,7 +19,7 @@ use crate::utils::{ }; use anyhow::{anyhow, bail, Context, Result}; -use inquire::{Confirm, Select, Text}; +use inquire::{Confirm, Select}; use is_terminal::IsTerminal; use parking_lot::RwLock; use serde::Deserialize; @@ -363,8 +363,10 @@ impl Config { } pub fn exit_role(&mut self) -> Result<()> { + if self.session.is_none() { + self.restore_model()?; + } self.role = None; - self.restore_model()?; Ok(()) } @@ -386,6 +388,10 @@ impl Config { flags } + pub fn has_role_or_session(&self) -> bool { + self.role.is_some() || self.session.is_some() + } + pub fn set_temperature(&mut self, value: Option<f64>) { if let Some(session) = self.session.as_mut() { session.set_temperature(value); @@ -437,8 +443,7 @@ impl Config { } pub fn set_model(&mut self, value: &str) -> Result<()> { - let models = list_chat_models(self); - let model = Model::find(&models, value); + let model = Model::find(&list_chat_models(self), value); match model { None => bail!("No model '{}'", value), Some(model) => { @@ -740,23 +745,10 @@ impl Config { pub fn exit_session(&mut self) -> Result<()> { if let Some(mut session) = self.session.take() { + let is_repl = self.working_mode == WorkingMode::Repl; + let sessions_dir = Self::sessions_dir()?; + session.exit(&sessions_dir, is_repl)?; self.last_message = None; - let save_session = session.save_session(); - if session.dirty && save_session != Some(false) { - if save_session.is_none() { - if self.working_mode != WorkingMode::Repl { - return Ok(()); - } - let ans = Confirm::new("Save session?").with_default(false).prompt()?; - if !ans { - return Ok(()); - } - while session.is_temp() { - session.name = Text::new("Session name:").prompt()?; - } - } - Self::save_session_to_file(&mut session)?; - } self.restore_model()?; } Ok(()) @@ -767,7 +759,8 @@ impl Config { if !name.is_empty() { session.name = name.to_string(); } - Self::save_session_to_file(session)?; + let sessions_dir = Self::sessions_dir()?; + session.save(&sessions_dir)?; } Ok(()) } @@ -1032,20 +1025,6 @@ impl Config { .with_context(|| format!("Failed to create/append {}", path.display())) } - fn save_session_to_file(session: &mut Session) -> Result<()> { - let session_path = Self::session_file(session.name())?; - let sessions_dir = session_path - .parent() - .ok_or_else(|| anyhow!("Unable to save session file to {}", session_path.display()))?; - if !sessions_dir.exists() { - create_dir_all(sessions_dir).with_context(|| { - format!("Failed to create session_dir '{}'", sessions_dir.display()) - })?; - } - session.save(&session_path)?; - Ok(()) - } - fn load_config_file(config_path: &Path) -> Result<Self> { let content = read_to_string(config_path) .with_context(|| format!("Failed to load config at {}", config_path.display()))?; @@ -1117,13 +1096,12 @@ impl Config { bail!("No available model"); } - let model_id = models[0].id(); - self.model_id.clone_from(&model_id); - model_id + models[0].id() } else { self.model_id.clone() }; self.set_model(&model_id)?; + self.model_id = model_id; Ok(()) } diff --git a/src/config/role.rs b/src/config/role.rs index 515810c..1b49bad 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -1,7 +1,7 @@ -use super::Input; +use super::{Config, Input}; use crate::{ - client::{Message, MessageContent, MessageRole, Model}, + client::{list_chat_models, Message, MessageContent, MessageRole, Model}, utils::{detect_os, detect_shell}, }; @@ -99,6 +99,12 @@ async function timeout(ms) { self.prompt.contains(INPUT_PLACEHOLDER) } + pub fn retrieve_model(&self, config: &Config) -> Option<Model> { + self.model_id + .as_ref() + .and_then(|model_id| Model::find(&list_chat_models(config), model_id)) + } + pub fn set_model(&mut self, model: &Model) { self.model_id = Some(model.id()); } diff --git a/src/config/session.rs b/src/config/session.rs index 14e0731..908cc5b 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -5,10 +5,11 @@ use crate::client::{Message, MessageContent, MessageRole}; use crate::render::MarkdownRender; use anyhow::{bail, Context, Result}; +use inquire::{Confirm, Text}; use serde::{Deserialize, Serialize}; use serde_json::json; use std::collections::HashMap; -use std::fs::{self, read_to_string}; +use std::fs::{self, create_dir_all, read_to_string}; use std::path::Path; pub const TEMP_SESSION_NAME: &str = "temp"; @@ -302,12 +303,40 @@ impl Session { self.dirty = true; } - pub fn save(&mut self, session_path: &Path) -> Result<()> { + pub fn exit(&mut self, sessions_dir: &Path, is_repl: bool) -> Result<()> { + let save_session = self.save_session(); + if self.dirty && save_session != Some(false) { + if save_session.is_none() { + if !is_repl { + return Ok(()); + } + let ans = Confirm::new("Save session?").with_default(false).prompt()?; + if !ans { + return Ok(()); + } + while self.is_temp() { + self.name = Text::new("Session name:").prompt()?; + } + } + self.save(sessions_dir)?; + } + Ok(()) + } + + pub fn save(&mut self, sessions_dir: &Path) -> Result<()> { + let mut session_path = sessions_dir.to_path_buf(); + session_path.push(format!("{}.yaml", self.name())); + if !sessions_dir.exists() { + create_dir_all(sessions_dir).with_context(|| { + format!("Failed to create session_dir '{}'", sessions_dir.display()) + })?; + } + self.path = Some(session_path.display().to_string()); let content = serde_yaml::to_string(&self) .with_context(|| format!("Failed to serde session {}", self.name))?; - fs::write(session_path, content).with_context(|| { + fs::write(&session_path, content).with_context(|| { format!( "Failed to write session {} to {}", self.name, diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 387d3d9..5ab9de5 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -372,9 +372,8 @@ pub fn split_vector_id(value: VectorID) -> (usize, usize) { } fn retrieve_embedding_model(config: &Config, model_id: &str) -> Result<Model> { - let models = list_embedding_models(config); - let model = - Model::find(&models, model_id).ok_or_else(|| anyhow!("No embedding model '{model_id}'"))?; + let model = Model::find(&list_embedding_models(config), model_id) + .ok_or_else(|| anyhow!("No embedding model '{model_id}'"))?; Ok(model) } diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 291bcd2..bf5fd8e 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -200,7 +200,7 @@ impl Repl { ".model" => match args { Some(name) => { self.config.write().set_model(name)?; - if self.config.read().state().is_empty() { + if !self.config.read().has_role_or_session() { self.config.write().set_model_id(); } } |
