From a0d421ffc83ba97a1f99f1ba9f85a1689941d171 Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 30 Jul 2024 15:04:50 +0800 Subject: feat: export agent variable as `LLM_AGENT_VAR_*` (#766) --- src/config/mod.rs | 8 ++------ src/function.rs | 28 ++++++++++++++++++++++++---- src/utils/mod.rs | 4 ++++ 3 files changed, 30 insertions(+), 10 deletions(-) diff --git a/src/config/mod.rs b/src/config/mod.rs index 8f719f4..81fa6f8 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -342,7 +342,7 @@ impl Config { } pub fn agent_config_dir(name: &str) -> Result { - match env::var(format!("{}_CONFIG_DIR", convert_env_prefix(name))) { + match env::var(format!("{}_CONFIG_DIR", normalize_env_name(name))) { Ok(value) => Ok(PathBuf::from(value)), Err(_) => Ok(Self::agents_config_dir()?.join(name)), } @@ -365,7 +365,7 @@ impl Config { } pub fn agent_functions_dir(name: &str) -> Result { - match env::var(format!("{}_FUNCTIONS_DIR", convert_env_prefix(name))) { + match env::var(format!("{}_FUNCTIONS_DIR", normalize_env_name(name))) { Ok(value) => Ok(PathBuf::from(value)), Err(_) => Ok(Self::agents_functions_dir()?.join(name)), } @@ -1904,7 +1904,3 @@ fn complete_option_bool(value: Option) -> Vec { None => vec!["true".to_string(), "false".to_string()], } } - -fn convert_env_prefix(value: &str) -> String { - value.replace('-', "_").to_ascii_uppercase() -} diff --git a/src/function.rs b/src/function.rs index 18a9751..69e2bdf 100644 --- a/src/function.rs +++ b/src/function.rs @@ -156,23 +156,44 @@ impl ToolCall { pub fn eval(&self, config: &GlobalConfig) -> Result { let function_name = self.name.clone(); - let (call_name, cmd_name, mut cmd_args) = match &config.read().agent { + let (call_name, cmd_name, mut cmd_args, mut envs) = match &config.read().agent { Some(agent) => match agent.functions().find(&function_name) { Some(function) => { if function.agent { + let envs: HashMap = agent + .variables() + .iter() + .map(|v| { + ( + format!("LLM_AGENT_VAR_{}", normalize_env_name(&v.name)), + v.value.clone(), + ) + }) + .collect(); ( format!("{}:{}", agent.name(), function_name), agent.name().to_string(), vec![function_name], + envs, ) } else { - (function_name.clone(), function_name, vec![]) + ( + function_name.clone(), + function_name, + vec![], + Default::default(), + ) } } None => bail!("Unexpected call {function_name} {}", self.arguments), }, None => match config.read().functions.contains(&function_name) { - true => (function_name.clone(), function_name, vec![]), + true => ( + function_name.clone(), + function_name, + vec![], + Default::default(), + ), false => bail!("Unexpected call: {function_name} {}", self.arguments), }, }; @@ -193,7 +214,6 @@ impl ToolCall { cmd_args.push(json_data.to_string()); let prompt = format!("Call {cmd_name} {}", cmd_args.join(" ")); - let mut envs = HashMap::new(); let bin_dir = Config::functions_bin_dir()?; if bin_dir.exists() { envs.insert("PATH".into(), prepend_env_path(&bin_dir)?); diff --git a/src/utils/mod.rs b/src/utils/mod.rs index a3aef22..6f89fbe 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -38,6 +38,10 @@ pub fn get_env_name(key: &str) -> String { format!("{}_{key}", env!("CARGO_CRATE_NAME"),).to_ascii_uppercase() } +pub fn normalize_env_name(value: &str) -> String { + value.replace('-', "_").to_ascii_uppercase() +} + pub fn estimate_token_length(text: &str) -> usize { let words: Vec<&str> = text.unicode_words().collect(); let mut output: f32 = 0.0; -- cgit v1.2.3