summaryrefslogtreecommitdiffstats
path: root/src/config.rs
blob: 3a850669970edf70933fcb2cea70ce9e7d305cf3 (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
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
use std::{
    env,
    fs::{create_dir_all, read_to_string, File, OpenOptions},
    io::Write,
    path::{Path, PathBuf},
    process::exit,
};

use anyhow::{anyhow, Result};
use inquire::{Confirm, Text};
use serde::Deserialize;

use crate::utils::now;

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";

#[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,
    /// Whether to disable highlight
    #[serde(default)]
    pub no_highlight: 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(is_interactive: bool) -> Result<Config> {
        let config_path = Config::config_file()?;
        if is_interactive && !config_path.exists() {
            create_config_file(&config_path)?;
        }
        let content = read_to_string(&config_path)
            .map_err(|err| anyhow!("Failed to load config at {}, {err}", config_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 find_role(&self, name: &str) -> Option<Role> {
        self.roles.iter().find(|v| v.name == name).cloned()
    }

    pub fn config_dir() -> 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() {
            create_dir_all(&path).map_err(|err| {
                anyhow!("Failed to create config dir at {}, {err}", path.display())
            })?;
        }
        Ok(path)
    }

    pub fn local_file(name: &str) -> Result<PathBuf> {
        let mut path = Self::config_dir()?;
        path.push(name);
        Ok(path)
    }

    pub fn open_message_file(&self) -> Result<Option<File>> {
        if !self.save {
            return Ok(None);
        }
        let path = Config::messages_file()?;
        let file: Option<File> = if self.save {
            let file = OpenOptions::new()
                .create(true)
                .append(true)
                .open(&path)
                .map_err(|err| anyhow!("Failed to create/append {}, {err}", path.display()))?;
            Some(file)
        } else {
            None
        };
        Ok(file)
    }

    pub fn save_message(
        file: Option<&mut File>,
        input: &str,
        output: &str,
        role_name: &Option<String>,
    ) {
        let role_name = match role_name {
            Some(v) => format!("({v})"),
            None => String::new(),
        };
        let timestamp = format!("[{}]", now());
        if let (false, Some(file)) = (output.is_empty(), file) {
            let _ = file.write_all(
                format!(
                    "# CHAT:{timestamp} {role_name}\n{}\n\n--------\n{}\n--------\n\n",
                    input.trim(),
                    output.trim(),
                )
                .as_bytes(),
            );
        }
    }

    pub fn config_file() -> Result<PathBuf> {
        Self::local_file(CONFIG_FILE_NAME)
    }

    pub fn roles_file() -> Result<PathBuf> {
        Self::local_file(ROLES_FILE_NAME)
    }

    pub fn history_file() -> Result<PathBuf> {
        Self::local_file(HISTORY_FILE_NAME)
    }

    pub fn messages_file() -> Result<PathBuf> {
        Self::local_file(MESSAGE_FILE_NAME)
    }

    fn load_roles(&mut self) -> Result<()> {
        let path = Self::roles_file()?;
        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,
}

fn create_config_file(config_path: &Path) -> Result<()> {
    let confirm_map_err = |_| anyhow!("Error with questionnaire, try again later");
    let text_map_err = |_| anyhow!("An error happened when asking for your key, try again later.");
    let ans = Confirm::new("No config file, create a new one?")
        .with_default(true)
        .prompt()
        .map_err(confirm_map_err)?;
    if !ans {
        exit(0);
    }
    let api_key = Text::new("Openai API Key:")
        .prompt()
        .map_err(text_map_err)?;
    let mut raw_config = format!("api_key: {api_key}\n");

    let ans = Confirm::new("Use proxy?")
        .with_default(false)
        .prompt()
        .map_err(confirm_map_err)?;
    if ans {
        let proxy = Text::new("Set proxy:").prompt().map_err(text_map_err)?;
        raw_config.push_str(&format!("proxy: {proxy}\n"));
    }

    let ans = Confirm::new("Save chat messages")
        .with_default(false)
        .prompt()
        .map_err(confirm_map_err)?;
    if ans {
        raw_config.push_str("save: true\n");
    }

    std::fs::write(config_path, raw_config)
        .map_err(|err| anyhow!("Failed to write to config file, {err}"))?;
    Ok(())
}