diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/config.rs | 50 | ||||
| -rw-r--r-- | src/main.rs | 28 |
2 files changed, 50 insertions, 28 deletions
diff --git a/src/config.rs b/src/config.rs index c99a6ef..3c2c633 100644 --- a/src/config.rs +++ b/src/config.rs @@ -1,26 +1,33 @@ use std::{ - fs::read_to_string, + env, + fs::{self, read_to_string}, path::{Path, PathBuf}, }; use anyhow::{anyhow, Result}; use serde::Deserialize; +pub const CONFIG_FILE_NAME: &str = "config.yaml"; +pub const ROLES_FILE_NAME: &str = "roles.yaml"; +pub const HISTORY_FILE_NAME: &str = "history.txt"; +pub const MESSAGE_FILE_NAME: &str = "messages.md"; + #[derive(Debug, Clone, Deserialize)] pub struct Config { /// Openai api key pub api_key: String, /// What sampling temperature to use, between 0 and 2 pub temperature: Option<f64>, - /// Specify a file path to save chat messages to - pub save_path: Option<PathBuf>, + /// Whether to persistently save chat messages + #[serde(default)] + pub save: bool, /// Set proxy pub proxy: Option<String>, /// Used only for debugging #[serde(default)] pub dry_run: bool, /// Predefined roles - #[serde(default)] + #[serde(default, skip_serializing)] pub roles: Vec<Role>, } @@ -28,10 +35,41 @@ impl Config { pub fn init(path: &Path) -> Result<Config> { let content = read_to_string(path) .map_err(|err| anyhow!("Failed to load config at {}, {err}", path.display()))?; - let config: Config = - toml::from_str(&content).map_err(|err| anyhow!("Invalid config, {err}"))?; + let mut config: Config = + serde_yaml::from_str(&content).map_err(|err| anyhow!("Invalid config, {err}"))?; + config.load_roles()?; Ok(config) } + pub fn local_file(name: &str) -> Result<PathBuf> { + let env_name = format!( + "{}_CONFIG_DIR", + env!("CARGO_CRATE_NAME").to_ascii_uppercase() + ); + let mut path = match env::var(env_name) { + Ok(v) => PathBuf::from(v), + Err(_) => dirs::config_dir().ok_or_else(|| anyhow!("Not found config dir"))?, + }; + path.push(env!("CARGO_CRATE_NAME")); + if !path.exists() { + fs::create_dir_all(&path).map_err(|err| { + anyhow!("Failed to create config dir at {}, {err}", path.display()) + })?; + } + path.push(name); + Ok(path) + } + fn load_roles(&mut self) -> Result<()> { + let path = Self::local_file(ROLES_FILE_NAME)?; + if !path.exists() { + return Ok(()); + } + let content = read_to_string(&path) + .map_err(|err| anyhow!("Failed to load roles at {}, {err}", path.display()))?; + let roles: Vec<Role> = + serde_yaml::from_str(&content).map_err(|err| anyhow!("Invalid roles config, {err}"))?; + self.roles = roles; + Ok(()) + } } #[derive(Debug, Clone, Deserialize)] diff --git a/src/main.rs b/src/main.rs index 18be063..0f19adf 100644 --- a/src/main.rs +++ b/src/main.rs @@ -3,11 +3,10 @@ mod config; use std::fs::{File, OpenOptions}; use std::io::{stdout, Write}; use std::path::Path; -use std::path::PathBuf; use std::process::exit; use std::time::Duration; -use config::{Config, Role}; +use config::{Config, Role, CONFIG_FILE_NAME, HISTORY_FILE_NAME, MESSAGE_FILE_NAME}; use anyhow::{anyhow, Result}; use clap::{Arg, ArgAction, Command}; @@ -77,7 +76,7 @@ fn start() -> Result<()> { .collect::<Vec<String>>() .join(" ") }); - let config_path = get_config_path()?; + let config_path = Config::local_file(CONFIG_FILE_NAME)?; if !config_path.exists() && text.is_none() { create_config_file(&config_path)?; } @@ -137,7 +136,7 @@ fn run_repl( ]), ); let history = Box::new( - FileBackedHistory::with_file(1000, get_history_path()?) + FileBackedHistory::with_file(1000, Config::local_file(HISTORY_FILE_NAME)?) .map_err(|err| anyhow!("Failed to setup history file, {err}"))?, ); let edit_mode = Box::new(Emacs::new(keybindings)); @@ -150,17 +149,12 @@ fn run_repl( let mut trigged_ctrlc = false; let mut output = String::new(); let mut role: Option<Role> = None; - let mut save_file: Option<File> = if let Some(path) = &config.save_path { + let mut save_file: Option<File> = if config.save { let file = OpenOptions::new() .create(true) .append(true) - .open(path) - .map_err(|err| { - anyhow!( - "Failed to create/append save_file at {}, {err}", - path.display() - ) - })?; + .open(Config::local_file(MESSAGE_FILE_NAME)?) + .map_err(|err| anyhow!("Failed to create/append save_file, {err}"))?; Some(file) } else { None @@ -440,16 +434,6 @@ fn dump<T: ToString>(text: T, newlines: usize) { stdout().flush().unwrap(); } -fn get_config_path() -> Result<PathBuf> { - let config_dir = dirs::home_dir().ok_or_else(|| anyhow!("No home dir"))?; - Ok(config_dir.join(format!(".{}.toml", env!("CARGO_CRATE_NAME")))) -} - -fn get_history_path() -> Result<PathBuf> { - let config_dir = dirs::home_dir().ok_or_else(|| anyhow!("No home dir"))?; - Ok(config_dir.join(format!(".{}_history", env!("CARGO_CRATE_NAME")))) -} - fn print_repl_title() { println!("Welcome to aichat {}", env!("CARGO_PKG_VERSION")); println!("Type \".help\" for more information."); |
