diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/common.rs | 2 | ||||
| -rw-r--r-- | src/client/model_info.rs | 2 | ||||
| -rw-r--r-- | src/config/mod.rs | 72 | ||||
| -rw-r--r-- | src/config/session.rs | 6 | ||||
| -rw-r--r-- | src/main.rs | 8 | ||||
| -rw-r--r-- | src/repl/completer.rs | 94 | ||||
| -rw-r--r-- | src/repl/highlighter.rs | 22 | ||||
| -rw-r--r-- | src/repl/mod.rs | 90 |
8 files changed, 191 insertions, 105 deletions
diff --git a/src/client/common.rs b/src/client/common.rs index a7844f3..0dc637b 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -101,7 +101,7 @@ macro_rules! register_client { anyhow::bail!("Unknown client {}", client) } - pub fn all_models(config: &$crate::config::Config) -> Vec<$crate::client::ModelInfo> { + pub fn list_models(config: &$crate::config::Config) -> Vec<$crate::client::ModelInfo> { config .clients .iter() diff --git a/src/client/model_info.rs b/src/client/model_info.rs index 7a52e63..9e74951 100644 --- a/src/client/model_info.rs +++ b/src/client/model_info.rs @@ -45,7 +45,7 @@ impl ModelInfo { self } - pub fn full_name(&self) -> String { + pub fn id(&self) -> String { format!("{}:{}", self.client, self.name) } diff --git a/src/config/mod.rs b/src/config/mod.rs index da8e9b7..da197b9 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -5,7 +5,7 @@ use self::role::Role; use self::session::{Session, TEMP_SESSION_NAME}; use crate::client::{ - all_models, create_client_config, list_client_types, ClientConfig, ExtraConfig, Message, + create_client_config, list_client_types, list_models, ClientConfig, ExtraConfig, Message, ModelInfo, OpenAIClient, SendData, }; use crate::render::{MarkdownRender, RenderOptions}; @@ -35,16 +35,6 @@ const ROLES_FILE_NAME: &str = "roles.yaml"; const MESSAGES_FILE_NAME: &str = "messages.md"; const SESSIONS_DIR_NAME: &str = "sessions"; -const SET_COMPLETIONS: [&str; 7] = [ - ".set temperature", - ".set save true", - ".set save false", - ".set highlight true", - ".set highlight false", - ".set dry_run true", - ".set dry_run false", -]; - const CLIENTS_FIELD: &str = "clients"; #[derive(Debug, Clone, Deserialize)] @@ -311,10 +301,11 @@ impl Config { } pub fn set_model(&mut self, value: &str) -> Result<()> { - let models = all_models(self); + let models = list_models(self); let mut model_info = None; + let value = value.trim_end_matches(':'); if value.contains(':') { - if let Some(model) = models.iter().find(|v| v.full_name() == value) { + if let Some(model) = models.iter().find(|v| v.id() == value) { model_info = Some(model.clone()); } } else if let Some(model) = models.iter().find(|v| v.client == value) { @@ -345,7 +336,7 @@ impl Config { .clone() .map_or_else(|| String::from("no"), |v| v.to_string()); let items = vec![ - ("model", self.model_info.full_name()), + ("model", self.model_info.id()), ("temperature", temperature), ("dry_run", self.dry_run.to_string()), ("save", self.save.to_string()), @@ -402,22 +393,27 @@ impl Config { .unwrap_or_default() } - pub fn repl_completions(&self) -> Vec<String> { - let mut completion: Vec<String> = self - .roles - .iter() - .map(|v| format!(".role {}", v.name)) + pub fn repl_complete(&self, cmd: &str, args: &str) -> Vec<String> { + let possible_values = match cmd { + ".role" => self.roles.iter().map(|v| v.name.clone()).collect(), + ".model" => list_models(self).into_iter().map(|v| v.id()).collect(), + ".session" => self.list_sessions(), + ".set" => { + vec![ + "temperature ".into(), + format!("save {}", !self.save), + format!("highlight {}", !self.highlight), + format!("dry_run {}", !self.dry_run), + ] + } + _ => vec![], + }; + let mut possible_values: Vec<String> = possible_values + .into_iter() + .filter(|v| v.starts_with(args)) .collect(); - - completion.extend(SET_COMPLETIONS.map(std::string::ToString::to_string)); - completion.extend( - all_models(self) - .iter() - .map(|v| format!(".model {}", v.full_name())), - ); - let sessions = self.list_sessions().unwrap_or_default(); - completion.extend(sessions.iter().map(|v| format!(".session {}", v))); - completion + possible_values.sort_unstable(); + possible_values } pub fn update(&mut self, data: &str) -> Result<()> { @@ -541,21 +537,23 @@ impl Config { Ok(()) } - pub fn list_sessions(&self) -> Result<Vec<String>> { - let sessions_dir = Self::sessions_dir()?; + pub fn list_sessions(&self) -> Vec<String> { + let sessions_dir = match Self::sessions_dir() { + Ok(dir) => dir, + Err(_) => return vec![], + }; match read_dir(&sessions_dir) { Ok(rd) => { let mut names = vec![]; - for entry in rd { - let entry = entry?; + for entry in rd.flatten() { let name = entry.file_name(); if let Some(name) = name.to_string_lossy().strip_suffix(".yaml") { names.push(name.to_string()); } } - Ok(names) + names } - Err(_) => Ok(vec![]), + Err(_) => vec![], } } @@ -665,12 +663,12 @@ impl Config { let model = match &self.model { Some(v) => v.clone(), None => { - let models = all_models(self); + let models = list_models(self); if models.is_empty() { bail!("No available model"); } - models[0].full_name() + models[0].id() } }; self.set_model(&model)?; diff --git a/src/config/session.rs b/src/config/session.rs index 446b682..92e8c2a 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -33,7 +33,7 @@ impl Session { pub fn new(name: &str, model_info: ModelInfo, role: Option<Role>) -> Self { let temperature = role.as_ref().and_then(|v| v.temperature); Self { - model: model_info.full_name(), + model: model_info.id(), temperature, messages: vec![], name: name.to_string(), @@ -103,7 +103,7 @@ impl Session { items.push(("path", path.to_string())); } - items.push(("model", self.model_info.full_name())); + items.push(("model", self.model_info.id())); if let Some(temperature) = self.temperature() { items.push(("temperature", temperature.to_string())); @@ -165,7 +165,7 @@ impl Session { } pub fn set_model(&mut self, model_info: ModelInfo) -> Result<()> { - self.model = model_info.full_name(); + self.model = model_info.id(); self.model_info = model_info; Ok(()) } diff --git a/src/main.rs b/src/main.rs index 2bc3d63..eae9a4a 100644 --- a/src/main.rs +++ b/src/main.rs @@ -12,7 +12,7 @@ use crate::config::{Config, GlobalConfig}; use anyhow::Result; use clap::Parser; -use client::{all_models, init_client}; +use client::{init_client, list_models}; use crossbeam::sync::WaitGroup; use is_terminal::IsTerminal; use parking_lot::RwLock; @@ -36,13 +36,13 @@ fn main() -> Result<()> { exit(0); } if cli.list_models { - for model in all_models(&config.read()) { - println!("{}", model.full_name()); + for model in list_models(&config.read()) { + println!("{}", model.id()); } exit(0); } if cli.list_sessions { - let sessions = config.read().list_sessions()?.join("\n"); + let sessions = config.read().list_sessions().join("\n"); println!("{sessions}"); exit(0); } diff --git a/src/repl/completer.rs b/src/repl/completer.rs new file mode 100644 index 0000000..ea24be8 --- /dev/null +++ b/src/repl/completer.rs @@ -0,0 +1,94 @@ +use std::collections::HashMap; + +use super::{parse_command, REPL_COMMANDS}; + +use crate::config::GlobalConfig; + +use reedline::{Completer, Span, Suggestion}; + +impl Completer for ReplCompleter { + fn complete(&mut self, line: &str, pos: usize) -> Vec<Suggestion> { + let mut suggestions = vec![]; + if line.len() != pos { + return suggestions; + } + let line = &line[0..pos]; + if let Some((cmd, args)) = parse_command(line) { + let commands: Vec<_> = self + .commands + .iter() + .filter(|(cmd_name, _)| match args { + Some(args) => cmd_name.starts_with(&format!("{cmd} {args}")), + None => cmd_name.starts_with(cmd), + }) + .collect(); + + if args.is_some() || line.ends_with(' ') { + let args = args.unwrap_or_default(); + let start = line.chars().take_while(|c| *c == ' ').count() + cmd.len() + 1; + let span = Span::new(start, pos); + suggestions.extend( + self.config + .read() + .repl_complete(cmd, args) + .iter() + .map(|name| create_suggestion(name.clone(), None, span)), + ) + } + + if suggestions.is_empty() { + let start = line.chars().take_while(|c| *c == ' ').count(); + let span = Span::new(start, pos); + suggestions.extend(commands.iter().map(|(name, desc)| { + let has_group = self.groups.get(name).map(|v| *v > 1).unwrap_or_default(); + let name = if has_group { + name.to_string() + } else { + format!("{name} ") + }; + create_suggestion(name, Some(desc.to_string()), span) + })) + } + } + suggestions + } +} + +pub struct ReplCompleter { + config: GlobalConfig, + commands: Vec<(&'static str, &'static str)>, + groups: HashMap<&'static str, usize>, +} + +impl ReplCompleter { + pub fn new(config: &GlobalConfig) -> Self { + let mut groups = HashMap::new(); + + let mut commands = REPL_COMMANDS.to_vec(); + commands.sort_by(|(a, _), (b, _)| a.cmp(b)); + + for (name, _) in REPL_COMMANDS.iter() { + if let Some(count) = groups.get(name) { + groups.insert(*name, count + 1); + } else { + groups.insert(*name, 1); + } + } + + Self { + config: config.clone(), + commands, + groups, + } + } +} + +fn create_suggestion(value: String, description: Option<String>, span: Span) -> Suggestion { + Suggestion { + value, + description, + extra: None, + span, + append_whitespace: false, + } +} diff --git a/src/repl/highlighter.rs b/src/repl/highlighter.rs index 37df659..7d3c197 100644 --- a/src/repl/highlighter.rs +++ b/src/repl/highlighter.rs @@ -1,18 +1,18 @@ +use super::REPL_COMMANDS; + use crate::config::GlobalConfig; use nu_ansi_term::{Color, Style}; use reedline::{Highlighter, StyledText}; pub struct ReplHighlighter { - external_commands: Vec<String>, config: GlobalConfig, } impl ReplHighlighter { - pub fn new(external_commands: Vec<String>, config: GlobalConfig) -> Self { + pub fn new(config: &GlobalConfig) -> Self { Self { - external_commands, - config, + config: config.clone(), } } } @@ -28,17 +28,11 @@ impl Highlighter for ReplHighlighter { let mut styled_text = StyledText::new(); - if self - .external_commands - .clone() - .iter() - .any(|x| line.contains(x)) - { - let matches: Vec<&str> = self - .external_commands + if REPL_COMMANDS.iter().any(|(cmd, _)| line.contains(cmd)) { + let matches: Vec<&str> = REPL_COMMANDS .iter() - .filter(|c| line.contains(*c)) - .map(std::ops::Deref::deref) + .filter(|(cmd, _)| line.contains(*cmd)) + .map(|(cmd, _)| *cmd) .collect(); let longest_match = matches.iter().fold(String::new(), |acc, &item| { if item.len() > acc.len() { diff --git a/src/repl/mod.rs b/src/repl/mod.rs index c19058b..6dba38e 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -1,7 +1,9 @@ +mod completer; mod highlighter; mod prompt; mod validator; +use self::completer::ReplCompleter; use self::highlighter::ReplHighlighter; use self::prompt::ReplPrompt; use self::validator::ReplValidator; @@ -19,8 +21,8 @@ use lazy_static::lazy_static; use reedline::Signal; use reedline::{ default_emacs_keybindings, default_vi_insert_keybindings, default_vi_normal_keybindings, - ColumnarMenu, DefaultCompleter, EditMode, Emacs, KeyCode, KeyModifiers, Keybindings, Reedline, - ReedlineEvent, ReedlineMenu, Vi, + ColumnarMenu, EditMode, Emacs, KeyCode, KeyModifiers, Keybindings, Reedline, ReedlineEvent, + ReedlineMenu, Vi, }; use std::cell::RefCell; use std::io::Read; @@ -45,7 +47,7 @@ const REPL_COMMANDS: [(&str, &str); 14] = [ ]; lazy_static! { - static ref COMMAND_RE: Regex = Regex::new(r"^\s*(\.\S+)\s*").unwrap(); + static ref COMMAND_RE: Regex = Regex::new(r"^\s*(\.\S*)\s*").unwrap(); static ref EDIT_RE: Regex = Regex::new(r"^\s*\.edit\s*").unwrap(); } @@ -59,36 +61,7 @@ pub struct Repl { impl Repl { pub fn init(config: GlobalConfig) -> Result<Self> { - let commands: Vec<String> = REPL_COMMANDS - .into_iter() - .map(|(v, _)| v.to_string()) - .collect(); - - let completer = Self::create_completer(&config, &commands); - let highlighter = ReplHighlighter::new(commands, config.clone()); - let menu = Self::create_menu(); - let edit_mode: Box<dyn EditMode> = if config.read().keybindings.is_vi() { - let mut normal_keybindings = default_vi_normal_keybindings(); - let mut insert_keybindings = default_vi_insert_keybindings(); - Self::extra_keybindings(&mut normal_keybindings); - Self::extra_keybindings(&mut insert_keybindings); - Box::new(Vi::new(insert_keybindings, normal_keybindings)) - } else { - let mut keybindings = default_emacs_keybindings(); - Self::extra_keybindings(&mut keybindings); - Box::new(Emacs::new(keybindings)) - }; - let mut editor = Reedline::create() - .with_completer(Box::new(completer)) - .with_highlighter(Box::new(highlighter)) - .with_menu(menu) - .with_edit_mode(edit_mode) - .with_quick_completions(true) - .with_partial_completions(true) - .with_validator(Box::new(ReplValidator)) - .with_ansi_colors(true); - - editor.enable_bracketed_paste()?; + let editor = Self::create_editor(&config)?; let prompt = ReplPrompt::new(config.clone()); @@ -281,13 +254,24 @@ Type ".help" for more information. ) } - fn create_completer(config: &GlobalConfig, commands: &[String]) -> DefaultCompleter { - let mut completion = commands.to_vec(); - completion.extend(config.read().repl_completions()); - let mut completer = - DefaultCompleter::with_inclusions(&['.', '-', '_', ':']).set_min_word_len(2); - completer.insert(completion.clone()); - completer + fn create_editor(config: &GlobalConfig) -> Result<Reedline> { + let completer = ReplCompleter::new(config); + let highlighter = ReplHighlighter::new(config); + let menu = Self::create_menu(); + let edit_mode = Self::create_edit_mode(config); + let mut editor = Reedline::create() + .with_completer(Box::new(completer)) + .with_highlighter(Box::new(highlighter)) + .with_menu(menu) + .with_edit_mode(edit_mode) + .with_quick_completions(true) + .with_partial_completions(true) + .with_validator(Box::new(ReplValidator)) + .with_ansi_colors(true); + + editor.enable_bracketed_paste()?; + + Ok(editor) } fn extra_keybindings(keybindings: &mut Keybindings) { @@ -306,6 +290,21 @@ Type ".help" for more information. ); } + fn create_edit_mode(config: &GlobalConfig) -> Box<dyn EditMode> { + let edit_mode: Box<dyn EditMode> = if config.read().keybindings.is_vi() { + let mut normal_keybindings = default_vi_normal_keybindings(); + let mut insert_keybindings = default_vi_insert_keybindings(); + Self::extra_keybindings(&mut normal_keybindings); + Self::extra_keybindings(&mut insert_keybindings); + Box::new(Vi::new(insert_keybindings, normal_keybindings)) + } else { + let mut keybindings = default_emacs_keybindings(); + Self::extra_keybindings(&mut keybindings); + Box::new(Emacs::new(keybindings)) + }; + edit_mode + } + fn create_menu() -> ReedlineMenu { let completion_menu = ColumnarMenu::default().with_name(MENU_NAME); ReedlineMenu::EngineCompleter(Box::new(completion_menu)) @@ -343,15 +342,15 @@ Press Ctrl+C to abort readline, Ctrl+D to exit the REPL"###, } fn parse_command(line: &str) -> Option<(&str, Option<&str>)> { - if let Ok(Some(captures)) = COMMAND_RE.captures(line) { - if let Some(cmd) = captures.get(1) { - let cmd = cmd.as_str(); + match COMMAND_RE.captures(line) { + Ok(Some(captures)) => { + let cmd = captures.get(1)?.as_str(); let args = line[captures[0].len()..].trim(); let args = if args.is_empty() { None } else { Some(args) }; - return Some((cmd, args)); + Some((cmd, args)) } + _ => None, } - None } #[cfg(test)] mod tests { @@ -359,6 +358,7 @@ mod tests { #[test] fn test_process_command_line() { + assert_eq!(parse_command(" ."), Some((".", None))); assert_eq!(parse_command(" .role"), Some((".role", None))); assert_eq!(parse_command(" .role "), Some((".role", None))); assert_eq!( |
