Coverage Report

Created: 2026-10-01 15:31

next uncovered line (L), next uncovered region (R), next uncovered branch (B)
/build/source/src/bin/nativelink.rs
Line
Count
Source
1
// Copyright 2024 The NativeLink Authors. All rights reserved.
2
//
3
// Licensed under the Functional Source License, Version 1.1, Apache 2.0 Future License (the "License");
4
// you may not use this file except in compliance with the License.
5
// You may obtain a copy of the License at
6
//
7
//    See LICENSE file for details
8
//
9
// Unless required by applicable law or agreed to in writing, software
10
// distributed under the License is distributed on an "AS IS" BASIS,
11
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
// See the License for the specific language governing permissions and
13
// limitations under the License.
14
15
// The worker's action future is a long chain of combinators; proving it
16
// Send for the spawn walks the whole chain, and the default limit of 128 is
17
// a few steps short of it on the sanitizer toolchain.
18
#![recursion_limit = "256"]
19
20
use core::net::SocketAddr;
21
use core::time::Duration;
22
use std::collections::{HashMap, HashSet};
23
use std::io::ErrorKind;
24
use std::sync::Arc;
25
26
use async_lock::Mutex as AsyncMutex;
27
use axum::Router;
28
use axum::http::Uri;
29
use clap::Parser;
30
use futures::FutureExt;
31
use futures::future::{BoxFuture, Either, OptionFuture, TryFutureExt, try_join_all};
32
use hyper::StatusCode;
33
use hyper_util::rt::TokioTimer;
34
use hyper_util::rt::tokio::TokioIo;
35
use hyper_util::server::conn::auto;
36
use hyper_util::service::TowerToHyperService;
37
use mimalloc::MiMalloc;
38
use nativelink_config::cas_server::{
39
    CasConfig, CasStoreConfig, GlobalConfig, HttpCompressionAlgorithm, ListenerConfig,
40
    SchedulerConfig, ServerConfig, StoreConfig, WithInstanceName, WorkerConfig,
41
};
42
use nativelink_config::stores::ConfigDigestHashFunction;
43
use nativelink_error::{Code, Error, ResultExt, make_err, make_input_err};
44
use nativelink_scheduler::default_scheduler_factory::scheduler_factory;
45
use nativelink_service::ac_server::AcServer;
46
use nativelink_service::bep_server::BepServer;
47
use nativelink_service::bytestream_server::ByteStreamServer;
48
use nativelink_service::capabilities_server::CapabilitiesServer;
49
use nativelink_service::cas_server::CasServer;
50
use nativelink_service::execution_server::ExecutionServer;
51
use nativelink_service::fetch_server::FetchServer;
52
use nativelink_service::health_server::{HealthServer, health_paths};
53
use nativelink_service::push_server::PushServer;
54
use nativelink_service::wire_compression::RemoteCacheCompressionInstances;
55
use nativelink_service::worker_api_server::WorkerApiServer;
56
use nativelink_store::default_store_factory::store_factory;
57
use nativelink_store::store_manager::StoreManager;
58
use nativelink_util::common::fs::set_open_file_limit;
59
use nativelink_util::digest_hasher::{DigestHasherFunc, set_default_digest_hasher_func};
60
use nativelink_util::health_utils::HealthRegistryBuilder;
61
use nativelink_util::origin_event_publisher::OriginEventPublisher;
62
#[cfg(target_family = "unix")]
63
use nativelink_util::shutdown_guard::Priority;
64
use nativelink_util::shutdown_guard::ShutdownGuard;
65
use nativelink_util::store_trait::{
66
    DEFAULT_DIGEST_SIZE_HEALTH_CHECK_CFG, set_default_digest_size_health_check,
67
};
68
use nativelink_util::task::TaskExecutor;
69
use nativelink_util::telemetry::init_tracing;
70
use nativelink_util::{background_spawn, fs, spawn};
71
use nativelink_worker::local_worker::{WorkerRegistration, new_local_worker};
72
use rustls_pki_types::pem::PemObject;
73
use rustls_pki_types::{CertificateRevocationListDer, PrivateKeyDer};
74
use tokio::net::{TcpListener, TcpSocket};
75
use tokio::select;
76
#[cfg(target_family = "unix")]
77
use tokio::signal::unix::{SignalKind, signal};
78
use tokio::sync::oneshot::Sender;
79
use tokio::sync::{broadcast, mpsc, oneshot};
80
use tokio_rustls::TlsAcceptor;
81
use tokio_rustls::rustls::pki_types::CertificateDer;
82
use tokio_rustls::rustls::server::WebPkiClientVerifier;
83
use tokio_rustls::rustls::{RootCertStore, ServerConfig as TlsServerConfig};
84
use tonic::codec::CompressionEncoding;
85
use tonic::service::Routes;
86
use tracing::{error, error_span, info, trace_span, warn};
87
88
#[global_allocator]
89
static GLOBAL: MiMalloc = MiMalloc;
90
91
/// Note: This must be kept in sync with the documentation in `AdminConfig::path`.
92
const DEFAULT_ADMIN_API_PATH: &str = "/admin";
93
94
// Note: This must be kept in sync with the documentation in `HealthConfig::path`.
95
96
// Note: This must be kept in sync with the documentation in
97
// `OriginEventsConfig::max_event_queue_size`.
98
const DEFAULT_MAX_QUEUE_EVENTS: usize = 0x0001_0000;
99
100
/// Broadcast Channel Capacity
101
/// Note: The actual capacity may be greater than the provided capacity.
102
const BROADCAST_CAPACITY: usize = 1;
103
104
0
fn install_default_rustls_crypto_provider() {
105
0
    drop(tokio_rustls::rustls::crypto::ring::default_provider().install_default());
106
0
}
107
108
/// Bind a [`TcpListener`] with `IP_FREEBIND` set.
109
0
fn bind_freebind(socket_addr: SocketAddr) -> Result<TcpListener, std::io::Error> {
110
0
    let socket = match socket_addr {
111
0
        SocketAddr::V4(_) => TcpSocket::new_v4(),
112
0
        SocketAddr::V6(_) => TcpSocket::new_v6(),
113
0
    }?;
114
0
    fs::set_freebind(&socket)?;
115
0
    socket.bind(socket_addr)?;
116
0
    socket.listen(1024)
117
0
}
118
119
/// Backend for bazel remote execution / cache API.
120
#[derive(Parser, Debug)]
121
#[clap(
122
    author = "Trace Machina, Inc. <nativelink@tracemachina.com>",
123
    version,
124
    about,
125
    long_about = None
126
)]
127
struct Args {
128
    /// Config file to use.
129
    #[clap(value_parser)]
130
    config_file: String,
131
}
132
133
trait RoutesExt {
134
    fn add_optional_service<S>(self, svc: Option<S>) -> Self
135
    where
136
        S: tower::Service<
137
                axum::http::Request<tonic::body::Body>,
138
                Error = core::convert::Infallible,
139
            > + tonic::server::NamedService
140
            + Clone
141
            + Send
142
            + Sync
143
            + 'static,
144
        S::Response: axum::response::IntoResponse,
145
        S::Future: Send + 'static;
146
}
147
148
impl RoutesExt for Routes {
149
0
    fn add_optional_service<S>(mut self, svc: Option<S>) -> Self
150
0
    where
151
0
        S: tower::Service<
152
0
                axum::http::Request<tonic::body::Body>,
153
0
                Error = core::convert::Infallible,
154
0
            > + tonic::server::NamedService
155
0
            + Clone
156
0
            + Send
157
0
            + Sync
158
0
            + 'static,
159
0
        S::Response: axum::response::IntoResponse,
160
0
        S::Future: Send + 'static,
161
    {
162
0
        if let Some(svc) = svc {
163
0
            self = self.add_service(svc);
164
0
        }
165
0
        self
166
0
    }
167
}
168
169
/// If this value changes update the documentation in the config definition.
170
const DEFAULT_MAX_DECODING_MESSAGE_SIZE: usize = 4 * 1024 * 1024;
171
172
macro_rules! service_setup {
173
    ($service: expr, $http_config: ident) => {{
174
        let mut service = $service;
175
        let max_decoding_message_size = if $http_config.max_decoding_message_size == 0 {
176
            DEFAULT_MAX_DECODING_MESSAGE_SIZE
177
        } else {
178
            $http_config.max_decoding_message_size
179
        };
180
        service = service.max_decoding_message_size(max_decoding_message_size);
181
        let send_algo = &$http_config.compression.send_compression_algorithm;
182
        if let Some(encoding) = into_encoding(send_algo.unwrap_or(HttpCompressionAlgorithm::None)) {
183
            service = service.send_compressed(encoding);
184
        }
185
        for encoding in $http_config
186
            .compression
187
            .accepted_compression_algorithms
188
            .iter()
189
            // Filter None values.
190
0
            .filter_map(|from: &HttpCompressionAlgorithm| into_encoding(*from))
191
        {
192
            service = service.accept_compressed(encoding);
193
        }
194
        service
195
    }};
196
}
197
198
/// A JSON body with its content type, for the admin API's listings.
199
0
const fn json_response(
200
0
    body: String,
201
0
) -> ([(axum::http::header::HeaderName, &'static str); 1], String) {
202
0
    (
203
0
        [(axum::http::header::CONTENT_TYPE, "application/json")],
204
0
        body,
205
0
    )
206
0
}
207
208
0
async fn inner_main(
209
0
    cfg: CasConfig,
210
0
    shutdown_tx: broadcast::Sender<ShutdownGuard>,
211
0
    scheduler_shutdown_tx: Sender<()>,
212
0
) -> Result<(), Error> {
213
0
    const fn into_encoding(from: HttpCompressionAlgorithm) -> Option<CompressionEncoding> {
214
0
        match from {
215
0
            HttpCompressionAlgorithm::Gzip => Some(CompressionEncoding::Gzip),
216
0
            HttpCompressionAlgorithm::None => None,
217
        }
218
0
    }
219
220
0
    let health_registry_builder =
221
0
        Arc::new(AsyncMutex::new(HealthRegistryBuilder::new("nativelink")));
222
0
    let mut worker_registrations: Vec<Arc<WorkerRegistration>> = Vec::new();
223
224
0
    let store_manager = Arc::new(StoreManager::new());
225
    {
226
0
        let mut health_registry_lock = health_registry_builder.lock().await;
227
228
0
        for StoreConfig { name, spec } in cfg.stores {
229
0
            let health_component_name = format!("stores/{name}");
230
0
            let mut health_register_store =
231
0
                health_registry_lock.sub_builder(&health_component_name);
232
0
            let store = store_factory(&spec, &store_manager, Some(&mut health_register_store))
233
0
                .await
234
0
                .err_tip(|| format!("Failed to create store '{name}'"))?;
235
0
            store_manager
236
0
                .add_store(&name, store)
237
0
                .err_tip(|| format!("Failed to add store '{name}'"))?;
238
        }
239
0
        store_manager.run_post_init().await?;
240
241
        // Workers start after the listeners are up, but the health registry
242
        // is built with the listeners, so their registration flags are
243
        // made and registered here and handed to the workers later.
244
0
        let mut workers_health = health_registry_lock.sub_builder("workers");
245
0
        for (i, worker_cfg) in cfg.workers.iter().flatten().enumerate() {
246
0
            let WorkerConfig::Local(local_worker_cfg) = worker_cfg;
247
0
            let name = if local_worker_cfg.name.is_empty() {
248
0
                format!("worker_{i}")
249
            } else {
250
0
                local_worker_cfg.name.clone()
251
            };
252
0
            let registration = WorkerRegistration::new(&name);
253
0
            workers_health
254
0
                .sub_builder(&name)
255
0
                .register_indicator(registration.clone());
256
0
            worker_registrations.push(registration);
257
        }
258
    }
259
260
0
    let mut root_futures: Vec<BoxFuture<Result<(), Error>>> = Vec::new();
261
0
    let (single_use_complete_tx, single_use_complete_rx) = tokio::sync::oneshot::channel();
262
0
    let mut single_use_complete_tx = Some(single_use_complete_tx);
263
0
    let single_use_enabled = cfg.workers.as_ref().is_some_and(|workers| {
264
0
        workers.iter().any(|worker| match worker {
265
0
            WorkerConfig::Local(worker) => worker.single_use,
266
0
        })
267
0
    });
268
269
0
    let maybe_origin_event_tx = cfg
270
0
        .experimental_origin_events
271
0
        .as_ref()
272
0
        .map(|origin_events_cfg| {
273
0
            let mut max_queued_events = origin_events_cfg.max_event_queue_size;
274
0
            if max_queued_events == 0 {
275
0
                max_queued_events = DEFAULT_MAX_QUEUE_EVENTS;
276
0
            }
277
0
            let (tx, rx) = mpsc::channel(max_queued_events);
278
0
            let store_name = origin_events_cfg.publisher.store.as_str();
279
0
            let store = store_manager.get_store(store_name).err_tip(|| {
280
0
                format!("Could not get store {store_name} for origin event publisher")
281
0
            })?;
282
283
0
            root_futures.push(Box::pin(
284
0
                OriginEventPublisher::new(store, rx, shutdown_tx.clone())
285
0
                    .run()
286
0
                    .map(Ok),
287
0
            ));
288
289
0
            Ok::<_, Error>(tx)
290
0
        })
291
0
        .transpose()?;
292
293
0
    let mut action_schedulers = HashMap::new();
294
0
    let mut worker_schedulers = HashMap::new();
295
0
    for SchedulerConfig { name, spec } in cfg.schedulers.iter().flatten() {
296
0
        let (maybe_action_scheduler, maybe_worker_scheduler) =
297
0
            scheduler_factory(spec, &store_manager, maybe_origin_event_tx.as_ref())
298
0
                .await
299
0
                .err_tip(|| format!("Failed to create scheduler '{name}'"))?;
300
0
        if let Some(action_scheduler) = maybe_action_scheduler {
301
0
            action_schedulers.insert(name.clone(), action_scheduler.clone());
302
0
        }
303
0
        if let Some(worker_scheduler) = maybe_worker_scheduler {
304
0
            worker_schedulers.insert(name.clone(), worker_scheduler.clone());
305
0
        }
306
    }
307
308
0
    let server_cfgs: Vec<ServerConfig> = cfg.servers.into_iter().collect();
309
310
    // The capabilities service advertises chunking support for CAS instances
311
    // that may be served from a different server block (e.g. behind an L7
312
    // router), so collect the CAS configs across all blocks.
313
0
    let all_cas_configs: Vec<WithInstanceName<CasStoreConfig>> = server_cfgs
314
0
        .iter()
315
0
        .filter_map(|server_cfg| server_cfg.services.as_ref())
316
0
        .filter_map(|services| services.cas.as_deref())
317
0
        .flatten()
318
0
        .cloned()
319
0
        .collect();
320
321
0
    for server_cfg in server_cfgs {
322
0
        let services = server_cfg
323
0
            .services
324
0
            .err_tip(|| "'services' must be configured")?;
325
326
        // Currently we only support http as our socket type.
327
0
        let ListenerConfig::Http(http_config) = server_cfg.listener;
328
329
0
        let execution_server = services
330
0
            .execution
331
0
            .as_ref()
332
0
            .map(|cfg| ExecutionServer::new(cfg, &action_schedulers, &store_manager))
333
0
            .transpose()
334
0
            .err_tip(|| "Could not create Execution service")?;
335
336
0
        let capabilities_configs = services.capabilities.as_deref().unwrap_or_default();
337
0
        let remote_cache_compression_instances =
338
0
            RemoteCacheCompressionInstances::from_capabilities_configs(capabilities_configs);
339
340
0
        let tonic_services = Routes::builder()
341
0
            .routes()
342
0
            .add_optional_service(
343
0
                services
344
0
                    .ac
345
0
                    .map_or(Ok(None), |cfg| {
346
0
                        AcServer::new(&cfg, &store_manager)
347
0
                            .map(|v| Some(service_setup!(v.into_service(), http_config)))
348
0
                    })
349
0
                    .err_tip(|| "Could not create AC service")?,
350
            )
351
0
            .add_optional_service(
352
0
                services
353
0
                    .cas
354
0
                    .as_deref()
355
0
                    .map_or(Ok(None), |cfg| {
356
0
                        CasServer::new(cfg, &store_manager, &remote_cache_compression_instances)
357
0
                            .map(|v| Some(service_setup!(v.into_service(), http_config)))
358
0
                    })
359
0
                    .err_tip(|| "Could not create CAS service")?,
360
            )
361
0
            .add_optional_service(
362
0
                execution_server
363
0
                    .clone()
364
0
                    .map(|v| service_setup!(v.into_service(), http_config)),
365
            )
366
0
            .add_optional_service(
367
0
                execution_server.map(|v| service_setup!(v.into_operations_service(), http_config)),
368
            )
369
0
            .add_optional_service(
370
0
                services
371
0
                    .fetch
372
0
                    .map_or(Ok(None), |cfg| {
373
0
                        FetchServer::new(&cfg, &store_manager)
374
0
                            .map(|v| Some(service_setup!(v.into_service(), http_config)))
375
0
                    })
376
0
                    .err_tip(|| "Could not create Fetch service")?,
377
            )
378
0
            .add_optional_service(
379
0
                services
380
0
                    .push
381
0
                    .map_or(Ok(None), |cfg| {
382
0
                        PushServer::new(&cfg, &store_manager)
383
0
                            .map(|v| Some(service_setup!(v.into_service(), http_config)))
384
0
                    })
385
0
                    .err_tip(|| "Could not create Push service")?,
386
            )
387
0
            .add_optional_service(
388
0
                services
389
0
                    .bytestream
390
0
                    .map_or(Ok(None), |cfg| {
391
0
                        ByteStreamServer::new(
392
0
                            &cfg,
393
0
                            &store_manager,
394
0
                            &remote_cache_compression_instances,
395
                        )
396
0
                        .map(|v| Some(service_setup!(v.into_service(), http_config)))
397
0
                    })
398
0
                    .err_tip(|| "Could not create ByteStream service")?,
399
            )
400
0
            .add_optional_service(
401
0
                OptionFuture::from(services.capabilities.as_ref().map(|cfg| {
402
0
                    CapabilitiesServer::new(
403
0
                        cfg,
404
0
                        &action_schedulers,
405
0
                        &remote_cache_compression_instances,
406
0
                        &all_cas_configs,
407
                    )
408
0
                }))
409
0
                .await
410
0
                .map_or(Ok::<Option<CapabilitiesServer>, Error>(None), |server| {
411
0
                    Ok(Some(server?))
412
0
                })
413
0
                .err_tip(|| "Could not create Capabilities service")?
414
0
                .map(|v| service_setup!(v.into_service(), http_config)),
415
            )
416
0
            .add_optional_service(
417
0
                services
418
0
                    .worker_api
419
0
                    .map_or(Ok(None), |cfg| {
420
0
                        WorkerApiServer::new(&cfg, &worker_schedulers)
421
0
                            .map(|v| Some(service_setup!(v.into_service(), http_config)))
422
0
                    })
423
0
                    .err_tip(|| "Could not create WorkerApi service")?,
424
            )
425
0
            .add_optional_service(
426
0
                services
427
0
                    .experimental_bep
428
0
                    .map_or(Ok(None), |cfg| {
429
0
                        BepServer::new(&cfg, &store_manager)
430
0
                            .map(|v| Some(service_setup!(v.into_service(), http_config)))
431
0
                    })
432
0
                    .err_tip(|| "Could not create BEP service")?,
433
            );
434
435
0
        let health_registry = health_registry_builder.lock().await.build();
436
437
0
        let mut svc =
438
0
            tonic_services
439
0
                .into_axum_router()
440
0
                .layer(nativelink_util::telemetry::OtlpLayer::new(
441
0
                    server_cfg.experimental_identity_header.required,
442
                ));
443
444
0
        if let Some(health_cfg) = services.health {
445
0
            let (path, readiness_path) = health_paths(&health_cfg)?;
446
0
            svc = svc.route_service(
447
0
                &path,
448
0
                HealthServer::new(health_registry.clone(), &health_cfg),
449
            );
450
0
            svc = svc.route_service(
451
0
                &readiness_path,
452
0
                HealthServer::readiness(health_registry, &health_cfg),
453
            );
454
0
        }
455
456
0
        if let Some(admin_config) = services.admin {
457
0
            let path = if admin_config.path.is_empty() {
458
0
                DEFAULT_ADMIN_API_PATH
459
            } else {
460
0
                &admin_config.path
461
            };
462
0
            let worker_schedulers = Arc::new(worker_schedulers.clone());
463
0
            let listing_worker_schedulers = worker_schedulers.clone();
464
0
            let listing_action_schedulers = Arc::new(action_schedulers.clone());
465
            // Read-only views for the provisioner and operators: what the
466
            // scheduler has connected and what it has queued, with the
467
            // reservations the queued actions carry.
468
0
            let admin_router = Router::new()
469
0
                .route(
470
0
                    "/scheduler/{instance_name}/workers",
471
0
                    axum::routing::get(move |params: axum::extract::Path<String>| async move {
472
0
                        let instance_name = params.0;
473
0
                        let Some(scheduler) = listing_worker_schedulers.get(&instance_name) else {
474
0
                            return Err((
475
0
                                StatusCode::NOT_FOUND,
476
0
                                format!("No scheduler named '{instance_name}'"),
477
0
                            ));
478
                        };
479
0
                        let workers = scheduler.worker_snapshot().await;
480
0
                        nativelink_scheduler::admin::workers_json(&workers)
481
0
                            .map(json_response)
482
0
                            .map_err(|e| {
483
0
                                (StatusCode::INTERNAL_SERVER_ERROR, format!("Error: {e:?}"))
484
0
                            })
485
0
                    }),
486
                )
487
0
                .route(
488
0
                    "/scheduler/{instance_name}/demand",
489
0
                    axum::routing::get(move |params: axum::extract::Path<String>| async move {
490
0
                        let instance_name = params.0;
491
0
                        let Some(scheduler) = listing_action_schedulers.get(&instance_name) else {
492
0
                            return Err((
493
0
                                StatusCode::NOT_FOUND,
494
0
                                format!("No scheduler named '{instance_name}'"),
495
0
                            ));
496
                        };
497
0
                        nativelink_scheduler::admin::queued_demand_json(scheduler.as_ref())
498
0
                            .await
499
0
                            .map(json_response)
500
0
                            .map_err(|e| {
501
0
                                (StatusCode::INTERNAL_SERVER_ERROR, format!("Error: {e:?}"))
502
0
                            })
503
0
                    }),
504
                );
505
0
            svc = svc.nest_service(
506
0
                path,
507
0
                admin_router.route(
508
0
                    "/scheduler/{instance_name}/set_drain_worker/{worker_id}/{is_draining}",
509
0
                    axum::routing::post(
510
0
                        move |params: axum::extract::Path<(String, String, String)>| async move {
511
0
                            let (instance_name, worker_id, is_draining) = params.0;
512
0
                            (async move {
513
0
                                let is_draining = match is_draining.as_str() {
514
0
                                    "0" => false,
515
0
                                    "1" => true,
516
                                    _ => {
517
0
                                        return Err(make_err!(
518
0
                                            Code::Internal,
519
0
                                            "{} is neither 0 nor 1",
520
0
                                            is_draining
521
0
                                        ));
522
                                    }
523
                                };
524
0
                                worker_schedulers
525
0
                                    .get(&instance_name)
526
0
                                    .err_tip(|| {
527
0
                                        format!(
528
                                            "Can not get an instance with the name of '{instance_name}'",
529
                                        )
530
0
                                    })?
531
0
                                    .clone()
532
0
                                    .set_drain_worker(&worker_id.clone().into(), is_draining)
533
0
                                    .await?;
534
0
                                Ok::<_, Error>(format!("Draining worker {worker_id}"))
535
                            })
536
0
                            .await
537
0
                            .map_err(|e| {
538
0
                                Err::<String, _>((
539
0
                                    StatusCode::INTERNAL_SERVER_ERROR,
540
0
                                    format!("Error: {e:?}"),
541
0
                                ))
542
0
                            })
543
0
                        },
544
                    ),
545
                ),
546
            );
547
0
        }
548
549
        // This is the default service that executes if no other endpoint matches.
550
0
        svc = svc.fallback(|uri: Uri| async move {
551
0
            warn!("No route for {uri}");
552
0
            (StatusCode::NOT_FOUND, format!("No route for {uri}"))
553
0
        });
554
555
        // Configure our TLS acceptor if we have TLS configured.
556
0
        let maybe_tls_acceptor = http_config.tls.map_or(Ok(None), |tls_config| {
557
0
            fn read_cert(cert_file: &str) -> Result<Vec<CertificateDer<'static>>, Error> {
558
0
                let mut cert_reader = std::io::BufReader::new(
559
0
                    std::fs::File::open(cert_file)
560
0
                        .err_tip(|| format!("Could not open cert file {cert_file}"))?,
561
                );
562
0
                let certs = CertificateDer::pem_reader_iter(&mut cert_reader)
563
0
                    .collect::<Result<Vec<CertificateDer<'_>>, _>>()
564
0
                    .err_tip(|| format!("Could not extract certs from file {cert_file}"))?;
565
0
                Ok(certs)
566
0
            }
567
0
            let certs = read_cert(&tls_config.cert_file)?;
568
0
            let mut key_reader = std::io::BufReader::new(
569
0
                std::fs::File::open(&tls_config.key_file)
570
0
                    .err_tip(|| format!("Could not open key file {}", tls_config.key_file))?,
571
            );
572
0
            let key = match PrivateKeyDer::from_pem_reader(&mut key_reader)
573
0
                .err_tip(|| format!("Could not extract key(s) from file {}", tls_config.key_file))?
574
            {
575
0
                PrivateKeyDer::Pkcs8(key) => key.into(),
576
0
                PrivateKeyDer::Sec1(key) => key.into(),
577
0
                PrivateKeyDer::Pkcs1(key) => key.into(),
578
                _ => {
579
0
                    return Err(make_err!(
580
0
                        Code::Internal,
581
0
                        "No keys found in file {}",
582
0
                        tls_config.key_file
583
0
                    ));
584
                }
585
            };
586
0
            if PrivateKeyDer::from_pem_reader(&mut key_reader).is_ok() {
587
0
                return Err(make_err!(
588
0
                    Code::InvalidArgument,
589
0
                    "Expected 1 key in file {}",
590
0
                    tls_config.key_file
591
0
                ));
592
0
            }
593
0
            let verifier = if let Some(client_ca_file) = &tls_config.client_ca_file {
594
0
                let mut client_auth_roots = RootCertStore::empty();
595
0
                for cert in read_cert(client_ca_file)? {
596
0
                    client_auth_roots.add(cert).map_err(|e| {
597
0
                        Error::from_std_err(Code::Internal, &e).append("Could not read client CA")
598
0
                    })?;
599
                }
600
0
                let crls = if let Some(client_crl_file) = &tls_config.client_crl_file {
601
0
                    let mut crl_reader = std::io::BufReader::new(
602
0
                        std::fs::File::open(client_crl_file)
603
0
                            .err_tip(|| format!("Could not open CRL file {client_crl_file}"))?,
604
                    );
605
0
                    CertificateRevocationListDer::pem_reader_iter(&mut crl_reader)
606
0
                        .collect::<Result<_, _>>()
607
0
                        .err_tip(|| format!("Could not extract CRLs from file {client_crl_file}"))?
608
                } else {
609
0
                    Vec::new()
610
                };
611
0
                WebPkiClientVerifier::builder(Arc::new(client_auth_roots))
612
0
                    .with_crls(crls)
613
0
                    .build()
614
0
                    .map_err(|e| {
615
0
                        Error::from_std_err(Code::Internal, &e)
616
0
                            .append("Could not create WebPkiClientVerifier")
617
0
                    })?
618
            } else {
619
0
                WebPkiClientVerifier::no_client_auth()
620
            };
621
0
            let mut config = TlsServerConfig::builder()
622
0
                .with_client_cert_verifier(verifier)
623
0
                .with_single_cert(certs, key)
624
0
                .map_err(|e| {
625
0
                    Error::from_std_err(Code::Internal, &e)
626
0
                        .append("Could not create TlsServerConfig")
627
0
                })?;
628
629
0
            config.alpn_protocols.push("h2".into());
630
0
            Ok(Some(TlsAcceptor::from(Arc::new(config))))
631
0
        })?;
632
633
0
        let socket_addr = http_config
634
0
            .socket_address
635
0
            .parse::<SocketAddr>()
636
0
            .map_err(|e| {
637
0
                Error::from_std_err(Code::InvalidArgument, &e)
638
0
                    .append(format!("Invalid address '{}'", http_config.socket_address))
639
0
            })?;
640
0
        let tcp_listener = if http_config.freebind {
641
0
            bind_freebind(socket_addr)
642
        } else {
643
0
            TcpListener::bind(&socket_addr).await
644
        }
645
0
        .map_err(|e| match e.kind() {
646
0
            ErrorKind::AddrInUse => make_err!(
647
0
                Code::AlreadyExists,
648
                "Address '{socket_addr}' is already in use by another process.",
649
            ),
650
0
            ErrorKind::PermissionDenied => make_err!(
651
0
                Code::PermissionDenied,
652
                "Permission denied. You may need root privileges to bind to address '{socket_addr}'.",
653
            ),
654
            ErrorKind::InvalidInput => {
655
0
                make_input_err!("The provided address '{socket_addr}' is invalid.")
656
            }
657
0
            _ => Error::from_std_err(Code::Internal, &e)
658
0
                .append(format!("Failed to bind to socket address '{socket_addr}'")),
659
0
        })?;
660
0
        let mut http = auto::Builder::new(TaskExecutor::default());
661
0
        http.http2().timer(TokioTimer::new());
662
663
0
        let http_config = &http_config.advanced_http;
664
0
        if let Some(value) = http_config.http2_keep_alive_interval {
665
0
            http.http2()
666
0
                .keep_alive_interval(Duration::from_secs(u64::from(value)));
667
0
        }
668
669
0
        if let Some(value) = http_config.experimental_http2_max_pending_accept_reset_streams {
670
0
            http.http2()
671
0
                .max_pending_accept_reset_streams(usize::try_from(value).err_tip(
672
                    || "Could not convert experimental_http2_max_pending_accept_reset_streams",
673
0
                )?);
674
0
        }
675
0
        if let Some(value) = http_config.experimental_http2_initial_stream_window_size {
676
0
            http.http2().initial_stream_window_size(value);
677
0
        }
678
0
        if let Some(value) = http_config.experimental_http2_initial_connection_window_size {
679
0
            http.http2().initial_connection_window_size(value);
680
0
        }
681
0
        if let Some(value) = http_config.experimental_http2_adaptive_window {
682
0
            http.http2().adaptive_window(value);
683
0
        }
684
0
        if let Some(value) = http_config.experimental_http2_max_frame_size {
685
0
            http.http2().max_frame_size(value);
686
0
        }
687
0
        if let Some(value) = http_config.experimental_http2_max_concurrent_streams {
688
0
            http.http2().max_concurrent_streams(value);
689
0
        }
690
0
        if let Some(value) = http_config.experimental_http2_keep_alive_timeout_s {
691
0
            http.http2()
692
0
                .keep_alive_timeout(Duration::from_secs(u64::from(value)));
693
0
        }
694
0
        if let Some(value) = http_config.experimental_http2_max_send_buf_size {
695
0
            http.http2().max_send_buf_size(
696
0
                usize::try_from(value).err_tip(|| "Could not convert http2_max_send_buf_size")?,
697
            );
698
0
        }
699
0
        if http_config.experimental_http2_enable_connect_protocol == Some(true) {
700
0
            http.http2().enable_connect_protocol();
701
0
        }
702
0
        if let Some(value) = http_config.experimental_http2_max_header_list_size {
703
0
            http.http2().max_header_list_size(value);
704
0
        }
705
0
        info!("Ready, listening on {socket_addr}",);
706
0
        root_futures.push(Box::pin(async move {
707
            loop {
708
0
                select! {
709
0
                    accept_result = tcp_listener.accept() => {
710
0
                        match accept_result {
711
0
                            Ok((tcp_stream, remote_addr)) => {
712
0
                                info!(
713
                                    target: "nativelink::services",
714
                                    ?remote_addr,
715
                                    ?socket_addr,
716
                                    "Client connected"
717
                                );
718
719
0
                                let (http, svc, maybe_tls_acceptor) =
720
0
                                    (http.clone(), svc.clone(), maybe_tls_acceptor.clone());
721
722
0
                                background_spawn!(
723
                                    name: "http_connection",
724
0
                                    fut: error_span!(
725
                                        "http_connection",
726
                                        remote_addr = %remote_addr,
727
                                        socket_addr = %socket_addr,
728
0
                                    ).in_scope(|| async move {
729
0
                                        let serve_connection = if let Some(tls_acceptor) = maybe_tls_acceptor {
730
0
                                            match tls_acceptor.accept(tcp_stream).await {
731
0
                                                Ok(tls_stream) => Either::Left(http.serve_connection(
732
0
                                                    TokioIo::new(tls_stream),
733
0
                                                    TowerToHyperService::new(svc),
734
0
                                                )),
735
0
                                                Err(err) => {
736
0
                                                    error!(?err, "Failed to accept tls stream");
737
0
                                                    return;
738
                                                }
739
                                            }
740
                                        } else {
741
0
                                            Either::Right(http.serve_connection(
742
0
                                                TokioIo::new(tcp_stream),
743
0
                                                TowerToHyperService::new(svc),
744
0
                                            ))
745
                                        };
746
747
0
                                        if let Err(err) = serve_connection.await {
748
0
                                            error!(
749
                                                target: "nativelink::services",
750
                                                ?err,
751
                                                "Failed running service"
752
                                            );
753
0
                                        }
754
0
                                    }),
755
                                    target: "nativelink::services",
756
                                    ?remote_addr,
757
                                    ?socket_addr,
758
                                );
759
                            },
760
0
                            Err(err) => {
761
0
                                error!(?err, "Failed to accept tcp connection");
762
                            }
763
                        }
764
                    },
765
                }
766
            }
767
            // Unreachable
768
        }));
769
    }
770
771
    {
772
        // We start workers after our TcpListener is setup so if our worker connects to one
773
        // of these services it will be able to connect.
774
0
        let worker_cfgs = cfg.workers.unwrap_or_default();
775
0
        let mut worker_names = HashSet::with_capacity(worker_cfgs.len());
776
0
        for (i, worker_cfg) in worker_cfgs.into_iter().enumerate() {
777
0
            let spawn_fut = match worker_cfg {
778
0
                WorkerConfig::Local(local_worker_cfg) => {
779
0
                    let completion_tx = if local_worker_cfg.single_use {
780
0
                        single_use_complete_tx.take()
781
                    } else {
782
0
                        None
783
                    };
784
0
                    let fast_slow_store = store_manager
785
0
                        .get_store(&local_worker_cfg.cas_fast_slow_store)
786
0
                        .err_tip(|| {
787
0
                            format!(
788
                                "Failed to find store for cas_store_ref in worker config : {}",
789
                                local_worker_cfg.cas_fast_slow_store
790
                            )
791
0
                        })?;
792
793
0
                    let maybe_ac_store = if let Some(ac_store_ref) =
794
0
                        &local_worker_cfg.upload_action_result.ac_store
795
                    {
796
0
                        Some(store_manager.get_store(ac_store_ref).err_tip(|| {
797
0
                            format!("Failed to find store for ac_store in worker config : {ac_store_ref}")
798
0
                        })?)
799
                    } else {
800
0
                        None
801
                    };
802
                    // Note: Defaults to fast_slow_store if not specified. If this ever changes it must
803
                    // be updated in config documentation for the `historical_results_store` the field.
804
0
                    let historical_store = if let Some(cas_store_ref) = &local_worker_cfg
805
0
                        .upload_action_result
806
0
                        .historical_results_store
807
                    {
808
0
                        store_manager.get_store(cas_store_ref).err_tip(|| {
809
0
                                format!(
810
                                "Failed to find store for historical_results_store in worker config : {cas_store_ref}"
811
                            )
812
0
                            })?
813
                    } else {
814
0
                        fast_slow_store.clone()
815
                    };
816
0
                    let local_worker = new_local_worker(
817
0
                        Arc::new(local_worker_cfg),
818
0
                        fast_slow_store,
819
0
                        maybe_ac_store,
820
0
                        historical_store,
821
0
                    )
822
0
                    .await
823
0
                    .err_tip(|| "Could not make LocalWorker")?
824
0
                    .with_registration(worker_registrations[i].clone());
825
826
0
                    let name = if local_worker.name().is_empty() {
827
0
                        format!("worker_{i}")
828
                    } else {
829
0
                        local_worker.name().clone()
830
                    };
831
832
0
                    if worker_names.contains(&name) {
833
0
                        Err(make_input_err!(
834
0
                            "Duplicate worker name '{}' found in config",
835
0
                            name
836
0
                        ))?;
837
0
                    }
838
0
                    worker_names.insert(name.clone());
839
0
                    let shutdown_rx = shutdown_tx.subscribe();
840
0
                    let fut = trace_span!("worker_ctx", worker_name = %name)
841
0
                        .in_scope(|| local_worker.run(shutdown_rx));
842
0
                    spawn!(
843
                        "worker",
844
0
                        async move {
845
0
                            fut.await?;
846
0
                            if let Some(completion_tx) = completion_tx {
847
0
                                let _ = completion_tx.send(());
848
0
                            }
849
0
                            Ok::<(), Error>(())
850
0
                        },
851
                        ?name
852
                    )
853
                }
854
            };
855
0
            root_futures.push(Box::pin(spawn_fut.map_ok_or_else(|e| Err(e.into()), |v| v)));
856
        }
857
    }
858
859
    // Set up a shutdown handler for the worker schedulers.
860
0
    let mut shutdown_rx = shutdown_tx.subscribe();
861
0
    root_futures.push(Box::pin(async move {
862
0
        if let Ok(shutdown_guard) = shutdown_rx.recv().await {
863
0
            let _ = scheduler_shutdown_tx.send(());
864
0
            for (_name, scheduler) in worker_schedulers {
865
0
                scheduler.shutdown(shutdown_guard.clone()).await;
866
            }
867
0
        }
868
0
        Ok(())
869
0
    }));
870
871
0
    let result = select! {
872
0
        result = try_join_all(root_futures) => result.map(|_| ()),
873
0
        result = single_use_complete_rx, if single_use_enabled => {
874
0
            result.map_err(|err| make_err!(Code::Internal, "Single-use worker completion channel closed: {err}"))
875
        },
876
    };
877
0
    if let Err(e) = result {
878
0
        panic!("{e:?}");
879
0
    }
880
881
0
    Ok(())
882
0
}
883
884
0
fn get_config() -> Result<CasConfig, Error> {
885
0
    let args = Args::parse();
886
0
    CasConfig::try_from_json5_file(&args.config_file)
887
0
}
888
889
0
fn main() -> Result<(), Box<dyn core::error::Error>> {
890
0
    install_default_rustls_crypto_provider();
891
892
    // Set QoS to USER_INITIATED on the main thread *before* the tokio
893
    // runtime is built so the spawned worker threads inherit P-core
894
    // scheduling preference via pthread QoS inheritance on Apple
895
    // Silicon. `on_thread_start` below is a belt-and-suspenders hook
896
    // for any thread that misses the inherited class (e.g. tokio
897
    // blocking pool threads created lazily). No-op on non-macOS.
898
0
    let _ = nativelink_worker::qos::set_user_initiated();
899
900
    #[expect(clippy::disallowed_methods, reason = "starting main runtime")]
901
0
    let runtime = tokio::runtime::Builder::new_multi_thread()
902
0
        .on_thread_start(|| {
903
0
            let _ = nativelink_worker::qos::set_user_initiated();
904
0
        })
905
0
        .enable_all()
906
0
        .build()?;
907
908
    // The OTLP exporters need to run in a Tokio context
909
    // Do this first so all the other logging works
910
    #[expect(clippy::disallowed_methods, reason = "tracing init on main runtime")]
911
0
    runtime.block_on(async { tokio::spawn(async { init_tracing().await }).await? })?;
912
913
0
    let mut cfg = get_config()?;
914
915
0
    let global_cfg = if let Some(global_cfg) = &mut cfg.global {
916
0
        if global_cfg.max_open_files == 0 {
917
0
            global_cfg.max_open_files = fs::DEFAULT_OPEN_FILE_LIMIT;
918
0
        }
919
0
        if global_cfg.default_digest_size_health_check == 0 {
920
0
            global_cfg.default_digest_size_health_check = DEFAULT_DIGEST_SIZE_HEALTH_CHECK_CFG;
921
0
        }
922
923
0
        *global_cfg
924
    } else {
925
0
        GlobalConfig {
926
0
            max_open_files: fs::DEFAULT_OPEN_FILE_LIMIT,
927
0
            default_digest_hash_function: None,
928
0
            default_digest_size_health_check: DEFAULT_DIGEST_SIZE_HEALTH_CHECK_CFG,
929
0
        }
930
    };
931
0
    set_open_file_limit(global_cfg.max_open_files);
932
0
    set_default_digest_hasher_func(DigestHasherFunc::from(
933
0
        global_cfg
934
0
            .default_digest_hash_function
935
0
            .unwrap_or(ConfigDigestHashFunction::Sha256),
936
0
    ))?;
937
0
    set_default_digest_size_health_check(global_cfg.default_digest_size_health_check)?;
938
939
    // Initiates the shutdown process by broadcasting the shutdown signal via the `oneshot::Sender` to all listeners.
940
    // Each listener will perform its cleanup and then drop its `oneshot::Sender`, signaling completion.
941
    // Once all `oneshot::Sender` instances are dropped, the worker knows it can safely terminate.
942
0
    let (shutdown_tx, _) = broadcast::channel::<ShutdownGuard>(BROADCAST_CAPACITY);
943
    #[cfg(target_family = "unix")]
944
0
    let shutdown_tx_clone = shutdown_tx.clone();
945
    #[cfg(target_family = "unix")]
946
0
    let mut shutdown_guard = ShutdownGuard::default();
947
948
    #[allow(unused_variables)]
949
0
    let (scheduler_shutdown_tx, scheduler_shutdown_rx) = oneshot::channel();
950
951
    // SIGINT takes the SIGTERM path: a worker stopped from a terminal
952
    // drains like one stopped by its supervisor. A second SIGINT during
953
    // the drain exits at once, for the operator who meant it.
954
    #[cfg(target_family = "unix")]
955
    #[expect(clippy::disallowed_methods, reason = "signal handler on main runtime")]
956
0
    runtime.spawn(async move {
957
0
        let mut sigterm = signal(SignalKind::terminate()).expect("Failed to listen to SIGTERM");
958
0
        let mut sigint = signal(SignalKind::interrupt()).expect("Failed to listen to SIGINT");
959
0
        let exit_code = tokio::select! {
960
0
            _ = sigterm.recv() => { warn!("Process terminated via SIGTERM"); 143 }
961
0
            _ = sigint.recv() => { warn!("Process terminated via SIGINT, draining; send it again to exit at once"); 130 }
962
        };
963
0
        tokio::spawn(async move {
964
0
            sigint.recv().await;
965
0
            eprintln!("User terminated process via second SIGINT");
966
0
            std::process::exit(130);
967
        });
968
0
        drop(shutdown_tx_clone.send(shutdown_guard.clone()));
969
0
        scheduler_shutdown_rx
970
0
            .await
971
0
            .expect("Failed to receive scheduler shutdown");
972
0
        let () = shutdown_guard.wait_for(Priority::P0).await;
973
0
        warn!("Successfully shut down nativelink.");
974
0
        std::process::exit(exit_code);
975
    });
976
977
    #[expect(clippy::disallowed_methods, reason = "waiting on everything to finish")]
978
0
    runtime
979
0
        .block_on(async {
980
0
            trace_span!("main")
981
0
                .in_scope(|| async { inner_main(cfg, shutdown_tx, scheduler_shutdown_tx).await })
982
0
                .await
983
0
        })
984
0
        .err_tip(|| "main() function failed")?;
985
0
    Ok(())
986
0
}