1use crate::error::{AgentDbError, Result};
2use crate::schema::now_ms;
3use rusqlite::params;
4use rusqlite::Connection;
5use std::sync::{Arc, Mutex};
6
7#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
8pub struct DataLabel {
9 pub table_name: String,
10 pub record_id: String,
11 pub label: String,
12 pub tagged_by: Option<String>,
13 pub tagged_at: i64,
14}
15
16pub struct LabelStore {
17 conn: Arc<Mutex<Connection>>,
18}
19
20impl LabelStore {
21 pub(crate) fn new(conn: Arc<Mutex<Connection>>) -> Self {
22 Self { conn }
23 }
24
25 pub fn tag(
26 &self,
27 table_name: &str,
28 record_id: &str,
29 label: &str,
30 tagged_by: Option<&str>,
31 ) -> Result<()> {
32 let conn = self.conn.lock().unwrap();
33 let now = now_ms();
34 conn.execute(
35 "INSERT OR REPLACE INTO _adb_data_labels (table_name, record_id, label, tagged_by, tagged_at)
36 VALUES (?1, ?2, ?3, ?4, ?5)",
37 params![table_name, record_id, label, tagged_by, now],
38 )?;
39 Ok(())
40 }
41
42 pub fn untag(&self, table_name: &str, record_id: &str, label: &str) -> Result<()> {
43 let conn = self.conn.lock().unwrap();
44 conn.execute(
45 "DELETE FROM _adb_data_labels WHERE table_name = ?1 AND record_id = ?2 AND label = ?3",
46 params![table_name, record_id, label],
47 )?;
48 Ok(())
49 }
50
51 pub fn get_labels(&self, table_name: &str, record_id: &str) -> Result<Vec<DataLabel>> {
52 let conn = self.conn.lock().unwrap();
53 let mut stmt = conn.prepare(
54 "SELECT table_name, record_id, label, tagged_by, tagged_at
55 FROM _adb_data_labels
56 WHERE table_name = ?1 AND record_id = ?2
57 ORDER BY tagged_at",
58 )?;
59 let rows = stmt.query_map(params![table_name, record_id], parse_label_row)?;
60 rows.map(|r| r.map_err(AgentDbError::Sqlite)).collect()
61 }
62
63 pub fn find_by_label(&self, label: &str, limit: Option<usize>) -> Result<Vec<DataLabel>> {
64 let conn = self.conn.lock().unwrap();
65 let lim = limit.unwrap_or(100) as i64;
66 let mut stmt = conn.prepare(
67 "SELECT table_name, record_id, label, tagged_by, tagged_at
68 FROM _adb_data_labels
69 WHERE label = ?1
70 ORDER BY tagged_at DESC LIMIT ?2",
71 )?;
72 let rows = stmt.query_map(params![label, lim], parse_label_row)?;
73 rows.map(|r| r.map_err(AgentDbError::Sqlite)).collect()
74 }
75
76 pub fn has_label(&self, table_name: &str, record_id: &str, label: &str) -> Result<bool> {
77 let conn = self.conn.lock().unwrap();
78 let count: i64 = conn
79 .query_row(
80 "SELECT COUNT(*) FROM _adb_data_labels
81 WHERE table_name = ?1 AND record_id = ?2 AND label = ?3",
82 params![table_name, record_id, label],
83 |r| r.get(0),
84 )
85 .unwrap_or(0);
86 Ok(count > 0)
87 }
88
89 pub fn clear_record(&self, table_name: &str, record_id: &str) -> Result<()> {
90 let conn = self.conn.lock().unwrap();
91 conn.execute(
92 "DELETE FROM _adb_data_labels WHERE table_name = ?1 AND record_id = ?2",
93 params![table_name, record_id],
94 )?;
95 Ok(())
96 }
97}
98
99fn parse_label_row(row: &rusqlite::Row) -> rusqlite::Result<DataLabel> {
100 Ok(DataLabel {
101 table_name: row.get(0)?,
102 record_id: row.get(1)?,
103 label: row.get(2)?,
104 tagged_by: row.get(3)?,
105 tagged_at: row.get(4)?,
106 })
107}