1use aion_proto::{
4 ProtoActivityDescriptor, 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 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 worker_id = registration
84 .worker_id()
85 .ok_or_else(|| Status::internal("worker registration missing id"))?;
86 heartbeat
87 .register_connection(worker_id, std::time::Instant::now())
88 .map_err(|error| status_from_server_error(&error))?;
89 let authorized_namespace = registration
93 .namespaces()
94 .filter(|namespaces| !namespaces.is_empty())
95 .ok_or_else(|| Status::internal("worker registration missing namespace"))?
96 .iter()
97 .cloned()
98 .collect::<Vec<_>>()
99 .join(",");
100
101 task_tx
106 .try_send(Ok(register_ack_frame(
107 worker_id,
108 &authorized_namespace,
109 heartbeat_grace,
110 )))
111 .map_err(|_| Status::internal("worker response channel closed before RegisterAck"))?;
112
113 tokio::spawn(async move {
114 let write_handle = spawn_write_forwarder(worker_rx, task_tx.clone());
115
116 let teardown = StreamTeardown {
121 worker_id,
122 heartbeat: &heartbeat,
123 registry: ®istry,
124 pending: &pending,
125 drain: &drain,
126 };
127 let session = WorkerSession {
128 worker_id,
129 pending: &pending,
130 heartbeat: &heartbeat,
131 drain: &drain,
132 token_expires_at,
133 heartbeat_grace,
134 task_tx: task_tx.clone(),
135 };
136 if let Err(status) = process_inbound(inbound, session).await {
137 tracing::info!(
138 worker_id = ?worker_id,
139 %status,
140 "worker stream closed with status"
141 );
142 }
143
144 write_handle.abort();
145 drop(task_tx);
146 drop(teardown);
147 if let Err(error) = registration.deregister() {
151 tracing::error!(
152 worker_id = ?worker_id,
153 %error,
154 "worker deregistration failed during stream teardown"
155 );
156 }
157 });
158
159 Ok(Response::new(ReceiverStream::new(task_rx)))
160 }
161}
162
163fn spawn_write_forwarder(
180 mut worker_rx: mpsc::Receiver<WorkerMessage>,
181 task_tx: mpsc::Sender<Result<generated::ServerToWorker, Status>>,
182) -> tokio::task::JoinHandle<()> {
183 tokio::spawn(async move {
184 while let Some(message) = worker_rx.recv().await {
185 let msg = encode_server_to_worker(message);
186 if task_tx.send(Ok(msg)).await.is_err() {
187 return;
188 }
189 }
190 let _ = task_tx
191 .send(Err(Status::unavailable(
192 "worker was deregistered by the server (heartbeat window expired); \
193 reconnect and re-register",
194 )))
195 .await;
196 })
197}
198
199struct StreamTeardown<'a> {
208 worker_id: WorkerId,
209 heartbeat: &'a crate::worker::HeartbeatTracker,
210 registry: &'a crate::worker::ConnectedWorkerRegistry,
211 pending: &'a PendingActivities,
212 drain: &'a crate::shutdown::DrainState,
213}
214
215impl Drop for StreamTeardown<'_> {
216 fn drop(&mut self) {
217 teardown_worker_stream(
218 self.worker_id,
219 self.heartbeat,
220 self.registry,
221 self.pending,
222 self.drain,
223 );
224 }
225}
226
227fn teardown_worker_stream(
247 worker_id: WorkerId,
248 heartbeat: &crate::worker::HeartbeatTracker,
249 registry: &crate::worker::ConnectedWorkerRegistry,
250 pending: &PendingActivities,
251 drain: &crate::shutdown::DrainState,
252) {
253 if drain.is_draining() {
254 match heartbeat.park_disconnected_worker(worker_id, registry, pending) {
255 Ok(report) if report.tasks.is_empty() => {}
256 Ok(report) => {
257 tracing::info!(
258 worker_id = ?worker_id,
259 parked_tasks = report.tasks.len(),
260 "worker stream ended during drain; in-flight activities \
261 parked for restart recovery"
262 );
263 }
264 Err(error) => {
265 tracing::error!(
266 worker_id = ?worker_id,
267 %error,
268 "failed to park draining worker's in-flight activities"
269 );
270 }
271 }
272 } else {
273 match heartbeat.fail_disconnected_worker(worker_id, registry, pending) {
274 Ok(report) if report.tasks.is_empty() => {}
275 Ok(report) => {
276 tracing::warn!(
277 worker_id = ?worker_id,
278 failed_tasks = report.tasks.len(),
279 "worker disconnected with in-flight activities; \
280 surfaced as transport losses, to be re-dispatched \
281 attempt-neutrally"
282 );
283 }
284 Err(error) => {
285 tracing::error!(
286 worker_id = ?worker_id,
287 %error,
288 "failed to sweep disconnected worker's in-flight activities"
289 );
290 }
291 }
292 }
293 drain.notify_activity_drained();
296}
297
298struct WorkerSession<'a> {
299 worker_id: WorkerId,
300 pending: &'a PendingActivities,
301 heartbeat: &'a crate::worker::HeartbeatTracker,
302 drain: &'a crate::shutdown::DrainState,
303 token_expires_at: Option<u64>,
304 heartbeat_grace: std::time::Duration,
305 task_tx: mpsc::Sender<Result<generated::ServerToWorker, Status>>,
306}
307
308async fn process_inbound(
309 mut inbound: Streaming<generated::WorkerToServer>,
310 session: WorkerSession<'_>,
311) -> Result<(), Status> {
312 let mut expired_since: Option<std::time::Instant> = None;
313 while let Some(msg) = inbound.message().await? {
314 refresh_connection_lease(&session)?;
315 let Some(inner) = msg.message else {
316 continue;
317 };
318 match inner {
319 generated::worker_to_server::Message::Result(result) => {
320 let proto_result = decode_activity_result(result);
321 match ActivityCompletion::try_from(proto_result) {
322 Ok(completion) => {
323 let workflow_id = completion.workflow_id.clone();
324 let activity_id = completion.activity_id.clone();
325 match session.pending.complete_activity(completion) {
326 Ok(()) => {
327 if let Err(error) = session.heartbeat.complete_task(
328 session.worker_id,
329 &workflow_id,
330 &activity_id,
331 ) {
332 tracing::error!(
336 worker_id = ?session.worker_id,
337 workflow_id = %workflow_id,
338 activity_id = %activity_id,
339 %error,
340 "failed to clear in-flight tracking for completed activity"
341 );
342 }
343 session.drain.notify_activity_drained();
344 }
345 Err(error) => {
346 tracing::error!(
350 worker_id = ?session.worker_id,
351 workflow_id = %workflow_id,
352 activity_id = %activity_id,
353 %error,
354 "activity completion handoff failed"
355 );
356 }
357 }
358 let ack = result_ack_frame(&workflow_id, &activity_id);
365 if let Err(error) = session.task_tx.try_send(Ok(ack)) {
366 tracing::warn!(
367 worker_id = ?session.worker_id,
368 workflow_id = %workflow_id,
369 activity_id = %activity_id,
370 %error,
371 "result ack dropped: worker stream channel unavailable"
372 );
373 }
374 }
375 Err(error) => {
376 tracing::error!(
380 worker_id = ?session.worker_id,
381 %error,
382 "malformed activity result frame; no ack sent"
383 );
384 }
385 }
386 }
387 generated::worker_to_server::Message::Register(_) => {
388 warn_duplicate_registration(session.worker_id);
389 }
390 generated::worker_to_server::Message::Heartbeat(heartbeat_msg) => {
391 if heartbeat_msg.workflow_id.is_none() && heartbeat_msg.activity_id.is_none() {
395 continue;
396 }
397 if let Err(error) = session.heartbeat.record_heartbeat(
398 session.worker_id,
399 decode_heartbeat(heartbeat_msg),
400 std::time::Instant::now(),
401 ) {
402 if matches!(error, crate::ServerError::LockPoisoned { .. }) {
407 tracing::error!(
408 worker_id = ?session.worker_id,
409 %error,
410 "heartbeat tracker lock poisoned; liveness state untrustworthy"
411 );
412 } else {
413 tracing::warn!(
414 worker_id = ?session.worker_id,
415 %error,
416 "worker heartbeat rejected"
417 );
418 }
419 }
420 enforce_token_expiration(&session, &mut expired_since).await?;
421 }
422 }
423 }
424 Ok(())
425}
426
427fn refresh_connection_lease(session: &WorkerSession<'_>) -> Result<(), Status> {
428 session
429 .heartbeat
430 .record_connection_activity(session.worker_id, std::time::Instant::now())
431 .map(|_| ())
432 .map_err(|error| {
433 tracing::error!(
434 worker_id = ?session.worker_id,
435 %error,
436 "failed to advance worker connection lease"
437 );
438 status_from_server_error(&error)
439 })
440}
441
442fn warn_duplicate_registration(worker_id: WorkerId) {
443 tracing::warn!(
444 worker_id = ?worker_id,
445 "ignoring subsequent RegisterWorker message; \
446 only the first registration is accepted per stream"
447 );
448}
449
450async fn enforce_token_expiration(
451 session: &WorkerSession<'_>,
452 expired_since: &mut Option<std::time::Instant>,
453) -> Result<(), Status> {
454 if !token_expired(session.token_expires_at) {
455 return Ok(());
456 }
457 let first_expired = *expired_since.get_or_insert_with(std::time::Instant::now);
458 let _ = session
459 .task_tx
460 .send(Err(Status::unauthenticated(
461 "worker token expired; re-authentication required",
462 )))
463 .await;
464 if first_expired.elapsed() >= session.heartbeat_grace {
465 return Err(Status::unauthenticated("worker token expired"));
466 }
467 Ok(())
468}
469
470async fn worker_caller_from_metadata(
471 metadata: &tonic::metadata::MetadataMap,
472 state: &ServerState,
473) -> Result<CallerIdentity, Status> {
474 crate::api::grpc::caller_from_metadata(metadata, state).await
475}
476
477async fn token_expiration_from_metadata(
478 metadata: &tonic::metadata::MetadataMap,
479 state: &ServerState,
480) -> Result<Option<u64>, Status> {
481 if !state.runtime_config().auth.enabled {
482 return Ok(None);
483 }
484 #[cfg(feature = "auth")]
485 {
486 let bearer = metadata
487 .get("authorization")
488 .and_then(|value| value.to_str().ok())
489 .and_then(parse_bearer)
490 .ok_or_else(|| Status::unauthenticated("missing bearer token"))?;
491 let Some(cache) = state.jwks_cache() else {
492 return Err(Status::unauthenticated("invalid bearer token"));
493 };
494 return cache
495 .validate(&bearer)
496 .await
497 .map(|claims| Some(claims.expires_at()))
498 .map_err(|_error| Status::unauthenticated("invalid bearer token"));
499 }
500 #[cfg(not(feature = "auth"))]
501 {
502 let _ = metadata;
503 tokio::task::yield_now().await;
505 Ok(None)
506 }
507}
508
509#[cfg(feature = "auth")]
510fn parse_bearer(value: &str) -> Option<String> {
511 let token = value.strip_prefix("Bearer ")?.trim();
512 if token.is_empty() {
513 return None;
514 }
515 Some(token.to_owned())
516}
517
518fn token_expired(expires_at: Option<u64>) -> bool {
519 expires_at.is_some_and(|expires_at| {
520 #[cfg(feature = "auth")]
521 {
522 crate::auth::jwks::is_expired(expires_at)
523 }
524 #[cfg(not(feature = "auth"))]
525 {
526 let _ = expires_at;
527 false
528 }
529 })
530}
531
532fn status_from_server_error(error: &crate::ServerError) -> Status {
533 let wire = error.to_wire_error();
534 if wire.code == aion_proto::WireErrorCode::NamespaceDenied {
535 Status::permission_denied(wire.message)
536 } else {
537 Status::internal(wire.message)
538 }
539}
540
541fn register_ack_frame(
544 worker_id: WorkerId,
545 namespace: &str,
546 heartbeat_window: std::time::Duration,
547) -> generated::ServerToWorker {
548 generated::ServerToWorker {
549 message: Some(generated::server_to_worker::Message::RegisterAck(
550 generated::RegisterAck {
551 worker_id: worker_id.value(),
552 namespace: namespace.to_owned(),
553 heartbeat_window_ms: u64::try_from(heartbeat_window.as_millis())
554 .unwrap_or(u64::MAX),
555 },
556 )),
557 }
558}
559
560fn result_ack_frame(
562 workflow_id: &aion_core::WorkflowId,
563 activity_id: &aion_core::ActivityId,
564) -> generated::ServerToWorker {
565 generated::ServerToWorker {
566 message: Some(generated::server_to_worker::Message::ResultAck(
567 generated::ResultAck {
568 workflow_id: Some(generated::WorkflowId {
569 uuid: workflow_id.to_string(),
570 }),
571 activity_id: Some(generated::ActivityId {
572 sequence_position: activity_id.sequence_position(),
573 }),
574 },
575 )),
576 }
577}
578
579fn decode_register(r: generated::RegisterWorker) -> ProtoRegisterWorker {
580 ProtoRegisterWorker {
581 namespaces: r.namespaces,
582 activity_types: r.activity_types,
583 task_queue: r.task_queue,
584 node: r.node,
585 activities: r
586 .activities
587 .into_iter()
588 .map(|activity| ProtoActivityDescriptor {
589 name: activity.name,
590 input_schema_json: activity.input_schema_json,
591 output_schema_json: activity.output_schema_json,
592 })
593 .collect(),
594 identity: r.identity,
595 }
596}
597
598fn validate_worker_contracts(
599 state: &ServerState,
600 register: &ProtoRegisterWorker,
601) -> Result<(), Status> {
602 let advertised = register
603 .activities
604 .iter()
605 .map(|activity| {
606 let input_schema =
607 serde_json::from_str(&activity.input_schema_json).map_err(|error| {
608 Status::invalid_argument(format!(
609 "worker activity `{}` input_schema_json is invalid: {error}",
610 activity.name
611 ))
612 })?;
613 let output_schema =
614 serde_json::from_str(&activity.output_schema_json).map_err(|error| {
615 Status::invalid_argument(format!(
616 "worker activity `{}` output_schema_json is invalid: {error}",
617 activity.name
618 ))
619 })?;
620 Ok(aion_package::ActivityDescriptor {
621 name: activity.name.clone(),
622 input_schema,
623 output_schema,
624 })
625 })
626 .collect::<Result<Vec<_>, Status>>()?;
627 let Ok(engine) = state.engine() else {
633 tracing::warn!(
634 task_queue = %register.task_queue,
635 identity = %register.identity,
636 "worker contract check skipped: server state has no engine handle, \
637 so no deployed contracts exist to check against"
638 );
639 return Ok(());
640 };
641 let activity_types = register
645 .activity_types
646 .iter()
647 .cloned()
648 .collect::<std::collections::BTreeSet<_>>();
649 crate::worker::contracts::validate_worker_contracts(
650 &engine,
651 ®ister.task_queue,
652 crate::worker::registry::optional_node(®ister.node).as_deref(),
653 ®ister.identity,
654 crate::worker::contracts::WorkerAdvertisement {
655 activity_types: &activity_types,
656 contracts: &advertised,
657 },
658 )
659 .map_err(|error| match error {
660 crate::worker::contracts::ContractAdmissionError::Mismatch { .. } => {
661 Status::failed_precondition(error.to_string())
662 }
663 crate::worker::contracts::ContractAdmissionError::Catalog { .. } => {
664 Status::internal(error.to_string())
665 }
666 })
667}
668
669fn encode_server_to_worker(message: WorkerMessage) -> generated::ServerToWorker {
670 let message = match message {
671 WorkerMessage::ActivityTask(task) => {
672 generated::server_to_worker::Message::Task(encode_task(*task))
673 }
674 WorkerMessage::DrainRequest => {
675 generated::server_to_worker::Message::Drain(generated::DrainRequest {})
676 }
677 };
678 generated::ServerToWorker {
679 message: Some(message),
680 }
681}
682
683fn encode_task(task: aion_proto::ProtoActivityTask) -> generated::ActivityTask {
684 generated::ActivityTask {
685 workflow_id: task
686 .workflow_id
687 .map(|id| generated::WorkflowId { uuid: id.uuid }),
688 activity_id: task.activity_id.map(|id| generated::ActivityId {
689 sequence_position: id.sequence_position,
690 }),
691 activity_type: task.activity_type,
692 input: task.input.map(|p| generated::Payload {
693 content_type: p.content_type,
694 bytes: p.bytes,
695 }),
696 attempt: task.attempt,
697 labels: task.labels,
698 run_id: task.run_id.map(|id| generated::RunId { uuid: id.uuid }),
699 completion_token: task.completion_token,
700 idempotency_key: task.idempotency_key,
701 }
702}
703
704fn decode_activity_result(r: generated::ActivityResult) -> ProtoActivityResult {
705 ProtoActivityResult {
706 workflow_id: r
707 .workflow_id
708 .map(|id| aion_proto::ProtoWorkflowId { uuid: id.uuid }),
709 activity_id: r.activity_id.map(|id| aion_proto::ProtoActivityId {
710 sequence_position: id.sequence_position,
711 }),
712 outcome: r.outcome.map(decode_outcome),
713 run_id: r.run_id.map(|id| aion_proto::ProtoRunId { uuid: id.uuid }),
714 completion_token: r.completion_token,
715 }
716}
717
718fn decode_heartbeat(r: generated::Heartbeat) -> aion_proto::ProtoHeartbeat {
719 aion_proto::ProtoHeartbeat {
720 workflow_id: r
721 .workflow_id
722 .map(|id| aion_proto::ProtoWorkflowId { uuid: id.uuid }),
723 activity_id: r.activity_id.map(|id| aion_proto::ProtoActivityId {
724 sequence_position: id.sequence_position,
725 }),
726 progress: r.progress.map(|p| aion_proto::ProtoPayload {
727 content_type: p.content_type,
728 bytes: p.bytes,
729 }),
730 }
731}
732
733fn decode_outcome(
734 outcome: generated::activity_result::Outcome,
735) -> aion_proto::proto_activity_result::Outcome {
736 match outcome {
737 generated::activity_result::Outcome::Result(p) => {
738 aion_proto::proto_activity_result::Outcome::Result(aion_proto::ProtoPayload {
739 content_type: p.content_type,
740 bytes: p.bytes,
741 })
742 }
743 generated::activity_result::Outcome::Error(e) => {
744 aion_proto::proto_activity_result::Outcome::Error(aion_proto::ProtoActivityError {
745 kind: e.kind,
746 message: e.message,
747 details: e.details.map(|p| aion_proto::ProtoPayload {
748 content_type: p.content_type,
749 bytes: p.bytes,
750 }),
751 })
752 }
753 }
754}
755
756#[cfg(test)]
757mod tests {
758 use std::time::{Duration, Instant};
759
760 use aion_core::{ActivityId, ContentType, Payload, WorkflowId};
761
762 use crate::shutdown::DrainState;
763 use crate::worker::dispatch::{
764 ActivityCompletion, ActivityCompletionOutcome, ActivityCompletionSink,
765 };
766 use crate::worker::heartbeat::InFlightActivity;
767 use crate::worker::registry::ConnectedWorkerRegistry;
768 use crate::worker::{HeartbeatTracker, PendingActivities};
769
770 use super::teardown_worker_stream;
771
772 type TestError = Box<dyn std::error::Error>;
773
774 struct TeardownFixture {
778 registry: ConnectedWorkerRegistry,
779 tracker: HeartbeatTracker,
780 pending: PendingActivities,
781 drain: DrainState,
782 worker_id: crate::worker::registry::WorkerId,
783 workflow_id: WorkflowId,
784 activity_id: ActivityId,
785 completion_token: crate::worker::CompletionToken,
786 rx: std::sync::mpsc::Receiver<Result<String, String>>,
787 _registration: crate::worker::registry::WorkerRegistration,
790 }
791
792 fn fixture() -> Result<TeardownFixture, TestError> {
793 let registry = ConnectedWorkerRegistry::default();
794 let (tx, _rx) = tokio::sync::mpsc::channel(1);
795 let activity_types = [String::from("greet")];
796 let registration = registry.register("default", activity_types.iter(), tx)?;
797 let worker_id = registration
798 .worker_id()
799 .ok_or("test worker registration missing id")?;
800 let tracker = HeartbeatTracker::new(Duration::from_secs(5));
801 let pending = PendingActivities::default();
802 let workflow_id = WorkflowId::new_v4();
803 let activity_id = ActivityId::from_sequence_position(0);
804 let (completion_token, rx) =
805 pending.insert_for_test(workflow_id.clone(), activity_id.clone())?;
806 tracker.track_task(
807 worker_id,
808 InFlightActivity {
809 workflow_id: workflow_id.clone(),
810 activity_id: activity_id.clone(),
811 completion_token: completion_token.clone(),
812 },
813 Instant::now(),
814 )?;
815 Ok(TeardownFixture {
816 registry,
817 tracker,
818 pending,
819 drain: DrainState::default(),
820 worker_id,
821 workflow_id,
822 activity_id,
823 completion_token,
824 rx,
825 _registration: registration,
826 })
827 }
828
829 #[test]
833 fn teardown_under_drain_parks_instead_of_failing() -> Result<(), TestError> {
834 let fixture = fixture()?;
835 assert!(fixture.drain.begin());
836
837 teardown_worker_stream(
838 fixture.worker_id,
839 &fixture.tracker,
840 &fixture.registry,
841 &fixture.pending,
842 &fixture.drain,
843 );
844
845 let resolved = fixture.rx.recv_timeout(Duration::from_millis(200))?;
846 assert_eq!(
847 resolved,
848 Err(aion::PARKED_ACTIVITY_REASON.to_owned()),
849 "a drain teardown must resolve the waiter with the parked sentinel"
850 );
851 assert_eq!(fixture.tracker.in_flight_count()?, 0);
852 assert!(
853 !fixture.tracker.is_tracked(
854 fixture.worker_id,
855 &fixture.workflow_id,
856 &fixture.activity_id
857 )?,
858 "parking must retire the tracked entry"
859 );
860 Ok(())
861 }
862
863 #[test]
873 fn teardown_without_drain_fails_with_the_transport_domain_lost_worker_class()
874 -> Result<(), TestError> {
875 let fixture = fixture()?;
876
877 teardown_worker_stream(
878 fixture.worker_id,
879 &fixture.tracker,
880 &fixture.registry,
881 &fixture.pending,
882 &fixture.drain,
883 );
884
885 let resolved = fixture.rx.recv_timeout(Duration::from_millis(200))?;
886 let reason = resolved.err().ok_or("expected a lost-worker failure")?;
887 assert!(
888 reason.starts_with(crate::worker::WORKER_LOST_REASON_PREFIX),
889 "a mid-run teardown must surface the TRANSPORT-domain loss class, never the \
890 action's retry vocabulary: {reason}"
891 );
892 assert!(
893 reason.contains("lost before reporting activity result"),
894 "the failure must name worker loss: {reason}"
895 );
896 assert_eq!(fixture.tracker.in_flight_count()?, 0);
897 Ok(())
898 }
899
900 #[test]
904 fn stale_worker_completion_after_heartbeat_loss_does_not_resolve_retry() -> Result<(), TestError>
905 {
906 let fixture = fixture()?;
907 teardown_worker_stream(
908 fixture.worker_id,
909 &fixture.tracker,
910 &fixture.registry,
911 &fixture.pending,
912 &fixture.drain,
913 );
914 let first = fixture.rx.recv_timeout(Duration::from_millis(200))?;
915 assert!(
916 first
917 .err()
918 .is_some_and(|reason| reason.starts_with(crate::worker::WORKER_LOST_REASON_PREFIX)),
919 "worker A loss must release attempt 1 in the transport-loss class"
920 );
921
922 let (retry_token, retry_rx) = fixture
923 .pending
924 .insert_for_test(fixture.workflow_id.clone(), fixture.activity_id.clone())?;
925 let rejected = fixture.pending.complete_activity(ActivityCompletion {
926 workflow_id: fixture.workflow_id,
927 activity_id: fixture.activity_id,
928 run_id: None,
929 completion_token: fixture.completion_token,
930 outcome: ActivityCompletionOutcome::Succeeded(Payload::new(
931 ContentType::Json,
932 br#"{"worker":"A","stale":true}"#.to_vec(),
933 )),
934 });
935
936 assert!(matches!(
937 rejected,
938 Err(crate::ServerError::ActivityCompletionRejected { .. })
939 ));
940 drop(retry_token);
941 assert!(
942 retry_rx.recv_timeout(Duration::from_millis(50)).is_err(),
943 "worker A's late completion must be rejected instead of resolving worker B's retry"
944 );
945 Ok(())
946 }
947}