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