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 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 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 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}