use std::{collections::BTreeMap, fmt, sync::Arc};
use runifold_core::{
Checkpoint, CheckpointError, CheckpointErrorKind, CheckpointId, CheckpointStore, RunContext,
Usage,
};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::{StepId, WorkflowError, WorkflowOutcome};
const CHECKPOINT_KIND: &str = "runifold.workflow";
const CHECKPOINT_SCHEMA_VERSION: u32 = 3;
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
#[non_exhaustive]
pub enum WorkflowResumePolicy {
#[default]
RejectAmbiguous,
RetryInterruptedStep,
}
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
#[serde(tag = "state", rename_all = "snake_case")]
#[non_exhaustive]
pub enum WorkflowCheckpointPhase {
Ready,
StepInFlight {
step: StepId,
},
ParallelInFlight {
step: StepId,
branches: BTreeMap<StepId, ParallelBranchCheckpoint>,
},
RaceInFlight {
step: StepId,
branches: BTreeMap<StepId, ParallelBranchCheckpoint>,
},
Completed {
outcome: WorkflowOutcome,
},
}
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
#[serde(tag = "state", rename_all = "snake_case")]
#[non_exhaustive]
pub enum ParallelBranchCheckpoint {
InFlight,
Completed {
output: Value,
},
Failed {
message: String,
},
Cancelled,
}
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
pub struct WorkflowCheckpointState {
pub workflow: String,
pub workflow_version: u32,
pub layout: Vec<StepId>,
pub next_index: usize,
pub value: Value,
pub outputs: BTreeMap<StepId, Value>,
pub usage: Usage,
pub phase: WorkflowCheckpointPhase,
}
impl WorkflowCheckpointState {
pub(crate) fn outcome(&self) -> Option<WorkflowOutcome> {
match &self.phase {
WorkflowCheckpointPhase::Completed { outcome } => Some(outcome.clone()),
_ => None,
}
}
}
#[derive(Clone)]
pub struct WorkflowCheckpoint {
id: CheckpointId,
store: Arc<dyn CheckpointStore>,
}
impl WorkflowCheckpoint {
pub fn new(store: Arc<dyn CheckpointStore>) -> Self {
Self {
id: CheckpointId::new(),
store,
}
}
pub fn existing(id: CheckpointId, store: Arc<dyn CheckpointStore>) -> Self {
Self { id, store }
}
pub const fn id(&self) -> CheckpointId {
self.id
}
pub fn load(&self) -> Result<(Checkpoint, WorkflowCheckpointState), CheckpointError> {
let checkpoint = self.store.load(self.id)?;
if checkpoint.kind != CHECKPOINT_KIND
|| checkpoint.schema_version != CHECKPOINT_SCHEMA_VERSION
{
return Err(CheckpointError::new(
CheckpointErrorKind::InvalidPayload,
"checkpoint kind or schema version is not supported",
));
}
let state = serde_json::from_value(checkpoint.payload.clone()).map_err(|error| {
CheckpointError::new(CheckpointErrorKind::InvalidPayload, error.to_string())
})?;
Ok((checkpoint, state))
}
}
impl fmt::Debug for WorkflowCheckpoint {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("WorkflowCheckpoint")
.field("id", &self.id)
.finish_non_exhaustive()
}
}
pub(crate) struct WorkflowCheckpointCursor {
handle: WorkflowCheckpoint,
envelope: Checkpoint,
}
impl WorkflowCheckpointCursor {
pub(crate) fn create(
handle: &WorkflowCheckpoint,
run: &RunContext,
state: &WorkflowCheckpointState,
) -> Result<Self, WorkflowError> {
let envelope = Checkpoint::initial(
handle.id,
run.run_id(),
CHECKPOINT_KIND,
CHECKPOINT_SCHEMA_VERSION,
serialize(state)?,
);
handle.store.compare_and_swap(&envelope, None)?;
Ok(Self {
handle: handle.clone(),
envelope,
})
}
pub(crate) fn loaded(handle: &WorkflowCheckpoint, envelope: Checkpoint) -> Self {
Self {
handle: handle.clone(),
envelope,
}
}
pub(crate) fn save(&mut self, state: &WorkflowCheckpointState) -> Result<(), WorkflowError> {
let next = self.envelope.next(serialize(state)?)?;
self.handle
.store
.compare_and_swap(&next, Some(self.envelope.revision))?;
self.envelope = next;
Ok(())
}
}
fn serialize(state: &WorkflowCheckpointState) -> Result<Value, WorkflowError> {
serde_json::to_value(state).map_err(|error| {
CheckpointError::new(CheckpointErrorKind::InvalidPayload, error.to_string()).into()
})
}