Skip to main content

aion_worker/protocol/
session.rs

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