use std::collections::HashMap;
use std::time::Duration;
use taquba::{JobRecord, JobStatus, Queue, SettlementEffects, WorkerError};
use tracing::{debug, warn};
use crate::error::{Result, worker_error};
use crate::keys::{
HEADER_SIGNAL_DELIVERED, HEADER_SIGNAL_WAIT, RunId, signal_buf_kv_key, signal_delivered_kv_key,
signal_wait_kv_key,
};
use crate::runner::{StepError, StepRunner};
use crate::runtime::{RuntimeCore, RuntimeInner, StepEnqueueOpts, WorkflowRuntime};
use crate::terminal::TerminalHook;
use crate::worker::ClaimedStep;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SignalOutcome {
Delivered,
Buffered,
}
const SIGNAL_WAIT_READ_ATTEMPTS: u32 = 10;
const SIGNAL_WAIT_READ_INTERVAL: Duration = Duration::from_millis(25);
async fn remove_entry(queue: &Queue, key: &[u8], expected: &[u8]) {
if let Err(err) = queue.kv_compare_delete(key, expected).await {
debug!(key = %String::from_utf8_lossy(key), "signal entry removal failed: {err}");
}
}
impl<R: StepRunner, H: TerminalHook> WorkflowRuntime<R, H> {
pub async fn signal(&self, correlation_key: &str, payload: Vec<u8>) -> Result<SignalOutcome> {
let queue = &self.inner.core.queue;
let buf_key = signal_buf_kv_key(correlation_key);
queue.kv_put(&buf_key, &payload).await?;
let wait_key = signal_wait_kv_key(correlation_key);
for attempt in 0..SIGNAL_WAIT_READ_ATTEMPTS {
if attempt > 0 {
tokio::time::sleep(SIGNAL_WAIT_READ_INTERVAL).await;
}
let Some(waiter) = queue.view().kv_get(&wait_key).await? else {
continue;
};
let Ok(job_id) = std::str::from_utf8(&waiter).map(str::to_string) else {
remove_entry(queue, &wait_key, &waiter).await;
return Ok(SignalOutcome::Buffered);
};
return match queue.wake_scheduled(&job_id, Some(payload.clone())).await? {
taquba::WakeOutcome::Woken => {
remove_entry(queue, &buf_key, &payload).await;
remove_entry(queue, &wait_key, &waiter).await;
Ok(SignalOutcome::Delivered)
}
taquba::WakeOutcome::NotScheduled | taquba::WakeOutcome::NotFound => {
remove_entry(queue, &wait_key, &waiter).await;
Ok(SignalOutcome::Buffered)
}
};
}
Ok(SignalOutcome::Buffered)
}
pub async fn clear_signal(&self, correlation_key: &str) -> Result<bool> {
let queue = &self.inner.core.queue;
let buf_key = signal_buf_kv_key(correlation_key);
loop {
let Some(current) = queue.view().kv_get(&buf_key).await? else {
return Ok(false);
};
if queue.kv_compare_delete(&buf_key, ¤t).await? {
return Ok(true);
}
}
}
}
impl<R: StepRunner, H: TerminalHook> RuntimeInner<R, H> {
pub(crate) async fn advance_on_signal(
&self,
claimed: &ClaimedStep<'_>,
payload: Vec<u8>,
correlation_key: &str,
timeout: Duration,
input_hash: [u8; 32],
) -> std::result::Result<SettlementEffects, WorkerError> {
let wait_key = signal_wait_kv_key(correlation_key);
if let Some(existing) = self
.core
.queue
.view()
.kv_get(&wait_key)
.await
.map_err(worker_error)?
&& let Ok(existing_id) = std::str::from_utf8(&existing)
&& let Ok(Some(job)) = self.core.queue.view().job_record(existing_id).await
&& job.status == JobStatus::Scheduled
{
let message =
format!("a waiter is already registered for correlation key `{correlation_key}`");
return Err(self
.terminating_failure(claimed, StepError::permanent(message), input_hash)
.await);
}
let buf_key = signal_buf_kv_key(correlation_key);
match self
.core
.queue
.view()
.kv_get(&buf_key)
.await
.map_err(worker_error)?
{
Some(buffered) => {
let opts = StepEnqueueOpts {
reserved_headers: claimed
.reserved_headers_with((HEADER_SIGNAL_DELIVERED, "1".to_string())),
..claimed.next_step_opts()
};
let delivered_key =
signal_delivered_kv_key(&claimed.run_id, claimed.step_number + 1);
let buffered = buffered.to_vec();
let mut effects = self
.core
.advance_with_kv(claimed, payload, opts, |_| {
HashMap::from([(delivered_key, buffered)])
})
.await;
effects.kv_deletes.push(buf_key);
Ok(effects)
}
None => {
let opts = StepEnqueueOpts {
run_at: Some(self.core.run_at_after(timeout)),
reserved_headers: claimed
.reserved_headers_with((HEADER_SIGNAL_WAIT, correlation_key.to_string())),
..claimed.next_step_opts()
};
let effects = self
.core
.advance_with_kv(claimed, payload, opts, |job_id| {
HashMap::from([(wait_key, job_id.as_bytes().to_vec())])
})
.await;
Ok(effects)
}
}
}
}
impl RuntimeCore {
pub(crate) async fn resolve_step_signal(
&self,
job: &JobRecord,
run_id: &RunId,
step_number: u32,
) -> Result<(Option<Vec<u8>>, Vec<Vec<u8>>)> {
if job.headers.contains_key(HEADER_SIGNAL_DELIVERED) {
let delivered_key = signal_delivered_kv_key(run_id, step_number);
let payload = self
.queue
.view()
.kv_get(&delivered_key)
.await?
.map(|b| b.to_vec());
if payload.is_none() {
warn!(run_id = %run_id, step_number, "delivered signal record is missing");
}
return Ok((payload, vec![delivered_key]));
}
let Some(correlation_key) = job.headers.get(HEADER_SIGNAL_WAIT) else {
return Ok((None, Vec::new()));
};
let wait_key = signal_wait_kv_key(correlation_key);
remove_entry(&self.queue, &wait_key, job.id.as_bytes()).await;
if job.woken_at.is_some() {
return Ok((job.wake_payload.clone(), Vec::new()));
}
let delivered_key = signal_delivered_kv_key(run_id, step_number);
if let Some(prior) = self.queue.view().kv_get(&delivered_key).await? {
return Ok((Some(prior.to_vec()), vec![delivered_key]));
}
let buf_key = signal_buf_kv_key(correlation_key);
if let Some(buffered) = self.queue.view().kv_get(&buf_key).await? {
let buffered = buffered.to_vec();
self.queue.kv_put(&delivered_key, &buffered).await?;
remove_entry(&self.queue, &buf_key, &buffered).await;
return Ok((Some(buffered), vec![delivered_key]));
}
Ok((None, Vec::new()))
}
}