use std::fmt::Display;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::time::{Duration, SystemTime};
use serde::de::DeserializeOwned;
use serde_json::{Value, json};
use tracing::{Instrument, debug, field, info_span, warn};
use crate::effect::EffectFailure;
use crate::error::RuntimeError;
use crate::failure::{Disposition, FailureClass};
use crate::fault::FaultPoint;
use crate::id::{EffectId, EffectKey, EffectName, IdempotencyKey, LogicalKey};
use crate::retry::RetryPolicy;
use crate::runtime::{Interrupt, Runtime, jitter_sample, last_error};
use crate::state::{EffectStatus, Transition};
use crate::store::{EffectRecord, EffectStore, ErrorRecord, Lease, StoreError};
const MAX_ROUNDS: usize = 4;
#[derive(Clone, Debug)]
pub struct CompensationContext {
pub(crate) id: EffectId,
pub(crate) key: EffectKey,
pub(crate) attempt: u32,
pub(crate) reason: Option<String>,
}
impl CompensationContext {
pub fn effect_id(&self) -> EffectId {
self.id
}
pub fn key(&self) -> &EffectKey {
&self.key
}
pub fn idempotency_key(&self) -> IdempotencyKey {
self.key.compensation_idempotency_key()
}
pub fn attempt(&self) -> u32 {
self.attempt
}
pub fn reason(&self) -> Option<&str> {
self.reason.as_deref()
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub enum CompensationOutcome {
Compensated,
Failed(ErrorRecord),
InProgress {
id: EffectId,
},
NotCommitted {
id: EffectId,
status: EffectStatus,
},
}
type BoxFuture<T> = Pin<Box<dyn Future<Output = T> + Send>>;
pub(crate) type Compensator = Arc<
dyn Fn(
CompensationContext,
Option<Value>,
Option<Value>,
) -> BoxFuture<Result<(), EffectFailure>>
+ Send
+ Sync,
>;
pub(crate) struct CompensationSpec {
pub(crate) key: EffectKey,
pub(crate) reason: Option<String>,
pub(crate) actor: Option<String>,
pub(crate) retry: Option<RetryPolicy>,
pub(crate) attempt_timeout: Option<Duration>,
}
#[must_use = "a compensation does nothing until `run` is awaited"]
pub struct CompensationBuilder<'a, S> {
runtime: &'a Runtime<S>,
name: String,
key: String,
reason: Option<String>,
actor: Option<String>,
retry: Option<RetryPolicy>,
attempt_timeout: Option<Duration>,
}
impl<'a, S: EffectStore> CompensationBuilder<'a, S> {
pub(crate) fn new(runtime: &'a Runtime<S>, name: String, key: String) -> Self {
Self {
runtime,
name,
key,
reason: None,
actor: None,
retry: None,
attempt_timeout: None,
}
}
pub fn reason(mut self, reason: impl Into<String>) -> Self {
self.reason = Some(reason.into());
self
}
pub fn actor(mut self, actor: impl Into<String>) -> Self {
self.actor = Some(actor.into());
self
}
pub fn retry(mut self, policy: RetryPolicy) -> Self {
self.retry = Some(policy);
self
}
pub fn attempt_timeout(mut self, timeout: Duration) -> Self {
self.attempt_timeout = Some(timeout);
self
}
pub async fn run<T, F, Fut>(self, compensate: F) -> Result<CompensationOutcome, RuntimeError>
where
T: DeserializeOwned + Send + 'static,
F: Fn(CompensationContext, Option<T>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<(), EffectFailure>> + Send + 'static,
{
let key = EffectKey::new(EffectName::new(self.name)?, LogicalKey::new(self.key)?);
let compensate: Compensator = Arc::new(move |ctx, _input, output| {
match output.map(serde_json::from_value::<T>).transpose() {
Ok(output) => Box::pin(compensate(ctx, output)),
Err(e) => Box::pin(std::future::ready(Err(EffectFailure::permanent(format!(
"stored output does not match the compensation's type: {e}"
))))),
}
});
let spec = CompensationSpec {
key,
reason: self.reason,
actor: self.actor,
retry: self.retry,
attempt_timeout: self.attempt_timeout,
};
self.runtime.compensate_effect(spec, compensate).await
}
}
impl<S: EffectStore> Runtime<S> {
pub fn compensation(
&self,
name: impl Into<String>,
key: impl Display,
) -> CompensationBuilder<'_, S> {
CompensationBuilder::new(self, name.into(), key.to_string())
}
pub(crate) async fn compensate_effect(
&self,
spec: CompensationSpec,
compensate: Compensator,
) -> Result<CompensationOutcome, RuntimeError> {
let span = info_span!(
"agent_effect.compensate",
effect.name = %spec.key.name,
effect.logical_key = %spec.key.key,
effect.id = field::Empty,
);
let runtime = self.clone();
tokio::spawn(
async move { runtime.drive_compensation(spec, compensate).await }.instrument(span),
)
.await
.unwrap_or_else(|e| Err(RuntimeError::Internal(e.to_string())))
}
async fn drive_compensation(
&self,
spec: CompensationSpec,
compensate: Compensator,
) -> Result<CompensationOutcome, RuntimeError> {
let store = self.store();
let mut record =
store
.get_by_key(&spec.key)
.await?
.ok_or_else(|| RuntimeError::NoSuchEffect {
key: spec.key.to_string(),
})?;
tracing::Span::current().record("effect.id", field::display(record.id));
for _ in 0..MAX_ROUNDS {
if let Some(outcome) = observe_compensation(&record) {
return Ok(outcome);
}
let lease = match store
.acquire_lease(record.id, self.worker_id(), self.now(), self.lease_ttl())
.await
{
Ok(lease) => lease,
Err(StoreError::LeaseHeld { .. }) => {
return Ok(CompensationOutcome::InProgress { id: record.id });
}
Err(e) => return Err(e.into()),
};
let current = store
.get(record.id)
.await?
.ok_or(StoreError::NotFound(record.id))?;
let result = self
.compensate_leased(current, &lease, &spec, &compensate)
.await;
if let Err(e) = store.release_lease(&lease).await {
warn!(error = %e, "could not release lease; it will expire");
}
match result {
Ok(settled) => {
return Ok(observe_compensation(&settled)
.unwrap_or(CompensationOutcome::InProgress { id: settled.id }));
}
Err(Interrupt::LeaseLost) => {
record = store
.get(record.id)
.await?
.ok_or(StoreError::NotFound(record.id))?;
}
Err(Interrupt::Error(e)) => return Err(e),
}
}
Ok(CompensationOutcome::InProgress { id: record.id })
}
async fn resume_compensation(
&self,
record: &EffectRecord,
lease: &Lease,
actor: Option<&str>,
policy: RetryPolicy,
) -> Result<EffectRecord, Interrupt> {
let resumed = match record.next_attempt_at {
Some(at) => {
self.sleep_leased(lease, at).await?;
None
}
None if policy.allows_another(record.compensation_attempts) => {
Some(json!({ "resumed": true }))
}
None => {
return self
.transition_leased(record, lease, actor, Transition::CompensationFailed, |r| {
r.error = Some(ErrorRecord {
class: Some(FailureClass::Ambiguous),
message: "a compensation attempt was interrupted, and no retries \
are left"
.into(),
});
})
.await;
}
};
self.transition_leased(
record,
lease,
actor,
Transition::StartCompensationRetry,
|r| r.payload = resumed,
)
.await
}
async fn compensate_leased(
&self,
record: EffectRecord,
lease: &Lease,
spec: &CompensationSpec,
compensate: &Compensator,
) -> Result<EffectRecord, Interrupt> {
let actor = spec.actor.as_deref();
let policy = spec.retry.unwrap_or_else(|| self.default_retry());
let reason = spec.reason.as_ref().map(|r| json!({ "reason": r }));
let mut record = match record.status {
EffectStatus::Committed => {
self.transition_leased(&record, lease, actor, Transition::StartCompensation, |r| {
r.payload = reason;
})
.await?
}
EffectStatus::Compensating => {
let resumed = self
.resume_compensation(&record, lease, actor, policy)
.await?;
if resumed.status != EffectStatus::Compensating {
return Ok(resumed);
}
resumed
}
_ => return Ok(record),
};
self.checkpoint(FaultPoint::AfterCompensationStarted);
loop {
let failure = match self
.attempt_compensation(&record, lease, spec, compensate)
.await?
{
Ok(()) => {
return self
.transition_leased(
&record,
lease,
actor,
Transition::CompensationSucceeded,
|_| {},
)
.await;
}
Err(failure) => failure,
};
debug!(%failure, "compensation attempt failed");
let class = failure.class();
let error = failure.to_record();
let retryable = !matches!(class.disposition(), Disposition::Fail);
if !(retryable && policy.allows_another(record.compensation_attempts)) {
return self
.transition_leased(&record, lease, actor, Transition::CompensationFailed, |r| {
r.error = Some(error);
})
.await;
}
let retry_class = if class == FailureClass::Ambiguous {
FailureClass::Transient
} else {
class
};
let delay = policy.delay(
record.compensation_attempts.saturating_sub(1),
retry_class,
jitter_sample(),
);
let at = self.now() + delay;
record = self
.transition_leased(
&record,
lease,
actor,
Transition::ScheduleCompensationRetry,
|r| {
r.next_attempt_at = Some(at);
r.error = Some(error);
},
)
.await?;
self.sleep_leased(lease, at).await?;
record = self
.transition_leased(
&record,
lease,
actor,
Transition::StartCompensationRetry,
|_| {},
)
.await?;
}
}
async fn attempt_compensation(
&self,
record: &EffectRecord,
lease: &Lease,
spec: &CompensationSpec,
compensate: &Compensator,
) -> Result<Result<(), EffectFailure>, Interrupt> {
let ctx = CompensationContext {
id: record.id,
key: record.key.clone(),
attempt: record.compensation_attempts,
reason: spec.reason.clone(),
};
let mut task = tokio::spawn(compensate(ctx, record.input.clone(), record.output.clone()));
let joined = match spec.attempt_timeout {
None => self.with_lease(lease, &mut task).await?,
Some(limit) => {
match self
.with_lease(lease, tokio::time::timeout(limit, &mut task))
.await?
{
Ok(joined) => joined,
Err(_elapsed) => {
task.abort();
return Ok(Err(EffectFailure::ambiguous(format!(
"compensation attempt timed out after {limit:?}"
))));
}
}
}
};
Ok(joined.unwrap_or_else(|join_error| {
Err(EffectFailure::ambiguous(format!(
"compensation did not complete: {join_error}"
)))
}))
}
async fn sleep_leased(&self, lease: &Lease, until: SystemTime) -> Result<(), Interrupt> {
let wait = until.duration_since(self.now()).unwrap_or_default();
if wait.is_zero() {
return Ok(());
}
self.with_lease(lease, tokio::time::sleep(wait)).await
}
}
fn observe_compensation(record: &EffectRecord) -> Option<CompensationOutcome> {
let id = record.id;
match record.status {
EffectStatus::Compensated => Some(CompensationOutcome::Compensated),
EffectStatus::CompensationFailed => Some(CompensationOutcome::Failed(last_error(record))),
EffectStatus::Committed | EffectStatus::Compensating => None,
status => Some(CompensationOutcome::NotCommitted { id, status }),
}
}