1use async_trait::async_trait;
4use chrono::Utc;
5use rusqlite::{OptionalExtension, TransactionBehavior, params};
6use starweaver_context::ResumableState;
7use starweaver_core::{CheckpointId, RunId, SessionId};
8use starweaver_runtime::{AgentCheckpoint, AgentStreamRecord};
9use starweaver_session::{
10 ApprovalRecord, CheckpointRef, CompactRunTrace, CompactSessionTrace, DeferredToolRecord,
11 EnvironmentStateRef, RunRecord, RunStatus, SessionFilter, SessionRecord, SessionStatus,
12 SessionStore, SessionStoreError, SessionStoreResult, StreamCursorRef,
13};
14
15use super::{
16 LocalStore,
17 db::{
18 insert_approval_records_tx, insert_deferred_tool_records_tx, insert_raw_stream_records_tx,
19 insert_stream_cursor_tx, load_session_tx, next_sequence_tx, upsert_run_tx,
20 upsert_session_tx,
21 },
22};
23use crate::{CliError, config::CliConfig};
24
25#[derive(Clone, Debug)]
27pub struct LocalSessionStore {
28 config: CliConfig,
29}
30
31impl LocalSessionStore {
32 #[must_use]
34 pub const fn new(config: CliConfig) -> Self {
35 Self { config }
36 }
37
38 fn open_store(&self) -> SessionStoreResult<LocalStore> {
39 LocalStore::open(&self.config).map_err(session_failed_cli)
40 }
41}
42
43#[async_trait]
44impl SessionStore for LocalSessionStore {
45 async fn save_session(&self, mut session: SessionRecord) -> SessionStoreResult<()> {
46 session.updated_at = Utc::now();
47 let mut store = self.open_store()?;
48 let tx = store
49 .conn
50 .transaction_with_behavior(TransactionBehavior::Immediate)
51 .map_err(session_failed)?;
52 upsert_session_tx(&tx, &session).map_err(session_failed)?;
53 tx.commit().map_err(session_failed)
54 }
55
56 async fn load_session(&self, session_id: &SessionId) -> SessionStoreResult<SessionRecord> {
57 self.open_store()?
58 .load_session(session_id.as_str())
59 .map_err(session_failed_cli)
60 }
61
62 async fn list_sessions(&self, filter: SessionFilter) -> SessionStoreResult<Vec<SessionRecord>> {
63 let store = self.open_store()?;
64 let mut stmt = store
65 .conn
66 .prepare("SELECT record_json FROM sessions ORDER BY updated_at DESC")
67 .map_err(session_failed)?;
68 let rows = stmt
69 .query_map([], |row| row.get::<_, String>(0))
70 .map_err(session_failed)?;
71 let mut sessions = Vec::new();
72 for row in rows {
73 let session: SessionRecord =
74 serde_json::from_str(&row.map_err(session_failed)?).map_err(session_failed)?;
75 if filter.status.is_some_and(|status| session.status != status) {
76 continue;
77 }
78 if filter
79 .profile
80 .as_ref()
81 .is_some_and(|profile| session.profile.as_ref() != Some(profile))
82 {
83 continue;
84 }
85 if filter
86 .workspace
87 .as_ref()
88 .is_some_and(|workspace| session.workspace.as_ref() != Some(workspace))
89 {
90 continue;
91 }
92 sessions.push(session);
93 if filter.limit.is_some_and(|limit| sessions.len() >= limit) {
94 break;
95 }
96 }
97 Ok(sessions)
98 }
99
100 async fn update_session_status(
101 &self,
102 session_id: &SessionId,
103 status: SessionStatus,
104 ) -> SessionStoreResult<()> {
105 let mut session = self.load_session(session_id).await?;
106 session.status = status;
107 self.save_session(session).await
108 }
109
110 async fn save_context_state(
111 &self,
112 session_id: &SessionId,
113 state: ResumableState,
114 ) -> SessionStoreResult<()> {
115 let mut session = self.load_session(session_id).await?;
116 session.state = state;
117 self.save_session(session).await
118 }
119
120 async fn save_environment_state(
121 &self,
122 session_id: &SessionId,
123 environment_state: EnvironmentStateRef,
124 ) -> SessionStoreResult<()> {
125 let mut session = self.load_session(session_id).await?;
126 session.environment_state = Some(environment_state);
127 self.save_session(session).await
128 }
129
130 async fn append_run(&self, mut run: RunRecord) -> SessionStoreResult<()> {
131 let mut store = self.open_store()?;
132 let tx = store
133 .conn
134 .transaction_with_behavior(TransactionBehavior::Immediate)
135 .map_err(session_failed)?;
136 let mut session = load_session_tx(&tx, run.session_id.as_str()).map_err(session_failed)?;
137 run.updated_at = Utc::now();
138 if let Some(existing_sequence) =
139 existing_run_sequence(&tx, run.session_id.as_str(), run.run_id.as_str())?
140 {
141 run.sequence_no = existing_sequence;
142 } else if run.sequence_no == 0
143 || sequence_exists(&tx, run.session_id.as_str(), run.sequence_no)?
144 {
145 run.sequence_no =
146 next_sequence_tx(&tx, run.session_id.as_str()).map_err(session_failed)?;
147 }
148 apply_run_to_session(&mut session, &run);
149 upsert_run_tx(&tx, &run).map_err(session_failed)?;
150 upsert_session_tx(&tx, &session).map_err(session_failed)?;
151 tx.commit().map_err(session_failed)
152 }
153
154 async fn load_run(
155 &self,
156 session_id: &SessionId,
157 run_id: &RunId,
158 ) -> SessionStoreResult<RunRecord> {
159 self.open_store()?
160 .load_run(session_id.as_str(), run_id.as_str())
161 .map_err(session_failed_cli)
162 }
163
164 async fn list_runs(&self, session_id: &SessionId) -> SessionStoreResult<Vec<RunRecord>> {
165 let store = self.open_store()?;
166 let mut stmt = store
167 .conn
168 .prepare("SELECT record_json FROM runs WHERE session_id = ?1 ORDER BY sequence_no ASC")
169 .map_err(session_failed)?;
170 let rows = stmt
171 .query_map(params![session_id.as_str()], |row| row.get::<_, String>(0))
172 .map_err(session_failed)?;
173 let runs = collect_json_records(rows)?;
174 Ok(runs)
175 }
176
177 async fn update_run_status(
178 &self,
179 session_id: &SessionId,
180 run_id: &RunId,
181 status: RunStatus,
182 output_preview: Option<String>,
183 ) -> SessionStoreResult<()> {
184 let mut run = self.load_run(session_id, run_id).await?;
185 run.status = status;
186 run.output_preview = output_preview;
187 run.updated_at = Utc::now();
188 self.append_run(run).await
189 }
190
191 async fn append_checkpoint(
192 &self,
193 session_id: &SessionId,
194 checkpoint: AgentCheckpoint,
195 ) -> SessionStoreResult<()> {
196 let mut store = self.open_store()?;
197 let tx = store
198 .conn
199 .transaction_with_behavior(TransactionBehavior::Immediate)
200 .map_err(session_failed)?;
201 let checkpoint_id = checkpoint.checkpoint_id.clone();
202 let checkpoint_run_id = checkpoint.run_id.clone();
203 let checkpoint_node = checkpoint.node;
204 let checkpoint_node_label = format!("{checkpoint_node:?}");
205 let checkpoint_sequence = checkpoint.run_step;
206 let stream_cursor = checkpoint.resume.cursor.stream_cursor;
207 let checkpoint_metadata = checkpoint.metadata.clone();
208 let mut run = load_run_tx(&tx, session_id.as_str(), checkpoint_run_id.as_str())?;
209 tx.execute(
210 "INSERT OR REPLACE INTO checkpoints
211 (checkpoint_id, session_id, run_id, sequence_no, node, checkpoint_json, created_at)
212 VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)",
213 params![
214 checkpoint_id.as_str(),
215 session_id.as_str(),
216 checkpoint_run_id.as_str(),
217 i64::try_from(checkpoint_sequence).map_err(session_failed)?,
218 checkpoint_node_label,
219 serde_json::to_string(&checkpoint).map_err(session_failed)?,
220 Utc::now().to_rfc3339(),
221 ],
222 )
223 .map_err(session_failed)?;
224 run.latest_checkpoint = Some(CheckpointRef {
225 checkpoint_id,
226 run_id: checkpoint_run_id,
227 sequence: checkpoint_sequence,
228 node: format!("{checkpoint_node:?}"),
229 storage_ref: None,
230 stream_cursor,
231 created_at: Utc::now(),
232 metadata: checkpoint_metadata,
233 });
234 run.updated_at = Utc::now();
235 upsert_run_tx(&tx, &run).map_err(session_failed)?;
236 tx.commit().map_err(session_failed)
237 }
238
239 async fn load_checkpoints(
240 &self,
241 session_id: &SessionId,
242 run_id: &RunId,
243 ) -> SessionStoreResult<Vec<AgentCheckpoint>> {
244 let store = self.open_store()?;
245 let mut stmt = store
246 .conn
247 .prepare(
248 "SELECT checkpoint_json FROM checkpoints
249 WHERE session_id = ?1 AND run_id = ?2
250 ORDER BY sequence_no ASC, checkpoint_id ASC",
251 )
252 .map_err(session_failed)?;
253 let rows = stmt
254 .query_map(params![session_id.as_str(), run_id.as_str()], |row| {
255 row.get::<_, String>(0)
256 })
257 .map_err(session_failed)?;
258 let mut checkpoints = Vec::new();
259 for row in rows {
260 let json = row.map_err(session_failed)?;
261 if let Ok(checkpoint) = serde_json::from_str::<AgentCheckpoint>(&json) {
262 checkpoints.push(checkpoint);
263 }
264 }
265 Ok(checkpoints)
266 }
267
268 async fn append_stream_records(
269 &self,
270 session_id: &SessionId,
271 run_id: &RunId,
272 records: Vec<AgentStreamRecord>,
273 ) -> SessionStoreResult<()> {
274 let mut store = self.open_store()?;
275 let tx = store
276 .conn
277 .transaction_with_behavior(TransactionBehavior::Immediate)
278 .map_err(session_failed)?;
279 let mut run = load_run_tx(&tx, session_id.as_str(), run_id.as_str())?;
280 insert_raw_stream_records_tx(&tx, &run, &records).map_err(session_failed)?;
281 if let Some(sequence) = latest_raw_sequence(&tx, session_id.as_str(), run_id.as_str())? {
282 let cursor =
283 StreamCursorRef::new("raw_runtime", format!("run:{}", run_id.as_str()), sequence);
284 run.stream_cursors
285 .retain(|existing| existing.family != cursor.family);
286 run.stream_cursors.push(cursor.clone());
287 run.updated_at = Utc::now();
288 upsert_run_tx(&tx, &run).map_err(session_failed)?;
289 let mut session = load_session_tx(&tx, session_id.as_str()).map_err(session_failed)?;
290 upsert_session_cursor(&mut session, cursor);
291 upsert_session_tx(&tx, &session).map_err(session_failed)?;
292 }
293 tx.commit().map_err(session_failed)
294 }
295
296 async fn replay_stream_records(
297 &self,
298 session_id: &SessionId,
299 run_id: &RunId,
300 ) -> SessionStoreResult<Vec<AgentStreamRecord>> {
301 self.replay_stream_records_after(session_id, run_id, None)
302 .await
303 }
304
305 async fn replay_stream_records_after(
306 &self,
307 session_id: &SessionId,
308 run_id: &RunId,
309 after_sequence: Option<usize>,
310 ) -> SessionStoreResult<Vec<AgentStreamRecord>> {
311 let after = after_sequence.map_or(-1_i64, |value| i64::try_from(value).unwrap_or(i64::MAX));
312 let store = self.open_store()?;
313 let mut stmt = store
314 .conn
315 .prepare(
316 "SELECT record_json FROM raw_stream_records
317 WHERE session_id = ?1 AND run_id = ?2 AND sequence_no > ?3
318 ORDER BY sequence_no ASC",
319 )
320 .map_err(session_failed)?;
321 let rows = stmt
322 .query_map(
323 params![session_id.as_str(), run_id.as_str(), after],
324 |row| row.get::<_, String>(0),
325 )
326 .map_err(session_failed)?;
327 let records = collect_json_records(rows)?;
328 Ok(records)
329 }
330
331 async fn save_stream_cursor(
332 &self,
333 session_id: &SessionId,
334 run_id: &RunId,
335 cursor: StreamCursorRef,
336 ) -> SessionStoreResult<()> {
337 let mut store = self.open_store()?;
338 let tx = store
339 .conn
340 .transaction_with_behavior(TransactionBehavior::Immediate)
341 .map_err(session_failed)?;
342 let mut run = load_run_tx(&tx, session_id.as_str(), run_id.as_str())?;
343 run.stream_cursors
344 .retain(|existing| existing.family != cursor.family || existing.scope != cursor.scope);
345 run.stream_cursors.push(cursor.clone());
346 run.updated_at = Utc::now();
347 upsert_run_tx(&tx, &run).map_err(session_failed)?;
348 let mut session = load_session_tx(&tx, session_id.as_str()).map_err(session_failed)?;
349 upsert_session_cursor(&mut session, cursor.clone());
350 upsert_session_tx(&tx, &session).map_err(session_failed)?;
351 insert_stream_cursor_tx(&tx, &run, &cursor).map_err(session_failed)?;
352 tx.commit().map_err(session_failed)
353 }
354
355 async fn append_approval(&self, approval: ApprovalRecord) -> SessionStoreResult<()> {
356 let mut store = self.open_store()?;
357 let tx = store
358 .conn
359 .transaction_with_behavior(TransactionBehavior::Immediate)
360 .map_err(session_failed)?;
361 insert_approval_records_tx(&tx, &[approval]).map_err(session_failed)?;
362 tx.commit().map_err(session_failed)
363 }
364
365 async fn load_approvals(
366 &self,
367 session_id: &SessionId,
368 run_id: &RunId,
369 ) -> SessionStoreResult<Vec<ApprovalRecord>> {
370 self.open_store()?
371 .list_approvals(Some(session_id.as_str()), Some(run_id.as_str()))
372 .map_err(session_failed_cli)
373 }
374
375 async fn append_deferred_tool(&self, record: DeferredToolRecord) -> SessionStoreResult<()> {
376 let mut store = self.open_store()?;
377 let tx = store
378 .conn
379 .transaction_with_behavior(TransactionBehavior::Immediate)
380 .map_err(session_failed)?;
381 insert_deferred_tool_records_tx(&tx, &[record]).map_err(session_failed)?;
382 tx.commit().map_err(session_failed)
383 }
384
385 async fn load_deferred_tools(
386 &self,
387 session_id: &SessionId,
388 run_id: &RunId,
389 ) -> SessionStoreResult<Vec<DeferredToolRecord>> {
390 self.open_store()?
391 .list_deferred_tools(Some(session_id.as_str()), Some(run_id.as_str()))
392 .map_err(session_failed_cli)
393 }
394
395 async fn compact_run_trace(
396 &self,
397 session_id: &SessionId,
398 run_id: &RunId,
399 ) -> SessionStoreResult<CompactRunTrace> {
400 let store = self.open_store()?;
401 let run = store
402 .load_run(session_id.as_str(), run_id.as_str())
403 .map_err(session_failed_cli)?;
404 Ok(CompactRunTrace {
405 session_id: Some(session_id.clone()),
406 run_id: Some(run_id.clone()),
407 status: run.status,
408 checkpoints: checkpoint_ids(&store, session_id.as_str(), run_id.as_str())?,
409 approvals: pending_approval_count(&store, session_id.as_str(), run_id.as_str())?,
410 deferred_tools: pending_deferred_count(&store, session_id.as_str(), run_id.as_str())?,
411 latest_checkpoint: run
412 .latest_checkpoint
413 .as_ref()
414 .map(|checkpoint| checkpoint.checkpoint_id.clone()),
415 stream_cursor: latest_raw_sequence_ref(&store, session_id.as_str(), run_id.as_str())?,
416 stream_cursors: run.stream_cursors,
417 output_preview: run.output_preview,
418 trace_context: run.trace_context,
419 updated_at: Some(run.updated_at),
420 metadata: run.metadata,
421 })
422 }
423
424 async fn compact_session_trace(
425 &self,
426 session_id: &SessionId,
427 ) -> SessionStoreResult<CompactSessionTrace> {
428 let session = self.load_session(session_id).await?;
429 let runs = self.list_runs(session_id).await?;
430 let latest_run = runs.last();
431 Ok(CompactSessionTrace {
432 session_id: session.session_id,
433 title: session.title,
434 workspace: session.workspace,
435 profile: session.profile,
436 status: session.status,
437 runs: runs.len(),
438 latest_run_id: latest_run.map(|run| run.run_id.clone()),
439 last_output_preview: latest_run.and_then(|run| run.output_preview.clone()),
440 stream_cursors: session.stream_cursors,
441 trace_context: session.trace_context,
442 created_at: session.created_at,
443 updated_at: session.updated_at,
444 metadata: session.metadata,
445 })
446 }
447}
448
449fn load_run_tx(
450 tx: &rusqlite::Transaction<'_>,
451 session_id: &str,
452 run_id: &str,
453) -> SessionStoreResult<RunRecord> {
454 tx.query_row(
455 "SELECT record_json FROM runs WHERE session_id = ?1 AND run_id = ?2",
456 params![session_id, run_id],
457 |row| row.get::<_, String>(0),
458 )
459 .optional()
460 .map_err(session_failed)?
461 .map(|json| serde_json::from_str(&json).map_err(session_failed))
462 .transpose()?
463 .ok_or_else(|| SessionStoreError::NotFound(format!("{session_id}:{run_id}")))
464}
465
466fn existing_run_sequence(
467 tx: &rusqlite::Transaction<'_>,
468 session_id: &str,
469 run_id: &str,
470) -> SessionStoreResult<Option<usize>> {
471 tx.query_row(
472 "SELECT sequence_no FROM runs WHERE session_id = ?1 AND run_id = ?2",
473 params![session_id, run_id],
474 |row| row.get::<_, i64>(0),
475 )
476 .optional()
477 .map_err(session_failed)?
478 .map(|value| usize::try_from(value).map_err(session_failed))
479 .transpose()
480}
481
482fn sequence_exists(
483 tx: &rusqlite::Transaction<'_>,
484 session_id: &str,
485 sequence_no: usize,
486) -> SessionStoreResult<bool> {
487 let count = tx
488 .query_row(
489 "SELECT COUNT(*) FROM runs WHERE session_id = ?1 AND sequence_no = ?2",
490 params![
491 session_id,
492 i64::try_from(sequence_no).map_err(session_failed)?
493 ],
494 |row| row.get::<_, i64>(0),
495 )
496 .map_err(session_failed)?;
497 Ok(count > 0)
498}
499
500fn apply_run_to_session(session: &mut SessionRecord, run: &RunRecord) {
501 session.profile.clone_from(&run.profile);
502 session.head_run_id = Some(run.run_id.clone());
503 if run.status == RunStatus::Completed {
504 session.head_success_run_id = Some(run.run_id.clone());
505 }
506 if matches!(
507 run.status,
508 RunStatus::Queued | RunStatus::Running | RunStatus::Waiting
509 ) {
510 session.active_run_id = Some(run.run_id.clone());
511 } else if session.active_run_id.as_ref() == Some(&run.run_id) {
512 session.active_run_id = None;
513 }
514 session.updated_at = run.updated_at;
515}
516
517fn upsert_session_cursor(session: &mut SessionRecord, cursor: StreamCursorRef) {
518 session
519 .stream_cursors
520 .retain(|existing| existing.family != cursor.family || existing.scope != cursor.scope);
521 session.stream_cursors.push(cursor);
522 session.updated_at = Utc::now();
523}
524
525fn latest_raw_sequence(
526 tx: &rusqlite::Transaction<'_>,
527 session_id: &str,
528 run_id: &str,
529) -> SessionStoreResult<Option<usize>> {
530 tx.query_row(
531 "SELECT MAX(sequence_no) FROM raw_stream_records WHERE session_id = ?1 AND run_id = ?2",
532 params![session_id, run_id],
533 |row| row.get::<_, Option<i64>>(0),
534 )
535 .map_err(session_failed)?
536 .map(|value| usize::try_from(value).map_err(session_failed))
537 .transpose()
538}
539
540fn latest_raw_sequence_ref(
541 store: &LocalStore,
542 session_id: &str,
543 run_id: &str,
544) -> SessionStoreResult<Option<usize>> {
545 store
546 .conn
547 .query_row(
548 "SELECT MAX(sequence_no) FROM raw_stream_records WHERE session_id = ?1 AND run_id = ?2",
549 params![session_id, run_id],
550 |row| row.get::<_, Option<i64>>(0),
551 )
552 .map_err(session_failed)?
553 .map(|value| usize::try_from(value).map_err(session_failed))
554 .transpose()
555}
556
557fn checkpoint_ids(
558 store: &LocalStore,
559 session_id: &str,
560 run_id: &str,
561) -> SessionStoreResult<Vec<CheckpointId>> {
562 let mut stmt = store
563 .conn
564 .prepare(
565 "SELECT checkpoint_id FROM checkpoints
566 WHERE session_id = ?1 AND run_id = ?2
567 ORDER BY sequence_no ASC, checkpoint_id ASC",
568 )
569 .map_err(session_failed)?;
570 let rows = stmt
571 .query_map(params![session_id, run_id], |row| row.get::<_, String>(0))
572 .map_err(session_failed)?;
573 rows.collect::<Result<Vec<_>, _>>()
574 .map_err(session_failed)
575 .map(|ids| ids.into_iter().map(CheckpointId::from_string).collect())
576}
577
578fn pending_approval_count(
579 store: &LocalStore,
580 session_id: &str,
581 run_id: &str,
582) -> SessionStoreResult<usize> {
583 count_rows(
584 store,
585 "SELECT COUNT(*) FROM approvals WHERE session_id = ?1 AND run_id = ?2 AND status = 'pending'",
586 session_id,
587 run_id,
588 )
589}
590
591fn pending_deferred_count(
592 store: &LocalStore,
593 session_id: &str,
594 run_id: &str,
595) -> SessionStoreResult<usize> {
596 count_rows(
597 store,
598 "SELECT COUNT(*) FROM deferred_tools
599 WHERE session_id = ?1 AND run_id = ?2
600 AND status IN ('pending', 'running', 'waiting')",
601 session_id,
602 run_id,
603 )
604}
605
606fn count_rows(
607 store: &LocalStore,
608 sql: &str,
609 session_id: &str,
610 run_id: &str,
611) -> SessionStoreResult<usize> {
612 let count = store
613 .conn
614 .query_row(sql, params![session_id, run_id], |row| row.get::<_, i64>(0))
615 .map_err(session_failed)?;
616 usize::try_from(count).map_err(session_failed)
617}
618
619fn collect_json_records<T>(
620 rows: rusqlite::MappedRows<'_, impl FnMut(&rusqlite::Row<'_>) -> rusqlite::Result<String>>,
621) -> SessionStoreResult<Vec<T>>
622where
623 T: serde::de::DeserializeOwned,
624{
625 rows.collect::<Result<Vec<_>, _>>()
626 .map_err(session_failed)?
627 .into_iter()
628 .map(|json| serde_json::from_str(&json).map_err(session_failed))
629 .collect()
630}
631
632fn session_failed(error: impl std::fmt::Display) -> SessionStoreError {
633 SessionStoreError::Failed(error.to_string())
634}
635
636fn session_failed_cli(error: CliError) -> SessionStoreError {
637 match error {
638 CliError::NotFound(id) => SessionStoreError::NotFound(id),
639 error => SessionStoreError::Failed(error.to_string()),
640 }
641}