summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-07-30 15:04:50 +0800
committerGitHub <noreply@github.com>2024-07-30 15:04:50 +0800
commita0d421ffc83ba97a1f99f1ba9f85a1689941d171 (patch)
tree15af76aab2baf1a2843fb98860b5a3c7099239a9 /src
parentcc74be5d2163fbb1e9d57f85d5546857c6602446 (diff)
downloadaichat-a0d421ffc83ba97a1f99f1ba9f85a1689941d171.tar.gz
feat: export agent variable as `LLM_AGENT_VAR_*` (#766)
Diffstat (limited to 'src')
-rw-r--r--src/config/mod.rs8
-rw-r--r--src/function.rs28
-rw-r--r--src/utils/mod.rs4
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<PathBuf> {
- 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<PathBuf> {
- 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<bool>) -> Vec<String> {
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<Value> {
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<String, String> = 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;