summaryrefslogtreecommitdiffstats
path: root/src/config.rs
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/config.rs
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/config.rs')
-rw-r--r--src/config.rs50
1 files changed, 44 insertions, 6 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)]