1use crate::conversations::ConversationStore;
2use crate::error::Result;
3use crate::fts::FullTextStore;
4use crate::hybrid::{HybridQuery, HybridResult, HybridStore};
5use crate::memory::MemoryGraph;
6use crate::schema;
7use crate::traces::TraceStore;
8use crate::vectors::VectorStore;
9use crate::workflows::WorkflowStore;
10use rusqlite::Connection;
11use std::sync::{Arc, Mutex};
12
13pub struct AgentDB {
15 conn: Arc<Mutex<Connection>>,
16}
17
18impl AgentDB {
19 pub fn open(path: &str) -> Result<Self> {
21 let conn = Connection::open(path)?;
22 schema::bootstrap(&conn)?;
23 schema::check_version(&conn)?;
24 Ok(Self {
25 conn: Arc::new(Mutex::new(conn)),
26 })
27 }
28
29 pub fn vectors(&self) -> VectorStore {
31 VectorStore::new(Arc::clone(&self.conn))
32 }
33
34 pub fn memory(&self) -> MemoryGraph {
36 MemoryGraph::new(Arc::clone(&self.conn))
37 }
38
39 pub fn fts(&self) -> FullTextStore {
41 FullTextStore::new(Arc::clone(&self.conn))
42 }
43
44 pub fn conversations(&self) -> ConversationStore {
46 ConversationStore::new(Arc::clone(&self.conn))
47 }
48
49 pub fn workflows(&self) -> WorkflowStore {
51 WorkflowStore::new(Arc::clone(&self.conn))
52 }
53
54 pub fn traces(&self) -> TraceStore {
56 TraceStore::new(Arc::clone(&self.conn))
57 }
58
59 pub fn hybrid_query(&self, q: HybridQuery) -> Result<Vec<HybridResult>> {
61 let dim: usize = {
62 let conn = self.conn.lock().unwrap();
63 conn.query_row(
64 "SELECT dim FROM _adb_collections WHERE name = ?1",
65 rusqlite::params![q.collection],
66 |r| r.get::<_, i64>(0).map(|v| v as usize),
67 )
68 .unwrap_or(q.embedding.len())
69 };
70 let col = self.vectors().collection(q.collection, dim)?;
71 let store = HybridStore::new(Arc::clone(&self.conn));
72 store.query(q, &col)
73 }
74
75 pub fn execute(&self, sql: &str) -> Result<usize> {
77 let conn = self.conn.lock().unwrap();
78 Ok(conn.execute(sql, [])?)
79 }
80
81 pub fn execute_params(&self, sql: &str, params: &[&dyn rusqlite::ToSql]) -> Result<usize> {
83 let conn = self.conn.lock().unwrap();
84 Ok(conn.execute(sql, params)?)
85 }
86
87 pub fn transaction<F, T>(&self, f: F) -> Result<T>
105 where
106 F: FnOnce(&rusqlite::Transaction) -> Result<T>,
107 {
108 let mut conn = self.conn.lock().unwrap();
109 let tx = conn.transaction()?;
110 let result = f(&tx)?;
111 tx.commit()?;
112 Ok(result)
113 }
114
115 pub fn execute_batch(&self, sql: &str) -> Result<()> {
131 let mut conn = self.conn.lock().unwrap();
132 let tx = conn.transaction()?;
133 tx.execute_batch(sql)?;
134 tx.commit()?;
135 Ok(())
136 }
137
138 pub fn query_json(&self, sql: &str) -> Result<Vec<serde_json::Value>> {
140 let conn = self.conn.lock().unwrap();
141 let mut stmt = conn.prepare(sql)?;
142 let col_names: Vec<String> = stmt.column_names().iter().map(|s| s.to_string()).collect();
143 let rows = stmt.query_map([], |row| {
144 let mut map = serde_json::Map::new();
145 for (i, name) in col_names.iter().enumerate() {
146 let val: rusqlite::types::Value = row.get(i)?;
147 map.insert(name.clone(), rusqlite_value_to_json(val));
148 }
149 Ok(serde_json::Value::Object(map))
150 })?;
151 rows.map(|r| r.map_err(crate::error::AgentDbError::Sqlite))
152 .collect()
153 }
154
155 pub fn close(self) -> Result<()> {
157 let collections = self.vectors().list_collections()?;
158 for (name, dim, _) in collections {
159 let col = self.vectors().collection(&name, dim)?;
160 let is_dirty: i64 = {
161 let conn = self.conn.lock().unwrap();
162 conn.query_row(
163 "SELECT COALESCE(
164 (SELECT is_dirty FROM _adb_hnsw_index
165 WHERE collection_id =
166 (SELECT id FROM _adb_collections WHERE name = ?1)
167 ), 0)",
168 rusqlite::params![name],
169 |r| r.get(0),
170 )
171 .unwrap_or(0)
172 };
173 if is_dirty == 1 {
174 col.reindex()?;
175 }
176 }
177 Ok(())
178 }
179
180 pub fn stats(&self) -> Result<DbStats> {
182 let conn = self.conn.lock().unwrap();
183 let collections: i64 =
184 conn.query_row("SELECT COUNT(*) FROM _adb_collections", [], |r| r.get(0))?;
185 let vectors: i64 = conn.query_row(
186 "SELECT COALESCE(SUM(count), 0) FROM _adb_collections",
187 [],
188 |r| r.get(0),
189 )?;
190 let nodes: i64 = conn.query_row("SELECT COUNT(*) FROM _adb_nodes", [], |r| r.get(0))?;
191 let edges: i64 = conn.query_row("SELECT COUNT(*) FROM _adb_edges", [], |r| r.get(0))?;
192 let conversations: i64 =
193 conn.query_row("SELECT COUNT(*) FROM _adb_conversations", [], |r| r.get(0))?;
194 let messages: i64 =
195 conn.query_row("SELECT COUNT(*) FROM _adb_messages", [], |r| r.get(0))?;
196 let workflows: i64 =
197 conn.query_row("SELECT COUNT(*) FROM _adb_workflows", [], |r| r.get(0))?;
198 let workflow_steps: i64 =
199 conn.query_row("SELECT COUNT(*) FROM _adb_workflow_steps", [], |r| r.get(0))?;
200 let traces: i64 =
201 conn.query_row("SELECT COUNT(*) FROM _adb_traces", [], |r| r.get(0))?;
202 Ok(DbStats {
203 collections,
204 vectors,
205 nodes,
206 edges,
207 conversations,
208 messages,
209 workflows,
210 workflow_steps,
211 traces,
212 })
213 }
214}
215
216#[derive(Debug)]
218pub struct DbStats {
219 pub collections: i64,
221 pub vectors: i64,
223 pub nodes: i64,
225 pub edges: i64,
227 pub conversations: i64,
229 pub messages: i64,
231 pub workflows: i64,
233 pub workflow_steps: i64,
235 pub traces: i64,
237}
238
239fn rusqlite_value_to_json(val: rusqlite::types::Value) -> serde_json::Value {
240 match val {
241 rusqlite::types::Value::Null => serde_json::Value::Null,
242 rusqlite::types::Value::Integer(i) => serde_json::Value::Number(i.into()),
243 rusqlite::types::Value::Real(f) => serde_json::Number::from_f64(f)
244 .map(serde_json::Value::Number)
245 .unwrap_or(serde_json::Value::Null),
246 rusqlite::types::Value::Text(s) => serde_json::Value::String(s),
247 rusqlite::types::Value::Blob(b) => {
248 serde_json::Value::String(format!("<blob {} bytes>", b.len()))
249 }
250 }
251}