summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
Diffstat (limited to 'src/config')
-rw-r--r--src/config/agent.rs15
-rw-r--r--src/config/mod.rs52
2 files changed, 38 insertions, 29 deletions
diff --git a/src/config/agent.rs b/src/config/agent.rs
index a6257ca..82f887c 100644
--- a/src/config/agent.rs
+++ b/src/config/agent.rs
@@ -148,12 +148,7 @@ impl RoleLike for Agent {
}
fn use_tools(&self) -> Option<String> {
- let common_tools = &self.definition.common_tools;
- if common_tools.is_empty() {
- None
- } else {
- Some(common_tools.join(","))
- }
+ self.config.use_tools.clone()
}
fn set_model(&mut self, model: &Model) {
@@ -169,7 +164,9 @@ impl RoleLike for Agent {
self.config.top_p = value;
}
- fn set_use_tools(&mut self, _value: Option<String>) {}
+ fn set_use_tools(&mut self, value: Option<String>) {
+ self.config.use_tools = value;
+ }
}
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
@@ -182,6 +179,8 @@ pub struct AgentConfig {
#[serde(skip_serializing_if = "Option::is_none")]
pub top_p: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
+ use_tools: Option<String>,
+ #[serde(skip_serializing_if = "Option::is_none")]
pub dangerously_functions_filter: Option<String>,
}
@@ -206,8 +205,6 @@ pub struct AgentDefinition {
pub conversation_starters: Vec<String>,
#[serde(default)]
pub documents: Vec<String>,
- #[serde(default)]
- pub common_tools: Vec<String>,
}
impl AgentDefinition {
diff --git a/src/config/mod.rs b/src/config/mod.rs
index d2804ed..eeffc87 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -555,9 +555,6 @@ impl Config {
self.function_calling = value;
}
"use_tools" => {
- if self.agent.is_some() {
- bail!("This action cannot be performed within an agent.")
- }
let value = parse_value(value)?;
self.set_use_tools(value);
}
@@ -1059,15 +1056,14 @@ impl Config {
pub fn select_functions(&self, model: &Model, role: &Role) -> Option<Vec<FunctionDeclaration>> {
let mut functions = vec![];
if self.function_calling {
- let use_tools = role.use_tools();
- let declaration_names: HashSet<String> = self
- .functions
- .declarations()
- .iter()
- .map(|v| v.name.to_string())
- .collect();
- if let Some(use_tools) = use_tools {
+ if let Some(use_tools) = role.use_tools() {
let mut tool_names: HashSet<String> = Default::default();
+ let declaration_names: HashSet<String> = self
+ .functions
+ .declarations()
+ .iter()
+ .map(|v| v.name.to_string())
+ .collect();
for item in use_tools.split(',') {
let item = item.trim();
if item == "all" {
@@ -1096,15 +1092,31 @@ impl Config {
}
})
.collect();
- if let Some(agent) = &self.agent {
- let agent_functions = agent.functions().declarations().to_vec();
- functions = [agent_functions, functions].concat();
- }
- if !model.supports_function_calling() {
- functions.clear();
- if *IS_STDOUT_TERMINAL {
- eprintln!("{}", warning_text("WARNING: This LLM or client does not support function calling, despite the context requiring it."));
- }
+ }
+
+ if let Some(agent) = &self.agent {
+ let mut agent_functions = agent.functions().declarations().to_vec();
+ let tool_names: HashSet<String> = agent_functions
+ .iter()
+ .filter_map(|v| {
+ if v.agent {
+ None
+ } else {
+ Some(v.name.to_string())
+ }
+ })
+ .collect();
+ agent_functions.extend(
+ functions
+ .into_iter()
+ .filter(|v| !tool_names.contains(&v.name)),
+ );
+ functions = agent_functions;
+ }
+ if !functions.is_empty() && !model.supports_function_calling() {
+ functions.clear();
+ if *IS_STDOUT_TERMINAL {
+ eprintln!("{}", warning_text("WARNING: This LLM or client does not support function calling, despite the context requiring it."));
}
}
};