1use async_trait::async_trait;
2use std::collections::HashMap;
3
4use crate::error::Result;
5use crate::types::{ToolRequest, ToolResponse, ToolSpec};
6
7#[async_trait]
9pub trait Tool: Send + Sync {
10 fn spec(&self) -> ToolSpec;
12
13 async fn invoke(&self, req: ToolRequest) -> Result<ToolResponse>;
15}
16
17#[derive(Default)]
19pub struct ToolCatalog {
20 tools: parking_lot::RwLock<HashMap<String, Box<dyn Tool>>>,
21}
22
23impl ToolCatalog {
24 pub fn new() -> Self {
26 Self {
27 tools: parking_lot::RwLock::new(HashMap::new()),
28 }
29 }
30
31 pub fn register(&self, tool: Box<dyn Tool>) -> Result<()> {
33 let spec = tool.spec();
34 let mut tools = self.tools.write();
35 tools.insert(spec.name.clone(), tool);
36 Ok(())
37 }
38
39 pub fn lookup(&self, name: &str) -> Option<ToolSpec> {
41 let tools = self.tools.read();
42 tools.get(name).map(|tool| tool.spec())
43 }
44
45 pub fn specs(&self) -> Vec<ToolSpec> {
47 let tools = self.tools.read();
48 tools.values().map(|tool| tool.spec()).collect()
49 }
50
51 pub async fn invoke(&self, name: &str, req: ToolRequest) -> Result<ToolResponse> {
53 let tool = {
54 let tools = self.tools.read();
55 tools.get(name).map(|t| t.spec().name.clone())
56 };
57
58 if tool.is_none() {
59 return Err(crate::error::AgentError::ToolNotFound(name.to_string()));
60 }
61
62 let tools = self.tools.read();
63 let tool = tools.get(name).unwrap();
64 tool.invoke(req).await
65 }
66}
67
68#[cfg(test)]
69mod tests {
70 use super::*;
71
72 struct EchoTool;
73
74 #[async_trait]
75 impl Tool for EchoTool {
76 fn spec(&self) -> ToolSpec {
77 ToolSpec {
78 name: "echo".to_string(),
79 description: "Echoes the input".to_string(),
80 input_schema: serde_json::json!({
81 "type": "object",
82 "properties": {
83 "input": {
84 "type": "string",
85 "description": "Text to echo"
86 }
87 },
88 "required": ["input"]
89 }),
90 examples: None,
91 }
92 }
93
94 async fn invoke(&self, req: ToolRequest) -> Result<ToolResponse> {
95 let input = req
96 .arguments
97 .get("input")
98 .and_then(|v| v.as_str())
99 .unwrap_or("");
100
101 Ok(ToolResponse {
102 content: input.to_string(),
103 metadata: None,
104 })
105 }
106 }
107
108 #[tokio::test]
109 async fn test_tool_catalog() {
110 let catalog = ToolCatalog::new();
111 catalog.register(Box::new(EchoTool)).unwrap();
112
113 let spec = catalog.lookup("echo");
114 assert!(spec.is_some());
115
116 let mut args = HashMap::new();
117 args.insert("input".to_string(), serde_json::json!("hello"));
118
119 let response = catalog
120 .invoke(
121 "echo",
122 ToolRequest {
123 session_id: "test".to_string(),
124 arguments: args,
125 },
126 )
127 .await
128 .unwrap();
129
130 assert_eq!(response.content, "hello");
131 }
132}