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(Box<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    /// Registers activity names together with their committed wire schemas.
96    ///
97    /// Session implementations predating contract handshakes remain useful as
98    /// test and custom transport adapters: the default validates through their
99    /// legacy registration path. Production transports override this method and
100    /// carry `activities` on the wire.
101    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    /// Opens the receive side of AW's `StreamWorker` RPC and yields pushed tasks.
112    fn receive_tasks(&mut self) -> WorkerTaskStream;
113
114    /// Reports successful activity output and echoes its opaque execution
115    /// generation via `WorkerToServer.result`.
116    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    /// Reports explicit activity failure and echoes its opaque execution
126    /// generation via `WorkerToServer.result`.
127    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    /// Sends cooperative progress via `WorkerToServer.heartbeat`.
137    async fn send_heartbeat(
138        &mut self,
139        workflow_id: WorkflowId,
140        activity_id: ActivityId,
141        progress: Option<Payload>,
142    ) -> Result<(), WorkerError>;
143
144    /// Heartbeat the worker connection while it is idle.
145    ///
146    /// The gRPC transport encodes this on the existing heartbeat path with no
147    /// task identifiers. Transport fakes may retain the no-op default.
148    async fn send_connection_heartbeat(&mut self) -> Result<(), WorkerError> {
149        Ok(())
150    }
151
152    /// Server-assigned liveness window from the `RegisterAck`, once registered.
153    ///
154    /// The serve loop derives its AUTOMATIC liveness-heartbeat cadence from
155    /// this window (see `serve_activity_tasks_until`): the server's heartbeat
156    /// sweeper expires any worker whose in-flight task goes longer than the
157    /// window without a heartbeat, so the runtime — not each handler — must
158    /// keep every in-flight activity beating. `None` (the default, and the
159    /// value for fake/unregistered sessions) disables the automatic pump.
160    fn heartbeat_window(&self) -> Option<std::time::Duration> {
161        None
162    }
163}
164
165/// Validates that every requested activity type has a registered handler.
166///
167/// # Errors
168///
169/// Returns [`WorkerError::Registration`] for the first missing handler name.
170pub 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/// Server-assigned registration facts carried by the `RegisterAck` frame.
187#[derive(Clone, Debug, PartialEq, Eq)]
188pub struct RegisteredSessionInfo {
189    /// Server-assigned stream identifier, for correlating worker logs with
190    /// server logs (`worker_id=3 lost`).
191    pub worker_id: u64,
192    /// The namespace the registration was authorized against.
193    pub namespace: String,
194    /// The server's operator-configured liveness window: an in-flight
195    /// activity must heartbeat at least this often or be declared lost.
196    pub heartbeat_window: std::time::Duration,
197}
198
199/// gRPC-backed [`WorkerSession`] using `aion-proto` generated tonic stubs.
200pub 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    /// Connects to the configured worker endpoint.
211    ///
212    /// Opaque credentials are accepted by [`WorkerConfig`] but the current AW
213    /// worker proto does not define a credential metadata convention, so no
214    /// authentication scheme is interpreted here.
215    ///
216    /// # Errors
217    ///
218    /// Returns [`WorkerError::Connect`] if tonic cannot create the channel.
219    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    /// Creates a session from an existing tonic channel.
235    #[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    /// Server-assigned registration facts from the `RegisterAck`, available
248    /// once [`WorkerSession::register`] has succeeded.
249    #[must_use]
250    pub const fn registered_info(&self) -> Option<&RegisteredSessionInfo> {
251        self.registered_info.as_ref()
252    }
253
254    /// Opens AW's `StreamWorker` RPC with `RegisterWorker` queued as the first
255    /// outbound frame and awaits the server's `RegisterAck`.
256    ///
257    /// The server reads `RegisterWorker` from the inbound stream *before* it
258    /// returns its response stream (and therefore before tonic receives
259    /// response headers), so the frame must already be queued when the RPC is
260    /// issued. Awaiting `stream_worker` before sending `RegisterWorker`
261    /// deadlocks: the client waits for headers the server withholds until it
262    /// has read the registration.
263    ///
264    /// Registration succeeds only when the server's `RegisterAck` — its
265    /// guaranteed first response frame — arrives. The ack wait is bounded by
266    /// the reconnect policy's `max_backoff` (the operator's own definition of
267    /// the longest tolerable pause); a timeout, a non-ack first frame, or a
268    /// stream that ends before the ack is a retryable registration failure.
269    /// Denials surface as the RPC's gRPC error status exactly as before.
270    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    /// Sends one frame with a per-send deadline of the reconnect policy's
342    /// `max_backoff`: a send that outlives the operator's longest tolerable
343    /// pause is, by that same definition, a dead session and surfaces as a
344    /// retryable transport error instead of hanging the worker forever.
345    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
371/// Maps the `StreamWorker` RPC's rejection status to the worker error taxonomy.
372///
373/// The server validates stream metadata (credentials) and the `RegisterWorker`
374/// frame before returning response headers, so both failure classes surface
375/// from the same await: `Unauthenticated` is a credential/handshake rejection,
376/// everything else is a registration outcome (`PermissionDenied` for an
377/// ungranted namespace, `Unavailable` for transient transport faults). Both
378/// shapes preserve the status for `WorkerError::grpc_status` / `is_retryable`.
379fn 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    // The `x-aion-namespaces` metadata reflects the worker's full namespace SET,
394    // comma-joined in advertised order. The server authorizes the worker for
395    // every namespace it registers under.
396    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        // OQ-5: the registration namespace SET is the SAME set advertised in the
445        // `x-aion-namespaces` auth metadata (`apply_auth_metadata`), so a worker
446        // registers into exactly the namespaces it is authorized for.
447        // `task_queue` is the disjoint pool/flavour selector within each
448        // namespace; `node` is the optional locality affinity (default hostname).
449        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            // The ack is consumed inside `open_registered_stream`; a second
597            // one mid-stream is a server ordering bug that must surface.
598            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    /// Brief test 16: a report send that never completes (server stopped
974    /// reading; outbound channel full) times out retryably at the reconnect
975    /// policy's `max_backoff` on a paused clock — the worker never hangs.
976    #[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        // Fill the channel so the next send blocks forever, modelling a
992        // server that stopped draining its receive side.
993        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}