use aion_core::{ActivityId, WorkflowId};
use aion_proto::{ProtoActivityId, ProtoCancelActivity, ProtoWorkflowId};
use tracing::{info, warn};
use crate::error::ServerError;
use crate::worker::declared_body_cancel::DeclaredCommandAttempts;
use crate::worker::heartbeat::HeartbeatTracker;
use crate::worker::intervention::AttemptKey;
use crate::worker::registry::{ConnectedWorkerRegistry, WorkerId, WorkerMessage};
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum CancelDelivery {
Requested,
WorkerGone,
TransportCannotCarry,
ChannelUnavailable,
}
impl CancelDelivery {
#[must_use]
pub const fn was_requested(self) -> bool {
matches!(self, Self::Requested)
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct CancelRequest {
pub worker_id: WorkerId,
pub workflow_id: WorkflowId,
pub activity_id: ActivityId,
pub attempt: u32,
pub delivery: CancelDelivery,
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct InFlightCancellation {
pub worker_requests: Vec<CancelRequest>,
pub declared_attempts: Vec<AttemptKey>,
}
impl InFlightCancellation {
#[must_use]
pub fn stopped_nothing(&self) -> bool {
self.worker_requests.is_empty() && self.declared_attempts.is_empty()
}
}
pub fn cancel_in_flight_activities(
tracker: &HeartbeatTracker,
registry: &ConnectedWorkerRegistry,
declared: &DeclaredCommandAttempts,
workflow_id: &WorkflowId,
) -> Result<InFlightCancellation, ServerError> {
let in_flight = tracker.in_flight_for_workflow(workflow_id)?;
let mut worker_requests = Vec::with_capacity(in_flight.len());
for liveness in in_flight {
let delivery = ask_worker_to_stop(
registry,
liveness.worker_id,
&liveness.workflow_id,
&liveness.activity_id,
)?;
worker_requests.push(CancelRequest {
worker_id: liveness.worker_id,
workflow_id: liveness.workflow_id,
activity_id: liveness.activity_id,
attempt: liveness.attempt,
delivery,
});
}
let declared_attempts = declared.cancel_workflow(workflow_id)?;
let cancellation = InFlightCancellation {
worker_requests,
declared_attempts,
};
report(workflow_id, &cancellation);
Ok(cancellation)
}
fn ask_worker_to_stop(
registry: &ConnectedWorkerRegistry,
worker_id: WorkerId,
workflow_id: &WorkflowId,
activity_id: &ActivityId,
) -> Result<CancelDelivery, ServerError> {
let Some(worker) = registry.worker_by_id(worker_id)? else {
return Ok(CancelDelivery::WorkerGone);
};
let Some(sender) = worker.sender() else {
return Ok(CancelDelivery::TransportCannotCarry);
};
let message = WorkerMessage::CancelActivity(ProtoCancelActivity {
workflow_id: Some(ProtoWorkflowId {
uuid: workflow_id.to_string(),
}),
activity_id: Some(ProtoActivityId {
sequence_position: activity_id.sequence_position(),
}),
});
if sender.try_send(message).is_ok() {
Ok(CancelDelivery::Requested)
} else {
Ok(CancelDelivery::ChannelUnavailable)
}
}
fn report(workflow_id: &WorkflowId, cancellation: &InFlightCancellation) {
if cancellation.stopped_nothing() {
return;
}
let requests = &cancellation.worker_requests;
let requested = requests
.iter()
.filter(|request| request.delivery.was_requested())
.count();
info!(
workflow_id = %workflow_id,
in_flight = requests.len(),
requested,
declared_attempts = cancellation.declared_attempts.len(),
"stopped the cancelled run's in-flight activities: asked their workers, and \
signalled the declared bodies this server was executing"
);
for request in requests
.iter()
.filter(|request| !request.delivery.was_requested())
{
warn!(
workflow_id = %workflow_id,
worker_id = request.worker_id.value(),
activity_id = request.activity_id.sequence_position(),
attempt = request.attempt,
outcome = ?request.delivery,
"could not ask a worker to stop a cancelled run's activity; \
the work may still be running"
);
}
}
#[cfg(test)]
mod tests {
use super::{
CancelDelivery, CancelRequest, DeclaredCommandAttempts, cancel_in_flight_activities,
};
use crate::error::ServerError;
use crate::worker::heartbeat::{HeartbeatTracker, InFlightActivity};
use crate::worker::registry::{
ConnectedWorkerRegistry, WorkerId, WorkerMessage, WorkerRegistration,
};
use aion_core::{ActivityId, WorkflowId};
use std::time::{Duration, Instant};
use tokio::sync::mpsc::{self, Receiver};
type TestResult = Result<(), Box<dyn std::error::Error>>;
const WINDOW: Duration = Duration::from_secs(5);
struct RegisteredWorker {
registry: ConnectedWorkerRegistry,
registration: WorkerRegistration,
worker_id: WorkerId,
received: Receiver<WorkerMessage>,
}
fn registry_with_worker(capacity: usize) -> Result<RegisteredWorker, ServerError> {
let registry = ConnectedWorkerRegistry::default();
let (sender, received) = mpsc::channel(capacity);
let activity_types = [String::from("work")];
let registration = registry.register("tenant-a", activity_types.iter(), sender)?;
let worker_id = registration
.worker_id()
.ok_or_else(|| ServerError::lock_poisoned("test worker registration"))?;
Ok(RegisteredWorker {
registry,
registration,
worker_id,
received,
})
}
fn track(
tracker: &HeartbeatTracker,
worker_id: WorkerId,
workflow_id: &WorkflowId,
position: u64,
) -> Result<(), ServerError> {
tracker.track_task(
worker_id,
InFlightActivity {
workflow_id: workflow_id.clone(),
activity_id: ActivityId::from_sequence_position(position),
attempt: 1,
completion_token: crate::worker::CompletionToken::for_test(),
},
Instant::now(),
)
}
#[tokio::test]
async fn a_tracked_activity_is_asked_to_stop_on_its_workers_stream() -> TestResult {
let mut worker = registry_with_worker(4)?;
let tracker = HeartbeatTracker::new(WINDOW);
let workflow_id = WorkflowId::new_v4();
track(&tracker, worker.worker_id, &workflow_id, 7)?;
let declared = DeclaredCommandAttempts::new(crate::shutdown::DrainState::default());
let requests =
cancel_in_flight_activities(&tracker, &worker.registry, &declared, &workflow_id)?
.worker_requests;
assert_eq!(requests.len(), 1, "one tracked activity, one ask");
assert_eq!(requests[0].delivery, CancelDelivery::Requested);
assert_eq!(
requests[0].activity_id,
ActivityId::from_sequence_position(7)
);
let Ok(WorkerMessage::CancelActivity(cancel)) = worker.received.try_recv() else {
return Err("the worker's stream did not carry a cancel".into());
};
assert_eq!(
cancel.workflow_id.map(|id| id.uuid),
Some(workflow_id.to_string()),
"the cancel must name the workflow it is cancelling"
);
assert_eq!(
cancel.activity_id.map(|id| id.sequence_position),
Some(7),
"the cancel must name the activity it is cancelling"
);
Ok(())
}
#[tokio::test]
async fn a_workflow_with_nothing_in_flight_asks_nobody() -> TestResult {
let worker = registry_with_worker(4)?;
let tracker = HeartbeatTracker::new(WINDOW);
let declared = DeclaredCommandAttempts::new(crate::shutdown::DrainState::default());
let requests = cancel_in_flight_activities(
&tracker,
&worker.registry,
&declared,
&WorkflowId::new_v4(),
)?
.worker_requests;
assert!(
requests.is_empty(),
"an untracked workflow must produce no asks, not a synthesized one"
);
Ok(())
}
#[tokio::test]
async fn another_workflows_activity_is_never_asked_to_stop() -> TestResult {
let mut worker = registry_with_worker(4)?;
let tracker = HeartbeatTracker::new(WINDOW);
let cancelled = WorkflowId::new_v4();
let bystander = WorkflowId::new_v4();
track(&tracker, worker.worker_id, &cancelled, 1)?;
track(&tracker, worker.worker_id, &bystander, 1)?;
let declared = DeclaredCommandAttempts::new(crate::shutdown::DrainState::default());
let requests =
cancel_in_flight_activities(&tracker, &worker.registry, &declared, &cancelled)?
.worker_requests;
assert_eq!(requests.len(), 1, "only the cancelled run's activity");
let Ok(WorkerMessage::CancelActivity(cancel)) = worker.received.try_recv() else {
return Err("the worker's stream did not carry a cancel".into());
};
assert_eq!(
cancel.workflow_id.map(|id| id.uuid),
Some(cancelled.to_string()),
"the bystander workflow must not be named in any cancel"
);
assert!(
worker.received.try_recv().is_err(),
"exactly one cancel; the bystander's activity was asked to stop too"
);
Ok(())
}
#[tokio::test]
async fn a_departed_worker_is_reported_rather_than_skipped() -> TestResult {
let worker = registry_with_worker(4)?;
let tracker = HeartbeatTracker::new(WINDOW);
let workflow_id = WorkflowId::new_v4();
track(&tracker, worker.worker_id, &workflow_id, 1)?;
worker.registration.deregister()?;
let declared = DeclaredCommandAttempts::new(crate::shutdown::DrainState::default());
let requests: Vec<CancelRequest> =
cancel_in_flight_activities(&tracker, &worker.registry, &declared, &workflow_id)?
.worker_requests;
assert_eq!(requests.len(), 1, "the entry is reported, not dropped");
assert_eq!(
requests[0].delivery,
CancelDelivery::WorkerGone,
"an unaskable worker must be named, never counted as asked"
);
assert!(
!requests[0].delivery.was_requested(),
"a departed worker was never asked"
);
Ok(())
}
#[cfg(feature = "liminal-transport")]
#[tokio::test]
async fn a_liminal_worker_is_reported_as_transport_cannot_carry() -> TestResult {
use crate::worker::liminal_transport::LiminalWorkerDelivery;
use crate::worker::registry::WorkerDelivery;
let registry = ConnectedWorkerRegistry::default();
let supervisor = liminal_server::server::connection::ConnectionSupervisor::new()?;
let activity_types = [String::from("work")];
let registration = registry.register_delivery(
[String::from("tenant-a")],
String::from("default"),
None,
None,
activity_types.iter(),
WorkerDelivery::Liminal(LiminalWorkerDelivery::new(supervisor, 7)),
)?;
let worker_id = registration
.worker_id()
.ok_or_else(|| ServerError::lock_poisoned("test worker registration"))?;
let tracker = HeartbeatTracker::new(WINDOW);
let workflow_id = WorkflowId::new_v4();
track(&tracker, worker_id, &workflow_id, 1)?;
let declared = DeclaredCommandAttempts::new(crate::shutdown::DrainState::default());
let requests = cancel_in_flight_activities(&tracker, ®istry, &declared, &workflow_id)?
.worker_requests;
assert_eq!(requests.len(), 1, "the entry is reported, not dropped");
assert_eq!(
requests[0].delivery,
CancelDelivery::TransportCannotCarry,
"a control-less transport must be named as such, not as gone or asked"
);
drop(registration);
Ok(())
}
#[tokio::test]
async fn a_closed_dispatch_channel_is_reported_as_unavailable() -> TestResult {
let worker = registry_with_worker(1)?;
let tracker = HeartbeatTracker::new(WINDOW);
let workflow_id = WorkflowId::new_v4();
track(&tracker, worker.worker_id, &workflow_id, 1)?;
drop(worker.received);
let declared = DeclaredCommandAttempts::new(crate::shutdown::DrainState::default());
let requests =
cancel_in_flight_activities(&tracker, &worker.registry, &declared, &workflow_id)?
.worker_requests;
assert_eq!(
requests[0].delivery,
CancelDelivery::ChannelUnavailable,
"a dead channel must not read as a delivered cancel"
);
Ok(())
}
}