@@ -18,18 +18,17 @@ use core::sync::atomic::{AtomicU64, Ordering};
1818use core:: time:: Duration ;
1919use std:: process:: Stdio ;
2020use std:: sync:: { Arc , Weak } ;
21-
2221use futures:: future:: BoxFuture ;
2322use futures:: stream:: FuturesUnordered ;
24- use futures:: { Future , FutureExt , StreamExt , TryFutureExt , select } ;
25- use nativelink_config:: cas_server:: LocalWorkerConfig ;
26- use nativelink_error:: { Code , Error , ResultExt , make_err , make_input_err } ;
23+ use futures:: { select , Future , FutureExt , StreamExt , TryFutureExt } ;
24+ use nativelink_config:: cas_server:: { ExecutionCompletionBehaviour , LocalWorkerConfig } ;
25+ use nativelink_error:: { make_err , make_input_err , Code , Error , ResultExt } ;
2726use nativelink_metric:: { MetricsComponent , RootMetricsComponent } ;
2827use nativelink_proto:: com:: github:: trace_machina:: nativelink:: remote_execution:: update_for_worker:: Update ;
2928use nativelink_proto:: com:: github:: trace_machina:: nativelink:: remote_execution:: worker_api_client:: WorkerApiClient ;
3029use nativelink_proto:: com:: github:: trace_machina:: nativelink:: remote_execution:: {
31- ExecuteComplete , ExecuteResult , GoingAwayRequest , KeepAliveRequest , UpdateForWorker ,
32- execute_result ,
30+ execute_result , ExecuteComplete , ExecuteResult , GoingAwayRequest , KeepAliveRequest ,
31+ UpdateForWorker ,
3332} ;
3433use nativelink_store:: fast_slow_store:: FastSlowStore ;
3534use nativelink_util:: action_messages:: { ActionResult , ActionStage , OperationId } ;
@@ -42,10 +41,11 @@ use nativelink_util::{spawn, tls_utils};
4241use opentelemetry:: context:: Context ;
4342use tokio:: process;
4443use tokio:: sync:: { broadcast, mpsc} ;
44+ use tokio:: sync:: broadcast:: { Receiver , Sender } ;
4545use tokio:: time:: sleep;
4646use tokio_stream:: wrappers:: UnboundedReceiverStream ;
4747use tonic:: Streaming ;
48- use tracing:: { Level , debug, error, event, info, info_span, instrument, warn} ;
48+ use tracing:: { debug, error, event, info, info_span, instrument, warn, Level } ;
4949
5050use crate :: running_actions_manager:: {
5151 ExecutionConfiguration , Metrics as RunningActionManagerMetrics , RunningAction ,
@@ -82,6 +82,7 @@ struct LocalWorkerImpl<'a, T: WorkerApiClientTrait + 'static, U: RunningActionsM
8282 // on by the scheduler.
8383 actions_in_transit : Arc < AtomicU64 > ,
8484 metrics : Arc < Metrics > ,
85+ shutdown_tx : Sender < ShutdownGuard > ,
8586}
8687
8788async fn preconditions_met ( precondition_script : Option < String > ) -> Result < ( ) , Error > {
@@ -123,6 +124,7 @@ impl<'a, T: WorkerApiClientTrait + 'static, U: RunningActionsManager> LocalWorke
123124 worker_id : String ,
124125 running_actions_manager : Arc < U > ,
125126 metrics : Arc < Metrics > ,
127+ shutdown_tx : Sender < ShutdownGuard > ,
126128 ) -> Self {
127129 Self {
128130 config,
@@ -135,6 +137,7 @@ impl<'a, T: WorkerApiClientTrait + 'static, U: RunningActionsManager> LocalWorke
135137 // on by the scheduler.
136138 actions_in_transit : Arc :: new ( AtomicU64 :: new ( 0 ) ) ,
137139 metrics,
140+ shutdown_tx,
138141 }
139142 }
140143
@@ -184,6 +187,8 @@ impl<'a, T: WorkerApiClientTrait + 'static, U: RunningActionsManager> LocalWorke
184187
185188 let ( add_future_channel, add_future_rx) = mpsc:: unbounded_channel ( ) ;
186189 let mut add_future_rx = UnboundedReceiverStream :: new ( add_future_rx) . fuse ( ) ;
190+ let ( inner_shutdown_channel, inner_shutdown_rx) = mpsc:: unbounded_channel ( ) ;
191+ let mut inner_shutdown_rx = UnboundedReceiverStream :: new ( inner_shutdown_rx) . fuse ( ) ;
187192
188193 let mut update_for_worker_stream = update_for_worker_stream. fuse ( ) ;
189194 // A notify which is triggered every time actions_in_flight is subtracted.
@@ -193,6 +198,9 @@ impl<'a, T: WorkerApiClientTrait + 'static, U: RunningActionsManager> LocalWorke
193198 let actions_in_flight = Arc :: new ( AtomicU64 :: new ( 0 ) ) ;
194199 // Set to true when shutting down, this stops any new StartAction.
195200 let mut shutting_down = false ;
201+ // Channel to signal when shutdown is complete (GoingAway sent, ready to exit).
202+ let ( shutdown_complete_tx, shutdown_complete_rx) = mpsc:: unbounded_channel :: < ( ) > ( ) ;
203+ let mut shutdown_complete_rx = UnboundedReceiverStream :: new ( shutdown_complete_rx) . fuse ( ) ;
196204
197205 loop {
198206 select ! {
@@ -342,6 +350,7 @@ impl<'a, T: WorkerApiClientTrait + 'static, U: RunningActionsManager> LocalWorke
342350 self . actions_in_transit. fetch_add( 1 , Ordering :: Release ) ;
343351
344352 let add_future_channel = add_future_channel. clone( ) ;
353+ let inner_shutdown_channel = inner_shutdown_channel. clone( ) ;
345354
346355 info_span!(
347356 "worker_start_action_ctx" ,
@@ -364,7 +373,16 @@ impl<'a, T: WorkerApiClientTrait + 'static, U: RunningActionsManager> LocalWorke
364373 error!( ?err, "Error executing action" ) ;
365374 }
366375 add_future_channel
367- . send( make_publish_future( res) . then( move |res| {
376+ . send( make_publish_future( res)
377+ . then( move |res| {
378+ match self . config. execution_completion_behaviour {
379+ ExecutionCompletionBehaviour :: OneShotAlways => {
380+ inner_shutdown_channel. send( ( ) ) . ok( ) ;
381+ }
382+ ExecutionCompletionBehaviour :: Default => {
383+ // Do nothing
384+ }
385+ }
368386 actions_in_flight. fetch_sub( 1 , Ordering :: Release ) ;
369387 actions_notify. notify_one( ) ;
370388 core:: future:: ready( res)
@@ -388,13 +406,23 @@ impl<'a, T: WorkerApiClientTrait + 'static, U: RunningActionsManager> LocalWorke
388406 let fut = res. err_tip( || "New future stream receives should never be closed" ) ?;
389407 futures. push( fut) ;
390408 } ,
409+ _ = inner_shutdown_rx. next( ) => {
410+ warn!( "Shutting down worker because of inner shutdown signal" , ) ;
411+ let guard = ShutdownGuard :: default ( ) ;
412+ drop( self . shutdown_tx. send( guard. clone( ) ) ) ;
413+ }
391414 res = futures. next( ) => res. err_tip( || "Keep-alive should always pending. Likely unable to send data to scheduler" ) ??,
415+ _ = shutdown_complete_rx. next( ) => {
416+ info!( "Shutdown complete, exiting worker loop" ) ;
417+ return Ok ( ( ) ) ;
418+ } ,
392419 complete_msg = shutdown_rx. recv( ) . fuse( ) => {
393420 warn!( "Worker loop received shutdown signal. Shutting down worker..." , ) ;
394421 let mut grpc_client = self . grpc_client. clone( ) ;
395422 let shutdown_guard = complete_msg. map_err( |e| make_err!( Code :: Internal , "Failed to receive shutdown message: {e:?}" ) ) ?;
396423 let actions_in_flight = actions_in_flight. clone( ) ;
397424 let actions_notify = actions_notify. clone( ) ;
425+ let shutdown_complete_tx = shutdown_complete_tx. clone( ) ;
398426 let shutdown_future = async move {
399427 // Wait for in-flight operations to be fully completed.
400428 while actions_in_flight. load( Ordering :: Acquire ) > 0 {
@@ -408,6 +436,8 @@ impl<'a, T: WorkerApiClientTrait + 'static, U: RunningActionsManager> LocalWorke
408436 }
409437 // Allow shutdown to occur now.
410438 drop( shutdown_guard) ;
439+ // Signal that shutdown is complete.
440+ let _ = shutdown_complete_tx. send( ( ) ) ;
411441 Ok :: <( ) , Error >( ( ) )
412442 } ;
413443 futures. push( shutdown_future. boxed( ) ) ;
@@ -638,7 +668,8 @@ impl<T: WorkerApiClientTrait + 'static, U: RunningActionsManager> LocalWorker<T,
638668 #[ instrument( skip( self ) , level = Level :: INFO ) ]
639669 pub async fn run (
640670 mut self ,
641- mut shutdown_rx : broadcast:: Receiver < ShutdownGuard > ,
671+ shutdown_tx : Sender < ShutdownGuard > ,
672+ mut shutdown_rx : Receiver < ShutdownGuard > ,
642673 ) -> Result < ( ) , Error > {
643674 let sleep_fn = self
644675 . sleep_fn
@@ -673,6 +704,7 @@ impl<T: WorkerApiClientTrait + 'static, U: RunningActionsManager> LocalWorker<T,
673704 worker_id,
674705 self . running_actions_manager . clone ( ) ,
675706 self . metrics . clone ( ) ,
707+ shutdown_tx. clone ( ) ,
676708 ) ,
677709 update_for_worker_stream,
678710 ) ,
@@ -683,30 +715,37 @@ impl<T: WorkerApiClientTrait + 'static, U: RunningActionsManager> LocalWorker<T,
683715 ) ;
684716
685717 // Now listen for connections and run all other services.
686- if let Err ( err) = inner. run ( update_for_worker_stream, & mut shutdown_rx) . await {
687- ' no_more_actions: {
688- // Ensure there are no actions in transit before we try to kill
689- // all our actions.
690- const ITERATIONS : usize = 1_000 ;
691-
692- const ERROR_MSG : & str = "Actions in transit did not reach zero before we disconnected from the scheduler" ;
693-
694- let sleep_duration = ACTIONS_IN_TRANSIT_TIMEOUT_S / ITERATIONS as f32 ;
695- for _ in 0 ..ITERATIONS {
696- if inner. actions_in_transit . load ( Ordering :: Acquire ) == 0 {
697- break ' no_more_actions;
718+ match inner. run ( update_for_worker_stream, & mut shutdown_rx) . await {
719+ Ok ( ( ) ) => {
720+ // Graceful shutdown completed, return without retrying.
721+ info ! ( "Worker completed graceful shutdown" ) ;
722+ return Ok ( ( ) ) ;
723+ }
724+ Err ( err) => {
725+ ' no_more_actions: {
726+ // Ensure there are no actions in transit before we try to kill
727+ // all our actions.
728+ const ITERATIONS : usize = 1_000 ;
729+
730+ const ERROR_MSG : & str = "Actions in transit did not reach zero before we disconnected from the scheduler" ;
731+
732+ let sleep_duration = ACTIONS_IN_TRANSIT_TIMEOUT_S / ITERATIONS as f32 ;
733+ for _ in 0 ..ITERATIONS {
734+ if inner. actions_in_transit . load ( Ordering :: Acquire ) == 0 {
735+ break ' no_more_actions;
736+ }
737+ ( sleep_fn_pin) ( Duration :: from_secs_f32 ( sleep_duration) ) . await ;
698738 }
699- ( sleep_fn_pin) ( Duration :: from_secs_f32 ( sleep_duration) ) . await ;
739+ error ! ( ERROR_MSG ) ;
740+ return Err ( err. append ( ERROR_MSG ) ) ;
700741 }
701- error ! ( ERROR_MSG ) ;
702- return Err ( err. append ( ERROR_MSG ) ) ;
703- }
704- error ! ( ?err, "Worker disconnected from scheduler" ) ;
705- // Kill off any existing actions because if we re-connect, we'll
706- // get some more and it might resource lock us.
707- self . running_actions_manager . kill_all ( ) . await ;
742+ error ! ( ?err, "Worker disconnected from scheduler" ) ;
743+ // Kill off any existing actions because if we re-connect, we'll
744+ // get some more and it might resource lock us.
745+ self . running_actions_manager . kill_all ( ) . await ;
708746
709- ( error_handler) ( err) . await ; // Try to connect again.
747+ ( error_handler) ( err) . await ; // Try to connect again.
748+ }
710749 }
711750 }
712751 // Unreachable.
0 commit comments