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::RegisterAck(_)) => {
656 Err(WorkerError::decode(SessionStateError {
659 message: String::from(
660 "protocol violation: RegisterAck received after registration completed",
661 ),
662 }))
663 }
664 None => Err(WorkerError::decode(SessionStateError {
665 message: String::from("server-to-worker message was empty"),
666 })),
667 }
668}
669
670fn decode_result_ack(
671 ack: aion_proto::generated::ResultAck,
672) -> Result<WorkerSessionEvent, WorkerError> {
673 let workflow_id = ack
674 .workflow_id
675 .ok_or_else(|| {
676 WorkerError::decode(SessionStateError {
677 message: String::from("result ack workflow_id is missing"),
678 })
679 })
680 .and_then(|id| {
681 WorkflowId::try_from(ProtoWorkflowId { uuid: id.uuid }).map_err(|source| {
682 WorkerError::decode(SessionStateError {
683 message: format!("result ack workflow_id is invalid: {source}"),
684 })
685 })
686 })?;
687 let activity_id = ack
688 .activity_id
689 .map(|id| ActivityId::from_sequence_position(id.sequence_position))
690 .ok_or_else(|| {
691 WorkerError::decode(SessionStateError {
692 message: String::from("result ack activity_id is missing"),
693 })
694 })?;
695 Ok(WorkerSessionEvent::ResultAck {
696 workflow_id,
697 activity_id,
698 })
699}
700
701fn generated_activity_result(value: ProtoActivityResult) -> aion_proto::generated::ActivityResult {
702 aion_proto::generated::ActivityResult {
703 workflow_id: value.workflow_id.map(generated_workflow_id),
704 activity_id: value.activity_id.map(generated_activity_id),
705 run_id: value.run_id.map(generated_run_id),
706 completion_token: value.completion_token,
707 outcome: value.outcome.map(|outcome| match outcome {
708 proto_activity_result::Outcome::Result(result) => {
709 aion_proto::generated::activity_result::Outcome::Result(generated_payload(result))
710 }
711 proto_activity_result::Outcome::Error(error) => {
712 aion_proto::generated::activity_result::Outcome::Error(generated_error(error))
713 }
714 }),
715 }
716}
717
718fn generated_heartbeat(value: ProtoHeartbeat) -> aion_proto::generated::Heartbeat {
719 aion_proto::generated::Heartbeat {
720 workflow_id: value.workflow_id.map(generated_workflow_id),
721 activity_id: value.activity_id.map(generated_activity_id),
722 progress: value.progress.map(generated_payload),
723 }
724}
725
726fn proto_task(value: aion_proto::generated::ActivityTask) -> ProtoActivityTask {
727 ProtoActivityTask {
728 workflow_id: value.workflow_id.map(proto_workflow_id),
729 activity_id: value.activity_id.map(proto_activity_id),
730 activity_type: value.activity_type,
731 input: value.input.map(proto_payload),
732 attempt: value.attempt,
733 labels: value.labels,
734 run_id: value.run_id.map(proto_run_id),
735 completion_token: value.completion_token,
736 idempotency_key: value.idempotency_key,
737 }
738}
739
740fn generated_payload(value: ProtoPayload) -> aion_proto::generated::Payload {
741 aion_proto::generated::Payload {
742 content_type: value.content_type,
743 bytes: value.bytes,
744 }
745}
746
747fn proto_payload(value: aion_proto::generated::Payload) -> ProtoPayload {
748 ProtoPayload {
749 content_type: value.content_type,
750 bytes: value.bytes,
751 }
752}
753
754fn generated_workflow_id(value: ProtoWorkflowId) -> aion_proto::generated::WorkflowId {
755 aion_proto::generated::WorkflowId { uuid: value.uuid }
756}
757
758fn proto_workflow_id(value: aion_proto::generated::WorkflowId) -> ProtoWorkflowId {
759 ProtoWorkflowId { uuid: value.uuid }
760}
761
762fn generated_run_id(value: ProtoRunId) -> aion_proto::generated::RunId {
763 aion_proto::generated::RunId { uuid: value.uuid }
764}
765
766fn proto_run_id(value: aion_proto::generated::RunId) -> ProtoRunId {
767 ProtoRunId { uuid: value.uuid }
768}
769
770fn generated_activity_id(value: ProtoActivityId) -> aion_proto::generated::ActivityId {
771 aion_proto::generated::ActivityId {
772 sequence_position: value.sequence_position,
773 }
774}
775
776fn proto_activity_id(value: aion_proto::generated::ActivityId) -> ProtoActivityId {
777 ProtoActivityId {
778 sequence_position: value.sequence_position,
779 }
780}
781
782fn generated_error(value: aion_proto::ProtoActivityError) -> aion_proto::generated::ActivityError {
783 aion_proto::generated::ActivityError {
784 kind: value.kind,
785 message: value.message,
786 details: value.details.map(generated_payload),
787 }
788}
789
790#[derive(thiserror::Error, Debug)]
791#[error("{message}")]
792struct SessionStateError {
793 message: String,
794}
795
796#[cfg(test)]
797mod tests {
798 use std::collections::BTreeSet;
799
800 use aion_proto::ProtoActivityTask;
801 use async_trait::async_trait;
802 use futures::{StreamExt, stream};
803
804 use super::{
805 WorkerSession, WorkerSessionEvent, WorkerTaskStream, apply_auth_metadata,
806 validate_activity_handlers,
807 };
808 use crate::error::WorkerError;
809 use crate::{ReconnectConfig, WorkerConfig};
810
811 #[derive(Default)]
812 struct FakeSession {
813 handshakes: Vec<(String, String)>,
814 registrations: Vec<Vec<String>>,
815 }
816
817 #[async_trait]
818 impl WorkerSession for FakeSession {
819 async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
820 self.handshakes
821 .push((config.task_queue.clone(), config.identity.clone()));
822 Ok(())
823 }
824
825 async fn register(
826 &mut self,
827 activity_types: Vec<String>,
828 available_handlers: &BTreeSet<String>,
829 ) -> Result<(), WorkerError> {
830 validate_activity_handlers(&activity_types, available_handlers)?;
831 self.registrations.push(activity_types);
832 Ok(())
833 }
834
835 fn receive_tasks(&mut self) -> WorkerTaskStream {
836 Box::pin(stream::iter([Ok(WorkerSessionEvent::Task(Box::new(
837 ProtoActivityTask {
838 workflow_id: None,
839 activity_id: None,
840 activity_type: String::from("charge-card"),
841 input: None,
842 attempt: 1,
843 labels: std::collections::HashMap::new(),
844 run_id: Some(aion_proto::ProtoRunId::from(aion_core::RunId::new_v4())),
845 completion_token: String::from("generation-1"),
846 idempotency_key: String::from("effect-key"),
847 },
848 )))]))
849 }
850
851 async fn report_result(
852 &mut self,
853 workflow_id: aion_core::WorkflowId,
854 activity_id: aion_core::ActivityId,
855 run_id: Option<aion_core::RunId>,
856 completion_token: String,
857 result: aion_core::Payload,
858 ) -> Result<(), WorkerError> {
859 drop((workflow_id, activity_id, run_id, completion_token, result));
860 Ok(())
861 }
862
863 async fn report_failure(
864 &mut self,
865 workflow_id: aion_core::WorkflowId,
866 activity_id: aion_core::ActivityId,
867 run_id: Option<aion_core::RunId>,
868 completion_token: String,
869 failure: aion_core::ActivityError,
870 ) -> Result<(), WorkerError> {
871 drop((workflow_id, activity_id, run_id, completion_token, failure));
872 Ok(())
873 }
874
875 async fn send_heartbeat(
876 &mut self,
877 workflow_id: aion_core::WorkflowId,
878 activity_id: aion_core::ActivityId,
879 progress: Option<aion_core::Payload>,
880 ) -> Result<(), WorkerError> {
881 drop((workflow_id, activity_id, progress));
882 Ok(())
883 }
884 }
885
886 #[test]
887 fn apply_auth_metadata_sets_worker_authorization_headers() -> Result<(), WorkerError> {
888 let config = WorkerConfig::builder()
889 .endpoint("http://127.0.0.1:50051")
890 .task_queue("payments")
891 .identity("worker-a")
892 .max_concurrency(4)
893 .reconnect_initial_backoff(std::time::Duration::from_millis(5))
894 .reconnect_max_backoff(std::time::Duration::from_millis(20))
895 .reconnect_max_attempts(3)
896 .namespace("payments")
897 .subject("worker-a")
898 .build()
899 .map_err(WorkerError::registration)?;
900 let mut metadata = tonic::metadata::MetadataMap::new();
901
902 apply_auth_metadata(&mut metadata, &config)?;
903
904 assert_eq!(
905 metadata
906 .get("x-aion-namespaces")
907 .and_then(|value| value.to_str().ok()),
908 Some("payments")
909 );
910 assert_eq!(
911 metadata
912 .get("x-aion-subject")
913 .and_then(|value| value.to_str().ok()),
914 Some("worker-a")
915 );
916 Ok(())
917 }
918
919 #[tokio::test]
920 async fn fake_session_records_handshake_and_registration() -> Result<(), WorkerError> {
921 let config = WorkerConfig::new(
922 "http://127.0.0.1:50051",
923 "payments",
924 "worker-a",
925 4,
926 ReconnectConfig::new(
927 std::time::Duration::from_millis(5),
928 std::time::Duration::from_millis(20),
929 3,
930 ),
931 None,
932 );
933 let activity_types = vec![String::from("charge-card"), String::from("send-email")];
934 let handlers = activity_types.iter().cloned().collect::<BTreeSet<_>>();
935 let mut session = FakeSession::default();
936
937 session.handshake(&config).await?;
938 session.register(activity_types.clone(), &handlers).await?;
939 let received = session.receive_tasks().next().await;
940
941 assert_eq!(
942 session.handshakes,
943 vec![(String::from("payments"), String::from("worker-a"))]
944 );
945 assert_eq!(session.registrations, vec![activity_types]);
946 assert!(received.is_some());
947
948 Ok(())
949 }
950
951 #[tokio::test]
952 async fn grpc_reports_echo_the_dispatched_completion_token() -> Result<(), WorkerError> {
953 let config = WorkerConfig::new(
954 "http://127.0.0.1:50051",
955 "payments",
956 "worker-a",
957 1,
958 ReconnectConfig::new(
959 std::time::Duration::from_millis(5),
960 std::time::Duration::from_millis(20),
961 3,
962 ),
963 None,
964 );
965 let (sender, mut receiver) = tokio::sync::mpsc::channel(2);
966 let mut session = super::GrpcWorkerSession {
967 config,
968 activity_types: Vec::new(),
969 client: None,
970 sender: Some(sender),
971 receiver: None,
972 registered_info: None,
973 };
974
975 session
976 .report_result(
977 aion_core::WorkflowId::new_v4(),
978 aion_core::ActivityId::from_sequence_position(1),
979 Some(aion_core::RunId::new_v4()),
980 String::from("success-generation"),
981 aion_core::Payload::new(aion_core::ContentType::Json, b"{}".to_vec()),
982 )
983 .await?;
984 session
985 .report_failure(
986 aion_core::WorkflowId::new_v4(),
987 aion_core::ActivityId::from_sequence_position(2),
988 Some(aion_core::RunId::new_v4()),
989 String::from("failure-generation"),
990 aion_core::ActivityError {
991 kind: aion_core::ActivityErrorKind::Terminal,
992 message: String::from("failed"),
993 details: None,
994 },
995 )
996 .await?;
997
998 let success = receiver.recv().await.ok_or_else(|| {
999 WorkerError::decode(super::SessionStateError {
1000 message: String::from("result report channel closed"),
1001 })
1002 })?;
1003 let failure = receiver.recv().await.ok_or_else(|| {
1004 WorkerError::decode(super::SessionStateError {
1005 message: String::from("failure report channel closed"),
1006 })
1007 })?;
1008 let success_token = match success.message {
1009 Some(aion_proto::generated::worker_to_server::Message::Result(result)) => {
1010 result.completion_token
1011 }
1012 _ => {
1013 return Err(WorkerError::decode(super::SessionStateError {
1014 message: String::from("success report did not emit an ActivityResult"),
1015 }));
1016 }
1017 };
1018 let failure_token = match failure.message {
1019 Some(aion_proto::generated::worker_to_server::Message::Result(result)) => {
1020 result.completion_token
1021 }
1022 _ => {
1023 return Err(WorkerError::decode(super::SessionStateError {
1024 message: String::from("failure report did not emit an ActivityResult"),
1025 }));
1026 }
1027 };
1028 assert_eq!(success_token, "success-generation");
1029 assert_eq!(failure_token, "failure-generation");
1030 Ok(())
1031 }
1032
1033 #[tokio::test(start_paused = true)]
1037 async fn report_send_times_out_retryably_at_max_backoff() -> Result<(), WorkerError> {
1038 let config = WorkerConfig::new(
1039 "http://127.0.0.1:50051",
1040 "payments",
1041 "worker-a",
1042 1,
1043 ReconnectConfig::new(
1044 std::time::Duration::from_millis(5),
1045 std::time::Duration::from_millis(20),
1046 3,
1047 ),
1048 None,
1049 );
1050 let (sender, receiver) = tokio::sync::mpsc::channel(1);
1051 sender
1054 .try_send(aion_proto::generated::WorkerToServer { message: None })
1055 .map_err(WorkerError::decode)?;
1056 let mut session = super::GrpcWorkerSession {
1057 config,
1058 activity_types: Vec::new(),
1059 client: None,
1060 sender: Some(sender),
1061 receiver: None,
1062 registered_info: None,
1063 };
1064
1065 let result = session
1066 .report_result(
1067 aion_core::WorkflowId::new_v4(),
1068 aion_core::ActivityId::from_sequence_position(1),
1069 Some(aion_core::RunId::new_v4()),
1070 String::from("generation-1"),
1071 aion_core::Payload::new(aion_core::ContentType::Json, b"{}".to_vec()),
1072 )
1073 .await;
1074
1075 let Err(error) = result else {
1076 return Err(WorkerError::Transport {
1077 source: tonic::Status::internal("a hung send must time out, not hang"),
1078 });
1079 };
1080 assert!(
1081 matches!(error, WorkerError::Transport { .. }),
1082 "send deadline elapse must be a retryable transport error: {error}"
1083 );
1084 assert!(error.is_retryable());
1085 assert!(
1086 error.to_string().contains("did not complete"),
1087 "the error must name the deadline: {error}"
1088 );
1089 drop(receiver);
1090 Ok(())
1091 }
1092
1093 #[test]
1094 fn registration_rejects_activity_without_handler() {
1095 let activity_types = vec![String::from("charge-card"), String::from("send-email")];
1096 let handlers = [String::from("charge-card")]
1097 .into_iter()
1098 .collect::<BTreeSet<_>>();
1099
1100 let result = validate_activity_handlers(&activity_types, &handlers);
1101 assert!(result.is_err());
1102 let error = match result {
1103 Ok(()) => return,
1104 Err(error) => error,
1105 };
1106
1107 assert_eq!(
1108 error.to_string(),
1109 "worker registration failed: activity type `send-email` has no registered handler"
1110 );
1111 }
1112}