summaryrefslogtreecommitdiffstats
path: root/src/config/agent.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/config/agent.rs')
-rw-r--r--src/config/agent.rs120
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(())
+}