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
18pub struct SqlitePersistenceProvider {
20 pool: SqlitePool,
21}
22
23impl SqlitePersistenceProvider {
24 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 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 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 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
198fn 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#[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 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 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 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 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 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 provider.ensure_store_exists().await.unwrap();
1118 provider.ensure_store_exists().await.unwrap();
1119 }
1120}