1use aion_proto::{
4 ProtoActivityDescriptor, ProtoActivityResult, ProtoRegisterWorker, ProtoWorkerInstanceIdentity,
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;
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 validate_worker_contracts(&self.state, ®ister)?;
68
69 let (task_tx, task_rx) = mpsc::channel::<Result<generated::ServerToWorker, Status>>(32);
70 let (worker_tx, worker_rx) = mpsc::channel(32);
71
72 let registration = self
73 .state
74 .worker_registry()
75 .accept_registration(self.state.namespace_guard(), &caller, ®ister, worker_tx)
76 .await
77 .map_err(|error| status_from_server_error(&error))?;
78
79 let pending = self.state.pending_activities().clone();
80 let heartbeat = self.state.heartbeat_tracker().clone();
81 let drain = self.state.drain_state().clone();
82 let registry = self.state.worker_registry().clone();
83 let liveness_waiters = self.state.grpc_liveness_waiters().clone();
84 let worker_id = registration
85 .worker_id()
86 .ok_or_else(|| Status::internal("worker registration missing id"))?;
87 heartbeat
88 .register_connection(worker_id, std::time::Instant::now())
89 .map_err(|error| status_from_server_error(&error))?;
90 let authorized_namespace = registration
94 .namespaces()
95 .filter(|namespaces| !namespaces.is_empty())
96 .ok_or_else(|| Status::internal("worker registration missing namespace"))?
97 .iter()
98 .cloned()
99 .collect::<Vec<_>>()
100 .join(",");
101
102 task_tx
107 .try_send(Ok(register_ack_frame(
108 worker_id,
109 &authorized_namespace,
110 heartbeat_grace,
111 )))
112 .map_err(|_| Status::internal("worker response channel closed before RegisterAck"))?;
113
114 tokio::spawn(async move {
115 let write_handle = spawn_write_forwarder(worker_rx, task_tx.clone());
116
117 let teardown = StreamTeardown {
122 worker_id,
123 heartbeat: &heartbeat,
124 registry: ®istry,
125 pending: &pending,
126 drain: &drain,
127 liveness_waiters: &liveness_waiters,
128 };
129 let session = WorkerSession {
130 worker_id,
131 pending: &pending,
132 heartbeat: &heartbeat,
133 drain: &drain,
134 token_expires_at,
135 heartbeat_grace,
136 task_tx: task_tx.clone(),
137 liveness_waiters: liveness_waiters.clone(),
138 };
139 if let Err(status) = process_inbound(inbound, session).await {
140 tracing::info!(
141 worker_id = ?worker_id,
142 %status,
143 "worker stream closed with status"
144 );
145 }
146
147 write_handle.abort();
148 drop(task_tx);
149 drop(teardown);
150 if let Err(error) = registration.deregister() {
154 tracing::error!(
155 worker_id = ?worker_id,
156 %error,
157 "worker deregistration failed during stream teardown"
158 );
159 }
160 });
161
162 Ok(Response::new(ReceiverStream::new(task_rx)))
163 }
164}
165
166fn spawn_write_forwarder(
183 mut worker_rx: mpsc::Receiver<WorkerMessage>,
184 task_tx: mpsc::Sender<Result<generated::ServerToWorker, Status>>,
185) -> tokio::task::JoinHandle<()> {
186 tokio::spawn(async move {
187 while let Some(message) = worker_rx.recv().await {
188 let msg = encode_server_to_worker(message);
189 if task_tx.send(Ok(msg)).await.is_err() {
190 return;
191 }
192 }
193 let _ = task_tx
194 .send(Err(Status::unavailable(
195 "worker was deregistered by the server (heartbeat window expired); \
196 reconnect and re-register",
197 )))
198 .await;
199 })
200}
201
202struct StreamTeardown<'a> {
211 worker_id: WorkerId,
212 heartbeat: &'a crate::worker::HeartbeatTracker,
213 registry: &'a crate::worker::ConnectedWorkerRegistry,
214 pending: &'a PendingActivities,
215 drain: &'a crate::shutdown::DrainState,
216 liveness_waiters: &'a crate::worker::GrpcLivenessWaiters,
220}
221
222impl Drop for StreamTeardown<'_> {
223 fn drop(&mut self) {
224 teardown_worker_stream(
225 self.worker_id,
226 self.heartbeat,
227 self.registry,
228 self.pending,
229 self.drain,
230 );
231 if let Err(error) = self.liveness_waiters.disarm(self.worker_id) {
240 tracing::error!(
241 worker_id = ?self.worker_id,
242 %error,
243 "gRPC liveness waiter map is poisoned; armed pings for departed workers can no \
244 longer be released"
245 );
246 }
247 }
248}
249
250fn teardown_worker_stream(
270 worker_id: WorkerId,
271 heartbeat: &crate::worker::HeartbeatTracker,
272 registry: &crate::worker::ConnectedWorkerRegistry,
273 pending: &PendingActivities,
274 drain: &crate::shutdown::DrainState,
275) {
276 if drain.is_draining() {
277 match heartbeat.park_disconnected_worker(worker_id, registry, pending) {
278 Ok(report) if report.tasks.is_empty() => {}
279 Ok(report) => {
280 tracing::info!(
281 worker_id = ?worker_id,
282 parked_tasks = report.tasks.len(),
283 "worker stream ended during drain; in-flight activities \
284 parked for restart recovery"
285 );
286 }
287 Err(error) => {
288 tracing::error!(
289 worker_id = ?worker_id,
290 %error,
291 "failed to park draining worker's in-flight activities"
292 );
293 }
294 }
295 } else {
296 match heartbeat.fail_disconnected_worker(worker_id, registry, pending) {
297 Ok(report) if report.tasks.is_empty() => {}
298 Ok(report) => {
299 tracing::warn!(
300 worker_id = ?worker_id,
301 failed_tasks = report.tasks.len(),
302 "worker disconnected with in-flight activities; \
303 surfaced as transport losses, to be re-dispatched \
304 attempt-neutrally"
305 );
306 }
307 Err(error) => {
308 tracing::error!(
309 worker_id = ?worker_id,
310 %error,
311 "failed to sweep disconnected worker's in-flight activities"
312 );
313 }
314 }
315 }
316 drain.notify_activity_drained();
319}
320
321struct WorkerSession<'a> {
322 worker_id: WorkerId,
323 pending: &'a PendingActivities,
324 heartbeat: &'a crate::worker::HeartbeatTracker,
325 drain: &'a crate::shutdown::DrainState,
326 token_expires_at: Option<u64>,
327 heartbeat_grace: std::time::Duration,
328 task_tx: mpsc::Sender<Result<generated::ServerToWorker, Status>>,
329 liveness_waiters: crate::worker::GrpcLivenessWaiters,
333}
334
335async fn process_inbound(
336 mut inbound: Streaming<generated::WorkerToServer>,
337 session: WorkerSession<'_>,
338) -> Result<(), Status> {
339 let mut expired_since: Option<std::time::Instant> = None;
340 while let Some(msg) = inbound.message().await? {
341 refresh_connection_lease(&session)?;
342 let Some(inner) = msg.message else {
343 continue;
344 };
345 match inner {
346 generated::worker_to_server::Message::Result(result) => {
347 let proto_result = decode_activity_result(result);
348 match ActivityCompletion::try_from(proto_result) {
349 Ok(completion) => {
350 let workflow_id = completion.workflow_id.clone();
351 let activity_id = completion.activity_id.clone();
352 let after_accept = || {
353 let _ = crate::worker::bridge::clear_completed_task_tracking(
359 session.heartbeat,
360 session.worker_id,
361 &workflow_id,
362 &activity_id,
363 );
364 session.drain.notify_activity_drained();
365 Ok(())
366 };
367 if let Err(error) = session
368 .pending
369 .complete_activity_after_accept(completion, after_accept)
370 {
371 tracing::error!(
378 worker_id = ?session.worker_id,
379 workflow_id = %workflow_id,
380 activity_id = %activity_id,
381 %error,
382 "activity completion rejected by execution-generation proof"
383 );
384 }
385 let ack = result_ack_frame(&workflow_id, &activity_id);
392 if let Err(error) = session.task_tx.try_send(Ok(ack)) {
393 tracing::warn!(
394 worker_id = ?session.worker_id,
395 workflow_id = %workflow_id,
396 activity_id = %activity_id,
397 %error,
398 "result ack dropped: worker stream channel unavailable"
399 );
400 }
401 }
402 Err(error) => {
403 tracing::error!(
407 worker_id = ?session.worker_id,
408 %error,
409 "malformed activity result frame; no ack sent"
410 );
411 }
412 }
413 }
414 generated::worker_to_server::Message::Register(_) => {
415 warn_duplicate_registration(session.worker_id);
416 }
417 generated::worker_to_server::Message::LivenessAnswer(answer) => {
423 deliver_liveness_answer(&session, answer.liveness_ping);
424 }
425 generated::worker_to_server::Message::Heartbeat(heartbeat_msg) => {
426 if heartbeat_msg.workflow_id.is_none() && heartbeat_msg.activity_id.is_none() {
430 continue;
431 }
432 if let Err(error) = session.heartbeat.record_heartbeat(
433 session.worker_id,
434 decode_heartbeat(heartbeat_msg),
435 std::time::Instant::now(),
436 ) {
437 if matches!(error, crate::ServerError::LockPoisoned { .. }) {
442 tracing::error!(
443 worker_id = ?session.worker_id,
444 %error,
445 "heartbeat tracker lock poisoned; liveness state untrustworthy"
446 );
447 } else {
448 tracing::warn!(
449 worker_id = ?session.worker_id,
450 %error,
451 "worker heartbeat rejected"
452 );
453 }
454 }
455 enforce_token_expiration(&session, &mut expired_since).await?;
456 }
457 }
458 }
459 Ok(())
460}
461
462fn deliver_liveness_answer(session: &WorkerSession<'_>, sequence: u64) {
472 match session.liveness_waiters.answer(session.worker_id, sequence) {
473 Ok(true) => {}
474 Ok(false) => tracing::warn!(
475 worker_id = ?session.worker_id,
476 liveness_ping = sequence,
477 "worker answered a liveness ping the server was no longer waiting for; the answer \
478 arrived after its probe cadence expired, or echoed a sequence that was never asked. \
479 It banks NOTHING toward the dispatch probation"
480 ),
481 Err(error) => tracing::error!(
482 worker_id = ?session.worker_id,
483 liveness_ping = sequence,
484 %error,
485 "gRPC liveness waiter map is poisoned; no gRPC worker on this server can clear its \
486 dispatch probation until the process is restarted"
487 ),
488 }
489}
490
491fn refresh_connection_lease(session: &WorkerSession<'_>) -> Result<(), Status> {
492 session
493 .heartbeat
494 .record_connection_activity(session.worker_id, std::time::Instant::now())
495 .map(|_| ())
496 .map_err(|error| {
497 tracing::error!(
498 worker_id = ?session.worker_id,
499 %error,
500 "failed to advance worker connection lease"
501 );
502 status_from_server_error(&error)
503 })
504}
505
506fn warn_duplicate_registration(worker_id: WorkerId) {
507 tracing::warn!(
508 worker_id = ?worker_id,
509 "ignoring subsequent RegisterWorker message; \
510 only the first registration is accepted per stream"
511 );
512}
513
514async fn enforce_token_expiration(
515 session: &WorkerSession<'_>,
516 expired_since: &mut Option<std::time::Instant>,
517) -> Result<(), Status> {
518 if !token_expired(session.token_expires_at) {
519 return Ok(());
520 }
521 let first_expired = *expired_since.get_or_insert_with(std::time::Instant::now);
522 let _ = session
523 .task_tx
524 .send(Err(Status::unauthenticated(
525 "worker token expired; re-authentication required",
526 )))
527 .await;
528 if first_expired.elapsed() >= session.heartbeat_grace {
529 return Err(Status::unauthenticated("worker token expired"));
530 }
531 Ok(())
532}
533
534async fn worker_caller_from_metadata(
535 metadata: &tonic::metadata::MetadataMap,
536 state: &ServerState,
537) -> Result<CallerIdentity, Status> {
538 crate::api::grpc::caller_from_metadata(metadata, state).await
539}
540
541async fn token_expiration_from_metadata(
542 metadata: &tonic::metadata::MetadataMap,
543 state: &ServerState,
544) -> Result<Option<u64>, Status> {
545 if !state.runtime_config().auth.enabled {
546 return Ok(None);
547 }
548 #[cfg(feature = "auth")]
549 {
550 let bearer = metadata
551 .get("authorization")
552 .and_then(|value| value.to_str().ok())
553 .and_then(parse_bearer)
554 .ok_or_else(|| Status::unauthenticated("missing bearer token"))?;
555 let Some(cache) = state.jwks_cache() else {
556 return Err(Status::unauthenticated("invalid bearer token"));
557 };
558 return cache
559 .validate(&bearer)
560 .await
561 .map(|claims| Some(claims.expires_at()))
562 .map_err(|_error| Status::unauthenticated("invalid bearer token"));
563 }
564 #[cfg(not(feature = "auth"))]
565 {
566 let _ = metadata;
567 tokio::task::yield_now().await;
569 Ok(None)
570 }
571}
572
573#[cfg(feature = "auth")]
574fn parse_bearer(value: &str) -> Option<String> {
575 let token = value.strip_prefix("Bearer ")?.trim();
576 if token.is_empty() {
577 return None;
578 }
579 Some(token.to_owned())
580}
581
582fn token_expired(expires_at: Option<u64>) -> bool {
583 expires_at.is_some_and(|expires_at| {
584 #[cfg(feature = "auth")]
585 {
586 crate::auth::jwks::is_expired(expires_at)
587 }
588 #[cfg(not(feature = "auth"))]
589 {
590 let _ = expires_at;
591 false
592 }
593 })
594}
595
596fn status_from_server_error(error: &crate::ServerError) -> Status {
597 let wire = error.to_wire_error();
598 if wire.code == aion_proto::WireErrorCode::NamespaceDenied {
599 Status::permission_denied(wire.message)
600 } else {
601 Status::internal(wire.message)
602 }
603}
604
605fn register_ack_frame(
608 worker_id: WorkerId,
609 namespace: &str,
610 heartbeat_window: std::time::Duration,
611) -> generated::ServerToWorker {
612 generated::ServerToWorker {
613 message: Some(generated::server_to_worker::Message::RegisterAck(
614 generated::RegisterAck {
615 worker_id: worker_id.value(),
616 namespace: namespace.to_owned(),
617 heartbeat_window_ms: u64::try_from(heartbeat_window.as_millis())
618 .unwrap_or(u64::MAX),
619 },
620 )),
621 }
622}
623
624fn result_ack_frame(
626 workflow_id: &aion_core::WorkflowId,
627 activity_id: &aion_core::ActivityId,
628) -> generated::ServerToWorker {
629 generated::ServerToWorker {
630 message: Some(generated::server_to_worker::Message::ResultAck(
631 generated::ResultAck {
632 workflow_id: Some(generated::WorkflowId {
633 uuid: workflow_id.to_string(),
634 }),
635 activity_id: Some(generated::ActivityId {
636 sequence_position: activity_id.sequence_position(),
637 }),
638 },
639 )),
640 }
641}
642
643fn decode_register(r: generated::RegisterWorker) -> ProtoRegisterWorker {
644 ProtoRegisterWorker {
645 namespaces: r.namespaces,
646 activity_types: r.activity_types,
647 task_queue: r.task_queue,
648 node: r.node,
649 activities: r
650 .activities
651 .into_iter()
652 .map(|activity| ProtoActivityDescriptor {
653 name: activity.name,
654 input_schema_json: activity.input_schema_json,
655 output_schema_json: activity.output_schema_json,
656 })
657 .collect(),
658 identity: r.identity,
659 instance: r.instance.map(|instance| ProtoWorkerInstanceIdentity {
660 deployment: instance.deployment,
661 instance_id: instance.instance_id,
662 }),
663 }
664}
665
666fn validate_worker_contracts(
667 state: &ServerState,
668 register: &ProtoRegisterWorker,
669) -> Result<(), Status> {
670 let advertised = register
671 .activities
672 .iter()
673 .map(|activity| {
674 let input_schema =
675 serde_json::from_str(&activity.input_schema_json).map_err(|error| {
676 Status::invalid_argument(format!(
677 "worker activity `{}` input_schema_json is invalid: {error}",
678 activity.name
679 ))
680 })?;
681 let output_schema =
682 serde_json::from_str(&activity.output_schema_json).map_err(|error| {
683 Status::invalid_argument(format!(
684 "worker activity `{}` output_schema_json is invalid: {error}",
685 activity.name
686 ))
687 })?;
688 Ok(aion_package::ActivityDescriptor {
689 name: activity.name.clone(),
690 input_schema,
691 output_schema,
692 })
693 })
694 .collect::<Result<Vec<_>, Status>>()?;
695 let Ok(engine) = state.engine() else {
701 tracing::warn!(
702 task_queue = %register.task_queue,
703 identity = %register.identity,
704 "worker contract check skipped: server state has no engine handle, \
705 so no deployed contracts exist to check against"
706 );
707 return Ok(());
708 };
709 let activity_types = register
713 .activity_types
714 .iter()
715 .cloned()
716 .collect::<std::collections::BTreeSet<_>>();
717 crate::worker::contracts::validate_worker_contracts(
718 &engine,
719 state.worker_registry().admission_audit(),
720 ®ister.task_queue,
721 crate::worker::registry::optional_node(®ister.node).as_deref(),
722 ®ister.identity,
723 crate::worker::contracts::WorkerAdvertisement {
724 activity_types: &activity_types,
725 contracts: &advertised,
726 },
727 )
728 .map_err(|error| match error {
729 crate::worker::contracts::ContractAdmissionError::Mismatch { .. } => {
730 Status::failed_precondition(error.to_string())
731 }
732 crate::worker::contracts::ContractAdmissionError::Catalog { .. } => {
733 Status::internal(error.to_string())
734 }
735 })
736}
737
738fn encode_server_to_worker(message: WorkerMessage) -> generated::ServerToWorker {
739 let message = match message {
740 WorkerMessage::ActivityTask(task) => {
741 generated::server_to_worker::Message::Task(encode_task(*task))
742 }
743 WorkerMessage::DrainRequest => {
744 generated::server_to_worker::Message::Drain(generated::DrainRequest {})
745 }
746 WorkerMessage::LivenessPing(ping) => {
750 generated::server_to_worker::Message::LivenessPing(generated::LivenessPing {
751 liveness_ping: ping.liveness_ping,
752 silence_window_ms: ping.silence_window_ms,
753 })
754 }
755 WorkerMessage::CancelActivity(cancel) => {
759 generated::server_to_worker::Message::CancelActivity(generated::CancelActivity {
760 workflow_id: cancel
761 .workflow_id
762 .map(|id| generated::WorkflowId { uuid: id.uuid }),
763 activity_id: cancel.activity_id.map(|id| generated::ActivityId {
764 sequence_position: id.sequence_position,
765 }),
766 })
767 }
768 };
769 generated::ServerToWorker {
770 message: Some(message),
771 }
772}
773
774fn encode_task(task: aion_proto::ProtoActivityTask) -> generated::ActivityTask {
775 generated::ActivityTask {
776 workflow_id: task
777 .workflow_id
778 .map(|id| generated::WorkflowId { uuid: id.uuid }),
779 activity_id: task.activity_id.map(|id| generated::ActivityId {
780 sequence_position: id.sequence_position,
781 }),
782 activity_type: task.activity_type,
783 input: task.input.map(|p| generated::Payload {
784 content_type: p.content_type,
785 bytes: p.bytes,
786 }),
787 attempt: task.attempt,
788 labels: task.labels,
789 run_id: task.run_id.map(|id| generated::RunId { uuid: id.uuid }),
790 completion_token: task.completion_token,
791 idempotency_key: task.idempotency_key,
792 }
793}
794
795fn decode_activity_result(r: generated::ActivityResult) -> ProtoActivityResult {
796 ProtoActivityResult {
797 workflow_id: r
798 .workflow_id
799 .map(|id| aion_proto::ProtoWorkflowId { uuid: id.uuid }),
800 activity_id: r.activity_id.map(|id| aion_proto::ProtoActivityId {
801 sequence_position: id.sequence_position,
802 }),
803 outcome: r.outcome.map(decode_outcome),
804 run_id: r.run_id.map(|id| aion_proto::ProtoRunId { uuid: id.uuid }),
805 completion_token: r.completion_token,
806 }
807}
808
809fn decode_heartbeat(r: generated::Heartbeat) -> aion_proto::ProtoHeartbeat {
810 aion_proto::ProtoHeartbeat {
811 workflow_id: r
812 .workflow_id
813 .map(|id| aion_proto::ProtoWorkflowId { uuid: id.uuid }),
814 activity_id: r.activity_id.map(|id| aion_proto::ProtoActivityId {
815 sequence_position: id.sequence_position,
816 }),
817 progress: r.progress.map(|p| aion_proto::ProtoPayload {
818 content_type: p.content_type,
819 bytes: p.bytes,
820 }),
821 }
822}
823
824fn decode_outcome(
825 outcome: generated::activity_result::Outcome,
826) -> aion_proto::proto_activity_result::Outcome {
827 match outcome {
828 generated::activity_result::Outcome::Result(p) => {
829 aion_proto::proto_activity_result::Outcome::Result(aion_proto::ProtoPayload {
830 content_type: p.content_type,
831 bytes: p.bytes,
832 })
833 }
834 generated::activity_result::Outcome::Error(e) => {
835 aion_proto::proto_activity_result::Outcome::Error(aion_proto::ProtoActivityError {
836 kind: e.kind,
837 message: e.message,
838 details: e.details.map(|p| aion_proto::ProtoPayload {
839 content_type: p.content_type,
840 bytes: p.bytes,
841 }),
842 })
843 }
844 }
845}
846
847#[cfg(test)]
848mod tests {
849 use std::time::{Duration, Instant};
850
851 use aion_core::{ActivityId, ContentType, Payload, RunId, WorkflowId};
852
853 use crate::shutdown::DrainState;
854 use crate::worker::dispatch::{
855 ActivityCompletion, ActivityCompletionOutcome, ActivityCompletionSink,
856 };
857 use crate::worker::heartbeat::InFlightActivity;
858 use crate::worker::registry::ConnectedWorkerRegistry;
859 use crate::worker::{HeartbeatTracker, PendingActivities};
860
861 use super::{decode_register, teardown_worker_stream};
862
863 type TestError = Box<dyn std::error::Error>;
864
865 #[test]
866 fn decode_register_maps_tag_seven_instance_without_changing_absent_registration() {
867 let generated = super::generated::RegisterWorker {
868 namespaces: vec!["orders".to_owned()],
869 activity_types: vec!["shell".to_owned()],
870 task_queue: "shell".to_owned(),
871 node: "node-a".to_owned(),
872 activities: Vec::new(),
873 identity: "build-a".to_owned(),
874 instance: Some(super::generated::WorkerInstanceIdentity {
875 deployment: "shells".to_owned(),
876 instance_id: "instance-1".to_owned(),
877 }),
878 };
879 let mapped = decode_register(generated.clone());
880 let instance = mapped.instance.as_ref();
881 assert_eq!(
882 instance.map(|value| value.deployment.as_str()),
883 Some("shells")
884 );
885 assert_eq!(
886 instance.map(|value| value.instance_id.as_str()),
887 Some("instance-1")
888 );
889
890 let mut absent = generated;
891 absent.instance = None;
892 let mapped_absent = decode_register(absent);
893 assert!(mapped_absent.instance.is_none());
894 assert_eq!(mapped_absent.identity, "build-a");
895 }
896
897 struct TeardownFixture {
901 registry: ConnectedWorkerRegistry,
902 tracker: HeartbeatTracker,
903 pending: PendingActivities,
904 drain: DrainState,
905 worker_id: crate::worker::registry::WorkerId,
906 workflow_id: WorkflowId,
907 run_id: RunId,
911 activity_id: ActivityId,
912 completion_token: crate::worker::CompletionToken,
913 rx: std::sync::mpsc::Receiver<Result<String, String>>,
914 _registration: crate::worker::registry::WorkerRegistration,
917 }
918
919 fn fixture() -> Result<TeardownFixture, TestError> {
920 let registry = ConnectedWorkerRegistry::default();
921 let (tx, _rx) = tokio::sync::mpsc::channel(1);
922 let activity_types = [String::from("greet")];
923 let registration = registry.register("default", activity_types.iter(), tx)?;
924 let worker_id = registration
925 .worker_id()
926 .ok_or("test worker registration missing id")?;
927 let tracker = HeartbeatTracker::new(Duration::from_secs(5));
928 let pending = PendingActivities::new(Duration::from_secs(5));
929 let workflow_id = WorkflowId::new_v4();
930 let run_id = RunId::new_v4();
931 let activity_id = ActivityId::from_sequence_position(0);
932 let (completion_token, rx) =
933 pending.insert_for_test(workflow_id.clone(), &run_id, activity_id.clone(), 1)?;
934 tracker.track_task(
935 worker_id,
936 InFlightActivity {
937 workflow_id: workflow_id.clone(),
938 activity_id: activity_id.clone(),
939 attempt: 1,
940 completion_token: completion_token.clone(),
941 },
942 Instant::now(),
943 )?;
944 Ok(TeardownFixture {
945 registry,
946 tracker,
947 pending,
948 drain: DrainState::default(),
949 worker_id,
950 workflow_id,
951 run_id,
952 activity_id,
953 completion_token,
954 rx,
955 _registration: registration,
956 })
957 }
958
959 #[test]
963 fn teardown_under_drain_parks_instead_of_failing() -> Result<(), TestError> {
964 let fixture = fixture()?;
965 assert!(fixture.drain.begin());
966
967 teardown_worker_stream(
968 fixture.worker_id,
969 &fixture.tracker,
970 &fixture.registry,
971 &fixture.pending,
972 &fixture.drain,
973 );
974
975 let resolved = fixture.rx.recv_timeout(Duration::from_millis(200))?;
976 assert_eq!(
977 resolved,
978 Err(aion::PARKED_ACTIVITY_REASON.to_owned()),
979 "a drain teardown must resolve the waiter with the parked sentinel"
980 );
981 assert_eq!(fixture.tracker.in_flight_count()?, 0);
982 assert!(
983 !fixture.tracker.is_tracked(
984 fixture.worker_id,
985 &fixture.workflow_id,
986 &fixture.activity_id
987 )?,
988 "parking must retire the tracked entry"
989 );
990 Ok(())
991 }
992
993 #[test]
1003 fn teardown_without_drain_fails_with_the_transport_domain_lost_worker_class()
1004 -> Result<(), TestError> {
1005 let fixture = fixture()?;
1006
1007 teardown_worker_stream(
1008 fixture.worker_id,
1009 &fixture.tracker,
1010 &fixture.registry,
1011 &fixture.pending,
1012 &fixture.drain,
1013 );
1014
1015 let resolved = fixture.rx.recv_timeout(Duration::from_millis(200))?;
1016 let reason = resolved.err().ok_or("expected a lost-worker failure")?;
1017 assert!(
1018 reason.starts_with(crate::worker::WORKER_LOST_REASON_PREFIX),
1019 "a mid-run teardown must surface the TRANSPORT-domain loss class, never the \
1020 action's retry vocabulary: {reason}"
1021 );
1022 assert!(
1023 reason.contains("lost before reporting activity result"),
1024 "the failure must name worker loss: {reason}"
1025 );
1026 assert_eq!(fixture.tracker.in_flight_count()?, 0);
1027 Ok(())
1028 }
1029
1030 #[test]
1034 fn stale_worker_completion_after_heartbeat_loss_does_not_resolve_retry() -> Result<(), TestError>
1035 {
1036 let fixture = fixture()?;
1037 teardown_worker_stream(
1038 fixture.worker_id,
1039 &fixture.tracker,
1040 &fixture.registry,
1041 &fixture.pending,
1042 &fixture.drain,
1043 );
1044 let first = fixture.rx.recv_timeout(Duration::from_millis(200))?;
1045 assert!(
1046 first
1047 .err()
1048 .is_some_and(|reason| reason.starts_with(crate::worker::WORKER_LOST_REASON_PREFIX)),
1049 "worker A loss must release attempt 1 in the transport-loss class"
1050 );
1051
1052 let (retry_token, retry_rx) = fixture.pending.insert_for_test(
1055 fixture.workflow_id.clone(),
1056 &fixture.run_id,
1057 fixture.activity_id.clone(),
1058 2,
1059 )?;
1060 let rejected = fixture.pending.complete_activity(ActivityCompletion {
1061 workflow_id: fixture.workflow_id,
1062 activity_id: fixture.activity_id,
1063 run_id: None,
1064 completion_token: fixture.completion_token,
1065 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
1066 ContentType::Json,
1067 br#"{"worker":"A","stale":true}"#.to_vec(),
1068 )),
1069 });
1070
1071 assert!(matches!(
1072 rejected,
1073 Err(crate::ServerError::ActivityCompletionRejected { .. })
1074 ));
1075 drop(retry_token);
1076 assert!(
1077 retry_rx.recv_timeout(Duration::from_millis(50)).is_err(),
1078 "worker A's late completion must be rejected instead of resolving worker B's retry"
1079 );
1080 Ok(())
1081 }
1082
1083 #[test]
1101 fn a_worker_loss_consumes_the_sibling_token_of_a_redelivered_attempt() -> Result<(), TestError>
1102 {
1103 let fixture = fixture()?;
1104 let sibling = fixture.pending.completion_fences().issue(
1108 &fixture.workflow_id,
1109 &fixture.run_id,
1110 &fixture.activity_id,
1111 1,
1112 )?;
1113
1114 teardown_worker_stream(
1115 fixture.worker_id,
1116 &fixture.tracker,
1117 &fixture.registry,
1118 &fixture.pending,
1119 &fixture.drain,
1120 );
1121
1122 let resolved = fixture.rx.recv_timeout(Duration::from_millis(200))?;
1123 let reason = resolved.err().ok_or("expected a lost-worker failure")?;
1124 assert!(
1125 reason.starts_with(crate::worker::WORKER_LOST_REASON_PREFIX),
1126 "the loss must resolve in the transport domain, which re-dispatches the same attempt \
1127 rather than faulting the action: {reason}"
1128 );
1129
1130 let refused = fixture.pending.complete_activity(ActivityCompletion {
1131 workflow_id: fixture.workflow_id.clone(),
1132 activity_id: fixture.activity_id.clone(),
1133 run_id: None,
1134 completion_token: sibling,
1135 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
1136 ContentType::Json,
1137 br#"{"worker":"B"}"#.to_vec(),
1138 )),
1139 });
1140 assert!(
1141 matches!(
1142 refused,
1143 Err(crate::ServerError::ActivityCompletionRejected {
1144 reason: crate::error::CompletionRejectionReason::NoCurrentGeneration,
1145 ..
1146 })
1147 ),
1148 "the first accepted completion — here the loss verdict — consumes the whole \
1149 generation, so the sibling worker's result is the duplicate: {refused:?}"
1150 );
1151 Ok(())
1152 }
1153}