Skip to main content

agentdb/
tools.rs

1use crate::error::{AgentDbError, Result};
2use crate::schema::now_ms;
3use rusqlite::params;
4use rusqlite::Connection;
5use serde_json::Value;
6use std::sync::{Arc, Mutex};
7use uuid::Uuid;
8
9#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
10pub struct Tool {
11    pub id: String,
12    pub name: String,
13    pub description: Option<String>,
14    pub parameters_schema: Option<Value>,
15    pub version: String,
16    pub created_at: i64,
17    pub updated_at: i64,
18}
19
20#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
21pub struct ToolCall {
22    pub id: String,
23    pub session_id: Option<String>,
24    pub tool_name: String,
25    pub arguments: Option<Value>,
26    pub result: Option<Value>,
27    pub error: Option<String>,
28    pub latency_ms: Option<i64>,
29    pub created_at: i64,
30}
31
32pub struct ToolStore {
33    conn: Arc<Mutex<Connection>>,
34}
35
36impl ToolStore {
37    pub(crate) fn new(conn: Arc<Mutex<Connection>>) -> Self {
38        Self { conn }
39    }
40
41    pub fn register_tool(
42        &self,
43        name: &str,
44        description: Option<&str>,
45        parameters_schema: Option<Value>,
46        version: Option<&str>,
47    ) -> Result<String> {
48        let id = Uuid::new_v4().to_string();
49        let conn = self.conn.lock().unwrap();
50        let schema_str = parameters_schema.as_ref().map(|v| v.to_string());
51        let ver = version.unwrap_or("1.0.0");
52        let now = now_ms();
53        conn.execute(
54            "INSERT INTO _adb_tools (id, name, description, parameters_schema, version, created_at, updated_at)
55             VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)
56             ON CONFLICT(name) DO UPDATE SET
57                 description = excluded.description,
58                 parameters_schema = excluded.parameters_schema,
59                 version = excluded.version,
60                 updated_at = excluded.updated_at",
61            params![id, name, description, schema_str, ver, now, now],
62        )?;
63        Ok(id)
64    }
65
66    pub fn get_tool(&self, name: &str) -> Result<Tool> {
67        let conn = self.conn.lock().unwrap();
68        conn.query_row(
69            "SELECT id, name, description, parameters_schema, version, created_at, updated_at
70             FROM _adb_tools WHERE name = ?1",
71            params![name],
72            |row| {
73                let schema_str: Option<String> = row.get(3)?;
74                Ok(Tool {
75                    id: row.get(0)?,
76                    name: row.get(1)?,
77                    description: row.get(2)?,
78                    parameters_schema: schema_str.and_then(|s| serde_json::from_str(&s).ok()),
79                    version: row.get(4)?,
80                    created_at: row.get(5)?,
81                    updated_at: row.get(6)?,
82                })
83            },
84        )
85        .map_err(|_| AgentDbError::InvalidArgument(format!("tool not found: {name}")))
86    }
87
88    pub fn list_tools(&self) -> Result<Vec<Tool>> {
89        let conn = self.conn.lock().unwrap();
90        let mut stmt = conn.prepare(
91            "SELECT id, name, description, parameters_schema, version, created_at, updated_at
92             FROM _adb_tools ORDER BY name",
93        )?;
94        let rows = stmt.query_map([], |row| {
95            let schema_str: Option<String> = row.get(3)?;
96            Ok(Tool {
97                id: row.get(0)?,
98                name: row.get(1)?,
99                description: row.get(2)?,
100                parameters_schema: schema_str.and_then(|s| serde_json::from_str(&s).ok()),
101                version: row.get(4)?,
102                created_at: row.get(5)?,
103                updated_at: row.get(6)?,
104            })
105        })?;
106        rows.map(|r| r.map_err(AgentDbError::Sqlite)).collect()
107    }
108
109    pub fn delete_tool(&self, name: &str) -> Result<()> {
110        let conn = self.conn.lock().unwrap();
111        conn.execute("DELETE FROM _adb_tools WHERE name = ?1", params![name])?;
112        Ok(())
113    }
114
115    pub fn log_tool_call(
116        &self,
117        session_id: Option<&str>,
118        tool_name: &str,
119        arguments: Option<Value>,
120        result: Option<Value>,
121        error: Option<&str>,
122        latency_ms: Option<i64>,
123    ) -> Result<String> {
124        let id = Uuid::new_v4().to_string();
125        let conn = self.conn.lock().unwrap();
126        let args_str = arguments.as_ref().map(|v| v.to_string());
127        let result_str = result.as_ref().map(|v| v.to_string());
128        let now = now_ms();
129        conn.execute(
130            "INSERT INTO _adb_tool_calls (id, session_id, tool_name, arguments, result, error, latency_ms, created_at)
131             VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)",
132            params![id, session_id, tool_name, args_str, result_str, error, latency_ms, now],
133        )?;
134        Ok(id)
135    }
136
137    pub fn get_tool_calls(
138        &self,
139        session_id: Option<&str>,
140        tool_name: Option<&str>,
141        limit: Option<usize>,
142    ) -> Result<Vec<ToolCall>> {
143        let conn = self.conn.lock().unwrap();
144        let sql = match (session_id, tool_name) {
145            (Some(_), Some(_)) => {
146                "SELECT id, session_id, tool_name, arguments, result, error, latency_ms, created_at
147                 FROM _adb_tool_calls WHERE session_id = ?1 AND tool_name = ?2
148                 ORDER BY created_at DESC LIMIT ?3"
149            }
150            (Some(_), None) => {
151                "SELECT id, session_id, tool_name, arguments, result, error, latency_ms, created_at
152                 FROM _adb_tool_calls WHERE session_id = ?1
153                 ORDER BY created_at DESC LIMIT ?3"
154            }
155            (None, Some(_)) => {
156                "SELECT id, session_id, tool_name, arguments, result, error, latency_ms, created_at
157                 FROM _adb_tool_calls WHERE tool_name = ?2
158                 ORDER BY created_at DESC LIMIT ?3"
159            }
160            (None, None) => {
161                "SELECT id, session_id, tool_name, arguments, result, error, latency_ms, created_at
162                 FROM _adb_tool_calls
163                 ORDER BY created_at DESC LIMIT ?3"
164            }
165        };
166        let lim = limit.unwrap_or(100) as i64;
167        let mut stmt = conn.prepare(sql)?;
168        let rows = stmt.query_map(
169            params![session_id.unwrap_or(""), tool_name.unwrap_or(""), lim],
170            parse_tool_call_row,
171        )?;
172        rows.map(|r| r.map_err(AgentDbError::Sqlite)).collect()
173    }
174}
175
176fn parse_tool_call_row(row: &rusqlite::Row) -> rusqlite::Result<ToolCall> {
177    let args_str: Option<String> = row.get(3)?;
178    let result_str: Option<String> = row.get(4)?;
179    Ok(ToolCall {
180        id: row.get(0)?,
181        session_id: row.get(1)?,
182        tool_name: row.get(2)?,
183        arguments: args_str.and_then(|s| serde_json::from_str(&s).ok()),
184        result: result_str.and_then(|s| serde_json::from_str(&s).ok()),
185        error: row.get(5)?,
186        latency_ms: row.get(6)?,
187        created_at: row.get(7)?,
188    })
189}