use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use tracing::{debug, warn};
use uuid::Uuid;
use super::TaskSpec;
use super::ledger::{TaskKey, TaskLedger, TaskRecord, TaskState, unix_now};
use super::protocol::ExecuteRequest;
use super::scheduler::{ComputeScheduler, Delegated, ScheduleError};
#[derive(Debug, Clone)]
pub struct TriggerRule {
pub id: String,
pub store_address: String,
pub spec: TaskSpec,
}
#[async_trait]
pub trait TaskDispatcher: Send + Sync + 'static {
async fn dispatch(&self, request: ExecuteRequest) -> Result<Delegated, DispatchError>;
}
#[derive(Debug, Clone, thiserror::Error)]
#[error("{message}")]
pub struct DispatchError {
pub permanent: bool,
pub message: String,
}
#[async_trait]
impl TaskDispatcher for ComputeScheduler {
async fn dispatch(&self, request: ExecuteRequest) -> Result<Delegated, DispatchError> {
self.execute(request).await.map_err(|e| DispatchError {
permanent: matches!(e, ScheduleError::TaskFailed { .. }),
message: e.to_string(),
})
}
}
#[derive(Debug, Clone)]
pub struct TriggerConfig {
pub max_attempts: u32,
pub task_deadline: Duration,
}
impl Default for TriggerConfig {
fn default() -> Self {
Self {
max_attempts: 3,
task_deadline: Duration::from_secs(300),
}
}
}
pub struct TriggerEngine {
rules: parking_lot::RwLock<Vec<TriggerRule>>,
ledger: TaskLedger,
dispatcher: Arc<dyn TaskDispatcher>,
config: TriggerConfig,
}
impl std::fmt::Debug for TriggerEngine {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TriggerEngine")
.field("rules", &self.rules.read().len())
.finish_non_exhaustive()
}
}
impl TriggerEngine {
pub fn new(
ledger: TaskLedger,
dispatcher: Arc<dyn TaskDispatcher>,
config: TriggerConfig,
) -> Self {
Self {
rules: parking_lot::RwLock::new(Vec::new()),
ledger,
dispatcher,
config,
}
}
pub fn on_replicated(&self, store_address: impl Into<String>, spec: TaskSpec) -> String {
let store_address = store_address.into();
let id = format!("rule-{}", Uuid::new_v4());
self.rules.write().push(TriggerRule {
id: id.clone(),
store_address,
spec,
});
id
}
pub fn add_rule(&self, rule: TriggerRule) {
self.rules.write().push(rule);
}
pub fn remove_rule(&self, id: &str) {
self.rules.write().retain(|r| r.id != id);
}
pub fn ledger(&self) -> &TaskLedger {
&self.ledger
}
pub async fn notify_replicated(
self: &Arc<Self>,
store_address: &str,
event_id: &[u8],
payload: &[u8],
) -> Vec<TaskKey> {
let matching: Vec<TriggerRule> = self
.rules
.read()
.iter()
.filter(|rule| store_address.starts_with(&rule.store_address))
.cloned()
.collect();
let mut claimed = Vec::new();
for rule in matching {
let key = dedup_key(&rule.id, event_id);
let record = TaskRecord::pending(
key.clone(),
rule.spec.wasm_hash,
rule.spec.entrypoint.clone(),
rule.spec.class,
rule.spec.limits,
payload.to_vec(),
rule.spec.required_model.clone(),
);
match self.ledger.claim(&record).await {
Ok(true) => {
debug!(rule = %rule.id, task = %key, store = store_address,
"compute trigger: task claimed");
let engine = self.clone();
let task_key = key.clone();
tokio::spawn(async move {
engine.run_task(&task_key).await;
});
claimed.push(key);
}
Ok(false) => {
debug!(rule = %rule.id, task = %key,
"compute trigger: already claimed elsewhere, skipping");
}
Err(e) => warn!(rule = %rule.id, "compute trigger: ledger claim failed: {e}"),
}
}
claimed
}
async fn run_task(&self, key: &str) {
let deadline = unix_now() + self.config.task_deadline.as_secs();
let record = match self.ledger.claim_for_dispatch(key, deadline).await {
Ok(Some(record)) => record,
Ok(None) => return, Err(e) => {
warn!(task = %key, "compute trigger: claim failed: {e}");
return;
}
};
let request = ExecuteRequest {
task_id: Uuid::new_v4(),
wasm_hash: record.wasm_hash,
entrypoint: record.entrypoint.clone(),
class: record.class,
limits: record.limits,
input: record.input.clone(),
required_model: record.required_model.clone(),
};
match self.dispatcher.dispatch(request).await {
Ok(delegated) => {
debug!(task = %key, executor = %delegated.executor.fmt_short(),
"compute trigger: task done");
let _ = self
.ledger
.update(key, |r| {
r.state = TaskState::Done {
executor: delegated.executor,
};
r.result = Some(delegated.completed.clone());
r.input = Vec::new();
})
.await;
}
Err(e) => {
let spent = record.attempts >= self.config.max_attempts;
let final_failure = e.permanent || spent;
warn!(task = %key, attempts = record.attempts, permanent = e.permanent,
"compute trigger: dispatch failed: {}", e.message);
let _ = self
.ledger
.update(key, |r| {
if final_failure {
r.state = TaskState::Failed { error: e.message };
r.input = Vec::new(); } else {
r.state = TaskState::Pending;
}
})
.await;
}
}
}
pub async fn requeue_due(self: &Arc<Self>) -> usize {
let due = match self.ledger.needing_dispatch().await {
Ok(due) => due,
Err(e) => {
warn!("compute trigger: requeue scan failed: {e}");
return 0;
}
};
let mut dispatched = 0;
for record in due {
if record.attempts >= self.config.max_attempts {
let _ = self
.ledger
.update(&record.key, |r| {
r.state = TaskState::Failed {
error: "retry budget exhausted".into(),
};
r.input = Vec::new(); })
.await;
continue;
}
let engine = self.clone();
let key = record.key.clone();
tokio::spawn(async move {
engine.run_task(&key).await;
});
dispatched += 1;
}
dispatched
}
pub fn spawn_requeue_loop(self: &Arc<Self>, every: Duration) -> tokio::task::JoinHandle<()> {
let engine = self.clone();
tokio::spawn(async move {
let mut ticker = tokio::time::interval(every);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
ticker.tick().await;
engine.requeue_due().await;
}
})
}
pub fn attach_event_bus(
self: &Arc<Self>,
event_bus: Arc<crate::p2p::EventBus>,
) -> tokio::task::JoinHandle<()> {
let engine = self.clone();
tokio::spawn(async move {
let mut receiver = match event_bus
.subscribe::<crate::stores::events::EventReplicated>()
.await
{
Ok(receiver) => receiver,
Err(e) => {
warn!("compute trigger: EventReplicated subscription failed: {e}");
return;
}
};
while let Ok(event) = receiver.recv().await {
let address = event.address.to_string();
for entry in &event.entries {
engine
.notify_replicated(&address, entry.hash().as_bytes(), entry.payload())
.await;
}
}
})
}
}
fn dedup_key(rule_id: &str, event_id: &[u8]) -> TaskKey {
let mut hasher = blake3::Hasher::new();
hasher.update(rule_id.as_bytes());
hasher.update(&[0x1f]);
hasher.update(event_id);
hex::encode(hasher.finalize().as_bytes())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::compute::{CompletedTask, ExecMetrics, Placement, ResourceLimits, TaskClass};
use iroh_blobs::Hash;
use std::sync::atomic::{AtomicU32, Ordering};
fn spec() -> TaskSpec {
TaskSpec {
wasm_hash: Hash::new(b"thumbnailer"),
entrypoint: "generate_thumbnail".into(),
class: TaskClass::Media,
limits: ResourceLimits::default(),
placement: Placement::BestAvailable,
required_model: None,
}
}
struct CountingDispatcher {
calls: AtomicU32,
}
#[async_trait]
impl TaskDispatcher for CountingDispatcher {
async fn dispatch(&self, request: ExecuteRequest) -> Result<Delegated, DispatchError> {
self.calls.fetch_add(1, Ordering::SeqCst);
Ok(Delegated {
executor: iroh::SecretKey::from_bytes(&[7u8; 32]).public(),
completed: CompletedTask {
output: request.input, metrics: ExecMetrics {
fuel_consumed: 1,
duration_ms: 1,
peak_memory_bytes: 65_536,
},
},
})
}
}
struct FailingDispatcher {
permanent: bool,
calls: AtomicU32,
}
#[async_trait]
impl TaskDispatcher for FailingDispatcher {
async fn dispatch(&self, _request: ExecuteRequest) -> Result<Delegated, DispatchError> {
self.calls.fetch_add(1, Ordering::SeqCst);
Err(DispatchError {
permanent: self.permanent,
message: "boom".into(),
})
}
}
async fn wait_for_state<F>(engine: &Arc<TriggerEngine>, key: &str, matches: F)
where
F: Fn(&TaskState) -> bool,
{
let deadline = tokio::time::Instant::now() + Duration::from_secs(10);
loop {
if let Ok(Some(record)) = engine.ledger().get(key).await
&& matches(&record.state)
{
return;
}
assert!(
tokio::time::Instant::now() < deadline,
"task {key} never reached the expected state"
);
tokio::time::sleep(Duration::from_millis(20)).await;
}
}
#[tokio::test]
async fn event_fires_rule_once_and_completes() {
let dispatcher = Arc::new(CountingDispatcher {
calls: AtomicU32::new(0),
});
let engine = Arc::new(TriggerEngine::new(
TaskLedger::in_memory(),
dispatcher.clone(),
TriggerConfig::default(),
));
engine.add_rule(TriggerRule {
id: "thumbnails".into(),
store_address: "/fotos".into(),
spec: spec(),
});
let claimed = engine
.notify_replicated("/fotos/album1", b"entry-hash-1", b"jpeg bytes")
.await;
assert_eq!(claimed.len(), 1);
let key = claimed[0].clone();
let again = engine
.notify_replicated("/fotos/album1", b"entry-hash-1", b"jpeg bytes")
.await;
assert!(again.is_empty(), "duplicate event must be deduped");
wait_for_state(&engine, &key, |s| matches!(s, TaskState::Done { .. })).await;
assert_eq!(dispatcher.calls.load(Ordering::SeqCst), 1);
let record = engine.ledger().get(&key).await.unwrap().unwrap();
assert!(
record.input.is_empty(),
"a terminal record drops its input so the requeue scan stays cheap"
);
assert_eq!(record.result.unwrap().output, b"jpeg bytes");
}
#[tokio::test]
async fn unmatched_address_fires_nothing() {
let engine = Arc::new(TriggerEngine::new(
TaskLedger::in_memory(),
Arc::new(CountingDispatcher {
calls: AtomicU32::new(0),
}),
TriggerConfig::default(),
));
engine.add_rule(TriggerRule {
id: "thumbnails".into(),
store_address: "/fotos".into(),
spec: spec(),
});
let claimed = engine.notify_replicated("/documentos", b"e1", b"pdf").await;
assert!(claimed.is_empty());
}
#[tokio::test]
async fn permanent_failure_is_terminal() {
let dispatcher = Arc::new(FailingDispatcher {
permanent: true,
calls: AtomicU32::new(0),
});
let engine = Arc::new(TriggerEngine::new(
TaskLedger::in_memory(),
dispatcher.clone(),
TriggerConfig::default(),
));
engine.add_rule(TriggerRule {
id: "r".into(),
store_address: "/fotos".into(),
spec: spec(),
});
let key = engine
.notify_replicated("/fotos", b"e1", b"x")
.await
.remove(0);
wait_for_state(&engine, &key, |s| matches!(s, TaskState::Failed { .. })).await;
assert_eq!(engine.requeue_due().await, 0);
assert_eq!(dispatcher.calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn transient_failure_retries_until_budget_then_fails() {
let dispatcher = Arc::new(FailingDispatcher {
permanent: false,
calls: AtomicU32::new(0),
});
let engine = Arc::new(TriggerEngine::new(
TaskLedger::in_memory(),
dispatcher.clone(),
TriggerConfig {
max_attempts: 2,
..TriggerConfig::default()
},
));
engine.add_rule(TriggerRule {
id: "r".into(),
store_address: "/fotos".into(),
spec: spec(),
});
let key = engine
.notify_replicated("/fotos", b"e1", b"x")
.await
.remove(0);
wait_for_state(&engine, &key, |s| matches!(s, TaskState::Pending)).await;
assert_eq!(engine.requeue_due().await, 1);
wait_for_state(&engine, &key, |s| matches!(s, TaskState::Failed { .. })).await;
assert_eq!(engine.requeue_due().await, 0, "failed tasks stay failed");
assert_eq!(dispatcher.calls.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn abandoned_running_task_is_requeued() {
let dispatcher = Arc::new(CountingDispatcher {
calls: AtomicU32::new(0),
});
let engine = Arc::new(TriggerEngine::new(
TaskLedger::in_memory(),
dispatcher.clone(),
TriggerConfig::default(),
));
let mut record = TaskRecord::pending(
"abandoned".into(),
Hash::new(b"wasm"),
"gdb_run".into(),
TaskClass::Media,
ResourceLimits::default(),
b"input".to_vec(),
None,
);
record.attempts = 1;
record.state = TaskState::Running {
deadline_unix: unix_now() - 10,
};
engine.ledger().claim(&record).await.unwrap();
assert_eq!(engine.requeue_due().await, 1);
wait_for_state(&engine, "abandoned", |s| {
matches!(s, TaskState::Done { .. })
})
.await;
assert_eq!(dispatcher.calls.load(Ordering::SeqCst), 1);
}
}