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 result.
648pub fn update_tool_message(
649    conn: &Connection,
650    session_id: &str,
651    tool_call_id: &str,
652    tool_info: &str,
653) -> SqliteResult<()> {
654    conn.execute(
655        "UPDATE messages SET tool_info = ?1 WHERE session_id = ?2 AND tool_call_id = ?3",
656        params![tool_info, session_id, tool_call_id],
657    )?;
658    Ok(())
659}
660
661// ============================================================================
662// Message conversion to/from ChatCompletionRequestMessage (new)
663// ============================================================================
664
665/// Convert stored MessageData to ChatCompletionRequestMessage
666pub fn message_to_chat_message(data: &MessageData) -> Result<ChatCompletionRequestMessage> {
667    match data.role.as_str() {
668        "user" => Ok(ChatCompletionRequestMessage::User(
669            ChatCompletionRequestUserMessage {
670                content: ChatCompletionRequestUserMessageContent::Text(data.content.clone()),
671                name: None,
672            }
673            .into(),
674        )),
675        "assistant" => {
676            // Try to parse tool_calls from tool_info
677            let tool_calls = if let Some(tool_info) = &data.tool_info {
678                if let serde_json::Value::Object(obj) = tool_info {
679                    // First try the "tool_calls" field (for complete tool calls)
680                    if let Some(serde_json::Value::Array(arr)) = obj.get("tool_calls") {
681                        use async_openai::types::chat::{
682                            ChatCompletionMessageToolCall, ChatCompletionMessageToolCalls,
683                        };
684                        let mut calls = Vec::new();
685                        for call_val in arr {
686                            if let Ok(call) =
687                                serde_json::from_value::<ChatCompletionMessageToolCall>(
688                                    call_val.clone(),
689                                )
690                            {
691                                calls.push(ChatCompletionMessageToolCalls::Function(call));
692                            }
693                        }
694                        if !calls.is_empty() {
695                            Some(calls)
696                        } else {
697                            None
698                        }
699                    } else {
700                        None
701                    }
702                } else {
703                    None
704                }
705            } else {
706                None
707            };
708
709            let content = if data.content.is_empty() {
710                None
711            } else {
712                Some(data.content.clone().into())
713            };
714
715            // Ensure assistant message has either content or tool_calls
716            if content.is_none() && tool_calls.is_none() {
717                // Skip this invalid message - it will cause API error
718                tracing::warn!("Skipping invalid assistant message: both content and tool_calls are None (message id: {:?})", data.id);
719                return Err(crate::error::AgentError::InternalError("Invalid assistant message".to_string()));
720            }
721
722            Ok(ChatCompletionRequestMessage::Assistant(
723                ChatCompletionRequestAssistantMessage {
724                    content,
725                    name: None,
726                    tool_calls,
727                    refusal: None,
728                    audio: None,
729                    #[allow(deprecated)]
730                    function_call: None,
731                }
732                .into(),
733            ))
734        }
735        "tool" => {
736            let tool_call_id = data.tool_call_id.clone().unwrap_or_default();
737            Ok(ChatCompletionRequestMessage::Tool(
738                ChatCompletionRequestToolMessage {
739                    content: data.content.clone().into(),
740                    tool_call_id,
741                }
742                .into(),
743            ))
744        }
745        // Ignore other roles for now (system is regenerated)
746        _ => Err(crate::error::AgentError::InternalError(format!(
747            "Unknown role: {}",
748            data.role
749        ))),
750    }
751}
752
753// ============================================================================
754// Message full-text search
755// ============================================================================
756
757/// A message search result from FTS full-text search.
758#[derive(Debug, Clone, Serialize)]
759pub struct MessageSearchResult {
760    pub message_id: i64,
761    pub session_id: String,
762    pub session_title: String,
763    pub role: String,
764    pub content_snippet: String,
765    pub created_at: String,
766}
767
768/// Filter for message search queries.
769#[derive(Debug, Clone, Default)]
770pub struct MessageSearchFilter<'a> {
771    /// Search within a specific session (None means all sessions).
772    pub session_id: Option<&'a str>,
773    /// Filter by message role: "user", "assistant", "tool".
774    pub role: Option<&'a str>,
775    /// Only messages created after this ISO 8601 timestamp.
776    pub since: Option<&'a str>,
777    /// Only messages created before this ISO 8601 timestamp.
778    pub until: Option<&'a str>,
779}
780
781/// Search messages using FTS5 full-text search.
782///
783/// `query` is an FTS5 match expression (supports prefix queries with `*`,
784/// phrase queries with quotes, AND/OR/NOT operators).
785///
786/// Returns results ordered by relevance (FTS5 bm25), with snippets.
787pub fn search_messages(
788    conn: &Connection,
789    query: &str,
790    filter: &MessageSearchFilter,
791    limit: usize,
792) -> SqliteResult<Vec<MessageSearchResult>> {
793    if query.trim().is_empty() {
794        return Ok(Vec::new());
795    }
796
797    let mut conditions: Vec<String> = Vec::new();
798    let mut params: Vec<ToSqlOutput> = Vec::new();
799
800    // FTS query is always the first param (bound to ?1)
801    params.push(ToSqlOutput::from(query));
802
803    // Session filter
804    if let Some(session_id) = filter.session_id {
805        conditions.push("m.session_id = ?".to_string());
806        params.push(ToSqlOutput::from(session_id));
807    }
808
809    // Role filter
810    if let Some(role) = filter.role {
811        conditions.push("m.role = ?".to_string());
812        params.push(ToSqlOutput::from(role));
813    }
814
815    // Since filter
816    if let Some(since) = filter.since {
817        conditions.push("m.created_at >= ?".to_string());
818        params.push(ToSqlOutput::from(since));
819    }
820
821    // Until filter
822    if let Some(until) = filter.until {
823        conditions.push("m.created_at <= ?".to_string());
824        params.push(ToSqlOutput::from(until));
825    }
826
827    let where_extra = if conditions.is_empty() {
828        String::new()
829    } else {
830        format!(" AND {}", conditions.join(" AND "))
831    };
832
833    let sql = format!(
834        "SELECT
835            m.id,
836            m.session_id,
837            s.title,
838            m.role,
839            snippet(messages_fts, 0, '<b>', '</b>', '...', 16),
840            m.created_at
841         FROM messages_fts
842         JOIN messages m ON m.id = messages_fts.rowid
843         JOIN sessions s ON s.id = m.session_id
844         WHERE messages_fts MATCH ?1{}
845         ORDER BY bm25(messages_fts)
846         LIMIT {}",
847        where_extra,
848        limit
849    );
850
851    let mut stmt = conn.prepare(&sql)?;
852    let param_refs: Vec<&dyn rusqlite::ToSql> = params.iter().map(|p| p as &dyn rusqlite::ToSql).collect();
853    let rows = stmt.query_map(param_refs.as_slice(), |row| {
854        Ok(MessageSearchResult {
855            message_id: row.get(0)?,
856            session_id: row.get(1)?,
857            session_title: row.get(2)?,
858            role: row.get(3)?,
859            content_snippet: row.get(4)?,
860            created_at: row.get(5)?,
861        })
862    })?;
863
864    rows.collect()
865}
866
867/// Load all messages for a session and convert to chat messages
868pub fn load_chat_messages(
869    conn: &Connection,
870    session_id: &str,
871) -> Result<Vec<ChatCompletionRequestMessage>> {
872    let messages = get_messages(conn, session_id)?;
873    tracing::debug!(
874        "load_chat_messages: session_id={}, loaded {} messages from DB",
875        session_id,
876        messages.len()
877    );
878
879    let mut result = Vec::with_capacity(messages.len());
880    for (idx, msg) in messages.iter().enumerate() {
881        match message_to_chat_message(&msg) {
882            Ok(chat_msg) => {
883                // Double-check that assistant messages are valid before adding
884                let is_valid = match &chat_msg {
885                    ChatCompletionRequestMessage::Assistant(assistant_msg) => {
886                        assistant_msg.content.is_some() || assistant_msg.tool_calls.is_some()
887                    }
888                    _ => true,
889                };
890
891                if is_valid {
892                    tracing::debug!(
893                        "  Message {}: role={}, content_len={}",
894                        idx,
895                        msg.role,
896                        msg.content.len()
897                    );
898                    result.push(chat_msg);
899                } else {
900                    tracing::warn!(
901                        "Skipping invalid assistant message {} (has neither content nor tool_calls)",
902                        idx
903                    );
904                }
905            }
906            Err(e) => {
907                tracing::warn!("Skipping invalid message {}: {}", idx, e);
908            }
909        }
910    }
911    tracing::debug!("load_chat_messages: successfully converted {} messages", result.len());
912    Ok(result)
913}
914
915// ============================================================================
916// Memory accessors (public)
917// ============================================================================
918
919/// Insert a new memory into the database.
920pub fn insert_memory(conn: &Connection, memory: &Memory) -> SqliteResult<()> {
921    let tags_str = if memory.tags.is_empty() {
922        None
923    } else {
924        Some(memory.tags.join(","))
925    };
926
927    conn.execute(
928        "INSERT INTO memories (
929            id, session_id, chat_id, memory_type, title, content, tags, is_active, created_at, updated_at
930        ) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10)",
931        params![
932            memory.id,
933            memory.session_id,
934            memory.chat_id,
935            memory.memory_type.as_str(),
936            memory.title,
937            memory.content,
938            tags_str,
939            memory.is_active,
940            memory.created_at,
941            memory.updated_at,
942        ],
943    )?;
944    Ok(())
945}
946
947/// Update an existing memory.
948pub fn update_memory(conn: &Connection, memory: &Memory) -> SqliteResult<()> {
949    let tags_str = if memory.tags.is_empty() {
950        None
951    } else {
952        Some(memory.tags.join(","))
953    };
954
955    conn.execute(
956        "UPDATE memories SET
957            title = ?1,
958            content = ?2,
959            tags = ?3,
960            memory_type = ?4,
961            updated_at = ?5
962        WHERE id = ?6",
963        params![
964            memory.title,
965            memory.content,
966            tags_str,
967            memory.memory_type.as_str(),
968            current_timestamp(),
969            memory.id,
970        ],
971    )?;
972    Ok(())
973}
974
975/// Deactivate a memory (soft delete).
976pub fn deactivate_memory(conn: &Connection, memory_id: &str) -> SqliteResult<()> {
977    conn.execute(
978        "UPDATE memories SET is_active = 0, updated_at = ?1 WHERE id = ?2",
979        params![current_timestamp(), memory_id],
980    )?;
981    Ok(())
982}
983
984/// Permanently delete a memory (hard delete).
985pub fn delete_memory_permanently(conn: &Connection, memory_id: &str) -> SqliteResult<()> {
986    conn.execute("DELETE FROM memories WHERE id = ?1", params![memory_id])?;
987    Ok(())
988}
989
990/// Get a single memory by ID.
991pub fn get_memory(conn: &Connection, memory_id: &str) -> SqliteResult<Option<Memory>> {
992    let mut stmt = conn.prepare(
993        "SELECT id, session_id, chat_id, memory_type, title, content, tags, is_active, created_at, updated_at
994         FROM memories WHERE id = ?1",
995    )?;
996
997    let mut rows = stmt.query_map(params![memory_id], map_memory_row)?;
998    rows.next().transpose()
999}
1000
1001/// Find memories by title (fuzzy match using LIKE).
1002pub fn find_memories_by_title(
1003    conn: &Connection,
1004    title_part: &str,
1005    filter: &MemoryFilter,
1006    limit: Option<usize>,
1007) -> SqliteResult<Vec<Memory>> {
1008    let (sql, params) = build_memory_query(Some(title_part), filter, limit);
1009    let mut stmt = conn.prepare(&sql)?;
1010
1011    let rows = stmt.query_map(rusqlite::params_from_iter(params), map_memory_row)?;
1012    rows.collect()
1013}
1014
1015/// List memories with optional filtering.
1016pub fn list_memories(
1017    conn: &Connection,
1018    filter: &MemoryFilter,
1019    limit: Option<usize>,
1020) -> SqliteResult<Vec<Memory>> {
1021    let (sql, params) = build_memory_query(None, filter, limit);
1022    let mut stmt = conn.prepare(&sql)?;
1023
1024    let rows = stmt.query_map(rusqlite::params_from_iter(params), map_memory_row)?;
1025    rows.collect()
1026}
1027
1028/// Recall memories using keyword search (title, content, tags).
1029pub fn recall_memories(
1030    conn: &Connection,
1031    query: &str,
1032    filter: &MemoryFilter,
1033    limit: usize,
1034) -> SqliteResult<Vec<Memory>> {
1035    // First try exact match in title
1036    let mut results = find_memories_by_title(conn, query, filter, Some(limit))?;
1037
1038    // If we have enough results, return them
1039    if results.len() >= limit {
1040        results.truncate(limit);
1041        return Ok(results);
1042    }
1043
1044    // Otherwise, do a broader search using LIKE on title or content
1045    let remaining = limit - results.len();
1046    let (sql, params) = build_recall_query(query, filter, remaining);
1047    let mut stmt = conn.prepare(&sql)?;
1048
1049    let rows = stmt.query_map(rusqlite::params_from_iter(params), map_memory_row)?;
1050    for row in rows {
1051        let memory = row?;
1052        if !results.iter().any(|m| m.id == memory.id) {
1053            results.push(memory);
1054        }
1055    }
1056
1057    results.truncate(limit);
1058    Ok(results)
1059}
1060
1061// ============================================================================
1062// Memory helpers (private)
1063// ============================================================================
1064
1065fn map_memory_row(row: &rusqlite::Row) -> SqliteResult<Memory> {
1066    let tags_str: Option<String> = row.get(6)?;
1067    let tags = tags_str
1068        .map(|s| {
1069            s.split(',')
1070                .map(|t| t.trim().to_string())
1071                .filter(|t| !t.is_empty())
1072                .collect()
1073        })
1074        .unwrap_or_default();
1075
1076    Ok(Memory {
1077        id: row.get(0)?,
1078        session_id: row.get(1)?,
1079        chat_id: row.get(2)?,
1080        memory_type: MemoryType::from_str(&row.get::<_, String>(3)?),
1081        title: row.get(4)?,
1082        content: row.get(5)?,
1083        tags,
1084        is_active: row.get::<_, i32>(7)? != 0,
1085        created_at: row.get(8)?,
1086        updated_at: row.get(9)?,
1087    })
1088}
1089
1090fn build_memory_query<'a>(
1091    title_search: Option<&'a str>,
1092    filter: &'a MemoryFilter,
1093    limit: Option<usize>,
1094) -> (String, Vec<rusqlite::types::ToSqlOutput<'a>>) {
1095    let mut conditions = Vec::new();
1096    let mut params: Vec<rusqlite::types::ToSqlOutput> = Vec::new();
1097
1098    // Active flag
1099    if filter.only_active {
1100        conditions.push("is_active = 1".to_string());
1101    }
1102
1103    // Type filter
1104    if let Some(memory_type) = &filter.memory_type {
1105        conditions.push("memory_type = ?".to_string());
1106        params.push(memory_type.as_str().into());
1107    }
1108
1109    // Session ID
1110    if let Some(session_id) = &filter.session_id {
1111        conditions.push("session_id = ?".to_string());
1112        params.push(session_id.as_str().into());
1113    }
1114
1115    // Chat ID
1116    if let Some(chat_id) = &filter.chat_id {
1117        conditions.push("chat_id = ?".to_string());
1118        params.push(chat_id.as_str().into());
1119    }
1120
1121    // Since time
1122    if let Some(since) = &filter.since {
1123        conditions.push("created_at >= ?".to_string());
1124        params.push(since.as_str().into());
1125    }
1126
1127    // Title search
1128    if let Some(title) = title_search {
1129        conditions.push("title LIKE ?".to_string());
1130        params.push(format!("%{}%", title).into());
1131    }
1132
1133    // Tags filter (any match)
1134    if let Some(tags) = &filter.tags {
1135        if !tags.is_empty() {
1136            let tag_conditions: Vec<_> = tags.iter().map(|_| "tags LIKE ?").collect();
1137            conditions.push(format!("({})", tag_conditions.join(" OR ")));
1138            for tag in tags {
1139                params.push(format!("%{}%", tag).into());
1140            }
1141        }
1142    }
1143
1144    let where_clause = if conditions.is_empty() {
1145        String::new()
1146    } else {
1147        format!("WHERE {}", conditions.join(" AND "))
1148    };
1149
1150    let limit_clause = limit.map(|l| format!("LIMIT {}", l)).unwrap_or_default();
1151
1152    let sql = format!(
1153        "SELECT id, session_id, chat_id, memory_type, title, content, tags, is_active, created_at, updated_at
1154         FROM memories
1155         {}
1156         ORDER BY created_at DESC
1157         {}",
1158        where_clause, limit_clause
1159    );
1160
1161    (sql, params)
1162}
1163
1164fn build_recall_query<'a>(
1165    query: &'a str,
1166    filter: &'a MemoryFilter,
1167    limit: usize,
1168) -> (String, Vec<rusqlite::types::ToSqlOutput<'a>>) {
1169    let mut conditions = Vec::new();
1170    let mut params: Vec<rusqlite::types::ToSqlOutput> = Vec::new();
1171
1172    // Active flag
1173    if filter.only_active {
1174        conditions.push("is_active = 1".to_string());
1175    }
1176
1177    // Type filter
1178    if let Some(memory_type) = &filter.memory_type {
1179        conditions.push("memory_type = ?".to_string());
1180        params.push(memory_type.as_str().into());
1181    }
1182
1183    // Session ID
1184    if let Some(session_id) = &filter.session_id {
1185        conditions.push("session_id = ?".to_string());
1186        params.push(session_id.as_str().into());
1187    }
1188
1189    // Chat ID
1190    if let Some(chat_id) = &filter.chat_id {
1191        conditions.push("chat_id = ?".to_string());
1192        params.push(chat_id.as_str().into());
1193    }
1194
1195    // Keyword search (match in title, content, or tags)
1196    conditions.push("(title LIKE ? OR content LIKE ? OR tags LIKE ?)".to_string());
1197    let pattern = format!("%{}%", query);
1198    params.push(pattern.clone().into());
1199    params.push(pattern.clone().into());
1200    params.push(pattern.into());
1201
1202    let where_clause = format!("WHERE {}", conditions.join(" AND "));
1203
1204    let sql = format!(
1205        "SELECT id, session_id, chat_id, memory_type, title, content, tags, is_active, created_at, updated_at
1206         FROM memories
1207         {}
1208         ORDER BY created_at DESC
1209         LIMIT {}",
1210        where_clause, limit
1211    );
1212
1213    (sql, params)
1214}
1215
1216#[cfg(test)]
1217mod tests {
1218    use super::*;
1219
1220    #[test]
1221    fn resolves_local_db_path() {
1222        let working_dir = PathBuf::from("project");
1223        let path = resolve_db_path(&working_dir, false).unwrap();
1224        assert_eq!(
1225            path,
1226            working_dir.join(ROBIT_DIR).join(MEMORY_DIR).join(DB_FILE)
1227        );
1228    }
1229
1230    #[test]
1231    fn session_crud() {
1232        let conn = Connection::open_in_memory().unwrap();
1233        init_db(&conn).unwrap();
1234
1235        insert_session(
1236            &conn,
1237            "test-123",
1238            None,
1239            "Test Session",
1240            "deepseek/deepseek-chat",
1241            "gui",
1242        )
1243        .unwrap();
1244
1245        let sessions = list_sessions(&conn, None).unwrap();
1246        assert_eq!(sessions.len(), 1);
1247        assert_eq!(sessions[0].id, "test-123");
1248        assert_eq!(sessions[0].title, "Test Session");
1249        assert_eq!(sessions[0].source, "gui");
1250        assert_eq!(sessions[0].chat_id, None);
1251        assert_eq!(sessions[0].status, "idle");
1252
1253        let session = get_session(&conn, "test-123").unwrap().unwrap();
1254        assert_eq!(session.title, "Test Session");
1255        assert_eq!(session.source, "gui");
1256
1257        update_session_title(&conn, "test-123", "Updated Title").unwrap();
1258        let updated = get_session(&conn, "test-123").unwrap().unwrap();
1259        assert_eq!(updated.title, "Updated Title");
1260
1261        delete_session(&conn, "test-123").unwrap();
1262        assert!(get_session(&conn, "test-123").unwrap().is_none());
1263        assert!(list_sessions(&conn, None).unwrap().is_empty());
1264    }
1265
1266    #[test]
1267    fn message_operations() {
1268        let conn = Connection::open_in_memory().unwrap();
1269        init_db(&conn).unwrap();
1270
1271        insert_session(&conn, "session-msg", None, "Chat Session", "model", "gui").unwrap();
1272        let user_id = insert_message(
1273            &conn,
1274            "session-msg",
1275            "user",
1276            "Hello Robit",
1277            None,
1278            None,
1279            None,
1280        )
1281        .unwrap();
1282        let assistant_id = insert_message(
1283            &conn,
1284            "session-msg",
1285            "assistant",
1286            "Hello! How can I help?",
1287            None,
1288            None,
1289            None,
1290        )
1291        .unwrap();
1292
1293        let messages = get_messages(&conn, "session-msg").unwrap();
1294        assert_eq!(messages.len(), 2);
1295        assert_eq!(messages[0].id, user_id);
1296        assert_eq!(messages[0].role, "user");
1297        assert_eq!(messages[0].content, "Hello Robit");
1298        assert_eq!(messages[1].id, assistant_id);
1299        assert_eq!(messages[1].role, "assistant");
1300        assert_eq!(messages[1].content, "Hello! How can I help?");
1301    }
1302
1303    #[test]
1304    fn empty_sessions() {
1305        let conn = Connection::open_in_memory().unwrap();
1306        init_db(&conn).unwrap();
1307
1308        let sessions = list_sessions(&conn, None).unwrap();
1309        assert_eq!(sessions.len(), 0);
1310    }
1311
1312    #[test]
1313    fn get_nonexistent_session() {
1314        let conn = Connection::open_in_memory().unwrap();
1315        init_db(&conn).unwrap();
1316
1317        let session = get_session(&conn, "nonexistent").unwrap();
1318        assert!(session.is_none());
1319    }
1320
1321    #[test]
1322    fn tool_message_update() {
1323        let conn = Connection::open_in_memory().unwrap();
1324        init_db(&conn).unwrap();
1325
1326        insert_session(&conn, "session-tool", None, "Tool Session", "model", "gui").unwrap();
1327        let initial = serde_json::json!({
1328            "tool_call_id": "tool-1",
1329            "name": "bash",
1330            "arguments": "{}",
1331            "status": "pending",
1332            "requires_confirm": true
1333        })
1334        .to_string();
1335        insert_message(
1336            &conn,
1337            "session-tool",
1338            "tool",
1339            "{}",
1340            Some("bash"),
1341            Some("tool-1"),
1342            Some(&initial),
1343        )
1344        .unwrap();
1345
1346        let updated = serde_json::json!({
1347            "tool_call_id": "tool-1",
1348            "status": "success",
1349            "output": "done"
1350        })
1351        .to_string();
1352        update_tool_message(&conn, "session-tool", "tool-1", &updated).unwrap();
1353
1354        let messages = get_messages(&conn, "session-tool").unwrap();
1355        assert_eq!(messages.len(), 1);
1356        assert_eq!(messages[0].tool_name.as_deref(), Some("bash"));
1357        assert_eq!(messages[0].tool_call_id.as_deref(), Some("tool-1"));
1358        assert_eq!(messages[0].tool_info.as_ref().unwrap()["status"], "success");
1359        assert_eq!(messages[0].tool_info.as_ref().unwrap()["output"], "done");
1360    }
1361
1362    #[test]
1363    fn chat_id_lookup_and_source_filter() {
1364        let conn = Connection::open_in_memory().unwrap();
1365        init_db(&conn).unwrap();
1366
1367        insert_session(&conn, "gui-1", None, "GUI Session", "model", "gui").unwrap();
1368        insert_session(
1369            &conn,
1370            "qq-1",
1371            Some("group:abc"),
1372            "技术讨论群",
1373            "model",
1374            "qq",
1375        )
1376        .unwrap();
1377        insert_session(
1378            &conn,
1379            "qq-2",
1380            Some("private:xyz"),
1381            "私聊",
1382            "model",
1383            "qq",
1384        )
1385        .unwrap();
1386
1387        // find_session_by_chat_id
1388        let found = find_session_by_chat_id(&conn, "group:abc").unwrap().unwrap();
1389        assert_eq!(found.id, "qq-1");
1390        assert_eq!(found.source, "qq");
1391        assert_eq!(found.chat_id.as_deref(), Some("group:abc"));
1392
1393        // chat_id lookup returns None for GUI sessions (NULL chat_id)
1394        assert!(find_session_by_chat_id(&conn, "does-not-exist")
1395            .unwrap()
1396            .is_none());
1397
1398        // source filter
1399        let qq_sessions = list_sessions(&conn, Some("qq")).unwrap();
1400        assert_eq!(qq_sessions.len(), 2);
1401        assert!(qq_sessions.iter().all(|s| s.source == "qq"));
1402
1403        let gui_sessions = list_sessions(&conn, Some("gui")).unwrap();
1404        assert_eq!(gui_sessions.len(), 1);
1405        assert_eq!(gui_sessions[0].id, "gui-1");
1406
1407        // no filter returns all
1408        assert_eq!(list_sessions(&conn, None).unwrap().len(), 3);
1409    }
1410
1411    #[test]
1412    fn chat_id_unique_per_chat() {
1413        let conn = Connection::open_in_memory().unwrap();
1414        init_db(&conn).unwrap();
1415
1416        insert_session(
1417            &conn,
1418            "qq-1",
1419            Some("group:abc"),
1420            "First",
1421            "model",
1422            "qq",
1423        )
1424        .unwrap();
1425        // Inserting a second session with the same chat_id must fail (unique index).
1426        let err = insert_session(&conn, "qq-2", Some("group:abc"), "Second", "model", "qq");
1427        assert!(err.is_err());
1428    }
1429
1430    #[test]
1431    fn migrates_legacy_v1_database() {
1432        let conn = Connection::open_in_memory().unwrap();
1433        // Simulate a legacy v1 database: old schema, no _schema_meta.
1434        conn.execute_batch(
1435            "CREATE TABLE sessions (
1436                id          TEXT PRIMARY KEY,
1437                title       TEXT NOT NULL,
1438                model       TEXT NOT NULL,
1439                created_at  TEXT NOT NULL,
1440                updated_at  TEXT NOT NULL,
1441                is_active   INTEGER DEFAULT 1
1442            );
1443            CREATE TABLE messages (
1444                id           INTEGER PRIMARY KEY AUTOINCREMENT,
1445                session_id   TEXT NOT NULL REFERENCES sessions(id),
1446                role         TEXT NOT NULL,
1447                content      TEXT NOT NULL,
1448                tool_name    TEXT,
1449                tool_call_id TEXT,
1450                tokens       INTEGER,
1451                created_at   TEXT NOT NULL
1452            );",
1453        )
1454        .unwrap();
1455        conn.execute(
1456            "INSERT INTO sessions (id, title, model, created_at, updated_at) \
1457             VALUES ('legacy-1', 'Legacy', 'model', '2020-01-01', '2020-01-01')",
1458            [],
1459        )
1460        .unwrap();
1461
1462        // Run init_db — it should detect v0 → ... actually no _schema_meta means version 0,
1463        // but tables already exist. create_all_tables uses IF NOT EXISTS so it's safe,
1464        // and version is written as current. The legacy row's source defaults to 'gui'.
1465        init_db(&conn).unwrap();
1466
1467        // Schema version is now current.
1468        let v: i32 = read_schema_version(&conn).unwrap();
1469        assert_eq!(v, CURRENT_SCHEMA_VERSION);
1470
1471        // New columns exist and legacy data is preserved.
1472        let session = get_session(&conn, "legacy-1").unwrap().unwrap();
1473        assert_eq!(session.title, "Legacy");
1474        assert_eq!(session.source, "gui");
1475        assert_eq!(session.chat_id, None);
1476    }
1477
1478    #[test]
1479    fn init_db_is_idempotent() {
1480        let conn = Connection::open_in_memory().unwrap();
1481        init_db(&conn).unwrap();
1482        // Running again on an already-current DB must not error.
1483        init_db(&conn).unwrap();
1484        assert_eq!(read_schema_version(&conn).unwrap(), CURRENT_SCHEMA_VERSION);
1485    }
1486
1487    // ========================================================================
1488    // Message search tests
1489    // ========================================================================
1490
1491    fn setup_search_test(conn: &Connection) {
1492        init_db(conn).unwrap();
1493        insert_session(conn, "sess-1", None, "Session One", "model", "gui").unwrap();
1494        insert_session(conn, "sess-2", None, "Session Two", "model", "gui").unwrap();
1495
1496        insert_message(conn, "sess-1", "user", "Hello, how do I write Rust code?", None, None, None).unwrap();
1497        insert_message(conn, "sess-1", "assistant", "To write Rust code, start with cargo new.", None, None, None).unwrap();
1498        insert_message(conn, "sess-1", "user", "What about Python?", None, None, None).unwrap();
1499        insert_message(conn, "sess-1", "assistant", "Python is also a great language.", None, None, None).unwrap();
1500
1501        insert_message(conn, "sess-2", "user", "How to deploy a Rust application?", None, None, None).unwrap();
1502        insert_message(conn, "sess-2", "assistant", "You can deploy Rust apps with Docker.", None, None, None).unwrap();
1503    }
1504
1505    #[test]
1506    fn search_messages_basic() {
1507        let conn = Connection::open_in_memory().unwrap();
1508        setup_search_test(&conn);
1509
1510        let filter = MessageSearchFilter {
1511            session_id: None,
1512            role: None,
1513            since: None,
1514            until: None,
1515        };
1516        let results = search_messages(&conn, "Rust", &filter, 10).unwrap();
1517        assert!(results.len() >= 2, "Expected at least 2 results for 'Rust', got {}", results.len());
1518
1519        // Verify snippet contains highlighting
1520        assert!(results[0].content_snippet.contains("<b>"));
1521    }
1522
1523    #[test]
1524    fn search_messages_session_filter() {
1525        let conn = Connection::open_in_memory().unwrap();
1526        setup_search_test(&conn);
1527
1528        let filter = MessageSearchFilter {
1529            session_id: Some("sess-1"),
1530            role: None,
1531            since: None,
1532            until: None,
1533        };
1534        let results = search_messages(&conn, "Rust", &filter, 10).unwrap();
1535        assert_eq!(results.len(), 2);
1536
1537        let filter2 = MessageSearchFilter {
1538            session_id: Some("sess-2"),
1539            role: None,
1540            since: None,
1541            until: None,
1542        };
1543        let results2 = search_messages(&conn, "Rust", &filter2, 10).unwrap();
1544        assert_eq!(results2.len(), 2);
1545    }
1546
1547    #[test]
1548    fn search_messages_role_filter() {
1549        let conn = Connection::open_in_memory().unwrap();
1550        setup_search_test(&conn);
1551
1552        let filter = MessageSearchFilter {
1553            session_id: Some("sess-1"),
1554            role: Some("user"),
1555            since: None,
1556            until: None,
1557        };
1558        let results = search_messages(&conn, "Rust", &filter, 10).unwrap();
1559        assert_eq!(results.len(), 1);
1560        assert_eq!(results[0].role, "user");
1561    }
1562
1563    #[test]
1564    fn search_messages_empty_query() {
1565        let conn = Connection::open_in_memory().unwrap();
1566        setup_search_test(&conn);
1567
1568        let filter = MessageSearchFilter::default();
1569        let results = search_messages(&conn, "", &filter, 10).unwrap();
1570        assert!(results.is_empty());
1571    }
1572
1573    #[test]
1574    fn search_messages_no_results() {
1575        let conn = Connection::open_in_memory().unwrap();
1576        setup_search_test(&conn);
1577
1578        let filter = MessageSearchFilter::default();
1579        let results = search_messages(&conn, "nonexistent_keyword_xyz", &filter, 10).unwrap();
1580        assert!(results.is_empty());
1581    }
1582
1583    #[test]
1584    fn search_messages_limit() {
1585        let conn = Connection::open_in_memory().unwrap();
1586        setup_search_test(&conn);
1587
1588        let filter = MessageSearchFilter {
1589            session_id: None,
1590            role: None,
1591            since: None,
1592            until: None,
1593        };
1594        let results = search_messages(&conn, "Rust", &filter, 2).unwrap();
1595        assert_eq!(results.len(), 2);
1596    }
1597
1598    #[test]
1599    fn search_messages_cross_session_has_session_title() {
1600        let conn = Connection::open_in_memory().unwrap();
1601        setup_search_test(&conn);
1602
1603        let filter = MessageSearchFilter {
1604            session_id: None,
1605            role: None,
1606            since: None,
1607            until: None,
1608        };
1609        let results = search_messages(&conn, "Rust", &filter, 10).unwrap();
1610        // All results should have a session title
1611        for r in &results {
1612            assert!(!r.session_title.is_empty());
1613        }
1614    }
1615
1616    #[test]
1617    fn migration_v3_to_v4_backfill() {
1618        let conn = Connection::open_in_memory().unwrap();
1619
1620        // Create v3 schema manually (no FTS)
1621        conn.execute_batch(
1622            "CREATE TABLE sessions (
1623                id          TEXT PRIMARY KEY,
1624                chat_id     TEXT,
1625                title       TEXT NOT NULL,
1626                model       TEXT NOT NULL,
1627                source      TEXT NOT NULL DEFAULT 'gui',
1628                created_at  TEXT NOT NULL,
1629                updated_at  TEXT NOT NULL,
1630                is_active   INTEGER DEFAULT 1
1631            );
1632
1633            CREATE TABLE messages (
1634                id           INTEGER PRIMARY KEY AUTOINCREMENT,
1635                session_id   TEXT NOT NULL REFERENCES sessions(id),
1636                role         TEXT NOT NULL,
1637                content      TEXT NOT NULL,
1638                tool_name    TEXT,
1639                tool_call_id TEXT,
1640                tool_info    TEXT,
1641                tokens       INTEGER,
1642                created_at   TEXT NOT NULL
1643            );
1644
1645            CREATE TABLE memories (
1646                id           TEXT PRIMARY KEY,
1647                session_id   TEXT,
1648                chat_id      TEXT,
1649                memory_type  TEXT NOT NULL,
1650                title        TEXT NOT NULL,
1651                content      TEXT NOT NULL,
1652                tags         TEXT,
1653                is_active    INTEGER DEFAULT 1,
1654                created_at   TEXT NOT NULL,
1655                updated_at   TEXT NOT NULL
1656            );
1657
1658            CREATE TABLE _schema_meta (
1659                key   TEXT PRIMARY KEY,
1660                value TEXT NOT NULL
1661            );
1662
1663            INSERT INTO _schema_meta (key, value) VALUES ('version', '3');
1664
1665            INSERT INTO sessions (id, title, model, source, created_at, updated_at)
1666            VALUES ('old-sess', 'Old Session', 'model', 'gui', '2025-01-01', '2025-01-01');
1667
1668            INSERT INTO messages (session_id, role, content, created_at)
1669            VALUES ('old-sess', 'user', 'This is a legacy message about testing', '2025-01-01');",
1670        )
1671        .unwrap();
1672
1673        // Run init_db — should migrate v3 -> v4 and backfill
1674        init_db(&conn).unwrap();
1675
1676        assert_eq!(read_schema_version(&conn).unwrap(), 4);
1677
1678        // Verify search works on the legacy message
1679        let filter = MessageSearchFilter {
1680            session_id: Some("old-sess"),
1681            role: None,
1682            since: None,
1683            until: None,
1684        };
1685        let results = search_messages(&conn, "legacy", &filter, 10).unwrap();
1686        assert_eq!(results.len(), 1);
1687        assert_eq!(results[0].session_title, "Old Session");
1688    }
1689}