summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-03-03 09:38:53 +0800
committersigoden <sigoden@gmail.com>2023-03-03 10:11:23 +0800
commit627a2d08ce0932241a21a281712e58427c15d478 (patch)
treebf7059a6c9c8e0df939888ba37dc4759640f30b7 /src
parent9e8a5481cf9aa4dfd75d5c55c003b3813ccbca1e (diff)
downloadaichat-627a2d08ce0932241a21a281712e58427c15d478.tar.gz
refactor: config path and config data
use yaml other than toml for configuration. use $AICHAT_CONFIG_DIR other than $HOME for config dir. split roles to seperate config file
Diffstat (limited to 'src')
-rw-r--r--src/config.rs50
-rw-r--r--src/main.rs28
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.");