diff options
| author | sigoden <sigoden@gmail.com> | 2024-07-09 09:36:50 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-07-09 09:36:50 +0800 |
| commit | 10a4c23c83b2101b049e3d94458310a9511eb18e (patch) | |
| tree | 0317be3d80890feb5d01e198e8cbb1ba3eccbccc /src/config | |
| parent | 8b0c648a73dc812b0fd22a60a2975de57b7dd02c (diff) | |
| download | aichat-10a4c23c83b2101b049e3d94458310a9511eb18e.tar.gz | |
feat: agent can reuse tools (#690)
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/agent.rs | 15 | ||||
| -rw-r--r-- | src/config/mod.rs | 52 |
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.")); } } }; |
