summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-11-05 17:16:01 +0800
committerGitHub <noreply@github.com>2024-11-05 17:16:01 +0800
commit42deaa082fe026b955f1bc11b5168a2309aad99b (patch)
tree46a3675c7f97848d5b51db984a0151cc53436c04 /src
parent0a324b6ddd4b89cc61a8dd15eeb06d60273308fa (diff)
downloadaichat-42deaa082fe026b955f1bc11b5168a2309aad99b.tar.gz
feat: support session-scoped agent variables (#969)
Diffstat (limited to 'src')
-rw-r--r--src/config/agent.rs147
-rw-r--r--src/config/mod.rs52
-rw-r--r--src/config/session.rs29
-rw-r--r--src/function.rs6
-rw-r--r--src/repl/mod.rs6
5 files changed, 167 insertions, 73 deletions
diff --git a/src/config/agent.rs b/src/config/agent.rs
index 39c913a..5d7278b 100644
--- a/src/config/agent.rs
+++ b/src/config/agent.rs
@@ -8,12 +8,18 @@ use std::{fs::read_to_string, path::Path};
use serde::{Deserialize, Serialize};
+const DEFAULT_AGENT_NAME: &str = "rag";
+
#[derive(Debug, Clone, Serialize)]
pub struct Agent {
name: String,
config: AgentConfig,
definition: AgentDefinition,
#[serde(skip)]
+ shared_variables: IndexMap<String, String>,
+ #[serde(skip)]
+ session_variables: Option<IndexMap<String, String>>,
+ #[serde(skip)]
functions: Functions,
#[serde(skip)]
rag: Option<Arc<Rag>>,
@@ -33,7 +39,7 @@ impl Agent {
bail!("Unknown agent `{name}`");
}
let functions_file_path = functions_dir.join("functions.json");
- let rag_path = Config::agent_rag_file(name, "rag")?;
+ let rag_path = Config::agent_rag_file(name, DEFAULT_AGENT_NAME)?;
let config_path = Config::agent_config_file(name)?;
let agent_config = if config_path.exists() {
AgentConfig::load(&config_path)?
@@ -41,9 +47,6 @@ impl Agent {
AgentConfig::new(&config.read())
};
let mut definition = AgentDefinition::load(&definition_file_path)?;
- init_variables(&mut definition.variables, &agent_config.variables)
- .context("Failed to init variables")?;
-
let functions = if functions_file_path.exists() {
Functions::init(&functions_file_path)?
} else {
@@ -60,7 +63,7 @@ impl Agent {
};
let rag = if rag_path.exists() {
- Some(Arc::new(Rag::load(config, "rag", &rag_path)?))
+ Some(Arc::new(Rag::load(config, DEFAULT_AGENT_NAME, &rag_path)?))
} else if !definition.documents.is_empty() {
let mut ans = false;
if *IS_STDOUT_TERMINAL {
@@ -93,16 +96,66 @@ impl Agent {
name: name.to_string(),
config: agent_config,
definition,
+ shared_variables: Default::default(),
+ session_variables: None,
functions,
rag,
model,
})
}
+ pub fn init_agent_variables(
+ agent_variables: &[AgentVariable],
+ variables: &IndexMap<String, String>,
+ ) -> Result<IndexMap<String, String>> {
+ let mut output = IndexMap::new();
+ if agent_variables.is_empty() {
+ return Ok(output);
+ }
+ let mut printed = false;
+ for agent_variable in agent_variables {
+ let key = agent_variable.name.clone();
+ match variables.get(&key) {
+ Some(value) => {
+ output.insert(key, value.clone());
+ }
+ None => {
+ if let Some(value) = agent_variable.default.clone() {
+ output.insert(key, value);
+ continue;
+ }
+ if *IS_STDOUT_TERMINAL {
+ if !printed {
+ println!("🚀 Init agent variables...");
+ printed = true;
+ }
+ let value = Text::new(&agent_variable.description)
+ .with_validator(|input: &str| {
+ if input.trim().is_empty() {
+ Ok(Validation::Invalid("This field is required".into()))
+ } else {
+ Ok(Validation::Valid)
+ }
+ })
+ .prompt()?;
+ output.insert(key, value);
+ } else {
+ bail!("Failed to init agent variables in non-interactive mode");
+ }
+ }
+ }
+ }
+ Ok(output)
+ }
+
pub fn export(&self) -> Result<String> {
let mut agent = self.clone();
agent.definition.instructions = self.interpolated_instructions();
let mut value = serde_json::json!(agent);
+ let variables = self.variables();
+ if !variables.is_empty() {
+ value["variables"] = serde_json::to_value(variables)?;
+ }
value["functions_dir"] = Config::agent_functions_dir(&self.name)?
.display()
.to_string()
@@ -140,7 +193,7 @@ impl Agent {
}
pub fn interpolated_instructions(&self) -> String {
- self.definition.interpolated_instructions()
+ self.definition.interpolated_instructions(self.variables())
}
pub fn agent_prelude(&self) -> Option<&str> {
@@ -151,18 +204,43 @@ impl Agent {
self.config.agent_prelude = value;
}
- pub fn variables(&self) -> &[AgentVariable] {
- &self.definition.variables
+ pub fn variables(&self) -> &IndexMap<String, String> {
+ match &self.session_variables {
+ Some(variables) => variables,
+ None => &self.shared_variables,
+ }
+ }
+
+ pub fn config_variables(&self) -> &IndexMap<String, String> {
+ &self.config.variables
+ }
+
+ pub fn shared_variables(&self) -> &IndexMap<String, String> {
+ &self.shared_variables
+ }
+
+ pub fn set_shared_variables(&mut self, shared_variables: IndexMap<String, String>) {
+ self.shared_variables = shared_variables;
+ }
+
+ pub fn set_session_variables(&mut self, session_variables: Option<IndexMap<String, String>>) {
+ self.session_variables = session_variables;
}
pub fn set_variable(&mut self, key: &str, value: &str) -> Result<()> {
- match self.definition.variables.iter_mut().find(|v| v.name == key) {
- Some(variable) => {
- variable.value = value.to_string();
- Ok(())
- }
- None => bail!("Unknown variable '{key}'"),
+ let variables = match self.session_variables.as_mut() {
+ Some(v) => v,
+ None => &mut self.shared_variables,
+ };
+ if !variables.contains_key(key) {
+ bail!("Unknown variable: '{key}'")
}
+ variables.insert(key.to_string(), value.to_string());
+ Ok(())
+ }
+
+ pub fn defined_variables(&self) -> &[AgentVariable] {
+ &self.definition.variables
}
}
@@ -296,10 +374,10 @@ impl AgentDefinition {
)
}
- fn interpolated_instructions(&self) -> String {
+ fn interpolated_instructions(&self, variables: &IndexMap<String, String>) -> String {
let mut output = self.instructions.clone();
- for variable in &self.variables {
- output = output.replace(&format!("{{{{{}}}}}", variable.name), &variable.value)
+ for (k, v) in variables {
+ output = output.replace(&format!("{{{{{k}}}}}"), v)
}
interpolate_variables(&mut output);
output
@@ -356,38 +434,3 @@ fn list_agents_impl() -> Result<Vec<String>> {
.collect();
Ok(agents)
}
-
-fn init_variables(
- variables: &mut [AgentVariable],
- config_variable: &IndexMap<String, String>,
-) -> Result<()> {
- if variables.is_empty() {
- return Ok(());
- }
- for variable in variables.iter_mut() {
- match config_variable.get(&variable.name) {
- Some(value) => variable.value = value.to_string(),
- None => {
- if let Some(value) = variable.default.clone() {
- variable.value = value;
- continue;
- }
- if *IS_STDOUT_TERMINAL {
- let value = Text::new(&variable.description)
- .with_validator(|input: &str| {
- if input.trim().is_empty() {
- Ok(Validation::Invalid("This field is required".into()))
- } else {
- Ok(Validation::Valid)
- }
- })
- .prompt()?;
- variable.value = value;
- } else {
- bail!("Failed to init agent variables in non-interactive mode");
- }
- }
- }
- }
- Ok(())
-}
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 6f99fd6..7673133 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -1039,6 +1039,7 @@ impl Config {
}
}
self.session = session;
+ self.init_agent_session_variables()?;
Ok(())
}
@@ -1058,6 +1059,9 @@ impl Config {
session.exit(&sessions_dir, self.working_mode.is_repl())?;
self.last_message = None;
}
+ if let Some(agent) = self.agent.as_mut() {
+ agent.set_session_variables(None);
+ }
Ok(())
}
@@ -1097,10 +1101,10 @@ impl Config {
pub fn empty_session(&mut self) -> Result<()> {
if let Some(session) = self.session.as_mut() {
- session.clear_messages();
if let Some(agent) = self.agent.as_ref() {
- session.set_agent(agent);
+ session.sync_agent(agent, false);
}
+ session.clear_messages();
} else {
bail!("No session")
}
@@ -1380,6 +1384,8 @@ impl Config {
config.write().agent = Some(agent);
if let Some(session) = session {
config.write().use_session(Some(&session))?;
+ } else {
+ config.write().init_agent_shared_variables()?;
}
Ok(())
}
@@ -1408,7 +1414,12 @@ impl Config {
let key = parts[0];
let value = parts[1];
match self.agent.as_mut() {
- Some(agent) => agent.set_variable(key, value)?,
+ Some(agent) => {
+ agent.set_variable(key, value)?;
+ if let Some(session) = self.session.as_mut() {
+ session.sync_agent(agent, true);
+ }
+ }
None => bail!("No agent"),
};
Ok(())
@@ -1563,7 +1574,7 @@ impl Config {
},
".variable" => match &self.agent {
Some(agent) => agent
- .variables()
+ .defined_variables()
.iter()
.map(|v| (v.name.clone(), Some(v.description.clone())))
.collect(),
@@ -1856,6 +1867,39 @@ impl Config {
.with_context(|| "Failed to save message")
}
+ fn init_agent_shared_variables(&mut self) -> Result<()> {
+ let agent = match self.agent.as_mut() {
+ Some(v) => v,
+ None => return Ok(()),
+ };
+ let new_variables =
+ Agent::init_agent_variables(agent.defined_variables(), agent.config_variables())?;
+ agent.set_shared_variables(new_variables);
+ Ok(())
+ }
+
+ fn init_agent_session_variables(&mut self) -> Result<()> {
+ let (agent, session) = match (self.agent.as_mut(), self.session.as_mut()) {
+ (Some(agent), Some(session)) => (agent, session),
+ _ => return Ok(()),
+ };
+ let config_variables = agent.config_variables();
+ let shared_variables = agent.shared_variables();
+ let mut all_variables = if shared_variables.is_empty() {
+ config_variables.clone()
+ } else {
+ shared_variables.clone()
+ };
+ all_variables.extend(session.agent_variables().clone());
+ let new_variables = Agent::init_agent_variables(agent.defined_variables(), &all_variables)?;
+ if shared_variables.is_empty() {
+ agent.set_shared_variables(new_variables.clone());
+ }
+ agent.set_session_variables(Some(new_variables));
+ session.sync_agent(agent, false);
+ Ok(())
+ }
+
fn open_message_file(&self) -> Result<File> {
let path = self.messages_file()?;
ensure_parent_exists(&path)?;
diff --git a/src/config/session.rs b/src/config/session.rs
index 7335128..8e82c9a 100644
--- a/src/config/session.rs
+++ b/src/config/session.rs
@@ -27,12 +27,15 @@ pub struct Session {
#[serde(skip_serializing_if = "Option::is_none")]
compress_threshold: Option<usize>,
+ #[serde(skip_serializing_if = "Option::is_none")]
+ role_name: Option<String>,
+ #[serde(skip_serializing_if = "IndexMap::is_empty")]
+ agent_variables: IndexMap<String, String>,
+
#[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")]
- role_name: Option<String>,
messages: Vec<Message>,
@@ -79,10 +82,6 @@ impl Session {
}
}
- if let Some(agent) = &config.agent {
- session.set_agent(agent);
- }
-
Ok(session)
}
@@ -249,16 +248,24 @@ impl Session {
self.dirty = true;
}
- pub fn set_agent(&mut self, agent: &Agent) {
- self.role_prompt
- .clone_from(&agent.interpolated_instructions());
- }
-
pub fn clear_role(&mut self) {
self.role_name = None;
self.role_prompt.clear();
}
+ pub fn sync_agent(&mut self, agent: &Agent, set_dirty: bool) {
+ self.role_name = None;
+ self.role_prompt = agent.interpolated_instructions();
+ self.agent_variables = agent.variables().clone();
+ if set_dirty {
+ self.dirty = true;
+ }
+ }
+
+ pub fn agent_variables(&self) -> &IndexMap<String, String> {
+ &self.agent_variables
+ }
+
pub fn set_save_session(&mut self, value: Option<bool>) {
if self.name == TEMP_SESSION_NAME {
return;
diff --git a/src/function.rs b/src/function.rs
index 3e0828a..dfc3de2 100644
--- a/src/function.rs
+++ b/src/function.rs
@@ -163,10 +163,10 @@ impl ToolCall {
let envs: HashMap<String, String> = agent
.variables()
.iter()
- .map(|v| {
+ .map(|(k, v)| {
(
- format!("LLM_AGENT_VAR_{}", normalize_env_name(&v.name)),
- v.value.clone(),
+ format!("LLM_AGENT_VAR_{}", normalize_env_name(k)),
+ v.clone(),
)
})
.collect();
diff --git a/src/repl/mod.rs b/src/repl/mod.rs
index f40bdd0..4482b84 100644
--- a/src/repl/mod.rs
+++ b/src/repl/mod.rs
@@ -63,7 +63,7 @@ lazy_static::lazy_static! {
ReplCommand::new(
".exit role",
"Leave the role",
- AssertState::True(StateFlags::ROLE),
+ AssertState::TrueFalse(StateFlags::ROLE, StateFlags::SESSION),
),
ReplCommand::new(
".session",
@@ -108,7 +108,7 @@ lazy_static::lazy_static! {
ReplCommand::new(
".edit rag-docs",
"Edit the RAG documents",
- AssertState::True(StateFlags::RAG),
+ AssertState::TrueFalse(StateFlags::RAG, StateFlags::AGENT),
),
ReplCommand::new(
".rebuild rag",
@@ -139,7 +139,7 @@ lazy_static::lazy_static! {
ReplCommand::new(
".variable",
"Set agent variable",
- AssertState::True(StateFlags::AGENT)
+ AssertState::TrueFalse(StateFlags::AGENT, StateFlags::SESSION)
),
ReplCommand::new(
".info agent",