Skip to content

Commit a701eb4

Browse files
committed
Introduce execution_completion_behaviour: one_shot_always for workers.
1 parent db44f6a commit a701eb4

5 files changed

Lines changed: 272 additions & 39 deletions

File tree

‎nativelink-config/src/cas_server.rs‎

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -803,6 +803,17 @@ pub struct LocalWorkerConfig {
803803
/// them from CAS for every action.
804804
/// Default: None (directory cache disabled)
805805
pub directory_cache: Option<DirectoryCacheConfig>,
806+
807+
#[serde(default)]
808+
pub execution_completion_behaviour: ExecutionCompletionBehaviour,
809+
}
810+
811+
#[derive(Deserialize, Serialize, Debug, Default, Copy, Clone)]
812+
#[serde(rename_all = "snake_case")]
813+
pub enum ExecutionCompletionBehaviour {
814+
#[default]
815+
Default,
816+
OneShotAlways,
806817
}
807818

808819
#[derive(Deserialize, Serialize, Debug, Clone)]

‎nativelink-worker/src/local_worker.rs‎

Lines changed: 69 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -18,18 +18,17 @@ use core::sync::atomic::{AtomicU64, Ordering};
1818
use core::time::Duration;
1919
use std::process::Stdio;
2020
use std::sync::{Arc, Weak};
21-
2221
use futures::future::BoxFuture;
2322
use 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};
2726
use nativelink_metric::{MetricsComponent, RootMetricsComponent};
2827
use nativelink_proto::com::github::trace_machina::nativelink::remote_execution::update_for_worker::Update;
2928
use nativelink_proto::com::github::trace_machina::nativelink::remote_execution::worker_api_client::WorkerApiClient;
3029
use 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
};
3433
use nativelink_store::fast_slow_store::FastSlowStore;
3534
use nativelink_util::action_messages::{ActionResult, ActionStage, OperationId};
@@ -42,10 +41,11 @@ use nativelink_util::{spawn, tls_utils};
4241
use opentelemetry::context::Context;
4342
use tokio::process;
4443
use tokio::sync::{broadcast, mpsc};
44+
use tokio::sync::broadcast::{Receiver, Sender};
4545
use tokio::time::sleep;
4646
use tokio_stream::wrappers::UnboundedReceiverStream;
4747
use 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

5050
use 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

8788
async 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.

‎nativelink-worker/tests/local_worker_test.rs‎

Lines changed: 121 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@ use core::time::Duration;
1616
use std::collections::HashMap;
1717
use std::env;
1818
use std::ffi::OsString;
19-
use std::io::Write;
19+
use std::io::{Write};
2020
#[cfg(target_family = "unix")]
2121
use std::os::unix::fs::OpenOptionsExt;
2222
use std::path::PathBuf;
@@ -29,7 +29,7 @@ mod utils {
2929
}
3030

3131
use hyper::body::Frame;
32-
use nativelink_config::cas_server::{LocalWorkerConfig, WorkerProperty};
32+
use nativelink_config::cas_server::{ExecutionCompletionBehaviour, LocalWorkerConfig, WorkerProperty};
3333
use nativelink_config::stores::{
3434
FastSlowSpec, FilesystemSpec, MemorySpec, StoreDirection, StoreSpec,
3535
};
@@ -421,6 +421,125 @@ async fn simple_worker_start_action_test() -> Result<(), Error> {
421421
Ok(())
422422
}
423423

424+
#[nativelink_test]
425+
async fn one_shot_shutdowns_worker_test() -> Result<(), Error> {
426+
let config = LocalWorkerConfig {
427+
execution_completion_behaviour: ExecutionCompletionBehaviour::OneShotAlways,
428+
..Default::default()
429+
};
430+
let mut test_context = setup_local_worker_with_config(config).await;
431+
let streaming_response = test_context.maybe_streaming_response.take().unwrap();
432+
433+
{
434+
let props = test_context
435+
.client
436+
.expect_connect_worker(Ok(streaming_response))
437+
.await;
438+
assert_eq!(props, ConnectWorkerRequest::default());
439+
}
440+
441+
let expected_worker_id = "foobar".to_string();
442+
443+
let tx_stream = test_context.maybe_tx_stream.take().unwrap();
444+
{
445+
// First initialize our worker by sending the response to the connection request.
446+
tx_stream
447+
.send(Frame::data(
448+
encode_stream_proto(&UpdateForWorker {
449+
update: Some(Update::ConnectionResult(ConnectionResult {
450+
worker_id: expected_worker_id.clone(),
451+
})),
452+
})
453+
.unwrap(),
454+
))
455+
.await
456+
.map_err(|e| make_input_err!("Could not send : {:?}", e))?;
457+
}
458+
459+
let action_digest = DigestInfo::new([3u8; 32], 10);
460+
let action_info = ActionInfo {
461+
command_digest: DigestInfo::new([1u8; 32], 10),
462+
input_root_digest: DigestInfo::new([2u8; 32], 10),
463+
timeout: Duration::from_secs(1),
464+
platform_properties: HashMap::new(),
465+
priority: 0,
466+
load_timestamp: SystemTime::UNIX_EPOCH,
467+
insert_timestamp: SystemTime::UNIX_EPOCH,
468+
unique_qualifier: ActionUniqueQualifier::Uncacheable(ActionUniqueKey {
469+
instance_name: INSTANCE_NAME.to_string(),
470+
digest_function: DigestHasherFunc::Sha256,
471+
digest: action_digest,
472+
}),
473+
};
474+
475+
{
476+
// Send execution request.
477+
tx_stream
478+
.send(Frame::data(
479+
encode_stream_proto(&UpdateForWorker {
480+
update: Some(Update::StartAction(StartExecute {
481+
execute_request: Some((&action_info).into()),
482+
operation_id: String::new(),
483+
queued_timestamp: None,
484+
platform: Some(Platform::default()),
485+
worker_id: expected_worker_id.clone(),
486+
})),
487+
})
488+
.unwrap(),
489+
))
490+
.await
491+
.map_err(|e| make_input_err!("Could not send : {:?}", e))?;
492+
}
493+
494+
let running_action = Arc::new(MockRunningAction::new());
495+
496+
let action_result = ActionResult {
497+
output_files: vec![],
498+
output_folders: vec![],
499+
output_file_symlinks: vec![],
500+
output_directory_symlinks: vec![],
501+
exit_code: 5,
502+
stdout_digest: DigestInfo::new([21u8; 32], 10),
503+
stderr_digest: DigestInfo::new([22u8; 32], 10),
504+
execution_metadata: ExecutionMetadata {
505+
worker: expected_worker_id.clone(),
506+
queued_timestamp: SystemTime::UNIX_EPOCH,
507+
worker_start_timestamp: SystemTime::UNIX_EPOCH,
508+
worker_completed_timestamp: SystemTime::UNIX_EPOCH,
509+
input_fetch_start_timestamp: SystemTime::UNIX_EPOCH,
510+
input_fetch_completed_timestamp: SystemTime::UNIX_EPOCH,
511+
execution_start_timestamp: SystemTime::UNIX_EPOCH,
512+
execution_completed_timestamp: SystemTime::UNIX_EPOCH,
513+
output_upload_start_timestamp: SystemTime::UNIX_EPOCH,
514+
output_upload_completed_timestamp: SystemTime::UNIX_EPOCH,
515+
},
516+
server_logs: HashMap::new(),
517+
error: None,
518+
message: String::new(),
519+
};
520+
521+
// Send and wait for response from create_and_add_action to RunningActionsManager.
522+
test_context
523+
.actions_manager
524+
.expect_create_and_add_action(Ok(running_action.clone()))
525+
.await;
526+
527+
528+
// Now the RunningAction needs to send a series of state updates. This shortcuts them
529+
// into a single call (shortcut for prepare, execute, upload, collect_results, cleanup).
530+
running_action
531+
.simple_expect_get_finished_result(Ok(action_result.clone()))
532+
.await?;
533+
534+
test_context.client.expect_execution_response(Ok(())).await;
535+
536+
test_context.client
537+
.expect_going_away(Ok(()))
538+
.await;
539+
540+
Ok(())
541+
}
542+
424543
#[nativelink_test]
425544
async fn new_local_worker_creates_work_directory_test() -> Result<(), Error> {
426545
let cas_store = Store::new(FastSlowStore::new(

0 commit comments

Comments
 (0)