use std::sync::Arc;
use async_trait::async_trait;
use aion_core::{ActivityId, RunId, WorkflowId};
use aion_proto::ProtoActivityTask;
use super::delivery_intent::SharedDeliveryIntent;
use super::intervention::{AttemptKey, AttemptOwnerIndex};
use super::liminal_transport::{AttemptOwnerGuard, DispatchRequest, LiminalCompletionSource};
use super::registry::{WorkerDelivery, WorkerHandle};
use super::task_delivery::{DeliveryAccepted, LivenessTracking, TaskDelivery, WorkerTaskDelivery};
pub struct LiminalTaskDelivery {
completion: Arc<LiminalCompletionSource>,
attempt_owners: Option<AttemptOwnerIndex>,
completion_tracking: Option<CompletionTracking>,
}
#[derive(Clone)]
struct CompletionTracking {
heartbeat_tracker: crate::worker::heartbeat::HeartbeatTracker,
registry: crate::worker::ConnectedWorkerRegistry,
}
impl std::fmt::Debug for LiminalTaskDelivery {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("LiminalTaskDelivery")
.field("attempt_owners", &self.attempt_owners.is_some())
.finish_non_exhaustive()
}
}
impl LiminalTaskDelivery {
#[must_use]
pub fn new(completion: Arc<LiminalCompletionSource>) -> Self {
Self {
completion,
attempt_owners: None,
completion_tracking: None,
}
}
#[must_use]
pub fn with_completion_tracking(
mut self,
heartbeat_tracker: crate::worker::heartbeat::HeartbeatTracker,
registry: crate::worker::ConnectedWorkerRegistry,
) -> Self {
self.completion_tracking = Some(CompletionTracking {
heartbeat_tracker,
registry,
});
self
}
#[must_use]
pub fn with_attempt_owners(mut self, attempt_owners: AttemptOwnerIndex) -> Self {
self.attempt_owners = Some(attempt_owners);
self
}
}
struct AbandonedDispatchGuard<'a> {
tracking: &'a CompletionTracking,
worker_id: crate::worker::WorkerId,
workflow_id: &'a WorkflowId,
activity_id: &'a ActivityId,
armed: bool,
}
impl AbandonedDispatchGuard<'_> {
fn disarm(&mut self) {
self.armed = false;
}
}
impl Drop for AbandonedDispatchGuard<'_> {
fn drop(&mut self) {
if !self.armed {
return;
}
let retired = crate::worker::bridge::clear_completed_task_tracking(
&self.tracking.heartbeat_tracker,
&self.tracking.registry,
self.worker_id,
self.workflow_id,
self.activity_id,
);
if retired {
tracing::warn!(
worker_id = ?self.worker_id,
workflow_id = %self.workflow_id,
activity_id = %self.activity_id,
"liminal delivery was abandoned before the worker replied; its tracked slot has \
been returned and the row re-drives"
);
}
}
}
struct Resolved {
workflow_id: WorkflowId,
run_id: RunId,
activity_id: ActivityId,
attempt: u32,
}
impl Resolved {
fn from_task(task: &ProtoActivityTask) -> Result<Self, &'static str> {
let workflow_id = task
.workflow_id
.clone()
.ok_or("task carries no workflow id")?
.try_into()
.map_err(|_| "task carries a malformed workflow id")?;
let run_id: RunId = task
.run_id
.clone()
.ok_or("activity run id is missing; refusing unfenced external effect")?
.try_into()
.map_err(|_| "task carries a malformed run id")?;
let activity_id: ActivityId = task
.activity_id
.ok_or("task carries no activity id")?
.into();
Ok(Self {
workflow_id,
run_id,
activity_id,
attempt: task.attempt,
})
}
fn attempt_key(&self) -> AttemptKey {
AttemptKey::new(
self.workflow_id.clone(),
self.run_id.clone(),
self.activity_id.clone(),
self.attempt,
)
}
fn bind_owner(
&self,
owners: Option<&AttemptOwnerIndex>,
worker: super::registry::WorkerId,
) -> Option<AttemptOwnerGuard> {
owners.map(|owners| AttemptOwnerGuard::bind(owners.clone(), self.attempt_key(), worker))
}
fn ordinal(&self) -> u64 {
self.activity_id.sequence_position()
}
}
fn undelivered_by(error: &crate::error::ServerError) -> TaskDelivery {
if error.is_worker_dispatch_unservable() {
TaskDelivery::unservable(format!("liminal dispatch is unservable: {error}"))
} else {
TaskDelivery::failed(format!("liminal dispatch failed: {error}"))
}
}
fn request_for_task(task: &ProtoActivityTask, resolved: &Resolved) -> DispatchRequest {
DispatchRequest {
activity_type: task.activity_type.clone(),
workflow_id: resolved.workflow_id.clone(),
ordinal: resolved.ordinal(),
run_id: Some(resolved.run_id.clone()),
completion_token: task.completion_token.clone(),
idempotency_key: task.idempotency_key.clone(),
input: task
.input
.as_ref()
.map(|payload| payload.bytes.clone())
.unwrap_or_default(),
attempt: resolved.attempt,
labels: task
.labels
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect(),
heartbeat_window_ms: LivenessTracking::NotTrackedPerTask.heartbeat_window_ms(),
}
}
#[async_trait]
impl WorkerTaskDelivery for LiminalTaskDelivery {
async fn deliver(
&self,
worker: &WorkerHandle,
task: &ProtoActivityTask,
intent: &SharedDeliveryIntent,
accepted: &dyn DeliveryAccepted,
) -> TaskDelivery {
let resolved = match Resolved::from_task(task) {
Ok(resolved) => resolved,
Err(reason) => return TaskDelivery::failed(reason),
};
let delivery = match worker.delivery() {
WorkerDelivery::Liminal(delivery) => delivery.clone(),
WorkerDelivery::Grpc(_) => {
return TaskDelivery::failed(
"selected worker is not delivered over liminal; the liminal transport cannot \
reach it",
);
}
};
let _owner_guard = resolved.bind_owner(self.attempt_owners.as_ref(), worker.id());
let request = request_for_task(task, &resolved);
let mut abandoned =
self.completion_tracking
.as_ref()
.map(|tracking| AbandonedDispatchGuard {
tracking,
worker_id: worker.id(),
workflow_id: &resolved.workflow_id,
activity_id: &resolved.activity_id,
armed: true,
});
let push_delivery = delivery.clone();
let pushed =
tokio::task::spawn_blocking(move || push_delivery.push_dispatch(&request)).await;
let awaiter = match pushed {
Ok(Ok(awaiter)) => awaiter,
Ok(Err(error)) => {
if let Some(guard) = abandoned.as_mut() {
guard.disarm();
}
return undelivered_by(&error);
}
Err(error) => {
if let Some(guard) = abandoned.as_mut() {
guard.disarm();
}
return TaskDelivery::failed(format!("dispatch task join failed: {error}"));
}
};
accepted.accepted().await;
let waiting_intent = Arc::clone(intent);
let dispatched = tokio::task::spawn_blocking(move || {
super::liminal_transport::receive_bridge_reply(&awaiter, || {
waiting_intent.still_wanted()
})
})
.await;
let response = match dispatched {
Ok(Ok(Some(response))) => response,
Ok(Ok(None)) => {
return TaskDelivery::failed("delivery wait abandoned before worker reply");
}
Ok(Err(error)) => return undelivered_by(&error),
Err(error) => {
return TaskDelivery::failed(format!("dispatch task join failed: {error}"));
}
};
if let Err(error) = self.completion.deliver(&response) {
return TaskDelivery::failed(format!(
"worker replied but the completion could not be recorded: {error}"
));
}
if let Some(tracking) = &self.completion_tracking {
let _ = crate::worker::bridge::clear_completed_task_tracking(
&tracking.heartbeat_tracker,
&tracking.registry,
worker.id(),
&resolved.workflow_id,
&resolved.activity_id,
);
}
if let Some(guard) = abandoned.as_mut() {
guard.disarm();
}
TaskDelivery::Delivered
}
}
#[cfg(test)]
mod tests {
use crate::worker::registry::RegistrationOptions;
#[derive(Default)]
struct NeverAccepted {
reached: std::sync::atomic::AtomicBool,
}
#[async_trait::async_trait]
impl super::DeliveryAccepted for NeverAccepted {
async fn accepted(&self) {
self.reached
.store(true, std::sync::atomic::Ordering::SeqCst);
}
}
use std::sync::Arc;
use aion_core::{ActivityId, InterventionCapabilities, RunId, WorkflowId};
use aion_proto::{ProtoActivityId, ProtoActivityTask, ProtoWorkflowId};
use uuid::Uuid;
use crate::error::ServerError;
use crate::worker::bridge::OutboxDeliveryCallback;
use crate::worker::delivery_intent::{AlwaysWanted, SharedDeliveryIntent};
use crate::worker::intervention::AttemptOwnerIndex;
use crate::worker::liminal_transport::LiminalCompletionSource;
use crate::worker::registry::{ConnectedWorkerRegistry, WorkerDelivery};
use crate::worker::task_delivery::{TaskDelivery, WorkerTaskDelivery};
use super::{LiminalTaskDelivery, Resolved};
struct NoopCallback;
impl OutboxDeliveryCallback for NoopCallback {
fn deliver_completion(
&self,
_workflow_id: &WorkflowId,
_activity_id: &ActivityId,
_run_id: Option<&RunId>,
_result: String,
) -> Result<bool, ServerError> {
Ok(false)
}
fn deliver_failure(
&self,
_workflow_id: &WorkflowId,
_activity_id: &ActivityId,
_run_id: Option<&RunId>,
_reason: String,
) -> Result<bool, ServerError> {
Ok(false)
}
}
const WORKFLOW: u128 = 0x51;
fn task_without_a_run() -> ProtoActivityTask {
ProtoActivityTask {
workflow_id: Some(ProtoWorkflowId::from(WorkflowId::new(Uuid::from_u128(
WORKFLOW,
)))),
activity_id: Some(ProtoActivityId::from(ActivityId::from_sequence_position(3))),
activity_type: String::from("agent"),
input: None,
attempt: 1,
labels: std::collections::HashMap::new(),
run_id: None,
completion_token: String::from("token"),
idempotency_key: String::from("key"),
}
}
fn liminal_delivery(owners: &AttemptOwnerIndex) -> LiminalTaskDelivery {
LiminalTaskDelivery::new(Arc::new(LiminalCompletionSource::new(Arc::new(
NoopCallback,
))))
.with_attempt_owners(owners.clone())
}
#[tokio::test]
async fn run_resolution_precedes_attempt_owner_binding()
-> Result<(), Box<dyn std::error::Error>> {
let owners = AttemptOwnerIndex::new();
let delivery = liminal_delivery(&owners);
let registry = ConnectedWorkerRegistry::default();
let (sender, _receiver) = tokio::sync::mpsc::channel(1);
let types = [String::from("agent")];
let registration = registry.register_delivery(
[String::from("default")],
String::from("default"),
None,
types.iter(),
WorkerDelivery::Grpc(sender),
RegistrationOptions::identified(
"agent-1",
crate::worker::UNBOUNDED_SENDER_WORKER_CONCURRENCY,
)
.with_intervention_capabilities(InterventionCapabilities::none()),
)?;
let worker_id = registration
.worker_id()
.ok_or("a registration must assign a worker id")?;
let worker = registry
.worker_by_id(worker_id)?
.ok_or("the worker just registered must be readable")?;
let intent: SharedDeliveryIntent = Arc::new(AlwaysWanted);
let accepted = NeverAccepted::default();
let outcome = delivery
.deliver(&worker, &task_without_a_run(), &intent, &accepted)
.await;
assert!(
!accepted.reached.load(std::sync::atomic::Ordering::SeqCst),
"a task refused by its missing run must never reach the accept point"
);
match outcome {
TaskDelivery::Delivered => {
return Err("a task with no run must never be delivered".into());
}
TaskDelivery::Undeliverable(undeliverable) => {
assert!(
undeliverable.reason().contains("run id is missing"),
"a run-less task must be refused BY ITS MISSING RUN, not by the transport; \
got: {}",
undeliverable.reason()
);
assert!(
!undeliverable.deregisters_worker(),
"a malformed task is not evidence that the worker is gone"
);
}
}
let bound = owners.attempts_for_workflow(&WorkflowId::new(Uuid::from_u128(WORKFLOW)));
assert!(
bound.is_empty(),
"the attempt-owner index must be untouched when the run refused; a binding here \
means the bind was lifted above the run resolution, and an intervention could \
resolve the wrong generation's worker. Found: {bound:?}"
);
Ok(())
}
fn task_with_a_run() -> ProtoActivityTask {
ProtoActivityTask {
run_id: Some(aion_proto::ProtoRunId::from(RunId::new(Uuid::from_u128(
0x52,
)))),
..task_without_a_run()
}
}
#[test]
fn a_delivery_abandoned_during_its_push_returns_the_tracked_slot()
-> Result<(), Box<dyn std::error::Error>> {
use std::time::{Duration, Instant};
use crate::worker::envelope::CompletionToken;
use crate::worker::heartbeat::{HeartbeatTracker, InFlightActivity};
use crate::worker::liminal_transport::LiminalWorkerDelivery;
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_all()
.max_blocking_threads(1)
.build()?;
runtime.block_on(async {
let (open_gate, gate) = std::sync::mpsc::channel::<()>();
let occupant = tokio::task::spawn_blocking(move || gate.recv());
let registry = ConnectedWorkerRegistry::default();
let tracker = HeartbeatTracker::new(Duration::from_secs(30));
let supervisor = liminal_server::server::connection::ConnectionSupervisor::new()?;
let types = [String::from("agent")];
let registration = registry.register_delivery(
[String::from("default")],
String::from("default"),
None,
types.iter(),
WorkerDelivery::Liminal(LiminalWorkerDelivery::new(supervisor, 7)),
RegistrationOptions::identified("agent-1", 1)
.with_intervention_capabilities(InterventionCapabilities::none()),
)?;
let worker_id = registration
.worker_id()
.ok_or("a registration must assign a worker id")?;
let worker = registry
.worker_by_id(worker_id)?
.ok_or("the worker just registered must be readable")?;
let workflow_id = WorkflowId::new(Uuid::from_u128(WORKFLOW));
let activity_id = ActivityId::from_sequence_position(3);
tracker.track_task(
worker_id,
InFlightActivity {
workflow_id: workflow_id.clone(),
activity_id: activity_id.clone(),
attempt: 1,
completion_token: CompletionToken::for_test(),
},
Instant::now(),
®istry,
None,
)?;
assert_eq!(
registry.in_flight_for_worker(worker_id)?,
1,
"precondition: the tracked dispatch holds the worker's slot"
);
let owners = AttemptOwnerIndex::new();
let delivery = liminal_delivery(&owners)
.with_completion_tracking(tracker.clone(), registry.clone());
let intent: SharedDeliveryIntent = Arc::new(AlwaysWanted);
let accepted = NeverAccepted::default();
let task = task_with_a_run();
{
let mut future = Box::pin(delivery.deliver(&worker, &task, &intent, &accepted));
let first_poll = futures::poll!(future.as_mut());
assert!(
first_poll.is_pending(),
"the push is queued behind the occupied blocking thread, so the first poll \
must suspend AT the push; a ready outcome means the window this test \
holds open was not held"
);
drop(future);
}
assert!(
!accepted.reached.load(std::sync::atomic::Ordering::SeqCst),
"a delivery abandoned during its push never reached the accept point"
);
assert_eq!(
registry.in_flight_for_worker(worker_id)?,
0,
"a delivery abandoned during its push must return the tracked slot; a slot \
still held means the guard was armed AFTER the push and the cancellation \
window is open again"
);
assert!(
!tracker.complete_task(worker_id, &workflow_id, &activity_id, ®istry)?,
"the guard retired the entry itself; the caller's later untrack must find \
nothing left to retire"
);
open_gate.send(())?;
occupant.await??;
Ok::<(), Box<dyn std::error::Error>>(())
})
}
#[test]
fn identity_resolution_refuses_a_task_with_no_run() -> Result<(), Box<dyn std::error::Error>> {
let Err(refusal) = Resolved::from_task(&task_without_a_run()) else {
return Err("a task with no run must not resolve".into());
};
assert!(
refusal.contains("run id is missing"),
"the refusal must name the run so an operator is sent to the right remedy: {refusal}"
);
Ok(())
}
}