diff --git a/atom/compass/audit/sync_sites.json b/atom/compass/audit/sync_sites.json index e7d5bfd3c5..832b0f6716 100644 --- a/atom/compass/audit/sync_sites.json +++ b/atom/compass/audit/sync_sites.json @@ -2483,19 +2483,19 @@ "anchor": "pub const DEFAULT_WORKER_HTTP_TIMEOUT_SECS", "category": "C1", "peer": "deployment", - "why": "the router's thirty-second bound on one worker request, compiled in; it fires whenever simulated time runs slower than real" + "why": "the default timeout of the client the router uses only for HTTP health checks; each check sets its own timeout from --health-check-timeout-secs, which replaces this one, so it bounds no request" }, { - "file": "atom/mesh/src/core/worker_manager.rs", - "line": 24, - "anchor": "const REQUEST_TIMEOUT", + "file": "atom/mesh/src/cliargs.rs", + "line": 334, + "anchor": "pub worker_request_timeout_secs", "category": "C1", "peer": "deployment", - "why": "the router's five-second bound on a fan-out request, compiled in; it fires whenever simulated time runs slower than real" + "why": "the router's five-second default bound on its /get_load and /flush_cache requests to workers; raise it from the command line when simulated time runs slower than real" }, { "file": "atom/mesh/src/cliargs.rs", - "line": 423, + "line": 430, "anchor": "pub disable_health_check", "category": "C1", "peer": "deployment", @@ -2503,7 +2503,7 @@ }, { "file": "atom/mesh/src/cliargs.rs", - "line": 398, + "line": 405, "anchor": "pub disable_circuit_breaker", "category": "C1", "peer": "deployment", diff --git a/atom/mesh/src/cliargs.rs b/atom/mesh/src/cliargs.rs index 7ad7c83e7f..40f8e1fa76 100644 --- a/atom/mesh/src/cliargs.rs +++ b/atom/mesh/src/cliargs.rs @@ -1,4 +1,4 @@ -use std::sync::Arc; +use std::sync::{atomic::Ordering, Arc}; use clap::{ArgAction, Parser, Subcommand, ValueEnum}; @@ -8,7 +8,10 @@ use crate::{ HealthCheckConfig, MetricsConfig, PolicyConfig, RetryConfig, RouterConfig, RoutingMode, TokenizerCacheConfig, }, - core::ConnectionMode, + core::{ + worker_manager::{DEFAULT_WORKER_REQUEST_TIMEOUT_SECS, WORKER_REQUEST_TIMEOUT_SECS}, + ConnectionMode, + }, observability::metrics::PrometheusConfig, routers::atom_standalone::AtomStandaloneRuntime, server::{ServerConfig, ServerTlsConfig}, @@ -326,6 +329,10 @@ pub struct CliArgs { #[arg(long, default_value_t = 1800, help_heading = "Request Handling")] pub request_timeout_secs: u64, + /// Timeout in seconds for the router's own /get_load and /flush_cache requests to workers + #[arg(long, default_value_t = DEFAULT_WORKER_REQUEST_TIMEOUT_SECS, value_parser = clap::value_parser!(u64).range(1..), help_heading = "Request Handling")] + pub worker_request_timeout_secs: u64, + /// Grace period in seconds to wait for in-flight requests during shutdown #[arg(long, default_value_t = 180, help_heading = "Request Handling")] pub shutdown_grace_period_secs: u64, @@ -551,6 +558,7 @@ impl CliArgs { prefill_urls: Vec<(String, Option)>, ) -> ConfigResult { self.validate_tls_args()?; + WORKER_REQUEST_TIMEOUT_SECS.store(self.worker_request_timeout_secs, Ordering::Relaxed); // Determine routing mode based on PD disaggregation flag let mode = if self.pd_disaggregation { @@ -799,6 +807,7 @@ impl Default for CliArgs { prometheus_duration_buckets: Vec::new(), request_id_headers: Vec::new(), request_timeout_secs: 1800, + worker_request_timeout_secs: DEFAULT_WORKER_REQUEST_TIMEOUT_SECS, shutdown_grace_period_secs: 180, max_payload_size: 536_870_912, max_concurrent_requests: -1, diff --git a/atom/mesh/src/core/worker_manager.rs b/atom/mesh/src/core/worker_manager.rs index dad8f6c8cf..9393818cee 100644 --- a/atom/mesh/src/core/worker_manager.rs +++ b/atom/mesh/src/core/worker_manager.rs @@ -2,7 +2,14 @@ //! //! Provides worker lifecycle operations and fan-out request utilities. -use std::{collections::HashMap, sync::Arc, time::Duration}; +use std::{ + collections::HashMap, + sync::{ + atomic::{AtomicU64, Ordering}, + Arc, + }, + time::Duration, +}; use futures::{ future, @@ -21,7 +28,18 @@ use crate::{ protocols::worker_spec::{FlushCacheResult, WorkerLoadInfo, WorkerLoadsResult}, }; -const REQUEST_TIMEOUT: Duration = Duration::from_secs(5); +pub const DEFAULT_WORKER_REQUEST_TIMEOUT_SECS: u64 = 5; + +/// Timeout for the `/flush_cache` and `/get_load` requests below; set from +/// `--worker-request-timeout-secs` when the router config is built. +/// Process-wide: the last router config built in a process wins. +pub static WORKER_REQUEST_TIMEOUT_SECS: AtomicU64 = + AtomicU64::new(DEFAULT_WORKER_REQUEST_TIMEOUT_SECS); + +fn request_timeout() -> Duration { + Duration::from_secs(WORKER_REQUEST_TIMEOUT_SECS.load(Ordering::Relaxed)) +} + const MAX_CONCURRENT: usize = 32; /// Result of a fan-out request to a single worker @@ -47,7 +65,7 @@ async fn fan_out( let method = method.clone(); async move { - let mut req = client.request(method, &full_url).timeout(REQUEST_TIMEOUT); + let mut req = client.request(method, &full_url).timeout(request_timeout()); if let Some(key) = api_key { req = req.bearer_auth(key); } @@ -194,7 +212,7 @@ impl WorkerManager { api_key: Option<&str>, ) -> isize { let load_url = format!("{}/get_load", url); - let mut req = client.get(&load_url).timeout(REQUEST_TIMEOUT); + let mut req = client.get(&load_url).timeout(request_timeout()); if let Some(key) = api_key { req = req.bearer_auth(key); } @@ -340,3 +358,72 @@ impl Drop for LoadMonitor { } } } + +#[cfg(test)] +mod tests { + use std::time::Instant; + + use clap::Parser; + + use super::*; + use crate::{cliargs::CliArgs, core::BasicWorkerBuilder}; + + #[tokio::test] + async fn worker_request_timeout_option_outlasts_a_40s_worker() { + let slow = Duration::from_secs(40); + let app = axum::Router::new() + .route( + "/get_load", + axum::routing::get(move || async move { + tokio::time::sleep(slow).await; + axum::Json(serde_json::json!([{ "num_tokens": 7 }])) + }), + ) + .route( + "/flush_cache", + axum::routing::post(move || tokio::time::sleep(slow)), + ); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let url = format!("http://{}", listener.local_addr().unwrap()); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + + let registry = WorkerRegistry::new(); + registry.register(Arc::new(BasicWorkerBuilder::new(&url).build())); + let client = reqwest::Client::new(); + + for (flags, answered) in [ + (&[][..], false), + (&["--worker-request-timeout-secs", "60"][..], true), + ] { + let args = CliArgs::parse_from(["atomesh"].iter().chain(flags)); + args.to_router_config(vec![]).unwrap(); + let start = Instant::now(); + let (loads, flush) = tokio::join!( + WorkerManager::get_all_worker_loads(®istry, &client), + WorkerManager::flush_cache_all(®istry, &client), + ); + let elapsed = start.elapsed(); + eprintln!( + "flags={flags:?} load={} flushed={} elapsed={elapsed:?}", + loads.loads[0].load, + flush.successful.len() + ); + assert_eq!(loads.loads[0].load, if answered { 7 } else { -1 }); + assert_eq!(flush.successful.len(), answered as usize); + assert_eq!(elapsed >= slow, answered); + } + } + + #[test] + fn worker_request_timeout_defaults_to_5s_and_refuses_zero() { + let parse = |flags: &[&str]| CliArgs::try_parse_from(["atomesh"].iter().chain(flags)); + assert_eq!(parse(&[]).unwrap().worker_request_timeout_secs, 5); + assert_eq!( + parse(&["--worker-request-timeout-secs", "1"]) + .unwrap() + .worker_request_timeout_secs, + 1 + ); + assert!(parse(&["--worker-request-timeout-secs", "0"]).is_err()); + } +}