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    /// Carried on the wire as `CancelActivity` and decoded by
48    /// `decode_cancel_activity`. The runtime handles it without forcing task
49    /// termination: the activity's cancellation token is tripped and the
50    /// handler is given the chance to stop on its own terms.
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::CancelActivity(cancel)) => {
656            decode_cancel_activity(cancel)
657        }
658        Some(aion_proto::generated::server_to_worker::Message::RegisterAck(_)) => {
659            // The ack is consumed inside `open_registered_stream`; a second
660            // one mid-stream is a server ordering bug that must surface.
661            Err(WorkerError::decode(SessionStateError {
662                message: String::from(
663                    "protocol violation: RegisterAck received after registration completed",
664                ),
665            }))
666        }
667        None => Err(WorkerError::decode(SessionStateError {
668            message: String::from("server-to-worker message was empty"),
669        })),
670    }
671}
672
673/// Decode a server-initiated activity cancellation (#233).
674///
675/// This is the seam that makes the worker's cancellation path REACHABLE: until
676/// it existed, `WorkerSessionEvent::Cancel` was never constructed anywhere in
677/// production, so the receive loop's cancel arm — and the handle, in-flight
678/// registry, and process-group kill behind it — could not run at all.
679///
680/// Both ids are REQUIRED. A cancel that cannot name what it is cancelling is
681/// refused rather than defaulted: the in-flight registry is addressed by the
682/// `(workflow, activity)` pair together, so a missing half would silently
683/// address nothing and the refusal would be invisible — which is this whole
684/// defect's shape.
685fn decode_cancel_activity(
686    cancel: aion_proto::generated::CancelActivity,
687) -> Result<WorkerSessionEvent, WorkerError> {
688    let workflow_id = cancel
689        .workflow_id
690        .ok_or_else(|| {
691            WorkerError::decode(SessionStateError {
692                message: String::from("cancel activity workflow_id is missing"),
693            })
694        })
695        .and_then(|id| {
696            WorkflowId::try_from(ProtoWorkflowId { uuid: id.uuid }).map_err(|source| {
697                WorkerError::decode(SessionStateError {
698                    message: format!("cancel activity workflow_id is invalid: {source}"),
699                })
700            })
701        })?;
702    let activity_id = cancel
703        .activity_id
704        .map(|id| ActivityId::from_sequence_position(id.sequence_position))
705        .ok_or_else(|| {
706            WorkerError::decode(SessionStateError {
707                message: String::from("cancel activity activity_id is missing"),
708            })
709        })?;
710    Ok(WorkerSessionEvent::Cancel {
711        workflow_id,
712        activity_id,
713    })
714}
715
716fn decode_result_ack(
717    ack: aion_proto::generated::ResultAck,
718) -> Result<WorkerSessionEvent, WorkerError> {
719    let workflow_id = ack
720        .workflow_id
721        .ok_or_else(|| {
722            WorkerError::decode(SessionStateError {
723                message: String::from("result ack workflow_id is missing"),
724            })
725        })
726        .and_then(|id| {
727            WorkflowId::try_from(ProtoWorkflowId { uuid: id.uuid }).map_err(|source| {
728                WorkerError::decode(SessionStateError {
729                    message: format!("result ack workflow_id is invalid: {source}"),
730                })
731            })
732        })?;
733    let activity_id = ack
734        .activity_id
735        .map(|id| ActivityId::from_sequence_position(id.sequence_position))
736        .ok_or_else(|| {
737            WorkerError::decode(SessionStateError {
738                message: String::from("result ack activity_id is missing"),
739            })
740        })?;
741    Ok(WorkerSessionEvent::ResultAck {
742        workflow_id,
743        activity_id,
744    })
745}
746
747fn generated_activity_result(value: ProtoActivityResult) -> aion_proto::generated::ActivityResult {
748    aion_proto::generated::ActivityResult {
749        workflow_id: value.workflow_id.map(generated_workflow_id),
750        activity_id: value.activity_id.map(generated_activity_id),
751        run_id: value.run_id.map(generated_run_id),
752        completion_token: value.completion_token,
753        outcome: value.outcome.map(|outcome| match outcome {
754            proto_activity_result::Outcome::Result(result) => {
755                aion_proto::generated::activity_result::Outcome::Result(generated_payload(result))
756            }
757            proto_activity_result::Outcome::Error(error) => {
758                aion_proto::generated::activity_result::Outcome::Error(generated_error(error))
759            }
760        }),
761    }
762}
763
764fn generated_heartbeat(value: ProtoHeartbeat) -> aion_proto::generated::Heartbeat {
765    aion_proto::generated::Heartbeat {
766        workflow_id: value.workflow_id.map(generated_workflow_id),
767        activity_id: value.activity_id.map(generated_activity_id),
768        progress: value.progress.map(generated_payload),
769    }
770}
771
772fn proto_task(value: aion_proto::generated::ActivityTask) -> ProtoActivityTask {
773    ProtoActivityTask {
774        workflow_id: value.workflow_id.map(proto_workflow_id),
775        activity_id: value.activity_id.map(proto_activity_id),
776        activity_type: value.activity_type,
777        input: value.input.map(proto_payload),
778        attempt: value.attempt,
779        labels: value.labels,
780        run_id: value.run_id.map(proto_run_id),
781        completion_token: value.completion_token,
782        idempotency_key: value.idempotency_key,
783    }
784}
785
786fn generated_payload(value: ProtoPayload) -> aion_proto::generated::Payload {
787    aion_proto::generated::Payload {
788        content_type: value.content_type,
789        bytes: value.bytes,
790    }
791}
792
793fn proto_payload(value: aion_proto::generated::Payload) -> ProtoPayload {
794    ProtoPayload {
795        content_type: value.content_type,
796        bytes: value.bytes,
797    }
798}
799
800fn generated_workflow_id(value: ProtoWorkflowId) -> aion_proto::generated::WorkflowId {
801    aion_proto::generated::WorkflowId { uuid: value.uuid }
802}
803
804fn proto_workflow_id(value: aion_proto::generated::WorkflowId) -> ProtoWorkflowId {
805    ProtoWorkflowId { uuid: value.uuid }
806}
807
808fn generated_run_id(value: ProtoRunId) -> aion_proto::generated::RunId {
809    aion_proto::generated::RunId { uuid: value.uuid }
810}
811
812fn proto_run_id(value: aion_proto::generated::RunId) -> ProtoRunId {
813    ProtoRunId { uuid: value.uuid }
814}
815
816fn generated_activity_id(value: ProtoActivityId) -> aion_proto::generated::ActivityId {
817    aion_proto::generated::ActivityId {
818        sequence_position: value.sequence_position,
819    }
820}
821
822fn proto_activity_id(value: aion_proto::generated::ActivityId) -> ProtoActivityId {
823    ProtoActivityId {
824        sequence_position: value.sequence_position,
825    }
826}
827
828fn generated_error(value: aion_proto::ProtoActivityError) -> aion_proto::generated::ActivityError {
829    aion_proto::generated::ActivityError {
830        kind: value.kind,
831        message: value.message,
832        details: value.details.map(generated_payload),
833    }
834}
835
836#[derive(thiserror::Error, Debug)]
837#[error("{message}")]
838struct SessionStateError {
839    message: String,
840}
841
842#[cfg(test)]
843mod tests {
844    use std::collections::BTreeSet;
845
846    use aion_proto::ProtoActivityTask;
847    use async_trait::async_trait;
848    use futures::{StreamExt, stream};
849
850    use super::{
851        WorkerSession, WorkerSessionEvent, WorkerTaskStream, apply_auth_metadata,
852        decode_server_message, validate_activity_handlers,
853    };
854    use crate::error::WorkerError;
855    use crate::{ReconnectConfig, WorkerConfig};
856
857    #[derive(Default)]
858    struct FakeSession {
859        handshakes: Vec<(String, String)>,
860        registrations: Vec<Vec<String>>,
861    }
862
863    #[async_trait]
864    impl WorkerSession for FakeSession {
865        async fn handshake(&mut self, config: &WorkerConfig) -> Result<(), WorkerError> {
866            self.handshakes
867                .push((config.task_queue.clone(), config.identity.clone()));
868            Ok(())
869        }
870
871        async fn register(
872            &mut self,
873            activity_types: Vec<String>,
874            available_handlers: &BTreeSet<String>,
875        ) -> Result<(), WorkerError> {
876            validate_activity_handlers(&activity_types, available_handlers)?;
877            self.registrations.push(activity_types);
878            Ok(())
879        }
880
881        fn receive_tasks(&mut self) -> WorkerTaskStream {
882            Box::pin(stream::iter([Ok(WorkerSessionEvent::Task(Box::new(
883                ProtoActivityTask {
884                    workflow_id: None,
885                    activity_id: None,
886                    activity_type: String::from("charge-card"),
887                    input: None,
888                    attempt: 1,
889                    labels: std::collections::HashMap::new(),
890                    run_id: Some(aion_proto::ProtoRunId::from(aion_core::RunId::new_v4())),
891                    completion_token: String::from("generation-1"),
892                    idempotency_key: String::from("effect-key"),
893                },
894            )))]))
895        }
896
897        async fn report_result(
898            &mut self,
899            workflow_id: aion_core::WorkflowId,
900            activity_id: aion_core::ActivityId,
901            run_id: Option<aion_core::RunId>,
902            completion_token: String,
903            result: aion_core::Payload,
904        ) -> Result<(), WorkerError> {
905            drop((workflow_id, activity_id, run_id, completion_token, result));
906            Ok(())
907        }
908
909        async fn report_failure(
910            &mut self,
911            workflow_id: aion_core::WorkflowId,
912            activity_id: aion_core::ActivityId,
913            run_id: Option<aion_core::RunId>,
914            completion_token: String,
915            failure: aion_core::ActivityError,
916        ) -> Result<(), WorkerError> {
917            drop((workflow_id, activity_id, run_id, completion_token, failure));
918            Ok(())
919        }
920
921        async fn send_heartbeat(
922            &mut self,
923            workflow_id: aion_core::WorkflowId,
924            activity_id: aion_core::ActivityId,
925            progress: Option<aion_core::Payload>,
926        ) -> Result<(), WorkerError> {
927            drop((workflow_id, activity_id, progress));
928            Ok(())
929        }
930    }
931
932    #[test]
933    fn apply_auth_metadata_sets_worker_authorization_headers() -> Result<(), WorkerError> {
934        let config = WorkerConfig::builder()
935            .endpoint("http://127.0.0.1:50051")
936            .task_queue("payments")
937            .identity("worker-a")
938            .max_concurrency(4)
939            .reconnect_initial_backoff(std::time::Duration::from_millis(5))
940            .reconnect_max_backoff(std::time::Duration::from_millis(20))
941            .reconnect_max_attempts(3)
942            .namespace("payments")
943            .subject("worker-a")
944            .build()
945            .map_err(WorkerError::registration)?;
946        let mut metadata = tonic::metadata::MetadataMap::new();
947
948        apply_auth_metadata(&mut metadata, &config)?;
949
950        assert_eq!(
951            metadata
952                .get("x-aion-namespaces")
953                .and_then(|value| value.to_str().ok()),
954            Some("payments")
955        );
956        assert_eq!(
957            metadata
958                .get("x-aion-subject")
959                .and_then(|value| value.to_str().ok()),
960            Some("worker-a")
961        );
962        Ok(())
963    }
964
965    #[tokio::test]
966    async fn fake_session_records_handshake_and_registration() -> Result<(), WorkerError> {
967        let config = WorkerConfig::new(
968            "http://127.0.0.1:50051",
969            "payments",
970            "worker-a",
971            4,
972            ReconnectConfig::new(
973                std::time::Duration::from_millis(5),
974                std::time::Duration::from_millis(20),
975                3,
976            ),
977            None,
978        );
979        let activity_types = vec![String::from("charge-card"), String::from("send-email")];
980        let handlers = activity_types.iter().cloned().collect::<BTreeSet<_>>();
981        let mut session = FakeSession::default();
982
983        session.handshake(&config).await?;
984        session.register(activity_types.clone(), &handlers).await?;
985        let received = session.receive_tasks().next().await;
986
987        assert_eq!(
988            session.handshakes,
989            vec![(String::from("payments"), String::from("worker-a"))]
990        );
991        assert_eq!(session.registrations, vec![activity_types]);
992        assert!(received.is_some());
993
994        Ok(())
995    }
996
997    #[tokio::test]
998    async fn grpc_reports_echo_the_dispatched_completion_token() -> Result<(), WorkerError> {
999        let config = WorkerConfig::new(
1000            "http://127.0.0.1:50051",
1001            "payments",
1002            "worker-a",
1003            1,
1004            ReconnectConfig::new(
1005                std::time::Duration::from_millis(5),
1006                std::time::Duration::from_millis(20),
1007                3,
1008            ),
1009            None,
1010        );
1011        let (sender, mut receiver) = tokio::sync::mpsc::channel(2);
1012        let mut session = super::GrpcWorkerSession {
1013            config,
1014            activity_types: Vec::new(),
1015            client: None,
1016            sender: Some(sender),
1017            receiver: None,
1018            registered_info: None,
1019        };
1020
1021        session
1022            .report_result(
1023                aion_core::WorkflowId::new_v4(),
1024                aion_core::ActivityId::from_sequence_position(1),
1025                Some(aion_core::RunId::new_v4()),
1026                String::from("success-generation"),
1027                aion_core::Payload::new(aion_core::ContentType::Json, b"{}".to_vec()),
1028            )
1029            .await?;
1030        session
1031            .report_failure(
1032                aion_core::WorkflowId::new_v4(),
1033                aion_core::ActivityId::from_sequence_position(2),
1034                Some(aion_core::RunId::new_v4()),
1035                String::from("failure-generation"),
1036                aion_core::ActivityError {
1037                    kind: aion_core::ActivityErrorKind::Terminal,
1038                    message: String::from("failed"),
1039                    details: None,
1040                },
1041            )
1042            .await?;
1043
1044        let success = receiver.recv().await.ok_or_else(|| {
1045            WorkerError::decode(super::SessionStateError {
1046                message: String::from("result report channel closed"),
1047            })
1048        })?;
1049        let failure = receiver.recv().await.ok_or_else(|| {
1050            WorkerError::decode(super::SessionStateError {
1051                message: String::from("failure report channel closed"),
1052            })
1053        })?;
1054        let success_token = match success.message {
1055            Some(aion_proto::generated::worker_to_server::Message::Result(result)) => {
1056                result.completion_token
1057            }
1058            _ => {
1059                return Err(WorkerError::decode(super::SessionStateError {
1060                    message: String::from("success report did not emit an ActivityResult"),
1061                }));
1062            }
1063        };
1064        let failure_token = match failure.message {
1065            Some(aion_proto::generated::worker_to_server::Message::Result(result)) => {
1066                result.completion_token
1067            }
1068            _ => {
1069                return Err(WorkerError::decode(super::SessionStateError {
1070                    message: String::from("failure report did not emit an ActivityResult"),
1071                }));
1072            }
1073        };
1074        assert_eq!(success_token, "success-generation");
1075        assert_eq!(failure_token, "failure-generation");
1076        Ok(())
1077    }
1078
1079    /// Brief test 16: a report send that never completes (server stopped
1080    /// reading; outbound channel full) times out retryably at the reconnect
1081    /// policy's `max_backoff` on a paused clock — the worker never hangs.
1082    #[tokio::test(start_paused = true)]
1083    async fn report_send_times_out_retryably_at_max_backoff() -> Result<(), WorkerError> {
1084        let config = WorkerConfig::new(
1085            "http://127.0.0.1:50051",
1086            "payments",
1087            "worker-a",
1088            1,
1089            ReconnectConfig::new(
1090                std::time::Duration::from_millis(5),
1091                std::time::Duration::from_millis(20),
1092                3,
1093            ),
1094            None,
1095        );
1096        let (sender, receiver) = tokio::sync::mpsc::channel(1);
1097        // Fill the channel so the next send blocks forever, modelling a
1098        // server that stopped draining its receive side.
1099        sender
1100            .try_send(aion_proto::generated::WorkerToServer { message: None })
1101            .map_err(WorkerError::decode)?;
1102        let mut session = super::GrpcWorkerSession {
1103            config,
1104            activity_types: Vec::new(),
1105            client: None,
1106            sender: Some(sender),
1107            receiver: None,
1108            registered_info: None,
1109        };
1110
1111        let result = session
1112            .report_result(
1113                aion_core::WorkflowId::new_v4(),
1114                aion_core::ActivityId::from_sequence_position(1),
1115                Some(aion_core::RunId::new_v4()),
1116                String::from("generation-1"),
1117                aion_core::Payload::new(aion_core::ContentType::Json, b"{}".to_vec()),
1118            )
1119            .await;
1120
1121        let Err(error) = result else {
1122            return Err(WorkerError::Transport {
1123                source: tonic::Status::internal("a hung send must time out, not hang"),
1124            });
1125        };
1126        assert!(
1127            matches!(error, WorkerError::Transport { .. }),
1128            "send deadline elapse must be a retryable transport error: {error}"
1129        );
1130        assert!(error.is_retryable());
1131        assert!(
1132            error.to_string().contains("did not complete"),
1133            "the error must name the deadline: {error}"
1134        );
1135        drop(receiver);
1136        Ok(())
1137    }
1138
1139    #[test]
1140    fn registration_rejects_activity_without_handler() {
1141        let activity_types = vec![String::from("charge-card"), String::from("send-email")];
1142        let handlers = [String::from("charge-card")]
1143            .into_iter()
1144            .collect::<BTreeSet<_>>();
1145
1146        let result = validate_activity_handlers(&activity_types, &handlers);
1147        assert!(result.is_err());
1148        let error = match result {
1149            Ok(()) => return,
1150            Err(error) => error,
1151        };
1152
1153        assert_eq!(
1154            error.to_string(),
1155            "worker registration failed: activity type `send-email` has no registered handler"
1156        );
1157    }
1158
1159    /// Build a `ServerToWorker` carrying a cancel with the supplied halves, so
1160    /// each test states exactly which half it is withholding.
1161    fn cancel_frame(
1162        workflow_uuid: Option<String>,
1163        sequence_position: Option<u64>,
1164    ) -> aion_proto::generated::ServerToWorker {
1165        aion_proto::generated::ServerToWorker {
1166            message: Some(
1167                aion_proto::generated::server_to_worker::Message::CancelActivity(
1168                    aion_proto::generated::CancelActivity {
1169                        workflow_id: workflow_uuid
1170                            .map(|uuid| aion_proto::generated::WorkflowId { uuid }),
1171                        activity_id: sequence_position.map(|sequence_position| {
1172                            aion_proto::generated::ActivityId { sequence_position }
1173                        }),
1174                    },
1175                ),
1176            ),
1177        }
1178    }
1179
1180    type CancelTestResult = Result<(), Box<dyn std::error::Error>>;
1181
1182    /// Assert that a frame is refused with a message naming `expected_refusal`.
1183    fn assert_cancel_refused(
1184        frame: aion_proto::generated::ServerToWorker,
1185        expected_refusal: &str,
1186    ) -> CancelTestResult {
1187        let error = match decode_server_message(frame) {
1188            Ok(event) => {
1189                return Err(format!("an invalid cancel was accepted as {event:?}").into());
1190            }
1191            Err(error) => error,
1192        };
1193
1194        assert!(
1195            error.to_string().contains(expected_refusal),
1196            "refusal did not name the defective half: {error}"
1197        );
1198        Ok(())
1199    }
1200
1201    #[test]
1202    fn cancel_frame_decodes_to_the_pair_it_names() -> CancelTestResult {
1203        let workflow_id = aion_core::WorkflowId::new_v4();
1204
1205        let event = decode_server_message(cancel_frame(Some(workflow_id.to_string()), Some(7)))?;
1206
1207        // The whole point of the frame is that BOTH halves arrive intact: the
1208        // in-flight registry is addressed by the pair, so a decode that kept
1209        // one and lost the other would cancel nothing while reporting success.
1210        match event {
1211            WorkerSessionEvent::Cancel {
1212                workflow_id: decoded_workflow,
1213                activity_id,
1214            } => {
1215                assert_eq!(decoded_workflow, workflow_id);
1216                assert_eq!(activity_id.sequence_position(), 7);
1217            }
1218            other => {
1219                return Err(format!("cancel frame decoded to the wrong event: {other:?}").into());
1220            }
1221        }
1222        Ok(())
1223    }
1224
1225    #[test]
1226    fn cancel_without_a_workflow_id_is_refused() -> CancelTestResult {
1227        assert_cancel_refused(cancel_frame(None, Some(7)), "workflow_id is missing")
1228    }
1229
1230    #[test]
1231    fn cancel_without_an_activity_id_is_refused() -> CancelTestResult {
1232        let workflow_id = aion_core::WorkflowId::new_v4();
1233        assert_cancel_refused(
1234            cancel_frame(Some(workflow_id.to_string()), None),
1235            "activity_id is missing",
1236        )
1237    }
1238
1239    #[test]
1240    fn cancel_carrying_an_unparseable_workflow_id_is_refused() -> CancelTestResult {
1241        assert_cancel_refused(
1242            cancel_frame(Some(String::from("not-a-uuid")), Some(7)),
1243            "workflow_id is invalid",
1244        )
1245    }
1246}