1use 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#[derive(Clone, Debug, Eq, PartialEq)]
18pub struct ScheduledActivity {
19 pub namespace: String,
22 pub task_queue: String,
26 pub activity_type: String,
29 pub node: Option<String>,
36 pub workflow_id: WorkflowId,
38 pub activity_id: ActivityId,
40 pub run_id: Option<RunId>,
42 pub input: Payload,
44 pub attempt: u32,
47 pub labels: BTreeMap<String, String>,
50}
51
52impl ScheduledActivity {
53 #[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#[derive(Clone, Debug)]
70pub struct ActivityDispatcher {
71 registry: ConnectedWorkerRegistry,
72 drain_state: DrainState,
73}
74
75impl ActivityDispatcher {
76 #[must_use]
78 pub fn new(registry: ConnectedWorkerRegistry) -> Self {
79 Self {
80 registry,
81 drain_state: DrainState::default(),
82 }
83 }
84
85 #[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 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 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#[derive(Clone, Debug, Eq, PartialEq)]
194pub enum ActivityCompletionOutcome {
195 Succeeded(Payload),
197 Failed(ActivityError),
199}
200
201#[derive(Clone, Debug, Eq, PartialEq)]
203pub struct ActivityCompletion {
204 pub workflow_id: WorkflowId,
206 pub activity_id: ActivityId,
208 pub run_id: Option<RunId>,
210 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
253pub trait ActivityCompletionSink {
255 fn complete_activity(&self, completion: ActivityCompletion) -> Result<(), ServerError>;
261}
262
263pub 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#[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}