summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--src/client/common.rs2
-rw-r--r--src/client/model_info.rs2
-rw-r--r--src/config/mod.rs72
-rw-r--r--src/config/session.rs6
-rw-r--r--src/main.rs8
-rw-r--r--src/repl/completer.rs94
-rw-r--r--src/repl/highlighter.rs22
-rw-r--r--src/repl/mod.rs90
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!(