Skip to main content

aion_worker/protocol/
session.rs

1//! `WorkerSession` trait and gRPC-backed implementation.
2
3use 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
22/// Boxed receive stream returned by worker sessions.
23pub type WorkerTaskStream =
24    Pin<Box<dyn Stream<Item = Result<WorkerSessionEvent, WorkerError>> + Send>>;
25
26/// Event pushed by the worker session receive stream.
27#[derive(Clone, Debug, PartialEq, Eq)]
28pub enum WorkerSessionEvent {
29    /// A new activity task to execute.
30    Task(ProtoActivityTask),
31    /// Server-initiated drain: the server is going away (restart, deploy,
32    /// rebalance). The worker finishes in-flight work, reports what it can,
33    /// stops expecting new tasks, and reconnects after the schedule's initial
34    /// backoff. A drain frame latches for the session: the eventual stream
35    /// end — clean or abrupt — is drain-class and consumes no drop budget.
36    Drain,
37    /// The server consumed the identified `ActivityResult` frame; the worker
38    /// may stop re-reporting it. Clears the matching unacked-tracker entry.
39    ResultAck {
40        /// Workflow owning the acknowledged result.
41        workflow_id: WorkflowId,
42        /// Activity whose result was acknowledged.
43        activity_id: ActivityId,
44    },
45    /// Cooperative cancellation for an in-flight activity.
46    ///
47    /// The current AW worker proto in this worktree does not yet carry this
48    /// frame, but fake sessions can emit it and the runtime handles it without
49    /// forcing task termination. When AW lands the wire variant,
50    /// `decode_server_message` should map it to this event.
51    Cancel {
52        /// Workflow owning the activity.
53        workflow_id: WorkflowId,
54        /// Activity to mark cancelled.
55        activity_id: ActivityId,
56    },
57}
58
59/// Transport abstraction for the AW-owned worker protocol.
60///
61/// The current `aion-proto` worker endpoint is `WorkerProtocol::StreamWorker`,
62/// a single bidirectional gRPC stream. These methods intentionally present the
63/// worker conversation as handshake/register/receive/report/heartbeat phases so
64/// execution machinery can be tested against fakes and never touches generated
65/// stubs directly. If AW changes the wire shape, this trait adapts in this module.
66#[async_trait]
67pub trait WorkerSession: Send {
68    /// Performs the worker handshake for the configured namespace, task queue,
69    /// and identity.
70    ///
71    /// Maps to transport/channel establishment for AW's `StreamWorker` RPC. The
72    /// wire carries the genuine `namespace` (correctness boundary) and
73    /// `task_queue` (pool selector) as disjoint registration fields; it has no
74    /// identity field, so identity is retained at this SDK boundary until the
75    /// wire adds a corresponding shape.
76    async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError>;
77
78    /// Registers activity-type names implemented by this worker.
79    ///
80    /// Maps to opening AW's `StreamWorker` RPC with `RegisterWorker` queued as
81    /// the mandatory first frame and then awaiting the server's `RegisterAck`
82    /// — the guaranteed first frame on the response stream. Registration
83    /// succeeds only when the ack arrives; a denial fails the RPC with a gRPC
84    /// error status (`PermissionDenied` / `Unauthenticated`), and an ack that
85    /// does not arrive within the reconnect policy's `max_backoff` is a
86    /// retryable registration failure. The caller supplies
87    /// `available_handlers` so registration can be rejected before serving if
88    /// any requested name lacks a handler.
89    async fn register(
90        &mut self,
91        activity_types: Vec<String>,
92        available_handlers: &BTreeSet<String>,
93    ) -> Result<(), WorkerError>;
94
95    /// Opens the receive side of AW's `StreamWorker` RPC and yields pushed tasks.
96    fn receive_tasks(&mut self) -> WorkerTaskStream;
97
98    /// Reports successful activity output via `WorkerToServer.result`.
99    async fn report_result(
100        &mut self,
101        workflow_id: WorkflowId,
102        activity_id: ActivityId,
103        run_id: Option<RunId>,
104        result: Payload,
105    ) -> Result<(), WorkerError>;
106
107    /// Reports explicit activity failure via `WorkerToServer.result`.
108    async fn report_failure(
109        &mut self,
110        workflow_id: WorkflowId,
111        activity_id: ActivityId,
112        run_id: Option<RunId>,
113        failure: ActivityError,
114    ) -> Result<(), WorkerError>;
115
116    /// Sends cooperative progress via `WorkerToServer.heartbeat`.
117    async fn send_heartbeat(
118        &mut self,
119        workflow_id: WorkflowId,
120        activity_id: ActivityId,
121        progress: Option<Payload>,
122    ) -> Result<(), WorkerError>;
123
124    /// Server-assigned liveness window from the `RegisterAck`, once registered.
125    ///
126    /// The serve loop derives its AUTOMATIC liveness-heartbeat cadence from
127    /// this window (see `serve_activity_tasks_until`): the server's heartbeat
128    /// sweeper expires any worker whose in-flight task goes longer than the
129    /// window without a heartbeat, so the runtime — not each handler — must
130    /// keep every in-flight activity beating. `None` (the default, and the
131    /// value for fake/unregistered sessions) disables the automatic pump.
132    fn heartbeat_window(&self) -> Option<std::time::Duration> {
133        None
134    }
135}
136
137/// Validates that every requested activity type has a registered handler.
138///
139/// # Errors
140///
141/// Returns [`WorkerError::Registration`] for the first missing handler name.
142pub fn validate_activity_handlers(
143    activity_types: &[String],
144    available_handlers: &BTreeSet<String>,
145) -> Result<(), WorkerError> {
146    if let Some(activity_type) = activity_types
147        .iter()
148        .find(|activity_type| !available_handlers.contains(*activity_type))
149    {
150        return Err(WorkerError::registration(MissingActivityHandler {
151            activity_type: activity_type.clone(),
152        }));
153    }
154
155    Ok(())
156}
157
158/// Server-assigned registration facts carried by the `RegisterAck` frame.
159#[derive(Clone, Debug, PartialEq, Eq)]
160pub struct RegisteredSessionInfo {
161    /// Server-assigned stream identifier, for correlating worker logs with
162    /// server logs (`worker_id=3 lost`).
163    pub worker_id: u64,
164    /// The namespace the registration was authorized against.
165    pub namespace: String,
166    /// The server's operator-configured liveness window: an in-flight
167    /// activity must heartbeat at least this often or be declared lost.
168    pub heartbeat_window: std::time::Duration,
169}
170
171/// gRPC-backed [`WorkerSession`] using `aion-proto` generated tonic stubs.
172pub struct GrpcWorkerSession {
173    config: WorkerConfig,
174    activity_types: Vec<String>,
175    client: Option<GeneratedClient>,
176    sender: Option<mpsc::Sender<aion_proto::generated::WorkerToServer>>,
177    receiver: Option<tonic::codec::Streaming<aion_proto::generated::ServerToWorker>>,
178    registered_info: Option<RegisteredSessionInfo>,
179}
180
181impl GrpcWorkerSession {
182    /// Connects to the configured worker endpoint.
183    ///
184    /// Opaque credentials are accepted by [`WorkerConfig`] but the current AW
185    /// worker proto does not define a credential metadata convention, so no
186    /// authentication scheme is interpreted here.
187    ///
188    /// # Errors
189    ///
190    /// Returns [`WorkerError::Connect`] if tonic cannot create the channel.
191    pub async fn connect(config: WorkerConfig) -> Result<Self, WorkerError> {
192        let client = GeneratedClient::connect(config.endpoint.clone())
193            .await
194            .map_err(|source| WorkerError::Connect { source })?;
195
196        Ok(Self {
197            config,
198            activity_types: Vec::new(),
199            client: Some(client),
200            sender: None,
201            receiver: None,
202            registered_info: None,
203        })
204    }
205
206    /// Creates a session from an existing tonic channel.
207    #[must_use]
208    pub fn from_channel(config: WorkerConfig, channel: Channel) -> Self {
209        Self {
210            config,
211            activity_types: Vec::new(),
212            client: Some(GeneratedClient::new(channel)),
213            sender: None,
214            receiver: None,
215            registered_info: None,
216        }
217    }
218
219    /// Server-assigned registration facts from the `RegisterAck`, available
220    /// once [`WorkerSession::register`] has succeeded.
221    #[must_use]
222    pub const fn registered_info(&self) -> Option<&RegisteredSessionInfo> {
223        self.registered_info.as_ref()
224    }
225
226    /// Opens AW's `StreamWorker` RPC with `RegisterWorker` queued as the first
227    /// outbound frame and awaits the server's `RegisterAck`.
228    ///
229    /// The server reads `RegisterWorker` from the inbound stream *before* it
230    /// returns its response stream (and therefore before tonic receives
231    /// response headers), so the frame must already be queued when the RPC is
232    /// issued. Awaiting `stream_worker` before sending `RegisterWorker`
233    /// deadlocks: the client waits for headers the server withholds until it
234    /// has read the registration.
235    ///
236    /// Registration succeeds only when the server's `RegisterAck` — its
237    /// guaranteed first response frame — arrives. The ack wait is bounded by
238    /// the reconnect policy's `max_backoff` (the operator's own definition of
239    /// the longest tolerable pause); a timeout, a non-ack first frame, or a
240    /// stream that ends before the ack is a retryable registration failure.
241    /// Denials surface as the RPC's gRPC error status exactly as before.
242    async fn open_registered_stream(
243        &mut self,
244        register: aion_proto::generated::RegisterWorker,
245    ) -> Result<(), WorkerError> {
246        let client = self.client.as_mut().ok_or_else(|| {
247            WorkerError::registration(SessionStateError {
248                message: String::from("worker session has not completed its handshake"),
249            })
250        })?;
251        let (sender, outbound) = mpsc::channel(16);
252        sender
253            .try_send(aion_proto::generated::WorkerToServer {
254                message: Some(aion_proto::generated::worker_to_server::Message::Register(
255                    register,
256                )),
257            })
258            .map_err(|_| {
259                WorkerError::registration(SessionStateError {
260                    message: String::from(
261                        "could not queue RegisterWorker as the first stream frame",
262                    ),
263                })
264            })?;
265        let mut request = Request::new(ReceiverStream::new(outbound));
266        apply_auth_metadata(request.metadata_mut(), &self.config)?;
267        let response = client
268            .stream_worker(request)
269            .await
270            .map_err(registration_denial_error)?;
271        let mut receiver = response.into_inner();
272
273        let first = tokio::time::timeout(self.config.reconnect.max_backoff, receiver.message())
274            .await
275            .map_err(|_| {
276                WorkerError::registration(SessionStateError {
277                    message: format!(
278                        "server did not acknowledge registration within {:?}",
279                        self.config.reconnect.max_backoff
280                    ),
281                })
282            })?
283            .map_err(registration_denial_error)?;
284        let ack = match first.and_then(|frame| frame.message) {
285            Some(aion_proto::generated::server_to_worker::Message::RegisterAck(ack)) => ack,
286            Some(_) => {
287                return Err(WorkerError::decode(SessionStateError {
288                    message: String::from(
289                        "protocol violation: server sent a non-RegisterAck frame before \
290                         acknowledging registration",
291                    ),
292                }));
293            }
294            None => {
295                return Err(WorkerError::registration(SessionStateError {
296                    message: String::from(
297                        "server ended the stream before acknowledging registration",
298                    ),
299                }));
300            }
301        };
302
303        self.registered_info = Some(RegisteredSessionInfo {
304            worker_id: ack.worker_id,
305            namespace: ack.namespace,
306            heartbeat_window: std::time::Duration::from_millis(ack.heartbeat_window_ms),
307        });
308        self.sender = Some(sender);
309        self.receiver = Some(receiver);
310        Ok(())
311    }
312
313    /// Sends one frame with a per-send deadline of the reconnect policy's
314    /// `max_backoff`: a send that outlives the operator's longest tolerable
315    /// pause is, by that same definition, a dead session and surfaces as a
316    /// retryable transport error instead of hanging the worker forever.
317    async fn send_to_server(
318        &self,
319        message: aion_proto::generated::worker_to_server::Message,
320    ) -> Result<(), WorkerError> {
321        let sender = self.sender.as_ref().ok_or_else(|| {
322            WorkerError::registration(SessionStateError {
323                message: String::from("worker stream has not been opened"),
324            })
325        })?;
326        let send = sender.send(aion_proto::generated::WorkerToServer {
327            message: Some(message),
328        });
329        tokio::time::timeout(self.config.reconnect.max_backoff, send)
330            .await
331            .map_err(|_| WorkerError::Transport {
332                source: tonic::Status::unavailable(format!(
333                    "worker stream send did not complete within {:?}",
334                    self.config.reconnect.max_backoff
335                )),
336            })?
337            .map_err(|source| WorkerError::Transport {
338                source: tonic::Status::unavailable(format!("worker stream send failed: {source}")),
339            })
340    }
341}
342
343/// Maps the `StreamWorker` RPC's rejection status to the worker error taxonomy.
344///
345/// The server validates stream metadata (credentials) and the `RegisterWorker`
346/// frame before returning response headers, so both failure classes surface
347/// from the same await: `Unauthenticated` is a credential/handshake rejection,
348/// everything else is a registration outcome (`PermissionDenied` for an
349/// ungranted namespace, `Unavailable` for transient transport faults). Both
350/// shapes preserve the status for `WorkerError::grpc_status` / `is_retryable`.
351fn registration_denial_error(status: tonic::Status) -> WorkerError {
352    if status.code() == tonic::Code::Unauthenticated {
353        WorkerError::Handshake { source: status }
354    } else {
355        WorkerError::Registration {
356            source: Box::new(status),
357        }
358    }
359}
360
361fn apply_auth_metadata(
362    metadata: &mut tonic::metadata::MetadataMap,
363    config: &WorkerConfig,
364) -> Result<(), WorkerError> {
365    // The `x-aion-namespaces` metadata reflects the worker's full namespace SET,
366    // comma-joined in advertised order. The server authorizes the worker for
367    // every namespace it registers under.
368    let namespaces_value = config.namespaces.join(",");
369    let namespace =
370        MetadataValue::try_from(namespaces_value.as_str()).map_err(|_| WorkerError::Handshake {
371            source: tonic::Status::invalid_argument(
372                "worker namespaces are not valid gRPC metadata",
373            ),
374        })?;
375    let subject =
376        MetadataValue::try_from(config.subject.as_str()).map_err(|_| WorkerError::Handshake {
377            source: tonic::Status::invalid_argument("worker subject is not valid gRPC metadata"),
378        })?;
379    metadata.insert("x-aion-namespaces", namespace);
380    metadata.insert("x-aion-subject", subject);
381    Ok(())
382}
383
384#[async_trait]
385impl WorkerSession for GrpcWorkerSession {
386    async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
387        self.config = config.clone();
388        if self.client.is_none() {
389            self.client = Some(
390                GeneratedClient::connect(self.config.endpoint.clone())
391                    .await
392                    .map_err(|source| WorkerError::Connect { source })?,
393            );
394        }
395        Ok(())
396    }
397
398    async fn register(
399        &mut self,
400        activity_types: Vec<String>,
401        available_handlers: &BTreeSet<String>,
402    ) -> Result<(), WorkerError> {
403        validate_activity_handlers(&activity_types, available_handlers)?;
404        self.activity_types.clone_from(&activity_types);
405
406        // OQ-5: the registration namespace SET is the SAME set advertised in the
407        // `x-aion-namespaces` auth metadata (`apply_auth_metadata`), so a worker
408        // registers into exactly the namespaces it is authorized for.
409        // `task_queue` is the disjoint pool/flavour selector within each
410        // namespace; `node` is the optional locality affinity (default hostname).
411        let register = aion_proto::generated::RegisterWorker {
412            namespaces: self.config.namespaces.clone(),
413            activity_types,
414            task_queue: self.config.task_queue.clone(),
415            node: self.config.node.clone(),
416        };
417        self.open_registered_stream(register).await
418    }
419
420    fn receive_tasks(&mut self) -> WorkerTaskStream {
421        match self.receiver.take() {
422            Some(receiver) => Box::pin(receiver.filter_map(|message| async move {
423                Some(match message {
424                    Ok(server_message) => decode_server_message(server_message),
425                    Err(source) => Err(WorkerError::Transport { source }),
426                })
427            })),
428            None => Box::pin(futures::stream::iter([Err(WorkerError::Transport {
429                source: tonic::Status::failed_precondition(
430                    "worker receive stream has not been opened",
431                ),
432            })])),
433        }
434    }
435
436    async fn report_result(
437        &mut self,
438        workflow_id: WorkflowId,
439        activity_id: ActivityId,
440        run_id: Option<RunId>,
441        result: Payload,
442    ) -> Result<(), WorkerError> {
443        let result = ProtoActivityResult {
444            workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
445            activity_id: Some(ProtoActivityId::from(activity_id)),
446            run_id: run_id.map(ProtoRunId::from),
447            outcome: Some(proto_activity_result::Outcome::Result(ProtoPayload::from(
448                result,
449            ))),
450        };
451        self.send_to_server(aion_proto::generated::worker_to_server::Message::Result(
452            generated_activity_result(result),
453        ))
454        .await
455    }
456
457    async fn report_failure(
458        &mut self,
459        workflow_id: WorkflowId,
460        activity_id: ActivityId,
461        run_id: Option<RunId>,
462        failure: ActivityError,
463    ) -> Result<(), WorkerError> {
464        let result = ProtoActivityResult {
465            workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
466            activity_id: Some(ProtoActivityId::from(activity_id)),
467            run_id: run_id.map(ProtoRunId::from),
468            outcome: Some(proto_activity_result::Outcome::Error(failure.into())),
469        };
470        self.send_to_server(aion_proto::generated::worker_to_server::Message::Result(
471            generated_activity_result(result),
472        ))
473        .await
474    }
475
476    async fn send_heartbeat(
477        &mut self,
478        workflow_id: WorkflowId,
479        activity_id: ActivityId,
480        progress: Option<Payload>,
481    ) -> Result<(), WorkerError> {
482        let heartbeat = ProtoHeartbeat {
483            workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
484            activity_id: Some(ProtoActivityId::from(activity_id)),
485            progress: progress.map(ProtoPayload::from),
486        };
487        self.send_to_server(aion_proto::generated::worker_to_server::Message::Heartbeat(
488            generated_heartbeat(heartbeat),
489        ))
490        .await
491    }
492
493    fn heartbeat_window(&self) -> Option<std::time::Duration> {
494        self.registered_info
495            .as_ref()
496            .map(|info| info.heartbeat_window)
497    }
498}
499
500fn decode_server_message(
501    message: aion_proto::generated::ServerToWorker,
502) -> Result<WorkerSessionEvent, WorkerError> {
503    match message.message {
504        Some(aion_proto::generated::server_to_worker::Message::Task(task)) => {
505            Ok(WorkerSessionEvent::Task(proto_task(task)))
506        }
507        Some(aion_proto::generated::server_to_worker::Message::Drain(_)) => {
508            Ok(WorkerSessionEvent::Drain)
509        }
510        Some(aion_proto::generated::server_to_worker::Message::ResultAck(ack)) => {
511            decode_result_ack(ack)
512        }
513        Some(aion_proto::generated::server_to_worker::Message::RegisterAck(_)) => {
514            // The ack is consumed inside `open_registered_stream`; a second
515            // one mid-stream is a server ordering bug that must surface.
516            Err(WorkerError::decode(SessionStateError {
517                message: String::from(
518                    "protocol violation: RegisterAck received after registration completed",
519                ),
520            }))
521        }
522        None => Err(WorkerError::decode(SessionStateError {
523            message: String::from("server-to-worker message was empty"),
524        })),
525    }
526}
527
528fn decode_result_ack(
529    ack: aion_proto::generated::ResultAck,
530) -> Result<WorkerSessionEvent, WorkerError> {
531    let workflow_id = ack
532        .workflow_id
533        .ok_or_else(|| {
534            WorkerError::decode(SessionStateError {
535                message: String::from("result ack workflow_id is missing"),
536            })
537        })
538        .and_then(|id| {
539            WorkflowId::try_from(ProtoWorkflowId { uuid: id.uuid }).map_err(|source| {
540                WorkerError::decode(SessionStateError {
541                    message: format!("result ack workflow_id is invalid: {source}"),
542                })
543            })
544        })?;
545    let activity_id = ack
546        .activity_id
547        .map(|id| ActivityId::from_sequence_position(id.sequence_position))
548        .ok_or_else(|| {
549            WorkerError::decode(SessionStateError {
550                message: String::from("result ack activity_id is missing"),
551            })
552        })?;
553    Ok(WorkerSessionEvent::ResultAck {
554        workflow_id,
555        activity_id,
556    })
557}
558
559fn generated_activity_result(value: ProtoActivityResult) -> aion_proto::generated::ActivityResult {
560    aion_proto::generated::ActivityResult {
561        workflow_id: value.workflow_id.map(generated_workflow_id),
562        activity_id: value.activity_id.map(generated_activity_id),
563        run_id: value.run_id.map(generated_run_id),
564        outcome: value.outcome.map(|outcome| match outcome {
565            proto_activity_result::Outcome::Result(result) => {
566                aion_proto::generated::activity_result::Outcome::Result(generated_payload(result))
567            }
568            proto_activity_result::Outcome::Error(error) => {
569                aion_proto::generated::activity_result::Outcome::Error(generated_error(error))
570            }
571        }),
572    }
573}
574
575fn generated_heartbeat(value: ProtoHeartbeat) -> aion_proto::generated::Heartbeat {
576    aion_proto::generated::Heartbeat {
577        workflow_id: value.workflow_id.map(generated_workflow_id),
578        activity_id: value.activity_id.map(generated_activity_id),
579        progress: value.progress.map(generated_payload),
580    }
581}
582
583fn proto_task(value: aion_proto::generated::ActivityTask) -> ProtoActivityTask {
584    ProtoActivityTask {
585        workflow_id: value.workflow_id.map(proto_workflow_id),
586        activity_id: value.activity_id.map(proto_activity_id),
587        activity_type: value.activity_type,
588        input: value.input.map(proto_payload),
589        attempt: value.attempt,
590        labels: value.labels,
591        run_id: value.run_id.map(proto_run_id),
592    }
593}
594
595fn generated_payload(value: ProtoPayload) -> aion_proto::generated::Payload {
596    aion_proto::generated::Payload {
597        content_type: value.content_type,
598        bytes: value.bytes,
599    }
600}
601
602fn proto_payload(value: aion_proto::generated::Payload) -> ProtoPayload {
603    ProtoPayload {
604        content_type: value.content_type,
605        bytes: value.bytes,
606    }
607}
608
609fn generated_workflow_id(value: ProtoWorkflowId) -> aion_proto::generated::WorkflowId {
610    aion_proto::generated::WorkflowId { uuid: value.uuid }
611}
612
613fn proto_workflow_id(value: aion_proto::generated::WorkflowId) -> ProtoWorkflowId {
614    ProtoWorkflowId { uuid: value.uuid }
615}
616
617fn generated_run_id(value: ProtoRunId) -> aion_proto::generated::RunId {
618    aion_proto::generated::RunId { uuid: value.uuid }
619}
620
621fn proto_run_id(value: aion_proto::generated::RunId) -> ProtoRunId {
622    ProtoRunId { uuid: value.uuid }
623}
624
625fn generated_activity_id(value: ProtoActivityId) -> aion_proto::generated::ActivityId {
626    aion_proto::generated::ActivityId {
627        sequence_position: value.sequence_position,
628    }
629}
630
631fn proto_activity_id(value: aion_proto::generated::ActivityId) -> ProtoActivityId {
632    ProtoActivityId {
633        sequence_position: value.sequence_position,
634    }
635}
636
637fn generated_error(value: aion_proto::ProtoActivityError) -> aion_proto::generated::ActivityError {
638    aion_proto::generated::ActivityError {
639        kind: value.kind,
640        message: value.message,
641        details: value.details.map(generated_payload),
642    }
643}
644
645#[derive(thiserror::Error, Debug)]
646#[error("{message}")]
647struct SessionStateError {
648    message: String,
649}
650
651#[cfg(test)]
652mod tests {
653    use std::collections::BTreeSet;
654
655    use aion_proto::ProtoActivityTask;
656    use async_trait::async_trait;
657    use futures::{StreamExt, stream};
658
659    use super::{
660        WorkerSession, WorkerSessionEvent, WorkerTaskStream, apply_auth_metadata,
661        validate_activity_handlers,
662    };
663    use crate::error::WorkerError;
664    use crate::{ReconnectConfig, WorkerConfig};
665
666    #[derive(Default)]
667    struct FakeSession {
668        handshakes: Vec<(String, String)>,
669        registrations: Vec<Vec<String>>,
670    }
671
672    #[async_trait]
673    impl WorkerSession for FakeSession {
674        async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
675            self.handshakes
676                .push((config.task_queue.clone(), config.identity.clone()));
677            Ok(())
678        }
679
680        async fn register(
681            &mut self,
682            activity_types: Vec<String>,
683            available_handlers: &BTreeSet<String>,
684        ) -> Result<(), WorkerError> {
685            validate_activity_handlers(&activity_types, available_handlers)?;
686            self.registrations.push(activity_types);
687            Ok(())
688        }
689
690        fn receive_tasks(&mut self) -> WorkerTaskStream {
691            Box::pin(stream::iter([Ok(WorkerSessionEvent::Task(
692                ProtoActivityTask {
693                    workflow_id: None,
694                    activity_id: None,
695                    activity_type: String::from("charge-card"),
696                    input: None,
697                    attempt: 1,
698                    labels: std::collections::HashMap::new(),
699                    run_id: None,
700                },
701            ))]))
702        }
703
704        async fn report_result(
705            &mut self,
706            workflow_id: aion_core::WorkflowId,
707            activity_id: aion_core::ActivityId,
708            run_id: Option<aion_core::RunId>,
709            result: aion_core::Payload,
710        ) -> Result<(), WorkerError> {
711            drop((workflow_id, activity_id, run_id, result));
712            Ok(())
713        }
714
715        async fn report_failure(
716            &mut self,
717            workflow_id: aion_core::WorkflowId,
718            activity_id: aion_core::ActivityId,
719            run_id: Option<aion_core::RunId>,
720            failure: aion_core::ActivityError,
721        ) -> Result<(), WorkerError> {
722            drop((workflow_id, activity_id, run_id, failure));
723            Ok(())
724        }
725
726        async fn send_heartbeat(
727            &mut self,
728            workflow_id: aion_core::WorkflowId,
729            activity_id: aion_core::ActivityId,
730            progress: Option<aion_core::Payload>,
731        ) -> Result<(), WorkerError> {
732            drop((workflow_id, activity_id, progress));
733            Ok(())
734        }
735    }
736
737    #[test]
738    fn apply_auth_metadata_sets_worker_authorization_headers() -> Result<(), WorkerError> {
739        let config = WorkerConfig::builder()
740            .endpoint("http://127.0.0.1:50051")
741            .task_queue("payments")
742            .identity("worker-a")
743            .max_concurrency(4)
744            .reconnect_initial_backoff(std::time::Duration::from_millis(5))
745            .reconnect_max_backoff(std::time::Duration::from_millis(20))
746            .reconnect_max_attempts(3)
747            .namespace("payments")
748            .subject("worker-a")
749            .build()
750            .map_err(WorkerError::registration)?;
751        let mut metadata = tonic::metadata::MetadataMap::new();
752
753        apply_auth_metadata(&mut metadata, &config)?;
754
755        assert_eq!(
756            metadata
757                .get("x-aion-namespaces")
758                .and_then(|value| value.to_str().ok()),
759            Some("payments")
760        );
761        assert_eq!(
762            metadata
763                .get("x-aion-subject")
764                .and_then(|value| value.to_str().ok()),
765            Some("worker-a")
766        );
767        Ok(())
768    }
769
770    #[tokio::test]
771    async fn fake_session_records_handshake_and_registration() -> Result<(), WorkerError> {
772        let config = WorkerConfig::new(
773            "http://127.0.0.1:50051",
774            "payments",
775            "worker-a",
776            4,
777            ReconnectConfig::new(
778                std::time::Duration::from_millis(5),
779                std::time::Duration::from_millis(20),
780                3,
781            ),
782            None,
783        );
784        let activity_types = vec![String::from("charge-card"), String::from("send-email")];
785        let handlers = activity_types.iter().cloned().collect::<BTreeSet<_>>();
786        let mut session = FakeSession::default();
787
788        session.handshake(&config).await?;
789        session.register(activity_types.clone(), &handlers).await?;
790        let received = session.receive_tasks().next().await;
791
792        assert_eq!(
793            session.handshakes,
794            vec![(String::from("payments"), String::from("worker-a"))]
795        );
796        assert_eq!(session.registrations, vec![activity_types]);
797        assert!(received.is_some());
798
799        Ok(())
800    }
801
802    /// Brief test 16: a report send that never completes (server stopped
803    /// reading; outbound channel full) times out retryably at the reconnect
804    /// policy's `max_backoff` on a paused clock — the worker never hangs.
805    #[tokio::test(start_paused = true)]
806    async fn report_send_times_out_retryably_at_max_backoff() -> Result<(), WorkerError> {
807        let config = WorkerConfig::new(
808            "http://127.0.0.1:50051",
809            "payments",
810            "worker-a",
811            1,
812            ReconnectConfig::new(
813                std::time::Duration::from_millis(5),
814                std::time::Duration::from_millis(20),
815                3,
816            ),
817            None,
818        );
819        let (sender, receiver) = tokio::sync::mpsc::channel(1);
820        // Fill the channel so the next send blocks forever, modelling a
821        // server that stopped draining its receive side.
822        sender
823            .try_send(aion_proto::generated::WorkerToServer { message: None })
824            .map_err(WorkerError::decode)?;
825        let mut session = super::GrpcWorkerSession {
826            config,
827            activity_types: Vec::new(),
828            client: None,
829            sender: Some(sender),
830            receiver: None,
831            registered_info: None,
832        };
833
834        let result = session
835            .report_result(
836                aion_core::WorkflowId::new_v4(),
837                aion_core::ActivityId::from_sequence_position(1),
838                None,
839                aion_core::Payload::new(aion_core::ContentType::Json, b"{}".to_vec()),
840            )
841            .await;
842
843        let Err(error) = result else {
844            return Err(WorkerError::Transport {
845                source: tonic::Status::internal("a hung send must time out, not hang"),
846            });
847        };
848        assert!(
849            matches!(error, WorkerError::Transport { .. }),
850            "send deadline elapse must be a retryable transport error: {error}"
851        );
852        assert!(error.is_retryable());
853        assert!(
854            error.to_string().contains("did not complete"),
855            "the error must name the deadline: {error}"
856        );
857        drop(receiver);
858        Ok(())
859    }
860
861    #[test]
862    fn registration_rejects_activity_without_handler() {
863        let activity_types = vec![String::from("charge-card"), String::from("send-email")];
864        let handlers = [String::from("charge-card")]
865            .into_iter()
866            .collect::<BTreeSet<_>>();
867
868        let result = validate_activity_handlers(&activity_types, &handlers);
869        assert!(result.is_err());
870        let error = match result {
871            Ok(()) => return,
872            Err(error) => error,
873        };
874
875        assert_eq!(
876            error.to_string(),
877            "worker registration failed: activity type `send-email` has no registered handler"
878        );
879    }
880}