Skip to main content

ravenclaws/
persistence.rs

1//! # Conversation Persistence (SQLite backend)
2//!
3//! Provides SQLite-backed storage for conversation history so agents survive
4//! pod restarts without losing context. Supports configurable retention policies
5//! (time-based, count-based, token-budget-based).
6//!
7//! ## Architecture
8//!
9//! - `ConversationStore` — manages a SQLite database with sessions and messages tables
10//! - Sessions are identified by a session ID (UUID or user-provided)
11//! - Messages are stored with role, content, timestamp, and token count
12//! - Retention policies are applied on read (not on write) for simplicity
13//!
14//! ## Usage
15//!
16//! ```rust,no_run
17//! use ravenclaws::persistence::ConversationStore;
18//!
19//! # async fn example() -> Result<(), Box<dyn std::error::Error>> {
20//! let store = ConversationStore::open(":memory:")?;
21//! store.create_session("session-1", "You are a helpful assistant.")?;
22//! store.add_message("session-1", "user", "Hello!", None)?;
23//! let history = store.get_history("session-1", None)?;
24//! # Ok(())
25//! # }
26//! ```
27
28use rusqlite::{params, Connection, OptionalExtension, Result as SqlResult};
29use serde::{Deserialize, Serialize};
30use std::path::Path;
31use std::time::{Duration, SystemTime, UNIX_EPOCH};
32
33/// A stored conversation message
34#[derive(Debug, Clone, Serialize, Deserialize)]
35pub struct StoredMessage {
36    /// Message role (system, user, assistant, tool)
37    pub role: String,
38    /// Message content
39    pub content: String,
40    /// Unix timestamp when the message was created
41    pub created_at: u64,
42    /// Optional token count for budget tracking
43    pub token_count: Option<u64>,
44}
45
46/// A stored conversation session
47#[derive(Debug, Clone, Serialize, Deserialize)]
48pub struct StoredSession {
49    /// Unique session identifier
50    pub session_id: String,
51    /// Human-readable title (auto-generated or user-set)
52    pub title: String,
53    /// System prompt used for this session
54    pub system_prompt: String,
55    /// Unix timestamp when the session was created
56    pub created_at: u64,
57    /// Unix timestamp of the last activity
58    pub updated_at: u64,
59    /// Total token count across all messages
60    pub total_tokens: u64,
61    /// Number of messages in the session
62    pub message_count: u64,
63}
64
65/// Retention policy for pruning old conversations
66#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
67pub enum RetentionPolicy {
68    /// Keep messages newer than this duration
69    TimeBased(Duration),
70    /// Keep at most this many messages (oldest removed first)
71    CountBased(usize),
72    /// Keep messages until total tokens exceed this budget
73    TokenBudget(u64),
74    /// No retention limit
75    Unlimited,
76}
77
78impl RetentionPolicy {
79    /// Apply this policy to a list of messages, returning the pruned list
80    pub fn apply(&self, messages: &mut Vec<StoredMessage>) {
81        match self {
82            RetentionPolicy::TimeBased(duration) => {
83                let cutoff = SystemTime::now()
84                    .duration_since(UNIX_EPOCH)
85                    .unwrap_or_default()
86                    .as_secs()
87                    - duration.as_secs();
88                messages.retain(|m| m.created_at >= cutoff);
89            }
90            RetentionPolicy::CountBased(max) => {
91                if messages.len() > *max {
92                    // Keep the most recent `max` messages
93                    let keep = messages.split_off(messages.len() - max);
94                    *messages = keep;
95                }
96            }
97            RetentionPolicy::TokenBudget(budget) => {
98                let mut total: u64 = 0;
99                // Keep messages from newest to oldest until budget is exceeded
100                messages.reverse();
101                messages.retain(|m| {
102                    let tokens = m.token_count.unwrap_or(0);
103                    if total + tokens <= *budget {
104                        total += tokens;
105                        true
106                    } else {
107                        false
108                    }
109                });
110                messages.reverse();
111            }
112            RetentionPolicy::Unlimited => {
113                // No pruning
114            }
115        }
116    }
117}
118
119/// SQLite-backed conversation store
120#[derive(Debug)]
121pub struct ConversationStore {
122    conn: Connection,
123}
124
125impl ConversationStore {
126    /// Open or create a SQLite database at the given path.
127    /// Use `:memory:` for an in-memory database (useful for testing).
128    pub fn open<P: AsRef<Path>>(path: P) -> SqlResult<Self> {
129        let conn = Connection::open(path)?;
130        let store = Self { conn };
131        store.initialize_tables()?;
132        Ok(store)
133    }
134
135    /// Initialize the database schema
136    fn initialize_tables(&self) -> SqlResult<()> {
137        self.conn.execute_batch(
138            "
139            CREATE TABLE IF NOT EXISTS sessions (
140                session_id   TEXT PRIMARY KEY,
141                title        TEXT NOT NULL DEFAULT '',
142                system_prompt TEXT NOT NULL DEFAULT '',
143                created_at   INTEGER NOT NULL,
144                updated_at   INTEGER NOT NULL,
145                total_tokens INTEGER NOT NULL DEFAULT 0,
146                message_count INTEGER NOT NULL DEFAULT 0
147            );
148
149            CREATE TABLE IF NOT EXISTS messages (
150                id          INTEGER PRIMARY KEY AUTOINCREMENT,
151                session_id  TEXT NOT NULL,
152                role        TEXT NOT NULL,
153                content     TEXT NOT NULL,
154                created_at  INTEGER NOT NULL,
155                token_count INTEGER DEFAULT NULL,
156                FOREIGN KEY (session_id) REFERENCES sessions(session_id) ON DELETE CASCADE
157            );
158
159            CREATE INDEX IF NOT EXISTS idx_messages_session_id ON messages(session_id);
160            CREATE INDEX IF NOT EXISTS idx_messages_created_at ON messages(created_at);
161            CREATE INDEX IF NOT EXISTS idx_sessions_updated_at ON sessions(updated_at);
162            ",
163        )?;
164        Ok(())
165    }
166
167    /// Create a new conversation session
168    pub fn create_session(&self, session_id: &str, system_prompt: &str) -> SqlResult<()> {
169        let now = SystemTime::now()
170            .duration_since(UNIX_EPOCH)
171            .unwrap_or_default()
172            .as_secs();
173        self.conn.execute(
174            "INSERT OR IGNORE INTO sessions (session_id, system_prompt, created_at, updated_at)
175             VALUES (?1, ?2, ?3, ?3)",
176            params![session_id, system_prompt, now],
177        )?;
178        Ok(())
179    }
180
181    /// Delete a session and all its messages
182    pub fn delete_session(&self, session_id: &str) -> SqlResult<()> {
183        self.conn.execute(
184            "DELETE FROM messages WHERE session_id = ?1",
185            params![session_id],
186        )?;
187        self.conn.execute(
188            "DELETE FROM sessions WHERE session_id = ?1",
189            params![session_id],
190        )?;
191        Ok(())
192    }
193
194    /// List all sessions, ordered by most recently updated first
195    pub fn list_sessions(&self) -> SqlResult<Vec<StoredSession>> {
196        let mut stmt = self.conn.prepare(
197            "SELECT session_id, title, system_prompt, created_at, updated_at, total_tokens, message_count
198             FROM sessions ORDER BY updated_at DESC",
199        )?;
200        let sessions = stmt
201            .query_map([], |row| {
202                Ok(StoredSession {
203                    session_id: row.get(0)?,
204                    title: row.get(1)?,
205                    system_prompt: row.get(2)?,
206                    created_at: row.get(3)?,
207                    updated_at: row.get(4)?,
208                    total_tokens: row.get(5)?,
209                    message_count: row.get(6)?,
210                })
211            })?
212            .collect::<SqlResult<Vec<_>>>()?;
213        Ok(sessions)
214    }
215
216    /// Add a message to a session
217    pub fn add_message(
218        &self,
219        session_id: &str,
220        role: &str,
221        content: &str,
222        token_count: Option<u64>,
223    ) -> SqlResult<()> {
224        let now = SystemTime::now()
225            .duration_since(UNIX_EPOCH)
226            .unwrap_or_default()
227            .as_secs();
228
229        // Insert the message
230        self.conn.execute(
231            "INSERT INTO messages (session_id, role, content, created_at, token_count)
232             VALUES (?1, ?2, ?3, ?4, ?5)",
233            params![session_id, role, content, now, token_count],
234        )?;
235
236        // Update session metadata
237        self.conn.execute(
238            "UPDATE sessions SET
239                updated_at = ?1,
240                total_tokens = total_tokens + ?2,
241                message_count = message_count + 1
242             WHERE session_id = ?3",
243            params![now, token_count.unwrap_or(0), session_id],
244        )?;
245
246        Ok(())
247    }
248
249    /// Get message history for a session, optionally applying a retention policy
250    pub fn get_history(
251        &self,
252        session_id: &str,
253        policy: Option<RetentionPolicy>,
254    ) -> SqlResult<Vec<StoredMessage>> {
255        let mut stmt = self.conn.prepare(
256            "SELECT role, content, created_at, token_count
257             FROM messages WHERE session_id = ?1
258             ORDER BY created_at ASC",
259        )?;
260
261        let mut messages: Vec<StoredMessage> = stmt
262            .query_map(params![session_id], |row| {
263                Ok(StoredMessage {
264                    role: row.get(0)?,
265                    content: row.get(1)?,
266                    created_at: row.get(2)?,
267                    token_count: row.get(3)?,
268                })
269            })?
270            .collect::<SqlResult<Vec<_>>>()?;
271
272        // Apply retention policy if specified
273        if let Some(policy) = policy {
274            policy.apply(&mut messages);
275        }
276
277        Ok(messages)
278    }
279
280    /// Get the number of messages in a session
281    pub fn message_count(&self, session_id: &str) -> SqlResult<u64> {
282        let count: u64 = self
283            .conn
284            .query_row(
285                "SELECT COUNT(*) FROM messages WHERE session_id = ?1",
286                params![session_id],
287                |row| row.get(0),
288            )
289            .unwrap_or(0);
290        Ok(count)
291    }
292
293    /// Get the total token count for a session
294    pub fn total_tokens(&self, session_id: &str) -> SqlResult<u64> {
295        let total: u64 = self
296            .conn
297            .query_row(
298                "SELECT COALESCE(SUM(token_count), 0) FROM messages WHERE session_id = ?1",
299                params![session_id],
300                |row| row.get(0),
301            )
302            .unwrap_or(0);
303        Ok(total)
304    }
305
306    /// Prune old sessions based on a retention policy applied to session age
307    pub fn prune_sessions(&self, max_age: Duration) -> SqlResult<u64> {
308        let cutoff = SystemTime::now()
309            .duration_since(UNIX_EPOCH)
310            .unwrap_or_default()
311            .as_secs()
312            - max_age.as_secs();
313
314        // Find sessions to delete
315        let sessions: Vec<String> = self
316            .conn
317            .prepare("SELECT session_id FROM sessions WHERE updated_at < ?1")?
318            .query_map(params![cutoff], |row| row.get(0))?
319            .collect::<SqlResult<Vec<_>>>()?;
320
321        let count = sessions.len() as u64;
322        for session_id in &sessions {
323            self.delete_session(session_id)?;
324        }
325
326        Ok(count)
327    }
328
329    /// Convert stored messages to `ChatMessage` format for the LLM
330    pub fn to_chat_messages(
331        &self,
332        session_id: &str,
333        policy: Option<RetentionPolicy>,
334    ) -> SqlResult<Vec<crate::llm::ChatMessage>> {
335        let stored = self.get_history(session_id, policy)?;
336        Ok(stored
337            .into_iter()
338            .map(|m| crate::llm::ChatMessage {
339                role: m.role,
340                content: m.content,
341                content_parts: None,
342            })
343            .collect())
344    }
345
346    /// Import messages from a `ConversationMemory` into a session
347    pub fn import_memory(
348        &self,
349        session_id: &str,
350        memory: &crate::agent::ConversationMemory,
351        system_prompt: &str,
352    ) -> SqlResult<()> {
353        self.create_session(session_id, system_prompt)?;
354
355        for msg in memory.history() {
356            self.add_message(session_id, &msg.role, &msg.content, None)?;
357        }
358
359        Ok(())
360    }
361
362    /// Set an explicit title for a session.
363    pub fn set_title(&self, session_id: &str, title: &str) -> SqlResult<()> {
364        let now = SystemTime::now()
365            .duration_since(UNIX_EPOCH)
366            .unwrap_or_default()
367            .as_secs();
368        self.conn.execute(
369            "UPDATE sessions SET title = ?1, updated_at = ?2 WHERE session_id = ?3",
370            params![title, now, session_id],
371        )?;
372        Ok(())
373    }
374
375    /// Get the title of a session (empty string if untitled).
376    pub fn get_title(&self, session_id: &str) -> SqlResult<String> {
377        let title: String = self
378            .conn
379            .query_row(
380                "SELECT title FROM sessions WHERE session_id = ?1",
381                params![session_id],
382                |row| row.get(0),
383            )
384            .unwrap_or_default();
385        Ok(title)
386    }
387
388    /// Auto-title a session from its first user message (truncated to `max_len`).
389    ///
390    /// Returns the assigned title, or `None` if the session has no user message.
391    pub fn auto_title(&self, session_id: &str, max_len: usize) -> SqlResult<Option<String>> {
392        let first: Option<String> = self
393            .conn
394            .query_row(
395                "SELECT content FROM messages WHERE session_id = ?1 AND role = 'user'
396                 ORDER BY created_at ASC LIMIT 1",
397                params![session_id],
398                |row| row.get(0),
399            )
400            .optional()?;
401
402        let Some(first) = first else {
403            return Ok(None);
404        };
405
406        let title = truncate_to_char_boundary(&first, max_len);
407        self.set_title(session_id, &title)?;
408        Ok(Some(title))
409    }
410
411    /// Search conversations by keyword across titles, system prompts, and message
412    /// content. Returns matching session IDs (deduplicated), most recently
413    /// updated first.
414    pub fn search_conversations(&self, query: &str) -> SqlResult<Vec<StoredSession>> {
415        let pattern = format!("%{}%", query);
416        let mut stmt = self.conn.prepare(
417            "SELECT DISTINCT s.session_id, s.title, s.system_prompt, s.created_at, s.updated_at,
418                    s.total_tokens, s.message_count
419             FROM sessions s
420             LEFT JOIN messages m ON m.session_id = s.session_id
421             WHERE s.title LIKE ?1 COLLATE NOCASE
422                OR s.system_prompt LIKE ?1 COLLATE NOCASE
423                OR m.content LIKE ?1 COLLATE NOCASE
424             ORDER BY s.updated_at DESC",
425        )?;
426
427        let results = stmt
428            .query_map(params![pattern], |row| {
429                Ok(StoredSession {
430                    session_id: row.get(0)?,
431                    title: row.get(1)?,
432                    system_prompt: row.get(2)?,
433                    created_at: row.get(3)?,
434                    updated_at: row.get(4)?,
435                    total_tokens: row.get(5)?,
436                    message_count: row.get(6)?,
437                })
438            })?
439            .collect::<SqlResult<Vec<_>>>()?;
440        Ok(results)
441    }
442}
443
444/// Truncate a string to at most `max_len` bytes, breaking on a UTF-8 character
445/// boundary (never splitting a multi-byte character).
446fn truncate_to_char_boundary(s: &str, max_len: usize) -> String {
447    if s.len() <= max_len {
448        return s.to_string();
449    }
450    let mut end = max_len;
451    while end > 0 && !s.is_char_boundary(end) {
452        end -= 1;
453    }
454    s[..end].to_string()
455}
456
457/// A single long-term memory entry (key-value store).
458#[derive(Debug, Clone, Serialize, Deserialize)]
459pub struct MemoryEntry {
460    /// Memory key (unique within its scope)
461    pub key: String,
462    /// Memory value
463    pub value: String,
464    /// Scope (e.g. "user", "global", "project:<id>")
465    pub scope: String,
466    /// Unix timestamp when the memory was created
467    pub created_at: u64,
468    /// Unix timestamp of the last update
469    pub updated_at: u64,
470}
471
472/// Long-term memory store backed by SQLite.
473///
474/// Provides key-value persistence with upsert semantics and scoping, so agents
475/// can retain durable facts about users and projects across sessions — the
476/// "long-term memory" primitive used by OpenClaw/Manus-style assistants.
477///
478/// # Usage
479///
480/// ```rust,no_run
481/// use ravenclaws::persistence::MemoryStore;
482///
483/// let store = MemoryStore::open(":memory:").expect("open memory store");
484/// store.set("user", "favorite_color", "blue").expect("set");
485/// assert_eq!(store.get("user", "favorite_color").unwrap(), Some("blue".to_string()));
486/// ```
487#[derive(Debug)]
488pub struct MemoryStore {
489    conn: Connection,
490}
491
492impl MemoryStore {
493    /// Open or create a SQLite database at the given path for memory storage.
494    /// Use `:memory:` for an in-memory database (useful for testing).
495    pub fn open<P: AsRef<Path>>(path: P) -> SqlResult<Self> {
496        let conn = Connection::open(path)?;
497        let store = Self { conn };
498        store.initialize_tables()?;
499        Ok(store)
500    }
501
502    fn initialize_tables(&self) -> SqlResult<()> {
503        self.conn.execute_batch(
504            "
505            CREATE TABLE IF NOT EXISTS memories (
506                key        TEXT NOT NULL,
507                scope      TEXT NOT NULL,
508                value      TEXT NOT NULL,
509                created_at INTEGER NOT NULL,
510                updated_at INTEGER NOT NULL,
511                PRIMARY KEY (scope, key)
512            );
513
514            CREATE INDEX IF NOT EXISTS idx_memories_scope ON memories(scope);
515            ",
516        )?;
517        Ok(())
518    }
519
520    /// Set (upsert) a memory value for a key within a scope.
521    pub fn set(&self, scope: &str, key: &str, value: &str) -> SqlResult<()> {
522        let now = SystemTime::now()
523            .duration_since(UNIX_EPOCH)
524            .unwrap_or_default()
525            .as_secs();
526        self.conn.execute(
527            "INSERT INTO memories (key, scope, value, created_at, updated_at)
528             VALUES (?1, ?2, ?3, ?4, ?4)
529             ON CONFLICT(scope, key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at",
530            params![key, scope, value, now],
531        )?;
532        Ok(())
533    }
534
535    /// Get a memory value for a key within a scope.
536    pub fn get(&self, scope: &str, key: &str) -> SqlResult<Option<String>> {
537        let value: Option<String> = self
538            .conn
539            .query_row(
540                "SELECT value FROM memories WHERE scope = ?1 AND key = ?2",
541                params![scope, key],
542                |row| row.get(0),
543            )
544            .optional()?;
545        Ok(value)
546    }
547
548    /// Delete a memory entry.
549    pub fn delete(&self, scope: &str, key: &str) -> SqlResult<()> {
550        self.conn.execute(
551            "DELETE FROM memories WHERE scope = ?1 AND key = ?2",
552            params![scope, key],
553        )?;
554        Ok(())
555    }
556
557    /// List all memories in a scope (optionally all scopes if `scope` is `None`).
558    pub fn list(&self, scope: Option<&str>) -> SqlResult<Vec<MemoryEntry>> {
559        let mut stmt = if scope.is_some() {
560            self.conn.prepare(
561                "SELECT key, value, scope, created_at, updated_at
562                 FROM memories WHERE scope = ?1 ORDER BY updated_at DESC",
563            )?
564        } else {
565            self.conn.prepare(
566                "SELECT key, value, scope, created_at, updated_at
567                 FROM memories ORDER BY scope, updated_at DESC",
568            )?
569        };
570
571        let entries = if scope.is_some() {
572            stmt.query_map(params![scope], |row| {
573                Ok(MemoryEntry {
574                    key: row.get(0)?,
575                    value: row.get(1)?,
576                    scope: row.get(2)?,
577                    created_at: row.get(3)?,
578                    updated_at: row.get(4)?,
579                })
580            })?
581            .collect::<SqlResult<Vec<_>>>()?
582        } else {
583            stmt.query_map([], |row| {
584                Ok(MemoryEntry {
585                    key: row.get(0)?,
586                    value: row.get(1)?,
587                    scope: row.get(2)?,
588                    created_at: row.get(3)?,
589                    updated_at: row.get(4)?,
590                })
591            })?
592            .collect::<SqlResult<Vec<_>>>()?
593        };
594
595        Ok(entries)
596    }
597}
598
599#[cfg(test)]
600mod tests {
601    use super::*;
602    use std::time::Duration;
603
604    fn create_test_store() -> ConversationStore {
605        ConversationStore::open(":memory:").expect("Failed to create in-memory store")
606    }
607
608    #[test]
609    fn test_create_and_list_sessions() {
610        let store = create_test_store();
611        store.create_session("test-1", "You are helpful.").unwrap();
612        store.create_session("test-2", "You are a poet.").unwrap();
613
614        let sessions = store.list_sessions().unwrap();
615        assert_eq!(sessions.len(), 2);
616        assert_eq!(sessions[0].session_id, "test-2"); // most recent first
617        assert_eq!(sessions[1].session_id, "test-1");
618    }
619
620    #[test]
621    fn test_add_and_get_messages() {
622        let store = create_test_store();
623        store
624            .create_session("session-1", "You are helpful.")
625            .unwrap();
626        store
627            .add_message("session-1", "user", "Hello!", Some(5))
628            .unwrap();
629        store
630            .add_message("session-1", "assistant", "Hi there!", Some(10))
631            .unwrap();
632
633        let history = store.get_history("session-1", None).unwrap();
634        assert_eq!(history.len(), 2);
635        assert_eq!(history[0].role, "user");
636        assert_eq!(history[0].content, "Hello!");
637        assert_eq!(history[0].token_count, Some(5));
638        assert_eq!(history[1].role, "assistant");
639        assert_eq!(history[1].content, "Hi there!");
640        assert_eq!(history[1].token_count, Some(10));
641    }
642
643    #[test]
644    fn test_message_count_and_tokens() {
645        let store = create_test_store();
646        store
647            .create_session("session-1", "You are helpful.")
648            .unwrap();
649        store
650            .add_message("session-1", "user", "Hello!", Some(5))
651            .unwrap();
652        store
653            .add_message("session-1", "assistant", "Hi!", Some(3))
654            .unwrap();
655
656        assert_eq!(store.message_count("session-1").unwrap(), 2);
657        assert_eq!(store.total_tokens("session-1").unwrap(), 8);
658    }
659
660    #[test]
661    fn test_delete_session() {
662        let store = create_test_store();
663        store
664            .create_session("session-1", "You are helpful.")
665            .unwrap();
666        store
667            .add_message("session-1", "user", "Hello!", None)
668            .unwrap();
669
670        store.delete_session("session-1").unwrap();
671        let sessions = store.list_sessions().unwrap();
672        assert_eq!(sessions.len(), 0);
673        assert_eq!(store.message_count("session-1").unwrap(), 0);
674    }
675
676    #[test]
677    fn test_retention_policy_time_based() {
678        let mut messages = vec![
679            StoredMessage {
680                role: "user".into(),
681                content: "old".into(),
682                created_at: 1000,
683                token_count: None,
684            },
685            StoredMessage {
686                role: "user".into(),
687                content: "new".into(),
688                created_at: u64::MAX,
689                token_count: None,
690            },
691        ];
692
693        // Keep messages newer than 1 hour
694        let policy = RetentionPolicy::TimeBased(Duration::from_secs(3600));
695        policy.apply(&mut messages);
696
697        // Only the "new" message (with far-future timestamp) should remain
698        assert_eq!(messages.len(), 1);
699        assert_eq!(messages[0].content, "new");
700    }
701
702    #[test]
703    fn test_retention_policy_count_based() {
704        let mut messages: Vec<StoredMessage> = (0..10)
705            .map(|i| StoredMessage {
706                role: "user".into(),
707                content: format!("msg-{}", i),
708                created_at: i as u64,
709                token_count: None,
710            })
711            .collect();
712
713        let policy = RetentionPolicy::CountBased(3);
714        policy.apply(&mut messages);
715
716        assert_eq!(messages.len(), 3);
717        assert_eq!(messages[0].content, "msg-7");
718        assert_eq!(messages[2].content, "msg-9");
719    }
720
721    #[test]
722    fn test_retention_policy_token_budget() {
723        let mut messages = vec![
724            StoredMessage {
725                role: "user".into(),
726                content: "a".into(),
727                created_at: 1,
728                token_count: Some(100),
729            },
730            StoredMessage {
731                role: "user".into(),
732                content: "b".into(),
733                created_at: 2,
734                token_count: Some(50),
735            },
736            StoredMessage {
737                role: "user".into(),
738                content: "c".into(),
739                created_at: 3,
740                token_count: Some(30),
741            },
742        ];
743
744        // Budget of 80 tokens — should keep newest messages up to 80 tokens
745        let policy = RetentionPolicy::TokenBudget(80);
746        policy.apply(&mut messages);
747
748        // From newest: c(30) + b(50) = 80, a(100) exceeds budget
749        assert_eq!(messages.len(), 2);
750        assert_eq!(messages[0].content, "b");
751        assert_eq!(messages[1].content, "c");
752    }
753
754    #[test]
755    fn test_retention_policy_unlimited() {
756        let mut messages = vec![
757            StoredMessage {
758                role: "user".into(),
759                content: "a".into(),
760                created_at: 1,
761                token_count: None,
762            },
763            StoredMessage {
764                role: "user".into(),
765                content: "b".into(),
766                created_at: 2,
767                token_count: None,
768            },
769        ];
770
771        let policy = RetentionPolicy::Unlimited;
772        policy.apply(&mut messages);
773        assert_eq!(messages.len(), 2);
774    }
775
776    #[test]
777    fn test_prune_sessions() {
778        let store = create_test_store();
779        store.create_session("old-session", "Old.").unwrap();
780        store.create_session("new-session", "New.").unwrap();
781
782        // Manually set old session's updated_at to the past
783        let past = 1000; // year 1970
784        store
785            .conn
786            .execute(
787                "UPDATE sessions SET updated_at = ?1 WHERE session_id = 'old-session'",
788                params![past],
789            )
790            .unwrap();
791
792        let pruned = store.prune_sessions(Duration::from_secs(3600)).unwrap();
793        assert_eq!(pruned, 1);
794
795        let sessions = store.list_sessions().unwrap();
796        assert_eq!(sessions.len(), 1);
797        assert_eq!(sessions[0].session_id, "new-session");
798    }
799
800    #[test]
801    fn test_to_chat_messages() {
802        let store = create_test_store();
803        store.create_session("s1", "System prompt.").unwrap();
804        store
805            .add_message("s1", "system", "System prompt.", None)
806            .unwrap();
807        store.add_message("s1", "user", "Hello!", None).unwrap();
808
809        let chat_msgs = store.to_chat_messages("s1", None).unwrap();
810        assert_eq!(chat_msgs.len(), 2);
811        assert_eq!(chat_msgs[0].role, "system");
812        assert_eq!(chat_msgs[1].content, "Hello!");
813    }
814
815    #[test]
816    fn test_import_memory() {
817        let store = create_test_store();
818        let mut memory = crate::agent::ConversationMemory::new("System prompt.", 0);
819        memory.add_user_message("Hello!");
820        memory.add_assistant_message("Hi there!");
821
822        store
823            .import_memory("imported-session", &memory, "System prompt.")
824            .unwrap();
825
826        let history = store.get_history("imported-session", None).unwrap();
827        assert_eq!(history.len(), 3); // system + user + assistant
828        assert_eq!(history[0].content, "System prompt.");
829        assert_eq!(history[1].content, "Hello!");
830        assert_eq!(history[2].content, "Hi there!");
831    }
832
833    #[test]
834    fn test_session_metadata_updates() {
835        let store = create_test_store();
836        store.create_session("s1", "Helpful assistant.").unwrap();
837
838        store.add_message("s1", "user", "Hi", Some(3)).unwrap();
839        store
840            .add_message("s1", "assistant", "Hello!", Some(5))
841            .unwrap();
842
843        let sessions = store.list_sessions().unwrap();
844        assert_eq!(sessions.len(), 1);
845        assert_eq!(sessions[0].message_count, 2);
846        assert_eq!(sessions[0].total_tokens, 8);
847    }
848
849    #[test]
850    fn test_nonexistent_session_returns_empty() {
851        let store = create_test_store();
852        let history = store.get_history("nonexistent", None).unwrap();
853        assert!(history.is_empty());
854        assert_eq!(store.message_count("nonexistent").unwrap(), 0);
855        assert_eq!(store.total_tokens("nonexistent").unwrap(), 0);
856    }
857
858    // ── Auto-title tests ───────────────────────────────────────────────────
859
860    #[test]
861    fn test_auto_title_from_first_user_message() {
862        let store = create_test_store();
863        store.create_session("s1", "System.").unwrap();
864        store
865            .add_message("s1", "user", "Hello there friend", None)
866            .unwrap();
867        store.add_message("s1", "assistant", "Hi!", None).unwrap();
868
869        let title = store.auto_title("s1", 40).unwrap().unwrap();
870        assert_eq!(title, "Hello there friend");
871        assert_eq!(store.get_title("s1").unwrap(), "Hello there friend");
872    }
873
874    #[test]
875    fn test_auto_title_truncates_to_max_len() {
876        let store = create_test_store();
877        store.create_session("s1", "System.").unwrap();
878        store
879            .add_message("s1", "user", "This is a very long first message", None)
880            .unwrap();
881
882        let title = store.auto_title("s1", 10).unwrap().unwrap();
883        assert_eq!(title, "This is a ");
884        assert_eq!(store.get_title("s1").unwrap(), "This is a ");
885    }
886
887    #[test]
888    fn test_auto_title_no_user_message_returns_none() {
889        let store = create_test_store();
890        store.create_session("s1", "System.").unwrap();
891        assert!(store.auto_title("s1", 40).unwrap().is_none());
892    }
893
894    #[test]
895    fn test_set_and_get_title() {
896        let store = create_test_store();
897        store.create_session("s1", "System.").unwrap();
898        store.set_title("s1", "My custom title").unwrap();
899        assert_eq!(store.get_title("s1").unwrap(), "My custom title");
900    }
901
902    // ── Search tests ───────────────────────────────────────────────────────
903
904    #[test]
905    fn test_search_by_message_content() {
906        let store = create_test_store();
907        store.create_session("s1", "System.").unwrap();
908        store
909            .add_message("s1", "user", "The capital of Norway", None)
910            .unwrap();
911        store.create_session("s2", "System.").unwrap();
912        store
913            .add_message("s2", "user", "Something unrelated", None)
914            .unwrap();
915
916        let results = store.search_conversations("Norway").unwrap();
917        assert_eq!(results.len(), 1);
918        assert_eq!(results[0].session_id, "s1");
919    }
920
921    #[test]
922    fn test_search_by_title() {
923        let store = create_test_store();
924        store.create_session("s1", "System.").unwrap();
925        store.set_title("s1", "Deployment checklist").unwrap();
926        store.create_session("s2", "System.").unwrap();
927
928        let results = store.search_conversations("deployment").unwrap();
929        assert_eq!(results.len(), 1);
930        assert_eq!(results[0].session_id, "s1");
931    }
932
933    #[test]
934    fn test_search_no_match_returns_empty() {
935        let store = create_test_store();
936        store.create_session("s1", "System.").unwrap();
937        store.add_message("s1", "user", "hello", None).unwrap();
938
939        let results = store.search_conversations("zzzzz").unwrap();
940        assert!(results.is_empty());
941    }
942
943    // ── MemoryStore tests ──────────────────────────────────────────────────
944
945    fn create_test_memory_store() -> MemoryStore {
946        MemoryStore::open(":memory:").expect("Failed to create in-memory memory store")
947    }
948
949    #[test]
950    fn test_memory_set_and_get() {
951        let store = create_test_memory_store();
952        store.set("user", "name", "Alice").unwrap();
953        assert_eq!(
954            store.get("user", "name").unwrap(),
955            Some("Alice".to_string())
956        );
957        assert_eq!(store.get("user", "missing").unwrap(), None);
958    }
959
960    #[test]
961    fn test_memory_upsert() {
962        let store = create_test_memory_store();
963        store.set("user", "name", "Alice").unwrap();
964        store.set("user", "name", "Bob").unwrap();
965        assert_eq!(store.get("user", "name").unwrap(), Some("Bob".to_string()));
966    }
967
968    #[test]
969    fn test_memory_scoped() {
970        let store = create_test_memory_store();
971        store.set("user", "name", "Alice").unwrap();
972        store.set("project:1", "name", "Bob").unwrap();
973        // Same key, different scopes are independent
974        assert_eq!(
975            store.get("user", "name").unwrap(),
976            Some("Alice".to_string())
977        );
978        assert_eq!(
979            store.get("project:1", "name").unwrap(),
980            Some("Bob".to_string())
981        );
982    }
983
984    #[test]
985    fn test_memory_delete() {
986        let store = create_test_memory_store();
987        store.set("user", "name", "Alice").unwrap();
988        store.delete("user", "name").unwrap();
989        assert_eq!(store.get("user", "name").unwrap(), None);
990    }
991
992    #[test]
993    fn test_memory_list() {
994        let store = create_test_memory_store();
995        store.set("user", "a", "1").unwrap();
996        store.set("user", "b", "2").unwrap();
997        store.set("global", "c", "3").unwrap();
998
999        let user_entries = store.list(Some("user")).unwrap();
1000        assert_eq!(user_entries.len(), 2);
1001
1002        let all_entries = store.list(None).unwrap();
1003        assert_eq!(all_entries.len(), 3);
1004    }
1005}