diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-02 19:27:41 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-02 19:27:41 +0800 |
| commit | 71f2e94579511d7524f5534377001ab3f02a9597 (patch) | |
| tree | f97471dd33e2db76eaa2ff9c31889cc909ae2e80 | |
| parent | 571d1022f628cb7d2a3125664bd3293bac4471b5 (diff) | |
| download | aichat-71f2e94579511d7524f5534377001ab3f02a9597.tar.gz | |
refactor: switch to bitflags State (#557)
| -rw-r--r-- | Cargo.lock | 1 | ||||
| -rw-r--r-- | Cargo.toml | 1 | ||||
| -rw-r--r-- | src/config/mod.rs | 85 | ||||
| -rw-r--r-- | src/repl/completer.rs | 2 | ||||
| -rw-r--r-- | src/repl/mod.rs | 69 |
5 files changed, 72 insertions, 86 deletions
@@ -38,6 +38,7 @@ dependencies = [ "aws-smithy-eventstream", "base64 0.22.1", "bincode", + "bitflags 2.5.0", "bstr", "bytes", "chrono", @@ -59,6 +59,7 @@ unicode-segmentation = "1.11.0" num_cpus = "1.16.0" threadpool = "1.8.1" json-patch = { version = "2.0.0", default-features = false } +bitflags = "2.5.0" [dependencies.reqwest] version = "0.12.0" diff --git a/src/config/mod.rs b/src/config/mod.rs index 7d4c853..5ae6bec 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -327,22 +327,19 @@ impl Config { Ok(()) } - pub fn state(&self) -> State { + pub fn state(&self) -> StateFlags { + let mut flags = StateFlags::empty(); if let Some(session) = &self.session { if session.is_empty() { - if self.role.is_some() { - State::EmptySessionWithRole - } else { - State::EmptySession - } + flags |= StateFlags::SESSION_EMPTY; } else { - State::Session + flags |= StateFlags::SESSION; } - } else if self.role.is_some() { - State::Role - } else { - State::Normal } + if self.role.is_some() { + flags |= StateFlags::ROLE + } + flags } pub fn set_temperature(&mut self, value: Option<f64>) { @@ -1046,60 +1043,24 @@ pub enum WorkingMode { Serve, } -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -pub enum State { - Normal, - Role, - EmptySession, - EmptySessionWithRole, - Session, -} - -impl State { - pub fn all() -> Vec<Self> { - vec![ - Self::Normal, - Self::Role, - Self::EmptySession, - Self::EmptySessionWithRole, - Self::Session, - ] - } - - pub fn in_session() -> Vec<Self> { - vec![ - Self::EmptySession, - Self::EmptySessionWithRole, - Self::Session, - ] - } - - pub fn not_in_session() -> Vec<Self> { - let excludes: HashSet<_> = Self::in_session().into_iter().collect(); - Self::all() - .into_iter() - .filter(|v| !excludes.contains(v)) - .collect() - } - - pub fn unable_change_role() -> Vec<Self> { - vec![Self::Session] - } - - pub fn able_change_role() -> Vec<Self> { - let excludes: HashSet<_> = Self::unable_change_role().into_iter().collect(); - Self::all() - .into_iter() - .filter(|v| !excludes.contains(v)) - .collect() +bitflags::bitflags! { + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] + pub struct StateFlags: u32 { + const ROLE = 1 << 0; + const SESSION_EMPTY = 1 << 1; + const SESSION = 1 << 2; } +} - pub fn in_role() -> Vec<Self> { - vec![Self::Role, Self::EmptySessionWithRole] - } +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum AssertState { + True(StateFlags), + False(StateFlags), +} - pub fn is_normal(&self) -> bool { - self == &Self::Normal +impl AssertState { + pub fn any() -> Self { + AssertState::False(StateFlags::empty()) } } diff --git a/src/repl/completer.rs b/src/repl/completer.rs index 134cdbe..d026d67 100644 --- a/src/repl/completer.rs +++ b/src/repl/completer.rs @@ -33,7 +33,7 @@ impl Completer for ReplCompleter { .commands .iter() .filter(|cmd| { - if !cmd.is_valid(&state) { + if !cmd.is_valid(state) { return false; } let line = parts diff --git a/src/repl/mod.rs b/src/repl/mod.rs index b78d39d..ba506dd 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::send_stream; -use crate::config::{GlobalConfig, Input, InputContext, State}; +use crate::config::{AssertState, GlobalConfig, Input, InputContext, StateFlags}; use crate::function::need_send_call_results; use crate::render::render_error; use crate::utils::{create_abort_signal, set_text, AbortSignal}; @@ -34,42 +34,62 @@ const MENU_NAME: &str = "completion_menu"; lazy_static! { static ref REPL_COMMANDS: [ReplCommand; 16] = [ - ReplCommand::new(".help", "Show this help message", State::all()), - ReplCommand::new(".info", "View system info", State::all()), - ReplCommand::new(".model", "Change the current LLM", State::all()), + ReplCommand::new(".help", "Show this help message", AssertState::any()), + ReplCommand::new(".info", "View system info", AssertState::any()), + ReplCommand::new(".model", "Change the current LLM", AssertState::any()), ReplCommand::new( ".prompt", "Create a temporary role using a prompt", - State::able_change_role() + AssertState::False(StateFlags::SESSION) ), ReplCommand::new( ".role", "Switch to a specific role", - State::able_change_role() + AssertState::False(StateFlags::SESSION) + ), + ReplCommand::new( + ".info role", + "View role info", + AssertState::True(StateFlags::ROLE), + ), + ReplCommand::new( + ".exit role", + "Leave the role", + AssertState::True(StateFlags::ROLE), + ), + ReplCommand::new( + ".session", + "Begin a chat session", + AssertState::False(StateFlags::SESSION_EMPTY | StateFlags::SESSION), + ), + ReplCommand::new( + ".info session", + "View session info", + AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION), ), - ReplCommand::new(".info role", "View role info", State::in_role(),), - ReplCommand::new(".exit role", "Leave the role", State::in_role(),), - ReplCommand::new(".session", "Begin a chat session", State::not_in_session(),), - ReplCommand::new(".info session", "View session info", State::in_session(),), ReplCommand::new( ".save session", "Save the chat to file", - State::in_session(), + AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION) ), ReplCommand::new( ".clear messages", "Erase messages in the current session", - State::unable_change_role() + AssertState::True(StateFlags::SESSION) ), ReplCommand::new( ".exit session", "End the current session", - State::in_session(), + AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION) ), - ReplCommand::new(".file", "Include files with the message", State::all()), - ReplCommand::new(".set", "Adjust settings", State::all()), - ReplCommand::new(".copy", "Copy the last response", State::all()), - ReplCommand::new(".exit", "Exit the REPL", State::all()), + ReplCommand::new( + ".file", + "Include files with the message", + AssertState::any() + ), + ReplCommand::new(".set", "Adjust settings", AssertState::any()), + ReplCommand::new(".copy", "Copy the last response", AssertState::any()), + ReplCommand::new(".exit", "Exit the REPL", AssertState::any()), ]; 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(); @@ -165,7 +185,7 @@ impl Repl { ".model" => match args { Some(name) => { self.config.write().set_model(name)?; - if self.config.read().state().is_normal() { + if self.config.read().state().is_empty() { self.config.write().set_model_id(); } } @@ -352,20 +372,23 @@ Type ".help" for additional help. pub struct ReplCommand { name: &'static str, description: &'static str, - valid_states: Vec<State>, + state: AssertState, } impl ReplCommand { - fn new(name: &'static str, desc: &'static str, valid_states: Vec<State>) -> Self { + fn new(name: &'static str, desc: &'static str, state: AssertState) -> Self { Self { name, description: desc, - valid_states, + state, } } - fn is_valid(&self, state: &State) -> bool { - self.valid_states.contains(state) + fn is_valid(&self, flags: StateFlags) -> bool { + match self.state { + AssertState::True(check_flags) => check_flags & flags != StateFlags::empty(), + AssertState::False(check_flags) => check_flags & flags == StateFlags::empty(), + } } } |
