Skip to main content

agentdb/
context.rs

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    /// Build a context window for a session, filling up to `max_tokens`.
54    ///
55    /// Returns entries ordered by priority (desc), then relevance (desc),
56    /// stopping when the running token sum would exceed `max_tokens`.
57    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}