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