Skip to main content

aion_server/api/
worker_grpc.rs

1//! tonic `WorkerProtocol` service — bidirectional stream handler.
2
3use aion_proto::{
4    ProtoActivityResult, ProtoRegisterWorker,
5    generated::{
6        self,
7        worker_protocol_server::{WorkerProtocol, WorkerProtocolServer},
8    },
9};
10use tokio::sync::mpsc;
11use tokio_stream::wrappers::ReceiverStream;
12use tonic::{Request, Response, Status, Streaming};
13
14use crate::worker::PendingActivities;
15use crate::worker::dispatch::{ActivityCompletion, ActivityCompletionSink};
16use crate::worker::registry::{WorkerId, WorkerMessage};
17use crate::{CallerIdentity, ServerState};
18
19/// Cloneable tonic implementation for the worker bidirectional stream.
20#[derive(Clone)]
21pub struct WorkerGrpcService {
22    state: ServerState,
23}
24
25impl WorkerGrpcService {
26    /// Build a tonic worker service from shared server state.
27    #[must_use]
28    pub const fn new(state: ServerState) -> Self {
29        Self { state }
30    }
31}
32
33/// Construct the generated tonic server wrapper for the worker protocol.
34#[must_use]
35pub fn worker_service(state: ServerState) -> WorkerProtocolServer<WorkerGrpcService> {
36    WorkerProtocolServer::new(WorkerGrpcService::new(state))
37}
38
39#[tonic::async_trait]
40impl WorkerProtocol for WorkerGrpcService {
41    type StreamWorkerStream = ReceiverStream<Result<generated::ServerToWorker, Status>>;
42
43    async fn stream_worker(
44        &self,
45        request: Request<Streaming<generated::WorkerToServer>>,
46    ) -> Result<Response<Self::StreamWorkerStream>, Status> {
47        let metadata = request.metadata().clone();
48        let caller = worker_caller_from_metadata(&metadata, &self.state).await?;
49        let token_expires_at = token_expiration_from_metadata(&metadata, &self.state).await?;
50        let heartbeat_grace = self.state.runtime_config().worker.heartbeat_window;
51        let mut inbound = request.into_inner();
52
53        let first = inbound
54            .message()
55            .await?
56            .and_then(|msg| msg.message)
57            .ok_or_else(|| Status::invalid_argument("first message must be RegisterWorker"))?;
58
59        let register = match first {
60            generated::worker_to_server::Message::Register(r) => decode_register(r),
61            _ => {
62                return Err(Status::invalid_argument(
63                    "first message must be RegisterWorker",
64                ));
65            }
66        };
67
68        let (task_tx, task_rx) = mpsc::channel::<Result<generated::ServerToWorker, Status>>(32);
69        let (worker_tx, mut worker_rx) = mpsc::channel(32);
70
71        let registration = self
72            .state
73            .worker_registry()
74            .accept_registration(self.state.namespace_guard(), &caller, &register, worker_tx)
75            .await
76            .map_err(|error| status_from_server_error(&error))?;
77
78        let pending = self.state.pending_activities().clone();
79        let heartbeat = self.state.heartbeat_tracker().clone();
80        let drain = self.state.drain_state().clone();
81        let registry = self.state.worker_registry().clone();
82        let worker_id = registration
83            .worker_id()
84            .ok_or_else(|| Status::internal("worker registration missing id"))?;
85        // A worker serves a SET of namespaces; the ack echoes them joined in
86        // stable order purely for the worker's logs (the RegisterAck namespace
87        // field is informational, not a routing input).
88        let authorized_namespace = registration
89            .namespaces()
90            .filter(|namespaces| !namespaces.is_empty())
91            .ok_or_else(|| Status::internal("worker registration missing namespace"))?
92            .iter()
93            .cloned()
94            .collect::<Vec<_>>()
95            .join(",");
96
97        // RegisterAck ordering guarantee: the ack is enqueued on `task_tx`
98        // BEFORE the write forwarder that copies dispatched tasks onto the
99        // same channel is spawned, so no task frame can precede it on the
100        // wire. This is a structural ordering proof, not a timing hope.
101        task_tx
102            .try_send(Ok(register_ack_frame(
103                worker_id,
104                &authorized_namespace,
105                heartbeat_grace,
106            )))
107            .map_err(|_| Status::internal("worker response channel closed before RegisterAck"))?;
108
109        tokio::spawn(async move {
110            let write_handle = tokio::spawn({
111                let task_tx = task_tx.clone();
112                async move {
113                    while let Some(message) = worker_rx.recv().await {
114                        let msg = encode_server_to_worker(message);
115                        if task_tx.send(Ok(msg)).await.is_err() {
116                            break;
117                        }
118                    }
119                }
120            });
121
122            // Armed BEFORE the inbound loop runs: the sweep in its `Drop`
123            // fires on every exit from this task — clean stream end, stream
124            // error, token expiry, even a panic unwinding `process_inbound`.
125            // The unbounded dispatch wait depends on it.
126            let teardown = StreamTeardown {
127                worker_id,
128                heartbeat: &heartbeat,
129                registry: &registry,
130                pending: &pending,
131                drain: &drain,
132            };
133            let session = WorkerSession {
134                worker_id,
135                pending: &pending,
136                heartbeat: &heartbeat,
137                drain: &drain,
138                token_expires_at,
139                heartbeat_grace,
140                task_tx: task_tx.clone(),
141            };
142            if let Err(status) = process_inbound(inbound, session).await {
143                tracing::info!(
144                    worker_id = ?worker_id,
145                    %status,
146                    "worker stream closed with status"
147                );
148            }
149
150            write_handle.abort();
151            drop(task_tx);
152            drop(teardown);
153            // The teardown sweep already deregistered the stream; consuming
154            // the registration here is an idempotent no-op that still
155            // surfaces a poisoned-lock error loudly.
156            if let Err(error) = registration.deregister() {
157                tracing::error!(
158                    worker_id = ?worker_id,
159                    %error,
160                    "worker deregistration failed during stream teardown"
161                );
162            }
163        });
164
165        Ok(Response::new(ReceiverStream::new(task_rx)))
166    }
167}
168
169/// Drop guard that fails a torn-down worker stream's in-flight activities
170/// back to the engine.
171///
172/// A guard rather than a call site so the sweep cannot be skipped by any
173/// exit from the stream task — including a panic unwinding the inbound
174/// loop, which would otherwise leave every dispatch blocked on that worker
175/// waiting forever.
176struct StreamTeardown<'a> {
177    worker_id: WorkerId,
178    heartbeat: &'a crate::worker::HeartbeatTracker,
179    registry: &'a crate::worker::ConnectedWorkerRegistry,
180    pending: &'a PendingActivities,
181    drain: &'a crate::shutdown::DrainState,
182}
183
184impl Drop for StreamTeardown<'_> {
185    fn drop(&mut self) {
186        teardown_worker_stream(
187            self.worker_id,
188            self.heartbeat,
189            self.registry,
190            self.pending,
191            self.drain,
192        );
193    }
194}
195
196/// Fail a torn-down worker stream's in-flight activities back to the engine.
197///
198/// The stream is the worker's liveness. When it ends — process death,
199/// network disconnect, expired token — every activity still assigned to
200/// this worker must be failed back through the completion sink as a
201/// retryable lost-worker error. The activity dispatch wait is unbounded by
202/// design (the engine imposes no activity timeout), so this sweep is what
203/// unblocks dispatches whose worker died mid-activity; the engine's retry
204/// policy then decides re-dispatch.
205fn teardown_worker_stream(
206    worker_id: WorkerId,
207    heartbeat: &crate::worker::HeartbeatTracker,
208    registry: &crate::worker::ConnectedWorkerRegistry,
209    pending: &PendingActivities,
210    drain: &crate::shutdown::DrainState,
211) {
212    match heartbeat.fail_disconnected_worker(worker_id, registry, pending) {
213        Ok(report) if report.tasks.is_empty() => {}
214        Ok(report) => {
215            tracing::warn!(
216                worker_id = ?worker_id,
217                failed_tasks = report.tasks.len(),
218                "worker disconnected with in-flight activities; \
219                 surfaced as retryable lost-worker failures"
220            );
221        }
222        Err(error) => {
223            tracing::error!(
224                worker_id = ?worker_id,
225                %error,
226                "failed to sweep disconnected worker's in-flight activities"
227            );
228        }
229    }
230    // In-flight accounting may have just reached zero; wake any drain
231    // waiter so shutdown does not sit out its full timeout.
232    drain.notify_activity_drained();
233}
234
235struct WorkerSession<'a> {
236    worker_id: WorkerId,
237    pending: &'a PendingActivities,
238    heartbeat: &'a crate::worker::HeartbeatTracker,
239    drain: &'a crate::shutdown::DrainState,
240    token_expires_at: Option<u64>,
241    heartbeat_grace: std::time::Duration,
242    task_tx: mpsc::Sender<Result<generated::ServerToWorker, Status>>,
243}
244
245async fn process_inbound(
246    mut inbound: Streaming<generated::WorkerToServer>,
247    session: WorkerSession<'_>,
248) -> Result<(), Status> {
249    let mut expired_since: Option<std::time::Instant> = None;
250    while let Some(msg) = inbound.message().await? {
251        let Some(inner) = msg.message else {
252            continue;
253        };
254        match inner {
255            generated::worker_to_server::Message::Result(result) => {
256                let proto_result = decode_activity_result(result);
257                match ActivityCompletion::try_from(proto_result) {
258                    Ok(completion) => {
259                        let workflow_id = completion.workflow_id.clone();
260                        let activity_id = completion.activity_id.clone();
261                        if let Err(error) = session.heartbeat.complete_task(
262                            session.worker_id,
263                            &workflow_id,
264                            &activity_id,
265                        ) {
266                            // A poisoned liveness tracker would also break
267                            // the lost-worker sweep the unbounded dispatch
268                            // wait relies on — never swallow it.
269                            tracing::error!(
270                                worker_id = ?session.worker_id,
271                                workflow_id = %workflow_id,
272                                activity_id = %activity_id,
273                                %error,
274                                "failed to clear in-flight tracking for completed activity"
275                            );
276                        }
277                        session.drain.notify_activity_drained();
278                        if let Err(error) = session.pending.complete_activity(completion) {
279                            tracing::error!(
280                                worker_id = ?session.worker_id,
281                                workflow_id = %workflow_id,
282                                activity_id = %activity_id,
283                                %error,
284                                "activity completion handoff failed"
285                            );
286                        }
287                        // Ack every well-formed result frame — including
288                        // duplicates with no pending waiter; their re-report
289                        // obligation is equally discharged. `try_send`: a
290                        // worker that stopped draining its receive side must
291                        // not wedge the inbound loop; a dropped ack is
292                        // recovered by the next-session re-report.
293                        let ack = result_ack_frame(&workflow_id, &activity_id);
294                        if let Err(error) = session.task_tx.try_send(Ok(ack)) {
295                            tracing::warn!(
296                                worker_id = ?session.worker_id,
297                                workflow_id = %workflow_id,
298                                activity_id = %activity_id,
299                                %error,
300                                "result ack dropped: worker stream channel unavailable"
301                            );
302                        }
303                    }
304                    Err(error) => {
305                        // Malformed result: no ids to ack with. Loud, never
306                        // silent — the worker's entry will re-report and
307                        // re-fail visibly each session.
308                        tracing::error!(
309                            worker_id = ?session.worker_id,
310                            %error,
311                            "malformed activity result frame; no ack sent"
312                        );
313                    }
314                }
315            }
316            generated::worker_to_server::Message::Register(_) => {
317                tracing::warn!(
318                    worker_id = ?session.worker_id,
319                    "ignoring subsequent RegisterWorker message; \
320                     only the first registration is accepted per stream"
321                );
322            }
323            generated::worker_to_server::Message::Heartbeat(heartbeat_msg) => {
324                if let Err(error) = session.heartbeat.record_heartbeat(
325                    session.worker_id,
326                    decode_heartbeat(heartbeat_msg),
327                    std::time::Instant::now(),
328                ) {
329                    // Malformed frames and heartbeats for untracked tasks
330                    // are worker-side defects worth surfacing; a poisoned
331                    // tracker lock is a server-side corruption signal that
332                    // must never vanish silently.
333                    if matches!(error, crate::ServerError::LockPoisoned { .. }) {
334                        tracing::error!(
335                            worker_id = ?session.worker_id,
336                            %error,
337                            "heartbeat tracker lock poisoned; liveness state untrustworthy"
338                        );
339                    } else {
340                        tracing::warn!(
341                            worker_id = ?session.worker_id,
342                            %error,
343                            "worker heartbeat rejected"
344                        );
345                    }
346                }
347                if token_expired(session.token_expires_at) {
348                    let first_expired = *expired_since.get_or_insert_with(std::time::Instant::now);
349                    let _ = session
350                        .task_tx
351                        .send(Err(Status::unauthenticated(
352                            "worker token expired; re-authentication required",
353                        )))
354                        .await;
355                    if first_expired.elapsed() >= session.heartbeat_grace {
356                        return Err(Status::unauthenticated("worker token expired"));
357                    }
358                }
359            }
360        }
361    }
362    Ok(())
363}
364
365async fn worker_caller_from_metadata(
366    metadata: &tonic::metadata::MetadataMap,
367    state: &ServerState,
368) -> Result<CallerIdentity, Status> {
369    crate::api::grpc::caller_from_metadata(metadata, state).await
370}
371
372async fn token_expiration_from_metadata(
373    metadata: &tonic::metadata::MetadataMap,
374    state: &ServerState,
375) -> Result<Option<u64>, Status> {
376    if !state.runtime_config().auth.enabled {
377        return Ok(None);
378    }
379    #[cfg(feature = "auth")]
380    {
381        let bearer = metadata
382            .get("authorization")
383            .and_then(|value| value.to_str().ok())
384            .and_then(parse_bearer)
385            .ok_or_else(|| Status::unauthenticated("missing bearer token"))?;
386        let Some(cache) = state.jwks_cache() else {
387            return Err(Status::unauthenticated("invalid bearer token"));
388        };
389        return cache
390            .validate(&bearer)
391            .await
392            .map(|claims| Some(claims.expires_at()))
393            .map_err(|_error| Status::unauthenticated("invalid bearer token"));
394    }
395    #[cfg(not(feature = "auth"))]
396    {
397        let _ = metadata;
398        // Yield to preserve the async signature required by the auth-feature branch.
399        tokio::task::yield_now().await;
400        Ok(None)
401    }
402}
403
404#[cfg(feature = "auth")]
405fn parse_bearer(value: &str) -> Option<String> {
406    let token = value.strip_prefix("Bearer ")?.trim();
407    if token.is_empty() {
408        return None;
409    }
410    Some(token.to_owned())
411}
412
413fn token_expired(expires_at: Option<u64>) -> bool {
414    expires_at.is_some_and(|expires_at| {
415        #[cfg(feature = "auth")]
416        {
417            crate::auth::jwks::is_expired(expires_at)
418        }
419        #[cfg(not(feature = "auth"))]
420        {
421            let _ = expires_at;
422            false
423        }
424    })
425}
426
427fn status_from_server_error(error: &crate::ServerError) -> Status {
428    let wire = error.to_wire_error();
429    if wire.code == aion_proto::WireErrorCode::NamespaceDenied {
430        Status::permission_denied(wire.message)
431    } else {
432        Status::internal(wire.message)
433    }
434}
435
436/// Build the positive registration acknowledgement frame — the guaranteed
437/// first frame on every successful worker response stream.
438fn register_ack_frame(
439    worker_id: WorkerId,
440    namespace: &str,
441    heartbeat_window: std::time::Duration,
442) -> generated::ServerToWorker {
443    generated::ServerToWorker {
444        message: Some(generated::server_to_worker::Message::RegisterAck(
445            generated::RegisterAck {
446                worker_id: worker_id.value(),
447                namespace: namespace.to_owned(),
448                heartbeat_window_ms: u64::try_from(heartbeat_window.as_millis())
449                    .unwrap_or(u64::MAX),
450            },
451        )),
452    }
453}
454
455/// Build the per-result acknowledgement frame for a consumed `ActivityResult`.
456fn result_ack_frame(
457    workflow_id: &aion_core::WorkflowId,
458    activity_id: &aion_core::ActivityId,
459) -> generated::ServerToWorker {
460    generated::ServerToWorker {
461        message: Some(generated::server_to_worker::Message::ResultAck(
462            generated::ResultAck {
463                workflow_id: Some(generated::WorkflowId {
464                    uuid: workflow_id.to_string(),
465                }),
466                activity_id: Some(generated::ActivityId {
467                    sequence_position: activity_id.sequence_position(),
468                }),
469            },
470        )),
471    }
472}
473
474fn decode_register(r: generated::RegisterWorker) -> ProtoRegisterWorker {
475    ProtoRegisterWorker {
476        namespaces: r.namespaces,
477        activity_types: r.activity_types,
478        task_queue: r.task_queue,
479        node: r.node,
480    }
481}
482
483fn encode_server_to_worker(message: WorkerMessage) -> generated::ServerToWorker {
484    let message = match message {
485        WorkerMessage::ActivityTask(task) => {
486            generated::server_to_worker::Message::Task(encode_task(task))
487        }
488        WorkerMessage::DrainRequest => {
489            generated::server_to_worker::Message::Drain(generated::DrainRequest {})
490        }
491    };
492    generated::ServerToWorker {
493        message: Some(message),
494    }
495}
496
497fn encode_task(task: aion_proto::ProtoActivityTask) -> generated::ActivityTask {
498    generated::ActivityTask {
499        workflow_id: task
500            .workflow_id
501            .map(|id| generated::WorkflowId { uuid: id.uuid }),
502        activity_id: task.activity_id.map(|id| generated::ActivityId {
503            sequence_position: id.sequence_position,
504        }),
505        activity_type: task.activity_type,
506        input: task.input.map(|p| generated::Payload {
507            content_type: p.content_type,
508            bytes: p.bytes,
509        }),
510        attempt: task.attempt,
511        labels: task.labels,
512        run_id: task.run_id.map(|id| generated::RunId { uuid: id.uuid }),
513    }
514}
515
516fn decode_activity_result(r: generated::ActivityResult) -> ProtoActivityResult {
517    ProtoActivityResult {
518        workflow_id: r
519            .workflow_id
520            .map(|id| aion_proto::ProtoWorkflowId { uuid: id.uuid }),
521        activity_id: r.activity_id.map(|id| aion_proto::ProtoActivityId {
522            sequence_position: id.sequence_position,
523        }),
524        outcome: r.outcome.map(decode_outcome),
525        run_id: r.run_id.map(|id| aion_proto::ProtoRunId { uuid: id.uuid }),
526    }
527}
528
529fn decode_heartbeat(r: generated::Heartbeat) -> aion_proto::ProtoHeartbeat {
530    aion_proto::ProtoHeartbeat {
531        workflow_id: r
532            .workflow_id
533            .map(|id| aion_proto::ProtoWorkflowId { uuid: id.uuid }),
534        activity_id: r.activity_id.map(|id| aion_proto::ProtoActivityId {
535            sequence_position: id.sequence_position,
536        }),
537        progress: r.progress.map(|p| aion_proto::ProtoPayload {
538            content_type: p.content_type,
539            bytes: p.bytes,
540        }),
541    }
542}
543
544fn decode_outcome(
545    outcome: generated::activity_result::Outcome,
546) -> aion_proto::proto_activity_result::Outcome {
547    match outcome {
548        generated::activity_result::Outcome::Result(p) => {
549            aion_proto::proto_activity_result::Outcome::Result(aion_proto::ProtoPayload {
550                content_type: p.content_type,
551                bytes: p.bytes,
552            })
553        }
554        generated::activity_result::Outcome::Error(e) => {
555            aion_proto::proto_activity_result::Outcome::Error(aion_proto::ProtoActivityError {
556                kind: e.kind,
557                message: e.message,
558                details: e.details.map(|p| aion_proto::ProtoPayload {
559                    content_type: p.content_type,
560                    bytes: p.bytes,
561                }),
562            })
563        }
564    }
565}