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