Skip to main content

rushai_core/store/
mod.rs

1mod db;
2
3use std::path::Path;
4use std::time::{SystemTime, UNIX_EPOCH};
5
6use rushai_protocol::{MessageId, Part, Role, SessionId};
7use rusqlite::{OptionalExtension, Row, params};
8use thiserror::Error;
9
10use db::Db;
11
12#[derive(Debug, Error)]
13pub enum StoreError {
14    #[error(transparent)]
15    Sqlite(#[from] rusqlite::Error),
16    #[error("migration failed: {0}")]
17    Migration(#[from] rusqlite_migration::Error),
18    #[error("corrupt message parts: {0}")]
19    Parts(#[from] serde_json::Error),
20    #[error("invalid role {0:?} stored in database")]
21    InvalidRole(String),
22    #[error("store thread is gone")]
23    Closed,
24    #[error("failed to start store thread: {0}")]
25    Thread(#[from] std::io::Error),
26}
27
28#[derive(Debug, Clone, PartialEq)]
29pub struct Session {
30    pub id: SessionId,
31    pub parent: Option<SessionId>,
32    pub title: String,
33    pub summary_message_id: Option<MessageId>,
34    pub cost: f64,
35    pub prompt_tokens: u64,
36    pub completion_tokens: u64,
37    pub created_at: i64,
38    pub updated_at: i64,
39}
40
41#[derive(Debug, Clone, PartialEq)]
42pub struct StoredMessage {
43    pub id: MessageId,
44    pub session: SessionId,
45    pub role: Role,
46    pub provider: Option<String>,
47    pub model: Option<String>,
48    pub parts: Vec<Part>,
49    pub is_summary: bool,
50    pub created_at: i64,
51}
52
53pub struct Store {
54    db: Db,
55}
56
57impl Store {
58    pub fn open(path: impl AsRef<Path>) -> Result<Self, StoreError> {
59        Ok(Self {
60            db: Db::open(Some(path.as_ref().to_path_buf()))?,
61        })
62    }
63
64    pub fn open_in_memory() -> Result<Self, StoreError> {
65        Ok(Self {
66            db: Db::open(None)?,
67        })
68    }
69
70    pub async fn create_session(
71        &self,
72        title: String,
73        parent: Option<SessionId>,
74    ) -> Result<Session, StoreError> {
75        let now = now_ms();
76        let session = Session {
77            id: SessionId::new(),
78            parent,
79            title,
80            summary_message_id: None,
81            cost: 0.0,
82            prompt_tokens: 0,
83            completion_tokens: 0,
84            created_at: now,
85            updated_at: now,
86        };
87        let row = session.clone();
88        self.db
89            .call(move |conn| {
90                conn.execute(
91                    "INSERT INTO sessions (id, parent_session_id, title, cost, prompt_tokens, \
92                     completion_tokens, created_at, updated_at) \
93                     VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8)",
94                    params![
95                        row.id.as_str(),
96                        row.parent.as_ref().map(SessionId::as_str),
97                        row.title,
98                        row.cost,
99                        row.prompt_tokens as i64,
100                        row.completion_tokens as i64,
101                        row.created_at,
102                        row.updated_at,
103                    ],
104                )?;
105                Ok(())
106            })
107            .await?;
108        Ok(session)
109    }
110
111    pub async fn session(&self, id: &SessionId) -> Result<Option<Session>, StoreError> {
112        let id = id.clone();
113        self.db
114            .call(move |conn| {
115                conn.query_row(
116                    &format!("{SESSION_SELECT} WHERE id = ?1"),
117                    params![id.as_str()],
118                    session_from_row,
119                )
120                .optional()
121                .map_err(Into::into)
122            })
123            .await
124    }
125
126    /// All sessions, most recently updated first.
127    pub async fn sessions(&self) -> Result<Vec<Session>, StoreError> {
128        self.db
129            .call(move |conn| {
130                let mut stmt = conn.prepare(&format!(
131                    "{SESSION_SELECT} ORDER BY updated_at DESC, id DESC"
132                ))?;
133                let rows = stmt.query_map([], session_from_row)?;
134                rows.collect::<Result<Vec<_>, _>>().map_err(Into::into)
135            })
136            .await
137    }
138
139    pub async fn delete_session(&self, id: &SessionId) -> Result<(), StoreError> {
140        let id = id.clone();
141        self.db
142            .call(move |conn| {
143                conn.execute("DELETE FROM sessions WHERE id = ?1", params![id.as_str()])?;
144                Ok(())
145            })
146            .await
147    }
148
149    /// Insert or update a message and touch the session's updated_at.
150    pub async fn save_message(&self, message: &StoredMessage) -> Result<(), StoreError> {
151        let parts = serde_json::to_string(&message.parts)?;
152        let row = message.clone();
153        self.db
154            .call(move |conn| {
155                let tx = conn.transaction()?;
156                tx.execute(
157                    "INSERT INTO messages (id, session_id, role, provider, model, parts, \
158                     is_summary, created_at) VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8) \
159                     ON CONFLICT (id) DO UPDATE SET parts = excluded.parts, \
160                     is_summary = excluded.is_summary",
161                    params![
162                        row.id.as_str(),
163                        row.session.as_str(),
164                        role_str(row.role),
165                        row.provider,
166                        row.model,
167                        parts,
168                        row.is_summary,
169                        row.created_at,
170                    ],
171                )?;
172                tx.execute(
173                    "UPDATE sessions SET updated_at = ?1 WHERE id = ?2",
174                    params![now_ms(), row.session.as_str()],
175                )?;
176                tx.commit()?;
177                Ok(())
178            })
179            .await
180    }
181
182    /// Messages for a session in creation order.
183    pub async fn messages(&self, session: &SessionId) -> Result<Vec<StoredMessage>, StoreError> {
184        let session = session.clone();
185        self.db
186            .call(move |conn| {
187                let mut stmt = conn.prepare(
188                    "SELECT id, session_id, role, provider, model, parts, is_summary, created_at \
189                     FROM messages WHERE session_id = ?1 ORDER BY created_at, id",
190                )?;
191                let rows = stmt.query_map(params![session.as_str()], message_from_row)?;
192                rows.collect::<Result<Vec<_>, _>>()?
193                    .into_iter()
194                    .map(RawMessage::parse)
195                    .collect()
196            })
197            .await
198    }
199}
200
201const SESSION_SELECT: &str = "SELECT id, parent_session_id, title, summary_message_id, cost, \
202                              prompt_tokens, completion_tokens, created_at, updated_at \
203                              FROM sessions";
204
205fn session_from_row(row: &Row<'_>) -> rusqlite::Result<Session> {
206    Ok(Session {
207        id: SessionId::from(row.get::<_, String>(0)?),
208        parent: row.get::<_, Option<String>>(1)?.map(SessionId::from),
209        title: row.get(2)?,
210        summary_message_id: row.get::<_, Option<String>>(3)?.map(MessageId::from),
211        cost: row.get(4)?,
212        prompt_tokens: row.get::<_, i64>(5)? as u64,
213        completion_tokens: row.get::<_, i64>(6)? as u64,
214        created_at: row.get(7)?,
215        updated_at: row.get(8)?,
216    })
217}
218
219struct RawMessage {
220    id: String,
221    session: String,
222    role: String,
223    provider: Option<String>,
224    model: Option<String>,
225    parts: String,
226    is_summary: bool,
227    created_at: i64,
228}
229
230fn message_from_row(row: &Row<'_>) -> rusqlite::Result<RawMessage> {
231    Ok(RawMessage {
232        id: row.get(0)?,
233        session: row.get(1)?,
234        role: row.get(2)?,
235        provider: row.get(3)?,
236        model: row.get(4)?,
237        parts: row.get(5)?,
238        is_summary: row.get(6)?,
239        created_at: row.get(7)?,
240    })
241}
242
243impl RawMessage {
244    fn parse(self) -> Result<StoredMessage, StoreError> {
245        let role = match self.role.as_str() {
246            "user" => Role::User,
247            "assistant" => Role::Assistant,
248            other => return Err(StoreError::InvalidRole(other.to_owned())),
249        };
250        Ok(StoredMessage {
251            id: MessageId::from(self.id),
252            session: SessionId::from(self.session),
253            role,
254            provider: self.provider,
255            model: self.model,
256            parts: serde_json::from_str(&self.parts)?,
257            is_summary: self.is_summary,
258            created_at: self.created_at,
259        })
260    }
261}
262
263fn role_str(role: Role) -> &'static str {
264    match role {
265        Role::User => "user",
266        Role::Assistant => "assistant",
267    }
268}
269
270fn now_ms() -> i64 {
271    SystemTime::now()
272        .duration_since(UNIX_EPOCH)
273        .expect("system clock before unix epoch")
274        .as_millis() as i64
275}