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