1use 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#[derive(Clone, Debug, Eq, PartialEq)]
16pub struct ScheduledActivity {
17 pub namespace: String,
19 pub activity_type: String,
21 pub workflow_id: WorkflowId,
23 pub activity_id: ActivityId,
25 pub input: Payload,
27 pub attempt: u32,
30}
31
32impl ScheduledActivity {
33 #[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#[derive(Clone, Debug)]
48pub struct ActivityDispatcher {
49 registry: ConnectedWorkerRegistry,
50 drain_state: DrainState,
51}
52
53impl ActivityDispatcher {
54 #[must_use]
56 pub fn new(registry: ConnectedWorkerRegistry) -> Self {
57 Self {
58 registry,
59 drain_state: DrainState::default(),
60 }
61 }
62
63 #[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 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#[derive(Clone, Debug, Eq, PartialEq)]
164pub enum ActivityCompletionOutcome {
165 Succeeded(Payload),
167 Failed(ActivityError),
169}
170
171#[derive(Clone, Debug, Eq, PartialEq)]
173pub struct ActivityCompletion {
174 pub workflow_id: WorkflowId,
176 pub activity_id: ActivityId,
178 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
216pub trait ActivityCompletionSink {
218 fn complete_activity(&self, completion: ActivityCompletion) -> Result<(), ServerError>;
224}
225
226pub 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#[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}