Skip to main content

starweaver_cli/local_store/
session_store.rs

1//! `SessionStore` adapter over the CLI local store.
2
3use 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/// Shared session store adapter backed by the CLI local `SQLite` store.
26#[derive(Clone, Debug)]
27pub struct LocalSessionStore {
28    config: CliConfig,
29}
30
31impl LocalSessionStore {
32    /// Create a local session store adapter from resolved CLI config.
33    #[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}