1use std::collections::BTreeSet;
4use std::pin::Pin;
5
6use aion_core::{ActivityError, ActivityId, Payload, RunId, WorkflowId};
7use aion_proto::{
8 ProtoActivityId, ProtoActivityResult, ProtoActivityTask, ProtoHeartbeat, ProtoPayload,
9 ProtoRunId, ProtoWorkflowId, proto_activity_result,
10};
11use async_trait::async_trait;
12use futures::{Stream, StreamExt};
13use tokio::sync::mpsc;
14use tokio_stream::wrappers::ReceiverStream;
15use tonic::{Request, metadata::MetadataValue, transport::Channel};
16
17use crate::config::WorkerConfig;
18use crate::error::{MissingActivityHandler, WorkerError};
19
20type GeneratedClient = aion_proto::generated::worker_protocol_client::WorkerProtocolClient<Channel>;
21
22pub type WorkerTaskStream =
24 Pin<Box<dyn Stream<Item = Result<WorkerSessionEvent, WorkerError>> + Send>>;
25
26#[derive(Clone, Debug, PartialEq, Eq)]
28pub enum WorkerSessionEvent {
29 Task(Box<ProtoActivityTask>),
31 Drain,
37 ResultAck {
40 workflow_id: WorkflowId,
42 activity_id: ActivityId,
44 },
45 Cancel {
52 workflow_id: WorkflowId,
54 activity_id: ActivityId,
56 },
57 LivenessPing {
64 sequence: u64,
67 silence_window: std::time::Duration,
71 },
72}
73
74#[async_trait]
82pub trait WorkerSession: Send {
83 async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError>;
92
93 async fn register(
105 &mut self,
106 activity_types: Vec<String>,
107 available_handlers: &BTreeSet<String>,
108 ) -> Result<(), WorkerError>;
109
110 async fn register_with_contract(
117 &mut self,
118 activity_types: Vec<String>,
119 activities: Vec<aion_package::ActivityDescriptor>,
120 available_handlers: &BTreeSet<String>,
121 ) -> Result<(), WorkerError> {
122 drop(activities);
123 self.register(activity_types, available_handlers).await
124 }
125
126 fn receive_tasks(&mut self) -> WorkerTaskStream;
128
129 async fn report_result(
132 &mut self,
133 workflow_id: WorkflowId,
134 activity_id: ActivityId,
135 run_id: Option<RunId>,
136 completion_token: String,
137 result: Payload,
138 ) -> Result<(), WorkerError>;
139
140 async fn report_failure(
143 &mut self,
144 workflow_id: WorkflowId,
145 activity_id: ActivityId,
146 run_id: Option<RunId>,
147 completion_token: String,
148 failure: ActivityError,
149 ) -> Result<(), WorkerError>;
150
151 async fn send_heartbeat(
153 &mut self,
154 workflow_id: WorkflowId,
155 activity_id: ActivityId,
156 progress: Option<Payload>,
157 ) -> Result<(), WorkerError>;
158
159 async fn send_connection_heartbeat(&mut self) -> Result<(), WorkerError> {
164 Ok(())
165 }
166
167 async fn answer_liveness_ping(&mut self, sequence: u64) -> Result<(), WorkerError> {
185 tracing::debug!(
186 liveness_ping = sequence,
187 "this session's transport carries no liveness-answer frame, so nothing was sent. A \
188 server that probed over this transport will judge the ping unanswered — the honest \
189 verdict for a link this worker cannot answer on"
190 );
191 Ok(())
192 }
193
194 fn heartbeat_window(&self) -> Option<std::time::Duration> {
203 None
204 }
205}
206
207pub fn validate_activity_handlers(
213 activity_types: &[String],
214 available_handlers: &BTreeSet<String>,
215) -> Result<(), WorkerError> {
216 if let Some(activity_type) = activity_types
217 .iter()
218 .find(|activity_type| !available_handlers.contains(*activity_type))
219 {
220 return Err(WorkerError::registration(MissingActivityHandler {
221 activity_type: activity_type.clone(),
222 }));
223 }
224
225 Ok(())
226}
227
228#[derive(Clone, Debug, PartialEq, Eq)]
230pub struct RegisteredSessionInfo {
231 pub worker_id: u64,
234 pub namespace: String,
236 pub heartbeat_window: std::time::Duration,
239}
240
241pub struct GrpcWorkerSession {
243 config: WorkerConfig,
244 activity_types: Vec<String>,
245 client: Option<GeneratedClient>,
246 sender: Option<mpsc::Sender<aion_proto::generated::WorkerToServer>>,
247 receiver: Option<tonic::codec::Streaming<aion_proto::generated::ServerToWorker>>,
248 registered_info: Option<RegisteredSessionInfo>,
249}
250
251impl GrpcWorkerSession {
252 pub async fn connect(config: WorkerConfig) -> Result<Self, WorkerError> {
262 let client = GeneratedClient::connect(config.endpoint.clone())
263 .await
264 .map_err(|source| WorkerError::Connect { source })?;
265
266 Ok(Self {
267 config,
268 activity_types: Vec::new(),
269 client: Some(client),
270 sender: None,
271 receiver: None,
272 registered_info: None,
273 })
274 }
275
276 #[must_use]
278 pub fn from_channel(config: WorkerConfig, channel: Channel) -> Self {
279 Self {
280 config,
281 activity_types: Vec::new(),
282 client: Some(GeneratedClient::new(channel)),
283 sender: None,
284 receiver: None,
285 registered_info: None,
286 }
287 }
288
289 #[must_use]
292 pub const fn registered_info(&self) -> Option<&RegisteredSessionInfo> {
293 self.registered_info.as_ref()
294 }
295
296 async fn open_registered_stream(
313 &mut self,
314 register: aion_proto::generated::RegisterWorker,
315 ) -> Result<(), WorkerError> {
316 let client = self.client.as_mut().ok_or_else(|| {
317 WorkerError::registration(SessionStateError {
318 message: String::from("worker session has not completed its handshake"),
319 })
320 })?;
321 let (sender, outbound) = mpsc::channel(16);
322 sender
323 .try_send(aion_proto::generated::WorkerToServer {
324 message: Some(aion_proto::generated::worker_to_server::Message::Register(
325 register,
326 )),
327 })
328 .map_err(|_| {
329 WorkerError::registration(SessionStateError {
330 message: String::from(
331 "could not queue RegisterWorker as the first stream frame",
332 ),
333 })
334 })?;
335 let mut request = Request::new(ReceiverStream::new(outbound));
336 apply_auth_metadata(request.metadata_mut(), &self.config)?;
337 let response = client
338 .stream_worker(request)
339 .await
340 .map_err(registration_denial_error)?;
341 let mut receiver = response.into_inner();
342
343 let first = tokio::time::timeout(self.config.reconnect.max_backoff, receiver.message())
344 .await
345 .map_err(|_| {
346 WorkerError::registration(SessionStateError {
347 message: format!(
348 "server did not acknowledge registration within {:?}",
349 self.config.reconnect.max_backoff
350 ),
351 })
352 })?
353 .map_err(registration_denial_error)?;
354 let ack = match first.and_then(|frame| frame.message) {
355 Some(aion_proto::generated::server_to_worker::Message::RegisterAck(ack)) => ack,
356 Some(_) => {
357 return Err(WorkerError::decode(SessionStateError {
358 message: String::from(
359 "protocol violation: server sent a non-RegisterAck frame before \
360 acknowledging registration",
361 ),
362 }));
363 }
364 None => {
365 return Err(WorkerError::registration(SessionStateError {
366 message: String::from(
367 "server ended the stream before acknowledging registration",
368 ),
369 }));
370 }
371 };
372
373 self.registered_info = Some(RegisteredSessionInfo {
374 worker_id: ack.worker_id,
375 namespace: ack.namespace,
376 heartbeat_window: std::time::Duration::from_millis(ack.heartbeat_window_ms),
377 });
378 self.sender = Some(sender);
379 self.receiver = Some(receiver);
380 Ok(())
381 }
382
383 async fn send_to_server(
388 &self,
389 message: aion_proto::generated::worker_to_server::Message,
390 ) -> Result<(), WorkerError> {
391 let sender = self.sender.as_ref().ok_or_else(|| {
392 WorkerError::registration(SessionStateError {
393 message: String::from("worker stream has not been opened"),
394 })
395 })?;
396 let send = sender.send(aion_proto::generated::WorkerToServer {
397 message: Some(message),
398 });
399 tokio::time::timeout(self.config.reconnect.max_backoff, send)
400 .await
401 .map_err(|_| WorkerError::Transport {
402 source: tonic::Status::unavailable(format!(
403 "worker stream send did not complete within {:?}",
404 self.config.reconnect.max_backoff
405 )),
406 })?
407 .map_err(|source| WorkerError::Transport {
408 source: tonic::Status::unavailable(format!("worker stream send failed: {source}")),
409 })
410 }
411}
412
413fn registration_denial_error(status: tonic::Status) -> WorkerError {
422 if status.code() == tonic::Code::Unauthenticated {
423 WorkerError::Handshake { source: status }
424 } else {
425 WorkerError::Registration {
426 source: Box::new(status),
427 }
428 }
429}
430
431fn apply_auth_metadata(
432 metadata: &mut tonic::metadata::MetadataMap,
433 config: &WorkerConfig,
434) -> Result<(), WorkerError> {
435 let namespaces_value = config.namespaces.join(",");
439 let namespace =
440 MetadataValue::try_from(namespaces_value.as_str()).map_err(|_| WorkerError::Handshake {
441 source: tonic::Status::invalid_argument(
442 "worker namespaces are not valid gRPC metadata",
443 ),
444 })?;
445 let subject =
446 MetadataValue::try_from(config.subject.as_str()).map_err(|_| WorkerError::Handshake {
447 source: tonic::Status::invalid_argument("worker subject is not valid gRPC metadata"),
448 })?;
449 metadata.insert("x-aion-namespaces", namespace);
450 metadata.insert("x-aion-subject", subject);
451 Ok(())
452}
453
454#[async_trait]
455impl WorkerSession for GrpcWorkerSession {
456 async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
457 self.config = config.clone();
458 if self.client.is_none() {
459 self.client = Some(
460 GeneratedClient::connect(self.config.endpoint.clone())
461 .await
462 .map_err(|source| WorkerError::Connect { source })?,
463 );
464 }
465 Ok(())
466 }
467
468 async fn register(
469 &mut self,
470 activity_types: Vec<String>,
471 available_handlers: &BTreeSet<String>,
472 ) -> Result<(), WorkerError> {
473 self.register_with_contract(activity_types, Vec::new(), available_handlers)
474 .await
475 }
476
477 async fn register_with_contract(
478 &mut self,
479 activity_types: Vec<String>,
480 activities: Vec<aion_package::ActivityDescriptor>,
481 available_handlers: &BTreeSet<String>,
482 ) -> Result<(), WorkerError> {
483 validate_activity_handlers(&activity_types, available_handlers)?;
484 self.activity_types.clone_from(&activity_types);
485
486 let activities = activities
492 .into_iter()
493 .map(|activity| {
494 Ok(aion_proto::generated::ActivityDescriptor {
495 name: activity.name,
496 input_schema_json: serde_json::to_string(&activity.input_schema)
497 .map_err(WorkerError::encode)?,
498 output_schema_json: serde_json::to_string(&activity.output_schema)
499 .map_err(WorkerError::encode)?,
500 })
501 })
502 .collect::<Result<Vec<_>, WorkerError>>()?;
503 let register = aion_proto::generated::RegisterWorker {
504 namespaces: self.config.namespaces.clone(),
505 activity_types,
506 task_queue: self.config.task_queue.clone(),
507 node: self.config.node.clone(),
508 activities,
509 identity: self.config.identity.clone(),
510 instance: None,
511 };
512 self.open_registered_stream(register).await
513 }
514
515 fn receive_tasks(&mut self) -> WorkerTaskStream {
516 match self.receiver.take() {
517 Some(receiver) => Box::pin(receiver.filter_map(|message| async move {
518 Some(match message {
519 Ok(server_message) => decode_server_message(server_message),
520 Err(source) => Err(WorkerError::Transport { source }),
521 })
522 })),
523 None => Box::pin(futures::stream::iter([Err(WorkerError::Transport {
524 source: tonic::Status::failed_precondition(
525 "worker receive stream has not been opened",
526 ),
527 })])),
528 }
529 }
530
531 async fn report_result(
532 &mut self,
533 workflow_id: WorkflowId,
534 activity_id: ActivityId,
535 run_id: Option<RunId>,
536 completion_token: String,
537 result: Payload,
538 ) -> Result<(), WorkerError> {
539 let run_id = run_id.ok_or_else(|| {
540 WorkerError::decode(SessionStateError {
541 message: String::from(
542 "activity result run_id is missing; refusing an incomplete fenced report",
543 ),
544 })
545 })?;
546 let result = ProtoActivityResult {
547 workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
548 activity_id: Some(ProtoActivityId::from(activity_id)),
549 run_id: Some(ProtoRunId::from(run_id)),
550 outcome: Some(proto_activity_result::Outcome::Result(ProtoPayload::from(
551 result,
552 ))),
553 completion_token,
554 };
555 self.send_to_server(aion_proto::generated::worker_to_server::Message::Result(
556 generated_activity_result(result),
557 ))
558 .await
559 }
560
561 async fn report_failure(
562 &mut self,
563 workflow_id: WorkflowId,
564 activity_id: ActivityId,
565 run_id: Option<RunId>,
566 completion_token: String,
567 failure: ActivityError,
568 ) -> Result<(), WorkerError> {
569 let run_id = run_id.ok_or_else(|| {
570 WorkerError::decode(SessionStateError {
571 message: String::from(
572 "activity failure run_id is missing; refusing an incomplete fenced report",
573 ),
574 })
575 })?;
576 let result = ProtoActivityResult {
577 workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
578 activity_id: Some(ProtoActivityId::from(activity_id)),
579 run_id: Some(ProtoRunId::from(run_id)),
580 outcome: Some(proto_activity_result::Outcome::Error(failure.into())),
581 completion_token,
582 };
583 self.send_to_server(aion_proto::generated::worker_to_server::Message::Result(
584 generated_activity_result(result),
585 ))
586 .await
587 }
588
589 async fn send_heartbeat(
590 &mut self,
591 workflow_id: WorkflowId,
592 activity_id: ActivityId,
593 progress: Option<Payload>,
594 ) -> Result<(), WorkerError> {
595 let heartbeat = ProtoHeartbeat {
596 workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
597 activity_id: Some(ProtoActivityId::from(activity_id)),
598 progress: progress.map(ProtoPayload::from),
599 };
600 self.send_to_server(aion_proto::generated::worker_to_server::Message::Heartbeat(
601 generated_heartbeat(heartbeat),
602 ))
603 .await
604 }
605
606 async fn send_connection_heartbeat(&mut self) -> Result<(), WorkerError> {
607 let heartbeat = ProtoHeartbeat {
608 workflow_id: None,
609 activity_id: None,
610 progress: None,
611 };
612 self.send_to_server(aion_proto::generated::worker_to_server::Message::Heartbeat(
613 generated_heartbeat(heartbeat),
614 ))
615 .await
616 }
617
618 async fn answer_liveness_ping(&mut self, sequence: u64) -> Result<(), WorkerError> {
619 self.send_to_server(
620 aion_proto::generated::worker_to_server::Message::LivenessAnswer(
621 aion_proto::generated::LivenessAnswer {
622 liveness_ping: sequence,
623 },
624 ),
625 )
626 .await
627 }
628
629 fn heartbeat_window(&self) -> Option<std::time::Duration> {
630 self.registered_info
631 .as_ref()
632 .map(|info| info.heartbeat_window)
633 }
634}
635
636fn decode_server_message(
637 message: aion_proto::generated::ServerToWorker,
638) -> Result<WorkerSessionEvent, WorkerError> {
639 match message.message {
640 Some(aion_proto::generated::server_to_worker::Message::Task(task)) => {
641 Ok(WorkerSessionEvent::Task(Box::new(proto_task(task))))
642 }
643 Some(aion_proto::generated::server_to_worker::Message::Drain(_)) => {
644 Ok(WorkerSessionEvent::Drain)
645 }
646 Some(aion_proto::generated::server_to_worker::Message::ResultAck(ack)) => {
647 decode_result_ack(ack)
648 }
649 Some(aion_proto::generated::server_to_worker::Message::LivenessPing(ping)) => {
650 Ok(WorkerSessionEvent::LivenessPing {
651 sequence: ping.liveness_ping,
652 silence_window: std::time::Duration::from_millis(ping.silence_window_ms),
653 })
654 }
655 Some(aion_proto::generated::server_to_worker::Message::CancelActivity(cancel)) => {
656 decode_cancel_activity(cancel)
657 }
658 Some(aion_proto::generated::server_to_worker::Message::RegisterAck(_)) => {
659 Err(WorkerError::decode(SessionStateError {
662 message: String::from(
663 "protocol violation: RegisterAck received after registration completed",
664 ),
665 }))
666 }
667 None => Err(WorkerError::decode(SessionStateError {
668 message: String::from("server-to-worker message was empty"),
669 })),
670 }
671}
672
673fn decode_cancel_activity(
686 cancel: aion_proto::generated::CancelActivity,
687) -> Result<WorkerSessionEvent, WorkerError> {
688 let workflow_id = cancel
689 .workflow_id
690 .ok_or_else(|| {
691 WorkerError::decode(SessionStateError {
692 message: String::from("cancel activity workflow_id is missing"),
693 })
694 })
695 .and_then(|id| {
696 WorkflowId::try_from(ProtoWorkflowId { uuid: id.uuid }).map_err(|source| {
697 WorkerError::decode(SessionStateError {
698 message: format!("cancel activity workflow_id is invalid: {source}"),
699 })
700 })
701 })?;
702 let activity_id = cancel
703 .activity_id
704 .map(|id| ActivityId::from_sequence_position(id.sequence_position))
705 .ok_or_else(|| {
706 WorkerError::decode(SessionStateError {
707 message: String::from("cancel activity activity_id is missing"),
708 })
709 })?;
710 Ok(WorkerSessionEvent::Cancel {
711 workflow_id,
712 activity_id,
713 })
714}
715
716fn decode_result_ack(
717 ack: aion_proto::generated::ResultAck,
718) -> Result<WorkerSessionEvent, WorkerError> {
719 let workflow_id = ack
720 .workflow_id
721 .ok_or_else(|| {
722 WorkerError::decode(SessionStateError {
723 message: String::from("result ack workflow_id is missing"),
724 })
725 })
726 .and_then(|id| {
727 WorkflowId::try_from(ProtoWorkflowId { uuid: id.uuid }).map_err(|source| {
728 WorkerError::decode(SessionStateError {
729 message: format!("result ack workflow_id is invalid: {source}"),
730 })
731 })
732 })?;
733 let activity_id = ack
734 .activity_id
735 .map(|id| ActivityId::from_sequence_position(id.sequence_position))
736 .ok_or_else(|| {
737 WorkerError::decode(SessionStateError {
738 message: String::from("result ack activity_id is missing"),
739 })
740 })?;
741 Ok(WorkerSessionEvent::ResultAck {
742 workflow_id,
743 activity_id,
744 })
745}
746
747fn generated_activity_result(value: ProtoActivityResult) -> aion_proto::generated::ActivityResult {
748 aion_proto::generated::ActivityResult {
749 workflow_id: value.workflow_id.map(generated_workflow_id),
750 activity_id: value.activity_id.map(generated_activity_id),
751 run_id: value.run_id.map(generated_run_id),
752 completion_token: value.completion_token,
753 outcome: value.outcome.map(|outcome| match outcome {
754 proto_activity_result::Outcome::Result(result) => {
755 aion_proto::generated::activity_result::Outcome::Result(generated_payload(result))
756 }
757 proto_activity_result::Outcome::Error(error) => {
758 aion_proto::generated::activity_result::Outcome::Error(generated_error(error))
759 }
760 }),
761 }
762}
763
764fn generated_heartbeat(value: ProtoHeartbeat) -> aion_proto::generated::Heartbeat {
765 aion_proto::generated::Heartbeat {
766 workflow_id: value.workflow_id.map(generated_workflow_id),
767 activity_id: value.activity_id.map(generated_activity_id),
768 progress: value.progress.map(generated_payload),
769 }
770}
771
772fn proto_task(value: aion_proto::generated::ActivityTask) -> ProtoActivityTask {
773 ProtoActivityTask {
774 workflow_id: value.workflow_id.map(proto_workflow_id),
775 activity_id: value.activity_id.map(proto_activity_id),
776 activity_type: value.activity_type,
777 input: value.input.map(proto_payload),
778 attempt: value.attempt,
779 labels: value.labels,
780 run_id: value.run_id.map(proto_run_id),
781 completion_token: value.completion_token,
782 idempotency_key: value.idempotency_key,
783 }
784}
785
786fn generated_payload(value: ProtoPayload) -> aion_proto::generated::Payload {
787 aion_proto::generated::Payload {
788 content_type: value.content_type,
789 bytes: value.bytes,
790 }
791}
792
793fn proto_payload(value: aion_proto::generated::Payload) -> ProtoPayload {
794 ProtoPayload {
795 content_type: value.content_type,
796 bytes: value.bytes,
797 }
798}
799
800fn generated_workflow_id(value: ProtoWorkflowId) -> aion_proto::generated::WorkflowId {
801 aion_proto::generated::WorkflowId { uuid: value.uuid }
802}
803
804fn proto_workflow_id(value: aion_proto::generated::WorkflowId) -> ProtoWorkflowId {
805 ProtoWorkflowId { uuid: value.uuid }
806}
807
808fn generated_run_id(value: ProtoRunId) -> aion_proto::generated::RunId {
809 aion_proto::generated::RunId { uuid: value.uuid }
810}
811
812fn proto_run_id(value: aion_proto::generated::RunId) -> ProtoRunId {
813 ProtoRunId { uuid: value.uuid }
814}
815
816fn generated_activity_id(value: ProtoActivityId) -> aion_proto::generated::ActivityId {
817 aion_proto::generated::ActivityId {
818 sequence_position: value.sequence_position,
819 }
820}
821
822fn proto_activity_id(value: aion_proto::generated::ActivityId) -> ProtoActivityId {
823 ProtoActivityId {
824 sequence_position: value.sequence_position,
825 }
826}
827
828fn generated_error(value: aion_proto::ProtoActivityError) -> aion_proto::generated::ActivityError {
829 aion_proto::generated::ActivityError {
830 kind: value.kind,
831 message: value.message,
832 details: value.details.map(generated_payload),
833 }
834}
835
836#[derive(thiserror::Error, Debug)]
837#[error("{message}")]
838struct SessionStateError {
839 message: String,
840}
841
842#[cfg(test)]
843mod tests {
844 use std::collections::BTreeSet;
845
846 use aion_proto::ProtoActivityTask;
847 use async_trait::async_trait;
848 use futures::{StreamExt, stream};
849
850 use super::{
851 WorkerSession, WorkerSessionEvent, WorkerTaskStream, apply_auth_metadata,
852 decode_server_message, validate_activity_handlers,
853 };
854 use crate::error::WorkerError;
855 use crate::{ReconnectConfig, WorkerConfig};
856
857 #[derive(Default)]
858 struct FakeSession {
859 handshakes: Vec<(String, String)>,
860 registrations: Vec<Vec<String>>,
861 }
862
863 #[async_trait]
864 impl WorkerSession for FakeSession {
865 async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
866 self.handshakes
867 .push((config.task_queue.clone(), config.identity.clone()));
868 Ok(())
869 }
870
871 async fn register(
872 &mut self,
873 activity_types: Vec<String>,
874 available_handlers: &BTreeSet<String>,
875 ) -> Result<(), WorkerError> {
876 validate_activity_handlers(&activity_types, available_handlers)?;
877 self.registrations.push(activity_types);
878 Ok(())
879 }
880
881 fn receive_tasks(&mut self) -> WorkerTaskStream {
882 Box::pin(stream::iter([Ok(WorkerSessionEvent::Task(Box::new(
883 ProtoActivityTask {
884 workflow_id: None,
885 activity_id: None,
886 activity_type: String::from("charge-card"),
887 input: None,
888 attempt: 1,
889 labels: std::collections::HashMap::new(),
890 run_id: Some(aion_proto::ProtoRunId::from(aion_core::RunId::new_v4())),
891 completion_token: String::from("generation-1"),
892 idempotency_key: String::from("effect-key"),
893 },
894 )))]))
895 }
896
897 async fn report_result(
898 &mut self,
899 workflow_id: aion_core::WorkflowId,
900 activity_id: aion_core::ActivityId,
901 run_id: Option<aion_core::RunId>,
902 completion_token: String,
903 result: aion_core::Payload,
904 ) -> Result<(), WorkerError> {
905 drop((workflow_id, activity_id, run_id, completion_token, result));
906 Ok(())
907 }
908
909 async fn report_failure(
910 &mut self,
911 workflow_id: aion_core::WorkflowId,
912 activity_id: aion_core::ActivityId,
913 run_id: Option<aion_core::RunId>,
914 completion_token: String,
915 failure: aion_core::ActivityError,
916 ) -> Result<(), WorkerError> {
917 drop((workflow_id, activity_id, run_id, completion_token, failure));
918 Ok(())
919 }
920
921 async fn send_heartbeat(
922 &mut self,
923 workflow_id: aion_core::WorkflowId,
924 activity_id: aion_core::ActivityId,
925 progress: Option<aion_core::Payload>,
926 ) -> Result<(), WorkerError> {
927 drop((workflow_id, activity_id, progress));
928 Ok(())
929 }
930 }
931
932 #[test]
933 fn apply_auth_metadata_sets_worker_authorization_headers() -> Result<(), WorkerError> {
934 let config = WorkerConfig::builder()
935 .endpoint("http://127.0.0.1:50051")
936 .task_queue("payments")
937 .identity("worker-a")
938 .max_concurrency(4)
939 .reconnect_initial_backoff(std::time::Duration::from_millis(5))
940 .reconnect_max_backoff(std::time::Duration::from_millis(20))
941 .reconnect_max_attempts(3)
942 .namespace("payments")
943 .subject("worker-a")
944 .build()
945 .map_err(WorkerError::registration)?;
946 let mut metadata = tonic::metadata::MetadataMap::new();
947
948 apply_auth_metadata(&mut metadata, &config)?;
949
950 assert_eq!(
951 metadata
952 .get("x-aion-namespaces")
953 .and_then(|value| value.to_str().ok()),
954 Some("payments")
955 );
956 assert_eq!(
957 metadata
958 .get("x-aion-subject")
959 .and_then(|value| value.to_str().ok()),
960 Some("worker-a")
961 );
962 Ok(())
963 }
964
965 #[tokio::test]
966 async fn fake_session_records_handshake_and_registration() -> Result<(), WorkerError> {
967 let config = WorkerConfig::new(
968 "http://127.0.0.1:50051",
969 "payments",
970 "worker-a",
971 4,
972 ReconnectConfig::new(
973 std::time::Duration::from_millis(5),
974 std::time::Duration::from_millis(20),
975 3,
976 ),
977 None,
978 );
979 let activity_types = vec![String::from("charge-card"), String::from("send-email")];
980 let handlers = activity_types.iter().cloned().collect::<BTreeSet<_>>();
981 let mut session = FakeSession::default();
982
983 session.handshake(&config).await?;
984 session.register(activity_types.clone(), &handlers).await?;
985 let received = session.receive_tasks().next().await;
986
987 assert_eq!(
988 session.handshakes,
989 vec![(String::from("payments"), String::from("worker-a"))]
990 );
991 assert_eq!(session.registrations, vec![activity_types]);
992 assert!(received.is_some());
993
994 Ok(())
995 }
996
997 #[tokio::test]
998 async fn grpc_reports_echo_the_dispatched_completion_token() -> Result<(), WorkerError> {
999 let config = WorkerConfig::new(
1000 "http://127.0.0.1:50051",
1001 "payments",
1002 "worker-a",
1003 1,
1004 ReconnectConfig::new(
1005 std::time::Duration::from_millis(5),
1006 std::time::Duration::from_millis(20),
1007 3,
1008 ),
1009 None,
1010 );
1011 let (sender, mut receiver) = tokio::sync::mpsc::channel(2);
1012 let mut session = super::GrpcWorkerSession {
1013 config,
1014 activity_types: Vec::new(),
1015 client: None,
1016 sender: Some(sender),
1017 receiver: None,
1018 registered_info: None,
1019 };
1020
1021 session
1022 .report_result(
1023 aion_core::WorkflowId::new_v4(),
1024 aion_core::ActivityId::from_sequence_position(1),
1025 Some(aion_core::RunId::new_v4()),
1026 String::from("success-generation"),
1027 aion_core::Payload::new(aion_core::ContentType::Json, b"{}".to_vec()),
1028 )
1029 .await?;
1030 session
1031 .report_failure(
1032 aion_core::WorkflowId::new_v4(),
1033 aion_core::ActivityId::from_sequence_position(2),
1034 Some(aion_core::RunId::new_v4()),
1035 String::from("failure-generation"),
1036 aion_core::ActivityError {
1037 kind: aion_core::ActivityErrorKind::Terminal,
1038 message: String::from("failed"),
1039 details: None,
1040 },
1041 )
1042 .await?;
1043
1044 let success = receiver.recv().await.ok_or_else(|| {
1045 WorkerError::decode(super::SessionStateError {
1046 message: String::from("result report channel closed"),
1047 })
1048 })?;
1049 let failure = receiver.recv().await.ok_or_else(|| {
1050 WorkerError::decode(super::SessionStateError {
1051 message: String::from("failure report channel closed"),
1052 })
1053 })?;
1054 let success_token = match success.message {
1055 Some(aion_proto::generated::worker_to_server::Message::Result(result)) => {
1056 result.completion_token
1057 }
1058 _ => {
1059 return Err(WorkerError::decode(super::SessionStateError {
1060 message: String::from("success report did not emit an ActivityResult"),
1061 }));
1062 }
1063 };
1064 let failure_token = match failure.message {
1065 Some(aion_proto::generated::worker_to_server::Message::Result(result)) => {
1066 result.completion_token
1067 }
1068 _ => {
1069 return Err(WorkerError::decode(super::SessionStateError {
1070 message: String::from("failure report did not emit an ActivityResult"),
1071 }));
1072 }
1073 };
1074 assert_eq!(success_token, "success-generation");
1075 assert_eq!(failure_token, "failure-generation");
1076 Ok(())
1077 }
1078
1079 #[tokio::test(start_paused = true)]
1083 async fn report_send_times_out_retryably_at_max_backoff() -> Result<(), WorkerError> {
1084 let config = WorkerConfig::new(
1085 "http://127.0.0.1:50051",
1086 "payments",
1087 "worker-a",
1088 1,
1089 ReconnectConfig::new(
1090 std::time::Duration::from_millis(5),
1091 std::time::Duration::from_millis(20),
1092 3,
1093 ),
1094 None,
1095 );
1096 let (sender, receiver) = tokio::sync::mpsc::channel(1);
1097 sender
1100 .try_send(aion_proto::generated::WorkerToServer { message: None })
1101 .map_err(WorkerError::decode)?;
1102 let mut session = super::GrpcWorkerSession {
1103 config,
1104 activity_types: Vec::new(),
1105 client: None,
1106 sender: Some(sender),
1107 receiver: None,
1108 registered_info: None,
1109 };
1110
1111 let result = session
1112 .report_result(
1113 aion_core::WorkflowId::new_v4(),
1114 aion_core::ActivityId::from_sequence_position(1),
1115 Some(aion_core::RunId::new_v4()),
1116 String::from("generation-1"),
1117 aion_core::Payload::new(aion_core::ContentType::Json, b"{}".to_vec()),
1118 )
1119 .await;
1120
1121 let Err(error) = result else {
1122 return Err(WorkerError::Transport {
1123 source: tonic::Status::internal("a hung send must time out, not hang"),
1124 });
1125 };
1126 assert!(
1127 matches!(error, WorkerError::Transport { .. }),
1128 "send deadline elapse must be a retryable transport error: {error}"
1129 );
1130 assert!(error.is_retryable());
1131 assert!(
1132 error.to_string().contains("did not complete"),
1133 "the error must name the deadline: {error}"
1134 );
1135 drop(receiver);
1136 Ok(())
1137 }
1138
1139 #[test]
1140 fn registration_rejects_activity_without_handler() {
1141 let activity_types = vec![String::from("charge-card"), String::from("send-email")];
1142 let handlers = [String::from("charge-card")]
1143 .into_iter()
1144 .collect::<BTreeSet<_>>();
1145
1146 let result = validate_activity_handlers(&activity_types, &handlers);
1147 assert!(result.is_err());
1148 let error = match result {
1149 Ok(()) => return,
1150 Err(error) => error,
1151 };
1152
1153 assert_eq!(
1154 error.to_string(),
1155 "worker registration failed: activity type `send-email` has no registered handler"
1156 );
1157 }
1158
1159 fn cancel_frame(
1162 workflow_uuid: Option<String>,
1163 sequence_position: Option<u64>,
1164 ) -> aion_proto::generated::ServerToWorker {
1165 aion_proto::generated::ServerToWorker {
1166 message: Some(
1167 aion_proto::generated::server_to_worker::Message::CancelActivity(
1168 aion_proto::generated::CancelActivity {
1169 workflow_id: workflow_uuid
1170 .map(|uuid| aion_proto::generated::WorkflowId { uuid }),
1171 activity_id: sequence_position.map(|sequence_position| {
1172 aion_proto::generated::ActivityId { sequence_position }
1173 }),
1174 },
1175 ),
1176 ),
1177 }
1178 }
1179
1180 type CancelTestResult = Result<(), Box<dyn std::error::Error>>;
1181
1182 fn assert_cancel_refused(
1184 frame: aion_proto::generated::ServerToWorker,
1185 expected_refusal: &str,
1186 ) -> CancelTestResult {
1187 let error = match decode_server_message(frame) {
1188 Ok(event) => {
1189 return Err(format!("an invalid cancel was accepted as {event:?}").into());
1190 }
1191 Err(error) => error,
1192 };
1193
1194 assert!(
1195 error.to_string().contains(expected_refusal),
1196 "refusal did not name the defective half: {error}"
1197 );
1198 Ok(())
1199 }
1200
1201 #[test]
1202 fn cancel_frame_decodes_to_the_pair_it_names() -> CancelTestResult {
1203 let workflow_id = aion_core::WorkflowId::new_v4();
1204
1205 let event = decode_server_message(cancel_frame(Some(workflow_id.to_string()), Some(7)))?;
1206
1207 match event {
1211 WorkerSessionEvent::Cancel {
1212 workflow_id: decoded_workflow,
1213 activity_id,
1214 } => {
1215 assert_eq!(decoded_workflow, workflow_id);
1216 assert_eq!(activity_id.sequence_position(), 7);
1217 }
1218 other => {
1219 return Err(format!("cancel frame decoded to the wrong event: {other:?}").into());
1220 }
1221 }
1222 Ok(())
1223 }
1224
1225 #[test]
1226 fn cancel_without_a_workflow_id_is_refused() -> CancelTestResult {
1227 assert_cancel_refused(cancel_frame(None, Some(7)), "workflow_id is missing")
1228 }
1229
1230 #[test]
1231 fn cancel_without_an_activity_id_is_refused() -> CancelTestResult {
1232 let workflow_id = aion_core::WorkflowId::new_v4();
1233 assert_cancel_refused(
1234 cancel_frame(Some(workflow_id.to_string()), None),
1235 "activity_id is missing",
1236 )
1237 }
1238
1239 #[test]
1240 fn cancel_carrying_an_unparseable_workflow_id_is_refused() -> CancelTestResult {
1241 assert_cancel_refused(
1242 cancel_frame(Some(String::from("not-a-uuid")), Some(7)),
1243 "workflow_id is invalid",
1244 )
1245 }
1246}