Skip to main content

wfe_sqlite/
lib.rs

1//! wfe-sqlite — SQLite persistence provider for the WFE workflow engine.
2use std::collections::HashMap;
3
4use async_trait::async_trait;
5use chrono::{DateTime, Utc};
6use sqlx::sqlite::{SqliteConnectOptions, SqlitePoolOptions};
7use sqlx::{Row, SqlitePool};
8
9use wfe_core::models::{
10    CommandName, Event, EventSubscription, ExecutionError, ExecutionPointer, ScheduledCommand,
11    WorkflowInstance, WorkflowStatus,
12};
13use wfe_core::traits::{
14    EventRepository, PersistenceProvider, ScheduledCommandRepository, SubscriptionRepository,
15    WorkflowRepository,
16};
17use wfe_core::{Result, WfeError};
18
19/// SQLite-backed persistence provider for the WFE workflow engine.
20pub struct SqlitePersistenceProvider {
21    pool: SqlitePool,
22}
23
24impl SqlitePersistenceProvider {
25    /// Create a new provider connected to the given SQLite database URL.
26    ///
27    /// For in-memory databases, pass `":memory:"`.
28    /// The schema tables are created automatically.
29    pub async fn new(database_url: &str) -> std::result::Result<Self, Box<dyn std::error::Error>> {
30        let options: SqliteConnectOptions = database_url
31            .parse::<SqliteConnectOptions>()?
32            .create_if_missing(true)
33            .journal_mode(sqlx::sqlite::SqliteJournalMode::Wal);
34
35        let max_connections = if database_url.contains(":memory:") {
36            1
37        } else {
38            4
39        };
40
41        let pool = SqlitePoolOptions::new()
42            .max_connections(max_connections)
43            .connect_with(options)
44            .await?;
45
46        // Enable WAL mode and foreign keys
47        sqlx::query("PRAGMA foreign_keys = ON")
48            .execute(&pool)
49            .await?;
50
51        let provider = Self { pool };
52        provider.ensure_store_exists().await?;
53        Ok(provider)
54    }
55
56    /// Run the DDL statements to create tables and indexes.
57    async fn create_tables(&self) -> std::result::Result<(), sqlx::Error> {
58        sqlx::query(
59            "CREATE TABLE IF NOT EXISTS workflows (
60                id TEXT PRIMARY KEY,
61                name TEXT NOT NULL UNIQUE,
62                root_workflow_id TEXT,
63                definition_id TEXT NOT NULL,
64                version INTEGER NOT NULL,
65                description TEXT,
66                reference TEXT,
67                status TEXT NOT NULL,
68                data TEXT NOT NULL,
69                next_execution INTEGER,
70                create_time TEXT NOT NULL,
71                complete_time TEXT
72            )",
73        )
74        .execute(&self.pool)
75        .await?;
76
77        // Per-definition monotonic counter used to generate human-friendly
78        // instance names of the form `{definition_id}-{N}`.
79        sqlx::query(
80            "CREATE TABLE IF NOT EXISTS definition_sequences (
81                definition_id TEXT PRIMARY KEY,
82                next_num INTEGER NOT NULL
83            )",
84        )
85        .execute(&self.pool)
86        .await?;
87
88        sqlx::query(
89            "CREATE TABLE IF NOT EXISTS execution_pointers (
90                id TEXT PRIMARY KEY,
91                workflow_id TEXT NOT NULL,
92                step_id INTEGER NOT NULL,
93                active INTEGER NOT NULL DEFAULT 1,
94                status TEXT NOT NULL,
95                sleep_until TEXT,
96                persistence_data TEXT,
97                start_time TEXT,
98                end_time TEXT,
99                event_name TEXT,
100                event_key TEXT,
101                event_published INTEGER NOT NULL DEFAULT 0,
102                event_data TEXT,
103                step_name TEXT,
104                retry_count INTEGER NOT NULL DEFAULT 0,
105                children TEXT NOT NULL DEFAULT '[]',
106                context_item TEXT,
107                predecessor_id TEXT,
108                outcome TEXT,
109                scope TEXT NOT NULL DEFAULT '[]',
110                extension_attributes TEXT NOT NULL DEFAULT '{}',
111                FOREIGN KEY (workflow_id) REFERENCES workflows(id) ON DELETE CASCADE
112            )",
113        )
114        .execute(&self.pool)
115        .await?;
116
117        sqlx::query(
118            "CREATE TABLE IF NOT EXISTS events (
119                id TEXT PRIMARY KEY,
120                event_name TEXT NOT NULL,
121                event_key TEXT NOT NULL,
122                event_data TEXT NOT NULL,
123                event_time TEXT NOT NULL,
124                is_processed INTEGER NOT NULL DEFAULT 0
125            )",
126        )
127        .execute(&self.pool)
128        .await?;
129
130        sqlx::query(
131            "CREATE TABLE IF NOT EXISTS event_subscriptions (
132                id TEXT PRIMARY KEY,
133                workflow_id TEXT NOT NULL,
134                step_id INTEGER NOT NULL,
135                execution_pointer_id TEXT NOT NULL,
136                event_name TEXT NOT NULL,
137                event_key TEXT NOT NULL,
138                subscribe_as_of TEXT NOT NULL,
139                subscription_data TEXT,
140                external_token TEXT,
141                external_worker_id TEXT,
142                external_token_expiry TEXT,
143                terminated INTEGER NOT NULL DEFAULT 0
144            )",
145        )
146        .execute(&self.pool)
147        .await?;
148
149        sqlx::query(
150            "CREATE TABLE IF NOT EXISTS execution_errors (
151                id INTEGER PRIMARY KEY AUTOINCREMENT,
152                error_time TEXT NOT NULL,
153                workflow_id TEXT NOT NULL,
154                execution_pointer_id TEXT NOT NULL,
155                message TEXT NOT NULL
156            )",
157        )
158        .execute(&self.pool)
159        .await?;
160
161        sqlx::query(
162            "CREATE TABLE IF NOT EXISTS scheduled_commands (
163                id INTEGER PRIMARY KEY AUTOINCREMENT,
164                command_name TEXT NOT NULL,
165                data TEXT NOT NULL,
166                execute_time INTEGER NOT NULL,
167                UNIQUE(command_name, data)
168            )",
169        )
170        .execute(&self.pool)
171        .await?;
172
173        // Indexes
174        sqlx::query(
175            "CREATE INDEX IF NOT EXISTS idx_workflows_next_execution ON workflows(next_execution)",
176        )
177        .execute(&self.pool)
178        .await?;
179        sqlx::query("CREATE INDEX IF NOT EXISTS idx_workflows_status ON workflows(status)")
180            .execute(&self.pool)
181            .await?;
182        sqlx::query("CREATE INDEX IF NOT EXISTS idx_execution_pointers_workflow_id ON execution_pointers(workflow_id)")
183            .execute(&self.pool)
184            .await?;
185        sqlx::query(
186            "CREATE INDEX IF NOT EXISTS idx_events_name_key ON events(event_name, event_key)",
187        )
188        .execute(&self.pool)
189        .await?;
190        sqlx::query("CREATE INDEX IF NOT EXISTS idx_events_is_processed ON events(is_processed)")
191            .execute(&self.pool)
192            .await?;
193        sqlx::query("CREATE INDEX IF NOT EXISTS idx_events_event_time ON events(event_time)")
194            .execute(&self.pool)
195            .await?;
196        sqlx::query("CREATE INDEX IF NOT EXISTS idx_event_subscriptions_name_key ON event_subscriptions(event_name, event_key)")
197            .execute(&self.pool)
198            .await?;
199        sqlx::query("CREATE INDEX IF NOT EXISTS idx_event_subscriptions_workflow_id ON event_subscriptions(workflow_id)")
200            .execute(&self.pool)
201            .await?;
202        sqlx::query("CREATE INDEX IF NOT EXISTS idx_scheduled_commands_execute_time ON scheduled_commands(execute_time)")
203            .execute(&self.pool)
204            .await?;
205
206        Ok(())
207    }
208}
209
210// ─── Helpers ───────────────────────────────────────────────────────────────
211
212fn to_persistence_err(e: sqlx::Error) -> WfeError {
213    WfeError::Persistence(e.to_string())
214}
215
216fn dt_to_string(dt: &DateTime<Utc>) -> String {
217    dt.to_rfc3339()
218}
219
220fn opt_dt_to_string(dt: &Option<DateTime<Utc>>) -> Option<String> {
221    dt.as_ref().map(dt_to_string)
222}
223
224fn string_to_dt(s: &str) -> std::result::Result<DateTime<Utc>, WfeError> {
225    s.parse::<DateTime<Utc>>()
226        .map_err(|e| WfeError::Persistence(format!("Failed to parse datetime '{s}': {e}")))
227}
228
229fn string_to_opt_dt(s: &Option<String>) -> std::result::Result<Option<DateTime<Utc>>, WfeError> {
230    match s {
231        Some(s) => Ok(Some(string_to_dt(s)?)),
232        None => Ok(None),
233    }
234}
235
236fn row_to_workflow(
237    row: &sqlx::sqlite::SqliteRow,
238    pointers: Vec<ExecutionPointer>,
239) -> std::result::Result<WorkflowInstance, WfeError> {
240    let status_str: String = row.try_get("status").map_err(to_persistence_err)?;
241    let status: WorkflowStatus = serde_json::from_str(&format!("\"{status_str}\""))
242        .map_err(|e| WfeError::Persistence(format!("Failed to deserialize WorkflowStatus: {e}")))?;
243
244    let data_str: String = row.try_get("data").map_err(to_persistence_err)?;
245    let data: serde_json::Value = serde_json::from_str(&data_str)
246        .map_err(|e| WfeError::Persistence(format!("Failed to deserialize data: {e}")))?;
247
248    let create_time_str: String = row.try_get("create_time").map_err(to_persistence_err)?;
249    let complete_time_str: Option<String> =
250        row.try_get("complete_time").map_err(to_persistence_err)?;
251
252    Ok(WorkflowInstance {
253        id: row.try_get("id").map_err(to_persistence_err)?,
254        name: row.try_get("name").map_err(to_persistence_err)?,
255        root_workflow_id: row
256            .try_get("root_workflow_id")
257            .map_err(to_persistence_err)?,
258        workflow_definition_id: row.try_get("definition_id").map_err(to_persistence_err)?,
259        version: row
260            .try_get::<i64, _>("version")
261            .map_err(to_persistence_err)? as u32,
262        description: row.try_get("description").map_err(to_persistence_err)?,
263        reference: row.try_get("reference").map_err(to_persistence_err)?,
264        execution_pointers: pointers,
265        next_execution: row.try_get("next_execution").map_err(to_persistence_err)?,
266        status,
267        data,
268        create_time: string_to_dt(&create_time_str)?,
269        complete_time: string_to_opt_dt(&complete_time_str)?,
270    })
271}
272
273fn row_to_pointer(
274    row: &sqlx::sqlite::SqliteRow,
275) -> std::result::Result<ExecutionPointer, WfeError> {
276    let status_str: String = row.try_get("status").map_err(to_persistence_err)?;
277    let status: wfe_core::models::PointerStatus =
278        serde_json::from_str(&format!("\"{status_str}\"")).map_err(|e| {
279            WfeError::Persistence(format!("Failed to deserialize PointerStatus: {e}"))
280        })?;
281
282    let persistence_data_str: Option<String> = row
283        .try_get("persistence_data")
284        .map_err(to_persistence_err)?;
285    let persistence_data: Option<serde_json::Value> = persistence_data_str
286        .as_deref()
287        .map(serde_json::from_str)
288        .transpose()
289        .map_err(|e| {
290            WfeError::Persistence(format!("Failed to deserialize persistence_data: {e}"))
291        })?;
292
293    let event_data_str: Option<String> = row.try_get("event_data").map_err(to_persistence_err)?;
294    let event_data: Option<serde_json::Value> = event_data_str
295        .as_deref()
296        .map(serde_json::from_str)
297        .transpose()
298        .map_err(|e| WfeError::Persistence(format!("Failed to deserialize event_data: {e}")))?;
299
300    let context_item_str: Option<String> =
301        row.try_get("context_item").map_err(to_persistence_err)?;
302    let context_item: Option<serde_json::Value> = context_item_str
303        .as_deref()
304        .map(serde_json::from_str)
305        .transpose()
306        .map_err(|e| WfeError::Persistence(format!("Failed to deserialize context_item: {e}")))?;
307
308    let outcome_str: Option<String> = row.try_get("outcome").map_err(to_persistence_err)?;
309    let outcome: Option<serde_json::Value> = outcome_str
310        .as_deref()
311        .map(serde_json::from_str)
312        .transpose()
313        .map_err(|e| WfeError::Persistence(format!("Failed to deserialize outcome: {e}")))?;
314
315    let children_str: String = row.try_get("children").map_err(to_persistence_err)?;
316    let children: Vec<String> = serde_json::from_str(&children_str)
317        .map_err(|e| WfeError::Persistence(format!("Failed to deserialize children: {e}")))?;
318
319    let scope_str: String = row.try_get("scope").map_err(to_persistence_err)?;
320    let scope: Vec<String> = serde_json::from_str(&scope_str)
321        .map_err(|e| WfeError::Persistence(format!("Failed to deserialize scope: {e}")))?;
322
323    let ext_str: String = row
324        .try_get("extension_attributes")
325        .map_err(to_persistence_err)?;
326    let extension_attributes: HashMap<String, serde_json::Value> = serde_json::from_str(&ext_str)
327        .map_err(|e| {
328        WfeError::Persistence(format!("Failed to deserialize extension_attributes: {e}"))
329    })?;
330
331    let sleep_until_str: Option<String> = row.try_get("sleep_until").map_err(to_persistence_err)?;
332    let start_time_str: Option<String> = row.try_get("start_time").map_err(to_persistence_err)?;
333    let end_time_str: Option<String> = row.try_get("end_time").map_err(to_persistence_err)?;
334
335    Ok(ExecutionPointer {
336        id: row.try_get("id").map_err(to_persistence_err)?,
337        step_id: row
338            .try_get::<i64, _>("step_id")
339            .map_err(to_persistence_err)? as usize,
340        active: row
341            .try_get::<bool, _>("active")
342            .map_err(to_persistence_err)?,
343        status,
344        sleep_until: string_to_opt_dt(&sleep_until_str)?,
345        persistence_data,
346        start_time: string_to_opt_dt(&start_time_str)?,
347        end_time: string_to_opt_dt(&end_time_str)?,
348        event_name: row.try_get("event_name").map_err(to_persistence_err)?,
349        event_key: row.try_get("event_key").map_err(to_persistence_err)?,
350        event_published: row
351            .try_get::<bool, _>("event_published")
352            .map_err(to_persistence_err)?,
353        event_data,
354        step_name: row.try_get("step_name").map_err(to_persistence_err)?,
355        retry_count: row
356            .try_get::<i64, _>("retry_count")
357            .map_err(to_persistence_err)? as u32,
358        children,
359        context_item,
360        predecessor_id: row.try_get("predecessor_id").map_err(to_persistence_err)?,
361        outcome,
362        scope,
363        extension_attributes,
364    })
365}
366
367fn row_to_event(row: &sqlx::sqlite::SqliteRow) -> std::result::Result<Event, WfeError> {
368    let event_data_str: String = row.try_get("event_data").map_err(to_persistence_err)?;
369    let event_data: serde_json::Value = serde_json::from_str(&event_data_str)
370        .map_err(|e| WfeError::Persistence(format!("Failed to deserialize event_data: {e}")))?;
371
372    let event_time_str: String = row.try_get("event_time").map_err(to_persistence_err)?;
373
374    Ok(Event {
375        id: row.try_get("id").map_err(to_persistence_err)?,
376        event_name: row.try_get("event_name").map_err(to_persistence_err)?,
377        event_key: row.try_get("event_key").map_err(to_persistence_err)?,
378        event_data,
379        event_time: string_to_dt(&event_time_str)?,
380        is_processed: row
381            .try_get::<bool, _>("is_processed")
382            .map_err(to_persistence_err)?,
383    })
384}
385
386fn row_to_subscription(
387    row: &sqlx::sqlite::SqliteRow,
388) -> std::result::Result<EventSubscription, WfeError> {
389    let subscribe_as_of_str: String = row.try_get("subscribe_as_of").map_err(to_persistence_err)?;
390
391    let subscription_data_str: Option<String> = row
392        .try_get("subscription_data")
393        .map_err(to_persistence_err)?;
394    let subscription_data: Option<serde_json::Value> = subscription_data_str
395        .as_deref()
396        .map(serde_json::from_str)
397        .transpose()
398        .map_err(|e| {
399            WfeError::Persistence(format!("Failed to deserialize subscription_data: {e}"))
400        })?;
401
402    let external_token_expiry_str: Option<String> = row
403        .try_get("external_token_expiry")
404        .map_err(to_persistence_err)?;
405
406    Ok(EventSubscription {
407        id: row.try_get("id").map_err(to_persistence_err)?,
408        workflow_id: row.try_get("workflow_id").map_err(to_persistence_err)?,
409        step_id: row
410            .try_get::<i64, _>("step_id")
411            .map_err(to_persistence_err)? as usize,
412        execution_pointer_id: row
413            .try_get("execution_pointer_id")
414            .map_err(to_persistence_err)?,
415        event_name: row.try_get("event_name").map_err(to_persistence_err)?,
416        event_key: row.try_get("event_key").map_err(to_persistence_err)?,
417        subscribe_as_of: string_to_dt(&subscribe_as_of_str)?,
418        subscription_data,
419        external_token: row.try_get("external_token").map_err(to_persistence_err)?,
420        external_worker_id: row
421            .try_get("external_worker_id")
422            .map_err(to_persistence_err)?,
423        external_token_expiry: string_to_opt_dt(&external_token_expiry_str)?,
424    })
425}
426
427// ─── Trait implementations ─────────────────────────────────────────────────
428
429#[async_trait]
430impl WorkflowRepository for SqlitePersistenceProvider {
431    async fn create_new_workflow(&self, instance: &WorkflowInstance) -> Result<String> {
432        let id = if instance.id.is_empty() {
433            uuid::Uuid::new_v4().to_string()
434        } else {
435            instance.id.clone()
436        };
437        // Fall back to the UUID when the caller didn't assign a human name.
438        // Production callers go through `WorkflowHost::start_workflow` which
439        // always fills this in, but test fixtures and external callers
440        // shouldn't trip the UNIQUE constraint.
441        let name = if instance.name.is_empty() {
442            id.clone()
443        } else {
444            instance.name.clone()
445        };
446
447        let status_str = serde_json::to_value(instance.status)
448            .map_err(|e| WfeError::Persistence(e.to_string()))?
449            .as_str()
450            .unwrap_or("Runnable")
451            .to_string();
452        let data_str = serde_json::to_string(&instance.data)
453            .map_err(|e| WfeError::Persistence(e.to_string()))?;
454        let create_time_str = dt_to_string(&instance.create_time);
455        let complete_time_str = opt_dt_to_string(&instance.complete_time);
456
457        let mut tx = self.pool.begin().await.map_err(to_persistence_err)?;
458
459        sqlx::query(
460            "INSERT INTO workflows (id, name, root_workflow_id, definition_id, version, description, reference, status, data, next_execution, create_time, complete_time)
461             VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12)",
462        )
463        .bind(&id)
464        .bind(&name)
465        .bind(&instance.root_workflow_id)
466        .bind(&instance.workflow_definition_id)
467        .bind(instance.version as i64)
468        .bind(&instance.description)
469        .bind(&instance.reference)
470        .bind(&status_str)
471        .bind(&data_str)
472        .bind(instance.next_execution)
473        .bind(&create_time_str)
474        .bind(&complete_time_str)
475        .execute(&mut *tx)
476        .await
477        .map_err(to_persistence_err)?;
478
479        for ptr in &instance.execution_pointers {
480            insert_pointer(&mut tx, &id, ptr).await?;
481        }
482
483        tx.commit().await.map_err(to_persistence_err)?;
484        Ok(id)
485    }
486
487    async fn persist_workflow(&self, instance: &WorkflowInstance) -> Result<()> {
488        let status_str = serde_json::to_value(instance.status)
489            .map_err(|e| WfeError::Persistence(e.to_string()))?
490            .as_str()
491            .unwrap_or("Runnable")
492            .to_string();
493        let data_str = serde_json::to_string(&instance.data)
494            .map_err(|e| WfeError::Persistence(e.to_string()))?;
495        let complete_time_str = opt_dt_to_string(&instance.complete_time);
496
497        let mut tx = self.pool.begin().await.map_err(to_persistence_err)?;
498
499        sqlx::query(
500            "UPDATE workflows SET name = ?1, root_workflow_id = ?2, definition_id = ?3,
501             version = ?4, description = ?5, reference = ?6, status = ?7, data = ?8,
502             next_execution = ?9, complete_time = ?10
503             WHERE id = ?11",
504        )
505        .bind(&instance.name)
506        .bind(&instance.root_workflow_id)
507        .bind(&instance.workflow_definition_id)
508        .bind(instance.version as i64)
509        .bind(&instance.description)
510        .bind(&instance.reference)
511        .bind(&status_str)
512        .bind(&data_str)
513        .bind(instance.next_execution)
514        .bind(&complete_time_str)
515        .bind(&instance.id)
516        .execute(&mut *tx)
517        .await
518        .map_err(to_persistence_err)?;
519
520        // Replace all pointers
521        sqlx::query("DELETE FROM execution_pointers WHERE workflow_id = ?1")
522            .bind(&instance.id)
523            .execute(&mut *tx)
524            .await
525            .map_err(to_persistence_err)?;
526
527        for ptr in &instance.execution_pointers {
528            insert_pointer(&mut tx, &instance.id, ptr).await?;
529        }
530
531        tx.commit().await.map_err(to_persistence_err)?;
532        Ok(())
533    }
534
535    async fn persist_workflow_with_subscriptions(
536        &self,
537        instance: &WorkflowInstance,
538        subscriptions: &[EventSubscription],
539    ) -> Result<()> {
540        let status_str = serde_json::to_value(instance.status)
541            .map_err(|e| WfeError::Persistence(e.to_string()))?
542            .as_str()
543            .unwrap_or("Runnable")
544            .to_string();
545        let data_str = serde_json::to_string(&instance.data)
546            .map_err(|e| WfeError::Persistence(e.to_string()))?;
547        let complete_time_str = opt_dt_to_string(&instance.complete_time);
548
549        let mut tx = self.pool.begin().await.map_err(to_persistence_err)?;
550
551        sqlx::query(
552            "UPDATE workflows SET name = ?1, root_workflow_id = ?2, definition_id = ?3,
553             version = ?4, description = ?5, reference = ?6, status = ?7, data = ?8,
554             next_execution = ?9, complete_time = ?10
555             WHERE id = ?11",
556        )
557        .bind(&instance.name)
558        .bind(&instance.root_workflow_id)
559        .bind(&instance.workflow_definition_id)
560        .bind(instance.version as i64)
561        .bind(&instance.description)
562        .bind(&instance.reference)
563        .bind(&status_str)
564        .bind(&data_str)
565        .bind(instance.next_execution)
566        .bind(&complete_time_str)
567        .bind(&instance.id)
568        .execute(&mut *tx)
569        .await
570        .map_err(to_persistence_err)?;
571
572        sqlx::query("DELETE FROM execution_pointers WHERE workflow_id = ?1")
573            .bind(&instance.id)
574            .execute(&mut *tx)
575            .await
576            .map_err(to_persistence_err)?;
577
578        for ptr in &instance.execution_pointers {
579            insert_pointer(&mut tx, &instance.id, ptr).await?;
580        }
581
582        for sub in subscriptions {
583            insert_subscription(&mut tx, sub).await?;
584        }
585
586        tx.commit().await.map_err(to_persistence_err)?;
587        Ok(())
588    }
589
590    async fn get_runnable_instances(&self, as_at: DateTime<Utc>) -> Result<Vec<String>> {
591        let as_at_millis = as_at.timestamp_millis();
592        let rows = sqlx::query(
593            "SELECT id FROM workflows WHERE status = 'Runnable' AND next_execution IS NOT NULL AND next_execution <= ?1",
594        )
595        .bind(as_at_millis)
596        .fetch_all(&self.pool)
597        .await
598        .map_err(to_persistence_err)?;
599
600        let ids = rows
601            .iter()
602            .map(|r| r.try_get("id").map_err(to_persistence_err))
603            .collect::<Result<Vec<String>>>()?;
604        Ok(ids)
605    }
606
607    async fn get_workflow_instance(&self, id: &str) -> Result<WorkflowInstance> {
608        let row = sqlx::query("SELECT * FROM workflows WHERE id = ?1")
609            .bind(id)
610            .fetch_optional(&self.pool)
611            .await
612            .map_err(to_persistence_err)?
613            .ok_or_else(|| WfeError::WorkflowNotFound(id.to_string()))?;
614
615        let pointer_rows = sqlx::query("SELECT * FROM execution_pointers WHERE workflow_id = ?1")
616            .bind(id)
617            .fetch_all(&self.pool)
618            .await
619            .map_err(to_persistence_err)?;
620
621        let pointers = pointer_rows
622            .iter()
623            .map(row_to_pointer)
624            .collect::<Result<Vec<ExecutionPointer>>>()?;
625
626        row_to_workflow(&row, pointers)
627    }
628
629    async fn get_workflow_instance_by_name(&self, name: &str) -> Result<WorkflowInstance> {
630        let row = sqlx::query("SELECT id FROM workflows WHERE name = ?1")
631            .bind(name)
632            .fetch_optional(&self.pool)
633            .await
634            .map_err(to_persistence_err)?
635            .ok_or_else(|| WfeError::WorkflowNotFound(name.to_string()))?;
636        let id: String = row.try_get("id").map_err(to_persistence_err)?;
637        self.get_workflow_instance(&id).await
638    }
639
640    async fn next_definition_sequence(&self, definition_id: &str) -> Result<u64> {
641        // SQLite doesn't support `INSERT ... ON CONFLICT ... RETURNING` prior
642        // to 3.35, but sqlx bundles a new-enough build. Emulate an atomic
643        // increment via UPSERT + RETURNING so concurrent callers don't collide.
644        let row = sqlx::query(
645            "INSERT INTO definition_sequences (definition_id, next_num)
646             VALUES (?1, 1)
647             ON CONFLICT(definition_id) DO UPDATE
648               SET next_num = next_num + 1
649             RETURNING next_num",
650        )
651        .bind(definition_id)
652        .fetch_one(&self.pool)
653        .await
654        .map_err(to_persistence_err)?;
655        let next: i64 = row.try_get("next_num").map_err(to_persistence_err)?;
656        Ok(next as u64)
657    }
658
659    async fn get_workflow_instances(&self, ids: &[String]) -> Result<Vec<WorkflowInstance>> {
660        if ids.is_empty() {
661            return Ok(Vec::new());
662        }
663
664        let mut result = Vec::with_capacity(ids.len());
665        for id in ids {
666            match self.get_workflow_instance(id).await {
667                Ok(w) => result.push(w),
668                Err(WfeError::WorkflowNotFound(_)) => continue,
669                Err(e) => return Err(e),
670            }
671        }
672        Ok(result)
673    }
674}
675
676async fn insert_pointer(
677    tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
678    workflow_id: &str,
679    ptr: &ExecutionPointer,
680) -> Result<()> {
681    let status_str = serde_json::to_value(ptr.status)
682        .map_err(|e| WfeError::Persistence(e.to_string()))?
683        .as_str()
684        .unwrap_or("Pending")
685        .to_string();
686    let persistence_data_str = ptr
687        .persistence_data
688        .as_ref()
689        .map(serde_json::to_string)
690        .transpose()
691        .map_err(|e| WfeError::Persistence(e.to_string()))?;
692    let event_data_str = ptr
693        .event_data
694        .as_ref()
695        .map(serde_json::to_string)
696        .transpose()
697        .map_err(|e| WfeError::Persistence(e.to_string()))?;
698    let context_item_str = ptr
699        .context_item
700        .as_ref()
701        .map(serde_json::to_string)
702        .transpose()
703        .map_err(|e| WfeError::Persistence(e.to_string()))?;
704    let outcome_str = ptr
705        .outcome
706        .as_ref()
707        .map(serde_json::to_string)
708        .transpose()
709        .map_err(|e| WfeError::Persistence(e.to_string()))?;
710    let children_str =
711        serde_json::to_string(&ptr.children).map_err(|e| WfeError::Persistence(e.to_string()))?;
712    let scope_str =
713        serde_json::to_string(&ptr.scope).map_err(|e| WfeError::Persistence(e.to_string()))?;
714    let ext_str = serde_json::to_string(&ptr.extension_attributes)
715        .map_err(|e| WfeError::Persistence(e.to_string()))?;
716
717    let sleep_until_str = opt_dt_to_string(&ptr.sleep_until);
718    let start_time_str = opt_dt_to_string(&ptr.start_time);
719    let end_time_str = opt_dt_to_string(&ptr.end_time);
720
721    sqlx::query(
722        "INSERT INTO execution_pointers
723         (id, workflow_id, step_id, active, status, sleep_until, persistence_data, start_time,
724          end_time, event_name, event_key, event_published, event_data, step_name, retry_count,
725          children, context_item, predecessor_id, outcome, scope, extension_attributes)
726         VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18, ?19, ?20, ?21)",
727    )
728    .bind(&ptr.id)
729    .bind(workflow_id)
730    .bind(ptr.step_id as i64)
731    .bind(ptr.active)
732    .bind(&status_str)
733    .bind(&sleep_until_str)
734    .bind(&persistence_data_str)
735    .bind(&start_time_str)
736    .bind(&end_time_str)
737    .bind(&ptr.event_name)
738    .bind(&ptr.event_key)
739    .bind(ptr.event_published)
740    .bind(&event_data_str)
741    .bind(&ptr.step_name)
742    .bind(ptr.retry_count as i64)
743    .bind(&children_str)
744    .bind(&context_item_str)
745    .bind(&ptr.predecessor_id)
746    .bind(&outcome_str)
747    .bind(&scope_str)
748    .bind(&ext_str)
749    .execute(&mut **tx)
750    .await
751    .map_err(to_persistence_err)?;
752
753    Ok(())
754}
755
756async fn insert_subscription(
757    tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
758    sub: &EventSubscription,
759) -> Result<()> {
760    let subscribe_as_of_str = dt_to_string(&sub.subscribe_as_of);
761    let subscription_data_str = sub
762        .subscription_data
763        .as_ref()
764        .map(serde_json::to_string)
765        .transpose()
766        .map_err(|e| WfeError::Persistence(e.to_string()))?;
767    let external_token_expiry_str = opt_dt_to_string(&sub.external_token_expiry);
768
769    sqlx::query(
770        "INSERT INTO event_subscriptions
771         (id, workflow_id, step_id, execution_pointer_id, event_name, event_key,
772          subscribe_as_of, subscription_data, external_token, external_worker_id,
773          external_token_expiry, terminated)
774         VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, 0)",
775    )
776    .bind(&sub.id)
777    .bind(&sub.workflow_id)
778    .bind(sub.step_id as i64)
779    .bind(&sub.execution_pointer_id)
780    .bind(&sub.event_name)
781    .bind(&sub.event_key)
782    .bind(&subscribe_as_of_str)
783    .bind(&subscription_data_str)
784    .bind(&sub.external_token)
785    .bind(&sub.external_worker_id)
786    .bind(&external_token_expiry_str)
787    .execute(&mut **tx)
788    .await
789    .map_err(to_persistence_err)?;
790
791    Ok(())
792}
793
794#[async_trait]
795impl SubscriptionRepository for SqlitePersistenceProvider {
796    async fn create_event_subscription(&self, subscription: &EventSubscription) -> Result<String> {
797        let id = if subscription.id.is_empty() {
798            uuid::Uuid::new_v4().to_string()
799        } else {
800            subscription.id.clone()
801        };
802
803        let mut stored = subscription.clone();
804        stored.id = id.clone();
805
806        let mut tx = self.pool.begin().await.map_err(to_persistence_err)?;
807        insert_subscription(&mut tx, &stored).await?;
808        tx.commit().await.map_err(to_persistence_err)?;
809        Ok(id)
810    }
811
812    async fn get_subscriptions(
813        &self,
814        event_name: &str,
815        event_key: &str,
816        as_of: DateTime<Utc>,
817    ) -> Result<Vec<EventSubscription>> {
818        let as_of_str = dt_to_string(&as_of);
819        let rows = sqlx::query(
820            "SELECT * FROM event_subscriptions
821             WHERE event_name = ?1 AND event_key = ?2 AND subscribe_as_of <= ?3 AND terminated = 0",
822        )
823        .bind(event_name)
824        .bind(event_key)
825        .bind(&as_of_str)
826        .fetch_all(&self.pool)
827        .await
828        .map_err(to_persistence_err)?;
829
830        rows.iter().map(row_to_subscription).collect()
831    }
832
833    async fn terminate_subscription(&self, subscription_id: &str) -> Result<()> {
834        let result = sqlx::query("UPDATE event_subscriptions SET terminated = 1 WHERE id = ?1")
835            .bind(subscription_id)
836            .execute(&self.pool)
837            .await
838            .map_err(to_persistence_err)?;
839
840        if result.rows_affected() == 0 {
841            return Err(WfeError::SubscriptionNotFound(subscription_id.to_string()));
842        }
843        Ok(())
844    }
845
846    async fn get_subscription(&self, subscription_id: &str) -> Result<EventSubscription> {
847        let row = sqlx::query("SELECT * FROM event_subscriptions WHERE id = ?1")
848            .bind(subscription_id)
849            .fetch_optional(&self.pool)
850            .await
851            .map_err(to_persistence_err)?
852            .ok_or_else(|| WfeError::SubscriptionNotFound(subscription_id.to_string()))?;
853
854        row_to_subscription(&row)
855    }
856
857    async fn get_first_open_subscription(
858        &self,
859        event_name: &str,
860        event_key: &str,
861        as_of: DateTime<Utc>,
862    ) -> Result<Option<EventSubscription>> {
863        let as_of_str = dt_to_string(&as_of);
864        let row = sqlx::query(
865            "SELECT * FROM event_subscriptions
866             WHERE event_name = ?1 AND event_key = ?2 AND subscribe_as_of <= ?3
867               AND terminated = 0 AND external_token IS NULL
868             LIMIT 1",
869        )
870        .bind(event_name)
871        .bind(event_key)
872        .bind(&as_of_str)
873        .fetch_optional(&self.pool)
874        .await
875        .map_err(to_persistence_err)?;
876
877        match row {
878            Some(r) => Ok(Some(row_to_subscription(&r)?)),
879            None => Ok(None),
880        }
881    }
882
883    async fn set_subscription_token(
884        &self,
885        subscription_id: &str,
886        token: &str,
887        worker_id: &str,
888        expiry: DateTime<Utc>,
889    ) -> Result<bool> {
890        let expiry_str = dt_to_string(&expiry);
891
892        // Only set token if external_token is currently NULL
893        let result = sqlx::query(
894            "UPDATE event_subscriptions
895             SET external_token = ?1, external_worker_id = ?2, external_token_expiry = ?3
896             WHERE id = ?4 AND external_token IS NULL",
897        )
898        .bind(token)
899        .bind(worker_id)
900        .bind(&expiry_str)
901        .bind(subscription_id)
902        .execute(&self.pool)
903        .await
904        .map_err(to_persistence_err)?;
905
906        if result.rows_affected() == 0 {
907            // Check if the subscription exists at all
908            let exists = sqlx::query("SELECT 1 FROM event_subscriptions WHERE id = ?1")
909                .bind(subscription_id)
910                .fetch_optional(&self.pool)
911                .await
912                .map_err(to_persistence_err)?;
913            if exists.is_none() {
914                return Err(WfeError::SubscriptionNotFound(subscription_id.to_string()));
915            }
916            return Ok(false);
917        }
918        Ok(true)
919    }
920
921    async fn clear_subscription_token(&self, subscription_id: &str, token: &str) -> Result<()> {
922        let result = sqlx::query(
923            "UPDATE event_subscriptions
924             SET external_token = NULL, external_worker_id = NULL, external_token_expiry = NULL
925             WHERE id = ?1 AND external_token = ?2",
926        )
927        .bind(subscription_id)
928        .bind(token)
929        .execute(&self.pool)
930        .await
931        .map_err(to_persistence_err)?;
932
933        if result.rows_affected() == 0 {
934            return Err(WfeError::SubscriptionNotFound(subscription_id.to_string()));
935        }
936        Ok(())
937    }
938}
939
940#[async_trait]
941impl EventRepository for SqlitePersistenceProvider {
942    async fn create_event(&self, event: &Event) -> Result<String> {
943        let id = if event.id.is_empty() {
944            uuid::Uuid::new_v4().to_string()
945        } else {
946            event.id.clone()
947        };
948
949        let event_data_str = serde_json::to_string(&event.event_data)
950            .map_err(|e| WfeError::Persistence(e.to_string()))?;
951        let event_time_str = dt_to_string(&event.event_time);
952
953        sqlx::query(
954            "INSERT INTO events (id, event_name, event_key, event_data, event_time, is_processed)
955             VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
956        )
957        .bind(&id)
958        .bind(&event.event_name)
959        .bind(&event.event_key)
960        .bind(&event_data_str)
961        .bind(&event_time_str)
962        .bind(event.is_processed)
963        .execute(&self.pool)
964        .await
965        .map_err(to_persistence_err)?;
966
967        Ok(id)
968    }
969
970    async fn get_event(&self, id: &str) -> Result<Event> {
971        let row = sqlx::query("SELECT * FROM events WHERE id = ?1")
972            .bind(id)
973            .fetch_optional(&self.pool)
974            .await
975            .map_err(to_persistence_err)?
976            .ok_or_else(|| WfeError::EventNotFound(id.to_string()))?;
977
978        row_to_event(&row)
979    }
980
981    async fn get_runnable_events(&self, as_at: DateTime<Utc>) -> Result<Vec<String>> {
982        let as_at_str = dt_to_string(&as_at);
983        let rows = sqlx::query("SELECT id FROM events WHERE is_processed = 0 AND event_time <= ?1")
984            .bind(&as_at_str)
985            .fetch_all(&self.pool)
986            .await
987            .map_err(to_persistence_err)?;
988
989        rows.iter()
990            .map(|r| r.try_get("id").map_err(to_persistence_err))
991            .collect()
992    }
993
994    async fn get_events(
995        &self,
996        event_name: &str,
997        event_key: &str,
998        as_of: DateTime<Utc>,
999    ) -> Result<Vec<String>> {
1000        let as_of_str = dt_to_string(&as_of);
1001        let rows = sqlx::query(
1002            "SELECT id FROM events WHERE event_name = ?1 AND event_key = ?2 AND event_time <= ?3",
1003        )
1004        .bind(event_name)
1005        .bind(event_key)
1006        .bind(&as_of_str)
1007        .fetch_all(&self.pool)
1008        .await
1009        .map_err(to_persistence_err)?;
1010
1011        rows.iter()
1012            .map(|r| r.try_get("id").map_err(to_persistence_err))
1013            .collect()
1014    }
1015
1016    async fn mark_event_processed(&self, id: &str) -> Result<()> {
1017        let result = sqlx::query("UPDATE events SET is_processed = 1 WHERE id = ?1")
1018            .bind(id)
1019            .execute(&self.pool)
1020            .await
1021            .map_err(to_persistence_err)?;
1022
1023        if result.rows_affected() == 0 {
1024            return Err(WfeError::EventNotFound(id.to_string()));
1025        }
1026        Ok(())
1027    }
1028
1029    async fn mark_event_unprocessed(&self, id: &str) -> Result<()> {
1030        let result = sqlx::query("UPDATE events SET is_processed = 0 WHERE id = ?1")
1031            .bind(id)
1032            .execute(&self.pool)
1033            .await
1034            .map_err(to_persistence_err)?;
1035
1036        if result.rows_affected() == 0 {
1037            return Err(WfeError::EventNotFound(id.to_string()));
1038        }
1039        Ok(())
1040    }
1041}
1042
1043#[async_trait]
1044impl ScheduledCommandRepository for SqlitePersistenceProvider {
1045    fn supports_scheduled_commands(&self) -> bool {
1046        true
1047    }
1048
1049    async fn schedule_command(&self, command: &ScheduledCommand) -> Result<()> {
1050        let command_name_str = serde_json::to_value(&command.command_name)
1051            .map_err(|e| WfeError::Persistence(e.to_string()))?
1052            .as_str()
1053            .unwrap_or("")
1054            .to_string();
1055
1056        sqlx::query(
1057            "INSERT OR IGNORE INTO scheduled_commands (command_name, data, execute_time)
1058             VALUES (?1, ?2, ?3)",
1059        )
1060        .bind(&command_name_str)
1061        .bind(&command.data)
1062        .bind(command.execute_time)
1063        .execute(&self.pool)
1064        .await
1065        .map_err(to_persistence_err)?;
1066
1067        Ok(())
1068    }
1069
1070    async fn process_commands(
1071        &self,
1072        as_of: DateTime<Utc>,
1073        handler: &(
1074             dyn Fn(
1075            ScheduledCommand,
1076        )
1077            -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<()>> + Send>>
1078                 + Send
1079                 + Sync
1080         ),
1081    ) -> Result<()> {
1082        let as_of_millis = as_of.timestamp_millis();
1083
1084        // Fetch due commands
1085        let rows = sqlx::query(
1086            "SELECT id, command_name, data, execute_time FROM scheduled_commands WHERE execute_time <= ?1",
1087        )
1088        .bind(as_of_millis)
1089        .fetch_all(&self.pool)
1090        .await
1091        .map_err(to_persistence_err)?;
1092
1093        let mut commands: Vec<(i64, ScheduledCommand)> = Vec::new();
1094        for row in &rows {
1095            let db_id: i64 = row.try_get("id").map_err(to_persistence_err)?;
1096            let command_name_str: String =
1097                row.try_get("command_name").map_err(to_persistence_err)?;
1098            let command_name: CommandName =
1099                serde_json::from_str(&format!("\"{command_name_str}\"")).map_err(|e| {
1100                    WfeError::Persistence(format!("Failed to deserialize CommandName: {e}"))
1101                })?;
1102            let data: String = row.try_get("data").map_err(to_persistence_err)?;
1103            let execute_time: i64 = row.try_get("execute_time").map_err(to_persistence_err)?;
1104
1105            commands.push((
1106                db_id,
1107                ScheduledCommand {
1108                    command_name,
1109                    data,
1110                    execute_time,
1111                },
1112            ));
1113        }
1114
1115        // Process each command then delete it
1116        for (db_id, cmd) in commands {
1117            handler(cmd).await?;
1118            sqlx::query("DELETE FROM scheduled_commands WHERE id = ?1")
1119                .bind(db_id)
1120                .execute(&self.pool)
1121                .await
1122                .map_err(to_persistence_err)?;
1123        }
1124        Ok(())
1125    }
1126}
1127
1128#[async_trait]
1129impl PersistenceProvider for SqlitePersistenceProvider {
1130    async fn persist_errors(&self, errors: &[ExecutionError]) -> Result<()> {
1131        for error in errors {
1132            let error_time_str = dt_to_string(&error.error_time);
1133            sqlx::query(
1134                "INSERT INTO execution_errors (error_time, workflow_id, execution_pointer_id, message)
1135                 VALUES (?1, ?2, ?3, ?4)",
1136            )
1137            .bind(&error_time_str)
1138            .bind(&error.workflow_id)
1139            .bind(&error.execution_pointer_id)
1140            .bind(&error.message)
1141            .execute(&self.pool)
1142            .await
1143            .map_err(to_persistence_err)?;
1144        }
1145        Ok(())
1146    }
1147
1148    async fn ensure_store_exists(&self) -> Result<()> {
1149        self.create_tables()
1150            .await
1151            .map_err(|e| WfeError::Persistence(e.to_string()))
1152    }
1153}
1154
1155#[cfg(test)]
1156mod tests {
1157    use super::*;
1158
1159    #[tokio::test]
1160    async fn schema_creation_idempotent() {
1161        let provider = SqlitePersistenceProvider::new(":memory:").await.unwrap();
1162        // Call ensure_store_exists again — should not fail
1163        provider.ensure_store_exists().await.unwrap();
1164        provider.ensure_store_exists().await.unwrap();
1165    }
1166}