diff options
| author | sigoden <sigoden@gmail.com> | 2024-07-09 21:22:51 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-07-09 21:22:51 +0800 |
| commit | a9268b600fb378400795fbaf98bf923372b4c19a (patch) | |
| tree | ea0f3d79da0ae02041ad1313e9764a4750e8a3f1 /src/config/agent.rs | |
| parent | 138c90b58bf90be84de3fedbbdf8e05bf0820a11 (diff) | |
| download | aichat-a9268b600fb378400795fbaf98bf923372b4c19a.tar.gz | |
feat: support agent variables (#692)
Diffstat (limited to 'src/config/agent.rs')
| -rw-r--r-- | src/config/agent.rs | 120 |
1 files changed, 117 insertions, 3 deletions
diff --git a/src/config/agent.rs b/src/config/agent.rs index 82f887c..7a065e7 100644 --- a/src/config/agent.rs +++ b/src/config/agent.rs @@ -3,7 +3,11 @@ use super::*; use crate::{client::Model, function::Functions}; use anyhow::{Context, Result}; -use std::{fs::read_to_string, path::Path}; +use inquire::{validator::Validation, Text}; +use std::{ + fs::{self, read_to_string}, + path::Path, +}; use serde::{Deserialize, Serialize}; @@ -29,8 +33,13 @@ impl Agent { let functions_dir = Config::agent_functions_dir(name)?; let definition_file_path = functions_dir.join("index.yaml"); let functions_file_path = functions_dir.join("functions.json"); + let variables_path = Config::agent_variables_file(name)?; let rag_path = Config::agent_rag_file(name)?; - let definition = AgentDefinition::load(&definition_file_path)?; + + let mut definition = AgentDefinition::load(&definition_file_path)?; + init_variables(&variables_path, &mut definition.variables) + .context("Failed to init variables")?; + let functions = if functions_file_path.exists() { Functions::init(&functions_file_path)? } else { @@ -50,6 +59,7 @@ impl Agent { None => config.current_model().clone(), } }; + let rag = if rag_path.exists() { Some(Arc::new(Rag::load(config, "rag", &rag_path)?)) } else if !definition.documents.is_empty() { @@ -91,6 +101,10 @@ impl Agent { .display() .to_string() .into(); + value["variables_file"] = Config::agent_variables_file(&self.name)? + .display() + .to_string() + .into(); let data = serde_yaml::to_string(&value)?; Ok(data) } @@ -122,11 +136,28 @@ impl Agent { pub fn conversation_staters(&self) -> &[String] { &self.definition.conversation_starters } + + pub fn variables(&self) -> &[AgentVariable] { + &self.definition.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(); + let variables_path = Config::agent_variables_file(&self.name)?; + save_variables(&variables_path, self.variables())?; + Ok(()) + } + None => bail!("Unknown variable '{key}'"), + } + } } impl RoleLike for Agent { fn to_role(&self) -> Role { - let mut role = Role::new("", &self.definition.instructions); + let prompt = self.definition.interpolated_instructions(); + let mut role = Role::new("", &prompt); role.sync(self); role } @@ -202,6 +233,8 @@ pub struct AgentDefinition { pub version: String, pub instructions: String, #[serde(default)] + pub variables: Vec<AgentVariable>, + #[serde(default)] pub conversation_starters: Vec<String>, #[serde(default)] pub documents: Vec<String>, @@ -244,6 +277,24 @@ impl AgentDefinition { {description}{starters}"# ) } + + fn interpolated_instructions(&self) -> String { + let mut output = self.instructions.clone(); + for variable in &self.variables { + output = output.replace(&format!("{{{{{}}}}}", variable.name), &variable.value) + } + output + } +} + +#[derive(Debug, Clone, Default, Deserialize, Serialize)] +pub struct AgentVariable { + pub name: String, + pub description: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub default: Option<String>, + #[serde(skip_deserializing, default)] + pub value: String, } pub fn list_agents() -> Vec<String> { @@ -266,3 +317,66 @@ fn list_agents_impl() -> Result<Vec<String>> { .collect(); Ok(agents) } + +fn init_variables(variables_path: &Path, variables: &mut [AgentVariable]) -> Result<()> { + if variables.is_empty() { + return Ok(()); + } + let variable_values = if variables_path.exists() { + let content = read_to_string(variables_path).with_context(|| { + format!( + "Failed to read variables from '{}'", + variables_path.display() + ) + })?; + let variable_values: IndexMap<String, String> = serde_yaml::from_str(&content)?; + variable_values + } else { + Default::default() + }; + let mut initialized = false; + for variable in variables.iter_mut() { + match variable_values.get(&variable.name) { + Some(value) => variable.value = value.to_string(), + None => { + if !initialized { + println!("The agent has the variables and is initializing them..."); + initialized = true; + } + if *IS_STDOUT_TERMINAL { + let mut text = + 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) + } + }); + if let Some(default) = &variable.default { + text = text.with_default(default); + } + let value = text.prompt()?; + variable.value = value; + } else { + bail!("Failed to init agent variables in the script mode."); + } + } + } + } + if initialized { + save_variables(variables_path, variables)?; + } + Ok(()) +} + +fn save_variables(variables_path: &Path, variables: &[AgentVariable]) -> Result<()> { + ensure_parent_exists(variables_path)?; + let variable_values: IndexMap<String, String> = variables + .iter() + .map(|v| (v.name.clone(), v.value.clone())) + .collect(); + let content = serde_yaml::to_string(&variable_values)?; + fs::write(variables_path, content) + .with_context(|| format!("Failed to save variables to '{}'", variables_path.display()))?; + Ok(()) +} |
