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}
58
59#[async_trait]
67pub trait WorkerSession: Send {
68 async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError>;
77
78 async fn register(
90 &mut self,
91 activity_types: Vec<String>,
92 available_handlers: &BTreeSet<String>,
93 ) -> Result<(), WorkerError>;
94
95 async fn register_with_contract(
102 &mut self,
103 activity_types: Vec<String>,
104 activities: Vec<aion_package::ActivityDescriptor>,
105 available_handlers: &BTreeSet<String>,
106 ) -> Result<(), WorkerError> {
107 drop(activities);
108 self.register(activity_types, available_handlers).await
109 }
110
111 fn receive_tasks(&mut self) -> WorkerTaskStream;
113
114 async fn report_result(
117 &mut self,
118 workflow_id: WorkflowId,
119 activity_id: ActivityId,
120 run_id: Option<RunId>,
121 completion_token: String,
122 result: Payload,
123 ) -> Result<(), WorkerError>;
124
125 async fn report_failure(
128 &mut self,
129 workflow_id: WorkflowId,
130 activity_id: ActivityId,
131 run_id: Option<RunId>,
132 completion_token: String,
133 failure: ActivityError,
134 ) -> Result<(), WorkerError>;
135
136 async fn send_heartbeat(
138 &mut self,
139 workflow_id: WorkflowId,
140 activity_id: ActivityId,
141 progress: Option<Payload>,
142 ) -> Result<(), WorkerError>;
143
144 async fn send_connection_heartbeat(&mut self) -> Result<(), WorkerError> {
149 Ok(())
150 }
151
152 fn heartbeat_window(&self) -> Option<std::time::Duration> {
161 None
162 }
163}
164
165pub fn validate_activity_handlers(
171 activity_types: &[String],
172 available_handlers: &BTreeSet<String>,
173) -> Result<(), WorkerError> {
174 if let Some(activity_type) = activity_types
175 .iter()
176 .find(|activity_type| !available_handlers.contains(*activity_type))
177 {
178 return Err(WorkerError::registration(MissingActivityHandler {
179 activity_type: activity_type.clone(),
180 }));
181 }
182
183 Ok(())
184}
185
186#[derive(Clone, Debug, PartialEq, Eq)]
188pub struct RegisteredSessionInfo {
189 pub worker_id: u64,
192 pub namespace: String,
194 pub heartbeat_window: std::time::Duration,
197}
198
199pub struct GrpcWorkerSession {
201 config: WorkerConfig,
202 activity_types: Vec<String>,
203 client: Option<GeneratedClient>,
204 sender: Option<mpsc::Sender<aion_proto::generated::WorkerToServer>>,
205 receiver: Option<tonic::codec::Streaming<aion_proto::generated::ServerToWorker>>,
206 registered_info: Option<RegisteredSessionInfo>,
207}
208
209impl GrpcWorkerSession {
210 pub async fn connect(config: WorkerConfig) -> Result<Self, WorkerError> {
220 let client = GeneratedClient::connect(config.endpoint.clone())
221 .await
222 .map_err(|source| WorkerError::Connect { source })?;
223
224 Ok(Self {
225 config,
226 activity_types: Vec::new(),
227 client: Some(client),
228 sender: None,
229 receiver: None,
230 registered_info: None,
231 })
232 }
233
234 #[must_use]
236 pub fn from_channel(config: WorkerConfig, channel: Channel) -> Self {
237 Self {
238 config,
239 activity_types: Vec::new(),
240 client: Some(GeneratedClient::new(channel)),
241 sender: None,
242 receiver: None,
243 registered_info: None,
244 }
245 }
246
247 #[must_use]
250 pub const fn registered_info(&self) -> Option<&RegisteredSessionInfo> {
251 self.registered_info.as_ref()
252 }
253
254 async fn open_registered_stream(
271 &mut self,
272 register: aion_proto::generated::RegisterWorker,
273 ) -> Result<(), WorkerError> {
274 let client = self.client.as_mut().ok_or_else(|| {
275 WorkerError::registration(SessionStateError {
276 message: String::from("worker session has not completed its handshake"),
277 })
278 })?;
279 let (sender, outbound) = mpsc::channel(16);
280 sender
281 .try_send(aion_proto::generated::WorkerToServer {
282 message: Some(aion_proto::generated::worker_to_server::Message::Register(
283 register,
284 )),
285 })
286 .map_err(|_| {
287 WorkerError::registration(SessionStateError {
288 message: String::from(
289 "could not queue RegisterWorker as the first stream frame",
290 ),
291 })
292 })?;
293 let mut request = Request::new(ReceiverStream::new(outbound));
294 apply_auth_metadata(request.metadata_mut(), &self.config)?;
295 let response = client
296 .stream_worker(request)
297 .await
298 .map_err(registration_denial_error)?;
299 let mut receiver = response.into_inner();
300
301 let first = tokio::time::timeout(self.config.reconnect.max_backoff, receiver.message())
302 .await
303 .map_err(|_| {
304 WorkerError::registration(SessionStateError {
305 message: format!(
306 "server did not acknowledge registration within {:?}",
307 self.config.reconnect.max_backoff
308 ),
309 })
310 })?
311 .map_err(registration_denial_error)?;
312 let ack = match first.and_then(|frame| frame.message) {
313 Some(aion_proto::generated::server_to_worker::Message::RegisterAck(ack)) => ack,
314 Some(_) => {
315 return Err(WorkerError::decode(SessionStateError {
316 message: String::from(
317 "protocol violation: server sent a non-RegisterAck frame before \
318 acknowledging registration",
319 ),
320 }));
321 }
322 None => {
323 return Err(WorkerError::registration(SessionStateError {
324 message: String::from(
325 "server ended the stream before acknowledging registration",
326 ),
327 }));
328 }
329 };
330
331 self.registered_info = Some(RegisteredSessionInfo {
332 worker_id: ack.worker_id,
333 namespace: ack.namespace,
334 heartbeat_window: std::time::Duration::from_millis(ack.heartbeat_window_ms),
335 });
336 self.sender = Some(sender);
337 self.receiver = Some(receiver);
338 Ok(())
339 }
340
341 async fn send_to_server(
346 &self,
347 message: aion_proto::generated::worker_to_server::Message,
348 ) -> Result<(), WorkerError> {
349 let sender = self.sender.as_ref().ok_or_else(|| {
350 WorkerError::registration(SessionStateError {
351 message: String::from("worker stream has not been opened"),
352 })
353 })?;
354 let send = sender.send(aion_proto::generated::WorkerToServer {
355 message: Some(message),
356 });
357 tokio::time::timeout(self.config.reconnect.max_backoff, send)
358 .await
359 .map_err(|_| WorkerError::Transport {
360 source: tonic::Status::unavailable(format!(
361 "worker stream send did not complete within {:?}",
362 self.config.reconnect.max_backoff
363 )),
364 })?
365 .map_err(|source| WorkerError::Transport {
366 source: tonic::Status::unavailable(format!("worker stream send failed: {source}")),
367 })
368 }
369}
370
371fn registration_denial_error(status: tonic::Status) -> WorkerError {
380 if status.code() == tonic::Code::Unauthenticated {
381 WorkerError::Handshake { source: status }
382 } else {
383 WorkerError::Registration {
384 source: Box::new(status),
385 }
386 }
387}
388
389fn apply_auth_metadata(
390 metadata: &mut tonic::metadata::MetadataMap,
391 config: &WorkerConfig,
392) -> Result<(), WorkerError> {
393 let namespaces_value = config.namespaces.join(",");
397 let namespace =
398 MetadataValue::try_from(namespaces_value.as_str()).map_err(|_| WorkerError::Handshake {
399 source: tonic::Status::invalid_argument(
400 "worker namespaces are not valid gRPC metadata",
401 ),
402 })?;
403 let subject =
404 MetadataValue::try_from(config.subject.as_str()).map_err(|_| WorkerError::Handshake {
405 source: tonic::Status::invalid_argument("worker subject is not valid gRPC metadata"),
406 })?;
407 metadata.insert("x-aion-namespaces", namespace);
408 metadata.insert("x-aion-subject", subject);
409 Ok(())
410}
411
412#[async_trait]
413impl WorkerSession for GrpcWorkerSession {
414 async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
415 self.config = config.clone();
416 if self.client.is_none() {
417 self.client = Some(
418 GeneratedClient::connect(self.config.endpoint.clone())
419 .await
420 .map_err(|source| WorkerError::Connect { source })?,
421 );
422 }
423 Ok(())
424 }
425
426 async fn register(
427 &mut self,
428 activity_types: Vec<String>,
429 available_handlers: &BTreeSet<String>,
430 ) -> Result<(), WorkerError> {
431 self.register_with_contract(activity_types, Vec::new(), available_handlers)
432 .await
433 }
434
435 async fn register_with_contract(
436 &mut self,
437 activity_types: Vec<String>,
438 activities: Vec<aion_package::ActivityDescriptor>,
439 available_handlers: &BTreeSet<String>,
440 ) -> Result<(), WorkerError> {
441 validate_activity_handlers(&activity_types, available_handlers)?;
442 self.activity_types.clone_from(&activity_types);
443
444 let activities = activities
450 .into_iter()
451 .map(|activity| {
452 Ok(aion_proto::generated::ActivityDescriptor {
453 name: activity.name,
454 input_schema_json: serde_json::to_string(&activity.input_schema)
455 .map_err(WorkerError::encode)?,
456 output_schema_json: serde_json::to_string(&activity.output_schema)
457 .map_err(WorkerError::encode)?,
458 })
459 })
460 .collect::<Result<Vec<_>, WorkerError>>()?;
461 let register = aion_proto::generated::RegisterWorker {
462 namespaces: self.config.namespaces.clone(),
463 activity_types,
464 task_queue: self.config.task_queue.clone(),
465 node: self.config.node.clone(),
466 activities,
467 identity: self.config.identity.clone(),
468 instance: None,
469 };
470 self.open_registered_stream(register).await
471 }
472
473 fn receive_tasks(&mut self) -> WorkerTaskStream {
474 match self.receiver.take() {
475 Some(receiver) => Box::pin(receiver.filter_map(|message| async move {
476 Some(match message {
477 Ok(server_message) => decode_server_message(server_message),
478 Err(source) => Err(WorkerError::Transport { source }),
479 })
480 })),
481 None => Box::pin(futures::stream::iter([Err(WorkerError::Transport {
482 source: tonic::Status::failed_precondition(
483 "worker receive stream has not been opened",
484 ),
485 })])),
486 }
487 }
488
489 async fn report_result(
490 &mut self,
491 workflow_id: WorkflowId,
492 activity_id: ActivityId,
493 run_id: Option<RunId>,
494 completion_token: String,
495 result: Payload,
496 ) -> Result<(), WorkerError> {
497 let run_id = run_id.ok_or_else(|| {
498 WorkerError::decode(SessionStateError {
499 message: String::from(
500 "activity result run_id is missing; refusing an incomplete fenced report",
501 ),
502 })
503 })?;
504 let result = ProtoActivityResult {
505 workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
506 activity_id: Some(ProtoActivityId::from(activity_id)),
507 run_id: Some(ProtoRunId::from(run_id)),
508 outcome: Some(proto_activity_result::Outcome::Result(ProtoPayload::from(
509 result,
510 ))),
511 completion_token,
512 };
513 self.send_to_server(aion_proto::generated::worker_to_server::Message::Result(
514 generated_activity_result(result),
515 ))
516 .await
517 }
518
519 async fn report_failure(
520 &mut self,
521 workflow_id: WorkflowId,
522 activity_id: ActivityId,
523 run_id: Option<RunId>,
524 completion_token: String,
525 failure: ActivityError,
526 ) -> Result<(), WorkerError> {
527 let run_id = run_id.ok_or_else(|| {
528 WorkerError::decode(SessionStateError {
529 message: String::from(
530 "activity failure run_id is missing; refusing an incomplete fenced report",
531 ),
532 })
533 })?;
534 let result = ProtoActivityResult {
535 workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
536 activity_id: Some(ProtoActivityId::from(activity_id)),
537 run_id: Some(ProtoRunId::from(run_id)),
538 outcome: Some(proto_activity_result::Outcome::Error(failure.into())),
539 completion_token,
540 };
541 self.send_to_server(aion_proto::generated::worker_to_server::Message::Result(
542 generated_activity_result(result),
543 ))
544 .await
545 }
546
547 async fn send_heartbeat(
548 &mut self,
549 workflow_id: WorkflowId,
550 activity_id: ActivityId,
551 progress: Option<Payload>,
552 ) -> Result<(), WorkerError> {
553 let heartbeat = ProtoHeartbeat {
554 workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
555 activity_id: Some(ProtoActivityId::from(activity_id)),
556 progress: progress.map(ProtoPayload::from),
557 };
558 self.send_to_server(aion_proto::generated::worker_to_server::Message::Heartbeat(
559 generated_heartbeat(heartbeat),
560 ))
561 .await
562 }
563
564 async fn send_connection_heartbeat(&mut self) -> Result<(), WorkerError> {
565 let heartbeat = ProtoHeartbeat {
566 workflow_id: None,
567 activity_id: None,
568 progress: None,
569 };
570 self.send_to_server(aion_proto::generated::worker_to_server::Message::Heartbeat(
571 generated_heartbeat(heartbeat),
572 ))
573 .await
574 }
575
576 fn heartbeat_window(&self) -> Option<std::time::Duration> {
577 self.registered_info
578 .as_ref()
579 .map(|info| info.heartbeat_window)
580 }
581}
582
583fn decode_server_message(
584 message: aion_proto::generated::ServerToWorker,
585) -> Result<WorkerSessionEvent, WorkerError> {
586 match message.message {
587 Some(aion_proto::generated::server_to_worker::Message::Task(task)) => {
588 Ok(WorkerSessionEvent::Task(Box::new(proto_task(task))))
589 }
590 Some(aion_proto::generated::server_to_worker::Message::Drain(_)) => {
591 Ok(WorkerSessionEvent::Drain)
592 }
593 Some(aion_proto::generated::server_to_worker::Message::ResultAck(ack)) => {
594 decode_result_ack(ack)
595 }
596 Some(aion_proto::generated::server_to_worker::Message::RegisterAck(_)) => {
597 Err(WorkerError::decode(SessionStateError {
600 message: String::from(
601 "protocol violation: RegisterAck received after registration completed",
602 ),
603 }))
604 }
605 None => Err(WorkerError::decode(SessionStateError {
606 message: String::from("server-to-worker message was empty"),
607 })),
608 }
609}
610
611fn decode_result_ack(
612 ack: aion_proto::generated::ResultAck,
613) -> Result<WorkerSessionEvent, WorkerError> {
614 let workflow_id = ack
615 .workflow_id
616 .ok_or_else(|| {
617 WorkerError::decode(SessionStateError {
618 message: String::from("result ack workflow_id is missing"),
619 })
620 })
621 .and_then(|id| {
622 WorkflowId::try_from(ProtoWorkflowId { uuid: id.uuid }).map_err(|source| {
623 WorkerError::decode(SessionStateError {
624 message: format!("result ack workflow_id is invalid: {source}"),
625 })
626 })
627 })?;
628 let activity_id = ack
629 .activity_id
630 .map(|id| ActivityId::from_sequence_position(id.sequence_position))
631 .ok_or_else(|| {
632 WorkerError::decode(SessionStateError {
633 message: String::from("result ack activity_id is missing"),
634 })
635 })?;
636 Ok(WorkerSessionEvent::ResultAck {
637 workflow_id,
638 activity_id,
639 })
640}
641
642fn generated_activity_result(value: ProtoActivityResult) -> aion_proto::generated::ActivityResult {
643 aion_proto::generated::ActivityResult {
644 workflow_id: value.workflow_id.map(generated_workflow_id),
645 activity_id: value.activity_id.map(generated_activity_id),
646 run_id: value.run_id.map(generated_run_id),
647 completion_token: value.completion_token,
648 outcome: value.outcome.map(|outcome| match outcome {
649 proto_activity_result::Outcome::Result(result) => {
650 aion_proto::generated::activity_result::Outcome::Result(generated_payload(result))
651 }
652 proto_activity_result::Outcome::Error(error) => {
653 aion_proto::generated::activity_result::Outcome::Error(generated_error(error))
654 }
655 }),
656 }
657}
658
659fn generated_heartbeat(value: ProtoHeartbeat) -> aion_proto::generated::Heartbeat {
660 aion_proto::generated::Heartbeat {
661 workflow_id: value.workflow_id.map(generated_workflow_id),
662 activity_id: value.activity_id.map(generated_activity_id),
663 progress: value.progress.map(generated_payload),
664 }
665}
666
667fn proto_task(value: aion_proto::generated::ActivityTask) -> ProtoActivityTask {
668 ProtoActivityTask {
669 workflow_id: value.workflow_id.map(proto_workflow_id),
670 activity_id: value.activity_id.map(proto_activity_id),
671 activity_type: value.activity_type,
672 input: value.input.map(proto_payload),
673 attempt: value.attempt,
674 labels: value.labels,
675 run_id: value.run_id.map(proto_run_id),
676 completion_token: value.completion_token,
677 idempotency_key: value.idempotency_key,
678 }
679}
680
681fn generated_payload(value: ProtoPayload) -> aion_proto::generated::Payload {
682 aion_proto::generated::Payload {
683 content_type: value.content_type,
684 bytes: value.bytes,
685 }
686}
687
688fn proto_payload(value: aion_proto::generated::Payload) -> ProtoPayload {
689 ProtoPayload {
690 content_type: value.content_type,
691 bytes: value.bytes,
692 }
693}
694
695fn generated_workflow_id(value: ProtoWorkflowId) -> aion_proto::generated::WorkflowId {
696 aion_proto::generated::WorkflowId { uuid: value.uuid }
697}
698
699fn proto_workflow_id(value: aion_proto::generated::WorkflowId) -> ProtoWorkflowId {
700 ProtoWorkflowId { uuid: value.uuid }
701}
702
703fn generated_run_id(value: ProtoRunId) -> aion_proto::generated::RunId {
704 aion_proto::generated::RunId { uuid: value.uuid }
705}
706
707fn proto_run_id(value: aion_proto::generated::RunId) -> ProtoRunId {
708 ProtoRunId { uuid: value.uuid }
709}
710
711fn generated_activity_id(value: ProtoActivityId) -> aion_proto::generated::ActivityId {
712 aion_proto::generated::ActivityId {
713 sequence_position: value.sequence_position,
714 }
715}
716
717fn proto_activity_id(value: aion_proto::generated::ActivityId) -> ProtoActivityId {
718 ProtoActivityId {
719 sequence_position: value.sequence_position,
720 }
721}
722
723fn generated_error(value: aion_proto::ProtoActivityError) -> aion_proto::generated::ActivityError {
724 aion_proto::generated::ActivityError {
725 kind: value.kind,
726 message: value.message,
727 details: value.details.map(generated_payload),
728 }
729}
730
731#[derive(thiserror::Error, Debug)]
732#[error("{message}")]
733struct SessionStateError {
734 message: String,
735}
736
737#[cfg(test)]
738mod tests {
739 use std::collections::BTreeSet;
740
741 use aion_proto::ProtoActivityTask;
742 use async_trait::async_trait;
743 use futures::{StreamExt, stream};
744
745 use super::{
746 WorkerSession, WorkerSessionEvent, WorkerTaskStream, apply_auth_metadata,
747 validate_activity_handlers,
748 };
749 use crate::error::WorkerError;
750 use crate::{ReconnectConfig, WorkerConfig};
751
752 #[derive(Default)]
753 struct FakeSession {
754 handshakes: Vec<(String, String)>,
755 registrations: Vec<Vec<String>>,
756 }
757
758 #[async_trait]
759 impl WorkerSession for FakeSession {
760 async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
761 self.handshakes
762 .push((config.task_queue.clone(), config.identity.clone()));
763 Ok(())
764 }
765
766 async fn register(
767 &mut self,
768 activity_types: Vec<String>,
769 available_handlers: &BTreeSet<String>,
770 ) -> Result<(), WorkerError> {
771 validate_activity_handlers(&activity_types, available_handlers)?;
772 self.registrations.push(activity_types);
773 Ok(())
774 }
775
776 fn receive_tasks(&mut self) -> WorkerTaskStream {
777 Box::pin(stream::iter([Ok(WorkerSessionEvent::Task(Box::new(
778 ProtoActivityTask {
779 workflow_id: None,
780 activity_id: None,
781 activity_type: String::from("charge-card"),
782 input: None,
783 attempt: 1,
784 labels: std::collections::HashMap::new(),
785 run_id: Some(aion_proto::ProtoRunId::from(aion_core::RunId::new_v4())),
786 completion_token: String::from("generation-1"),
787 idempotency_key: String::from("effect-key"),
788 },
789 )))]))
790 }
791
792 async fn report_result(
793 &mut self,
794 workflow_id: aion_core::WorkflowId,
795 activity_id: aion_core::ActivityId,
796 run_id: Option<aion_core::RunId>,
797 completion_token: String,
798 result: aion_core::Payload,
799 ) -> Result<(), WorkerError> {
800 drop((workflow_id, activity_id, run_id, completion_token, result));
801 Ok(())
802 }
803
804 async fn report_failure(
805 &mut self,
806 workflow_id: aion_core::WorkflowId,
807 activity_id: aion_core::ActivityId,
808 run_id: Option<aion_core::RunId>,
809 completion_token: String,
810 failure: aion_core::ActivityError,
811 ) -> Result<(), WorkerError> {
812 drop((workflow_id, activity_id, run_id, completion_token, failure));
813 Ok(())
814 }
815
816 async fn send_heartbeat(
817 &mut self,
818 workflow_id: aion_core::WorkflowId,
819 activity_id: aion_core::ActivityId,
820 progress: Option<aion_core::Payload>,
821 ) -> Result<(), WorkerError> {
822 drop((workflow_id, activity_id, progress));
823 Ok(())
824 }
825 }
826
827 #[test]
828 fn apply_auth_metadata_sets_worker_authorization_headers() -> Result<(), WorkerError> {
829 let config = WorkerConfig::builder()
830 .endpoint("http://127.0.0.1:50051")
831 .task_queue("payments")
832 .identity("worker-a")
833 .max_concurrency(4)
834 .reconnect_initial_backoff(std::time::Duration::from_millis(5))
835 .reconnect_max_backoff(std::time::Duration::from_millis(20))
836 .reconnect_max_attempts(3)
837 .namespace("payments")
838 .subject("worker-a")
839 .build()
840 .map_err(WorkerError::registration)?;
841 let mut metadata = tonic::metadata::MetadataMap::new();
842
843 apply_auth_metadata(&mut metadata, &config)?;
844
845 assert_eq!(
846 metadata
847 .get("x-aion-namespaces")
848 .and_then(|value| value.to_str().ok()),
849 Some("payments")
850 );
851 assert_eq!(
852 metadata
853 .get("x-aion-subject")
854 .and_then(|value| value.to_str().ok()),
855 Some("worker-a")
856 );
857 Ok(())
858 }
859
860 #[tokio::test]
861 async fn fake_session_records_handshake_and_registration() -> Result<(), WorkerError> {
862 let config = WorkerConfig::new(
863 "http://127.0.0.1:50051",
864 "payments",
865 "worker-a",
866 4,
867 ReconnectConfig::new(
868 std::time::Duration::from_millis(5),
869 std::time::Duration::from_millis(20),
870 3,
871 ),
872 None,
873 );
874 let activity_types = vec![String::from("charge-card"), String::from("send-email")];
875 let handlers = activity_types.iter().cloned().collect::<BTreeSet<_>>();
876 let mut session = FakeSession::default();
877
878 session.handshake(&config).await?;
879 session.register(activity_types.clone(), &handlers).await?;
880 let received = session.receive_tasks().next().await;
881
882 assert_eq!(
883 session.handshakes,
884 vec![(String::from("payments"), String::from("worker-a"))]
885 );
886 assert_eq!(session.registrations, vec![activity_types]);
887 assert!(received.is_some());
888
889 Ok(())
890 }
891
892 #[tokio::test]
893 async fn grpc_reports_echo_the_dispatched_completion_token() -> Result<(), WorkerError> {
894 let config = WorkerConfig::new(
895 "http://127.0.0.1:50051",
896 "payments",
897 "worker-a",
898 1,
899 ReconnectConfig::new(
900 std::time::Duration::from_millis(5),
901 std::time::Duration::from_millis(20),
902 3,
903 ),
904 None,
905 );
906 let (sender, mut receiver) = tokio::sync::mpsc::channel(2);
907 let mut session = super::GrpcWorkerSession {
908 config,
909 activity_types: Vec::new(),
910 client: None,
911 sender: Some(sender),
912 receiver: None,
913 registered_info: None,
914 };
915
916 session
917 .report_result(
918 aion_core::WorkflowId::new_v4(),
919 aion_core::ActivityId::from_sequence_position(1),
920 Some(aion_core::RunId::new_v4()),
921 String::from("success-generation"),
922 aion_core::Payload::new(aion_core::ContentType::Json, b"{}".to_vec()),
923 )
924 .await?;
925 session
926 .report_failure(
927 aion_core::WorkflowId::new_v4(),
928 aion_core::ActivityId::from_sequence_position(2),
929 Some(aion_core::RunId::new_v4()),
930 String::from("failure-generation"),
931 aion_core::ActivityError {
932 kind: aion_core::ActivityErrorKind::Terminal,
933 message: String::from("failed"),
934 details: None,
935 },
936 )
937 .await?;
938
939 let success = receiver.recv().await.ok_or_else(|| {
940 WorkerError::decode(super::SessionStateError {
941 message: String::from("result report channel closed"),
942 })
943 })?;
944 let failure = receiver.recv().await.ok_or_else(|| {
945 WorkerError::decode(super::SessionStateError {
946 message: String::from("failure report channel closed"),
947 })
948 })?;
949 let success_token = match success.message {
950 Some(aion_proto::generated::worker_to_server::Message::Result(result)) => {
951 result.completion_token
952 }
953 _ => {
954 return Err(WorkerError::decode(super::SessionStateError {
955 message: String::from("success report did not emit an ActivityResult"),
956 }));
957 }
958 };
959 let failure_token = match failure.message {
960 Some(aion_proto::generated::worker_to_server::Message::Result(result)) => {
961 result.completion_token
962 }
963 _ => {
964 return Err(WorkerError::decode(super::SessionStateError {
965 message: String::from("failure report did not emit an ActivityResult"),
966 }));
967 }
968 };
969 assert_eq!(success_token, "success-generation");
970 assert_eq!(failure_token, "failure-generation");
971 Ok(())
972 }
973
974 #[tokio::test(start_paused = true)]
978 async fn report_send_times_out_retryably_at_max_backoff() -> Result<(), WorkerError> {
979 let config = WorkerConfig::new(
980 "http://127.0.0.1:50051",
981 "payments",
982 "worker-a",
983 1,
984 ReconnectConfig::new(
985 std::time::Duration::from_millis(5),
986 std::time::Duration::from_millis(20),
987 3,
988 ),
989 None,
990 );
991 let (sender, receiver) = tokio::sync::mpsc::channel(1);
992 sender
995 .try_send(aion_proto::generated::WorkerToServer { message: None })
996 .map_err(WorkerError::decode)?;
997 let mut session = super::GrpcWorkerSession {
998 config,
999 activity_types: Vec::new(),
1000 client: None,
1001 sender: Some(sender),
1002 receiver: None,
1003 registered_info: None,
1004 };
1005
1006 let result = session
1007 .report_result(
1008 aion_core::WorkflowId::new_v4(),
1009 aion_core::ActivityId::from_sequence_position(1),
1010 Some(aion_core::RunId::new_v4()),
1011 String::from("generation-1"),
1012 aion_core::Payload::new(aion_core::ContentType::Json, b"{}".to_vec()),
1013 )
1014 .await;
1015
1016 let Err(error) = result else {
1017 return Err(WorkerError::Transport {
1018 source: tonic::Status::internal("a hung send must time out, not hang"),
1019 });
1020 };
1021 assert!(
1022 matches!(error, WorkerError::Transport { .. }),
1023 "send deadline elapse must be a retryable transport error: {error}"
1024 );
1025 assert!(error.is_retryable());
1026 assert!(
1027 error.to_string().contains("did not complete"),
1028 "the error must name the deadline: {error}"
1029 );
1030 drop(receiver);
1031 Ok(())
1032 }
1033
1034 #[test]
1035 fn registration_rejects_activity_without_handler() {
1036 let activity_types = vec![String::from("charge-card"), String::from("send-email")];
1037 let handlers = [String::from("charge-card")]
1038 .into_iter()
1039 .collect::<BTreeSet<_>>();
1040
1041 let result = validate_activity_handlers(&activity_types, &handlers);
1042 assert!(result.is_err());
1043 let error = match result {
1044 Ok(()) => return,
1045 Err(error) => error,
1046 };
1047
1048 assert_eq!(
1049 error.to_string(),
1050 "worker registration failed: activity type `send-email` has no registered handler"
1051 );
1052 }
1053}