1use 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
20pub 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
25pub 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#[derive(Debug, Clone, Serialize)]
45pub struct SessionInfo {
46 pub id: String,
47 pub chat_id: Option<String>,
49 pub title: String,
50 pub model: String,
51 pub source: String,
53 pub status: String, pub created_at: String,
55 pub updated_at: String,
56}
57
58#[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#[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
81const CURRENT_SCHEMA_VERSION: i32 = 5;
83
84#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
90pub enum MemoryType {
91 Fact,
93 Preference,
95 Note,
97 Task,
99 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#[derive(Debug, Clone, Serialize, Deserialize)]
127pub struct Memory {
128 pub id: String,
130 pub session_id: Option<String>,
132 pub chat_id: Option<String>,
134 pub memory_type: MemoryType,
136 pub title: String,
138 pub content: String,
140 pub tags: Vec<String>,
142 pub is_active: bool,
144 pub created_at: String,
146 pub updated_at: String,
148}
149
150#[derive(Debug, Clone, Default)]
152pub struct MemoryFilter {
153 pub memory_type: Option<MemoryType>,
155 pub tags: Option<Vec<String>>,
157 pub session_id: Option<String>,
159 pub chat_id: Option<String>,
161 pub since: Option<String>,
163 pub only_active: bool,
165}
166
167impl Memory {
168 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 pub fn with_session_id(mut self, session_id: String) -> Self {
192 self.session_id = Some(session_id);
193 self
194 }
195
196 pub fn with_chat_id(mut self, chat_id: String) -> Self {
198 self.chat_id = Some(chat_id);
199 self
200 }
201}
202
203pub 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 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
228fn 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 if sessions_table_exists(conn)? {
247 Ok(1) } else {
249 Ok(0) }
251 }
252 Err(e) => Err(e),
253 }
254}
255
256fn 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
266fn 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
345fn 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
369fn 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
385fn 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
412fn 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
442fn migrate_v4_to_v5(conn: &Connection) -> SqliteResult<()> {
450 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
460fn 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
481pub 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
507pub 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
531pub 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
549pub 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
562pub fn activate_session(
564 conn: &Connection,
565 session_id: &str,
566 chat_id: &str,
567) -> SqliteResult<()> {
568 let now = current_timestamp();
569 conn.execute(
571 "UPDATE sessions SET is_active = 0 WHERE chat_id = ?1",
572 params![chat_id],
573 )?;
574 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
582pub 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
595fn 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
609pub 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
619pub 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
629pub 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
638pub 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
656pub 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
677pub 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
696pub 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 let tool_calls = if let Some(tool_info) = &data.tool_info {
713 if let serde_json::Value::Object(obj) = tool_info {
714 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 if content.is_none() && tool_calls.is_none() {
752 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 _ => Err(crate::error::AgentError::InternalError(format!(
782 "Unknown role: {}",
783 data.role
784 ))),
785 }
786}
787
788#[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#[derive(Debug, Clone, Default)]
805pub struct MessageSearchFilter<'a> {
806 pub session_id: Option<&'a str>,
808 pub role: Option<&'a str>,
810 pub since: Option<&'a str>,
812 pub until: Option<&'a str>,
814}
815
816pub 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 params.push(ToSqlOutput::from(query));
837
838 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 if let Some(role) = filter.role {
846 conditions.push("m.role = ?".to_string());
847 params.push(ToSqlOutput::from(role));
848 }
849
850 if let Some(since) = filter.since {
852 conditions.push("m.created_at >= ?".to_string());
853 params.push(ToSqlOutput::from(since));
854 }
855
856 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
902pub 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 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
950pub 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
982pub 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
1010pub 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
1019pub 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
1025pub 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
1036pub 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
1050pub 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
1063pub fn recall_memories(
1065 conn: &Connection,
1066 query: &str,
1067 filter: &MemoryFilter,
1068 limit: usize,
1069) -> SqliteResult<Vec<Memory>> {
1070 let mut results = find_memories_by_title(conn, query, filter, Some(limit))?;
1072
1073 if results.len() >= limit {
1075 results.truncate(limit);
1076 return Ok(results);
1077 }
1078
1079 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
1096fn 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 if filter.only_active {
1135 conditions.push("is_active = 1".to_string());
1136 }
1137
1138 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 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 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 if let Some(since) = &filter.since {
1158 conditions.push("created_at >= ?".to_string());
1159 params.push(since.as_str().into());
1160 }
1161
1162 if let Some(title) = title_search {
1164 conditions.push("title LIKE ?".to_string());
1165 params.push(format!("%{}%", title).into());
1166 }
1167
1168 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 if filter.only_active {
1209 conditions.push("is_active = 1".to_string());
1210 }
1211
1212 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 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 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 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 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 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 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 assert!(find_session_by_chat_id(&conn, "does-not-exist")
1456 .unwrap()
1457 .is_none());
1458
1459 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 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 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 delete_session(&conn, "qq-1").unwrap();
1499 insert_session(&conn, "qq-2", Some("group:abc"), "Second", "model", "qq").unwrap();
1500
1501 let err = insert_session(&conn, "qq-3", Some("group:abc"), "Third", "model", "qq");
1503 assert!(err.is_err());
1504
1505 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 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 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 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 init_db(&conn).unwrap();
1578
1579 let v: i32 = read_schema_version(&conn).unwrap();
1581 assert_eq!(v, CURRENT_SCHEMA_VERSION);
1582
1583 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 init_db(&conn).unwrap();
1596 assert_eq!(read_schema_version(&conn).unwrap(), CURRENT_SCHEMA_VERSION);
1597 }
1598
1599 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 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 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 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 init_db(&conn).unwrap();
1787
1788 assert_eq!(read_schema_version(&conn).unwrap(), CURRENT_SCHEMA_VERSION);
1789
1790 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}