1use crate::error::{AgentDbError, Result};
2use crate::schema::now_ms;
3use rusqlite::params;
4use rusqlite::Connection;
5use std::sync::{Arc, Mutex};
6use uuid::Uuid;
7
8#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
9pub struct ContextEntry {
10 pub id: String,
11 pub session_id: String,
12 pub source_type: String,
13 pub source_id: String,
14 pub content_preview: Option<String>,
15 pub token_count: i64,
16 pub relevance_score: f64,
17 pub priority: i64,
18 pub included_at: i64,
19}
20
21pub struct ContextStore {
22 conn: Arc<Mutex<Connection>>,
23}
24
25impl ContextStore {
26 pub(crate) fn new(conn: Arc<Mutex<Connection>>) -> Self {
27 Self { conn }
28 }
29
30 #[allow(clippy::too_many_arguments)]
31 pub fn add_entry(
32 &self,
33 session_id: &str,
34 source_type: &str,
35 source_id: &str,
36 content_preview: Option<&str>,
37 token_count: i64,
38 relevance_score: f64,
39 priority: i64,
40 ) -> Result<String> {
41 let id = Uuid::new_v4().to_string();
42 let conn = self.conn.lock().unwrap();
43 let now = now_ms();
44 conn.execute(
45 "INSERT INTO _adb_context_entries
46 (id, session_id, source_type, source_id, content_preview, token_count, relevance_score, priority, included_at)
47 VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)",
48 params![id, session_id, source_type, source_id, content_preview, token_count, relevance_score, priority, now],
49 )?;
50 Ok(id)
51 }
52
53 pub fn build_window(&self, session_id: &str, max_tokens: i64) -> Result<Vec<ContextEntry>> {
58 let conn = self.conn.lock().unwrap();
59 let mut stmt = conn.prepare(
60 "SELECT id, session_id, source_type, source_id, content_preview, token_count, relevance_score, priority, included_at
61 FROM _adb_context_entries
62 WHERE session_id = ?1
63 ORDER BY priority DESC, relevance_score DESC",
64 )?;
65 let rows = stmt.query_map(params![session_id], parse_context_row)?;
66 let mut result = Vec::new();
67 let mut running_tokens: i64 = 0;
68 for row in rows {
69 let entry = row.map_err(AgentDbError::Sqlite)?;
70 if running_tokens + entry.token_count > max_tokens {
71 continue;
72 }
73 running_tokens += entry.token_count;
74 result.push(entry);
75 }
76 Ok(result)
77 }
78
79 pub fn get_entries(&self, session_id: &str) -> Result<Vec<ContextEntry>> {
80 let conn = self.conn.lock().unwrap();
81 let mut stmt = conn.prepare(
82 "SELECT id, session_id, source_type, source_id, content_preview, token_count, relevance_score, priority, included_at
83 FROM _adb_context_entries
84 WHERE session_id = ?1
85 ORDER BY priority DESC, relevance_score DESC",
86 )?;
87 let rows = stmt.query_map(params![session_id], parse_context_row)?;
88 rows.map(|r| r.map_err(AgentDbError::Sqlite)).collect()
89 }
90
91 pub fn clear_session(&self, session_id: &str) -> Result<()> {
92 let conn = self.conn.lock().unwrap();
93 conn.execute(
94 "DELETE FROM _adb_context_entries WHERE session_id = ?1",
95 params![session_id],
96 )?;
97 Ok(())
98 }
99
100 pub fn remove_entry(&self, id: &str) -> Result<()> {
101 let conn = self.conn.lock().unwrap();
102 conn.execute(
103 "DELETE FROM _adb_context_entries WHERE id = ?1",
104 params![id],
105 )?;
106 Ok(())
107 }
108}
109
110fn parse_context_row(row: &rusqlite::Row) -> rusqlite::Result<ContextEntry> {
111 Ok(ContextEntry {
112 id: row.get(0)?,
113 session_id: row.get(1)?,
114 source_type: row.get(2)?,
115 source_id: row.get(3)?,
116 content_preview: row.get(4)?,
117 token_count: row.get(5)?,
118 relevance_score: row.get(6)?,
119 priority: row.get(7)?,
120 included_at: row.get(8)?,
121 })
122}