summaryrefslogtreecommitdiffstats
path: root/src/config/mod.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-10-28 21:39:17 +0800
committerGitHub <noreply@github.com>2023-10-28 21:39:17 +0800
commitbc44026ff8828f532e9eba138e2ba35f7c247271 (patch)
tree98c55ad31a543ff3cd56e297751d32ee9ca960a7 /src/config/mod.rs
parent1575d441724a27854c8f9fe1510a47a3050ae28d (diff)
downloadaichat-bc44026ff8828f532e9eba138e2ba35f7c247271.tar.gz
feat: enhance session/conversation (#162)
* feat: enhance session/conversation * updates * updates * cut version v0.9.0-rc2 * add .session name completion
Diffstat (limited to 'src/config/mod.rs')
-rw-r--r--src/config/mod.rs229
1 files changed, 167 insertions, 62 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 8e5157f..73c9444 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -1,10 +1,10 @@
-mod conversation;
mod message;
mod role;
+mod session;
-use self::conversation::Conversation;
use self::message::Message;
use self::role::Role;
+use self::session::{Session, TEMP_SESSION_NAME};
use crate::client::openai::{OpenAIClient, OpenAIConfig};
use crate::client::{all_clients, create_client_config, list_models, ClientConfig, ModelInfo};
@@ -12,12 +12,12 @@ use crate::config::message::num_tokens_from_messages;
use crate::utils::{get_env_name, now};
use anyhow::{anyhow, bail, Context, Result};
-use inquire::{Confirm, Select};
+use inquire::{Confirm, Select, Text};
use parking_lot::RwLock;
use serde::Deserialize;
use std::{
env,
- fs::{create_dir_all, read_to_string, File, OpenOptions},
+ fs::{create_dir_all, read_dir, read_to_string, remove_file, File, OpenOptions},
io::Write,
path::{Path, PathBuf},
process::exit,
@@ -27,7 +27,9 @@ use std::{
const CONFIG_FILE_NAME: &str = "config.yaml";
const ROLES_FILE_NAME: &str = "roles.yaml";
const HISTORY_FILE_NAME: &str = "history.txt";
-const MESSAGE_FILE_NAME: &str = "messages.md";
+const MESSAGES_FILE_NAME: &str = "messages.md";
+const SESSIONS_DIR_NAME: &str = "sessions";
+
const SET_COMPLETIONS: [&str; 7] = [
".set temperature",
".set save true",
@@ -46,14 +48,12 @@ pub struct Config {
pub model: Option<String>,
/// What sampling temperature to use, between 0 and 2
pub temperature: Option<f64>,
- /// Whether to persistently save chat messages
+ /// Whether to persistently save non-session chat messages
pub save: bool,
/// Whether to disable highlight
pub highlight: bool,
/// Used only for debugging
pub dry_run: bool,
- /// If set ture, start a conversation immediately upon repl
- pub conversation_first: bool,
/// If set true, use light theme
pub light_theme: bool,
/// Automatically copy the last output to the clipboard
@@ -68,11 +68,13 @@ pub struct Config {
/// Current selected role
#[serde(skip)]
pub role: Option<Role>,
- /// Current conversation
+ /// Current session
#[serde(skip)]
- pub conversation: Option<Conversation>,
+ pub session: Option<Session>,
#[serde(skip)]
pub model_info: ModelInfo,
+ #[serde(skip)]
+ pub last_message: Option<(String, String)>,
}
impl Default for Config {
@@ -83,15 +85,15 @@ impl Default for Config {
save: false,
highlight: true,
dry_run: false,
- conversation_first: false,
light_theme: false,
auto_copy: false,
keybindings: Default::default(),
- roles: vec![],
clients: vec![ClientConfig::OpenAI(OpenAIConfig::default())],
+ roles: vec![],
role: None,
- conversation: None,
+ session: None,
model_info: Default::default(),
+ last_message: None,
}
}
}
@@ -123,19 +125,14 @@ impl Config {
if let Some(name) = config.model.clone() {
config.set_model(&name)?;
}
+
config.merge_env_vars();
config.load_roles()?;
+ config.ensure_sessions_dir()?;
Ok(config)
}
- pub fn on_repl(&mut self) -> Result<()> {
- if self.conversation_first {
- self.start_conversation()?;
- }
- Ok(())
- }
-
pub fn get_role(&self, name: &str) -> Option<Role> {
self.roles.iter().find(|v| v.match_name(name)).map(|v| {
let mut role = v.clone();
@@ -156,13 +153,20 @@ impl Config {
Ok(path)
}
- pub fn local_file(name: &str) -> Result<PathBuf> {
+ pub fn local_path(name: &str) -> Result<PathBuf> {
let mut path = Self::config_dir()?;
path.push(name);
Ok(path)
}
- pub fn save_message(&self, input: &str, output: &str) -> Result<()> {
+ pub fn save_message(&mut self, input: &str, output: &str) -> Result<()> {
+ self.last_message = Some((input.to_string(), output.to_string()));
+
+ if let Some(session) = self.session.as_mut() {
+ session.add_message(input, output)?;
+ return Ok(());
+ }
+
if !self.save {
return Ok(());
}
@@ -193,30 +197,40 @@ impl Config {
}
pub fn config_file() -> Result<PathBuf> {
- Self::local_file(CONFIG_FILE_NAME)
+ Self::local_path(CONFIG_FILE_NAME)
}
pub fn roles_file() -> Result<PathBuf> {
let env_name = get_env_name("roles_file");
env::var(env_name).map_or_else(
- |_| Self::local_file(ROLES_FILE_NAME),
+ |_| Self::local_path(ROLES_FILE_NAME),
|value| Ok(PathBuf::from(value)),
)
}
pub fn history_file() -> Result<PathBuf> {
- Self::local_file(HISTORY_FILE_NAME)
+ Self::local_path(HISTORY_FILE_NAME)
}
pub fn messages_file() -> Result<PathBuf> {
- Self::local_file(MESSAGE_FILE_NAME)
+ Self::local_path(MESSAGES_FILE_NAME)
+ }
+
+ pub fn sessions_dir() -> Result<PathBuf> {
+ Self::local_path(SESSIONS_DIR_NAME)
+ }
+
+ pub fn session_file(name: &str) -> Result<PathBuf> {
+ let mut path = Self::sessions_dir()?;
+ path.push(&format!("{name}.yaml"));
+ Ok(path)
}
pub fn change_role(&mut self, name: &str) -> Result<String> {
match self.get_role(name) {
Some(role) => {
- if let Some(conversation) = self.conversation.as_mut() {
- conversation.update_role(&role)?;
+ if let Some(session) = self.session.as_mut() {
+ session.update_role(Some(role.clone()))?;
}
let output = serde_yaml::to_string(&role)
.unwrap_or_else(|_| "Unable to echo role details".into());
@@ -228,8 +242,8 @@ impl Config {
}
pub fn clear_role(&mut self) -> Result<()> {
- if let Some(conversation) = self.conversation.as_ref() {
- conversation.can_clear_role()?;
+ if let Some(session) = self.session.as_mut() {
+ session.update_role(None)?;
}
self.role = None;
Ok(())
@@ -237,8 +251,8 @@ impl Config {
pub fn add_prompt(&mut self, prompt: &str) -> Result<()> {
let role = Role::new(prompt, self.temperature);
- if let Some(conversation) = self.conversation.as_mut() {
- conversation.update_role(&role)?;
+ if let Some(session) = self.session.as_mut() {
+ session.update_role(Some(role.clone()))?;
}
self.role = Some(role);
Ok(())
@@ -253,8 +267,8 @@ impl Config {
pub fn echo_messages(&self, content: &str) -> String {
#[allow(clippy::option_if_let_else)]
- if let Some(conversation) = self.conversation.as_ref() {
- conversation.echo_messages(content)
+ if let Some(session) = self.session.as_ref() {
+ session.echo_messages(content)
} else if let Some(role) = self.role.as_ref() {
role.echo_messages(content)
} else {
@@ -264,8 +278,8 @@ impl Config {
pub fn build_messages(&self, content: &str) -> Result<Vec<Message>> {
#[allow(clippy::option_if_let_else)]
- let messages = if let Some(conversation) = self.conversation.as_ref() {
- conversation.build_emssages(content)
+ let messages = if let Some(session) = self.session.as_ref() {
+ session.build_emssages(content)
} else if let Some(role) = self.role.as_ref() {
role.build_messages(content)
} else {
@@ -282,28 +296,36 @@ impl Config {
pub fn set_model(&mut self, value: &str) -> Result<()> {
let models = list_models(self);
+ let mut model_info = None;
if value.contains(':') {
if let Some(model) = models.iter().find(|v| v.stringify() == value) {
- self.model_info = model.clone();
- return Ok(());
+ model_info = Some(model.clone());
}
} else if let Some(model) = models.iter().find(|v| v.client == value) {
- self.model_info = model.clone();
- return Ok(());
+ model_info = Some(model.clone());
+ }
+ match model_info {
+ None => bail!("Invalid model"),
+ Some(model_info) => {
+ if let Some(session) = self.session.as_mut() {
+ session.model = model_info.stringify();
+ }
+ self.model_info = model_info;
+ Ok(())
+ }
}
- bail!("Invalid model")
}
pub const fn get_reamind_tokens(&self) -> usize {
let mut tokens = self.model_info.max_tokens;
- if let Some(conversation) = self.conversation.as_ref() {
- tokens = tokens.saturating_sub(conversation.tokens);
+ if let Some(session) = self.session.as_ref() {
+ tokens = tokens.saturating_sub(session.tokens);
}
tokens
}
pub fn info(&self) -> Result<String> {
- let file_info = |path: &Path| {
+ let path_info = |path: &Path| {
let state = if path.exists() { "" } else { " ⚠️" };
format!("{}{state}", path.display())
};
@@ -311,14 +333,14 @@ impl Config {
.temperature
.map_or_else(|| String::from("-"), |v| v.to_string());
let items = vec![
- ("config_file", file_info(&Self::config_file()?)),
- ("roles_file", file_info(&Self::roles_file()?)),
- ("messages_file", file_info(&Self::messages_file()?)),
+ ("config_file", path_info(&Self::config_file()?)),
+ ("roles_file", path_info(&Self::roles_file()?)),
+ ("messages_file", path_info(&Self::messages_file()?)),
+ ("sessions_dir", path_info(&Self::sessions_dir()?)),
("model", self.model_info.stringify()),
("temperature", temperature),
("save", self.save.to_string()),
("highlight", self.highlight.to_string()),
- ("conversation_first", self.conversation_first.to_string()),
("light_theme", self.light_theme.to_string()),
("dry_run", self.dry_run.to_string()),
("keybindings", self.keybindings.stringify().into()),
@@ -343,6 +365,13 @@ impl Config {
.iter()
.map(|v| format!(".model {}", v.stringify())),
);
+ completion.extend(
+ list_models(self)
+ .iter()
+ .map(|v| format!(".model {}", v.stringify())),
+ );
+ let sessions = self.list_sessions().unwrap_or_default();
+ completion.extend(sessions.iter().map(|v| format!(".session {}", v)));
completion
}
@@ -380,28 +409,94 @@ impl Config {
Ok(())
}
- pub fn start_conversation(&mut self) -> Result<()> {
- if self.conversation.is_some() && self.get_reamind_tokens() > 0 {
- let ans = Confirm::new("Already in a conversation, start a new one?")
- .with_default(true)
- .prompt()?;
- if !ans {
- return Ok(());
+ pub fn start_session(&mut self, session: &Option<String>) -> Result<()> {
+ if self.session.is_some() {
+ bail!("Already in a session, please use '.clear session' to exit the session first?");
+ }
+ match session {
+ None => {
+ let session_file = Self::session_file(TEMP_SESSION_NAME)?;
+ if session_file.exists() {
+ remove_file(session_file)
+ .with_context(|| "Failed to clean previous session")?;
+ }
+ self.session = Some(Session::new(
+ TEMP_SESSION_NAME,
+ &self.model_info.stringify(),
+ self.role.clone(),
+ ));
+ }
+ Some(name) => {
+ let session_path = Self::session_file(name)?;
+ if !session_path.exists() {
+ self.session = Some(Session::new(
+ name,
+ &self.model_info.stringify(),
+ 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();
+ self.session = Some(session);
+ }
+ }
+ }
+ if let Some(session) = self.session.as_mut() {
+ if session.is_empty() {
+ if let Some((input, output)) = &self.last_message {
+ let ans = Confirm::new(
+ "Start a session that incorporates the last question and answer?",
+ )
+ .with_default(false)
+ .prompt()?;
+ if ans {
+ session.add_message(input, output)?;
+ }
+ }
}
}
- self.conversation = Some(Conversation::new(self.role.clone()));
Ok(())
}
- pub fn end_conversation(&mut self) {
- self.conversation = None;
+ pub fn end_session(&mut self) -> Result<()> {
+ if let Some(mut session) = self.session.take() {
+ self.last_message = None;
+ if session.should_save() {
+ let ans = Confirm::new("Save session?").with_default(true).prompt()?;
+ if !ans {
+ return Ok(());
+ }
+ let mut name = session.name.clone();
+ if session.is_temp() {
+ name = Text::new("Session name:").with_default(&name).prompt()?;
+ }
+ let session_path = Self::session_file(&name)?;
+ session.save(&session_path)?;
+ }
+ }
+ Ok(())
}
- pub fn save_conversation(&mut self, input: &str, output: &str) -> Result<()> {
- if let Some(conversation) = self.conversation.as_mut() {
- conversation.add_message(input, output)?;
+ pub fn list_sessions(&self) -> Result<Vec<String>> {
+ let sessions_dir = Self::sessions_dir()?;
+ match read_dir(&sessions_dir) {
+ Ok(rd) => {
+ let mut names = vec![];
+ for entry in rd {
+ let entry = entry?;
+ let name = entry.file_name();
+ if let Some(name) = name.to_string_lossy().strip_suffix(".yaml") {
+ names.push(name.to_string());
+ }
+ }
+ Ok(names)
+ }
+ Err(_) => Ok(vec![]),
}
- Ok(())
}
pub const fn get_render_options(&self) -> (bool, bool) {
@@ -463,6 +558,16 @@ impl Config {
}
}
+ fn ensure_sessions_dir(&self) -> Result<()> {
+ let sessions_dir = Self::sessions_dir()?;
+ if !sessions_dir.exists() {
+ create_dir_all(&sessions_dir).with_context(|| {
+ format!("Failed to create session_dir '{}'", sessions_dir.display())
+ })?;
+ }
+ Ok(())
+ }
+
fn compat_old_config(&mut self, config_path: &PathBuf) -> Result<()> {
let content = read_to_string(config_path)?;
let value: serde_json::Value = serde_yaml::from_str(&content)?;