1use 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#[derive(Clone)]
21pub struct WorkerGrpcService {
22 state: ServerState,
23}
24
25impl WorkerGrpcService {
26 #[must_use]
28 pub const fn new(state: ServerState) -> Self {
29 Self { state }
30 }
31}
32
33#[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, ®ister, 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 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 let teardown = StreamTeardown {
120 worker_id,
121 heartbeat: &heartbeat,
122 registry: ®istry,
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 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
162struct 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
189fn 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 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 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 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 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 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 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
429fn 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
448fn 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}