Skip to main content

robit_agent/
storage.rs

1//! Session and message storage helpers.
2
3use std::path::{Path, PathBuf};
4
5use async_openai::types::chat::{
6    ChatCompletionRequestAssistantMessage, ChatCompletionRequestMessage,
7    ChatCompletionRequestToolMessage, ChatCompletionRequestUserMessage,
8    ChatCompletionRequestUserMessageContent,
9};
10use rusqlite::{params, types::ToSqlOutput, Connection, Result as SqliteResult};
11use serde::{Deserialize, Serialize};
12
13use crate::datetime::current_timestamp;
14use crate::error::Result;
15
16const ROBIT_DIR: &str = ".robit";
17const MEMORY_DIR: &str = "memory";
18const DB_FILE: &str = "robit.db";
19
20/// Resolve the session database path for a working directory and storage scope.
21pub fn resolve_db_path(working_dir: &Path, global_storage: bool) -> Result<PathBuf> {
22    if global_storage {
23        let home = dirs::home_dir().ok_or_else(|| {
24            crate::error::AgentError::InternalError("Cannot determine home directory".to_string())
25        })?;
26        Ok(home.join(ROBIT_DIR).join(MEMORY_DIR).join(DB_FILE))
27    } else {
28        Ok(working_dir.join(ROBIT_DIR).join(MEMORY_DIR).join(DB_FILE))
29    }
30}
31
32/// Session metadata returned to frontends.
33#[derive(Debug, Clone, Serialize)]
34pub struct SessionInfo {
35    pub id: String,
36    /// Platform chat identifier (None for GUI/TUI, Some for Bot platforms).
37    pub chat_id: Option<String>,
38    pub title: String,
39    pub model: String,
40    /// Which frontend created the session: "gui" | "tui" | "qq" | "feishu".
41    pub source: String,
42    pub status: String, // "idle" | "ready" | "running"
43    pub created_at: String,
44    pub updated_at: String,
45}
46
47/// Message data returned to frontends.
48#[derive(Debug, Clone, Serialize, Deserialize)]
49pub struct MessageData {
50    pub id: i64,
51    pub role: String,
52    pub content: String,
53    pub tool_name: Option<String>,
54    pub tool_call_id: Option<String>,
55    pub tool_info: Option<serde_json::Value>,
56    pub created_at: String,
57}
58
59/// Tool call info for storage in message.
60#[derive(Debug, Clone, Serialize, Deserialize)]
61pub struct ToolCallInfoData {
62    pub tool_call_id: String,
63    pub name: String,
64    pub arguments: String,
65    pub status: String,
66    pub output: Option<String>,
67    pub requires_confirm: bool,
68}
69
70/// Current schema version. Increment when the schema changes.
71const CURRENT_SCHEMA_VERSION: i32 = 4;
72
73// ============================================================================
74// Memory data structures
75// ============================================================================
76
77/// Type of memory entry.
78#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
79pub enum MemoryType {
80    /// Objective fact (e.g., "User likes Rust").
81    Fact,
82    /// User preference (e.g., "Prefer deepseek-chat").
83    Preference,
84    /// Note or documentation (e.g., "Project directory structure").
85    Note,
86    /// Task record (e.g., "Completed XX last time").
87    Task,
88    /// Custom type with a name.
89    Custom(String),
90}
91
92impl MemoryType {
93    pub fn as_str(&self) -> &str {
94        match self {
95            MemoryType::Fact => "fact",
96            MemoryType::Preference => "preference",
97            MemoryType::Note => "note",
98            MemoryType::Task => "task",
99            MemoryType::Custom(s) => s,
100        }
101    }
102
103    pub fn from_str(s: &str) -> Self {
104        match s.to_lowercase().as_str() {
105            "fact" => MemoryType::Fact,
106            "preference" => MemoryType::Preference,
107            "note" => MemoryType::Note,
108            "task" => MemoryType::Task,
109            _ => MemoryType::Custom(s.to_string()),
110        }
111    }
112}
113
114/// A memory entry stored in the database.
115#[derive(Debug, Clone, Serialize, Deserialize)]
116pub struct Memory {
117    /// Unique ID (UUID v4).
118    pub id: String,
119    /// Associated session ID (optional).
120    pub session_id: Option<String>,
121    /// Associated chat ID (for Bot platforms, optional).
122    pub chat_id: Option<String>,
123    /// Type of memory.
124    pub memory_type: MemoryType,
125    /// Short title for the memory (for retrieval).
126    pub title: String,
127    /// Full content of the memory.
128    pub content: String,
129    /// Tags for categorization and filtering.
130    pub tags: Vec<String>,
131    /// Soft deletion flag.
132    pub is_active: bool,
133    /// ISO 8601 creation timestamp.
134    pub created_at: String,
135    /// ISO 8601 last update timestamp.
136    pub updated_at: String,
137}
138
139/// Filter for memory queries.
140#[derive(Debug, Clone, Default)]
141pub struct MemoryFilter {
142    /// Filter by memory type.
143    pub memory_type: Option<MemoryType>,
144    /// Filter by tags (any match).
145    pub tags: Option<Vec<String>>,
146    /// Filter by session ID.
147    pub session_id: Option<String>,
148    /// Filter by chat ID (Bot platforms).
149    pub chat_id: Option<String>,
150    /// Filter by creation time (only memories created after this time).
151    pub since: Option<String>,
152    /// Only return active memories (default: true).
153    pub only_active: bool,
154}
155
156impl Memory {
157    /// Create a new memory with the given details.
158    pub fn new(
159        title: String,
160        content: String,
161        memory_type: MemoryType,
162        tags: Vec<String>,
163    ) -> Self {
164        let now = current_timestamp();
165        Memory {
166            id: uuid::Uuid::new_v4().to_string(),
167            session_id: None,
168            chat_id: None,
169            memory_type,
170            title,
171            content,
172            tags,
173            is_active: true,
174            created_at: now.clone(),
175            updated_at: now,
176        }
177    }
178
179    /// Attach a session ID to this memory.
180    pub fn with_session_id(mut self, session_id: String) -> Self {
181        self.session_id = Some(session_id);
182        self
183    }
184
185    /// Attach a chat ID to this memory.
186    pub fn with_chat_id(mut self, chat_id: String) -> Self {
187        self.chat_id = Some(chat_id);
188        self
189    }
190}
191
192/// Initialize the database: create tables if needed, then run migrations.
193///
194/// This is the single entry point used by all frontends. It auto-detects
195/// fresh databases (version 0) and existing databases, running the migration
196/// chain to bring them up to [`CURRENT_SCHEMA_VERSION`].
197pub fn init_db(conn: &Connection) -> SqliteResult<()> {
198    ensure_meta_table(conn)?;
199
200    let version = read_schema_version(conn)?;
201
202    if version == 0 {
203        // Fresh database — create everything at the current version in one shot.
204        create_all_tables(conn)?;
205        write_schema_version(conn, CURRENT_SCHEMA_VERSION)?;
206        tracing::info!(
207            "Database initialized at schema v{}",
208            CURRENT_SCHEMA_VERSION
209        );
210        return Ok(());
211    }
212
213    migrate(conn, version, CURRENT_SCHEMA_VERSION)?;
214    Ok(())
215}
216
217/// Detect the schema version of an existing database.
218///
219/// Returns:
220/// - `0` for a truly fresh database (no `sessions` table, no recorded version).
221/// - `1` for a legacy v1 database that predates `_schema_meta` versioning
222///   (tables exist but no version row was ever written).
223/// - `N` for a versioned database, read from `_schema_meta`.
224fn read_schema_version(conn: &Connection) -> SqliteResult<i32> {
225    match conn.query_row(
226        "SELECT value FROM _schema_meta WHERE key = 'version'",
227        [],
228        |row| row.get::<_, String>(0),
229    ) {
230        Ok(v) => v.parse().map_err(|_| {
231            rusqlite::Error::InvalidParameterName(format!("Invalid schema version: {}", v))
232        }),
233        Err(rusqlite::Error::QueryReturnedNoRows) => {
234            // No recorded version — is this a fresh DB or a legacy v1 DB?
235            if sessions_table_exists(conn)? {
236                Ok(1) // tables exist but unversioned → legacy v1
237            } else {
238                Ok(0) // nothing exists → fresh
239            }
240        }
241        Err(e) => Err(e),
242    }
243}
244
245/// Check whether the `sessions` table exists.
246fn sessions_table_exists(conn: &Connection) -> SqliteResult<bool> {
247    let count: i64 = conn.query_row(
248        "SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = 'sessions'",
249        [],
250        |row| row.get(0),
251    )?;
252    Ok(count > 0)
253}
254
255/// Create all tables and indexes at the current schema version (fresh DBs only).
256fn create_all_tables(conn: &Connection) -> SqliteResult<()> {
257    conn.execute_batch(
258        "CREATE TABLE IF NOT EXISTS sessions (
259            id          TEXT PRIMARY KEY,
260            chat_id     TEXT,
261            title       TEXT NOT NULL,
262            model       TEXT NOT NULL,
263            source      TEXT NOT NULL DEFAULT 'gui',
264            created_at  TEXT NOT NULL,
265            updated_at  TEXT NOT NULL,
266            is_active   INTEGER DEFAULT 1
267        );
268
269        CREATE TABLE IF NOT EXISTS messages (
270            id           INTEGER PRIMARY KEY AUTOINCREMENT,
271            session_id   TEXT NOT NULL REFERENCES sessions(id),
272            role         TEXT NOT NULL,
273            content      TEXT NOT NULL,
274            tool_name    TEXT,
275            tool_call_id TEXT,
276            tool_info    TEXT,
277            tokens       INTEGER,
278            created_at   TEXT NOT NULL
279        );
280
281        CREATE INDEX IF NOT EXISTS idx_messages_session
282            ON messages(session_id);
283        CREATE INDEX IF NOT EXISTS idx_messages_created
284            ON messages(session_id, created_at);
285        CREATE UNIQUE INDEX IF NOT EXISTS idx_sessions_chat_id
286            ON sessions(chat_id) WHERE chat_id IS NOT NULL;
287
288        CREATE TABLE IF NOT EXISTS memories (
289            id           TEXT PRIMARY KEY,
290            session_id   TEXT,
291            chat_id      TEXT,
292            memory_type  TEXT NOT NULL,
293            title        TEXT NOT NULL,
294            content      TEXT NOT NULL,
295            tags         TEXT,
296            is_active    INTEGER DEFAULT 1,
297            created_at   TEXT NOT NULL,
298            updated_at   TEXT NOT NULL,
299
300            FOREIGN KEY (session_id) REFERENCES sessions(id) ON DELETE SET NULL
301        );
302
303        CREATE INDEX IF NOT EXISTS idx_memories_type ON memories(memory_type);
304        CREATE INDEX IF NOT EXISTS idx_memories_created ON memories(created_at DESC);
305        CREATE INDEX IF NOT EXISTS idx_memories_session ON memories(session_id);
306        CREATE INDEX IF NOT EXISTS idx_memories_chat ON memories(chat_id);
307        CREATE INDEX IF NOT EXISTS idx_memories_active ON memories(is_active) WHERE is_active = 1;
308
309        -- FTS5 virtual table for full-text search on messages
310        CREATE VIRTUAL TABLE IF NOT EXISTS messages_fts USING fts5(
311            content,
312            content='messages',
313            content_rowid='id',
314            tokenize='unicode61'
315        );
316
317        -- Triggers to keep FTS index in sync with messages table
318        CREATE TRIGGER IF NOT EXISTS messages_ai AFTER INSERT ON messages BEGIN
319            INSERT INTO messages_fts(rowid, content) VALUES (new.id, new.content);
320        END;
321
322        CREATE TRIGGER IF NOT EXISTS messages_ad AFTER DELETE ON messages BEGIN
323            INSERT INTO messages_fts(messages_fts, rowid, content) VALUES ('delete', old.id, old.content);
324        END;
325
326        CREATE TRIGGER IF NOT EXISTS messages_au AFTER UPDATE OF content ON messages BEGIN
327            INSERT INTO messages_fts(messages_fts, rowid, content) VALUES ('delete', old.id, old.content);
328            INSERT INTO messages_fts(rowid, content) VALUES (new.id, new.content);
329        END;",
330    )?;
331    Ok(())
332}
333
334/// Run the migration chain from `from` to `to`, one step at a time.
335fn migrate(conn: &Connection, from: i32, to: i32) -> SqliteResult<()> {
336    let mut current = from;
337    while current < to {
338        tracing::info!("Migrating database: v{} → v{}", current, current + 1);
339        match current {
340            1 => migrate_v1_to_v2(conn)?,
341            2 => migrate_v2_to_v3(conn)?,
342            3 => migrate_v3_to_v4(conn)?,
343            other => {
344                return Err(rusqlite::Error::InvalidParameterName(format!(
345                    "Unknown schema version: {}",
346                    other
347                )))
348            }
349        }
350        current += 1;
351        write_schema_version(conn, current)?;
352        tracing::info!("Database migrated to v{}", current);
353    }
354    Ok(())
355}
356
357/// v1 → v2: add `chat_id`, `source` to sessions; ensure `tool_info` on messages;
358/// add the partial unique index on `sessions.chat_id`.
359fn migrate_v1_to_v2(conn: &Connection) -> SqliteResult<()> {
360    let _ = conn.execute("ALTER TABLE sessions ADD COLUMN chat_id TEXT", []);
361    let _ = conn.execute(
362        "ALTER TABLE sessions ADD COLUMN source TEXT NOT NULL DEFAULT 'gui'",
363        [],
364    );
365    let _ = conn.execute("ALTER TABLE messages ADD COLUMN tool_info TEXT", []);
366    conn.execute_batch(
367        "CREATE UNIQUE INDEX IF NOT EXISTS idx_sessions_chat_id
368            ON sessions(chat_id) WHERE chat_id IS NOT NULL;",
369    )?;
370    Ok(())
371}
372
373/// v2 → v3: add `memories` table for long-term memory.
374fn migrate_v2_to_v3(conn: &Connection) -> SqliteResult<()> {
375    conn.execute_batch(
376        "CREATE TABLE IF NOT EXISTS memories (
377            id           TEXT PRIMARY KEY,
378            session_id   TEXT,
379            chat_id      TEXT,
380            memory_type  TEXT NOT NULL,
381            title        TEXT NOT NULL,
382            content      TEXT NOT NULL,
383            tags         TEXT,
384            is_active    INTEGER DEFAULT 1,
385            created_at   TEXT NOT NULL,
386            updated_at   TEXT NOT NULL,
387
388            FOREIGN KEY (session_id) REFERENCES sessions(id) ON DELETE SET NULL
389        );
390
391        CREATE INDEX IF NOT EXISTS idx_memories_type ON memories(memory_type);
392        CREATE INDEX IF NOT EXISTS idx_memories_created ON memories(created_at DESC);
393        CREATE INDEX IF NOT EXISTS idx_memories_session ON memories(session_id);
394        CREATE INDEX IF NOT EXISTS idx_memories_chat ON memories(chat_id);
395        CREATE INDEX IF NOT EXISTS idx_memories_active ON memories(is_active) WHERE is_active = 1;",
396    )?;
397    Ok(())
398}
399
400/// v3 → v4: add FTS5 virtual table for full-text search on messages.
401fn migrate_v3_to_v4(conn: &Connection) -> SqliteResult<()> {
402    conn.execute_batch(
403        "CREATE VIRTUAL TABLE IF NOT EXISTS messages_fts USING fts5(
404            content,
405            content='messages',
406            content_rowid='id',
407            tokenize='unicode61'
408        );
409
410        CREATE TRIGGER IF NOT EXISTS messages_ai AFTER INSERT ON messages BEGIN
411            INSERT INTO messages_fts(rowid, content) VALUES (new.id, new.content);
412        END;
413
414        CREATE TRIGGER IF NOT EXISTS messages_ad AFTER DELETE ON messages BEGIN
415            INSERT INTO messages_fts(messages_fts, rowid, content) VALUES ('delete', old.id, old.content);
416        END;
417
418        CREATE TRIGGER IF NOT EXISTS messages_au AFTER UPDATE OF content ON messages BEGIN
419            INSERT INTO messages_fts(messages_fts, rowid, content) VALUES ('delete', old.id, old.content);
420            INSERT INTO messages_fts(rowid, content) VALUES (new.id, new.content);
421        END;
422
423        -- Backfill existing messages into the FTS index
424        INSERT INTO messages_fts(rowid, content)
425        SELECT id, content FROM messages;",
426    )?;
427    Ok(())
428}
429
430// ============================================================================
431// Schema version helpers (private)
432// ============================================================================
433
434fn ensure_meta_table(conn: &Connection) -> SqliteResult<()> {
435    conn.execute_batch(
436        "CREATE TABLE IF NOT EXISTS _schema_meta (
437            key   TEXT PRIMARY KEY,
438            value TEXT NOT NULL
439        )",
440    )
441}
442
443fn write_schema_version(conn: &Connection, version: i32) -> SqliteResult<()> {
444    conn.execute(
445        "INSERT OR REPLACE INTO _schema_meta (key, value) VALUES ('version', ?1)",
446        rusqlite::params![version.to_string()],
447    )?;
448    Ok(())
449}
450
451// ============================================================================
452// Session and message accessors (public)
453// ============================================================================
454
455/// Insert a new session.
456///
457/// `chat_id` is `Some` only for Bot platforms (the platform chat identifier);
458/// pass `None` for GUI/TUI sessions. `source` records which frontend created
459/// the session (`"gui"`, `"tui"`, `"qq"`, `"feishu"`).
460pub fn insert_session(
461    conn: &Connection,
462    id: &str,
463    chat_id: Option<&str>,
464    title: &str,
465    model: &str,
466    source: &str,
467) -> SqliteResult<()> {
468    let now = current_timestamp();
469    conn.execute(
470        "INSERT INTO sessions (id, chat_id, title, model, source, created_at, updated_at) \
471         VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)",
472        params![id, chat_id, title, model, source, now, now],
473    )?;
474    Ok(())
475}
476
477/// List all active sessions, ordered by most recently updated.
478///
479/// Pass `Some(source)` to filter by frontend (e.g. `"qq"`); `None` returns
480/// sessions from all sources.
481pub fn list_sessions(
482    conn: &Connection,
483    source_filter: Option<&str>,
484) -> SqliteResult<Vec<SessionInfo>> {
485    let sql = if source_filter.is_some() {
486        "SELECT id, chat_id, title, model, source, created_at, updated_at \
487         FROM sessions WHERE is_active = 1 AND source = ?1 ORDER BY updated_at DESC"
488    } else {
489        "SELECT id, chat_id, title, model, source, created_at, updated_at \
490         FROM sessions WHERE is_active = 1 ORDER BY updated_at DESC"
491    };
492    let mut stmt = conn.prepare(sql)?;
493    let rows = if let Some(source) = source_filter {
494        stmt.query_map(params![source], map_session_row)?
495    } else {
496        stmt.query_map([], map_session_row)?
497    };
498    rows.collect()
499}
500
501/// Find an active session by its platform chat identifier.
502///
503/// Returns `None` for GUI/TUI sessions (where `chat_id` is NULL).
504pub fn find_session_by_chat_id(
505    conn: &Connection,
506    chat_id: &str,
507) -> SqliteResult<Option<SessionInfo>> {
508    let mut stmt = conn.prepare(
509        "SELECT id, chat_id, title, model, source, created_at, updated_at \
510         FROM sessions WHERE chat_id = ?1 AND is_active = 1",
511    )?;
512    let mut rows = stmt.query_map(params![chat_id], map_session_row)?;
513    match rows.next() {
514        Some(Ok(session)) => Ok(Some(session)),
515        _ => Ok(None),
516    }
517}
518
519/// List ALL sessions (including inactive ones) for a chat_id, ordered by most recent first.
520pub fn list_all_sessions_by_chat_id(
521    conn: &Connection,
522    chat_id: &str,
523) -> SqliteResult<Vec<SessionInfo>> {
524    let mut stmt = conn.prepare(
525        "SELECT id, chat_id, title, model, source, created_at, updated_at \
526         FROM sessions WHERE chat_id = ?1 ORDER BY updated_at DESC",
527    )?;
528    let rows = stmt.query_map(params![chat_id], map_session_row)?;
529    rows.collect()
530}
531
532/// Activate a specific session and deactivate others for the same chat_id.
533pub fn activate_session(
534    conn: &Connection,
535    session_id: &str,
536    chat_id: &str,
537) -> SqliteResult<()> {
538    let now = current_timestamp();
539    // First deactivate all sessions for this chat_id
540    conn.execute(
541        "UPDATE sessions SET is_active = 0 WHERE chat_id = ?1",
542        params![chat_id],
543    )?;
544    // Then activate the target session and update its timestamp
545    conn.execute(
546        "UPDATE sessions SET is_active = 1, updated_at = ?1 WHERE id = ?2",
547        params![now, session_id],
548    )?;
549    Ok(())
550}
551
552/// Get a single session by ID.
553pub fn get_session(conn: &Connection, id: &str) -> SqliteResult<Option<SessionInfo>> {
554    let mut stmt = conn.prepare(
555        "SELECT id, chat_id, title, model, source, created_at, updated_at \
556         FROM sessions WHERE id = ?1 AND is_active = 1",
557    )?;
558    let mut rows = stmt.query_map(params![id], map_session_row)?;
559    match rows.next() {
560        Some(Ok(session)) => Ok(Some(session)),
561        _ => Ok(None),
562    }
563}
564
565/// Row mapper shared by all session SELECT queries.
566fn map_session_row(row: &rusqlite::Row<'_>) -> SqliteResult<SessionInfo> {
567    Ok(SessionInfo {
568        id: row.get(0)?,
569        chat_id: row.get(1)?,
570        title: row.get(2)?,
571        model: row.get(3)?,
572        source: row.get(4)?,
573        status: "idle".to_string(),
574        created_at: row.get(5)?,
575        updated_at: row.get(6)?,
576    })
577}
578
579/// Update a session's title.
580pub fn update_session_title(conn: &Connection, id: &str, title: &str) -> SqliteResult<()> {
581    let now = current_timestamp();
582    conn.execute(
583        "UPDATE sessions SET title = ?1, updated_at = ?2 WHERE id = ?3",
584        params![title, now, id],
585    )?;
586    Ok(())
587}
588
589/// Update a session's updated_at timestamp.
590pub fn touch_session(conn: &Connection, id: &str) -> SqliteResult<()> {
591    let now = current_timestamp();
592    conn.execute(
593        "UPDATE sessions SET updated_at = ?1 WHERE id = ?2",
594        params![now, id],
595    )?;
596    Ok(())
597}
598
599/// Soft-delete a session.
600pub fn delete_session(conn: &Connection, id: &str) -> SqliteResult<()> {
601    conn.execute(
602        "UPDATE sessions SET is_active = 0 WHERE id = ?1",
603        params![id],
604    )?;
605    Ok(())
606}
607
608/// Insert a message into a session.
609pub fn insert_message(
610    conn: &Connection,
611    session_id: &str,
612    role: &str,
613    content: &str,
614    tool_name: Option<&str>,
615    tool_call_id: Option<&str>,
616    tool_info: Option<&str>,
617) -> SqliteResult<i64> {
618    let now = current_timestamp();
619    conn.execute(
620        "INSERT INTO messages (session_id, role, content, tool_name, tool_call_id, tool_info, created_at) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)",
621        params![session_id, role, content, tool_name, tool_call_id, tool_info, now],
622    )?;
623    Ok(conn.last_insert_rowid())
624}
625
626/// Get all messages for a session, ordered by creation time.
627pub fn get_messages(conn: &Connection, session_id: &str) -> SqliteResult<Vec<MessageData>> {
628    let mut stmt = conn.prepare(
629        "SELECT id, role, content, tool_name, tool_call_id, tool_info, created_at FROM messages WHERE session_id = ?1 ORDER BY id ASC"
630    )?;
631    let rows = stmt.query_map(params![session_id], |row| {
632        let tool_info_str: Option<String> = row.get(5)?;
633        let tool_info = tool_info_str.and_then(|s| serde_json::from_str(&s).ok());
634        Ok(MessageData {
635            id: row.get(0)?,
636            role: row.get(1)?,
637            content: row.get(2)?,
638            tool_name: row.get(3)?,
639            tool_call_id: row.get(4)?,
640            tool_info,
641            created_at: row.get(6)?,
642        })
643    })?;
644    rows.collect()
645}
646
647/// Update a tool message with its result.
648///
649/// `content` replaces the message content (initially the call arguments,
650/// replaced by the actual tool output once available) so that restored
651/// history carries the real result. `tool_info` is the UI-facing JSON state.
652pub fn update_tool_message(
653    conn: &Connection,
654    session_id: &str,
655    tool_call_id: &str,
656    content: &str,
657    tool_info: &str,
658) -> SqliteResult<()> {
659    conn.execute(
660        "UPDATE messages SET content = ?1, tool_info = ?2 WHERE session_id = ?3 AND tool_call_id = ?4",
661        params![content, tool_info, session_id, tool_call_id],
662    )?;
663    Ok(())
664}
665
666// ============================================================================
667// Message conversion to/from ChatCompletionRequestMessage (new)
668// ============================================================================
669
670/// Convert stored MessageData to ChatCompletionRequestMessage
671pub fn message_to_chat_message(data: &MessageData) -> Result<ChatCompletionRequestMessage> {
672    match data.role.as_str() {
673        "user" => Ok(ChatCompletionRequestMessage::User(
674            ChatCompletionRequestUserMessage {
675                content: ChatCompletionRequestUserMessageContent::Text(data.content.clone()),
676                name: None,
677            }
678            .into(),
679        )),
680        "assistant" => {
681            // Try to parse tool_calls from tool_info
682            let tool_calls = if let Some(tool_info) = &data.tool_info {
683                if let serde_json::Value::Object(obj) = tool_info {
684                    // First try the "tool_calls" field (for complete tool calls)
685                    if let Some(serde_json::Value::Array(arr)) = obj.get("tool_calls") {
686                        use async_openai::types::chat::{
687                            ChatCompletionMessageToolCall, ChatCompletionMessageToolCalls,
688                        };
689                        let mut calls = Vec::new();
690                        for call_val in arr {
691                            if let Ok(call) =
692                                serde_json::from_value::<ChatCompletionMessageToolCall>(
693                                    call_val.clone(),
694                                )
695                            {
696                                calls.push(ChatCompletionMessageToolCalls::Function(call));
697                            }
698                        }
699                        if !calls.is_empty() {
700                            Some(calls)
701                        } else {
702                            None
703                        }
704                    } else {
705                        None
706                    }
707                } else {
708                    None
709                }
710            } else {
711                None
712            };
713
714            let content = if data.content.is_empty() {
715                None
716            } else {
717                Some(data.content.clone().into())
718            };
719
720            // Ensure assistant message has either content or tool_calls
721            if content.is_none() && tool_calls.is_none() {
722                // Skip this invalid message - it will cause API error
723                tracing::warn!("Skipping invalid assistant message: both content and tool_calls are None (message id: {:?})", data.id);
724                return Err(crate::error::AgentError::InternalError("Invalid assistant message".to_string()));
725            }
726
727            Ok(ChatCompletionRequestMessage::Assistant(
728                ChatCompletionRequestAssistantMessage {
729                    content,
730                    name: None,
731                    tool_calls,
732                    refusal: None,
733                    audio: None,
734                    #[allow(deprecated)]
735                    function_call: None,
736                }
737                .into(),
738            ))
739        }
740        "tool" => {
741            let tool_call_id = data.tool_call_id.clone().unwrap_or_default();
742            Ok(ChatCompletionRequestMessage::Tool(
743                ChatCompletionRequestToolMessage {
744                    content: data.content.clone().into(),
745                    tool_call_id,
746                }
747                .into(),
748            ))
749        }
750        // Ignore other roles for now (system is regenerated)
751        _ => Err(crate::error::AgentError::InternalError(format!(
752            "Unknown role: {}",
753            data.role
754        ))),
755    }
756}
757
758// ============================================================================
759// Message full-text search
760// ============================================================================
761
762/// A message search result from FTS full-text search.
763#[derive(Debug, Clone, Serialize)]
764pub struct MessageSearchResult {
765    pub message_id: i64,
766    pub session_id: String,
767    pub session_title: String,
768    pub role: String,
769    pub content_snippet: String,
770    pub created_at: String,
771}
772
773/// Filter for message search queries.
774#[derive(Debug, Clone, Default)]
775pub struct MessageSearchFilter<'a> {
776    /// Search within a specific session (None means all sessions).
777    pub session_id: Option<&'a str>,
778    /// Filter by message role: "user", "assistant", "tool".
779    pub role: Option<&'a str>,
780    /// Only messages created after this ISO 8601 timestamp.
781    pub since: Option<&'a str>,
782    /// Only messages created before this ISO 8601 timestamp.
783    pub until: Option<&'a str>,
784}
785
786/// Search messages using FTS5 full-text search.
787///
788/// `query` is an FTS5 match expression (supports prefix queries with `*`,
789/// phrase queries with quotes, AND/OR/NOT operators).
790///
791/// Returns results ordered by relevance (FTS5 bm25), with snippets.
792pub fn search_messages(
793    conn: &Connection,
794    query: &str,
795    filter: &MessageSearchFilter,
796    limit: usize,
797) -> SqliteResult<Vec<MessageSearchResult>> {
798    if query.trim().is_empty() {
799        return Ok(Vec::new());
800    }
801
802    let mut conditions: Vec<String> = Vec::new();
803    let mut params: Vec<ToSqlOutput> = Vec::new();
804
805    // FTS query is always the first param (bound to ?1)
806    params.push(ToSqlOutput::from(query));
807
808    // Session filter
809    if let Some(session_id) = filter.session_id {
810        conditions.push("m.session_id = ?".to_string());
811        params.push(ToSqlOutput::from(session_id));
812    }
813
814    // Role filter
815    if let Some(role) = filter.role {
816        conditions.push("m.role = ?".to_string());
817        params.push(ToSqlOutput::from(role));
818    }
819
820    // Since filter
821    if let Some(since) = filter.since {
822        conditions.push("m.created_at >= ?".to_string());
823        params.push(ToSqlOutput::from(since));
824    }
825
826    // Until filter
827    if let Some(until) = filter.until {
828        conditions.push("m.created_at <= ?".to_string());
829        params.push(ToSqlOutput::from(until));
830    }
831
832    let where_extra = if conditions.is_empty() {
833        String::new()
834    } else {
835        format!(" AND {}", conditions.join(" AND "))
836    };
837
838    let sql = format!(
839        "SELECT
840            m.id,
841            m.session_id,
842            s.title,
843            m.role,
844            snippet(messages_fts, 0, '<b>', '</b>', '...', 16),
845            m.created_at
846         FROM messages_fts
847         JOIN messages m ON m.id = messages_fts.rowid
848         JOIN sessions s ON s.id = m.session_id
849         WHERE messages_fts MATCH ?1{}
850         ORDER BY bm25(messages_fts)
851         LIMIT {}",
852        where_extra,
853        limit
854    );
855
856    let mut stmt = conn.prepare(&sql)?;
857    let param_refs: Vec<&dyn rusqlite::ToSql> = params.iter().map(|p| p as &dyn rusqlite::ToSql).collect();
858    let rows = stmt.query_map(param_refs.as_slice(), |row| {
859        Ok(MessageSearchResult {
860            message_id: row.get(0)?,
861            session_id: row.get(1)?,
862            session_title: row.get(2)?,
863            role: row.get(3)?,
864            content_snippet: row.get(4)?,
865            created_at: row.get(5)?,
866        })
867    })?;
868
869    rows.collect()
870}
871
872/// Load all messages for a session and convert to chat messages
873pub fn load_chat_messages(
874    conn: &Connection,
875    session_id: &str,
876) -> Result<Vec<ChatCompletionRequestMessage>> {
877    let messages = get_messages(conn, session_id)?;
878    tracing::debug!(
879        "load_chat_messages: session_id={}, loaded {} messages from DB",
880        session_id,
881        messages.len()
882    );
883
884    let mut result = Vec::with_capacity(messages.len());
885    for (idx, msg) in messages.iter().enumerate() {
886        match message_to_chat_message(&msg) {
887            Ok(chat_msg) => {
888                // Double-check that assistant messages are valid before adding
889                let is_valid = match &chat_msg {
890                    ChatCompletionRequestMessage::Assistant(assistant_msg) => {
891                        assistant_msg.content.is_some() || assistant_msg.tool_calls.is_some()
892                    }
893                    _ => true,
894                };
895
896                if is_valid {
897                    tracing::trace!(
898                        "load_chat_messages:   message {}: role={}, content_len={}",
899                        idx,
900                        msg.role,
901                        msg.content.len()
902                    );
903                    result.push(chat_msg);
904                } else {
905                    tracing::warn!(
906                        "Skipping invalid assistant message {} (has neither content nor tool_calls)",
907                        idx
908                    );
909                }
910            }
911            Err(e) => {
912                tracing::warn!("Skipping invalid message {}: {}", idx, e);
913            }
914        }
915    }
916    tracing::debug!("load_chat_messages: successfully converted {} messages", result.len());
917    Ok(result)
918}
919
920// ============================================================================
921// Memory accessors (public)
922// ============================================================================
923
924/// Insert a new memory into the database.
925pub fn insert_memory(conn: &Connection, memory: &Memory) -> SqliteResult<()> {
926    let tags_str = if memory.tags.is_empty() {
927        None
928    } else {
929        Some(memory.tags.join(","))
930    };
931
932    conn.execute(
933        "INSERT INTO memories (
934            id, session_id, chat_id, memory_type, title, content, tags, is_active, created_at, updated_at
935        ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10)",
936        params![
937            memory.id,
938            memory.session_id,
939            memory.chat_id,
940            memory.memory_type.as_str(),
941            memory.title,
942            memory.content,
943            tags_str,
944            memory.is_active,
945            memory.created_at,
946            memory.updated_at,
947        ],
948    )?;
949    Ok(())
950}
951
952/// Update an existing memory.
953pub fn update_memory(conn: &Connection, memory: &Memory) -> SqliteResult<()> {
954    let tags_str = if memory.tags.is_empty() {
955        None
956    } else {
957        Some(memory.tags.join(","))
958    };
959
960    conn.execute(
961        "UPDATE memories SET
962            title = ?1,
963            content = ?2,
964            tags = ?3,
965            memory_type = ?4,
966            updated_at = ?5
967        WHERE id = ?6",
968        params![
969            memory.title,
970            memory.content,
971            tags_str,
972            memory.memory_type.as_str(),
973            current_timestamp(),
974            memory.id,
975        ],
976    )?;
977    Ok(())
978}
979
980/// Deactivate a memory (soft delete).
981pub fn deactivate_memory(conn: &Connection, memory_id: &str) -> SqliteResult<()> {
982    conn.execute(
983        "UPDATE memories SET is_active = 0, updated_at = ?1 WHERE id = ?2",
984        params![current_timestamp(), memory_id],
985    )?;
986    Ok(())
987}
988
989/// Permanently delete a memory (hard delete).
990pub fn delete_memory_permanently(conn: &Connection, memory_id: &str) -> SqliteResult<()> {
991    conn.execute("DELETE FROM memories WHERE id = ?1", params![memory_id])?;
992    Ok(())
993}
994
995/// Get a single memory by ID.
996pub fn get_memory(conn: &Connection, memory_id: &str) -> SqliteResult<Option<Memory>> {
997    let mut stmt = conn.prepare(
998        "SELECT id, session_id, chat_id, memory_type, title, content, tags, is_active, created_at, updated_at
999         FROM memories WHERE id = ?1",
1000    )?;
1001
1002    let mut rows = stmt.query_map(params![memory_id], map_memory_row)?;
1003    rows.next().transpose()
1004}
1005
1006/// Find memories by title (fuzzy match using LIKE).
1007pub fn find_memories_by_title(
1008    conn: &Connection,
1009    title_part: &str,
1010    filter: &MemoryFilter,
1011    limit: Option<usize>,
1012) -> SqliteResult<Vec<Memory>> {
1013    let (sql, params) = build_memory_query(Some(title_part), filter, limit);
1014    let mut stmt = conn.prepare(&sql)?;
1015
1016    let rows = stmt.query_map(rusqlite::params_from_iter(params), map_memory_row)?;
1017    rows.collect()
1018}
1019
1020/// List memories with optional filtering.
1021pub fn list_memories(
1022    conn: &Connection,
1023    filter: &MemoryFilter,
1024    limit: Option<usize>,
1025) -> SqliteResult<Vec<Memory>> {
1026    let (sql, params) = build_memory_query(None, filter, limit);
1027    let mut stmt = conn.prepare(&sql)?;
1028
1029    let rows = stmt.query_map(rusqlite::params_from_iter(params), map_memory_row)?;
1030    rows.collect()
1031}
1032
1033/// Recall memories using keyword search (title, content, tags).
1034pub fn recall_memories(
1035    conn: &Connection,
1036    query: &str,
1037    filter: &MemoryFilter,
1038    limit: usize,
1039) -> SqliteResult<Vec<Memory>> {
1040    // First try exact match in title
1041    let mut results = find_memories_by_title(conn, query, filter, Some(limit))?;
1042
1043    // If we have enough results, return them
1044    if results.len() >= limit {
1045        results.truncate(limit);
1046        return Ok(results);
1047    }
1048
1049    // Otherwise, do a broader search using LIKE on title or content
1050    let remaining = limit - results.len();
1051    let (sql, params) = build_recall_query(query, filter, remaining);
1052    let mut stmt = conn.prepare(&sql)?;
1053
1054    let rows = stmt.query_map(rusqlite::params_from_iter(params), map_memory_row)?;
1055    for row in rows {
1056        let memory = row?;
1057        if !results.iter().any(|m| m.id == memory.id) {
1058            results.push(memory);
1059        }
1060    }
1061
1062    results.truncate(limit);
1063    Ok(results)
1064}
1065
1066// ============================================================================
1067// Memory helpers (private)
1068// ============================================================================
1069
1070fn map_memory_row(row: &rusqlite::Row) -> SqliteResult<Memory> {
1071    let tags_str: Option<String> = row.get(6)?;
1072    let tags = tags_str
1073        .map(|s| {
1074            s.split(',')
1075                .map(|t| t.trim().to_string())
1076                .filter(|t| !t.is_empty())
1077                .collect()
1078        })
1079        .unwrap_or_default();
1080
1081    Ok(Memory {
1082        id: row.get(0)?,
1083        session_id: row.get(1)?,
1084        chat_id: row.get(2)?,
1085        memory_type: MemoryType::from_str(&row.get::<_, String>(3)?),
1086        title: row.get(4)?,
1087        content: row.get(5)?,
1088        tags,
1089        is_active: row.get::<_, i32>(7)? != 0,
1090        created_at: row.get(8)?,
1091        updated_at: row.get(9)?,
1092    })
1093}
1094
1095fn build_memory_query<'a>(
1096    title_search: Option<&'a str>,
1097    filter: &'a MemoryFilter,
1098    limit: Option<usize>,
1099) -> (String, Vec<rusqlite::types::ToSqlOutput<'a>>) {
1100    let mut conditions = Vec::new();
1101    let mut params: Vec<rusqlite::types::ToSqlOutput> = Vec::new();
1102
1103    // Active flag
1104    if filter.only_active {
1105        conditions.push("is_active = 1".to_string());
1106    }
1107
1108    // Type filter
1109    if let Some(memory_type) = &filter.memory_type {
1110        conditions.push("memory_type = ?".to_string());
1111        params.push(memory_type.as_str().into());
1112    }
1113
1114    // Session ID
1115    if let Some(session_id) = &filter.session_id {
1116        conditions.push("session_id = ?".to_string());
1117        params.push(session_id.as_str().into());
1118    }
1119
1120    // Chat ID
1121    if let Some(chat_id) = &filter.chat_id {
1122        conditions.push("chat_id = ?".to_string());
1123        params.push(chat_id.as_str().into());
1124    }
1125
1126    // Since time
1127    if let Some(since) = &filter.since {
1128        conditions.push("created_at >= ?".to_string());
1129        params.push(since.as_str().into());
1130    }
1131
1132    // Title search
1133    if let Some(title) = title_search {
1134        conditions.push("title LIKE ?".to_string());
1135        params.push(format!("%{}%", title).into());
1136    }
1137
1138    // Tags filter (any match)
1139    if let Some(tags) = &filter.tags {
1140        if !tags.is_empty() {
1141            let tag_conditions: Vec<_> = tags.iter().map(|_| "tags LIKE ?").collect();
1142            conditions.push(format!("({})", tag_conditions.join(" OR ")));
1143            for tag in tags {
1144                params.push(format!("%{}%", tag).into());
1145            }
1146        }
1147    }
1148
1149    let where_clause = if conditions.is_empty() {
1150        String::new()
1151    } else {
1152        format!("WHERE {}", conditions.join(" AND "))
1153    };
1154
1155    let limit_clause = limit.map(|l| format!("LIMIT {}", l)).unwrap_or_default();
1156
1157    let sql = format!(
1158        "SELECT id, session_id, chat_id, memory_type, title, content, tags, is_active, created_at, updated_at
1159         FROM memories
1160         {}
1161         ORDER BY created_at DESC
1162         {}",
1163        where_clause, limit_clause
1164    );
1165
1166    (sql, params)
1167}
1168
1169fn build_recall_query<'a>(
1170    query: &'a str,
1171    filter: &'a MemoryFilter,
1172    limit: usize,
1173) -> (String, Vec<rusqlite::types::ToSqlOutput<'a>>) {
1174    let mut conditions = Vec::new();
1175    let mut params: Vec<rusqlite::types::ToSqlOutput> = Vec::new();
1176
1177    // Active flag
1178    if filter.only_active {
1179        conditions.push("is_active = 1".to_string());
1180    }
1181
1182    // Type filter
1183    if let Some(memory_type) = &filter.memory_type {
1184        conditions.push("memory_type = ?".to_string());
1185        params.push(memory_type.as_str().into());
1186    }
1187
1188    // Session ID
1189    if let Some(session_id) = &filter.session_id {
1190        conditions.push("session_id = ?".to_string());
1191        params.push(session_id.as_str().into());
1192    }
1193
1194    // Chat ID
1195    if let Some(chat_id) = &filter.chat_id {
1196        conditions.push("chat_id = ?".to_string());
1197        params.push(chat_id.as_str().into());
1198    }
1199
1200    // Keyword search (match in title, content, or tags)
1201    conditions.push("(title LIKE ? OR content LIKE ? OR tags LIKE ?)".to_string());
1202    let pattern = format!("%{}%", query);
1203    params.push(pattern.clone().into());
1204    params.push(pattern.clone().into());
1205    params.push(pattern.into());
1206
1207    let where_clause = format!("WHERE {}", conditions.join(" AND "));
1208
1209    let sql = format!(
1210        "SELECT id, session_id, chat_id, memory_type, title, content, tags, is_active, created_at, updated_at
1211         FROM memories
1212         {}
1213         ORDER BY created_at DESC
1214         LIMIT {}",
1215        where_clause, limit
1216    );
1217
1218    (sql, params)
1219}
1220
1221#[cfg(test)]
1222mod tests {
1223    use super::*;
1224
1225    #[test]
1226    fn resolves_local_db_path() {
1227        let working_dir = PathBuf::from("project");
1228        let path = resolve_db_path(&working_dir, false).unwrap();
1229        assert_eq!(
1230            path,
1231            working_dir.join(ROBIT_DIR).join(MEMORY_DIR).join(DB_FILE)
1232        );
1233    }
1234
1235    #[test]
1236    fn session_crud() {
1237        let conn = Connection::open_in_memory().unwrap();
1238        init_db(&conn).unwrap();
1239
1240        insert_session(
1241            &conn,
1242            "test-123",
1243            None,
1244            "Test Session",
1245            "deepseek/deepseek-chat",
1246            "gui",
1247        )
1248        .unwrap();
1249
1250        let sessions = list_sessions(&conn, None).unwrap();
1251        assert_eq!(sessions.len(), 1);
1252        assert_eq!(sessions[0].id, "test-123");
1253        assert_eq!(sessions[0].title, "Test Session");
1254        assert_eq!(sessions[0].source, "gui");
1255        assert_eq!(sessions[0].chat_id, None);
1256        assert_eq!(sessions[0].status, "idle");
1257
1258        let session = get_session(&conn, "test-123").unwrap().unwrap();
1259        assert_eq!(session.title, "Test Session");
1260        assert_eq!(session.source, "gui");
1261
1262        update_session_title(&conn, "test-123", "Updated Title").unwrap();
1263        let updated = get_session(&conn, "test-123").unwrap().unwrap();
1264        assert_eq!(updated.title, "Updated Title");
1265
1266        delete_session(&conn, "test-123").unwrap();
1267        assert!(get_session(&conn, "test-123").unwrap().is_none());
1268        assert!(list_sessions(&conn, None).unwrap().is_empty());
1269    }
1270
1271    #[test]
1272    fn message_operations() {
1273        let conn = Connection::open_in_memory().unwrap();
1274        init_db(&conn).unwrap();
1275
1276        insert_session(&conn, "session-msg", None, "Chat Session", "model", "gui").unwrap();
1277        let user_id = insert_message(
1278            &conn,
1279            "session-msg",
1280            "user",
1281            "Hello Robit",
1282            None,
1283            None,
1284            None,
1285        )
1286        .unwrap();
1287        let assistant_id = insert_message(
1288            &conn,
1289            "session-msg",
1290            "assistant",
1291            "Hello! How can I help?",
1292            None,
1293            None,
1294            None,
1295        )
1296        .unwrap();
1297
1298        let messages = get_messages(&conn, "session-msg").unwrap();
1299        assert_eq!(messages.len(), 2);
1300        assert_eq!(messages[0].id, user_id);
1301        assert_eq!(messages[0].role, "user");
1302        assert_eq!(messages[0].content, "Hello Robit");
1303        assert_eq!(messages[1].id, assistant_id);
1304        assert_eq!(messages[1].role, "assistant");
1305        assert_eq!(messages[1].content, "Hello! How can I help?");
1306    }
1307
1308    #[test]
1309    fn empty_sessions() {
1310        let conn = Connection::open_in_memory().unwrap();
1311        init_db(&conn).unwrap();
1312
1313        let sessions = list_sessions(&conn, None).unwrap();
1314        assert_eq!(sessions.len(), 0);
1315    }
1316
1317    #[test]
1318    fn get_nonexistent_session() {
1319        let conn = Connection::open_in_memory().unwrap();
1320        init_db(&conn).unwrap();
1321
1322        let session = get_session(&conn, "nonexistent").unwrap();
1323        assert!(session.is_none());
1324    }
1325
1326    #[test]
1327    fn tool_message_update() {
1328        let conn = Connection::open_in_memory().unwrap();
1329        init_db(&conn).unwrap();
1330
1331        insert_session(&conn, "session-tool", None, "Tool Session", "model", "gui").unwrap();
1332        let initial = serde_json::json!({
1333            "tool_call_id": "tool-1",
1334            "name": "bash",
1335            "arguments": "{}",
1336            "status": "pending",
1337            "requires_confirm": true
1338        })
1339        .to_string();
1340        insert_message(
1341            &conn,
1342            "session-tool",
1343            "tool",
1344            "{}",
1345            Some("bash"),
1346            Some("tool-1"),
1347            Some(&initial),
1348        )
1349        .unwrap();
1350
1351        let updated = serde_json::json!({
1352            "tool_call_id": "tool-1",
1353            "status": "success",
1354            "output": "done"
1355        })
1356        .to_string();
1357        update_tool_message(&conn, "session-tool", "tool-1", "done", &updated).unwrap();
1358
1359        let messages = get_messages(&conn, "session-tool").unwrap();
1360        assert_eq!(messages.len(), 1);
1361        assert_eq!(messages[0].tool_name.as_deref(), Some("bash"));
1362        assert_eq!(messages[0].tool_call_id.as_deref(), Some("tool-1"));
1363        // content is replaced with the actual tool output on update
1364        assert_eq!(messages[0].content, "done");
1365        assert_eq!(messages[0].tool_info.as_ref().unwrap()["status"], "success");
1366        assert_eq!(messages[0].tool_info.as_ref().unwrap()["output"], "done");
1367    }
1368
1369    #[test]
1370    fn chat_id_lookup_and_source_filter() {
1371        let conn = Connection::open_in_memory().unwrap();
1372        init_db(&conn).unwrap();
1373
1374        insert_session(&conn, "gui-1", None, "GUI Session", "model", "gui").unwrap();
1375        insert_session(
1376            &conn,
1377            "qq-1",
1378            Some("group:abc"),
1379            "技术讨论群",
1380            "model",
1381            "qq",
1382        )
1383        .unwrap();
1384        insert_session(
1385            &conn,
1386            "qq-2",
1387            Some("private:xyz"),
1388            "私聊",
1389            "model",
1390            "qq",
1391        )
1392        .unwrap();
1393
1394        // find_session_by_chat_id
1395        let found = find_session_by_chat_id(&conn, "group:abc").unwrap().unwrap();
1396        assert_eq!(found.id, "qq-1");
1397        assert_eq!(found.source, "qq");
1398        assert_eq!(found.chat_id.as_deref(), Some("group:abc"));
1399
1400        // chat_id lookup returns None for GUI sessions (NULL chat_id)
1401        assert!(find_session_by_chat_id(&conn, "does-not-exist")
1402            .unwrap()
1403            .is_none());
1404
1405        // source filter
1406        let qq_sessions = list_sessions(&conn, Some("qq")).unwrap();
1407        assert_eq!(qq_sessions.len(), 2);
1408        assert!(qq_sessions.iter().all(|s| s.source == "qq"));
1409
1410        let gui_sessions = list_sessions(&conn, Some("gui")).unwrap();
1411        assert_eq!(gui_sessions.len(), 1);
1412        assert_eq!(gui_sessions[0].id, "gui-1");
1413
1414        // no filter returns all
1415        assert_eq!(list_sessions(&conn, None).unwrap().len(), 3);
1416    }
1417
1418    #[test]
1419    fn chat_id_unique_per_chat() {
1420        let conn = Connection::open_in_memory().unwrap();
1421        init_db(&conn).unwrap();
1422
1423        insert_session(
1424            &conn,
1425            "qq-1",
1426            Some("group:abc"),
1427            "First",
1428            "model",
1429            "qq",
1430        )
1431        .unwrap();
1432        // Inserting a second session with the same chat_id must fail (unique index).
1433        let err = insert_session(&conn, "qq-2", Some("group:abc"), "Second", "model", "qq");
1434        assert!(err.is_err());
1435    }
1436
1437    #[test]
1438    fn migrates_legacy_v1_database() {
1439        let conn = Connection::open_in_memory().unwrap();
1440        // Simulate a legacy v1 database: old schema, no _schema_meta.
1441        conn.execute_batch(
1442            "CREATE TABLE sessions (
1443                id          TEXT PRIMARY KEY,
1444                title       TEXT NOT NULL,
1445                model       TEXT NOT NULL,
1446                created_at  TEXT NOT NULL,
1447                updated_at  TEXT NOT NULL,
1448                is_active   INTEGER DEFAULT 1
1449            );
1450            CREATE TABLE messages (
1451                id           INTEGER PRIMARY KEY AUTOINCREMENT,
1452                session_id   TEXT NOT NULL REFERENCES sessions(id),
1453                role         TEXT NOT NULL,
1454                content      TEXT NOT NULL,
1455                tool_name    TEXT,
1456                tool_call_id TEXT,
1457                tokens       INTEGER,
1458                created_at   TEXT NOT NULL
1459            );",
1460        )
1461        .unwrap();
1462        conn.execute(
1463            "INSERT INTO sessions (id, title, model, created_at, updated_at) \
1464             VALUES ('legacy-1', 'Legacy', 'model', '2020-01-01', '2020-01-01')",
1465            [],
1466        )
1467        .unwrap();
1468
1469        // Run init_db — it should detect v0 → ... actually no _schema_meta means version 0,
1470        // but tables already exist. create_all_tables uses IF NOT EXISTS so it's safe,
1471        // and version is written as current. The legacy row's source defaults to 'gui'.
1472        init_db(&conn).unwrap();
1473
1474        // Schema version is now current.
1475        let v: i32 = read_schema_version(&conn).unwrap();
1476        assert_eq!(v, CURRENT_SCHEMA_VERSION);
1477
1478        // New columns exist and legacy data is preserved.
1479        let session = get_session(&conn, "legacy-1").unwrap().unwrap();
1480        assert_eq!(session.title, "Legacy");
1481        assert_eq!(session.source, "gui");
1482        assert_eq!(session.chat_id, None);
1483    }
1484
1485    #[test]
1486    fn init_db_is_idempotent() {
1487        let conn = Connection::open_in_memory().unwrap();
1488        init_db(&conn).unwrap();
1489        // Running again on an already-current DB must not error.
1490        init_db(&conn).unwrap();
1491        assert_eq!(read_schema_version(&conn).unwrap(), CURRENT_SCHEMA_VERSION);
1492    }
1493
1494    // ========================================================================
1495    // Message search tests
1496    // ========================================================================
1497
1498    fn setup_search_test(conn: &Connection) {
1499        init_db(conn).unwrap();
1500        insert_session(conn, "sess-1", None, "Session One", "model", "gui").unwrap();
1501        insert_session(conn, "sess-2", None, "Session Two", "model", "gui").unwrap();
1502
1503        insert_message(conn, "sess-1", "user", "Hello, how do I write Rust code?", None, None, None).unwrap();
1504        insert_message(conn, "sess-1", "assistant", "To write Rust code, start with cargo new.", None, None, None).unwrap();
1505        insert_message(conn, "sess-1", "user", "What about Python?", None, None, None).unwrap();
1506        insert_message(conn, "sess-1", "assistant", "Python is also a great language.", None, None, None).unwrap();
1507
1508        insert_message(conn, "sess-2", "user", "How to deploy a Rust application?", None, None, None).unwrap();
1509        insert_message(conn, "sess-2", "assistant", "You can deploy Rust apps with Docker.", None, None, None).unwrap();
1510    }
1511
1512    #[test]
1513    fn search_messages_basic() {
1514        let conn = Connection::open_in_memory().unwrap();
1515        setup_search_test(&conn);
1516
1517        let filter = MessageSearchFilter {
1518            session_id: None,
1519            role: None,
1520            since: None,
1521            until: None,
1522        };
1523        let results = search_messages(&conn, "Rust", &filter, 10).unwrap();
1524        assert!(results.len() >= 2, "Expected at least 2 results for 'Rust', got {}", results.len());
1525
1526        // Verify snippet contains highlighting
1527        assert!(results[0].content_snippet.contains("<b>"));
1528    }
1529
1530    #[test]
1531    fn search_messages_session_filter() {
1532        let conn = Connection::open_in_memory().unwrap();
1533        setup_search_test(&conn);
1534
1535        let filter = MessageSearchFilter {
1536            session_id: Some("sess-1"),
1537            role: None,
1538            since: None,
1539            until: None,
1540        };
1541        let results = search_messages(&conn, "Rust", &filter, 10).unwrap();
1542        assert_eq!(results.len(), 2);
1543
1544        let filter2 = MessageSearchFilter {
1545            session_id: Some("sess-2"),
1546            role: None,
1547            since: None,
1548            until: None,
1549        };
1550        let results2 = search_messages(&conn, "Rust", &filter2, 10).unwrap();
1551        assert_eq!(results2.len(), 2);
1552    }
1553
1554    #[test]
1555    fn search_messages_role_filter() {
1556        let conn = Connection::open_in_memory().unwrap();
1557        setup_search_test(&conn);
1558
1559        let filter = MessageSearchFilter {
1560            session_id: Some("sess-1"),
1561            role: Some("user"),
1562            since: None,
1563            until: None,
1564        };
1565        let results = search_messages(&conn, "Rust", &filter, 10).unwrap();
1566        assert_eq!(results.len(), 1);
1567        assert_eq!(results[0].role, "user");
1568    }
1569
1570    #[test]
1571    fn search_messages_empty_query() {
1572        let conn = Connection::open_in_memory().unwrap();
1573        setup_search_test(&conn);
1574
1575        let filter = MessageSearchFilter::default();
1576        let results = search_messages(&conn, "", &filter, 10).unwrap();
1577        assert!(results.is_empty());
1578    }
1579
1580    #[test]
1581    fn search_messages_no_results() {
1582        let conn = Connection::open_in_memory().unwrap();
1583        setup_search_test(&conn);
1584
1585        let filter = MessageSearchFilter::default();
1586        let results = search_messages(&conn, "nonexistent_keyword_xyz", &filter, 10).unwrap();
1587        assert!(results.is_empty());
1588    }
1589
1590    #[test]
1591    fn search_messages_limit() {
1592        let conn = Connection::open_in_memory().unwrap();
1593        setup_search_test(&conn);
1594
1595        let filter = MessageSearchFilter {
1596            session_id: None,
1597            role: None,
1598            since: None,
1599            until: None,
1600        };
1601        let results = search_messages(&conn, "Rust", &filter, 2).unwrap();
1602        assert_eq!(results.len(), 2);
1603    }
1604
1605    #[test]
1606    fn search_messages_cross_session_has_session_title() {
1607        let conn = Connection::open_in_memory().unwrap();
1608        setup_search_test(&conn);
1609
1610        let filter = MessageSearchFilter {
1611            session_id: None,
1612            role: None,
1613            since: None,
1614            until: None,
1615        };
1616        let results = search_messages(&conn, "Rust", &filter, 10).unwrap();
1617        // All results should have a session title
1618        for r in &results {
1619            assert!(!r.session_title.is_empty());
1620        }
1621    }
1622
1623    #[test]
1624    fn migration_v3_to_v4_backfill() {
1625        let conn = Connection::open_in_memory().unwrap();
1626
1627        // Create v3 schema manually (no FTS)
1628        conn.execute_batch(
1629            "CREATE TABLE sessions (
1630                id          TEXT PRIMARY KEY,
1631                chat_id     TEXT,
1632                title       TEXT NOT NULL,
1633                model       TEXT NOT NULL,
1634                source      TEXT NOT NULL DEFAULT 'gui',
1635                created_at  TEXT NOT NULL,
1636                updated_at  TEXT NOT NULL,
1637                is_active   INTEGER DEFAULT 1
1638            );
1639
1640            CREATE TABLE messages (
1641                id           INTEGER PRIMARY KEY AUTOINCREMENT,
1642                session_id   TEXT NOT NULL REFERENCES sessions(id),
1643                role         TEXT NOT NULL,
1644                content      TEXT NOT NULL,
1645                tool_name    TEXT,
1646                tool_call_id TEXT,
1647                tool_info    TEXT,
1648                tokens       INTEGER,
1649                created_at   TEXT NOT NULL
1650            );
1651
1652            CREATE TABLE memories (
1653                id           TEXT PRIMARY KEY,
1654                session_id   TEXT,
1655                chat_id      TEXT,
1656                memory_type  TEXT NOT NULL,
1657                title        TEXT NOT NULL,
1658                content      TEXT NOT NULL,
1659                tags         TEXT,
1660                is_active    INTEGER DEFAULT 1,
1661                created_at   TEXT NOT NULL,
1662                updated_at   TEXT NOT NULL
1663            );
1664
1665            CREATE TABLE _schema_meta (
1666                key   TEXT PRIMARY KEY,
1667                value TEXT NOT NULL
1668            );
1669
1670            INSERT INTO _schema_meta (key, value) VALUES ('version', '3');
1671
1672            INSERT INTO sessions (id, title, model, source, created_at, updated_at)
1673            VALUES ('old-sess', 'Old Session', 'model', 'gui', '2025-01-01', '2025-01-01');
1674
1675            INSERT INTO messages (session_id, role, content, created_at)
1676            VALUES ('old-sess', 'user', 'This is a legacy message about testing', '2025-01-01');",
1677        )
1678        .unwrap();
1679
1680        // Run init_db — should migrate v3 -> v4 and backfill
1681        init_db(&conn).unwrap();
1682
1683        assert_eq!(read_schema_version(&conn).unwrap(), 4);
1684
1685        // Verify search works on the legacy message
1686        let filter = MessageSearchFilter {
1687            session_id: Some("old-sess"),
1688            role: None,
1689            since: None,
1690            until: None,
1691        };
1692        let results = search_messages(&conn, "legacy", &filter, 10).unwrap();
1693        assert_eq!(results.len(), 1);
1694        assert_eq!(results[0].session_title, "Old Session");
1695    }
1696}