summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-09-14 18:05:29 +0800
committerGitHub <noreply@github.com>2024-09-14 18:05:29 +0800
commit6211d01a648e941fc69954d0855bcdcef98f27b9 (patch)
tree88d3a44b8e6906bbcb39917ea72e5cc3aa5087d6
parent5a26c59a12e25cc96c6eb44c64395e673207135c (diff)
downloadaichat-6211d01a648e941fc69954d0855bcdcef98f27b9.tar.gz
feat: add `.save agent-config` repl command (#870)
-rw-r--r--config.example.yaml14
-rw-r--r--src/config/agent.rs35
-rw-r--r--src/config/mod.rs127
-rw-r--r--src/main.rs2
-rw-r--r--src/repl/mod.rs48
5 files changed, 135 insertions, 91 deletions
diff --git a/config.example.yaml b/config.example.yaml
index a1caea2..3b08856 100644
--- a/config.example.yaml
+++ b/config.example.yaml
@@ -11,6 +11,13 @@ editor: null # Specifies the command used to edit input buff
wrap: no # Controls text wrapping (no, auto, <max-width>)
wrap_code: false # Enables or disables wrapping of code blocks
+# ---- function-calling ----
+# Visit https://github.com/sigoden/llm-functions for setup instructions
+function_calling: true # Enables or disables function calling (Globally).
+mapping_tools: # Alias for a tool or toolset
+ fs: 'fs_cat,fs_ls,fs_mkdir,fs_rm,fs_write'
+use_tools: null # Which tools to use by default. (e.g. 'fs,web_search')
+
# ---- prelude ----
prelude: null # Set a default role or session to start with (e.g. role:<name>, session:<name>)
repl_prelude: null # Overrides the `prelude` setting specifically for conversations started in REPL
@@ -26,13 +33,6 @@ summarize_prompt: 'Summarize the discussion briefly in 200 words or less to use
# Text prompt used for including the summary of the entire session
summary_prompt: 'This is a summary of the chat history as a recap: '
-# ---- function-calling ----
-# Visit https://github.com/sigoden/llm-functions for setup instructions
-function_calling: true # Enables or disables function calling (Globally).
-mapping_tools: # Alias for a tool or toolset
- fs: 'fs_cat,fs_ls,fs_mkdir,fs_rm,fs_write'
-use_tools: null # Which tools to use by default. (e.g. 'fs,web_search')
-
# ---- RAG ----
# See [RAG-Guide](https://github.com/sigoden/aichat/wiki/RAG-Guide) for more details.
rag_embedding_model: null # Specifies the embedding model to use
diff --git a/src/config/agent.rs b/src/config/agent.rs
index 36af15e..813df3a 100644
--- a/src/config/agent.rs
+++ b/src/config/agent.rs
@@ -39,7 +39,7 @@ impl Agent {
let agent_config = if config_path.exists() {
AgentConfig::load(&config_path)?
} else {
- AgentConfig::default()
+ AgentConfig::new(&config.read())
};
let mut definition = AgentDefinition::load(&definition_file_path)?;
init_variables(&variables_path, &mut definition.variables)
@@ -91,6 +91,18 @@ impl Agent {
})
}
+ pub fn save_config(&self) -> Result<()> {
+ let config_path = Config::agent_config_file(&self.name)?;
+ ensure_parent_exists(&config_path)?;
+ let content = serde_yaml::to_string(&self.config)?;
+ fs::write(&config_path, content).with_context(|| {
+ format!("Failed to save agent config to '{}'", config_path.display())
+ })?;
+
+ println!("✨ Saved agent config to '{}'", config_path.display());
+ Ok(())
+ }
+
pub fn export(&self) -> Result<String> {
let mut agent = self.clone();
agent.definition.instructions = self.interpolated_instructions();
@@ -143,6 +155,10 @@ impl Agent {
self.config.agent_prelude.as_deref()
}
+ pub fn set_agent_prelude(&mut self, value: Option<String>) {
+ self.config.agent_prelude = value;
+ }
+
pub fn variables(&self) -> &[AgentVariable] {
&self.definition.variables
}
@@ -208,22 +224,23 @@ impl RoleLike for Agent {
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
pub struct AgentConfig {
- #[serde(
- rename(serialize = "model", deserialize = "model"),
- skip_serializing_if = "Option::is_none"
- )]
+ #[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>,
- #[serde(skip_serializing_if = "Option::is_none")]
pub use_tools: Option<String>,
- #[serde(skip_serializing_if = "Option::is_none")]
pub agent_prelude: Option<String>,
}
impl AgentConfig {
+ pub fn new(config: &Config) -> Self {
+ Self {
+ use_tools: config.use_tools.clone(),
+ agent_prelude: config.agent_prelude.clone(),
+ ..Default::default()
+ }
+ }
+
pub fn load(path: &Path) -> Result<Self> {
let contents = read_to_string(path)
.with_context(|| format!("Failed to read agent config file at '{}'", path.display()))?;
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 11c6d5c..4f7cfd7 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -99,6 +99,10 @@ pub struct Config {
pub wrap: Option<String>,
pub wrap_code: bool,
+ pub function_calling: bool,
+ pub mapping_tools: IndexMap<String, String>,
+ pub use_tools: Option<String>,
+
pub prelude: Option<String>,
pub repl_prelude: Option<String>,
pub agent_prelude: Option<String>,
@@ -108,10 +112,6 @@ pub struct Config {
pub summarize_prompt: Option<String>,
pub summary_prompt: Option<String>,
- pub function_calling: bool,
- pub mapping_tools: IndexMap<String, String>,
- pub use_tools: Option<String>,
-
pub rag_embedding_model: Option<String>,
pub rag_reranker_model: Option<String>,
pub rag_top_k: usize,
@@ -166,6 +166,10 @@ impl Default for Config {
wrap: None,
wrap_code: false,
+ function_calling: true,
+ mapping_tools: Default::default(),
+ use_tools: None,
+
prelude: None,
repl_prelude: None,
agent_prelude: None,
@@ -175,10 +179,6 @@ impl Default for Config {
summarize_prompt: None,
summary_prompt: None,
- function_calling: true,
- mapping_tools: Default::default(),
- use_tools: None,
-
rag_embedding_model: None,
rag_reranker_model: None,
rag_top_k: 4,
@@ -402,7 +402,7 @@ impl Config {
self.serve_addr.clone().unwrap_or_else(|| SERVE_ADDR.into())
}
- pub fn log(is_serve: bool) -> Result<(LevelFilter, Option<PathBuf>)> {
+ pub fn log_config(is_serve: bool) -> Result<(LevelFilter, Option<PathBuf>)> {
let log_level = env::var(get_env_name("log_level"))
.ok()
.and_then(|v| v.parse().ok())
@@ -513,10 +513,14 @@ impl Config {
.wrap
.clone()
.map_or_else(|| String::from("no"), |v| v.to_string());
- let (rag_reranker_model, rag_top_k) = match self.rag.as_ref() {
+ let (rag_reranker_model, rag_top_k) = match &self.rag {
Some(rag) => rag.get_config(),
None => (self.rag_reranker_model.clone(), self.rag_top_k),
};
+ let agent_prelude = match &self.agent {
+ Some(agent) => agent.agent_prelude(),
+ None => self.agent_prelude.as_deref(),
+ };
let role = self.extract_role();
let mut items = vec![
("model", role.model().id()),
@@ -535,10 +539,11 @@ impl Config {
("keybindings", self.keybindings.clone()),
("wrap", wrap),
("wrap_code", self.wrap_code.to_string()),
- ("save_session", format_option_value(&self.save_session)),
- ("compress_threshold", self.compress_threshold.to_string()),
("function_calling", self.function_calling.to_string()),
("use_tools", format_option_value(&role.use_tools())),
+ ("agent_prelude", format_option_value(&agent_prelude)),
+ ("save_session", format_option_value(&self.save_session)),
+ ("compress_threshold", self.compress_threshold.to_string()),
(
"rag_reranker_model",
format_option_value(&rag_reranker_model),
@@ -554,7 +559,7 @@ impl Config {
("functions_dir", display_path(&Self::functions_dir()?)),
("messages_file", display_path(&self.messages_file()?)),
];
- if let Ok((_, Some(log_path))) = Self::log(self.working_mode.is_serve()) {
+ if let Ok((_, Some(log_path))) = Self::log_config(self.working_mode.is_serve()) {
items.push(("log_path", display_path(&log_path)));
}
let output = items
@@ -597,14 +602,6 @@ impl Config {
let value = value.parse().with_context(|| "Invalid value")?;
config.write().save = value;
}
- "rag_reranker_model" => {
- let value = parse_value(value)?;
- Self::set_rag_reranker_model(config, value)?;
- }
- "rag_top_k" => {
- let value = value.parse().with_context(|| "Invalid value")?;
- Self::set_rag_top_k(config, value)?;
- }
"function_calling" => {
let value = value.parse().with_context(|| "Invalid value")?;
if value && config.write().functions.is_empty() {
@@ -616,6 +613,10 @@ impl Config {
let value = parse_value(value)?;
config.write().set_use_tools(value);
}
+ "agent_prelude" => {
+ let value = parse_value(value)?;
+ config.write().set_agent_prelude(value);
+ }
"save_session" => {
let value = parse_value(value)?;
config.write().set_save_session(value);
@@ -624,6 +625,14 @@ impl Config {
let value = parse_value(value)?;
config.write().set_compress_threshold(value);
}
+ "rag_reranker_model" => {
+ let value = parse_value(value)?;
+ Self::set_rag_reranker_model(config, value)?;
+ }
+ "rag_top_k" => {
+ let value = value.parse().with_context(|| "Invalid value")?;
+ Self::set_rag_top_k(config, value)?;
+ }
"highlight" => {
let value = value.parse().with_context(|| "Invalid value")?;
config.write().highlight = value;
@@ -638,7 +647,7 @@ impl Config {
"roles" => (Self::roles_dir()?, Some(".md")),
"sessions" => (config.read().sessions_dir()?, Some(".yaml")),
"rags" => (Self::rags_dir()?, Some(".yaml")),
- "agents-config" => (Self::agents_config_dir()?, None),
+ "agents" => (Self::agents_config_dir()?, None),
_ => bail!("Unknown kind '{kind}'"),
};
let names = match read_dir(&dir) {
@@ -722,6 +731,13 @@ impl Config {
}
}
+ pub fn set_agent_prelude(&mut self, value: Option<String>) {
+ match self.agent.as_mut() {
+ Some(agent) => agent.set_agent_prelude(value),
+ None => self.agent_prelude = value,
+ }
+ }
+
pub fn set_save_session(&mut self, value: Option<bool>) {
if let Some(session) = self.session.as_mut() {
session.set_save_session(value);
@@ -1269,13 +1285,9 @@ impl Config {
bail!("Already in a agent, please run '.exit agent' first to exit the current agent.");
}
let agent = Agent::init(config, name, abort_signal).await?;
- let session = session.map(|v| v.to_string()).or_else(|| {
- agent
- .agent_prelude()
- .map(|v| v.to_string())
- .or_else(|| config.read().agent_prelude.clone())
- .and_then(|v| if v.is_empty() { None } else { Some(v) })
- });
+ let session = session
+ .map(|v| v.to_string())
+ .or_else(|| agent.agent_prelude().map(|v| v.to_string()));
config.write().rag = agent.rag();
config.write().agent = Some(agent);
if let Some(session) = session {
@@ -1314,6 +1326,14 @@ impl Config {
Ok(())
}
+ pub fn save_agent_config(&mut self) -> Result<()> {
+ let agent = match &self.agent {
+ Some(v) => v,
+ None => bail!("No agent"),
+ };
+ agent.save_config()
+ }
+
pub fn exit_agent(&mut self) -> Result<()> {
self.exit_session()?;
if self.agent.take().is_some() {
@@ -1472,10 +1492,11 @@ impl Config {
"dry_run",
"stream",
"save",
- "save_session",
- "compress_threshold",
"function_calling",
"use_tools",
+ "agent_prelude",
+ "save_session",
+ "compress_threshold",
"rag_reranker_model",
"rag_top_k",
"highlight",
@@ -1486,9 +1507,7 @@ impl Config {
.map(|v| (format!("{v} "), None))
.collect()
}
- ".delete" => {
- map_completion_values(vec!["roles", "sessions", "rags", "agents-config"])
- }
+ ".delete" => map_completion_values(vec!["roles", "sessions", "rags", "agents"]),
_ => vec![],
};
filter = args[0]
@@ -1501,14 +1520,6 @@ impl Config {
"dry_run" => complete_bool(self.dry_run),
"stream" => complete_bool(self.stream),
"save" => complete_bool(self.save),
- "save_session" => {
- let save_session = if let Some(session) = &self.session {
- session.save_session()
- } else {
- self.save_session
- };
- complete_option_bool(save_session)
- }
"function_calling" => complete_bool(self.function_calling),
"use_tools" => {
let mut prefix = String::new();
@@ -1529,6 +1540,14 @@ impl Config {
.map(|v| format!("{prefix}{v}"))
.collect()
}
+ "save_session" => {
+ let save_session = if let Some(session) = &self.session {
+ session.save_session()
+ } else {
+ self.save_session
+ };
+ complete_option_bool(save_session)
+ }
"rag_reranker_model" => list_reranker_models(self).iter().map(|v| v.id()).collect(),
"highlight" => complete_bool(self.highlight),
_ => vec![],
@@ -1840,6 +1859,18 @@ impl Config {
self.wrap_code = v;
}
+ if let Some(Some(v)) = read_env_bool("function_calling") {
+ self.function_calling = v;
+ }
+ if let Ok(v) = env::var(get_env_name("mapping_tools")) {
+ if let Ok(v) = serde_json::from_str(&v) {
+ self.mapping_tools = v;
+ }
+ }
+ if let Some(v) = read_env_value::<String>("use_tools") {
+ self.use_tools = v;
+ }
+
if let Some(v) = read_env_value::<String>("prelude") {
self.prelude = v;
}
@@ -1863,18 +1894,6 @@ impl Config {
self.summary_prompt = v;
}
- if let Some(Some(v)) = read_env_bool("function_calling") {
- self.function_calling = v;
- }
- if let Ok(v) = env::var(get_env_name("mapping_tools")) {
- if let Ok(v) = serde_json::from_str(&v) {
- self.mapping_tools = v;
- }
- }
- if let Some(v) = read_env_value::<String>("use_tools") {
- self.use_tools = v;
- }
-
if let Some(v) = read_env_value::<String>("rag_embedding_model") {
self.rag_embedding_model = v;
}
diff --git a/src/main.rs b/src/main.rs
index b4b76b6..45d2430 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -317,7 +317,7 @@ async fn create_input(
}
fn setup_logger(is_serve: bool) -> Result<()> {
- let (log_level, log_path) = Config::log(is_serve)?;
+ let (log_level, log_path) = Config::log_config(is_serve)?;
if log_level == LevelFilter::Off {
return Ok(());
}
diff --git a/src/repl/mod.rs b/src/repl/mod.rs
index 2f6c76f..09fa74d 100644
--- a/src/repl/mod.rs
+++ b/src/repl/mod.rs
@@ -31,7 +31,7 @@ lazy_static::lazy_static! {
const MENU_NAME: &str = "completion_menu";
lazy_static::lazy_static! {
- static ref REPL_COMMANDS: [ReplCommand; 31] = [
+ static ref REPL_COMMANDS: [ReplCommand; 32] = [
ReplCommand::new(".help", "Show this help message", AssertState::pass()),
ReplCommand::new(".info", "View system info", AssertState::pass()),
ReplCommand::new(".model", "Change the current LLM", AssertState::pass()),
@@ -58,7 +58,7 @@ lazy_static::lazy_static! {
ReplCommand::new(
".save role",
"Save the current role to file",
- AssertState::True(StateFlags::ROLE)
+ AssertState::TrueFalse(StateFlags::ROLE, StateFlags::SESSION_EMPTY | StateFlags::SESSION),
),
ReplCommand::new(
".exit role",
@@ -71,6 +71,11 @@ lazy_static::lazy_static! {
AssertState::False(StateFlags::SESSION_EMPTY | StateFlags::SESSION),
),
ReplCommand::new(
+ ".clear messages",
+ "Erase messages in the current session",
+ AssertState::True(StateFlags::SESSION)
+ ),
+ ReplCommand::new(
".info session",
"View session info",
AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION),
@@ -81,11 +86,6 @@ lazy_static::lazy_static! {
AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION)
),
ReplCommand::new(
- ".clear messages",
- "Erase messages in the current session",
- AssertState::True(StateFlags::SESSION)
- ),
- ReplCommand::new(
".save session",
"Save the current session to file",
AssertState::True(StateFlags::SESSION_EMPTY | StateFlags::SESSION)
@@ -101,13 +101,13 @@ lazy_static::lazy_static! {
AssertState::False(StateFlags::AGENT)
),
ReplCommand::new(
- ".info rag",
- "View RAG info",
+ ".rebuild rag",
+ "Rebuild the RAG to sync document changes",
AssertState::True(StateFlags::RAG),
),
ReplCommand::new(
- ".rebuild rag",
- "Rebuild the RAG to sync document changes",
+ ".info rag",
+ "View RAG info",
AssertState::True(StateFlags::RAG),
),
ReplCommand::new(
@@ -117,11 +117,6 @@ lazy_static::lazy_static! {
),
ReplCommand::new(".agent", "Use a agent", AssertState::bare()),
ReplCommand::new(
- ".info agent",
- "View agent info",
- AssertState::True(StateFlags::AGENT),
- ),
- ReplCommand::new(
".starter",
"Use the conversation starter",
AssertState::True(StateFlags::AGENT)
@@ -132,6 +127,16 @@ lazy_static::lazy_static! {
AssertState::True(StateFlags::AGENT)
),
ReplCommand::new(
+ ".save agent-config",
+ "Save the current agent config to file",
+ AssertState::True(StateFlags::AGENT)
+ ),
+ ReplCommand::new(
+ ".info agent",
+ "View agent info",
+ AssertState::True(StateFlags::AGENT),
+ ),
+ ReplCommand::new(
".exit agent",
"Leave the agent",
AssertState::True(StateFlags::AGENT)
@@ -148,7 +153,7 @@ lazy_static::lazy_static! {
AssertState::pass()
),
ReplCommand::new(".set", "Adjust runtime configuration", AssertState::pass()),
- ReplCommand::new(".delete", "Delete roles/sessions/RAGs/agents-config", AssertState::pass()),
+ ReplCommand::new(".delete", "Delete roles/sessions/RAGs/agents", AssertState::pass()),
ReplCommand::new(".copy", "Copy the last response", AssertState::pass()),
ReplCommand::new(".exit", "Exit the REPL", AssertState::pass()),
];
@@ -326,8 +331,11 @@ impl Repl {
Some(("session", name)) => {
self.config.write().save_session(name)?;
}
+ Some(("agent-config", _)) => {
+ self.config.write().save_agent_config()?;
+ }
_ => {
- println!(r#"Usage: .save <role|session> [name]"#)
+ println!(r#"Usage: .save <role|session|aegnt-config> [name]"#)
}
}
}
@@ -343,7 +351,7 @@ impl Repl {
self.config.write().edit_session()?;
}
_ => {
- println!(r#"Usage: .edit session"#)
+ println!(r#"Usage: .edit <role|session>"#)
}
}
}
@@ -398,7 +406,7 @@ impl Repl {
Config::delete(&self.config, args)?;
}
_ => {
- println!("Usage: .delete [roles|sessions|rags|agents-config]")
+ println!("Usage: .delete [roles|sessions|rags|agents]")
}
},
".copy" => {