use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use parking_lot::RwLock;
use serde::{Deserialize, Serialize};
use super::super::super::channel::ChannelName;
use super::super::super::redex::{Redex, RedexError, RedexFileConfig};
use super::super::adapter::CortexAdapter;
use super::super::config::CortexAdapterConfig;
use super::super::error::CortexAdapterError;
use super::super::meta::{compute_checksum_with_meta, EventMeta, EVENT_META_SIZE};
use super::super::watermark::WatermarkingFold;
use super::dispatch::{
DISPATCH_TASK_ADVANCED, DISPATCH_TASK_CANCEL_REQUESTED, DISPATCH_TASK_DELETED,
DISPATCH_TASK_LINKED, DISPATCH_TASK_RETRIED, DISPATCH_TASK_SUBMITTED,
DISPATCH_TASK_TRANSITIONED, WORKFLOW_CHANNEL,
};
use super::fold::WorkflowFold;
use super::state::{StatusCounts, WorkflowState};
use super::types::{
AdvancedPayload, CancelRequestedPayload, DeletedPayload, LinkedPayload, RetriedPayload,
SubmittedPayload, TaskId, TaskState, TaskStatus, TransitionedPayload,
};
#[derive(Serialize, Deserialize)]
struct WorkflowSnapshotPayload {
app_seq: u64,
inner: Vec<u8>,
}
pub struct WorkflowAdapter {
inner: CortexAdapter<WorkflowState>,
origin_hash: u64,
app_seq: Arc<AtomicU64>,
}
impl WorkflowAdapter {
pub async fn open(redex: &Redex, origin_hash: u64) -> Result<Self, CortexAdapterError> {
Self::open_with_config(redex, origin_hash, RedexFileConfig::default()).await
}
pub async fn open_with_config(
redex: &Redex,
origin_hash: u64,
redex_config: RedexFileConfig,
) -> Result<Self, CortexAdapterError> {
let name = ChannelName::new(WORKFLOW_CHANNEL)
.map_err(|e| CortexAdapterError::Redex(RedexError::Channel(e.to_string())))?;
let app_seq = Arc::new(AtomicU64::new(0));
let fold = WatermarkingFold::new(WorkflowFold, app_seq.clone(), origin_hash);
let inner = CortexAdapter::open(
redex,
&name,
redex_config.clone(),
CortexAdapterConfig::default(),
fold,
WorkflowState::new(),
)?;
let file = redex.open_file(&name, redex_config)?;
let next_seq = file.next_seq();
if next_seq > 0 {
inner.wait_for_seq(next_seq - 1).await.map_err(|folded| {
CortexAdapterError::FoldStoppedBeforeSeq {
wanted: next_seq - 1,
folded_through: folded,
}
})?;
}
Ok(Self {
inner,
origin_hash,
app_seq,
})
}
pub fn submit(&self, id: TaskId) -> Result<u64, CortexAdapterError> {
self.ingest_typed(DISPATCH_TASK_SUBMITTED, &SubmittedPayload { id })
}
pub fn transition(&self, id: TaskId, status: TaskStatus) -> Result<u64, CortexAdapterError> {
self.ingest_typed(
DISPATCH_TASK_TRANSITIONED,
&TransitionedPayload { id, status },
)
}
pub fn start(&self, id: TaskId) -> Result<u64, CortexAdapterError> {
self.transition(id, TaskStatus::Running)
}
pub fn wait(&self, id: TaskId) -> Result<u64, CortexAdapterError> {
self.transition(id, TaskStatus::Waiting)
}
pub fn block(&self, id: TaskId) -> Result<u64, CortexAdapterError> {
self.transition(id, TaskStatus::Blocked)
}
pub fn complete(&self, id: TaskId) -> Result<u64, CortexAdapterError> {
self.transition(id, TaskStatus::Done)
}
pub fn fail(&self, id: TaskId) -> Result<u64, CortexAdapterError> {
self.transition(id, TaskStatus::Failed)
}
pub fn advance(&self, id: TaskId) -> Result<u64, CortexAdapterError> {
self.ingest_typed(DISPATCH_TASK_ADVANCED, &AdvancedPayload { id })
}
pub fn retry(&self, id: TaskId) -> Result<u64, CortexAdapterError> {
self.ingest_typed(DISPATCH_TASK_RETRIED, &RetriedPayload { id })
}
pub fn delete(&self, id: TaskId) -> Result<u64, CortexAdapterError> {
self.ingest_typed(DISPATCH_TASK_DELETED, &DeletedPayload { id })
}
pub fn link(&self, parent: TaskId, child: TaskId) -> Result<u64, CortexAdapterError> {
self.ingest_typed(DISPATCH_TASK_LINKED, &LinkedPayload { parent, child })
}
pub fn request_cancel(&self, id: TaskId) -> Result<u64, CortexAdapterError> {
self.ingest_typed(
DISPATCH_TASK_CANCEL_REQUESTED,
&CancelRequestedPayload { id },
)
}
pub fn state(&self) -> Arc<RwLock<WorkflowState>> {
self.inner.state()
}
pub fn get(&self, id: TaskId) -> Option<TaskState> {
self.inner.state().read().get(id)
}
pub fn is_cancel_requested(&self, id: TaskId) -> bool {
self.inner.state().read().is_cancel_requested(id)
}
pub fn subtree(&self, id: TaskId) -> Vec<TaskId> {
self.inner.state().read().subtree(id)
}
pub fn status_counts(&self) -> StatusCounts {
self.inner.state().read().status_counts()
}
pub async fn wait_for_seq(&self, seq: u64) -> Result<(), Option<u64>> {
self.inner.wait_for_seq(seq).await
}
pub fn snapshot(&self) -> Result<(Vec<u8>, Option<u64>), CortexAdapterError> {
let (inner, last_seq) = self.inner.snapshot()?;
let payload = WorkflowSnapshotPayload {
app_seq: self.app_seq.load(Ordering::Acquire),
inner,
};
let bytes = postcard::to_allocvec(&payload).map_err(|e| {
CortexAdapterError::Redex(RedexError::Encode(format!("workflow snapshot wrap: {e}")))
})?;
Ok((bytes, last_seq))
}
pub async fn open_from_snapshot(
redex: &Redex,
origin_hash: u64,
state_bytes: &[u8],
last_seq: Option<u64>,
) -> Result<Self, CortexAdapterError> {
Self::open_from_snapshot_with_config(
redex,
origin_hash,
RedexFileConfig::default(),
state_bytes,
last_seq,
)
.await
}
pub async fn open_from_snapshot_with_config(
redex: &Redex,
origin_hash: u64,
redex_config: RedexFileConfig,
state_bytes: &[u8],
last_seq: Option<u64>,
) -> Result<Self, CortexAdapterError> {
let payload: WorkflowSnapshotPayload = postcard::from_bytes(state_bytes).map_err(|e| {
CortexAdapterError::Redex(RedexError::Encode(format!("workflow snapshot unwrap: {e}")))
})?;
let name = ChannelName::new(WORKFLOW_CHANNEL)
.map_err(|e| CortexAdapterError::Redex(RedexError::Channel(e.to_string())))?;
let app_seq = Arc::new(AtomicU64::new(payload.app_seq));
let fold = WatermarkingFold::new(WorkflowFold, app_seq.clone(), origin_hash);
let inner = CortexAdapter::open_from_snapshot(
redex,
&name,
redex_config.clone(),
CortexAdapterConfig::default(),
fold,
&payload.inner,
last_seq,
)?;
let file = redex.open_file(&name, redex_config)?;
let next_seq = file.next_seq();
if next_seq > 0 {
inner.wait_for_seq(next_seq - 1).await.map_err(|folded| {
CortexAdapterError::FoldStoppedBeforeSeq {
wanted: next_seq - 1,
folded_through: folded,
}
})?;
}
Ok(Self {
inner,
origin_hash,
app_seq,
})
}
fn ingest_typed<T: serde::Serialize>(
&self,
dispatch: u8,
payload: &T,
) -> Result<u64, CortexAdapterError> {
let app_seq = self.app_seq.fetch_add(1, Ordering::AcqRel);
let mut meta = EventMeta::new(dispatch, 0, self.origin_hash, app_seq, 0);
let mut buf = Vec::with_capacity(EVENT_META_SIZE + 64);
buf.extend_from_slice(&meta.to_bytes());
buf = postcard::to_extend(payload, buf)
.map_err(|e| CortexAdapterError::Redex(RedexError::Encode(e.to_string())))?;
let tail = &buf[EVENT_META_SIZE..];
meta.checksum = compute_checksum_with_meta(&meta, tail);
EventMeta::patch_checksum(&mut buf, meta.checksum);
self.inner.ingest_prebuilt(&buf)
}
}
#[cfg(test)]
mod tests {
use super::*;
const ORIGIN: u64 = 0x0F10_0001;
async fn open() -> (Redex, WorkflowAdapter) {
let redex = Redex::new();
let adapter = WorkflowAdapter::open(&redex, ORIGIN).await.unwrap();
(redex, adapter)
}
#[tokio::test]
async fn submit_then_transitions_fold_into_state() {
let (_redex, wf) = open().await;
wf.submit(1).unwrap();
wf.start(1).unwrap();
wf.advance(1).unwrap(); wf.retry(1).unwrap(); let seq = wf.complete(1).unwrap();
wf.wait_for_seq(seq).await.unwrap();
let st = wf.get(1).expect("task present");
assert_eq!(st.step, 1);
assert_eq!(st.attempts, 1);
assert_eq!(st.status, TaskStatus::Done);
assert!(st.status.is_terminal());
}
#[tokio::test]
async fn advance_resets_attempts() {
let (_redex, wf) = open().await;
wf.submit(7).unwrap();
wf.retry(7).unwrap();
wf.retry(7).unwrap(); assert_eq!(
{
wf.wait_for_seq(wf.advance(7).unwrap()).await.unwrap();
wf.get(7).unwrap().attempts
},
0
);
assert_eq!(wf.get(7).unwrap().step, 1);
}
#[tokio::test]
async fn transition_on_unknown_id_is_a_noop() {
let (_redex, wf) = open().await;
let seq = wf.start(42).unwrap(); wf.wait_for_seq(seq).await.unwrap();
assert!(wf.get(42).is_none());
}
#[tokio::test]
async fn terminal_tasks_cannot_be_resurrected() {
let (_redex, wf) = open().await;
wf.submit(1).unwrap();
wf.complete(1).unwrap();
wf.start(1).unwrap();
let seq = wf.retry(1).unwrap();
wf.wait_for_seq(seq).await.unwrap();
assert_eq!(wf.get(1).unwrap().status, TaskStatus::Done);
assert_eq!(
wf.get(1).unwrap().attempts,
0,
"retry didn't bump a Done task"
);
wf.submit(2).unwrap();
wf.fail(2).unwrap();
let seq = wf.start(2).unwrap();
wf.wait_for_seq(seq).await.unwrap();
assert_eq!(wf.get(2).unwrap().status, TaskStatus::Failed);
let seq = wf.retry(2).unwrap();
wf.wait_for_seq(seq).await.unwrap();
assert_eq!(wf.get(2).unwrap().status, TaskStatus::Running);
assert_eq!(wf.get(2).unwrap().attempts, 1);
}
#[tokio::test]
async fn delete_reclaims_the_task() {
let (_redex, wf) = open().await;
wf.submit(3).unwrap();
let seq = wf.delete(3).unwrap();
wf.wait_for_seq(seq).await.unwrap();
assert!(wf.get(3).is_none());
}
#[tokio::test]
async fn delete_cascades_over_the_linked_subtree() {
let (_redex, wf) = open().await;
for id in [1, 2, 3, 4, 10, 11] {
wf.submit(id).unwrap();
}
wf.link(1, 2).unwrap();
wf.link(1, 3).unwrap();
wf.link(3, 4).unwrap();
wf.link(10, 11).unwrap();
let seq = wf.link(10, 11).unwrap(); wf.wait_for_seq(seq).await.unwrap();
assert_eq!(wf.subtree(1), vec![1, 2, 3, 4]);
assert_eq!(wf.state().read().children_of(10), &[11]);
let seq = wf.delete(1).unwrap();
wf.wait_for_seq(seq).await.unwrap();
for id in [1, 2, 3, 4] {
assert!(wf.get(id).is_none(), "subtree member {id} reclaimed");
}
assert!(wf.get(10).is_some());
assert!(wf.get(11).is_some());
}
#[tokio::test]
async fn delete_detaches_child_from_parent_lineage() {
let (_redex, wf) = open().await;
for id in [1, 2, 3] {
wf.submit(id).unwrap();
}
wf.link(1, 2).unwrap();
let seq = wf.link(1, 3).unwrap();
wf.wait_for_seq(seq).await.unwrap();
let seq = wf.delete(2).unwrap();
wf.wait_for_seq(seq).await.unwrap();
assert!(wf.get(2).is_none());
assert_eq!(wf.state().read().children_of(1), &[3]);
assert_eq!(wf.subtree(1), vec![1, 3]);
}
#[tokio::test]
async fn reopen_replays_to_identical_state() {
let redex = Redex::new();
let wf = WorkflowAdapter::open(&redex, ORIGIN).await.unwrap();
wf.submit(1).unwrap();
wf.start(1).unwrap();
wf.advance(1).unwrap();
wf.submit(2).unwrap();
let seq = wf.fail(2).unwrap();
wf.wait_for_seq(seq).await.unwrap();
let resumed = WorkflowAdapter::open(&redex, 0x0F10_0002).await.unwrap();
assert_eq!(resumed.get(1), wf.get(1));
assert_eq!(resumed.get(2), wf.get(2));
assert_eq!(
resumed.get(1).unwrap(),
TaskState {
step: 1,
status: TaskStatus::Running,
attempts: 0
}
);
assert_eq!(resumed.get(2).unwrap().status, TaskStatus::Failed);
}
#[tokio::test]
async fn cancel_signal_observed_then_cleared_on_delete() {
let (_redex, wf) = open().await;
wf.submit(1).unwrap();
wf.start(1).unwrap();
let seq = wf.request_cancel(1).unwrap();
wf.wait_for_seq(seq).await.unwrap();
assert!(wf.is_cancel_requested(1));
assert_eq!(wf.get(1).unwrap().status, TaskStatus::Running);
let seq = wf.fail(1).unwrap();
wf.wait_for_seq(seq).await.unwrap();
assert_eq!(wf.get(1).unwrap().status, TaskStatus::Failed);
let seq = wf.delete(1).unwrap();
wf.wait_for_seq(seq).await.unwrap();
assert!(!wf.is_cancel_requested(1));
assert!(wf.get(1).is_none());
}
#[tokio::test]
async fn fresh_submit_clears_a_stale_cancel() {
let (_redex, wf) = open().await;
wf.submit(1).unwrap();
wf.request_cancel(1).unwrap();
let seq = wf.submit(1).unwrap(); wf.wait_for_seq(seq).await.unwrap();
assert!(
!wf.is_cancel_requested(1),
"re-submit clears the stale cancel"
);
}
#[tokio::test]
async fn status_counts_roll_up() {
let (_redex, wf) = open().await;
wf.submit(1).unwrap(); wf.submit(2).unwrap();
wf.start(2).unwrap(); wf.submit(3).unwrap();
let seq = wf.complete(3).unwrap(); wf.wait_for_seq(seq).await.unwrap();
let c = wf.status_counts();
assert_eq!(c.submitted, 1);
assert_eq!(c.running, 1);
assert_eq!(c.done, 1);
assert_eq!(c.total(), 3);
}
#[tokio::test]
async fn snapshot_then_restore_reproduces_state() {
let redex = Redex::new();
let wf = WorkflowAdapter::open(&redex, ORIGIN).await.unwrap();
wf.submit(1).unwrap();
wf.start(1).unwrap();
wf.advance(1).unwrap();
wf.submit(2).unwrap();
let seq = wf.complete(2).unwrap();
wf.wait_for_seq(seq).await.unwrap();
let (bytes, last_seq) = wf.snapshot().unwrap();
let redex2 = Redex::new();
let restored = WorkflowAdapter::open_from_snapshot(&redex2, ORIGIN, &bytes, last_seq)
.await
.unwrap();
assert_eq!(restored.get(1), wf.get(1));
assert_eq!(restored.get(2), wf.get(2));
assert_eq!(restored.status_counts(), wf.status_counts());
}
}