summaryrefslogtreecommitdiffstats
path: root/src/config
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
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')
-rw-r--r--src/config/bot.rs255
-rw-r--r--src/config/input.rs130
-rw-r--r--src/config/mod.rs441
-rw-r--r--src/config/role.rs154
-rw-r--r--src/config/session.rs217
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;
+ }
+ }
+}