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