Skip to main content

rs_agent/tools/
mod.rs

1use async_trait::async_trait;
2use std::collections::HashMap;
3
4use crate::error::Result;
5use crate::types::{ToolRequest, ToolResponse, ToolSpec};
6
7/// Tool trait for defining custom tools
8#[async_trait]
9pub trait Tool: Send + Sync {
10    /// Returns the tool specification
11    fn spec(&self) -> ToolSpec;
12
13    /// Invokes the tool with the given request
14    async fn invoke(&self, req: ToolRequest) -> Result<ToolResponse>;
15}
16
17/// Tool catalog manages registered tools
18#[derive(Default)]
19pub struct ToolCatalog {
20    tools: parking_lot::RwLock<HashMap<String, Box<dyn Tool>>>,
21}
22
23impl ToolCatalog {
24    /// Creates a new empty tool catalog
25    pub fn new() -> Self {
26        Self {
27            tools: parking_lot::RwLock::new(HashMap::new()),
28        }
29    }
30
31    /// Registers a tool in the catalog
32    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    /// Looks up a tool by name
40    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    /// Returns all tool specifications
46    pub fn specs(&self) -> Vec<ToolSpec> {
47        let tools = self.tools.read();
48        tools.values().map(|tool| tool.spec()).collect()
49    }
50
51    /// Invokes a tool by name
52    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}