summaryrefslogtreecommitdiffstats
path: root/src/config.rs
blob: 3c2c633261058936fddd621d9513e1954c2074b4 (plain) (blame)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
use std::{
    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>,
    /// 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, skip_serializing)]
    pub roles: Vec<Role>,
}

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 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)]
pub struct Role {
    /// Role name
    pub name: String,
    /// Prompt text send to ai for setting up a role
    pub prompt: String,
}

impl Role {
    pub fn generate(&self, text: &str) -> String {
        format!("{} {}", self.prompt, text)
    }
}