summaryrefslogtreecommitdiffstats
path: root/src/function.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/function.rs')
-rw-r--r--src/function.rs49
1 files changed, 16 insertions, 33 deletions
diff --git a/src/function.rs b/src/function.rs
index fcb4c93..739ef0c 100644
--- a/src/function.rs
+++ b/src/function.rs
@@ -4,8 +4,7 @@ use crate::{
};
use anyhow::{anyhow, bail, Context, Result};
-use fancy_regex::Regex;
-use indexmap::{IndexMap, IndexSet};
+use indexmap::IndexMap;
use inquire::{validator::Validation, Text};
use serde::{Deserialize, Serialize};
use serde_json::{json, Value};
@@ -15,9 +14,7 @@ use std::{
path::Path,
};
-pub const SELECTED_ALL_FUNCTIONS: &str = ".*";
pub type ToolResults = (Vec<ToolResult>, String);
-pub type FunctionsFilter = String;
pub fn eval_tool_calls(config: &GlobalConfig, mut calls: Vec<ToolCall>) -> Result<Vec<ToolResult>> {
let mut output = vec![];
@@ -53,7 +50,6 @@ impl ToolResult {
#[derive(Debug, Clone, Default)]
pub struct Functions {
- names: IndexSet<String>,
declarations: Vec<FunctionDeclaration>,
}
@@ -72,35 +68,19 @@ impl Functions {
vec![]
};
- let names = declarations.iter().map(|v| v.name.clone()).collect();
-
- Ok(Self {
- names,
- declarations,
- })
+ Ok(Self { declarations })
}
- pub fn select(&self, filter: &FunctionsFilter) -> Option<Vec<FunctionDeclaration>> {
- let regex = Regex::new(&format!("^({filter})$")).ok()?;
- let output: Vec<FunctionDeclaration> = self
- .declarations
- .iter()
- .filter(|v| regex.is_match(&v.name).unwrap_or_default())
- .cloned()
- .collect();
- if output.is_empty() {
- None
- } else {
- Some(output)
- }
+ pub fn contains(&self, name: &str) -> bool {
+ self.declarations.iter().any(|v| v.name == name)
}
- pub fn contains(&self, name: &str) -> bool {
- self.names.contains(name)
+ pub fn declarations(&self) -> &[FunctionDeclaration] {
+ &self.declarations
}
pub fn is_empty(&self) -> bool {
- self.names.is_empty()
+ self.declarations.is_empty()
}
}
@@ -174,18 +154,21 @@ impl ToolCall {
let is_dangerously = config.read().is_dangerously_function(&function_name);
let (call_name, cmd_name, mut cmd_args) = match &config.read().agent {
Some(agent) => {
- if !agent.functions().contains(&function_name) {
+ if agent.functions().contains(&function_name) {
+ (
+ format!("{}:{}", agent.name(), function_name),
+ agent.name().to_string(),
+ vec![function_name],
+ )
+ } else if config.read().functions.contains(&function_name) {
+ (function_name.clone(), function_name, vec![])
+ } else {
bail!(
"Unexpected call: {} {function_name} {}",
agent.name(),
self.arguments
);
}
- (
- format!("{}:{}", agent.name(), function_name),
- agent.name().to_string(),
- vec![function_name],
- )
}
None => {
if !config.read().functions.contains(&function_name) {