Skip to main content

wfe_sqlite/
lib.rs

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