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    /// Transport liveness ping from the server (#197): answer it, promptly, or
58    /// this worker is not selected for dispatch.
59    ///
60    /// The RUNTIME answers it — never action code. The server is measuring
61    /// whether it can reach this worker's DISPATCH path, and an answer produced
62    /// by a handler would measure whether a handler happens to be running.
63    LivenessPing {
64        /// Sequence to echo back verbatim. An answer carrying anything else is
65        /// discarded by the server as a stale or fabricated echo.
66        sequence: u64,
67        /// How long this worker may hear NOTHING on this stream before it
68        /// should treat the link as dead — the server's own configured window,
69        /// carried on every ping so the worker holds no second copy of it.
70        silence_window: std::time::Duration,
71    },
72}
73
74/// Transport abstraction for the AW-owned worker protocol.
75///
76/// The current `aion-proto` worker endpoint is `WorkerProtocol::StreamWorker`,
77/// a single bidirectional gRPC stream. These methods intentionally present the
78/// worker conversation as handshake/register/receive/report/heartbeat phases so
79/// execution machinery can be tested against fakes and never touches generated
80/// stubs directly. If AW changes the wire shape, this trait adapts in this module.
81#[async_trait]
82pub trait WorkerSession: Send {
83    /// Performs the worker handshake for the configured namespace, task queue,
84    /// and identity.
85    ///
86    /// Maps to transport/channel establishment for AW's `StreamWorker` RPC. The
87    /// wire carries the genuine `namespace` (correctness boundary) and
88    /// `task_queue` (pool selector) as disjoint registration fields; it has no
89    /// identity field, so identity is retained at this SDK boundary until the
90    /// wire adds a corresponding shape.
91    async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError>;
92
93    /// Registers activity-type names implemented by this worker.
94    ///
95    /// Maps to opening AW's `StreamWorker` RPC with `RegisterWorker` queued as
96    /// the mandatory first frame and then awaiting the server's `RegisterAck`
97    /// — the guaranteed first frame on the response stream. Registration
98    /// succeeds only when the ack arrives; a denial fails the RPC with a gRPC
99    /// error status (`PermissionDenied` / `Unauthenticated`), and an ack that
100    /// does not arrive within the reconnect policy's `max_backoff` is a
101    /// retryable registration failure. The caller supplies
102    /// `available_handlers` so registration can be rejected before serving if
103    /// any requested name lacks a handler.
104    async fn register(
105        &mut self,
106        activity_types: Vec<String>,
107        available_handlers: &BTreeSet<String>,
108    ) -> Result<(), WorkerError>;
109
110    /// Registers activity names together with their committed wire schemas.
111    ///
112    /// Session implementations predating contract handshakes remain useful as
113    /// test and custom transport adapters: the default validates through their
114    /// legacy registration path. Production transports override this method and
115    /// carry `activities` on the wire.
116    async fn register_with_contract(
117        &mut self,
118        activity_types: Vec<String>,
119        activities: Vec<aion_package::ActivityDescriptor>,
120        available_handlers: &BTreeSet<String>,
121    ) -> Result<(), WorkerError> {
122        drop(activities);
123        self.register(activity_types, available_handlers).await
124    }
125
126    /// Opens the receive side of AW's `StreamWorker` RPC and yields pushed tasks.
127    fn receive_tasks(&mut self) -> WorkerTaskStream;
128
129    /// Reports successful activity output and echoes its opaque execution
130    /// generation via `WorkerToServer.result`.
131    async fn report_result(
132        &mut self,
133        workflow_id: WorkflowId,
134        activity_id: ActivityId,
135        run_id: Option<RunId>,
136        completion_token: String,
137        result: Payload,
138    ) -> Result<(), WorkerError>;
139
140    /// Reports explicit activity failure and echoes its opaque execution
141    /// generation via `WorkerToServer.result`.
142    async fn report_failure(
143        &mut self,
144        workflow_id: WorkflowId,
145        activity_id: ActivityId,
146        run_id: Option<RunId>,
147        completion_token: String,
148        failure: ActivityError,
149    ) -> Result<(), WorkerError>;
150
151    /// Sends cooperative progress via `WorkerToServer.heartbeat`.
152    async fn send_heartbeat(
153        &mut self,
154        workflow_id: WorkflowId,
155        activity_id: ActivityId,
156        progress: Option<Payload>,
157    ) -> Result<(), WorkerError>;
158
159    /// Heartbeat the worker connection while it is idle.
160    ///
161    /// The gRPC transport encodes this on the existing heartbeat path with no
162    /// task identifiers. Transport fakes may retain the no-op default.
163    async fn send_connection_heartbeat(&mut self) -> Result<(), WorkerError> {
164        Ok(())
165    }
166
167    /// Answer a server [`WorkerSessionEvent::LivenessPing`], echoing `sequence`
168    /// verbatim (#197).
169    ///
170    /// This is a RUNTIME obligation and the serve loop discharges it directly;
171    /// no handler is consulted and no concurrency permit is taken, because the
172    /// question is about the transport, not about the work.
173    ///
174    /// The default is a no-op for transports that carry no such frame (fakes,
175    /// and the liminal session, which answers its own ping/pong pair on its own
176    /// connection). That is not a silent failure: a session whose wire does
177    /// carry the ping and which does not answer it simply never clears its
178    /// dispatch probation, which is the honest verdict for a worker the server
179    /// cannot prove it reaches.
180    ///
181    /// # Errors
182    ///
183    /// Returns [`WorkerError::Transport`] when the answer cannot be sent.
184    async fn answer_liveness_ping(&mut self, sequence: u64) -> Result<(), WorkerError> {
185        tracing::debug!(
186            liveness_ping = sequence,
187            "this session's transport carries no liveness-answer frame, so nothing was sent. A \
188             server that probed over this transport will judge the ping unanswered — the honest \
189             verdict for a link this worker cannot answer on"
190        );
191        Ok(())
192    }
193
194    /// Server-assigned liveness window from the `RegisterAck`, once registered.
195    ///
196    /// The serve loop derives its AUTOMATIC liveness-heartbeat cadence from
197    /// this window (see `serve_activity_tasks_until`): the server's heartbeat
198    /// sweeper expires any worker whose in-flight task goes longer than the
199    /// window without a heartbeat, so the runtime — not each handler — must
200    /// keep every in-flight activity beating. `None` (the default, and the
201    /// value for fake/unregistered sessions) disables the automatic pump.
202    fn heartbeat_window(&self) -> Option<std::time::Duration> {
203        None
204    }
205}
206
207/// Validates that every requested activity type has a registered handler.
208///
209/// # Errors
210///
211/// Returns [`WorkerError::Registration`] for the first missing handler name.
212pub fn validate_activity_handlers(
213    activity_types: &[String],
214    available_handlers: &BTreeSet<String>,
215) -> Result<(), WorkerError> {
216    if let Some(activity_type) = activity_types
217        .iter()
218        .find(|activity_type| !available_handlers.contains(*activity_type))
219    {
220        return Err(WorkerError::registration(MissingActivityHandler {
221            activity_type: activity_type.clone(),
222        }));
223    }
224
225    Ok(())
226}
227
228/// Server-assigned registration facts carried by the `RegisterAck` frame.
229#[derive(Clone, Debug, PartialEq, Eq)]
230pub struct RegisteredSessionInfo {
231    /// Server-assigned stream identifier, for correlating worker logs with
232    /// server logs (`worker_id=3 lost`).
233    pub worker_id: u64,
234    /// The namespace the registration was authorized against.
235    pub namespace: String,
236    /// The server's operator-configured liveness window: an in-flight
237    /// activity must heartbeat at least this often or be declared lost.
238    pub heartbeat_window: std::time::Duration,
239}
240
241/// gRPC-backed [`WorkerSession`] using `aion-proto` generated tonic stubs.
242pub struct GrpcWorkerSession {
243    config: WorkerConfig,
244    activity_types: Vec<String>,
245    client: Option<GeneratedClient>,
246    sender: Option<mpsc::Sender<aion_proto::generated::WorkerToServer>>,
247    receiver: Option<tonic::codec::Streaming<aion_proto::generated::ServerToWorker>>,
248    registered_info: Option<RegisteredSessionInfo>,
249}
250
251impl GrpcWorkerSession {
252    /// Connects to the configured worker endpoint.
253    ///
254    /// Opaque credentials are accepted by [`WorkerConfig`] but the current AW
255    /// worker proto does not define a credential metadata convention, so no
256    /// authentication scheme is interpreted here.
257    ///
258    /// # Errors
259    ///
260    /// Returns [`WorkerError::Connect`] if tonic cannot create the channel.
261    pub async fn connect(config: WorkerConfig) -> Result<Self, WorkerError> {
262        let client = GeneratedClient::connect(config.endpoint.clone())
263            .await
264            .map_err(|source| WorkerError::Connect { source })?;
265
266        Ok(Self {
267            config,
268            activity_types: Vec::new(),
269            client: Some(client),
270            sender: None,
271            receiver: None,
272            registered_info: None,
273        })
274    }
275
276    /// Creates a session from an existing tonic channel.
277    #[must_use]
278    pub fn from_channel(config: WorkerConfig, channel: Channel) -> Self {
279        Self {
280            config,
281            activity_types: Vec::new(),
282            client: Some(GeneratedClient::new(channel)),
283            sender: None,
284            receiver: None,
285            registered_info: None,
286        }
287    }
288
289    /// Server-assigned registration facts from the `RegisterAck`, available
290    /// once [`WorkerSession::register`] has succeeded.
291    #[must_use]
292    pub const fn registered_info(&self) -> Option<&RegisteredSessionInfo> {
293        self.registered_info.as_ref()
294    }
295
296    /// Opens AW's `StreamWorker` RPC with `RegisterWorker` queued as the first
297    /// outbound frame and awaits the server's `RegisterAck`.
298    ///
299    /// The server reads `RegisterWorker` from the inbound stream *before* it
300    /// returns its response stream (and therefore before tonic receives
301    /// response headers), so the frame must already be queued when the RPC is
302    /// issued. Awaiting `stream_worker` before sending `RegisterWorker`
303    /// deadlocks: the client waits for headers the server withholds until it
304    /// has read the registration.
305    ///
306    /// Registration succeeds only when the server's `RegisterAck` — its
307    /// guaranteed first response frame — arrives. The ack wait is bounded by
308    /// the reconnect policy's `max_backoff` (the operator's own definition of
309    /// the longest tolerable pause); a timeout, a non-ack first frame, or a
310    /// stream that ends before the ack is a retryable registration failure.
311    /// Denials surface as the RPC's gRPC error status exactly as before.
312    async fn open_registered_stream(
313        &mut self,
314        register: aion_proto::generated::RegisterWorker,
315    ) -> Result<(), WorkerError> {
316        let client = self.client.as_mut().ok_or_else(|| {
317            WorkerError::registration(SessionStateError {
318                message: String::from("worker session has not completed its handshake"),
319            })
320        })?;
321        let (sender, outbound) = mpsc::channel(16);
322        sender
323            .try_send(aion_proto::generated::WorkerToServer {
324                message: Some(aion_proto::generated::worker_to_server::Message::Register(
325                    register,
326                )),
327            })
328            .map_err(|_| {
329                WorkerError::registration(SessionStateError {
330                    message: String::from(
331                        "could not queue RegisterWorker as the first stream frame",
332                    ),
333                })
334            })?;
335        let mut request = Request::new(ReceiverStream::new(outbound));
336        apply_auth_metadata(request.metadata_mut(), &self.config)?;
337        let response = client
338            .stream_worker(request)
339            .await
340            .map_err(registration_denial_error)?;
341        let mut receiver = response.into_inner();
342
343        let first = tokio::time::timeout(self.config.reconnect.max_backoff, receiver.message())
344            .await
345            .map_err(|_| {
346                WorkerError::registration(SessionStateError {
347                    message: format!(
348                        "server did not acknowledge registration within {:?}",
349                        self.config.reconnect.max_backoff
350                    ),
351                })
352            })?
353            .map_err(registration_denial_error)?;
354        let ack = match first.and_then(|frame| frame.message) {
355            Some(aion_proto::generated::server_to_worker::Message::RegisterAck(ack)) => ack,
356            Some(_) => {
357                return Err(WorkerError::decode(SessionStateError {
358                    message: String::from(
359                        "protocol violation: server sent a non-RegisterAck frame before \
360                         acknowledging registration",
361                    ),
362                }));
363            }
364            None => {
365                return Err(WorkerError::registration(SessionStateError {
366                    message: String::from(
367                        "server ended the stream before acknowledging registration",
368                    ),
369                }));
370            }
371        };
372
373        self.registered_info = Some(RegisteredSessionInfo {
374            worker_id: ack.worker_id,
375            namespace: ack.namespace,
376            heartbeat_window: std::time::Duration::from_millis(ack.heartbeat_window_ms),
377        });
378        self.sender = Some(sender);
379        self.receiver = Some(receiver);
380        Ok(())
381    }
382
383    /// Sends one frame with a per-send deadline of the reconnect policy's
384    /// `max_backoff`: a send that outlives the operator's longest tolerable
385    /// pause is, by that same definition, a dead session and surfaces as a
386    /// retryable transport error instead of hanging the worker forever.
387    async fn send_to_server(
388        &self,
389        message: aion_proto::generated::worker_to_server::Message,
390    ) -> Result<(), WorkerError> {
391        let sender = self.sender.as_ref().ok_or_else(|| {
392            WorkerError::registration(SessionStateError {
393                message: String::from("worker stream has not been opened"),
394            })
395        })?;
396        let send = sender.send(aion_proto::generated::WorkerToServer {
397            message: Some(message),
398        });
399        tokio::time::timeout(self.config.reconnect.max_backoff, send)
400            .await
401            .map_err(|_| WorkerError::Transport {
402                source: tonic::Status::unavailable(format!(
403                    "worker stream send did not complete within {:?}",
404                    self.config.reconnect.max_backoff
405                )),
406            })?
407            .map_err(|source| WorkerError::Transport {
408                source: tonic::Status::unavailable(format!("worker stream send failed: {source}")),
409            })
410    }
411}
412
413/// Maps the `StreamWorker` RPC's rejection status to the worker error taxonomy.
414///
415/// The server validates stream metadata (credentials) and the `RegisterWorker`
416/// frame before returning response headers, so both failure classes surface
417/// from the same await: `Unauthenticated` is a credential/handshake rejection,
418/// everything else is a registration outcome (`PermissionDenied` for an
419/// ungranted namespace, `Unavailable` for transient transport faults). Both
420/// shapes preserve the status for `WorkerError::grpc_status` / `is_retryable`.
421fn registration_denial_error(status: tonic::Status) -> WorkerError {
422    if status.code() == tonic::Code::Unauthenticated {
423        WorkerError::Handshake { source: status }
424    } else {
425        WorkerError::Registration {
426            source: Box::new(status),
427        }
428    }
429}
430
431fn apply_auth_metadata(
432    metadata: &mut tonic::metadata::MetadataMap,
433    config: &WorkerConfig,
434) -> Result<(), WorkerError> {
435    // The `x-aion-namespaces` metadata reflects the worker's full namespace SET,
436    // comma-joined in advertised order. The server authorizes the worker for
437    // every namespace it registers under.
438    let namespaces_value = config.namespaces.join(",");
439    let namespace =
440        MetadataValue::try_from(namespaces_value.as_str()).map_err(|_| WorkerError::Handshake {
441            source: tonic::Status::invalid_argument(
442                "worker namespaces are not valid gRPC metadata",
443            ),
444        })?;
445    let subject =
446        MetadataValue::try_from(config.subject.as_str()).map_err(|_| WorkerError::Handshake {
447            source: tonic::Status::invalid_argument("worker subject is not valid gRPC metadata"),
448        })?;
449    metadata.insert("x-aion-namespaces", namespace);
450    metadata.insert("x-aion-subject", subject);
451    Ok(())
452}
453
454#[async_trait]
455impl WorkerSession for GrpcWorkerSession {
456    async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
457        self.config = config.clone();
458        if self.client.is_none() {
459            self.client = Some(
460                GeneratedClient::connect(self.config.endpoint.clone())
461                    .await
462                    .map_err(|source| WorkerError::Connect { source })?,
463            );
464        }
465        Ok(())
466    }
467
468    async fn register(
469        &mut self,
470        activity_types: Vec<String>,
471        available_handlers: &BTreeSet<String>,
472    ) -> Result<(), WorkerError> {
473        self.register_with_contract(activity_types, Vec::new(), available_handlers)
474            .await
475    }
476
477    async fn register_with_contract(
478        &mut self,
479        activity_types: Vec<String>,
480        activities: Vec<aion_package::ActivityDescriptor>,
481        available_handlers: &BTreeSet<String>,
482    ) -> Result<(), WorkerError> {
483        validate_activity_handlers(&activity_types, available_handlers)?;
484        self.activity_types.clone_from(&activity_types);
485
486        // OQ-5: the registration namespace SET is the SAME set advertised in the
487        // `x-aion-namespaces` auth metadata (`apply_auth_metadata`), so a worker
488        // registers into exactly the namespaces it is authorized for.
489        // `task_queue` is the disjoint pool/flavour selector within each
490        // namespace; `node` is the optional locality affinity (default hostname).
491        let activities = activities
492            .into_iter()
493            .map(|activity| {
494                Ok(aion_proto::generated::ActivityDescriptor {
495                    name: activity.name,
496                    input_schema_json: serde_json::to_string(&activity.input_schema)
497                        .map_err(WorkerError::encode)?,
498                    output_schema_json: serde_json::to_string(&activity.output_schema)
499                        .map_err(WorkerError::encode)?,
500                })
501            })
502            .collect::<Result<Vec<_>, WorkerError>>()?;
503        let register = aion_proto::generated::RegisterWorker {
504            namespaces: self.config.namespaces.clone(),
505            activity_types,
506            task_queue: self.config.task_queue.clone(),
507            node: self.config.node.clone(),
508            activities,
509            identity: self.config.identity.clone(),
510            instance: None,
511        };
512        self.open_registered_stream(register).await
513    }
514
515    fn receive_tasks(&mut self) -> WorkerTaskStream {
516        match self.receiver.take() {
517            Some(receiver) => Box::pin(receiver.filter_map(|message| async move {
518                Some(match message {
519                    Ok(server_message) => decode_server_message(server_message),
520                    Err(source) => Err(WorkerError::Transport { source }),
521                })
522            })),
523            None => Box::pin(futures::stream::iter([Err(WorkerError::Transport {
524                source: tonic::Status::failed_precondition(
525                    "worker receive stream has not been opened",
526                ),
527            })])),
528        }
529    }
530
531    async fn report_result(
532        &mut self,
533        workflow_id: WorkflowId,
534        activity_id: ActivityId,
535        run_id: Option<RunId>,
536        completion_token: String,
537        result: Payload,
538    ) -> Result<(), WorkerError> {
539        let run_id = run_id.ok_or_else(|| {
540            WorkerError::decode(SessionStateError {
541                message: String::from(
542                    "activity result run_id is missing; refusing an incomplete fenced report",
543                ),
544            })
545        })?;
546        let result = ProtoActivityResult {
547            workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
548            activity_id: Some(ProtoActivityId::from(activity_id)),
549            run_id: Some(ProtoRunId::from(run_id)),
550            outcome: Some(proto_activity_result::Outcome::Result(ProtoPayload::from(
551                result,
552            ))),
553            completion_token,
554        };
555        self.send_to_server(aion_proto::generated::worker_to_server::Message::Result(
556            generated_activity_result(result),
557        ))
558        .await
559    }
560
561    async fn report_failure(
562        &mut self,
563        workflow_id: WorkflowId,
564        activity_id: ActivityId,
565        run_id: Option<RunId>,
566        completion_token: String,
567        failure: ActivityError,
568    ) -> Result<(), WorkerError> {
569        let run_id = run_id.ok_or_else(|| {
570            WorkerError::decode(SessionStateError {
571                message: String::from(
572                    "activity failure run_id is missing; refusing an incomplete fenced report",
573                ),
574            })
575        })?;
576        let result = ProtoActivityResult {
577            workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
578            activity_id: Some(ProtoActivityId::from(activity_id)),
579            run_id: Some(ProtoRunId::from(run_id)),
580            outcome: Some(proto_activity_result::Outcome::Error(failure.into())),
581            completion_token,
582        };
583        self.send_to_server(aion_proto::generated::worker_to_server::Message::Result(
584            generated_activity_result(result),
585        ))
586        .await
587    }
588
589    async fn send_heartbeat(
590        &mut self,
591        workflow_id: WorkflowId,
592        activity_id: ActivityId,
593        progress: Option<Payload>,
594    ) -> Result<(), WorkerError> {
595        let heartbeat = ProtoHeartbeat {
596            workflow_id: Some(ProtoWorkflowId::from(workflow_id)),
597            activity_id: Some(ProtoActivityId::from(activity_id)),
598            progress: progress.map(ProtoPayload::from),
599        };
600        self.send_to_server(aion_proto::generated::worker_to_server::Message::Heartbeat(
601            generated_heartbeat(heartbeat),
602        ))
603        .await
604    }
605
606    async fn send_connection_heartbeat(&mut self) -> Result<(), WorkerError> {
607        let heartbeat = ProtoHeartbeat {
608            workflow_id: None,
609            activity_id: None,
610            progress: None,
611        };
612        self.send_to_server(aion_proto::generated::worker_to_server::Message::Heartbeat(
613            generated_heartbeat(heartbeat),
614        ))
615        .await
616    }
617
618    async fn answer_liveness_ping(&mut self, sequence: u64) -> Result<(), WorkerError> {
619        self.send_to_server(
620            aion_proto::generated::worker_to_server::Message::LivenessAnswer(
621                aion_proto::generated::LivenessAnswer {
622                    liveness_ping: sequence,
623                },
624            ),
625        )
626        .await
627    }
628
629    fn heartbeat_window(&self) -> Option<std::time::Duration> {
630        self.registered_info
631            .as_ref()
632            .map(|info| info.heartbeat_window)
633    }
634}
635
636fn decode_server_message(
637    message: aion_proto::generated::ServerToWorker,
638) -> Result<WorkerSessionEvent, WorkerError> {
639    match message.message {
640        Some(aion_proto::generated::server_to_worker::Message::Task(task)) => {
641            Ok(WorkerSessionEvent::Task(Box::new(proto_task(task))))
642        }
643        Some(aion_proto::generated::server_to_worker::Message::Drain(_)) => {
644            Ok(WorkerSessionEvent::Drain)
645        }
646        Some(aion_proto::generated::server_to_worker::Message::ResultAck(ack)) => {
647            decode_result_ack(ack)
648        }
649        Some(aion_proto::generated::server_to_worker::Message::LivenessPing(ping)) => {
650            Ok(WorkerSessionEvent::LivenessPing {
651                sequence: ping.liveness_ping,
652                silence_window: std::time::Duration::from_millis(ping.silence_window_ms),
653            })
654        }
655        Some(aion_proto::generated::server_to_worker::Message::RegisterAck(_)) => {
656            // The ack is consumed inside `open_registered_stream`; a second
657            // one mid-stream is a server ordering bug that must surface.
658            Err(WorkerError::decode(SessionStateError {
659                message: String::from(
660                    "protocol violation: RegisterAck received after registration completed",
661                ),
662            }))
663        }
664        None => Err(WorkerError::decode(SessionStateError {
665            message: String::from("server-to-worker message was empty"),
666        })),
667    }
668}
669
670fn decode_result_ack(
671    ack: aion_proto::generated::ResultAck,
672) -> Result<WorkerSessionEvent, WorkerError> {
673    let workflow_id = ack
674        .workflow_id
675        .ok_or_else(|| {
676            WorkerError::decode(SessionStateError {
677                message: String::from("result ack workflow_id is missing"),
678            })
679        })
680        .and_then(|id| {
681            WorkflowId::try_from(ProtoWorkflowId { uuid: id.uuid }).map_err(|source| {
682                WorkerError::decode(SessionStateError {
683                    message: format!("result ack workflow_id is invalid: {source}"),
684                })
685            })
686        })?;
687    let activity_id = ack
688        .activity_id
689        .map(|id| ActivityId::from_sequence_position(id.sequence_position))
690        .ok_or_else(|| {
691            WorkerError::decode(SessionStateError {
692                message: String::from("result ack activity_id is missing"),
693            })
694        })?;
695    Ok(WorkerSessionEvent::ResultAck {
696        workflow_id,
697        activity_id,
698    })
699}
700
701fn generated_activity_result(value: ProtoActivityResult) -> aion_proto::generated::ActivityResult {
702    aion_proto::generated::ActivityResult {
703        workflow_id: value.workflow_id.map(generated_workflow_id),
704        activity_id: value.activity_id.map(generated_activity_id),
705        run_id: value.run_id.map(generated_run_id),
706        completion_token: value.completion_token,
707        outcome: value.outcome.map(|outcome| match outcome {
708            proto_activity_result::Outcome::Result(result) => {
709                aion_proto::generated::activity_result::Outcome::Result(generated_payload(result))
710            }
711            proto_activity_result::Outcome::Error(error) => {
712                aion_proto::generated::activity_result::Outcome::Error(generated_error(error))
713            }
714        }),
715    }
716}
717
718fn generated_heartbeat(value: ProtoHeartbeat) -> aion_proto::generated::Heartbeat {
719    aion_proto::generated::Heartbeat {
720        workflow_id: value.workflow_id.map(generated_workflow_id),
721        activity_id: value.activity_id.map(generated_activity_id),
722        progress: value.progress.map(generated_payload),
723    }
724}
725
726fn proto_task(value: aion_proto::generated::ActivityTask) -> ProtoActivityTask {
727    ProtoActivityTask {
728        workflow_id: value.workflow_id.map(proto_workflow_id),
729        activity_id: value.activity_id.map(proto_activity_id),
730        activity_type: value.activity_type,
731        input: value.input.map(proto_payload),
732        attempt: value.attempt,
733        labels: value.labels,
734        run_id: value.run_id.map(proto_run_id),
735        completion_token: value.completion_token,
736        idempotency_key: value.idempotency_key,
737    }
738}
739
740fn generated_payload(value: ProtoPayload) -> aion_proto::generated::Payload {
741    aion_proto::generated::Payload {
742        content_type: value.content_type,
743        bytes: value.bytes,
744    }
745}
746
747fn proto_payload(value: aion_proto::generated::Payload) -> ProtoPayload {
748    ProtoPayload {
749        content_type: value.content_type,
750        bytes: value.bytes,
751    }
752}
753
754fn generated_workflow_id(value: ProtoWorkflowId) -> aion_proto::generated::WorkflowId {
755    aion_proto::generated::WorkflowId { uuid: value.uuid }
756}
757
758fn proto_workflow_id(value: aion_proto::generated::WorkflowId) -> ProtoWorkflowId {
759    ProtoWorkflowId { uuid: value.uuid }
760}
761
762fn generated_run_id(value: ProtoRunId) -> aion_proto::generated::RunId {
763    aion_proto::generated::RunId { uuid: value.uuid }
764}
765
766fn proto_run_id(value: aion_proto::generated::RunId) -> ProtoRunId {
767    ProtoRunId { uuid: value.uuid }
768}
769
770fn generated_activity_id(value: ProtoActivityId) -> aion_proto::generated::ActivityId {
771    aion_proto::generated::ActivityId {
772        sequence_position: value.sequence_position,
773    }
774}
775
776fn proto_activity_id(value: aion_proto::generated::ActivityId) -> ProtoActivityId {
777    ProtoActivityId {
778        sequence_position: value.sequence_position,
779    }
780}
781
782fn generated_error(value: aion_proto::ProtoActivityError) -> aion_proto::generated::ActivityError {
783    aion_proto::generated::ActivityError {
784        kind: value.kind,
785        message: value.message,
786        details: value.details.map(generated_payload),
787    }
788}
789
790#[derive(thiserror::Error, Debug)]
791#[error("{message}")]
792struct SessionStateError {
793    message: String,
794}
795
796#[cfg(test)]
797mod tests {
798    use std::collections::BTreeSet;
799
800    use aion_proto::ProtoActivityTask;
801    use async_trait::async_trait;
802    use futures::{StreamExt, stream};
803
804    use super::{
805        WorkerSession, WorkerSessionEvent, WorkerTaskStream, apply_auth_metadata,
806        validate_activity_handlers,
807    };
808    use crate::error::WorkerError;
809    use crate::{ReconnectConfig, WorkerConfig};
810
811    #[derive(Default)]
812    struct FakeSession {
813        handshakes: Vec<(String, String)>,
814        registrations: Vec<Vec<String>>,
815    }
816
817    #[async_trait]
818    impl WorkerSession for FakeSession {
819        async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
820            self.handshakes
821                .push((config.task_queue.clone(), config.identity.clone()));
822            Ok(())
823        }
824
825        async fn register(
826            &mut self,
827            activity_types: Vec<String>,
828            available_handlers: &BTreeSet<String>,
829        ) -> Result<(), WorkerError> {
830            validate_activity_handlers(&activity_types, available_handlers)?;
831            self.registrations.push(activity_types);
832            Ok(())
833        }
834
835        fn receive_tasks(&mut self) -> WorkerTaskStream {
836            Box::pin(stream::iter([Ok(WorkerSessionEvent::Task(Box::new(
837                ProtoActivityTask {
838                    workflow_id: None,
839                    activity_id: None,
840                    activity_type: String::from("charge-card"),
841                    input: None,
842                    attempt: 1,
843                    labels: std::collections::HashMap::new(),
844                    run_id: Some(aion_proto::ProtoRunId::from(aion_core::RunId::new_v4())),
845                    completion_token: String::from("generation-1"),
846                    idempotency_key: String::from("effect-key"),
847                },
848            )))]))
849        }
850
851        async fn report_result(
852            &mut self,
853            workflow_id: aion_core::WorkflowId,
854            activity_id: aion_core::ActivityId,
855            run_id: Option<aion_core::RunId>,
856            completion_token: String,
857            result: aion_core::Payload,
858        ) -> Result<(), WorkerError> {
859            drop((workflow_id, activity_id, run_id, completion_token, result));
860            Ok(())
861        }
862
863        async fn report_failure(
864            &mut self,
865            workflow_id: aion_core::WorkflowId,
866            activity_id: aion_core::ActivityId,
867            run_id: Option<aion_core::RunId>,
868            completion_token: String,
869            failure: aion_core::ActivityError,
870        ) -> Result<(), WorkerError> {
871            drop((workflow_id, activity_id, run_id, completion_token, failure));
872            Ok(())
873        }
874
875        async fn send_heartbeat(
876            &mut self,
877            workflow_id: aion_core::WorkflowId,
878            activity_id: aion_core::ActivityId,
879            progress: Option<aion_core::Payload>,
880        ) -> Result<(), WorkerError> {
881            drop((workflow_id, activity_id, progress));
882            Ok(())
883        }
884    }
885
886    #[test]
887    fn apply_auth_metadata_sets_worker_authorization_headers() -> Result<(), WorkerError> {
888        let config = WorkerConfig::builder()
889            .endpoint("http://127.0.0.1:50051")
890            .task_queue("payments")
891            .identity("worker-a")
892            .max_concurrency(4)
893            .reconnect_initial_backoff(std::time::Duration::from_millis(5))
894            .reconnect_max_backoff(std::time::Duration::from_millis(20))
895            .reconnect_max_attempts(3)
896            .namespace("payments")
897            .subject("worker-a")
898            .build()
899            .map_err(WorkerError::registration)?;
900        let mut metadata = tonic::metadata::MetadataMap::new();
901
902        apply_auth_metadata(&mut metadata, &config)?;
903
904        assert_eq!(
905            metadata
906                .get("x-aion-namespaces")
907                .and_then(|value| value.to_str().ok()),
908            Some("payments")
909        );
910        assert_eq!(
911            metadata
912                .get("x-aion-subject")
913                .and_then(|value| value.to_str().ok()),
914            Some("worker-a")
915        );
916        Ok(())
917    }
918
919    #[tokio::test]
920    async fn fake_session_records_handshake_and_registration() -> Result<(), WorkerError> {
921        let config = WorkerConfig::new(
922            "http://127.0.0.1:50051",
923            "payments",
924            "worker-a",
925            4,
926            ReconnectConfig::new(
927                std::time::Duration::from_millis(5),
928                std::time::Duration::from_millis(20),
929                3,
930            ),
931            None,
932        );
933        let activity_types = vec![String::from("charge-card"), String::from("send-email")];
934        let handlers = activity_types.iter().cloned().collect::<BTreeSet<_>>();
935        let mut session = FakeSession::default();
936
937        session.handshake(&config).await?;
938        session.register(activity_types.clone(), &handlers).await?;
939        let received = session.receive_tasks().next().await;
940
941        assert_eq!(
942            session.handshakes,
943            vec![(String::from("payments"), String::from("worker-a"))]
944        );
945        assert_eq!(session.registrations, vec![activity_types]);
946        assert!(received.is_some());
947
948        Ok(())
949    }
950
951    #[tokio::test]
952    async fn grpc_reports_echo_the_dispatched_completion_token() -> Result<(), WorkerError> {
953        let config = WorkerConfig::new(
954            "http://127.0.0.1:50051",
955            "payments",
956            "worker-a",
957            1,
958            ReconnectConfig::new(
959                std::time::Duration::from_millis(5),
960                std::time::Duration::from_millis(20),
961                3,
962            ),
963            None,
964        );
965        let (sender, mut receiver) = tokio::sync::mpsc::channel(2);
966        let mut session = super::GrpcWorkerSession {
967            config,
968            activity_types: Vec::new(),
969            client: None,
970            sender: Some(sender),
971            receiver: None,
972            registered_info: None,
973        };
974
975        session
976            .report_result(
977                aion_core::WorkflowId::new_v4(),
978                aion_core::ActivityId::from_sequence_position(1),
979                Some(aion_core::RunId::new_v4()),
980                String::from("success-generation"),
981                aion_core::Payload::new(aion_core::ContentType::Json, b"{}".to_vec()),
982            )
983            .await?;
984        session
985            .report_failure(
986                aion_core::WorkflowId::new_v4(),
987                aion_core::ActivityId::from_sequence_position(2),
988                Some(aion_core::RunId::new_v4()),
989                String::from("failure-generation"),
990                aion_core::ActivityError {
991                    kind: aion_core::ActivityErrorKind::Terminal,
992                    message: String::from("failed"),
993                    details: None,
994                },
995            )
996            .await?;
997
998        let success = receiver.recv().await.ok_or_else(|| {
999            WorkerError::decode(super::SessionStateError {
1000                message: String::from("result report channel closed"),
1001            })
1002        })?;
1003        let failure = receiver.recv().await.ok_or_else(|| {
1004            WorkerError::decode(super::SessionStateError {
1005                message: String::from("failure report channel closed"),
1006            })
1007        })?;
1008        let success_token = match success.message {
1009            Some(aion_proto::generated::worker_to_server::Message::Result(result)) => {
1010                result.completion_token
1011            }
1012            _ => {
1013                return Err(WorkerError::decode(super::SessionStateError {
1014                    message: String::from("success report did not emit an ActivityResult"),
1015                }));
1016            }
1017        };
1018        let failure_token = match failure.message {
1019            Some(aion_proto::generated::worker_to_server::Message::Result(result)) => {
1020                result.completion_token
1021            }
1022            _ => {
1023                return Err(WorkerError::decode(super::SessionStateError {
1024                    message: String::from("failure report did not emit an ActivityResult"),
1025                }));
1026            }
1027        };
1028        assert_eq!(success_token, "success-generation");
1029        assert_eq!(failure_token, "failure-generation");
1030        Ok(())
1031    }
1032
1033    /// Brief test 16: a report send that never completes (server stopped
1034    /// reading; outbound channel full) times out retryably at the reconnect
1035    /// policy's `max_backoff` on a paused clock — the worker never hangs.
1036    #[tokio::test(start_paused = true)]
1037    async fn report_send_times_out_retryably_at_max_backoff() -> Result<(), WorkerError> {
1038        let config = WorkerConfig::new(
1039            "http://127.0.0.1:50051",
1040            "payments",
1041            "worker-a",
1042            1,
1043            ReconnectConfig::new(
1044                std::time::Duration::from_millis(5),
1045                std::time::Duration::from_millis(20),
1046                3,
1047            ),
1048            None,
1049        );
1050        let (sender, receiver) = tokio::sync::mpsc::channel(1);
1051        // Fill the channel so the next send blocks forever, modelling a
1052        // server that stopped draining its receive side.
1053        sender
1054            .try_send(aion_proto::generated::WorkerToServer { message: None })
1055            .map_err(WorkerError::decode)?;
1056        let mut session = super::GrpcWorkerSession {
1057            config,
1058            activity_types: Vec::new(),
1059            client: None,
1060            sender: Some(sender),
1061            receiver: None,
1062            registered_info: None,
1063        };
1064
1065        let result = session
1066            .report_result(
1067                aion_core::WorkflowId::new_v4(),
1068                aion_core::ActivityId::from_sequence_position(1),
1069                Some(aion_core::RunId::new_v4()),
1070                String::from("generation-1"),
1071                aion_core::Payload::new(aion_core::ContentType::Json, b"{}".to_vec()),
1072            )
1073            .await;
1074
1075        let Err(error) = result else {
1076            return Err(WorkerError::Transport {
1077                source: tonic::Status::internal("a hung send must time out, not hang"),
1078            });
1079        };
1080        assert!(
1081            matches!(error, WorkerError::Transport { .. }),
1082            "send deadline elapse must be a retryable transport error: {error}"
1083        );
1084        assert!(error.is_retryable());
1085        assert!(
1086            error.to_string().contains("did not complete"),
1087            "the error must name the deadline: {error}"
1088        );
1089        drop(receiver);
1090        Ok(())
1091    }
1092
1093    #[test]
1094    fn registration_rejects_activity_without_handler() {
1095        let activity_types = vec![String::from("charge-card"), String::from("send-email")];
1096        let handlers = [String::from("charge-card")]
1097            .into_iter()
1098            .collect::<BTreeSet<_>>();
1099
1100        let result = validate_activity_handlers(&activity_types, &handlers);
1101        assert!(result.is_err());
1102        let error = match result {
1103            Ok(()) => return,
1104            Err(error) => error,
1105        };
1106
1107        assert_eq!(
1108            error.to_string(),
1109            "worker registration failed: activity type `send-email` has no registered handler"
1110        );
1111    }
1112}