summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/config/input.rs23
-rw-r--r--src/config/mod.rs54
-rw-r--r--src/config/role.rs10
-rw-r--r--src/config/session.rs35
-rw-r--r--src/rag/mod.rs5
-rw-r--r--src/repl/mod.rs2
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();
}
}