diff options
| author | sigoden <sigoden@gmail.com> | 2023-11-01 10:28:54 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-11-01 10:28:54 +0800 |
| commit | 7f2210dbcab017ee30e8f95771e68494b4860c59 (patch) | |
| tree | 7a0c7bdc8f6930e262bf8cfdc8ffbd57aa843e40 /src/config/mod.rs | |
| parent | 64202758ec8693dd18951c554728bb5c6b7e5eeb (diff) | |
| download | aichat-7f2210dbcab017ee30e8f95771e68494b4860c59.tar.gz | |
refactor: clients/* and config.rs (#193)
- add register_clients macro to make it easier to add a new client
- no create_client_config, just add const PROMPTS
- move ModelInfo from clients/ to config/
- model's max_tokens are optional
- improve code quanity on config/mod.rs
- add/use macro config_get_fn
Diffstat (limited to 'src/config/mod.rs')
| -rw-r--r-- | src/config/mod.rs | 117 |
1 files changed, 68 insertions, 49 deletions
diff --git a/src/config/mod.rs b/src/config/mod.rs index 956bea0..18f1731 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -1,19 +1,20 @@ mod message; +mod model_info; mod role; mod session; pub use self::message::Message; +pub use self::model_info::ModelInfo; use self::role::Role; use self::session::{Session, TEMP_SESSION_NAME}; -use crate::client::openai::{OpenAIClient, OpenAIConfig}; use crate::client::{ - create_client_config, list_client_types, list_models, prompt_op_err, ClientConfig, ExtraConfig, - ModelInfo, SendData, + all_models, create_client_config, list_client_types, ClientConfig, ExtraConfig, OpenAIClient, + SendData, }; use crate::config::message::num_tokens_from_messages; use crate::render::RenderOptions; -use crate::utils::{get_env_name, light_theme_from_colorfgbg, now}; +use crate::utils::{get_env_name, light_theme_from_colorfgbg, now, prompt_op_err}; use anyhow::{anyhow, bail, Context, Result}; use inquire::{Confirm, Select, Text}; @@ -49,6 +50,8 @@ const SET_COMPLETIONS: [&str; 7] = [ ".set dry_run false", ]; +const CLIENTS_FIELD: &str = "clients"; + #[derive(Debug, Clone, Deserialize)] #[serde(default)] pub struct Config { @@ -61,7 +64,7 @@ pub struct Config { pub save: bool, /// Whether to disable highlight pub highlight: bool, - /// Used only for debugging + /// Dry-run flag pub dry_run: bool, /// Whether to use a light theme pub light_theme: bool, @@ -105,7 +108,7 @@ impl Default for Config { wrap_code: false, auto_copy: false, keybindings: Default::default(), - clients: vec![ClientConfig::OpenAI(OpenAIConfig::default())], + clients: vec![ClientConfig::default()], roles: vec![], role: None, session: None, @@ -145,11 +148,11 @@ impl Config { config.temperature = config.default_temperature; - config.set_model_info()?; - config.merge_env_vars(); config.load_roles()?; - config.ensure_sessions_dir()?; - config.detect_theme()?; + + config.setup_model_info()?; + config.setup_highlight(); + config.setup_light_theme()?; Ok(config) } @@ -296,8 +299,10 @@ impl Config { vec![message] }; let tokens = num_tokens_from_messages(&messages); - if tokens >= self.model_info.max_tokens { - bail!("Exceed max tokens limit") + if let Some(max_tokens) = self.model_info.max_tokens { + if tokens >= max_tokens { + bail!("Exceed max tokens limit") + } } Ok(messages) @@ -318,7 +323,7 @@ impl Config { } pub fn set_model(&mut self, value: &str) -> Result<()> { - let models = list_models(self); + let models = all_models(self); let mut model_info = None; if value.contains(':') { if let Some(model) = models.iter().find(|v| v.stringify() == value) { @@ -339,14 +344,6 @@ impl Config { } } - pub const fn get_reamind_tokens(&self) -> usize { - let mut tokens = self.model_info.max_tokens; - if let Some(session) = self.session.as_ref() { - tokens = tokens.saturating_sub(session.tokens); - } - tokens - } - pub fn info(&self) -> Result<String> { let path_info = |path: &Path| { let state = if path.exists() { "" } else { " ⚠️" }; @@ -390,12 +387,7 @@ impl Config { completion.extend(SET_COMPLETIONS.map(std::string::ToString::to_string)); completion.extend( - list_models(self) - .iter() - .map(|v| format!(".model {}", v.stringify())), - ); - completion.extend( - list_models(self) + all_models(self) .iter() .map(|v| format!(".model {}", v.stringify())), ); @@ -504,6 +496,14 @@ impl Config { name = Text::new("Session name:").with_default(&name).prompt()?; } let session_path = Self::session_file(&name)?; + let sessions_dir = session_path.parent().ok_or_else(|| { + anyhow!("Unable to save session file to {}", session_path.display()) + })?; + if !sessions_dir.exists() { + create_dir_all(sessions_dir).with_context(|| { + format!("Failed to create session_dir '{}'", sessions_dir.display()) + })?; + } session.save(&session_path)?; } } @@ -556,6 +556,24 @@ impl Config { Ok(RenderOptions::new(theme, wrap, self.wrap_code)) } + pub fn render_prompt_right(&self) -> String { + if let Some(session) = &self.session { + let tokens = session.tokens; + // 10000(%32) + match self.model_info.max_tokens { + Some(max_tokens) => { + let ratio = tokens as f32 / max_tokens as f32; + let percent = ratio * 100.0; + let percent = (percent * 100.0).round() / 100.0; + format!("{tokens}({percent}%)") + } + None => format!("{tokens}"), + } + } else { + String::new() + } + } + pub fn prepare_send_data(&self, content: &str, stream: bool) -> Result<SendData> { let messages = self.build_messages(content)?; Ok(SendData { @@ -585,11 +603,20 @@ impl Config { } fn load_config(config_path: &Path) -> Result<Self> { - let content = read_to_string(config_path) - .with_context(|| format!("Failed to load config at {}", config_path.display()))?; + let ctx = || format!("Failed to load config at {}", config_path.display()); + let content = read_to_string(config_path).with_context(ctx)?; let config: Self = serde_yaml::from_str(&content) - .with_context(|| format!("Invalid config at {}", config_path.display()))?; + .map_err(|err| { + let err_msg = err.to_string(); + if err_msg.starts_with(&format!("{}: ", CLIENTS_FIELD)) { + anyhow!("clients: invalid value") + } else { + anyhow!("{err_msg}") + } + }) + .with_context(ctx)?; + Ok(config) } @@ -606,11 +633,11 @@ impl Config { Ok(()) } - fn set_model_info(&mut self) -> Result<()> { + fn setup_model_info(&mut self) -> Result<()> { let model = match &self.model { Some(v) => v.clone(), None => { - let models = self::list_models(self); + let models = all_models(self); if models.is_empty() { bail!("No available model"); } @@ -622,7 +649,7 @@ impl Config { Ok(()) } - fn merge_env_vars(&mut self) { + fn setup_highlight(&mut self) { if let Ok(value) = env::var("NO_COLOR") { let mut no_color = false; set_bool(&mut no_color, &value); @@ -632,17 +659,7 @@ impl Config { } } - fn ensure_sessions_dir(&self) -> Result<()> { - let sessions_dir = Self::sessions_dir()?; - if !sessions_dir.exists() { - create_dir_all(&sessions_dir).with_context(|| { - format!("Failed to create session_dir '{}'", sessions_dir.display()) - })?; - } - Ok(()) - } - - fn detect_theme(&mut self) -> Result<()> { + fn setup_light_theme(&mut self) -> Result<()> { if self.light_theme { return Ok(()); } @@ -660,7 +677,7 @@ impl Config { fn compat_old_config(&mut self, config_path: &PathBuf) -> Result<()> { let content = read_to_string(config_path)?; let value: serde_json::Value = serde_yaml::from_str(&content)?; - if value.get("clients").is_some() { + if value.get(CLIENTS_FIELD).is_some() { return Ok(()); } @@ -725,16 +742,18 @@ fn create_config_file(config_path: &Path) -> Result<()> { exit(0); } - let client = Select::new("AI Platform:", list_client_types()) + let client = Select::new("Platform:", list_client_types()) .prompt() .map_err(prompt_op_err)?; - let mut raw_config = create_client_config(client)?; + let mut config = serde_json::json!({}); + config["model"] = client.into(); + config[CLIENTS_FIELD] = create_client_config(client)?; - raw_config.push_str(&format!("model: {client}\n")); + let config_data = serde_yaml::to_string(&config).with_context(|| "Failed to create config")?; ensure_parent_exists(config_path)?; - std::fs::write(config_path, raw_config).with_context(|| "Failed to write to config file")?; + std::fs::write(config_path, config_data).with_context(|| "Failed to write to config file")?; #[cfg(unix)] { use std::os::unix::prelude::PermissionsExt; |
