use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, OnceLock};
use aion_core::{ActivityId, RunId, WorkerAttribution, WorkflowId};
use async_trait::async_trait;
use super::envelope::{CompletionFences, CompletionToken};
use super::registry::WorkerHandle;
use super::task_delivery::DeliveryAccepted;
use crate::error::ServerError;
use crate::observability::Metrics;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct LeaseKey {
pub workflow_id: WorkflowId,
pub run_id: RunId,
pub activity_id: ActivityId,
pub attempt: u32,
}
impl std::fmt::Display for LeaseKey {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
formatter,
"workflow {} run {} activity {} attempt {}",
self.workflow_id, self.run_id, self.activity_id, self.attempt
)
}
}
#[async_trait]
pub trait ActivityLeaseRecorder: Send + Sync + 'static {
async fn record(&self, key: &LeaseKey, worker: WorkerAttribution) -> Result<(), String>;
}
pub struct EngineLeaseRecorder {
engine: Arc<aion::Engine>,
}
impl EngineLeaseRecorder {
#[must_use]
pub fn new(engine: Arc<aion::Engine>) -> Self {
Self { engine }
}
}
impl std::fmt::Debug for EngineLeaseRecorder {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.debug_struct("EngineLeaseRecorder").finish()
}
}
#[async_trait]
impl ActivityLeaseRecorder for EngineLeaseRecorder {
async fn record(&self, key: &LeaseKey, worker: WorkerAttribution) -> Result<(), String> {
self.engine
.record_activity_lease(
&key.workflow_id,
&key.run_id,
key.activity_id.clone(),
key.attempt,
worker,
)
.await
.map_err(|error| error.to_string())
}
}
#[derive(Clone, Debug, Default)]
pub struct LeaseRecordLedger {
failures: Arc<AtomicU64>,
}
impl LeaseRecordLedger {
#[must_use]
pub fn failures(&self) -> u64 {
self.failures.load(Ordering::SeqCst)
}
fn increment(&self) {
self.failures.fetch_add(1, Ordering::SeqCst);
}
}
struct Installed {
recorder: Arc<dyn ActivityLeaseRecorder>,
metrics: Option<Metrics>,
}
#[derive(Clone, Default)]
pub struct LeaseRecorderSeam {
installed: Arc<OnceLock<Installed>>,
ledger: LeaseRecordLedger,
}
impl std::fmt::Debug for LeaseRecorderSeam {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("LeaseRecorderSeam")
.field("installed", &self.installed.get().is_some())
.field("failures", &self.ledger.failures())
.finish()
}
}
impl LeaseRecorderSeam {
pub fn install(
&self,
recorder: Arc<dyn ActivityLeaseRecorder>,
metrics: Option<Metrics>,
) -> bool {
if self.installed.set(Installed { recorder, metrics }).is_err() {
tracing::warn!("activity lease recorder already installed; ignoring duplicate install");
return false;
}
true
}
#[must_use]
pub fn is_installed(&self) -> bool {
self.installed.get().is_some()
}
#[must_use]
pub const fn ledger(&self) -> &LeaseRecordLedger {
&self.ledger
}
pub async fn record(&self, key: &LeaseKey, worker: WorkerAttribution) {
let Some(installed) = self.installed.get() else {
self.lost(
key,
&worker,
None,
"no activity lease recorder is installed on this dispatcher",
);
return;
};
if let Err(reason) = installed.recorder.record(key, worker.clone()).await {
self.lost(key, &worker, installed.metrics.as_ref(), &reason);
}
}
pub fn record_blocking(
&self,
handle: Option<&tokio::runtime::Handle>,
key: &LeaseKey,
worker: WorkerAttribution,
) {
match tokio::runtime::Handle::try_current() {
Ok(ambient) => match ambient.runtime_flavor() {
tokio::runtime::RuntimeFlavor::MultiThread => {
let runner = handle.cloned().unwrap_or(ambient);
tokio::task::block_in_place(|| runner.block_on(self.record(key, worker)));
}
flavor => self.lost(
key,
&worker,
self.installed.get().and_then(|i| i.metrics.as_ref()),
&format!(
"the lease cannot be recorded from inside a {flavor:?} tokio runtime: \
blocking here would starve the recorder of its only thread"
),
),
},
Err(_) => match handle {
Some(runner) => runner.block_on(self.record(key, worker)),
None => self.lost(
key,
&worker,
self.installed.get().and_then(|i| i.metrics.as_ref()),
"no tokio runtime is reachable from this thread to record the lease on",
),
},
}
}
fn lost(
&self,
key: &LeaseKey,
worker: &WorkerAttribution,
metrics: Option<&Metrics>,
why: &str,
) {
self.ledger.increment();
if let Some(metrics) = metrics {
metrics.activity_lease_record_failed();
}
tracing::error!(
workflow_id = %key.workflow_id,
run_id = %key.run_id,
activity_id = %key.activity_id,
attempt = key.attempt,
worker_identity = %worker.identity,
task_queue = %worker.task_queue,
transport = worker.transport.name(),
lease_record_failures_total = self.ledger.failures(),
reason = why,
"activity lease was not recorded; the worker holds the attempt and the history reads \
it as unattributed"
);
}
}
#[must_use]
pub fn attribution_for(worker: &WorkerHandle) -> WorkerAttribution {
WorkerAttribution {
identity: worker.identity().to_owned(),
task_queue: worker.task_queue().to_owned(),
node: worker.node().map(str::to_owned),
deployment: worker
.instance()
.map(|instance| instance.deployment.clone()),
instance_id: worker
.instance()
.map(|instance| instance.instance_id.clone()),
transport: worker.delivery().transport(),
}
}
pub struct LeaseHandoff {
seam: LeaseRecorderSeam,
key: LeaseKey,
worker: WorkerAttribution,
fences: CompletionFences,
token: CompletionToken,
lease: aion::LeaseSignal,
settled: AtomicBool,
}
impl LeaseHandoff {
#[must_use]
pub fn key(&self) -> &LeaseKey {
&self.key
}
#[must_use]
pub fn token(&self) -> &CompletionToken {
&self.token
}
}
impl std::fmt::Debug for LeaseHandoff {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("LeaseHandoff")
.field("key", &self.key)
.field("settled", &self.settled.load(Ordering::SeqCst))
.finish_non_exhaustive()
}
}
impl LeaseHandoff {
pub fn arm(
seam: LeaseRecorderSeam,
key: LeaseKey,
worker: WorkerAttribution,
fences: CompletionFences,
token: CompletionToken,
lease: aion::LeaseSignal,
) -> Result<Self, ServerError> {
fences.arm_lease(&token)?;
Ok(Self {
seam,
key,
worker,
fences,
token,
lease,
settled: AtomicBool::new(false),
})
}
pub fn accepted_blocking(&self, handle: Option<&tokio::runtime::Handle>) {
if self.settled.load(Ordering::SeqCst) {
tracing::warn!(key = %self.key, "lease handoff accepted twice; the second is ignored");
return;
}
self.seam
.record_blocking(handle, &self.key, self.worker.clone());
self.lease.fire();
self.settle();
}
fn settle(&self) {
if self.settled.swap(true, Ordering::SeqCst) {
return;
}
if let Err(error) = self.fences.settle_lease(&self.token) {
tracing::error!(
key = %self.key,
%error,
"lease gate could not be settled; completions for this attempt may wait on a \
poisoned fence"
);
}
}
}
impl Drop for LeaseHandoff {
fn drop(&mut self) {
self.settle();
}
}
#[async_trait]
impl DeliveryAccepted for LeaseHandoff {
async fn accepted(&self) {
if self.settled.load(Ordering::SeqCst) {
tracing::warn!(key = %self.key, "lease handoff accepted twice; the second is ignored");
return;
}
self.seam.record(&self.key, self.worker.clone()).await;
self.lease.fire();
self.settle();
}
}
#[cfg(test)]
mod tests {
use super::*;
use aion_core::WorkerTransport;
use std::sync::Mutex;
use uuid::Uuid;
struct Recording(Mutex<Vec<(LeaseKey, WorkerAttribution)>>);
#[async_trait]
impl ActivityLeaseRecorder for Recording {
async fn record(&self, key: &LeaseKey, worker: WorkerAttribution) -> Result<(), String> {
self.0
.lock()
.map_err(|_| "recording lock poisoned".to_owned())?
.push((key.clone(), worker));
Ok(())
}
}
struct Faulting;
#[async_trait]
impl ActivityLeaseRecorder for Faulting {
async fn record(&self, _key: &LeaseKey, _worker: WorkerAttribution) -> Result<(), String> {
Err("store fault injected".to_owned())
}
}
fn key() -> LeaseKey {
LeaseKey {
workflow_id: WorkflowId::new(Uuid::new_v4()),
run_id: RunId::new(Uuid::new_v4()),
activity_id: ActivityId::from_sequence_position(3),
attempt: 1,
}
}
fn worker() -> WorkerAttribution {
WorkerAttribution {
identity: "w-1".to_owned(),
task_queue: "q".to_owned(),
node: None,
deployment: None,
instance_id: None,
transport: WorkerTransport::Grpc,
}
}
#[tokio::test]
async fn an_installed_recorder_receives_the_lease_and_the_ledger_stays_at_zero() {
let seam = LeaseRecorderSeam::default();
let recording = Arc::new(Recording(Mutex::new(Vec::new())));
assert!(seam.install(recording.clone(), None));
let key = key();
seam.record(&key, worker()).await;
let seen = recording
.0
.lock()
.map(|seen| seen.clone())
.unwrap_or_default();
assert_eq!(seen.len(), 1);
assert_eq!(seen[0].0, key);
assert_eq!(seen[0].1.identity, "w-1");
assert_eq!(seam.ledger().failures(), 0);
}
#[tokio::test]
async fn a_failed_record_counts_on_the_ledger_and_on_the_metrics_surface() {
let seam = LeaseRecorderSeam::default();
let metrics = Metrics::new().ok();
assert!(seam.install(Arc::new(Faulting), metrics.clone()));
seam.record(&key(), worker()).await;
seam.record(&key(), worker()).await;
assert_eq!(seam.ledger().failures(), 2);
if let Some(metrics) = metrics {
let text = String::from_utf8(metrics.encode().unwrap_or_default()).unwrap_or_default();
assert!(
text.contains("aion_activity_lease_record_failures_total 2"),
"the metrics surface must carry the same count; got:\n{text}"
);
}
}
#[tokio::test]
async fn an_uninstalled_seam_counts_the_loss_instead_of_skipping_it() {
let seam = LeaseRecorderSeam::default();
assert!(!seam.is_installed());
seam.record(&key(), worker()).await;
assert_eq!(seam.ledger().failures(), 1);
}
#[tokio::test]
async fn a_second_install_keeps_the_first() {
let seam = LeaseRecorderSeam::default();
let recording = Arc::new(Recording(Mutex::new(Vec::new())));
assert!(seam.install(recording.clone(), None));
assert!(!seam.install(Arc::new(Faulting), None));
seam.record(&key(), worker()).await;
assert_eq!(seam.ledger().failures(), 0);
assert_eq!(recording.0.lock().map_or(0, |s| s.len()), 1);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn a_handoff_arms_the_gate_and_settles_it_after_recording() -> Result<(), ServerError> {
let seam = LeaseRecorderSeam::default();
let recording = Arc::new(Recording(Mutex::new(Vec::new())));
seam.install(recording.clone(), None);
let fences = CompletionFences::default();
let k = key();
let token = fences.issue(&k.workflow_id, &k.run_id, &k.activity_id, k.attempt)?;
let handoff = LeaseHandoff::arm(
seam.clone(),
k,
worker(),
fences.clone(),
token.clone(),
aion::LeaseSignal::none(),
)?;
assert!(fences.lease_pending(&token)?);
handoff.accepted().await;
assert!(!fences.lease_pending(&token)?);
assert_eq!(recording.0.lock().map_or(0, |s| s.len()), 1);
Ok(())
}
#[tokio::test]
async fn dropping_an_unaccepted_handoff_settles_without_recording() -> Result<(), ServerError> {
let seam = LeaseRecorderSeam::default();
let recording = Arc::new(Recording(Mutex::new(Vec::new())));
seam.install(recording.clone(), None);
let fences = CompletionFences::default();
let k = key();
let token = fences.issue(&k.workflow_id, &k.run_id, &k.activity_id, k.attempt)?;
let handoff = LeaseHandoff::arm(
seam,
k,
worker(),
fences.clone(),
token.clone(),
aion::LeaseSignal::none(),
)?;
assert!(fences.lease_pending(&token)?);
drop(handoff);
assert!(!fences.lease_pending(&token)?);
assert_eq!(recording.0.lock().map_or(0, |s| s.len()), 0);
Ok(())
}
#[tokio::test]
async fn recording_from_a_current_thread_runtime_is_a_named_loss_not_a_deadlock() {
let seam = LeaseRecorderSeam::default();
seam.install(Arc::new(Recording(Mutex::new(Vec::new()))), None);
seam.record_blocking(None, &key(), worker());
assert_eq!(seam.ledger().failures(), 1);
}
#[test]
fn recording_from_a_plain_thread_runs_on_the_given_handle() {
let runtime = tokio::runtime::Builder::new_multi_thread()
.worker_threads(1)
.enable_all()
.build();
let Ok(runtime) = runtime else {
return;
};
let seam = LeaseRecorderSeam::default();
let recording = Arc::new(Recording(Mutex::new(Vec::new())));
seam.install(recording.clone(), None);
seam.record_blocking(Some(runtime.handle()), &key(), worker());
assert_eq!(seam.ledger().failures(), 0);
assert_eq!(recording.0.lock().map_or(0, |s| s.len()), 1);
}
}