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