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 };
469 self.open_registered_stream(register).await
470 }
471
472 fn receive_tasks(&mut self) -> WorkerTaskStream {
473 match self.receiver.take() {
474 Some(receiver) => Box::pin(receiver.filter_map(|message| async move {
475 Some(match message {
476 Ok(server_message) => decode_server_message(server_message),
477 Err(source) => Err(WorkerError::Transport { source }),
478 })
479 })),
480 None => Box::pin(futures::stream::iter([Err(WorkerError::Transport {
481 source: tonic::Status::failed_precondition(
482 "worker receive stream has not been opened",
483 ),
484 })])),
485 }
486 }
487
488 async fn report_result(
489 &mut self,
490 workflow_id: WorkflowId,
491 activity_id: ActivityId,
492 run_id: Option<RunId>,
493 completion_token: String,
494 result: Payload,
495 ) -> Result<(), WorkerError> {
496 let run_id = run_id.ok_or_else(|| {
497 WorkerError::decode(SessionStateError {
498 message: String::from(
499 "activity result run_id is missing; refusing an incomplete fenced report",
500 ),
501 })
502 })?;
503 let result = ProtoActivityResult {
504 workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
505 activity_id: Some(ProtoActivityId::from(activity_id)),
506 run_id: Some(ProtoRunId::from(run_id)),
507 outcome: Some(proto_activity_result::Outcome::Result(ProtoPayload::from(
508 result,
509 ))),
510 completion_token,
511 };
512 self.send_to_server(aion_proto::generated::worker_to_server::Message::Result(
513 generated_activity_result(result),
514 ))
515 .await
516 }
517
518 async fn report_failure(
519 &mut self,
520 workflow_id: WorkflowId,
521 activity_id: ActivityId,
522 run_id: Option<RunId>,
523 completion_token: String,
524 failure: ActivityError,
525 ) -> Result<(), WorkerError> {
526 let run_id = run_id.ok_or_else(|| {
527 WorkerError::decode(SessionStateError {
528 message: String::from(
529 "activity failure run_id is missing; refusing an incomplete fenced report",
530 ),
531 })
532 })?;
533 let result = ProtoActivityResult {
534 workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
535 activity_id: Some(ProtoActivityId::from(activity_id)),
536 run_id: Some(ProtoRunId::from(run_id)),
537 outcome: Some(proto_activity_result::Outcome::Error(failure.into())),
538 completion_token,
539 };
540 self.send_to_server(aion_proto::generated::worker_to_server::Message::Result(
541 generated_activity_result(result),
542 ))
543 .await
544 }
545
546 async fn send_heartbeat(
547 &mut self,
548 workflow_id: WorkflowId,
549 activity_id: ActivityId,
550 progress: Option<Payload>,
551 ) -> Result<(), WorkerError> {
552 let heartbeat = ProtoHeartbeat {
553 workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
554 activity_id: Some(ProtoActivityId::from(activity_id)),
555 progress: progress.map(ProtoPayload::from),
556 };
557 self.send_to_server(aion_proto::generated::worker_to_server::Message::Heartbeat(
558 generated_heartbeat(heartbeat),
559 ))
560 .await
561 }
562
563 async fn send_connection_heartbeat(&mut self) -> Result<(), WorkerError> {
564 let heartbeat = ProtoHeartbeat {
565 workflow_id: None,
566 activity_id: None,
567 progress: None,
568 };
569 self.send_to_server(aion_proto::generated::worker_to_server::Message::Heartbeat(
570 generated_heartbeat(heartbeat),
571 ))
572 .await
573 }
574
575 fn heartbeat_window(&self) -> Option<std::time::Duration> {
576 self.registered_info
577 .as_ref()
578 .map(|info| info.heartbeat_window)
579 }
580}
581
582fn decode_server_message(
583 message: aion_proto::generated::ServerToWorker,
584) -> Result<WorkerSessionEvent, WorkerError> {
585 match message.message {
586 Some(aion_proto::generated::server_to_worker::Message::Task(task)) => {
587 Ok(WorkerSessionEvent::Task(Box::new(proto_task(task))))
588 }
589 Some(aion_proto::generated::server_to_worker::Message::Drain(_)) => {
590 Ok(WorkerSessionEvent::Drain)
591 }
592 Some(aion_proto::generated::server_to_worker::Message::ResultAck(ack)) => {
593 decode_result_ack(ack)
594 }
595 Some(aion_proto::generated::server_to_worker::Message::RegisterAck(_)) => {
596 Err(WorkerError::decode(SessionStateError {
599 message: String::from(
600 "protocol violation: RegisterAck received after registration completed",
601 ),
602 }))
603 }
604 None => Err(WorkerError::decode(SessionStateError {
605 message: String::from("server-to-worker message was empty"),
606 })),
607 }
608}
609
610fn decode_result_ack(
611 ack: aion_proto::generated::ResultAck,
612) -> Result<WorkerSessionEvent, WorkerError> {
613 let workflow_id = ack
614 .workflow_id
615 .ok_or_else(|| {
616 WorkerError::decode(SessionStateError {
617 message: String::from("result ack workflow_id is missing"),
618 })
619 })
620 .and_then(|id| {
621 WorkflowId::try_from(ProtoWorkflowId { uuid: id.uuid }).map_err(|source| {
622 WorkerError::decode(SessionStateError {
623 message: format!("result ack workflow_id is invalid: {source}"),
624 })
625 })
626 })?;
627 let activity_id = ack
628 .activity_id
629 .map(|id| ActivityId::from_sequence_position(id.sequence_position))
630 .ok_or_else(|| {
631 WorkerError::decode(SessionStateError {
632 message: String::from("result ack activity_id is missing"),
633 })
634 })?;
635 Ok(WorkerSessionEvent::ResultAck {
636 workflow_id,
637 activity_id,
638 })
639}
640
641fn generated_activity_result(value: ProtoActivityResult) -> aion_proto::generated::ActivityResult {
642 aion_proto::generated::ActivityResult {
643 workflow_id: value.workflow_id.map(generated_workflow_id),
644 activity_id: value.activity_id.map(generated_activity_id),
645 run_id: value.run_id.map(generated_run_id),
646 completion_token: value.completion_token,
647 outcome: value.outcome.map(|outcome| match outcome {
648 proto_activity_result::Outcome::Result(result) => {
649 aion_proto::generated::activity_result::Outcome::Result(generated_payload(result))
650 }
651 proto_activity_result::Outcome::Error(error) => {
652 aion_proto::generated::activity_result::Outcome::Error(generated_error(error))
653 }
654 }),
655 }
656}
657
658fn generated_heartbeat(value: ProtoHeartbeat) -> aion_proto::generated::Heartbeat {
659 aion_proto::generated::Heartbeat {
660 workflow_id: value.workflow_id.map(generated_workflow_id),
661 activity_id: value.activity_id.map(generated_activity_id),
662 progress: value.progress.map(generated_payload),
663 }
664}
665
666fn proto_task(value: aion_proto::generated::ActivityTask) -> ProtoActivityTask {
667 ProtoActivityTask {
668 workflow_id: value.workflow_id.map(proto_workflow_id),
669 activity_id: value.activity_id.map(proto_activity_id),
670 activity_type: value.activity_type,
671 input: value.input.map(proto_payload),
672 attempt: value.attempt,
673 labels: value.labels,
674 run_id: value.run_id.map(proto_run_id),
675 completion_token: value.completion_token,
676 idempotency_key: value.idempotency_key,
677 }
678}
679
680fn generated_payload(value: ProtoPayload) -> aion_proto::generated::Payload {
681 aion_proto::generated::Payload {
682 content_type: value.content_type,
683 bytes: value.bytes,
684 }
685}
686
687fn proto_payload(value: aion_proto::generated::Payload) -> ProtoPayload {
688 ProtoPayload {
689 content_type: value.content_type,
690 bytes: value.bytes,
691 }
692}
693
694fn generated_workflow_id(value: ProtoWorkflowId) -> aion_proto::generated::WorkflowId {
695 aion_proto::generated::WorkflowId { uuid: value.uuid }
696}
697
698fn proto_workflow_id(value: aion_proto::generated::WorkflowId) -> ProtoWorkflowId {
699 ProtoWorkflowId { uuid: value.uuid }
700}
701
702fn generated_run_id(value: ProtoRunId) -> aion_proto::generated::RunId {
703 aion_proto::generated::RunId { uuid: value.uuid }
704}
705
706fn proto_run_id(value: aion_proto::generated::RunId) -> ProtoRunId {
707 ProtoRunId { uuid: value.uuid }
708}
709
710fn generated_activity_id(value: ProtoActivityId) -> aion_proto::generated::ActivityId {
711 aion_proto::generated::ActivityId {
712 sequence_position: value.sequence_position,
713 }
714}
715
716fn proto_activity_id(value: aion_proto::generated::ActivityId) -> ProtoActivityId {
717 ProtoActivityId {
718 sequence_position: value.sequence_position,
719 }
720}
721
722fn generated_error(value: aion_proto::ProtoActivityError) -> aion_proto::generated::ActivityError {
723 aion_proto::generated::ActivityError {
724 kind: value.kind,
725 message: value.message,
726 details: value.details.map(generated_payload),
727 }
728}
729
730#[derive(thiserror::Error, Debug)]
731#[error("{message}")]
732struct SessionStateError {
733 message: String,
734}
735
736#[cfg(test)]
737mod tests {
738 use std::collections::BTreeSet;
739
740 use aion_proto::ProtoActivityTask;
741 use async_trait::async_trait;
742 use futures::{StreamExt, stream};
743
744 use super::{
745 WorkerSession, WorkerSessionEvent, WorkerTaskStream, apply_auth_metadata,
746 validate_activity_handlers,
747 };
748 use crate::error::WorkerError;
749 use crate::{ReconnectConfig, WorkerConfig};
750
751 #[derive(Default)]
752 struct FakeSession {
753 handshakes: Vec<(String, String)>,
754 registrations: Vec<Vec<String>>,
755 }
756
757 #[async_trait]
758 impl WorkerSession for FakeSession {
759 async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
760 self.handshakes
761 .push((config.task_queue.clone(), config.identity.clone()));
762 Ok(())
763 }
764
765 async fn register(
766 &mut self,
767 activity_types: Vec<String>,
768 available_handlers: &BTreeSet<String>,
769 ) -> Result<(), WorkerError> {
770 validate_activity_handlers(&activity_types, available_handlers)?;
771 self.registrations.push(activity_types);
772 Ok(())
773 }
774
775 fn receive_tasks(&mut self) -> WorkerTaskStream {
776 Box::pin(stream::iter([Ok(WorkerSessionEvent::Task(Box::new(
777 ProtoActivityTask {
778 workflow_id: None,
779 activity_id: None,
780 activity_type: String::from("charge-card"),
781 input: None,
782 attempt: 1,
783 labels: std::collections::HashMap::new(),
784 run_id: Some(aion_proto::ProtoRunId::from(aion_core::RunId::new_v4())),
785 completion_token: String::from("generation-1"),
786 idempotency_key: String::from("effect-key"),
787 },
788 )))]))
789 }
790
791 async fn report_result(
792 &mut self,
793 workflow_id: aion_core::WorkflowId,
794 activity_id: aion_core::ActivityId,
795 run_id: Option<aion_core::RunId>,
796 completion_token: String,
797 result: aion_core::Payload,
798 ) -> Result<(), WorkerError> {
799 drop((workflow_id, activity_id, run_id, completion_token, result));
800 Ok(())
801 }
802
803 async fn report_failure(
804 &mut self,
805 workflow_id: aion_core::WorkflowId,
806 activity_id: aion_core::ActivityId,
807 run_id: Option<aion_core::RunId>,
808 completion_token: String,
809 failure: aion_core::ActivityError,
810 ) -> Result<(), WorkerError> {
811 drop((workflow_id, activity_id, run_id, completion_token, failure));
812 Ok(())
813 }
814
815 async fn send_heartbeat(
816 &mut self,
817 workflow_id: aion_core::WorkflowId,
818 activity_id: aion_core::ActivityId,
819 progress: Option<aion_core::Payload>,
820 ) -> Result<(), WorkerError> {
821 drop((workflow_id, activity_id, progress));
822 Ok(())
823 }
824 }
825
826 #[test]
827 fn apply_auth_metadata_sets_worker_authorization_headers() -> Result<(), WorkerError> {
828 let config = WorkerConfig::builder()
829 .endpoint("http://127.0.0.1:50051")
830 .task_queue("payments")
831 .identity("worker-a")
832 .max_concurrency(4)
833 .reconnect_initial_backoff(std::time::Duration::from_millis(5))
834 .reconnect_max_backoff(std::time::Duration::from_millis(20))
835 .reconnect_max_attempts(3)
836 .namespace("payments")
837 .subject("worker-a")
838 .build()
839 .map_err(WorkerError::registration)?;
840 let mut metadata = tonic::metadata::MetadataMap::new();
841
842 apply_auth_metadata(&mut metadata, &config)?;
843
844 assert_eq!(
845 metadata
846 .get("x-aion-namespaces")
847 .and_then(|value| value.to_str().ok()),
848 Some("payments")
849 );
850 assert_eq!(
851 metadata
852 .get("x-aion-subject")
853 .and_then(|value| value.to_str().ok()),
854 Some("worker-a")
855 );
856 Ok(())
857 }
858
859 #[tokio::test]
860 async fn fake_session_records_handshake_and_registration() -> Result<(), WorkerError> {
861 let config = WorkerConfig::new(
862 "http://127.0.0.1:50051",
863 "payments",
864 "worker-a",
865 4,
866 ReconnectConfig::new(
867 std::time::Duration::from_millis(5),
868 std::time::Duration::from_millis(20),
869 3,
870 ),
871 None,
872 );
873 let activity_types = vec![String::from("charge-card"), String::from("send-email")];
874 let handlers = activity_types.iter().cloned().collect::<BTreeSet<_>>();
875 let mut session = FakeSession::default();
876
877 session.handshake(&config).await?;
878 session.register(activity_types.clone(), &handlers).await?;
879 let received = session.receive_tasks().next().await;
880
881 assert_eq!(
882 session.handshakes,
883 vec![(String::from("payments"), String::from("worker-a"))]
884 );
885 assert_eq!(session.registrations, vec![activity_types]);
886 assert!(received.is_some());
887
888 Ok(())
889 }
890
891 #[tokio::test]
892 async fn grpc_reports_echo_the_dispatched_completion_token() -> Result<(), WorkerError> {
893 let config = WorkerConfig::new(
894 "http://127.0.0.1:50051",
895 "payments",
896 "worker-a",
897 1,
898 ReconnectConfig::new(
899 std::time::Duration::from_millis(5),
900 std::time::Duration::from_millis(20),
901 3,
902 ),
903 None,
904 );
905 let (sender, mut receiver) = tokio::sync::mpsc::channel(2);
906 let mut session = super::GrpcWorkerSession {
907 config,
908 activity_types: Vec::new(),
909 client: None,
910 sender: Some(sender),
911 receiver: None,
912 registered_info: None,
913 };
914
915 session
916 .report_result(
917 aion_core::WorkflowId::new_v4(),
918 aion_core::ActivityId::from_sequence_position(1),
919 Some(aion_core::RunId::new_v4()),
920 String::from("success-generation"),
921 aion_core::Payload::new(aion_core::ContentType::Json, b"{}".to_vec()),
922 )
923 .await?;
924 session
925 .report_failure(
926 aion_core::WorkflowId::new_v4(),
927 aion_core::ActivityId::from_sequence_position(2),
928 Some(aion_core::RunId::new_v4()),
929 String::from("failure-generation"),
930 aion_core::ActivityError {
931 kind: aion_core::ActivityErrorKind::Terminal,
932 message: String::from("failed"),
933 details: None,
934 },
935 )
936 .await?;
937
938 let success = receiver.recv().await.ok_or_else(|| {
939 WorkerError::decode(super::SessionStateError {
940 message: String::from("result report channel closed"),
941 })
942 })?;
943 let failure = receiver.recv().await.ok_or_else(|| {
944 WorkerError::decode(super::SessionStateError {
945 message: String::from("failure report channel closed"),
946 })
947 })?;
948 let success_token = match success.message {
949 Some(aion_proto::generated::worker_to_server::Message::Result(result)) => {
950 result.completion_token
951 }
952 _ => {
953 return Err(WorkerError::decode(super::SessionStateError {
954 message: String::from("success report did not emit an ActivityResult"),
955 }));
956 }
957 };
958 let failure_token = match failure.message {
959 Some(aion_proto::generated::worker_to_server::Message::Result(result)) => {
960 result.completion_token
961 }
962 _ => {
963 return Err(WorkerError::decode(super::SessionStateError {
964 message: String::from("failure report did not emit an ActivityResult"),
965 }));
966 }
967 };
968 assert_eq!(success_token, "success-generation");
969 assert_eq!(failure_token, "failure-generation");
970 Ok(())
971 }
972
973 #[tokio::test(start_paused = true)]
977 async fn report_send_times_out_retryably_at_max_backoff() -> Result<(), WorkerError> {
978 let config = WorkerConfig::new(
979 "http://127.0.0.1:50051",
980 "payments",
981 "worker-a",
982 1,
983 ReconnectConfig::new(
984 std::time::Duration::from_millis(5),
985 std::time::Duration::from_millis(20),
986 3,
987 ),
988 None,
989 );
990 let (sender, receiver) = tokio::sync::mpsc::channel(1);
991 sender
994 .try_send(aion_proto::generated::WorkerToServer { message: None })
995 .map_err(WorkerError::decode)?;
996 let mut session = super::GrpcWorkerSession {
997 config,
998 activity_types: Vec::new(),
999 client: None,
1000 sender: Some(sender),
1001 receiver: None,
1002 registered_info: None,
1003 };
1004
1005 let result = session
1006 .report_result(
1007 aion_core::WorkflowId::new_v4(),
1008 aion_core::ActivityId::from_sequence_position(1),
1009 Some(aion_core::RunId::new_v4()),
1010 String::from("generation-1"),
1011 aion_core::Payload::new(aion_core::ContentType::Json, b"{}".to_vec()),
1012 )
1013 .await;
1014
1015 let Err(error) = result else {
1016 return Err(WorkerError::Transport {
1017 source: tonic::Status::internal("a hung send must time out, not hang"),
1018 });
1019 };
1020 assert!(
1021 matches!(error, WorkerError::Transport { .. }),
1022 "send deadline elapse must be a retryable transport error: {error}"
1023 );
1024 assert!(error.is_retryable());
1025 assert!(
1026 error.to_string().contains("did not complete"),
1027 "the error must name the deadline: {error}"
1028 );
1029 drop(receiver);
1030 Ok(())
1031 }
1032
1033 #[test]
1034 fn registration_rejects_activity_without_handler() {
1035 let activity_types = vec![String::from("charge-card"), String::from("send-email")];
1036 let handlers = [String::from("charge-card")]
1037 .into_iter()
1038 .collect::<BTreeSet<_>>();
1039
1040 let result = validate_activity_handlers(&activity_types, &handlers);
1041 assert!(result.is_err());
1042 let error = match result {
1043 Ok(()) => return,
1044 Err(error) => error,
1045 };
1046
1047 assert_eq!(
1048 error.to_string(),
1049 "worker registration failed: activity type `send-email` has no registered handler"
1050 );
1051 }
1052}