diff options
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/message.rs | 9 | ||||
| -rw-r--r-- | src/config/mod.rs | 117 | ||||
| -rw-r--r-- | src/config/model_info.rs | 27 |
3 files changed, 103 insertions, 50 deletions
diff --git a/src/config/message.rs b/src/config/message.rs index d5fcff1..5882337 100644 --- a/src/config/message.rs +++ b/src/config/message.rs @@ -17,7 +17,7 @@ impl Message { } } -#[derive(Debug, Clone, Deserialize, Serialize)] +#[derive(Debug, Clone, Copy, Deserialize, Serialize)] #[serde(rename_all = "snake_case")] pub enum MessageRole { System, @@ -25,6 +25,13 @@ pub enum MessageRole { User, } +impl MessageRole { + #[allow(dead_code)] + pub fn is_system(&self) -> bool { + matches!(self, MessageRole::System) + } +} + pub fn num_tokens_from_messages(messages: &[Message]) -> usize { let mut num_tokens = 0; for message in messages.iter() { 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; diff --git a/src/config/model_info.rs b/src/config/model_info.rs new file mode 100644 index 0000000..1f8d6f0 --- /dev/null +++ b/src/config/model_info.rs @@ -0,0 +1,27 @@ +#[derive(Debug, Clone)] +pub struct ModelInfo { + pub client: String, + pub name: String, + pub max_tokens: Option<usize>, + pub index: usize, +} + +impl Default for ModelInfo { + fn default() -> Self { + ModelInfo::new("", "", None, 0) + } +} + +impl ModelInfo { + pub fn new(client: &str, name: &str, max_tokens: Option<usize>, index: usize) -> Self { + Self { + client: client.into(), + name: name.into(), + max_tokens, + index, + } + } + pub fn stringify(&self) -> String { + format!("{}:{}", self.client, self.name) + } +} |
