diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-11 11:00:12 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-11 11:00:12 +0800 |
| commit | bb867c4fcbc6f42770471f0a0117cd8909f660d7 (patch) | |
| tree | 6293f7f1108309160d1951f53f6429e9b004870d /src | |
| parent | 5635ca6a58fb4a590419335b098b7317285bfb82 (diff) | |
| download | aichat-bb867c4fcbc6f42770471f0a0117cd8909f660d7.tar.gz | |
feat: support bot (#579)
* feat: support bots
* refactor with RoleLike
* improve exiting session
* make bot works with rag
* refactor repl assert state
* add bot banner
* repl complete bots according bots.txt
* fix on windows
* remove threadpool executing function callings
* adjust repl left_prompt
* move bot config to global config.yaml
* `.bot` throw err if funciton callings is not configured
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/model.rs | 23 | ||||
| -rw-r--r-- | src/config/bot.rs | 255 | ||||
| -rw-r--r-- | src/config/input.rs | 130 | ||||
| -rw-r--r-- | src/config/mod.rs | 441 | ||||
| -rw-r--r-- | src/config/role.rs | 154 | ||||
| -rw-r--r-- | src/config/session.rs | 217 | ||||
| -rw-r--r-- | src/function.rs | 150 | ||||
| -rw-r--r-- | src/main.rs | 12 | ||||
| -rw-r--r-- | src/rag/mod.rs | 19 | ||||
| -rw-r--r-- | src/repl/completer.rs | 13 | ||||
| -rw-r--r-- | src/repl/mod.rs | 74 |
11 files changed, 999 insertions, 489 deletions
diff --git a/src/client/model.rs b/src/client/model.rs index 22ad4d9..d555232 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -1,8 +1,10 @@ use super::{ + list_chat_models, message::{Message, MessageContent}, EmbeddingsData, }; +use crate::config::Config; use crate::utils::{estimate_token_length, format_option_value}; use anyhow::{bail, Result}; @@ -41,9 +43,16 @@ impl Model { .collect() } - pub fn find(models: &[&Self], value: &str) -> Option<Self> { + pub fn retrieve(config: &Config, model_id: &str) -> Result<Self> { + match Self::find(&list_chat_models(config), model_id) { + Some(v) => Ok(v), + None => bail!("Invalid model '{model_id}'"), + } + } + + pub fn find(models: &[&Self], model_id: &str) -> Option<Self> { let mut model = None; - let (client_name, model_name) = match value.split_once(':') { + let (client_name, model_name) = match model_id.split_once(':') { Some((client_name, model_name)) => { if model_name.is_empty() { (client_name, None) @@ -51,11 +60,11 @@ impl Model { (client_name, Some(model_name)) } } - None => (value, None), + None => (model_id, None), }; match model_name { Some(model_name) => { - if let Some(found) = models.iter().find(|v| v.id() == value) { + if let Some(found) = models.iter().find(|v| v.id() == model_id) { model = Some((*found).clone()); } else if let Some(found) = models.iter().find(|v| v.client_name == client_name) { let mut found = (*found).clone(); @@ -73,7 +82,11 @@ impl Model { } pub fn id(&self) -> String { - format!("{}:{}", self.client_name, self.data.name) + if self.data.name.is_empty() { + self.client_name.to_string() + } else { + format!("{}:{}", self.client_name, self.data.name) + } } pub fn client_name(&self) -> &str { diff --git a/src/config/bot.rs b/src/config/bot.rs new file mode 100644 index 0000000..4d748bb --- /dev/null +++ b/src/config/bot.rs @@ -0,0 +1,255 @@ +use super::*; + +use crate::{ + client::Model, + function::{Functions, FUNCTION_ALL_MATCHER}, +}; + +use anyhow::{Context, Result}; +use std::{fs::read_to_string, path::Path}; + +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, Serialize)] +pub struct Bot { + name: String, + config: BotConfig, + definition: BotDefinition, + #[serde(skip)] + functions: Functions, + #[serde(skip)] + rag: Option<Arc<Rag>>, + #[serde(skip)] + model: Model, +} + +impl Bot { + pub async fn init( + config: &GlobalConfig, + name: &str, + abort_signal: AbortSignal, + ) -> Result<Self> { + let definition_path = Config::bot_definition_file(name)?; + let functions_path = Config::bot_functions_file(name)?; + let rag_path = Config::bot_rag_file(name)?; + let embeddings_dir = Config::bot_embeddings_dir(name)?; + let definition = BotDefinition::load(&definition_path)?; + let functions = if functions_path.exists() { + Functions::init(&functions_path)? + } else { + Functions::default() + }; + let bot_config = config + .read() + .bots + .iter() + .find(|v| v.name == name) + .cloned() + .unwrap_or_else(|| BotConfig::new(name)); + let model = { + let config = config.read(); + match bot_config.model_id.as_ref() { + Some(model_id) => Model::retrieve(&config, model_id)?, + None => config.current_model().clone(), + } + }; + + let render_options = config.read().get_render_options()?; + let mut markdown_render = MarkdownRender::init(render_options)?; + println!("{}", markdown_render.render(&definition.banner())); + + let rag = if rag_path.exists() { + Some(Arc::new(Rag::load(config, "rag", &rag_path)?)) + } else if embeddings_dir.is_dir() { + println!("The bot has an embeddings directory, RAG is initializing..."); + let ans = Confirm::new("The bot attached embeddings, init RAG?") + .with_default(true) + .prompt()?; + if ans { + let doc_path = embeddings_dir.display().to_string(); + Some(Arc::new( + Rag::init(config, "rag", &rag_path, &[doc_path], abort_signal).await?, + )) + } else { + None + } + } else { + None + }; + + Ok(Self { + name: name.to_string(), + config: bot_config, + definition, + functions, + rag, + model, + }) + } + + pub fn export(&self) -> Result<String> { + let mut value = serde_json::json!(self); + value["functions_dir"] = Config::bot_functions_dir(&self.name)? + .display() + .to_string() + .into(); + value["config_dir"] = Config::bot_config_dir(&self.name)? + .display() + .to_string() + .into(); + let data = serde_yaml::to_string(&value)?; + Ok(data) + } + + pub fn name(&self) -> &str { + &self.name + } + + pub fn functions(&self) -> &Functions { + &self.functions + } + + pub fn definition(&self) -> &BotDefinition { + &self.definition + } + + pub fn rag(&self) -> Option<Arc<Rag>> { + self.rag.clone() + } +} + +impl RoleLike for Bot { + fn to_role(&self) -> Role { + let mut role = Role::new("", &self.definition.instructions); + role.sync(self); + role + } + + fn model(&self) -> &Model { + &self.model + } + + fn temperature(&self) -> Option<f64> { + self.config.temperature + } + + fn top_p(&self) -> Option<f64> { + self.config.top_p + } + + fn function_matcher(&self) -> Option<String> { + if self.functions.is_empty() { + None + } else { + Some(FUNCTION_ALL_MATCHER.into()) + } + } + + fn set_model(&mut self, model: &Model) { + self.config.model_id = Some(model.id()); + self.model = model.clone(); + } + + fn set_temperature(&mut self, value: Option<f64>) { + self.config.temperature = value; + } + + fn set_top_p(&mut self, value: Option<f64>) { + self.config.top_p = value; + } + + fn set_function_matcher(&mut self, _value: Option<String>) {} +} + +#[derive(Debug, Clone, Default, Deserialize, Serialize)] +pub struct BotConfig { + pub name: String, + #[serde(rename(serialize = "model", deserialize = "model"))] + pub model_id: Option<String>, + #[serde(skip_serializing_if = "Option::is_none")] + pub temperature: Option<f64>, + #[serde(skip_serializing_if = "Option::is_none")] + pub top_p: Option<f64>, +} + +impl BotConfig { + pub fn new(name: &str) -> Self { + Self { + name: name.to_string(), + ..Default::default() + } + } +} + +#[derive(Debug, Clone, Default, Deserialize, Serialize)] +pub struct BotDefinition { + pub name: String, + #[serde(default)] + pub description: String, + #[serde(default)] + pub version: String, + pub instructions: String, + #[serde(default)] + pub conversation_starters: Vec<String>, +} + +impl BotDefinition { + pub fn load(path: &Path) -> Result<Self> { + let contents = read_to_string(path) + .with_context(|| format!("Failed to read bot index file at '{}'", path.display()))?; + let definition: Self = serde_yaml::from_str(&contents) + .with_context(|| format!("Failed to load bot at '{}'", path.display()))?; + Ok(definition) + } + + fn banner(&self) -> String { + let BotDefinition { + name, + description, + version, + conversation_starters, + .. + } = self; + let starters = if conversation_starters.is_empty() { + String::new() + } else { + let starters = conversation_starters + .iter() + .map(|v| format!("- {v}")) + .collect::<Vec<_>>() + .join("\n"); + format!( + r#" + +**Conversation Starters** +{starters}"# + ) + }; + format!( + r#"# {name} {version} +{description}{starters} +"# + ) + } +} + +pub fn list_bots() -> Vec<String> { + list_bots_impl().unwrap_or_default() +} + +fn list_bots_impl() -> Result<Vec<String>> { + let base_dir = Config::functions_dir()?; + let contents = read_to_string(base_dir.join("bots.txt"))?; + let bots = contents + .split('\n') + .filter_map(|line| { + let line = line.trim(); + if line.is_empty() { + None + } else { + Some(line.to_string()) + } + }) + .collect(); + Ok(bots) +} diff --git a/src/config/input.rs b/src/config/input.rs index 6403640..0e2aa7a 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -1,11 +1,11 @@ -use super::{role::Role, session::Session, GlobalConfig}; +use super::*; use crate::client::{ init_client, ChatCompletionsData, Client, ImageUrl, Message, MessageContent, MessageContentPart, MessageRole, Model, }; use crate::function::{ToolCallResult, ToolResults}; -use crate::utils::{base64_encode, sha256, warning_text, AbortSignal, IS_STDOUT_TERMINAL}; +use crate::utils::{base64_encode, sha256, AbortSignal}; use anyhow::{bail, Context, Result}; use fancy_regex::Regex; @@ -29,25 +29,28 @@ lazy_static! { pub struct Input { config: GlobalConfig, text: String, - patch_text: Option<String>, + patched_text: Option<String>, medias: Vec<String>, data_urls: HashMap<String, String>, tool_call: Option<ToolResults>, - rag: Option<String>, - context: InputContext, + rag_name: Option<String>, + role: Role, + with_session: bool, } impl Input { - pub fn from_str(config: &GlobalConfig, text: &str, context: Option<InputContext>) -> Self { + pub fn from_str(config: &GlobalConfig, text: &str, role: Option<Role>) -> Self { + let (role, with_session) = resolve_role(&config.read(), role); Self { config: config.clone(), text: text.to_string(), - patch_text: None, + patched_text: None, medias: Default::default(), data_urls: Default::default(), tool_call: None, - rag: None, - context: context.unwrap_or_else(|| InputContext::from_config(config)), + rag_name: None, + role, + with_session, } } @@ -55,7 +58,7 @@ impl Input { config: &GlobalConfig, text: &str, files: Vec<String>, - context: Option<InputContext>, + role: Option<Role>, ) -> Result<Self> { let mut texts = vec![text.to_string()]; let mut medias = vec![]; @@ -93,15 +96,17 @@ impl Input { } } + let (role, session) = resolve_role(&config.read(), role); Ok(Self { config: config.clone(), text: texts.join("\n"), - patch_text: None, + patched_text: None, medias, data_urls, tool_call: Default::default(), - rag: None, - context: context.unwrap_or_else(|| InputContext::from_config(config)), + rag_name: None, + role, + with_session: session, }) } @@ -114,7 +119,7 @@ impl Input { } pub fn text(&self) -> String { - match self.patch_text.clone() { + match self.patched_text.clone() { Some(text) => text, None => self.text.clone(), } @@ -134,19 +139,19 @@ impl Input { let top_k = self.config.read().rag_top_k; let embeddings = rag.search(&self.text, top_k, abort_signal).await?; let text = self.config.read().rag_template(&embeddings, &self.text); - self.patch_text = Some(text); - self.rag = Some(rag.name().to_string()); + self.patched_text = Some(text); + self.rag_name = Some(rag.name().to_string()); } } Ok(()) } - pub fn rag(&self) -> Option<&str> { - self.rag.as_deref() + pub fn rag_name(&self) -> Option<&str> { + self.rag_name.as_deref() } pub fn clear_patch_text(&mut self) { - self.patch_text.take(); + self.patched_text.take(); } pub fn merge_tool_call( @@ -164,20 +169,8 @@ impl Input { self } - pub fn model(&self) -> Model { - if let Some(session) = self.session(&self.config.read().session) { - return session.model.clone(); - } else if let Some(model) = self - .role() - .and_then(|v| v.retrieve_model(&self.config.read())) - { - return model; - } - self.config.read().model.clone() - } - pub fn create_client(&self) -> Result<Box<dyn Client>> { - init_client(&self.config, Some(self.model())) + init_client(&self.config, Some(self.role().model().clone())) } pub fn prepare_completion_data( @@ -190,35 +183,9 @@ impl Input { } let messages = self.build_messages()?; self.config.read().model.guard_max_input_tokens(&messages)?; - let (temperature, top_p) = if let Some(session) = self.session(&self.config.read().session) - { - (session.temperature(), session.top_p()) - } else if let Some(role) = self.role() { - (role.temperature, role.top_p) - } else { - let config = self.config.read(); - (config.temperature, config.top_p) - }; - let mut functions = None; - if self.config.read().function_calling { - let config = self.config.read(); - let function_matcher = if let Some(session) = self.session(&config.session) { - session.function_matcher() - } else if let Some(role) = self.role() { - role.function_matcher.as_deref() - } else { - None - }; - if let Some(function_matcher) = function_matcher { - functions = config.function.select(function_matcher); - if !model.supports_function_calling() { - functions = None; - if *IS_STDOUT_TERMINAL { - eprintln!("{}", warning_text("WARNING: the role or session includes functions, but the model or client does not support function calling.")); - } - } - } - }; + let temperature = self.role().temperature(); + let top_p = self.role().top_p(); + let functions = self.config.read().retrieve_functions(model, self.role()); Ok(ChatCompletionsData { messages, temperature, @@ -231,10 +198,8 @@ impl Input { pub fn build_messages(&self) -> Result<Vec<Message>> { let mut messages = if let Some(session) = self.session(&self.config.read().session) { session.build_messages(self) - } else if let Some(role) = self.role() { - role.build_messages(self) } else { - vec![Message::new(MessageRole::User, self.message_content())] + self.role().build_messages(self) }; if let Some(tool_results) = &self.tool_call { messages.push(Message::new( @@ -248,19 +213,17 @@ impl Input { pub fn echo_messages(&self) -> String { if let Some(session) = self.session(&self.config.read().session) { session.echo_messages(self) - } else if let Some(role) = self.role() { - role.echo_messages(self) } else { - self.render() + self.role().echo_messages(self) } } - pub fn role(&self) -> Option<&Role> { - self.context.role.as_ref() + pub fn role(&self) -> &Role { + &self.role } pub fn session<'a>(&self, session: &'a Option<Session>) -> Option<&'a Session> { - if self.context.session { + if self.with_session { session.as_ref() } else { None @@ -268,7 +231,7 @@ impl Input { } pub fn session_mut<'a>(&self, session: &'a mut Option<Session>) -> Option<&'a mut Session> { - if self.context.session { + if self.with_session { session.as_mut() } else { None @@ -337,27 +300,10 @@ impl Input { } } -#[derive(Debug, Clone, Default)] -pub struct InputContext { - role: Option<Role>, - session: bool, -} - -impl InputContext { - pub fn new(role: Option<Role>, session: bool) -> Self { - Self { role, session } - } - - pub fn from_config(config: &GlobalConfig) -> Self { - let config = config.read(); - InputContext::new(config.role.clone(), config.session.is_some()) - } - - pub fn role(role: Role) -> Self { - Self { - role: Some(role), - session: false, - } +fn resolve_role(config: &Config, role: Option<Role>) -> (Role, bool) { + match role { + Some(v) => (v, false), + None => (config.extract_role(), config.session.is_some()), } } diff --git a/src/config/mod.rs b/src/config/mod.rs index fbee667..b874b16 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -1,21 +1,23 @@ +mod bot; mod input; mod role; mod session; -pub use self::input::{Input, InputContext}; -pub use self::role::{Role, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE}; -use self::session::{Session, TEMP_SESSION_NAME}; +pub use self::bot::{list_bots, Bot, BotConfig}; +pub use self::input::Input; +pub use self::role::{Role, RoleLike, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE}; +use self::session::Session; use crate::client::{ create_client_config, list_chat_models, list_client_types, ClientConfig, Model, OPENAI_COMPATIBLE_PLATFORMS, }; -use crate::function::{Function, ToolCallResult}; -use crate::rag::{Rag, TEMP_RAG_NAME}; +use crate::function::{FunctionDeclaration, Functions, ToolCallResult}; +use crate::rag::Rag; use crate::render::{MarkdownRender, RenderOptions}; use crate::utils::{ format_option_value, fuzzy_match, get_env_name, light_theme_from_colorfgbg, now, render_prompt, - set_text, AbortSignal, IS_STDOUT_TERMINAL, + set_text, warning_text, AbortSignal, IS_STDOUT_TERMINAL, }; use anyhow::{anyhow, bail, Context, Result}; @@ -44,6 +46,15 @@ const MESSAGES_FILE_NAME: &str = "messages.md"; const SESSIONS_DIR_NAME: &str = "sessions"; const RAGS_DIR_NAME: &str = "rags"; const FUNCTIONS_DIR_NAME: &str = "functions"; +const FUNCTIONS_FILE_NAME: &str = "functions.json"; +const BOTS_DIR_NAME: &str = "bots"; +const BOT_DEFINITION_FILE_NAME: &str = "index.yaml"; +const BOT_EMBEDDINGS_DIR: &str = "embeddings"; +const BOT_RAG_FILE_NAME: &str = "rag.bin"; + +pub const TEMP_ROLE_NAME: &str = "%%"; +pub const TEMP_RAG_NAME: &str = "temp"; +pub const TEMP_SESSION_NAME: &str = "temp"; const CLIENTS_FIELD: &str = "clients"; @@ -51,7 +62,7 @@ const SUMMARIZE_PROMPT: &str = "Summarize the discussion briefly in 200 words or less to use as a prompt for future context."; const SUMMARY_PROMPT: &str = "This is a summary of the chat history as a recap: "; -const RAG_TEMPLATE: &str = r#"Answer the following question based only on the provided context: +const RAG_TEMPLATE: &str = r#"Answer the question based only on the provided context: <context> __CONTEXT__ </context> @@ -59,7 +70,7 @@ __CONTEXT__ Question: __INPUT__ "#; -const LEFT_PROMPT: &str = "{color.green}{?session {session}{?role /}}{role}{?rag #{rag}}{color.cyan}{?session )}{!session >}{color.reset} "; +const LEFT_PROMPT: &str = "{color.green}{?session {?bot {bot}#}{session}{?role /}}{!session {?bot {bot}}}{role}{?rag @{rag}}{color.cyan}{?session )}{!session >}{color.reset} "; const RIGHT_PROMPT: &str = "{color.purple}{?session {?consume_tokens {consume_tokens}({consume_percent}%)}{!consume_tokens {consume_tokens}}}{color.reset}"; #[derive(Debug, Clone, Deserialize)] @@ -91,6 +102,7 @@ pub struct Config { pub left_prompt: Option<String>, pub right_prompt: Option<String>, pub clients: Vec<ClientConfig>, + pub bots: Vec<BotConfig>, #[serde(skip)] pub roles: Vec<Role>, #[serde(skip)] @@ -100,9 +112,11 @@ pub struct Config { #[serde(skip)] pub rag: Option<Arc<Rag>>, #[serde(skip)] + pub bot: Option<Bot>, + #[serde(skip)] pub model: Model, #[serde(skip)] - pub function: Function, + pub functions: Functions, #[serde(skip)] pub working_mode: WorkingMode, #[serde(skip)] @@ -136,12 +150,14 @@ impl Default for Config { left_prompt: None, right_prompt: None, clients: vec![], + bots: vec![], roles: vec![], role: None, session: None, rag: None, + bot: None, model: Default::default(), - function: Default::default(), + functions: Default::default(), working_mode: WorkingMode::Command, last_message: None, } @@ -168,7 +184,7 @@ impl Config { config.set_wrap(&wrap)?; } - config.function = Function::init(&Self::functions_dir()?)?; + config.functions = Functions::init(&Self::functions_file()?)?; config.working_mode = working_mode; config.load_roles()?; @@ -211,7 +227,8 @@ impl Config { } pub fn retrieve_role(&self, name: &str) -> Result<Role> { - self.roles + let mut role = self + .roles .iter() .find(|v| v.match_name(name)) .map(|v| { @@ -219,7 +236,18 @@ impl Config { role.complete_prompt_args(name); role }) - .ok_or_else(|| anyhow!("Unknown role `{name}`")) + .ok_or_else(|| anyhow!("Unknown role `{name}`"))?; + + match role.model_id() { + Some(model_id) => { + if self.model.id() != model_id { + let model = Model::retrieve(self, model_id)?; + role.set_model(&model); + } + } + None => role.set_model(&self.model), + } + Ok(role) } pub fn config_dir() -> Result<PathBuf> { @@ -268,11 +296,20 @@ impl Config { let timestamp = now(); let summary = input.summary(); let input_markdown = input.render(); - let scope = match (input.role().map(|v| v.name.as_str()), input.rag()) { - (Some(role), Some(rag)) => format!(" ({role}#{rag})"), - (Some(role), _) => format!(" ({role})"), - (None, Some(rag)) => format!(" (#{rag})"), - _ => String::new(), + let scope = if self.bot.is_none() { + let role_name = if input.role().is_derived() { + None + } else { + Some(input.role().name()) + }; + match (role_name, input.rag_name()) { + (Some(role), Some(rag_name)) => format!(" ({role}#{rag_name})"), + (Some(role), _) => format!(" ({role})"), + (None, Some(rag_name)) => format!(" (#{rag_name})"), + _ => String::new(), + } + } else { + String::new() }; let output = format!("# CHAT: {summary} [{timestamp}]{scope}\n{input_markdown}\n--------\n{output}\n--------\n\n",); file.write_all(output.as_bytes()) @@ -299,17 +336,23 @@ impl Config { } } - pub fn messages_file() -> Result<PathBuf> { - match env::var(get_env_name("messages_file")) { - Ok(value) => Ok(PathBuf::from(value)), - Err(_) => Self::local_path(MESSAGES_FILE_NAME), + pub fn messages_file(&self) -> Result<PathBuf> { + match &self.bot { + None => match env::var(get_env_name("messages_file")) { + Ok(value) => Ok(PathBuf::from(value)), + Err(_) => Self::local_path(MESSAGES_FILE_NAME), + }, + Some(bot) => Ok(Self::bot_config_dir(bot.name())?.join(MESSAGES_FILE_NAME)), } } - pub fn sessions_dir() -> Result<PathBuf> { - match env::var(get_env_name("sessions_dir")) { - Ok(value) => Ok(PathBuf::from(value)), - Err(_) => Self::local_path(SESSIONS_DIR_NAME), + pub fn sessions_dir(&self) -> Result<PathBuf> { + match &self.bot { + None => match env::var(get_env_name("sessions_dir")) { + Ok(value) => Ok(PathBuf::from(value)), + Err(_) => Self::local_path(SESSIONS_DIR_NAME), + }, + Some(bot) => Ok(Self::bot_config_dir(bot.name())?.join(SESSIONS_DIR_NAME)), } } @@ -327,20 +370,66 @@ impl Config { } } - pub fn session_file(name: &str) -> Result<PathBuf> { - let mut path = Self::sessions_dir()?; - path.push(&format!("{name}.yaml")); - Ok(path) + pub fn functions_file() -> Result<PathBuf> { + Ok(Self::functions_dir()?.join(FUNCTIONS_FILE_NAME)) } - pub fn rag_file(name: &str) -> Result<PathBuf> { - let mut path = Self::rags_dir()?; - path.push(&format!("{name}.bin")); + pub fn functions_bin_dir() -> Result<PathBuf> { + Ok(Self::functions_dir()?.join("bin")) + } + + pub fn session_file(&self, name: &str) -> Result<PathBuf> { + Ok(self.sessions_dir()?.join(format!("{name}.yaml"))) + } + + pub fn rag_file(&self, name: &str) -> Result<PathBuf> { + let path = if self.bot.is_none() { + Self::rags_dir()?.join(format!("{name}.bin")) + } else { + Self::rags_dir()? + .join(BOTS_DIR_NAME) + .join(format!("{name}.bin")) + }; Ok(path) } + pub fn bots_dir() -> Result<PathBuf> { + match env::var(get_env_name("bots_config_dir")) { + Ok(value) => Ok(PathBuf::from(value)), + Err(_) => Self::local_path(BOTS_DIR_NAME), + } + } + + pub fn bot_config_dir(name: &str) -> Result<PathBuf> { + Ok(Self::bots_dir()?.join(name)) + } + + pub fn bot_rag_file(name: &str) -> Result<PathBuf> { + Ok(Self::bot_config_dir(name)?.join(BOT_RAG_FILE_NAME)) + } + + pub fn bots_functions_dir() -> Result<PathBuf> { + Ok(Self::functions_dir()?.join(BOTS_DIR_NAME)) + } + + pub fn bot_functions_dir(name: &str) -> Result<PathBuf> { + Ok(Self::bots_functions_dir()?.join(name)) + } + + pub fn bot_functions_file(name: &str) -> Result<PathBuf> { + Ok(Self::bot_functions_dir(name)?.join(FUNCTIONS_FILE_NAME)) + } + + pub fn bot_definition_file(name: &str) -> Result<PathBuf> { + Ok(Self::bot_functions_dir(name)?.join(BOT_DEFINITION_FILE_NAME)) + } + + pub fn bot_embeddings_dir(name: &str) -> Result<PathBuf> { + Ok(Self::bot_functions_dir(name)?.join(BOT_EMBEDDINGS_DIR)) + } + pub fn use_prompt(&mut self, prompt: &str) -> Result<()> { - let role = Role::temp(prompt); + let role = Role::new(TEMP_ROLE_NAME, prompt); self.use_role_obj(role) } @@ -350,22 +439,25 @@ impl Config { } pub fn use_role_obj(&mut self, role: Role) -> Result<()> { + if self.bot.is_some() { + bail!("Cannot perform this action because you are using a bot") + } if let Some(session) = self.session.as_mut() { session.guard_empty()?; - session.set_role_properties(&role); - } - if let Some(model_id) = &role.model_id { - self.set_model(model_id)?; + session.set_role(role); + } else { + self.role = Some(role); } - self.role = Some(role); Ok(()) } pub fn exit_role(&mut self) -> Result<()> { - if self.session.is_none() { - self.restore_model()?; + if self.role.is_some() { + if let Some(session) = self.session.as_mut() { + session.clear_role(); + } + self.role = None; } - self.role = None; Ok(()) } @@ -378,6 +470,9 @@ impl Config { flags |= StateFlags::SESSION; } } + if self.bot.is_some() { + flags |= StateFlags::BOT; + } if self.role.is_some() { flags |= StateFlags::ROLE; } @@ -387,27 +482,62 @@ impl Config { flags } - pub fn has_role_or_session(&self) -> bool { - self.role.is_some() || self.session.is_some() + pub fn current_model(&self) -> &Model { + if let Some(session) = self.session.as_ref() { + session.model() + } else if let Some(bot) = self.bot.as_ref() { + bot.model() + } else if let Some(role) = self.role.as_ref() { + role.model() + } else { + &self.model + } } - pub fn set_temperature(&mut self, value: Option<f64>) { - if let Some(session) = self.session.as_mut() { - session.set_temperature(value); - } else if let Some(role) = self.role.as_mut() { - role.set_temperature(value); + pub fn extract_role(&self) -> Role { + let mut role = if let Some(session) = self.session.as_ref() { + session.to_role() + } else if let Some(bot) = self.bot.as_ref() { + bot.to_role() + } else if let Some(role) = self.role.as_ref() { + role.clone() } else { - self.temperature = value; + let mut role = Role::default(); + role.batch_set(&self.model, self.temperature, self.top_p, None); + role + }; + if role.temperature().is_none() && self.temperature.is_some() { + role.set_temperature(self.temperature); } + if role.top_p().is_none() && self.top_p.is_some() { + role.set_top_p(self.top_p); + } + role } - pub fn set_top_p(&mut self, value: Option<f64>) { + pub fn role_like_mut(&mut self) -> Option<&mut dyn RoleLike> { if let Some(session) = self.session.as_mut() { - session.set_top_p(value); + Some(session) + } else if let Some(bot) = self.bot.as_mut() { + Some(bot) } else if let Some(role) = self.role.as_mut() { - role.set_top_p(value); + Some(role) } else { - self.top_p = value; + None + } + } + + pub fn set_temperature(&mut self, value: Option<f64>) { + match self.role_like_mut() { + Some(role_like) => role_like.set_temperature(value), + None => self.temperature = value, + } + } + + pub fn set_top_p(&mut self, value: Option<f64>) { + match self.role_like_mut() { + Some(role_like) => role_like.set_top_p(value), + None => self.top_p = value, } } @@ -441,46 +571,26 @@ impl Config { Ok(()) } - pub fn set_model(&mut self, value: &str) -> Result<()> { - let model = Model::find(&list_chat_models(self), value); - match model { - None => bail!("No model '{}'", value), - Some(model) => { - if let Some(session) = self.session.as_mut() { - session.set_model(&model); - } else if let Some(role) = self.role.as_mut() { - role.set_model(&model); - } + pub fn set_model(&mut self, model_id: &str) -> Result<()> { + let model = Model::retrieve(self, model_id)?; + match self.role_like_mut() { + Some(role_like) => role_like.set_model(&model), + None => { self.model = model; - Ok(()) } } + Ok(()) } - pub fn set_model_id(&mut self) { - self.model_id = self.model.id() - } - - pub fn restore_model(&mut self) -> Result<()> { - let origin_model_id = self.model_id.clone(); - self.set_model(&origin_model_id) - } - - pub fn system_info(&self) -> Result<String> { + pub fn sysinfo(&self) -> Result<String> { let display_path = |path: &Path| path.display().to_string(); let wrap = self .wrap .clone() .map_or_else(|| String::from("no"), |v| v.to_string()); - let (temperature, top_p) = if let Some(session) = &self.session { - (session.temperature(), session.top_p()) - } else if let Some(role) = &self.role { - (role.temperature, role.top_p) - } else { - (self.temperature, self.top_p) - }; + let role = self.extract_role(); let items = vec![ - ("model", self.model.id()), + ("model", role.model().id()), ( "max_output_tokens", self.model @@ -488,8 +598,8 @@ impl Config { .map(|v| format!("{v} (current model)")) .unwrap_or_else(|| "-".into()), ), - ("temperature", format_option_value(&temperature)), - ("top_p", format_option_value(&top_p)), + ("temperature", format_option_value(&role.temperature())), + ("top_p", format_option_value(&role.top_p())), ("rag_top_k", self.rag_top_k.to_string()), ("function_calling", self.function_calling.to_string()), ("compress_threshold", self.compress_threshold.to_string()), @@ -505,10 +615,11 @@ impl Config { ("prelude", format_option_value(&self.prelude)), ("config_file", display_path(&Self::config_file()?)), ("roles_file", display_path(&Self::roles_file()?)), - ("messages_file", display_path(&Self::messages_file()?)), - ("sessions_dir", display_path(&Self::sessions_dir()?)), - ("rags_dir", display_path(&Self::rags_dir()?)), ("functions_dir", display_path(&Self::functions_dir()?)), + ("rags_dir", display_path(&Self::rags_dir()?)), + ("bots_dir", display_path(&Self::bots_dir()?)), + ("sessions_dir", display_path(&self.sessions_dir()?)), + ("messages_file", display_path(&self.messages_file()?)), ]; let output = items .iter() @@ -530,7 +641,7 @@ impl Config { if let Some(session) = &self.session { let render_options = self.get_render_options()?; let mut markdown_render = MarkdownRender::init(render_options)?; - session.info(&mut markdown_render) + session.render(&mut markdown_render) } else { bail!("No session") } @@ -544,6 +655,14 @@ impl Config { } } + pub fn bot_info(&self) -> Result<String> { + if let Some(bot) = &self.bot { + bot.export() + } else { + bail!("No rag") + } + } + pub fn info(&self) -> Result<String> { if let Some(session) = &self.session { session.export() @@ -552,7 +671,7 @@ impl Config { } else if let Some(rag) = &self.rag { rag.export() } else { - self.system_info() + self.sysinfo() } } @@ -563,28 +682,25 @@ impl Config { .unwrap_or_default() } - pub fn repl_complete(&self, cmd: &str, args: &[&str]) -> Vec<(String, String)> { + pub fn repl_complete(&self, cmd: &str, args: &[&str]) -> Vec<(String, Option<String>)> { let (values, filter) = if args.len() == 1 { let values = match cmd { ".role" => self .roles .iter() - .map(|v| (v.name.clone(), String::new())) + .map(|v| (v.name().to_string(), None)) .collect(), ".model" => list_chat_models(self) .into_iter() - .map(|v| (v.id(), v.description())) + .map(|v| (v.id(), Some(v.description()))) .collect(), ".session" => self .list_sessions() .into_iter() - .map(|v| (v.clone(), String::new())) - .collect(), - ".rag" => self - .list_rags() - .into_iter() - .map(|v| (v.clone(), String::new())) + .map(|v| (v, None)) .collect(), + ".rag" => self.list_rags().into_iter().map(|v| (v, None)).collect(), + ".bot" => list_bots().into_iter().map(|v| (v, None)).collect(), ".set" => vec![ "max_output_tokens", "temperature", @@ -599,7 +715,7 @@ impl Config { "auto_copy", ] .into_iter() - .map(|v| (format!("{v} "), String::new())) + .map(|v| (format!("{v} "), None)) .collect(), _ => vec![], }; @@ -625,10 +741,7 @@ impl Config { "auto_copy" => complete_bool(self.auto_copy), _ => vec![], }; - ( - values.into_iter().map(|v| (v, String::new())).collect(), - args[1], - ) + (values.into_iter().map(|v| (v, None)).collect(), args[1]) } else { return vec![]; }; @@ -696,6 +809,30 @@ impl Config { Ok(()) } + pub fn retrieve_functions( + &self, + model: &Model, + role: &Role, + ) -> Option<Vec<FunctionDeclaration>> { + let mut functions = None; + if self.function_calling { + let function_matcher = role.function_matcher(); + if let Some(matcher) = function_matcher { + functions = match &self.bot { + Some(bot) => bot.functions().select(&matcher), + None => self.functions.select(&matcher), + }; + if !model.supports_function_calling() { + functions = None; + if *IS_STDOUT_TERMINAL { + eprintln!("{}", warning_text("WARNING: the role or session includes functions, but the model or client does not support function calling.")); + } + } + } + }; + functions + } + pub fn use_session(&mut self, session: Option<&str>) -> Result<()> { if self.session.is_some() { bail!( @@ -704,7 +841,7 @@ impl Config { } match session { None => { - let session_file = Self::session_file(TEMP_SESSION_NAME)?; + let session_file = self.session_file(TEMP_SESSION_NAME)?; if session_file.exists() { remove_file(session_file).with_context(|| { format!("Failed to cleanup previous '{TEMP_SESSION_NAME}' session") @@ -714,14 +851,12 @@ impl Config { self.session = Some(session); } Some(name) => { - let session_path = Self::session_file(name)?; + let session_path = self.session_file(name)?; if !session_path.exists() { self.session = Some(Session::new(self, name)); } else { - let session = Session::load(name, &session_path)?; - let model_id = session.model_id().to_string(); + let session = Session::load(self, name, &session_path)?; self.session = Some(session); - self.set_model(&model_id)?; } } } @@ -745,20 +880,19 @@ impl Config { pub fn exit_session(&mut self) -> Result<()> { if let Some(mut session) = self.session.take() { let is_repl = self.working_mode == WorkingMode::Repl; - let sessions_dir = Self::sessions_dir()?; + let sessions_dir = self.sessions_dir()?; session.exit(&sessions_dir, is_repl)?; self.last_message = None; - self.restore_model()?; } Ok(()) } pub fn save_session(&mut self, name: &str) -> Result<()> { + let sessions_dir = self.sessions_dir()?; if let Some(session) = self.session.as_mut() { if !name.is_empty() { - session.name = name.to_string(); + session.set_name(name); } - let sessions_dir = Self::sessions_dir()?; session.save(&sessions_dir)?; } Ok(()) @@ -772,7 +906,7 @@ impl Config { } pub fn list_sessions(&self) -> Vec<String> { - let sessions_dir = match Self::sessions_dir() { + let sessions_dir = match self.sessions_dir() { Ok(dir) => dir, Err(_) => return vec![], }; @@ -795,7 +929,7 @@ impl Config { pub fn should_compress_session(&mut self) -> bool { if let Some(session) = self.session.as_mut() { if session.need_compress(self.compress_threshold) { - session.compressing = true; + session.set_compressing(true); return true; } } @@ -816,13 +950,13 @@ impl Config { pub fn is_compressing_session(&self) -> bool { self.session .as_ref() - .map(|v| v.compressing) + .map(|v| v.compressing()) .unwrap_or_default() } pub fn end_compressing_session(&mut self) { if let Some(session) = self.session.as_mut() { - session.compressing = false; + session.set_compressing(false); } } @@ -831,23 +965,23 @@ impl Config { rag: Option<&str>, abort_signal: AbortSignal, ) -> Result<()> { - if config.read().rag.is_some() { - bail!("Already in a rag, please run '.exit rag' first to exit the current rag."); + if config.read().bot.is_some() { + bail!("Cannot perform this action because you are using a bot") } let rag = match rag { None => { - let rag_path = Self::rag_file(TEMP_RAG_NAME)?; + let rag_path = config.read().rag_file(TEMP_RAG_NAME)?; if rag_path.exists() { remove_file(&rag_path).with_context(|| { format!("Failed to cleanup previous '{TEMP_RAG_NAME}' rag") })?; } - Rag::init(config, TEMP_RAG_NAME, &rag_path, abort_signal).await? + Rag::init(config, TEMP_RAG_NAME, &rag_path, &[], abort_signal).await? } Some(name) => { - let rag_path = Self::rag_file(name)?; + let rag_path = config.read().rag_file(name)?; if !rag_path.exists() { - Rag::init(config, name, &rag_path, abort_signal).await? + Rag::init(config, name, &rag_path, &[], abort_signal).await? } else { Rag::load(config, name, &rag_path)? } @@ -894,6 +1028,29 @@ impl Config { .replace("__INPUT__", text) } + pub async fn use_bot( + config: &GlobalConfig, + name: &str, + abort_signal: AbortSignal, + ) -> Result<()> { + if !config.read().function_calling { + bail!("Before using the bot, please configure function calling first."); + } + if config.read().bot.is_some() { + bail!("Already in a bot, please run '.exit bot' first to exit the current bot."); + } + let bot = Bot::init(config, name, abort_signal).await?; + config.write().rag = bot.rag(); + config.write().bot = Some(bot); + Ok(()) + } + + pub fn exit_bot(&mut self) -> Result<()> { + self.rag.take(); + self.bot.take(); + Ok(()) + } + pub fn get_render_options(&self) -> Result<RenderOptions> { let theme = if self.highlight { let theme_mode = if self.light_theme { "light" } else { "dark" }; @@ -940,22 +1097,23 @@ impl Config { fn generate_prompt_context(&self) -> HashMap<&str, String> { let mut output = HashMap::new(); - output.insert("model", self.model.id()); - output.insert("client_name", self.model.client_name().to_string()); - output.insert("model_name", self.model.name().to_string()); + let role = self.extract_role(); + output.insert("model", role.model().id()); + output.insert("client_name", role.model().client_name().to_string()); + output.insert("model_name", role.model().name().to_string()); output.insert( "max_input_tokens", - self.model + role.model() .max_input_tokens() .unwrap_or_default() .to_string(), ); - if let Some(temperature) = self.temperature { + if let Some(temperature) = role.temperature() { if temperature != 0.0 { output.insert("temperature", temperature.to_string()); } } - if let Some(top_p) = self.top_p { + if let Some(top_p) = role.top_p() { if top_p != 0.0 { output.insert("top_p", top_p.to_string()); } @@ -974,13 +1132,13 @@ impl Config { if self.auto_copy { output.insert("auto_copy", "true".to_string()); } - if let Some(role) = &self.role { - output.insert("role", role.name.clone()); + if !role.is_derived() { + output.insert("role", role.name().to_string()); } if let Some(session) = &self.session { output.insert("session", session.name().to_string()); - output.insert("dirty", session.dirty.to_string()); - let (tokens, percent) = session.tokens_and_percent(); + output.insert("dirty", session.dirty().to_string()); + let (tokens, percent) = session.tokens_usage(); output.insert("consume_tokens", tokens.to_string()); output.insert("consume_percent", percent.to_string()); output.insert("user_messages_len", session.user_messages_len().to_string()); @@ -988,6 +1146,9 @@ impl Config { if let Some(rag) = &self.rag { output.insert("rag", rag.name().to_string()); } + if let Some(bot) = &self.bot { + output.insert("bot", bot.name().to_string()); + } if self.highlight { output.insert("color.reset", "\u{1b}[0m".to_string()); @@ -1015,7 +1176,7 @@ impl Config { } fn open_message_file(&self) -> Result<File> { - let path = Self::messages_file()?; + let path = self.messages_file()?; ensure_parent_exists(&path)?; OpenOptions::new() .create(true) @@ -1078,10 +1239,10 @@ impl Config { .with_context(|| format!("Failed to load roles at {}", path.display()))?; serde_yaml::from_str(&content).with_context(|| "Invalid roles config")? }; - let exist_roles: HashSet<_> = self.roles.iter().map(|v| v.name.clone()).collect(); + let exist_roles: HashSet<_> = self.roles.iter().map(|v| v.name().to_string()).collect(); let builtin_roles = Role::builtin(); for role in builtin_roles { - if !exist_roles.contains(&role.name) { + if !exist_roles.contains(role.name()) { self.roles.push(role); } } @@ -1165,6 +1326,7 @@ bitflags::bitflags! { const SESSION_EMPTY = 1 << 1; const SESSION = 1 << 2; const RAG = 1 << 3; + const BOT = 1 << 4; } } @@ -1172,12 +1334,17 @@ bitflags::bitflags! { pub enum AssertState { True(StateFlags), False(StateFlags), + TrueFalse(StateFlags, StateFlags), + Equal(StateFlags), } impl AssertState { - pub fn any() -> Self { + pub fn pass() -> Self { AssertState::False(StateFlags::empty()) } + pub fn bare() -> Self { + AssertState::Equal(StateFlags::empty()) + } } fn create_config_file(config_path: &Path) -> Result<()> { diff --git a/src/config/role.rs b/src/config/role.rs index 1b49bad..135bc50 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -1,46 +1,58 @@ -use super::{Config, Input}; +use super::*; use crate::{ - client::{list_chat_models, Message, MessageContent, MessageRole, Model}, + client::{Message, MessageContent, MessageRole, Model}, + function::FUNCTION_ALL_MATCHER, utils::{detect_os, detect_shell}, }; use anyhow::{Context, Result}; use serde::{Deserialize, Serialize}; -pub const TEMP_ROLE: &str = "%%"; pub const SHELL_ROLE: &str = "%shell%"; pub const EXPLAIN_SHELL_ROLE: &str = "%explain-shell%"; pub const CODE_ROLE: &str = "%code%"; pub const INPUT_PLACEHOLDER: &str = "__INPUT__"; -#[derive(Debug, Clone, Deserialize, Serialize)] +pub trait RoleLike { + fn to_role(&self) -> Role; + fn model(&self) -> &Model; + fn temperature(&self) -> Option<f64>; + fn top_p(&self) -> Option<f64>; + fn function_matcher(&self) -> Option<String>; + fn set_model(&mut self, model: &Model); + fn set_temperature(&mut self, value: Option<f64>); + fn set_top_p(&mut self, value: Option<f64>); + fn set_function_matcher(&mut self, value: Option<String>); +} + +#[derive(Debug, Clone, Default, Deserialize, Serialize)] pub struct Role { - pub name: String, - pub prompt: String, + name: String, + prompt: String, #[serde( rename(serialize = "model", deserialize = "model"), skip_serializing_if = "Option::is_none" )] - pub model_id: Option<String>, + model_id: Option<String>, #[serde(skip_serializing_if = "Option::is_none")] - pub temperature: Option<f64>, + temperature: Option<f64>, #[serde(skip_serializing_if = "Option::is_none")] - pub top_p: Option<f64>, + top_p: Option<f64>, #[serde(skip_serializing_if = "Option::is_none")] - pub function_matcher: Option<String>, + function_matcher: Option<String>, + + #[serde(skip)] + model: Model, } impl Role { - pub fn temp(prompt: &str) -> Self { + pub fn new(name: &str, prompt: &str) -> Self { Self { - name: TEMP_ROLE.into(), + name: name.into(), prompt: prompt.into(), - temperature: None, - model_id: None, - top_p: None, - function_matcher: None, + ..Default::default() } } @@ -71,16 +83,18 @@ async function timeout(ms) { .into(), None, ), - ("%functions%", String::new(), Some(".*".into())), + ( + "%functions%", + String::new(), + Some(FUNCTION_ALL_MATCHER.into()), + ), ] .into_iter() .map(|(name, prompt, function_matcher)| Self { name: name.into(), prompt, - model_id: None, - temperature: None, - top_p: None, function_matcher, + ..Default::default() }) .collect() } @@ -91,30 +105,55 @@ async function timeout(ms) { Ok(output.trim_end().to_string()) } - pub fn empty_prompt(&self) -> bool { - self.prompt.is_empty() + pub fn sync<T: RoleLike>(&mut self, role_like: &T) { + let model = role_like.model(); + let temperature = role_like.temperature(); + let top_p = role_like.top_p(); + let function_matcher = role_like.function_matcher(); + self.batch_set(model, temperature, top_p, function_matcher); } - pub fn embedded_prompt(&self) -> bool { - self.prompt.contains(INPUT_PLACEHOLDER) + pub fn batch_set( + &mut self, + model: &Model, + temperature: Option<f64>, + top_p: Option<f64>, + function_matcher: Option<String>, + ) { + self.set_model(model); + if temperature.is_some() { + self.set_temperature(temperature); + } + if top_p.is_some() { + self.set_top_p(top_p); + } + if function_matcher.is_some() { + self.set_function_matcher(function_matcher); + } } - pub fn retrieve_model(&self, config: &Config) -> Option<Model> { - self.model_id - .as_ref() - .and_then(|model_id| Model::find(&list_chat_models(config), model_id)) + pub fn is_derived(&self) -> bool { + self.name.is_empty() } - pub fn set_model(&mut self, model: &Model) { - self.model_id = Some(model.id()); + pub fn name(&self) -> &str { + &self.name } - pub fn set_temperature(&mut self, value: Option<f64>) { - self.temperature = value; + pub fn model_id(&self) -> Option<&str> { + self.model_id.as_deref() } - pub fn set_top_p(&mut self, value: Option<f64>) { - self.top_p = value; + pub fn prompt(&self) -> &str { + &self.prompt + } + + pub fn is_empty_prompt(&self) -> bool { + self.prompt.is_empty() + } + + pub fn is_embedded_prompt(&self) -> bool { + self.prompt.contains(INPUT_PLACEHOLDER) } pub fn complete_prompt_args(&mut self, name: &str) { @@ -134,9 +173,9 @@ async function timeout(ms) { pub fn echo_messages(&self, input: &Input) -> String { let input_markdown = input.render(); - if self.empty_prompt() { + if self.is_empty_prompt() { input_markdown - } else if self.embedded_prompt() { + } else if self.is_embedded_prompt() { self.prompt.replace(INPUT_PLACEHOLDER, &input_markdown) } else { format!("{}\n\n{}", self.prompt, input.render()) @@ -145,9 +184,9 @@ async function timeout(ms) { pub fn build_messages(&self, input: &Input) -> Vec<Message> { let mut content = input.message_content(); - if self.empty_prompt() { + if self.is_empty_prompt() { vec![Message::new(MessageRole::User, content)] - } else if self.embedded_prompt() { + } else if self.is_embedded_prompt() { content.merge_prompt(|v: &str| self.prompt.replace(INPUT_PLACEHOLDER, v)); vec![Message::new(MessageRole::User, content)] } else { @@ -173,6 +212,45 @@ async function timeout(ms) { } } +impl RoleLike for Role { + fn to_role(&self) -> Role { + self.clone() + } + + fn model(&self) -> &Model { + &self.model + } + + fn temperature(&self) -> Option<f64> { + self.temperature + } + + fn top_p(&self) -> Option<f64> { + self.top_p + } + + fn function_matcher(&self) -> Option<String> { + self.function_matcher.clone() + } + + fn set_model(&mut self, model: &Model) { + self.model_id = Some(model.id()); + self.model = model.clone(); + } + + fn set_temperature(&mut self, value: Option<f64>) { + self.temperature = value; + } + + fn set_top_p(&mut self, value: Option<f64>) { + self.top_p = value; + } + + fn set_function_matcher(&mut self, matcher: Option<String>) { + self.function_matcher = matcher; + } +} + fn complete_prompt_args(prompt: &str, name: &str) -> String { let mut prompt = prompt.trim().to_string(); for (i, arg) in name.split(':').skip(1).enumerate() { diff --git a/src/config/session.rs b/src/config/session.rs index 908cc5b..97843bd 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -1,19 +1,17 @@ -use super::input::resolve_data_url; -use super::{Config, Input, Model, Role}; +use super::input::*; +use super::*; use crate::client::{Message, MessageContent, MessageRole}; use crate::render::MarkdownRender; use anyhow::{bail, Context, Result}; -use inquire::{Confirm, Text}; +use inquire::{required, Confirm, Text}; use serde::{Deserialize, Serialize}; use serde_json::json; use std::collections::HashMap; use std::fs::{self, create_dir_all, read_to_string}; use std::path::Path; -pub const TEMP_SESSION_NAME: &str = "temp"; - #[derive(Debug, Clone, Default, Deserialize, Serialize)] pub struct Session { #[serde(rename(serialize = "model", deserialize = "model"))] @@ -26,68 +24,65 @@ pub struct Session { function_matcher: Option<String>, #[serde(skip_serializing_if = "Option::is_none")] save_session: Option<bool>, + #[serde(skip_serializing_if = "Option::is_none")] + compress_threshold: Option<usize>, + messages: Vec<Message>, #[serde(default, skip_serializing_if = "HashMap::is_empty")] data_urls: HashMap<String, String>, #[serde(default, skip_serializing_if = "Vec::is_empty")] compressed_messages: Vec<Message>, - #[serde(skip_serializing_if = "Option::is_none")] - compress_threshold: Option<usize>, + + #[serde(skip)] + model: Model, #[serde(skip)] - pub name: String, + role_prompt: String, #[serde(skip)] - pub path: Option<String>, + role_name: String, #[serde(skip)] - pub dirty: bool, + name: String, #[serde(skip)] - pub compressing: bool, + path: Option<String>, #[serde(skip)] - pub model: Model, + dirty: bool, + #[serde(skip)] + compressing: bool, } impl Session { pub fn new(config: &Config, name: &str) -> Self { - let name = if name.is_empty() { - TEMP_SESSION_NAME - } else { - name - }; let save_session = if name == TEMP_SESSION_NAME { None } else { config.save_session }; + let role = config.extract_role(); let mut session = Self { - model_id: config.model.id(), - temperature: config.temperature, - top_p: config.top_p, - function_matcher: None, - save_session, - messages: Default::default(), - compressed_messages: Default::default(), - compress_threshold: None, - data_urls: Default::default(), name: name.to_string(), - path: None, - dirty: false, - compressing: false, - model: config.model.clone(), + save_session, + ..Default::default() }; - if let Some(role) = &config.role { - session.set_role_properties(role); - } + session.set_role(role); + session.dirty = false; session } - pub fn load(name: &str, path: &Path) -> Result<Self> { + pub fn load(config: &Config, name: &str, path: &Path) -> Result<Self> { let content = read_to_string(path) .with_context(|| format!("Failed to load session {} at {}", name, path.display()))?; let mut session: Self = serde_yaml::from_str(&content).with_context(|| format!("Invalid session {}", name))?; + session.model = Model::retrieve(config, &session.model_id)?; session.name = name.to_string(); session.path = Some(path.display().to_string()); + if let Some(bot) = &config.bot { + session + .role_prompt + .clone_from(&bot.definition().instructions); + } + Ok(session) } @@ -95,20 +90,12 @@ impl Session { &self.name } - pub fn model_id(&self) -> &str { - &self.model_id + pub fn dirty(&self) -> bool { + self.dirty } - pub fn temperature(&self) -> Option<f64> { - self.temperature - } - - pub fn top_p(&self) -> Option<f64> { - self.top_p - } - - pub fn function_matcher(&self) -> Option<&str> { - self.function_matcher.as_deref() + pub fn compressing(&self) -> bool { + self.compressing } pub fn save_session(&self) -> Option<bool> { @@ -123,7 +110,7 @@ impl Session { } pub fn tokens(&self) -> usize { - self.model.total_tokens(&self.messages) + self.model().total_tokens(&self.messages) } pub fn user_messages_len(&self) -> usize { @@ -134,10 +121,9 @@ impl Session { if self.path.is_none() { bail!("Not found session '{}'", self.name) } - let (tokens, percent) = self.tokens_and_percent(); let mut data = json!({ "path": self.path, - "model": self.model_id(), + "model": self.model().id(), }); if let Some(temperature) = self.temperature() { data["temperature"] = temperature.into(); @@ -151,8 +137,9 @@ impl Session { if let Some(save_session) = self.save_session() { data["save_session"] = save_session.into(); } + let (tokens, percent) = self.tokens_usage(); data["total_tokens"] = tokens.into(); - if let Some(max_input_tokens) = self.model.max_input_tokens() { + if let Some(max_input_tokens) = self.model().max_input_tokens() { data["max_input_tokens"] = max_input_tokens.into(); } if percent != 0.0 { @@ -165,14 +152,14 @@ impl Session { Ok(output) } - pub fn info(&self, render: &mut MarkdownRender) -> Result<String> { + pub fn render(&self, render: &mut MarkdownRender) -> Result<String> { let mut items = vec![]; if let Some(path) = &self.path { items.push(("path", path.to_string())); } - items.push(("model", self.model.id())); + items.push(("model", self.model().id())); if let Some(temperature) = self.temperature() { items.push(("temperature", temperature.to_string())); @@ -182,7 +169,7 @@ impl Session { } if let Some(function_matcher) = self.function_matcher() { - items.push(("function_matcher", function_matcher.into())); + items.push(("function_matcher", function_matcher)); } if let Some(save_session) = self.save_session() { @@ -193,7 +180,7 @@ impl Session { items.push(("compress_threshold", compress_threshold.to_string())); } - if let Some(max_input_tokens) = self.model.max_input_tokens() { + if let Some(max_input_tokens) = self.model().max_input_tokens() { items.push(("max_input_tokens", max_input_tokens.to_string())); } @@ -228,13 +215,17 @@ impl Session { } } + if lines.last() == Some(&String::new()) { + lines.pop(); + } + let output = lines.join("\n"); Ok(output) } - pub fn tokens_and_percent(&self) -> (usize, f32) { + pub fn tokens_usage(&self) -> (usize, f32) { let tokens = self.tokens(); - let max_input_tokens = self.model.max_input_tokens().unwrap_or_default(); + let max_input_tokens = self.model().max_input_tokens().unwrap_or_default(); let percent = if max_input_tokens == 0 { 0.0 } else { @@ -244,28 +235,24 @@ impl Session { (tokens, percent) } - pub fn set_temperature(&mut self, value: Option<f64>) { - if self.temperature != value { - self.temperature = value; - self.dirty = true; - } - } - - pub fn set_top_p(&mut self, value: Option<f64>) { - if self.top_p != value { - self.top_p = value; - self.dirty = true; - } + pub fn set_name(&mut self, name: &str) { + self.name = name.to_string(); } - pub fn set_function_matcher(&mut self, function_matcher: Option<&str>) { - self.function_matcher = function_matcher.map(|v| v.to_string()); + pub fn set_role(&mut self, role: Role) { + self.model_id = role.model().id(); + self.temperature = role.temperature(); + self.top_p = role.top_p(); + self.function_matcher = role.function_matcher().map(|v| v.to_string()); + self.model = role.model().clone(); + self.role_name = role.name().to_string(); + self.role_prompt = role.prompt().to_string(); + self.dirty = true; } - pub fn set_role_properties(&mut self, role: &Role) { - self.set_temperature(role.temperature); - self.set_top_p(role.top_p); - self.set_function_matcher(role.function_matcher.as_deref()); + pub fn clear_role(&mut self) { + self.role_name.clear(); + self.role_prompt.clear(); } pub fn set_save_session(&mut self, value: Option<bool>) { @@ -285,13 +272,8 @@ impl Session { } } - pub fn set_model(&mut self, model: &Model) { - let model_id = model.id(); - if self.model_id != model_id { - self.model_id = model_id; - self.dirty = true; - } - self.model = model.clone(); + pub fn set_compressing(&mut self, compressing: bool) { + self.compressing = compressing; } pub fn compress(&mut self, prompt: String) { @@ -314,8 +296,10 @@ impl Session { if !ans { return Ok(()); } - while self.is_temp() { - self.name = Text::new("Session name:").prompt()?; + if self.is_temp() { + self.name = Text::new("Session name:") + .with_validator(required!("This field is required")) + .prompt()?; } } self.save(sessions_dir)?; @@ -344,6 +328,8 @@ impl Session { ) })?; + println!("✨ Saved session to '{}'", session_path.display()); + self.dirty = false; Ok(()) @@ -367,10 +353,8 @@ impl Session { pub fn add_message(&mut self, input: &Input, output: &str) -> Result<()> { let mut need_add_msg = true; if self.messages.is_empty() { - if let Some(role) = input.role() { - self.messages.extend(role.build_messages(input)); - need_add_msg = false; - } + self.messages.extend(input.role().build_messages(input)); + need_add_msg = false; } if need_add_msg { self.messages @@ -402,10 +386,8 @@ impl Session { let mut need_add_msg = true; let len = messages.len(); if len == 0 { - if let Some(role) = input.role() { - messages = role.build_messages(input); - need_add_msg = false; - } + messages = input.role().build_messages(input); + need_add_msg = false; } else if len == 1 && self.compressed_messages.len() >= 2 { messages .extend(self.compressed_messages[self.compressed_messages.len() - 2..].to_vec()); @@ -416,3 +398,56 @@ impl Session { messages } } + +impl RoleLike for Session { + fn to_role(&self) -> Role { + let mut role = Role::new(&self.role_name, &self.role_prompt); + role.sync(self); + role + } + + fn model(&self) -> &Model { + &self.model + } + + fn temperature(&self) -> Option<f64> { + self.temperature + } + + fn top_p(&self) -> Option<f64> { + self.top_p + } + + fn function_matcher(&self) -> Option<String> { + self.function_matcher.clone() + } + + fn set_model(&mut self, model: &Model) { + if self.model().id() != model.id() { + self.model_id = model.id(); + self.model = model.clone(); + self.dirty = true; + } + } + + fn set_temperature(&mut self, value: Option<f64>) { + if self.temperature != value { + self.temperature = value; + self.dirty = true; + } + } + + fn set_top_p(&mut self, value: Option<f64>) { + if self.top_p != value { + self.top_p = value; + self.dirty = true; + } + } + + fn set_function_matcher(&mut self, value: Option<String>) { + if self.function_matcher != value { + self.function_matcher = value; + self.dirty = true; + } + } +} diff --git a/src/function.rs b/src/function.rs index 29bb2c3..0896540 100644 --- a/src/function.rs +++ b/src/function.rs @@ -1,5 +1,5 @@ use crate::{ - config::GlobalConfig, + config::{Config, GlobalConfig}, utils::{ dimmed_text, get_env_bool, indent_text, run_command, run_command_with_output, warning_text, IS_STDOUT_TERMINAL, @@ -10,24 +10,15 @@ use anyhow::{anyhow, bail, Context, Result}; use fancy_regex::Regex; use indexmap::{IndexMap, IndexSet}; use inquire::{validator::Validation, Text}; -use lazy_static::lazy_static; use serde::{Deserialize, Serialize}; use serde_json::{json, Value}; use std::{ collections::{HashMap, HashSet}, fs, path::Path, - sync::mpsc::channel, }; -use threadpool::ThreadPool; - -const BIN_DIR_NAME: &str = "bin"; -const DECLARATIONS_FILE_PATH: &str = "functions.json"; - -lazy_static! { - static ref THREAD_POOL: ThreadPool = ThreadPool::new(num_cpus::get()); -} +pub const FUNCTION_ALL_MATCHER: &str = ".*"; pub type ToolResults = (Vec<ToolCallResult>, String); pub fn eval_tool_calls( @@ -39,28 +30,9 @@ pub fn eval_tool_calls( return Ok(output); } calls = ToolCall::dedup(calls); - let parallel = calls.len() > 1 && calls.iter().all(|v| !v.is_execute()); - if parallel { - let (tx, rx) = channel(); - let calls_len = calls.len(); - for (index, call) in calls.into_iter().enumerate() { - let tx = tx.clone(); - let config = config.clone(); - THREAD_POOL.execute(move || { - let result = call.eval(&config); - let _ = tx.send((index, call, result)); - }); - } - let mut list: Vec<(usize, ToolCall, Result<Value>)> = rx.iter().take(calls_len).collect(); - list.sort_by_key(|v| v.0); - for (_, call, result) in list { - output.push(ToolCallResult::new(call, result?)); - } - } else { - for call in calls { - let result = call.eval(config)?; - output.push(ToolCallResult::new(call, result)); - } + for call in calls { + let result = call.eval(config)?; + output.push(ToolCallResult::new(call, result)); } Ok(output) } @@ -82,33 +54,21 @@ impl ToolCallResult { } #[derive(Debug, Clone, Default)] -pub struct Function { +pub struct Functions { names: IndexSet<String>, declarations: Vec<FunctionDeclaration>, - #[cfg(windows)] - bin_dir: std::path::PathBuf, - env_path: Option<String>, } -impl Function { - pub fn init(functions_dir: &Path) -> Result<Self> { - let bin_dir = functions_dir.join(BIN_DIR_NAME); - let env_path = if bin_dir.exists() { - prepend_env_path(&bin_dir).ok() - } else { - None - }; - - let declarations_file = functions_dir.join(DECLARATIONS_FILE_PATH); - - let declarations: Vec<FunctionDeclaration> = if declarations_file.exists() { +impl Functions { + pub fn init(declarations_path: &Path) -> Result<Self> { + let declarations: Vec<FunctionDeclaration> = if declarations_path.exists() { let ctx = || { format!( "Failed to load function declarations at {}", - declarations_file.display() + declarations_path.display() ) }; - let content = fs::read_to_string(&declarations_file).with_context(ctx)?; + let content = fs::read_to_string(declarations_path).with_context(ctx)?; serde_json::from_str(&content).with_context(ctx)? } else { vec![] @@ -119,9 +79,6 @@ impl Function { Ok(Self { names: func_names, declarations, - #[cfg(windows)] - bin_dir, - env_path, }) } @@ -139,6 +96,14 @@ impl Function { Some(output) } } + + pub fn contains(&self, name: &str) -> bool { + self.names.contains(name) + } + + pub fn is_empty(&self) -> bool { + self.names.is_empty() + } } #[derive(Debug, Clone, Deserialize)] @@ -205,29 +170,54 @@ impl ToolCall { } pub fn eval(&self, config: &GlobalConfig) -> Result<Value> { - let name = self.name.clone(); - if !config.read().function.names.contains(&name) { - bail!("Unexpected call: {name} {}", self.arguments); - } - let arguments = if self.arguments.is_object() { + let function_name = self.name.clone(); + let (call_name, cmd_name, mut cmd_args) = match &config.read().bot { + Some(bot) => { + if !bot.functions().contains(&function_name) { + bail!( + "Unexpected call: {} {function_name} {}", + bot.name(), + self.arguments + ); + } + ( + format!("{}:{}", bot.name(), function_name), + bot.name().to_string(), + vec![function_name], + ) + } + None => { + if !config.read().functions.contains(&function_name) { + bail!("Unexpected call: {function_name} {}", self.arguments); + } + (function_name.clone(), function_name, vec![]) + } + }; + let json_data = if self.arguments.is_object() { self.arguments.clone() } else if let Some(arguments) = self.arguments.as_str() { - let args: Value = serde_json::from_str(arguments) - .map_err(|_| anyhow!("The {name} call has invalid arguments: {arguments}"))?; - args + let arguments: Value = serde_json::from_str(arguments).map_err(|_| { + anyhow!("The call '{call_name}' has invalid arguments: {arguments}") + })?; + arguments } else { - bail!("The {name} call has invalid arguments: {}", self.arguments); + bail!( + "The call '{call_name}' has invalid arguments: {}", + self.arguments + ); }; - let arguments = arguments.to_string(); - let prompt = format!("Call {name} '{arguments}'",); + cmd_args.push(json_data.to_string()); + let prompt = format!("Call {cmd_name} {}", cmd_args.join(" ")); let mut envs = HashMap::new(); - if let Some(env_path) = config.read().function.env_path.clone() { - envs.insert("PATH".into(), env_path); - }; + let bin_dir = Config::functions_bin_dir()?; + if bin_dir.exists() { + envs.insert("PATH".into(), prepend_env_path(&bin_dir)?); + } + #[cfg(windows)] - let name = polyfill_cmd_name(&name, &config.read().function.bin_dir); + let cmd_name = polyfill_cmd_name(&cmd_name, &bin_dir); let output = if self.is_execute() { if *IS_STDOUT_TERMINAL { @@ -243,13 +233,13 @@ impl ToolCall { .prompt()?; match answer.as_str() { "1" => { - let exit_code = run_command(&name, &[arguments], Some(envs))?; + let exit_code = run_command(&cmd_name, &cmd_args, Some(envs))?; if exit_code != 0 { bail!("Exit {exit_code}"); } Value::Null } - "2" => run_and_retrieve(&name, &arguments, envs, &prompt)?, + "2" => run_and_retrieve(&cmd_name, &cmd_args, envs, &prompt)?, _ => Value::Null, } } else { @@ -258,7 +248,7 @@ impl ToolCall { } } else { println!("{}", dimmed_text(&prompt)); - run_and_retrieve(&name, &arguments, envs, &prompt)? + run_and_retrieve(&cmd_name, &cmd_args, envs, &prompt)? }; Ok(output) @@ -274,12 +264,12 @@ impl ToolCall { } fn run_and_retrieve( - name: &str, - arguments: &str, + cmd_name: &str, + cmd_args: &[String], envs: HashMap<String, String>, prompt: &str, ) -> Result<Value> { - let (success, stdout, stderr) = run_command_with_output(name, &[arguments], Some(envs))?; + let (success, stdout, stderr) = run_command_with_output(cmd_name, cmd_args, Some(envs))?; if success { if !stderr.is_empty() { @@ -322,16 +312,16 @@ fn prepend_env_path(bin_dir: &Path) -> Result<String> { } #[cfg(windows)] -fn polyfill_cmd_name(name: &str, bin_dir: &std::path::Path) -> String { - let mut name = name.to_string(); +fn polyfill_cmd_name(cmd_name: &str, bin_dir: &std::path::Path) -> String { + let mut cmd_name = cmd_name.to_string(); if let Ok(exts) = std::env::var("PATHEXT") { if let Some(cmd_path) = exts .split(';') - .map(|ext| bin_dir.join(format!("{}{}", name, ext))) + .map(|ext| bin_dir.join(format!("{}{}", cmd_name, ext))) .find(|path| path.exists()) { - name = cmd_path.display().to_string(); + cmd_name = cmd_path.display().to_string(); } } - name + cmd_name } diff --git a/src/main.rs b/src/main.rs index e021f9e..459b50f 100644 --- a/src/main.rs +++ b/src/main.rs @@ -16,8 +16,7 @@ extern crate log; use crate::cli::Cli; use crate::client::{list_chat_models, send_stream, ChatCompletionsOutput}; use crate::config::{ - Config, GlobalConfig, Input, InputContext, WorkingMode, CODE_ROLE, EXPLAIN_SHELL_ROLE, - SHELL_ROLE, + Config, GlobalConfig, Input, WorkingMode, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE, }; use crate::function::{eval_tool_calls, need_send_call_results}; use crate::render::{render_error, MarkdownRender}; @@ -62,7 +61,7 @@ async fn main() -> Result<()> { .read() .roles .iter() - .for_each(|v| println!("{}", v.name)); + .for_each(|v| println!("{}", v.name())); return Ok(()); } if cli.list_models { @@ -99,9 +98,8 @@ async fn main() -> Result<()> { .write() .use_session(session.as_ref().map(|v| v.as_str()))?; } - if let Some(model) = &cli.model { - config.write().set_model(model)?; - config.write().set_model_id(); + if let Some(model_id) = &cli.model { + config.write().set_model(model_id)?; } if cli.save_session { config.write().set_save_session(Some(true)); @@ -241,7 +239,7 @@ async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) - } "📖 Explain" => { let role = config.read().retrieve_role(EXPLAIN_SHELL_ROLE)?; - let input = Input::from_str(config, &eval_str, Some(InputContext::role(role))); + let input = Input::from_str(config, &eval_str, Some(role)); let abort = create_abort_signal(); send_stream(&input, client.as_ref(), config, abort).await?; continue; diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 5ab9de5..029c59e 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -20,7 +20,6 @@ use std::fmt::Debug; use std::{io::BufReader, path::Path}; use tokio::sync::mpsc; -pub const TEMP_RAG_NAME: &str = "temp"; pub const CHUNK_OVERLAP: usize = 20; pub const SIMILARITY_THRESHOLD: f32 = 0.25; @@ -48,7 +47,8 @@ impl Rag { pub async fn init( config: &GlobalConfig, name: &str, - path: &Path, + save_path: &Path, + doc_paths: &[String], abort_signal: AbortSignal, ) -> Result<Self> { debug!("init rag: {name}"); @@ -56,9 +56,12 @@ impl Rag { let chunk_size = model.default_chunk_size(); let chunk_size = set_chunk_size(chunk_size)?; let data = RagData::new(&model.id(), chunk_size); - let mut rag = Self::create(config, name, path, data)?; - let paths = add_document_paths()?; - debug!("document paths: {paths:?}"); + let mut rag = Self::create(config, name, save_path, data)?; + let mut paths = doc_paths.to_vec(); + if paths.is_empty() { + paths = add_doc_paths()?; + }; + debug!("doc paths: {paths:?}"); let (stop_spinner_tx, set_spinner_message_tx) = run_spinner("Starting").await; tokio::select! { ret = rag.add_paths(&paths, Some(set_spinner_message_tx)) => { @@ -71,8 +74,8 @@ impl Rag { }, }; if !rag.is_temp() { - rag.save(path)?; - println!("✨ Saved rag to '{}'", path.display()); + rag.save(save_path)?; + println!("✨ Saved rag to '{}'", save_path.display()); } Ok(rag) } @@ -408,7 +411,7 @@ fn set_chunk_size(chunk_size: usize) -> Result<usize> { value.parse().map_err(|_| anyhow!("Invalid chunk_size")) } -fn add_document_paths() -> Result<Vec<String>> { +fn add_doc_paths() -> Result<Vec<String>> { let text = Text::new("Add document paths:") .with_validator(required!("This field is required")) .with_help_message("e.g. file1;dir2/;dir3/**/*.md") diff --git a/src/repl/completer.rs b/src/repl/completer.rs index d026d67..1d68792 100644 --- a/src/repl/completer.rs +++ b/src/repl/completer.rs @@ -49,13 +49,12 @@ impl Completer for ReplCompleter { if parts_len > 1 { let span = Span::new(parts[parts_len - 1].1, pos); let args: Vec<&str> = parts.iter().skip(1).map(|(v, _)| *v).collect(); - suggestions.extend( - self.config - .read() - .repl_complete(cmd, &args) - .iter() - .map(|(value, description)| create_suggestion(value, description, span)), - ) + suggestions.extend(self.config.read().repl_complete(cmd, &args).iter().map( + |(value, description)| { + let description = description.as_deref().unwrap_or_default(); + create_suggestion(value, description, span) + }, + )) } if suggestions.is_empty() { diff --git a/src/repl/mod.rs b/src/repl/mod.rs index bf5fd8e..dca5fc4 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::{AssertState, Config, GlobalConfig, Input, InputContext, StateFlags}; +use crate::config::{AssertState, Config, GlobalConfig, Input, StateFlags}; use crate::function::need_send_call_results; use crate::render::render_error; use crate::utils::{create_abort_signal, set_text, AbortSignal}; @@ -33,19 +33,19 @@ lazy_static! { const MENU_NAME: &str = "completion_menu"; lazy_static! { - static ref REPL_COMMANDS: [ReplCommand; 19] = [ - 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()), + static ref REPL_COMMANDS: [ReplCommand; 22] = [ + ReplCommand::new(".help", "Show this help message", AssertState::pass()), + ReplCommand::new(".info", "View system info", AssertState::pass()), + ReplCommand::new(".model", "Change the current LLM", AssertState::pass()), ReplCommand::new( ".prompt", "Create a temporary role using a prompt", - AssertState::False(StateFlags::SESSION) + AssertState::False(StateFlags::SESSION | StateFlags::BOT) ), ReplCommand::new( ".role", "Switch to a specific role", - AssertState::False(StateFlags::SESSION) + AssertState::False(StateFlags::SESSION | StateFlags::BOT) ), ReplCommand::new( ".info role", @@ -82,7 +82,11 @@ lazy_static! { "End the current session", AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION) ), - ReplCommand::new(".rag", "Init or use a rag", AssertState::any()), + ReplCommand::new( + ".rag", + "Init or use a rag", + AssertState::False(StateFlags::BOT) + ), ReplCommand::new( ".info rag", "View rag info", @@ -91,16 +95,27 @@ lazy_static! { ReplCommand::new( ".exit rag", "Leave the rag", - AssertState::True(StateFlags::RAG) + AssertState::TrueFalse(StateFlags::RAG, StateFlags::BOT), + ), + ReplCommand::new(".bot", "Use a bot", AssertState::bare()), + ReplCommand::new( + ".info bot", + "View bot info", + AssertState::True(StateFlags::BOT), + ), + ReplCommand::new( + ".exit bot", + "Leave the bot", + AssertState::True(StateFlags::BOT) ), ReplCommand::new( ".file", "Include files with the message", - AssertState::any() + AssertState::pass() ), - ReplCommand::new(".set", "Adjust settings", AssertState::any()), - ReplCommand::new(".copy", "Copy the last response", AssertState::any()), - ReplCommand::new(".exit", "Exit the REPL", AssertState::any()), + ReplCommand::new(".set", "Adjust settings", AssertState::pass()), + ReplCommand::new(".copy", "Copy the last response", AssertState::pass()), + ReplCommand::new(".exit", "Exit the REPL", AssertState::pass()), ]; 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(); @@ -191,18 +206,19 @@ impl Repl { let info = self.config.read().rag_info()?; println!("{}", info); } + Some("bot") => { + let info = self.config.read().bot_info()?; + println!("{}", info); + } Some(_) => unknown_command()?, None => { - let output = self.config.read().system_info()?; + let output = self.config.read().sysinfo()?; println!("{}", output); } }, ".model" => match args { Some(name) => { self.config.write().set_model(name)?; - if !self.config.read().has_role_or_session() { - self.config.write().set_model_id(); - } } None => println!("Usage: .model <name>"), }, @@ -216,11 +232,7 @@ impl Repl { Some(args) => match args.split_once(|c| c == '\n' || c == ' ') { Some((name, text)) => { let role = self.config.read().retrieve_role(name.trim())?; - let input = Input::from_str( - &self.config, - text.trim(), - Some(InputContext::role(role)), - ); + let input = Input::from_str(&self.config, text.trim(), Some(role)); ask(&self.config, self.abort_signal.clone(), input).await?; } None => { @@ -235,6 +247,12 @@ impl Repl { ".rag" => { Config::use_rag(&self.config, args, self.abort_signal.clone()).await?; } + ".bot" => match args { + Some(name) => { + Config::use_bot(&self.config, name, self.abort_signal.clone()).await?; + } + None => println!(r#"Usage: .bot <name>"#), + }, ".save" => { match args.map(|v| match v.split_once(' ') { Some((subcmd, args)) => (subcmd, args.trim()), @@ -280,6 +298,9 @@ impl Repl { Some("rag") => { self.config.write().exit_rag()?; } + Some("bot") => { + self.config.write().exit_bot()?; + } Some(_) => unknown_command()?, None => { return Ok(true); @@ -408,8 +429,13 @@ impl ReplCommand { 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(), + AssertState::True(true_flags) => true_flags & flags != StateFlags::empty(), + AssertState::False(false_flags) => false_flags & flags == StateFlags::empty(), + AssertState::TrueFalse(true_flags, false_flags) => { + (true_flags & flags != StateFlags::empty()) + && (false_flags & flags == StateFlags::empty()) + } + AssertState::Equal(check_flags) => check_flags == flags, } } } |
