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/config | |
| 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/config')
| -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 |
5 files changed, 839 insertions, 358 deletions
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; + } + } +} |
