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
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 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 let teardown = StreamTeardown {
127 worker_id,
128 heartbeat: &heartbeat,
129 registry: ®istry,
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 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
169struct 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
196fn 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 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 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 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 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 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 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
436fn 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
455fn 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}