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 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 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 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
209fn 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#[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 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 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 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 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 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 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 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 provider.ensure_store_exists().await.unwrap();
1163 provider.ensure_store_exists().await.unwrap();
1164 }
1165}