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/function.rs | |
| parent | 8b0c648a73dc812b0fd22a60a2975de57b7dd02c (diff) | |
| download | aichat-10a4c23c83b2101b049e3d94458310a9511eb18e.tar.gz | |
feat: agent can reuse tools (#690)
Diffstat (limited to 'src/function.rs')
| -rw-r--r-- | src/function.rs | 8 |
1 files changed, 7 insertions, 1 deletions
diff --git a/src/function.rs b/src/function.rs index 739ef0c..0016405 100644 --- a/src/function.rs +++ b/src/function.rs @@ -71,6 +71,10 @@ impl Functions { Ok(Self { declarations }) } + pub fn find(&self, name: &str) -> Option<&FunctionDeclaration> { + self.declarations.iter().find(|v| v.name == name) + } + pub fn contains(&self, name: &str) -> bool { self.declarations.iter().any(|v| v.name == name) } @@ -89,6 +93,8 @@ pub struct FunctionDeclaration { pub name: String, pub description: String, pub parameters: JsonSchema, + #[serde(skip_serializing, default)] + pub agent: bool, } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -154,7 +160,7 @@ 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 let Some(true) = agent.functions().find(&function_name).map(|v| v.agent) { ( format!("{}:{}", agent.name(), function_name), agent.name().to_string(), |
