summaryrefslogtreecommitdiffstats
path: root/src/config/mod.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-11 11:00:12 +0800
committerGitHub <noreply@github.com>2024-06-11 11:00:12 +0800
commitbb867c4fcbc6f42770471f0a0117cd8909f660d7 (patch)
tree6293f7f1108309160d1951f53f6429e9b004870d /src/config/mod.rs
parent5635ca6a58fb4a590419335b098b7317285bfb82 (diff)
downloadaichat-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/config/mod.rs')
-rw-r--r--src/config/mod.rs441
1 files changed, 304 insertions, 137 deletions
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<()> {