Skip to main content

aion_server/worker/
dispatch.rs

1//! Push dispatch for remote activity workers and result handoff to the engine contract.
2
3use std::collections::BTreeMap;
4
5use aion_core::{ActivityError, ActivityErrorKind, ActivityId, Payload, RunId, WorkflowId};
6use aion_proto::{
7    ProtoActivityId, ProtoActivityResult, ProtoActivityTask, ProtoPayload, ProtoRunId,
8    ProtoWorkflowId, WireError, proto_activity_result,
9};
10
11use crate::error::ServerError;
12use crate::shutdown::DrainState;
13use crate::worker::registry::{ConnectedWorkerRegistry, WorkerMessage};
14use tracing::{Instrument, info_span};
15
16/// Scheduled remote activity that must be placed with a connected worker.
17#[derive(Clone, Debug, Eq, PartialEq)]
18pub struct ScheduledActivity {
19    /// Namespace selected by the adapter boundary before dispatch — the
20    /// correctness/isolation boundary the activity may dispatch within.
21    pub namespace: String,
22    /// Task queue (pool/flavour) selected within the namespace. The worker-pool
23    /// address is `(namespace, task_queue)`; an empty value is normalized to the
24    /// named default pool by the registry lookup.
25    pub task_queue: String,
26    /// Activity type to match against worker registrations, *within* the
27    /// selected pool.
28    pub activity_type: String,
29    /// Optional node locality affinity. `Some(node)` pins this dispatch to
30    /// workers advertising that node (require semantics: it waits if none are
31    /// present, exactly like the no-worker path); `None` is unpinned and reaches
32    /// any worker in the `(namespace, task_queue)` pool — byte-identical to the
33    /// pre-NODE behaviour. Producers stamp `None` until SDK selection (NODE-4)
34    /// and the durable column (NODE-2) land.
35    pub node: Option<String>,
36    /// Owning workflow id.
37    pub workflow_id: WorkflowId,
38    /// Correlating activity id.
39    pub activity_id: ActivityId,
40    /// Concrete workflow run that staged this task, when known.
41    pub run_id: Option<RunId>,
42    /// Opaque activity input payload.
43    pub input: Payload,
44    /// One-based delivery attempt stamped by the dispatching engine seam.
45    /// Zero is malformed on the wire; producers must always stamp it.
46    pub attempt: u32,
47    /// Display labels the workflow attached to the activity. Display metadata
48    /// only — carried to the worker for its logs and the dashboard.
49    pub labels: BTreeMap<String, String>,
50}
51
52impl ScheduledActivity {
53    /// Build the wire task pushed to the worker stream.
54    #[must_use]
55    pub fn to_task(&self) -> ProtoActivityTask {
56        ProtoActivityTask {
57            workflow_id: Some(ProtoWorkflowId::from(self.workflow_id.clone())),
58            activity_id: Some(ProtoActivityId::from(self.activity_id.clone())),
59            activity_type: self.activity_type.clone(),
60            input: Some(ProtoPayload::from(self.input.clone())),
61            attempt: self.attempt,
62            labels: self.labels.clone().into_iter().collect(),
63            run_id: self.run_id.clone().map(ProtoRunId::from),
64        }
65    }
66}
67
68/// Push dispatcher backed by the connected-worker registry.
69#[derive(Clone, Debug)]
70pub struct ActivityDispatcher {
71    registry: ConnectedWorkerRegistry,
72    drain_state: DrainState,
73}
74
75impl ActivityDispatcher {
76    /// Build a dispatcher over the shared worker registry.
77    #[must_use]
78    pub fn new(registry: ConnectedWorkerRegistry) -> Self {
79        Self {
80            registry,
81            drain_state: DrainState::default(),
82        }
83    }
84
85    /// Share the server drain gate.
86    #[must_use]
87    pub fn with_drain_state(mut self, drain_state: DrainState) -> Self {
88        self.drain_state = drain_state;
89        self
90    }
91
92    /// Push a scheduled activity to a matching worker.
93    ///
94    /// # Errors
95    ///
96    /// Returns a typed dispatch error if no worker is available or the selected
97    /// stream is closed; returns lock poison if registry access cannot be trusted.
98    pub async fn dispatch(&self, activity: &ScheduledActivity) -> Result<(), ServerError> {
99        let span = info_span!(
100            "activity_dispatch",
101            operation = "activity_dispatch",
102            namespace = %activity.namespace,
103            task_queue = %activity.task_queue,
104            node = activity.node.as_deref(),
105            workflow_id = %activity.workflow_id,
106            activity_id = %activity.activity_id,
107            activity_type = %activity.activity_type,
108            worker_id = tracing::field::Empty,
109        );
110        let span_fields = span.clone();
111
112        async {
113            let workers = loop {
114                self.drain_state
115                    .ensure_accepting(&activity.namespace, &activity.activity_type)?;
116                let candidates = self.registry.workers_for(
117                    &activity.namespace,
118                    &activity.task_queue,
119                    &activity.activity_type,
120                    activity.node.as_deref(),
121                )?;
122                if !candidates.is_empty() {
123                    break candidates;
124                }
125                tracing::info!(
126                    namespace = %activity.namespace,
127                    task_queue = %activity.task_queue,
128                    node = activity.node.as_deref(),
129                    activity_type = %activity.activity_type,
130                    workflow_id = %activity.workflow_id,
131                    activity_id = %activity.activity_id,
132                    "no connected worker; waiting for a matching worker to register"
133                );
134                self.registry.wait_for_worker().await;
135            };
136
137            for worker in workers {
138                self.drain_state
139                    .ensure_accepting(&activity.namespace, &activity.activity_type)?;
140                span_fields.record("worker_id", format!("{:?}", worker.id()));
141                // The gRPC dispatch path only registers gRPC-delivery workers, so
142                // a worker here always carries a stream sender; a missing one means
143                // a non-gRPC-transport worker leaked into this path and cannot be
144                // served over it, so it is deregistered like a closed stream.
145                if let Some(sender) = worker.sender() {
146                    if sender
147                        .send(WorkerMessage::ActivityTask(activity.to_task()))
148                        .await
149                        .is_ok()
150                    {
151                        return Ok(());
152                    }
153                }
154                self.registry.deregister(worker.id())?;
155            }
156
157            Err(ServerError::worker_dispatch(
158                activity.namespace.clone(),
159                activity.activity_type.clone(),
160                format!(
161                    "all matching worker streams in task queue {} closed before task could be \
162                     delivered",
163                    activity.task_queue
164                ),
165            ))
166        }
167        .instrument(span)
168        .await
169        .inspect_err(|error| {
170            log_dispatch_error("activity_dispatch", activity, error);
171        })
172    }
173}
174
175fn log_dispatch_error(operation: &'static str, activity: &ScheduledActivity, error: &ServerError) {
176    let fields = error.trace_fields();
177    tracing::error!(
178        operation,
179        namespace = %activity.namespace,
180        task_queue = %activity.task_queue,
181        node = activity.node.as_deref(),
182        workflow_id = %activity.workflow_id,
183        activity_id = %activity.activity_id,
184        activity_type = %activity.activity_type,
185        error_type = %fields.error_type,
186        store_error_type = fields.store_error_type,
187        reason = %fields.reason,
188        "activity dispatch failed"
189    );
190}
191
192/// Decoded activity outcome reported by a worker.
193#[derive(Clone, Debug, Eq, PartialEq)]
194pub enum ActivityCompletionOutcome {
195    /// Activity completed successfully with an output payload.
196    Succeeded(Payload),
197    /// Activity failed, preserving retryability classification for the engine.
198    Failed(ActivityError),
199}
200
201/// Correlated activity completion handed to the engine-owned activity contract.
202#[derive(Clone, Debug, Eq, PartialEq)]
203pub struct ActivityCompletion {
204    /// Owning workflow id.
205    pub workflow_id: WorkflowId,
206    /// Correlating activity id.
207    pub activity_id: ActivityId,
208    /// Concrete workflow run echoed by the worker, when known.
209    pub run_id: Option<RunId>,
210    /// Worker-reported outcome.
211    pub outcome: ActivityCompletionOutcome,
212}
213
214impl TryFrom<ProtoActivityResult> for ActivityCompletion {
215    type Error = ServerError;
216
217    fn try_from(value: ProtoActivityResult) -> Result<Self, Self::Error> {
218        let workflow_id = value
219            .workflow_id
220            .ok_or_else(|| wire_error("activity result workflow id is missing"))
221            .and_then(|id| WorkflowId::try_from(id).map_err(ServerError::from))?;
222        let activity_id = value
223            .activity_id
224            .ok_or_else(|| wire_error("activity result activity id is missing"))
225            .map(ActivityId::from)?;
226        let run_id = value
227            .run_id
228            .map(|id| RunId::try_from(id).map_err(ServerError::from))
229            .transpose()?;
230        let outcome = match value.outcome {
231            Some(proto_activity_result::Outcome::Result(payload)) => {
232                ActivityCompletionOutcome::Succeeded(
233                    Payload::try_from(payload).map_err(ServerError::from)?,
234                )
235            }
236            Some(proto_activity_result::Outcome::Error(error)) => {
237                ActivityCompletionOutcome::Failed(
238                    ActivityError::try_from(error).map_err(ServerError::from)?,
239                )
240            }
241            None => return Err(wire_error("activity result outcome is missing")),
242        };
243
244        Ok(Self {
245            workflow_id,
246            activity_id,
247            run_id,
248            outcome,
249        })
250    }
251}
252
253/// Engine-owned activity completion contract used by the worker endpoint.
254pub trait ActivityCompletionSink {
255    /// Feed one worker-reported result into the engine activity contract.
256    ///
257    /// # Errors
258    ///
259    /// Returns [`ServerError`] when the engine rejects or cannot record the completion.
260    fn complete_activity(&self, completion: ActivityCompletion) -> Result<(), ServerError>;
261}
262
263/// Decode and hand a worker result to the engine-owned activity completion sink.
264///
265/// # Errors
266///
267/// Returns [`ServerError`] for malformed wire results or sink failures.
268pub fn handle_activity_result(
269    sink: &impl ActivityCompletionSink,
270    result: ProtoActivityResult,
271) -> Result<(), ServerError> {
272    sink.complete_activity(ActivityCompletion::try_from(result)?)
273}
274
275/// Build the retryable failure reported when a worker loses ownership of an in-flight task.
276///
277/// The retryable classification models worker loss as infrastructure failure: aion-server
278/// only reports the failure to the engine activity contract; the engine remains responsible
279/// for applying the activity retry policy.
280#[must_use]
281pub fn lost_worker_error(worker_id: crate::worker::registry::WorkerId) -> ActivityError {
282    ActivityError {
283        kind: ActivityErrorKind::Retryable,
284        message: format!("worker {worker_id:?} lost before reporting activity result"),
285        details: None,
286    }
287}
288
289fn wire_error(message: &'static str) -> ServerError {
290    ServerError::Wire {
291        wire: WireError::backend(message),
292    }
293}
294
295#[cfg(test)]
296mod tests {
297    use std::sync::Mutex;
298
299    use aion_core::{ActivityErrorKind, ContentType};
300    use aion_proto::{ProtoActivityError, ProtoActivityErrorKind};
301    use serde_json::json;
302    use uuid::Uuid;
303
304    use crate::worker::registry::ConnectedWorkerRegistry;
305
306    use super::*;
307
308    fn workflow_id() -> WorkflowId {
309        WorkflowId::new(Uuid::nil())
310    }
311
312    fn activity_id() -> ActivityId {
313        ActivityId::from_sequence_position(42)
314    }
315
316    fn payload(value: &serde_json::Value) -> Result<Payload, Box<dyn std::error::Error>> {
317        Ok(Payload::from_json(value)?)
318    }
319
320    #[tokio::test]
321    async fn dispatch_pushes_activity_task_with_correlation()
322    -> Result<(), Box<dyn std::error::Error>> {
323        let registry = ConnectedWorkerRegistry::default();
324        let (tx, mut rx) = tokio::sync::mpsc::channel(1);
325        let activity_types = [String::from("charge-card")];
326        let registration = registry.register("tenant-a", activity_types.iter(), tx)?;
327        let dispatcher = ActivityDispatcher::new(registry.clone());
328        let input = payload(&json!({"amount": 1200}))?;
329        let scheduled = ScheduledActivity {
330            namespace: String::from("tenant-a"),
331            task_queue: String::from("default"),
332            activity_type: String::from("charge-card"),
333            node: None,
334            workflow_id: workflow_id(),
335            activity_id: activity_id(),
336            run_id: None,
337            input: input.clone(),
338            attempt: 1,
339            labels: std::collections::BTreeMap::new(),
340        };
341
342        dispatcher.dispatch(&scheduled).await?;
343        let message = rx.recv().await.ok_or("expected pushed activity task")?;
344        let WorkerMessage::ActivityTask(task) = message else {
345            return Err("expected activity task message".into());
346        };
347
348        assert_eq!(task.workflow_id, Some(ProtoWorkflowId::from(workflow_id())));
349        assert_eq!(task.activity_id, Some(ProtoActivityId::from(activity_id())));
350        assert_eq!(task.activity_type, "charge-card");
351        assert_eq!(task.input, Some(ProtoPayload::from(input)));
352        assert_eq!(task.attempt, 1, "wire task must carry the stamped attempt");
353
354        registration.deregister()?;
355        Ok(())
356    }
357
358    #[tokio::test]
359    async fn dispatch_waits_for_worker_then_delivers() -> Result<(), Box<dyn std::error::Error>> {
360        let registry = ConnectedWorkerRegistry::default();
361        let dispatcher = ActivityDispatcher::new(registry.clone());
362        let scheduled = ScheduledActivity {
363            namespace: String::from("tenant-a"),
364            task_queue: String::from("default"),
365            activity_type: String::from("charge-card"),
366            node: None,
367            workflow_id: workflow_id(),
368            activity_id: activity_id(),
369            run_id: None,
370            input: Payload::new(ContentType::Json, b"{}".to_vec()),
371            attempt: 1,
372            labels: std::collections::BTreeMap::new(),
373        };
374
375        let dispatch_handle = tokio::spawn({
376            let dispatcher = dispatcher.clone();
377            let scheduled = scheduled.clone();
378            async move { dispatcher.dispatch(&scheduled).await }
379        });
380
381        tokio::time::sleep(std::time::Duration::from_millis(50)).await;
382        assert!(!dispatch_handle.is_finished(), "dispatch should be waiting");
383
384        let (tx, mut rx) = tokio::sync::mpsc::channel(1);
385        let activity_types = [String::from("charge-card")];
386        let _registration = registry.register("tenant-a", activity_types.iter(), tx)?;
387
388        dispatch_handle.await??;
389        assert!(rx.recv().await.is_some());
390        Ok(())
391    }
392
393    #[tokio::test]
394    async fn dispatch_skips_closed_worker_and_uses_next_match()
395    -> Result<(), Box<dyn std::error::Error>> {
396        let registry = ConnectedWorkerRegistry::default();
397        let (closed_tx, closed_rx) = tokio::sync::mpsc::channel(1);
398        let (live_tx, mut live_rx) = tokio::sync::mpsc::channel(1);
399        let activity_types = [String::from("charge-card")];
400        let closed_registration =
401            registry.register("tenant-a", activity_types.iter(), closed_tx)?;
402        let live_registration = registry.register("tenant-a", activity_types.iter(), live_tx)?;
403        drop(closed_rx);
404
405        let dispatcher = ActivityDispatcher::new(registry.clone());
406        let scheduled = ScheduledActivity {
407            namespace: String::from("tenant-a"),
408            task_queue: String::from("default"),
409            activity_type: String::from("charge-card"),
410            node: None,
411            workflow_id: workflow_id(),
412            activity_id: activity_id(),
413            run_id: None,
414            input: Payload::new(ContentType::Json, b"{}".to_vec()),
415            attempt: 1,
416            labels: std::collections::BTreeMap::new(),
417        };
418
419        dispatcher.dispatch(&scheduled).await?;
420
421        assert!(live_rx.recv().await.is_some());
422        assert_eq!(
423            registry
424                .workers_for("tenant-a", "default", "charge-card", None)?
425                .len(),
426            1
427        );
428
429        closed_registration.deregister()?;
430        live_registration.deregister()?;
431        Ok(())
432    }
433
434    #[derive(Default)]
435    struct RecordingSink {
436        completions: Mutex<Vec<ActivityCompletion>>,
437    }
438
439    impl ActivityCompletionSink for RecordingSink {
440        fn complete_activity(&self, completion: ActivityCompletion) -> Result<(), ServerError> {
441            self.completions
442                .lock()
443                .map_err(|_| ServerError::lock_poisoned("recording completion sink"))?
444                .push(completion);
445            Ok(())
446        }
447    }
448
449    #[test]
450    fn successful_activity_result_calls_completion_sink() -> Result<(), Box<dyn std::error::Error>>
451    {
452        let sink = RecordingSink::default();
453        let output = payload(&json!({"ok": true}))?;
454        let result = ProtoActivityResult {
455            workflow_id: Some(ProtoWorkflowId::from(workflow_id())),
456            activity_id: Some(ProtoActivityId::from(activity_id())),
457            run_id: None,
458            outcome: Some(proto_activity_result::Outcome::Result(ProtoPayload::from(
459                output.clone(),
460            ))),
461        };
462
463        handle_activity_result(&sink, result)?;
464        let completions = sink
465            .completions
466            .lock()
467            .map_err(|_| ServerError::lock_poisoned("recording completion sink"))?;
468
469        assert_eq!(completions.len(), 1);
470        assert_eq!(completions[0].workflow_id, workflow_id());
471        assert_eq!(completions[0].activity_id, activity_id());
472        assert_eq!(
473            completions[0].outcome,
474            ActivityCompletionOutcome::Succeeded(output)
475        );
476        Ok(())
477    }
478
479    #[test]
480    fn failed_activity_result_preserves_error_classification()
481    -> Result<(), Box<dyn std::error::Error>> {
482        let sink = RecordingSink::default();
483        let error = ProtoActivityError {
484            kind: ProtoActivityErrorKind::Retryable as i32,
485            message: String::from("temporary outage"),
486            details: Some(ProtoPayload::from(payload(
487                &json!({"retry_after_ms": 500}),
488            )?)),
489        };
490        let result = ProtoActivityResult {
491            workflow_id: Some(ProtoWorkflowId::from(workflow_id())),
492            activity_id: Some(ProtoActivityId::from(activity_id())),
493            run_id: None,
494            outcome: Some(proto_activity_result::Outcome::Error(error)),
495        };
496
497        handle_activity_result(&sink, result)?;
498        let completions = sink
499            .completions
500            .lock()
501            .map_err(|_| ServerError::lock_poisoned("recording completion sink"))?;
502
503        assert_eq!(completions.len(), 1);
504        match &completions[0].outcome {
505            ActivityCompletionOutcome::Failed(error) => {
506                assert_eq!(error.kind, ActivityErrorKind::Retryable);
507                assert!(error.is_retryable());
508            }
509            ActivityCompletionOutcome::Succeeded(_) => return Err("expected failed outcome".into()),
510        }
511        Ok(())
512    }
513}