summaryrefslogtreecommitdiffstats
path: root/src/repl
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-11-27 15:39:55 +0800
committerGitHub <noreply@github.com>2023-11-27 15:39:55 +0800
commit2508d56598a37844e369ab623ceb5bc7b2c78d38 (patch)
tree744986163dcfe09aade1ef61ef63ad21279666cc /src/repl
parent25e545474fcbd69dfb2a283981178491717cc6b5 (diff)
downloadaichat-2508d56598a37844e369ab623ceb5bc7b2c78d38.tar.gz
feat: state-aware completer (#251)
Diffstat (limited to 'src/repl')
-rw-r--r--src/repl/completer.rs30
-rw-r--r--src/repl/highlighter.rs6
-rw-r--r--src/repl/mod.rs84
3 files changed, 88 insertions, 32 deletions
diff --git a/src/repl/completer.rs b/src/repl/completer.rs
index f9006e2..02c34a6 100644
--- a/src/repl/completer.rs
+++ b/src/repl/completer.rs
@@ -1,4 +1,4 @@
-use super::REPL_COMMANDS;
+use super::{ReplCommand, REPL_COMMANDS};
use crate::config::GlobalConfig;
@@ -27,17 +27,22 @@ impl Completer for ReplCompleter {
return suggestions;
}
+ let state = self.config.read().get_state();
+
let commands: Vec<_> = self
.commands
.iter()
- .filter(|(cmd_name, _)| {
+ .filter(|cmd| {
+ if cmd.unavailable(&state) {
+ return false;
+ }
let line = parts
.iter()
.take(2)
.map(|(v, _)| *v)
.collect::<Vec<&str>>()
.join(" ");
- cmd_name.starts_with(&line)
+ cmd.name.starts_with(&line)
})
.collect();
@@ -55,14 +60,16 @@ impl Completer for ReplCompleter {
if suggestions.is_empty() {
let span = Span::new(cmd_start, pos);
- suggestions.extend(commands.iter().map(|(name, desc)| {
+ suggestions.extend(commands.iter().map(|cmd| {
+ let name = cmd.name;
+ let description = cmd.description;
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)
+ create_suggestion(name, Some(description.to_string()), span)
}))
}
suggestions
@@ -71,7 +78,7 @@ impl Completer for ReplCompleter {
pub struct ReplCompleter {
config: GlobalConfig,
- commands: Vec<(&'static str, &'static str)>,
+ commands: Vec<ReplCommand>,
groups: HashMap<&'static str, usize>,
}
@@ -79,14 +86,15 @@ 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));
+ let mut commands: Vec<ReplCommand> = REPL_COMMANDS.to_vec();
+ commands.sort_by(|a, b| a.name.cmp(b.name));
- for (name, _) in REPL_COMMANDS.iter() {
+ for cmd in REPL_COMMANDS.iter() {
+ let name = cmd.name;
if let Some(count) = groups.get(name) {
- groups.insert(*name, count + 1);
+ groups.insert(name, count + 1);
} else {
- groups.insert(*name, 1);
+ groups.insert(name, 1);
}
}
diff --git a/src/repl/highlighter.rs b/src/repl/highlighter.rs
index 7d3c197..65cb1cc 100644
--- a/src/repl/highlighter.rs
+++ b/src/repl/highlighter.rs
@@ -28,11 +28,11 @@ impl Highlighter for ReplHighlighter {
let mut styled_text = StyledText::new();
- if REPL_COMMANDS.iter().any(|(cmd, _)| line.contains(cmd)) {
+ if REPL_COMMANDS.iter().any(|cmd| line.contains(cmd.name)) {
let matches: Vec<&str> = REPL_COMMANDS
.iter()
- .filter(|(cmd, _)| line.contains(*cmd))
- .map(|(cmd, _)| *cmd)
+ .filter(|cmd| line.contains(cmd.name))
+ .map(|cmd| cmd.name)
.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 7cdb4a7..4ba2266 100644
--- a/src/repl/mod.rs
+++ b/src/repl/mod.rs
@@ -7,7 +7,7 @@ use self::highlighter::ReplHighlighter;
use self::prompt::ReplPrompt;
use crate::client::init_client;
-use crate::config::{GlobalConfig, Input};
+use crate::config::{GlobalConfig, Input, State};
use crate::render::{render_error, render_stream};
use crate::utils::{create_abort_signal, set_text, AbortSignal};
@@ -23,23 +23,50 @@ use reedline::{
const MENU_NAME: &str = "completion_menu";
-const REPL_COMMANDS: [(&str, &str); 13] = [
- (".help", "Print this help message"),
- (".info", "Print system info"),
- (".model", "Switch LLM model"),
- (".role", "Use a role"),
- (".info role", "Show role info"),
- (".exit role", "Leave current role"),
- (".session", "Start a context-aware chat session"),
- (".info session", "Show session info"),
- (".exit session", "End the current session"),
- (".file", "Attach files to the message and then submit it"),
- (".set", "Modify the configuration parameters"),
- (".copy", "Copy the last reply to the clipboard"),
- (".exit", "Exit the REPL"),
-];
-
lazy_static! {
+ static ref REPL_COMMANDS: [ReplCommand; 13] = [
+ ReplCommand::new(".help", "Print this help message", vec![]),
+ ReplCommand::new(".info", "Print system info", vec![]),
+ ReplCommand::new(".model", "Switch LLM model", vec![]),
+ ReplCommand::new(".role", "Use a role", vec![State::Session]),
+ ReplCommand::new(
+ ".info role",
+ "Show role info",
+ vec![State::Normal, State::EmptySession, State::Session]
+ ),
+ ReplCommand::new(
+ ".exit role",
+ "Leave current role",
+ vec![State::Normal, State::EmptySession, State::Session]
+ ),
+ ReplCommand::new(
+ ".session",
+ "Start a context-aware chat session",
+ vec![
+ State::EmptySession,
+ State::EmptySessionWithRole,
+ State::Session
+ ]
+ ),
+ ReplCommand::new(
+ ".info session",
+ "Show session info",
+ vec![State::Normal, State::Role]
+ ),
+ ReplCommand::new(
+ ".exit session",
+ "End the current session",
+ vec![State::Normal, State::Role]
+ ),
+ ReplCommand::new(
+ ".file",
+ "Attach files to the message and then submit it",
+ vec![]
+ ),
+ ReplCommand::new(".set", "Modify the configuration parameters", vec![]),
+ ReplCommand::new(".copy", "Copy the last reply to the clipboard", vec![]),
+ ReplCommand::new(".exit", "Exit the REPL", vec![]),
+ ];
static ref COMMAND_RE: Regex = Regex::new(r"^\s*(\.\S*)\s*").unwrap();
static ref MULTILINE_RE: Regex = Regex::new(r"(?s)^\s*:::\s*(.*)\s*:::\s*$").unwrap();
}
@@ -318,6 +345,27 @@ Type ".help" for more information.
}
}
+#[derive(Debug, Clone)]
+pub struct ReplCommand {
+ name: &'static str,
+ description: &'static str,
+ unavailable_states: Vec<State>,
+}
+
+impl ReplCommand {
+ fn new(name: &'static str, desc: &'static str, unavailable_states: Vec<State>) -> Self {
+ Self {
+ name,
+ description: desc,
+ unavailable_states,
+ }
+ }
+
+ fn unavailable(&self, state: &State) -> bool {
+ self.unavailable_states.contains(state)
+ }
+}
+
/// A default validator which checks for mismatched quotes and brackets
struct ReplValidator;
@@ -339,7 +387,7 @@ fn unknown_command() -> Result<()> {
fn dump_repl_help() {
let head = REPL_COMMANDS
.iter()
- .map(|(name, desc)| format!("{name:<24} {desc}"))
+ .map(|cmd| format!("{:<24} {}", cmd.name, cmd.description))
.collect::<Vec<String>>()
.join("\n");
println!(