From ebd6a3a96ef6d043d8eda2711c39f07b34d1ae79 Mon Sep 17 00:00:00 2001 From: Heran Lin Date: Fri, 4 Sep 2026 12:57:35 +0800 Subject: [PATCH 1/2] fix: improve worker launch logic --- crates/sail-common/src/config/application.rs | 1 + .../sail-common/src/config/application.yaml | 44 +++ crates/sail-common/src/utils/retry.rs | 174 ++++++++-- .../sail-execution/src/driver/actor/core.rs | 10 +- .../src/driver/actor/handler.rs | 143 +++++++- .../src/driver/actor/message.rs | 28 +- crates/sail-execution/src/driver/actor/mod.rs | 3 + .../src/driver/actor/options.rs | 2 + crates/sail-execution/src/driver/registry.rs | 6 +- .../src/driver/task_assigner/core.rs | 218 ++++++++++-- .../src/driver/task_assigner/mod.rs | 2 - .../src/driver/task_assigner/options.rs | 10 + .../src/driver/task_assigner/state.rs | 9 +- .../src/driver/worker_pool/core.rs | 57 ++-- .../src/driver/worker_pool/mod.rs | 1 + .../src/driver/worker_pool/state.rs | 16 + .../src/session_manager/actor/core.rs | 10 + .../src/session_manager/actor/handler.rs | 317 +++++++++++------- .../src/session_manager/actor/message.rs | 17 + .../src/session_manager/session.rs | 4 + 20 files changed, 843 insertions(+), 229 deletions(-) diff --git a/crates/sail-common/src/config/application.rs b/crates/sail-common/src/config/application.rs index 007fb3c1f4..ecc40f7d6d 100644 --- a/crates/sail-common/src/config/application.rs +++ b/crates/sail-common/src/config/application.rs @@ -209,6 +209,7 @@ pub struct ClusterConfig { pub worker_heartbeat_interval_secs: u64, pub worker_heartbeat_timeout_secs: u64, pub worker_launch_timeout_secs: u64, + pub worker_launch_retry_strategy: RetryStrategy, pub worker_task_slots: usize, pub task_launch_timeout_secs: u64, pub task_stream_buffer: usize, diff --git a/crates/sail-common/src/config/application.yaml b/crates/sail-common/src/config/application.yaml index 3f601d31a1..61cdbdb9ae 100644 --- a/crates/sail-common/src/config/application.yaml +++ b/crates/sail-common/src/config/application.yaml @@ -171,6 +171,50 @@ default: "120" description: The timeout in seconds for launching a worker. +- key: cluster.worker_launch_retry_strategy.type + type: string + default: "exponential_backoff" + description: | + The retry strategy for failed worker launches. + Valid values are `fixed` and `exponential_backoff`. + experimental: true + +- key: cluster.worker_launch_retry_strategy.fixed.max_count + type: number + default: "3" + description: The maximum number of worker launch retries using a fixed delay. + experimental: true + +- key: cluster.worker_launch_retry_strategy.fixed.delay_secs + type: number + default: "5" + description: The delay in seconds between worker launch retries using a fixed delay. + experimental: true + +- key: cluster.worker_launch_retry_strategy.exponential_backoff.max_count + type: number + default: "5" + description: The maximum number of worker launch retries using exponential backoff. + experimental: true + +- key: cluster.worker_launch_retry_strategy.exponential_backoff.initial_delay_secs + type: number + default: "1" + description: The initial delay in seconds before retrying a worker launch. + experimental: true + +- key: cluster.worker_launch_retry_strategy.exponential_backoff.max_delay_secs + type: number + default: "30" + description: The maximum delay in seconds before retrying a worker launch. + experimental: true + +- key: cluster.worker_launch_retry_strategy.exponential_backoff.factor + type: number + default: "2" + description: The factor by which the worker launch retry delay increases. + experimental: true + - key: cluster.worker_task_slots type: number default: "8" diff --git a/crates/sail-common/src/utils/retry.rs b/crates/sail-common/src/utils/retry.rs index e05945a468..0f97a7c6e4 100644 --- a/crates/sail-common/src/utils/retry.rs +++ b/crates/sail-common/src/utils/retry.rs @@ -21,19 +21,55 @@ pub enum RetryStrategy { }, } -struct ExponentialBackoffDelay { - delay: Duration, - max_delay: Duration, - factor: u32, +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RetryStep { + /// The retry number. Retries are numbered from one; the initial attempt is zero. + pub retry: usize, + pub delay: Duration, } -impl Iterator for ExponentialBackoffDelay { - type Item = Duration; +#[derive(Debug, Clone)] +pub struct RetrySchedule { + next: usize, + remaining: usize, + kind: RetryScheduleKind, +} + +#[derive(Debug, Clone)] +enum RetryScheduleKind { + Fixed { + delay: Duration, + }, + ExponentialBackoff { + delay: Duration, + max_delay: Duration, + factor: u32, + }, +} + +impl Iterator for RetrySchedule { + type Item = RetryStep; fn next(&mut self) -> Option { - let delay = self.delay; - self.delay = std::cmp::min(delay * self.factor, self.max_delay); - Some(delay) + if self.remaining == 0 { + return None; + } + let retry = self.next; + self.next += 1; + self.remaining -= 1; + let delay = match &mut self.kind { + RetryScheduleKind::Fixed { delay } => *delay, + RetryScheduleKind::ExponentialBackoff { + delay, + max_delay, + factor, + } => { + let current = *delay; + *delay = std::cmp::min(delay.saturating_mul(*factor), *max_delay); + current + } + }; + Some(RetryStep { retry, delay }) } } @@ -45,7 +81,7 @@ impl RetryStrategy { T: Send + 'static, E: std::fmt::Display + Send + 'static, { - let mut delay = self.delay(); + let mut retries = self.retries(); let mut attempt = 0; loop { let span = Span::enter_with_local_parent("RetryStrategy::run") @@ -55,33 +91,41 @@ impl RetryStrategy { x @ Ok(_) => return x, Err(e) => { warn!("retryable operation failed: {e}"); - if let Some(delay) = delay.next() { - tokio::time::sleep(delay).await; + if let Some(step) = retries.next() { + tokio::time::sleep(step.delay).await; + attempt = step.retry; } else { return Err(e); } } } - attempt += 1; } } - fn delay(&self) -> Box + Send> { + /// Returns a finite schedule containing only retries after the initial attempt. + /// + /// The first item has retry number one. If `max_count` is zero, the schedule is empty. + pub fn retries(&self) -> RetrySchedule { match self { Self::ExponentialBackoff { max_count, initial_delay, max_delay, factor, - } => Box::new( - ExponentialBackoffDelay { + } => RetrySchedule { + next: 1, + remaining: *max_count, + kind: RetryScheduleKind::ExponentialBackoff { delay: *initial_delay, max_delay: *max_delay, factor: *factor, - } - .take(*max_count), - ), - Self::Fixed { max_count, delay } => Box::new(std::iter::repeat_n(*delay, *max_count)), + }, + }, + Self::Fixed { max_count, delay } => RetrySchedule { + next: 1, + remaining: *max_count, + kind: RetryScheduleKind::Fixed { delay: *delay }, + }, } } } @@ -112,3 +156,93 @@ impl From<&config::RetryStrategy> for RetryStrategy { } } } + +#[cfg(test)] +mod tests { + use std::sync::Arc; + use std::sync::atomic::{AtomicUsize, Ordering}; + use std::time::Duration; + + use super::{RetryStep, RetryStrategy}; + + #[test] + fn fixed_schedule_contains_one_based_retries() { + let strategy = RetryStrategy::Fixed { + max_count: 3, + delay: Duration::from_secs(5), + }; + + assert_eq!( + strategy.retries().collect::>(), + vec![ + RetryStep { + retry: 1, + delay: Duration::from_secs(5), + }, + RetryStep { + retry: 2, + delay: Duration::from_secs(5), + }, + RetryStep { + retry: 3, + delay: Duration::from_secs(5), + }, + ] + ); + } + + #[test] + fn zero_max_count_has_no_retries() { + let strategy = RetryStrategy::Fixed { + max_count: 0, + delay: Duration::from_secs(5), + }; + + assert_eq!(strategy.retries().next(), None); + } + + #[test] + fn exponential_backoff_schedule_is_capped() { + let strategy = RetryStrategy::ExponentialBackoff { + max_count: 4, + initial_delay: Duration::from_secs(2), + max_delay: Duration::from_secs(5), + factor: 2, + }; + + let retries = strategy + .retries() + .map(|step| (step.retry, step.delay)) + .collect::>(); + assert_eq!( + retries, + vec![ + (1, Duration::from_secs(2)), + (2, Duration::from_secs(4)), + (3, Duration::from_secs(5)), + (4, Duration::from_secs(5)), + ] + ); + } + + #[tokio::test] + async fn run_performs_initial_attempt_and_scheduled_retries() { + let strategy = RetryStrategy::Fixed { + max_count: 2, + delay: Duration::ZERO, + }; + let calls = Arc::new(AtomicUsize::new(0)); + let result: Result<(), &str> = strategy + .run({ + let calls = Arc::clone(&calls); + move || { + calls.fetch_add(1, Ordering::Relaxed); + async { Err("failed") } + } + }) + .await; + + assert_eq!(result, Err("failed")); + assert_eq!(calls.load(Ordering::Relaxed), 3); + } +} diff --git a/crates/sail-execution/src/driver/actor/core.rs b/crates/sail-execution/src/driver/actor/core.rs index 85b82ee130..bb8e329015 100644 --- a/crates/sail-execution/src/driver/actor/core.rs +++ b/crates/sail-execution/src/driver/actor/core.rs @@ -50,6 +50,8 @@ impl Actor for DriverActor { task_assigner, task_runner: None, extensions: Default::default(), + activated: false, + worker_launch_retries_exhausted: false, task_sequences: HashMap::new(), shutdown_notifier: None, } @@ -129,7 +131,7 @@ impl Actor for DriverActor { message: DriverMessage, ) -> ActorAction { match message { - DriverMessage::Activate => self.handle_activate(ctx), + DriverMessage::Activate { result } => self.handle_activate(ctx, result), DriverMessage::RegisterWorker { worker_id, host, @@ -146,6 +148,12 @@ impl Actor for DriverActor { DriverMessage::ProbePendingWorker { worker_id } => { self.handle_probe_pending_worker(ctx, worker_id) } + DriverMessage::WorkerFailedToStart { worker_id, message } => { + self.handle_worker_failed_to_start(ctx, worker_id, message) + } + DriverMessage::RetryWorkerLaunch { worker_id, launch } => { + self.handle_retry_worker_launch(ctx, worker_id, launch) + } DriverMessage::ProbeIdleWorker { worker_id, instant } => { self.handle_probe_idle_worker(ctx, worker_id, instant) } diff --git a/crates/sail-execution/src/driver/actor/handler.rs b/crates/sail-execution/src/driver/actor/handler.rs index 7ec854242d..2d152feb31 100644 --- a/crates/sail-execution/src/driver/actor/handler.rs +++ b/crates/sail-execution/src/driver/actor/handler.rs @@ -5,7 +5,7 @@ use datafusion::arrow::datatypes::{Schema, SchemaRef}; use datafusion::execution::{SendableRecordBatchStream, TaskContext}; use datafusion::physical_plan::ExecutionPlan; use futures::TryStreamExt; -use log::{debug, info, warn}; +use log::{debug, error, info, warn}; use sail_celeborn::lifecycle::{LifecycleManagerActor, LocalLifecycleManager}; use sail_common::actor::{ActorAction, ActorContext, ActorHandle}; use sail_common_datafusion::error::CommonErrorCause; @@ -16,6 +16,7 @@ use tokio::time::Instant; use crate::driver::actor::DriverActor; use crate::driver::job_scheduler::{JobAction, TaskState}; use crate::driver::output::JobOutputItem; +use crate::driver::worker_pool::{WorkerLaunch, WorkerLaunchReason}; use crate::driver::{DriverMessage, TaskStatus}; use crate::error::{ExecutionError, ExecutionResult}; use crate::id::{JobId, TaskKey, TaskKeyDisplay, TaskStreamKey, TaskStreamKeyDisplay, WorkerId}; @@ -38,10 +39,25 @@ impl DriverActor { ActorAction::Continue } - pub(super) fn handle_activate(&mut self, ctx: &mut ActorContext) -> ActorAction { - info!("activating driver {}", self.options.driver_id); - for _ in 0..self.options.worker_initial_count { - self.worker_pool.start_worker(ctx); + pub(super) fn handle_activate( + &mut self, + ctx: &mut ActorContext, + result: oneshot::Sender>, + ) -> ActorAction { + let output = if self.activated { + Ok(()) + } else { + info!("activating driver {}", self.options.driver_id); + let count = self + .task_assigner + .request_initial_workers(self.options.worker_initial_count); + self.start_workers(ctx, count, WorkerLaunchReason::Initial) + .inspect(|_| { + self.activated = true; + }) + }; + if result.send(output).is_err() { + warn!("failed to send driver activation result"); } ActorAction::Continue } @@ -57,8 +73,10 @@ impl DriverActor { info!("worker {worker_id} is available at {host}:{port}"); let out = self.worker_pool.register_worker(ctx, worker_id, host, port); if out.is_ok() { + self.worker_launch_retries_exhausted = false; self.task_assigner.activate_worker(worker_id); self.run_tasks(ctx); + self.scale_up_workers(ctx); } if result.send(out).is_err() { warn!("failed to send worker registration result"); @@ -91,10 +109,52 @@ impl DriverActor { ctx: &mut ActorContext, worker_id: WorkerId, ) -> ActorAction { - if self.worker_pool.fail_worker_if_pending(worker_id) { - self.task_assigner.track_worker_failed_to_start(); - self.scale_up_workers(ctx); + self.handle_worker_launch_failure( + ctx, + worker_id, + "worker registration timeout".to_string(), + ); + ActorAction::Continue + } + + pub(super) fn handle_worker_failed_to_start( + &mut self, + ctx: &mut ActorContext, + worker_id: WorkerId, + message: String, + ) -> ActorAction { + self.handle_worker_launch_failure(ctx, worker_id, message); + ActorAction::Continue + } + + pub(super) fn handle_retry_worker_launch( + &mut self, + ctx: &mut ActorContext, + worker_id: WorkerId, + launch: WorkerLaunch, + ) -> ActorAction { + if !self.task_assigner.track_worker_failed_to_start(worker_id) { + return ActorAction::Continue; + } + let should_retry = match launch.reason { + WorkerLaunchReason::Initial => { + self.task_assigner + .request_initial_workers(self.options.worker_initial_count) + > 0 + } + WorkerLaunchReason::Demand => self.task_assigner.request_workers() > 0, + }; + if should_retry { + match self.worker_pool.start_worker(ctx, launch) { + Ok(worker_id) => self.task_assigner.track_pending_worker(worker_id), + Err(e) => { + error!("failed to retry worker launch: {e}"); + ctx.send(DriverMessage::Shutdown { result: None }); + return ActorAction::Continue; + } + } } + self.scale_up_workers(ctx); ActorAction::Continue } @@ -257,16 +317,16 @@ impl DriverActor { // once one registers (`handle_register_worker` runs pending tasks), // so reschedule the probe instead of failing. This keeps long, // many-stage jobs alive while the worker pool scales between stages. - // It cannot loop forever: a pending worker that never registers is - // failed at `worker_launch_timeout`, after which there are no pending - // workers and the task fails below. + // It cannot loop forever: each worker launch has a finite retry + // schedule. A failed worker remains pending here while waiting for + // its retry so the replacement capacity is not requested twice. // // Re-probe at `worker_launch_timeout` (capped by `task_launch_timeout`) // rather than a full `task_launch_timeout`: that is the window a // pending worker takes to register or be failed, so once the last // pending worker resolves the task fails promptly instead of waiting // another full launch window. - if self.worker_pool.has_pending_workers() { + if self.task_assigner.has_pending_workers() { let delay = self .options .worker_launch_timeout @@ -562,8 +622,63 @@ impl DriverActor { } fn scale_up_workers(&mut self, ctx: &mut ActorContext) { - for _ in 0..self.task_assigner.request_workers() { - self.worker_pool.start_worker(ctx); + if self.worker_launch_retries_exhausted { + if self.task_assigner.has_worker_demand() { + return; + } + self.worker_launch_retries_exhausted = false; + } + let count = self.task_assigner.request_workers(); + if let Err(e) = self.start_workers(ctx, count, WorkerLaunchReason::Demand) { + error!("failed to request workers: {e}"); + ctx.send(DriverMessage::Shutdown { result: None }); + } + } + + fn start_workers( + &mut self, + ctx: &mut ActorContext, + count: usize, + reason: WorkerLaunchReason, + ) -> ExecutionResult<()> { + for _ in 0..count { + let launch = WorkerLaunch { + reason, + attempt: 0, + retries: self.options.worker_launch_retry_strategy.retries(), + }; + let worker_id = self.worker_pool.start_worker(ctx, launch)?; + self.task_assigner.track_pending_worker(worker_id); + } + Ok(()) + } + + fn handle_worker_launch_failure( + &mut self, + ctx: &mut ActorContext, + worker_id: WorkerId, + message: String, + ) { + let Some(mut launch) = self.worker_pool.fail_worker_if_pending(worker_id, message) else { + return; + }; + if let Some(step) = launch.retries.next() { + launch.attempt = step.retry; + warn!( + "scheduling worker {worker_id} launch retry {} in {:?}", + step.retry, step.delay, + ); + ctx.send_with_delay( + DriverMessage::RetryWorkerLaunch { worker_id, launch }, + step.delay, + ); + } else { + self.task_assigner.track_worker_failed_to_start(worker_id); + self.worker_launch_retries_exhausted = self.task_assigner.has_worker_demand(); + warn!( + "worker {worker_id} launch retries exhausted after attempt {}", + launch.attempt + ); } } } diff --git a/crates/sail-execution/src/driver/actor/message.rs b/crates/sail-execution/src/driver/actor/message.rs index 8835ed594b..e3032d719a 100644 --- a/crates/sail-execution/src/driver/actor/message.rs +++ b/crates/sail-execution/src/driver/actor/message.rs @@ -14,12 +14,15 @@ use tokio::sync::oneshot; use tokio::time::Instant; use crate::driver::r#gen; +use crate::driver::worker_pool::WorkerLaunch; use crate::error::ExecutionResult; use crate::id::{JobId, TaskKey, TaskStreamKey, WorkerId}; use crate::stream::reader::TaskStreamSource; pub enum DriverMessage { - Activate, + Activate { + result: oneshot::Sender>, + }, RegisterWorker { worker_id: WorkerId, host: String, @@ -36,6 +39,15 @@ pub enum DriverMessage { ProbePendingWorker { worker_id: WorkerId, }, + WorkerFailedToStart { + worker_id: WorkerId, + message: String, + }, + RetryWorkerLaunch { + /// The failed worker whose reserved capacity this retry replaces. + worker_id: WorkerId, + launch: WorkerLaunch, + }, ProbeIdleWorker { worker_id: WorkerId, instant: Instant, @@ -127,11 +139,13 @@ impl From for r#gen::TaskStatus { impl SpanAssociation for DriverMessage { fn name(&self) -> Cow<'static, str> { let name = match self { - DriverMessage::Activate => "Activate", + DriverMessage::Activate { .. } => "Activate", DriverMessage::RegisterWorker { .. } => "RegisterWorker", DriverMessage::WorkerHeartbeat { .. } => "WorkerHeartbeat", DriverMessage::WorkerKnownPeers { .. } => "WorkerKnownPeers", DriverMessage::ProbePendingWorker { .. } => "ProbePendingWorker", + DriverMessage::WorkerFailedToStart { .. } => "WorkerFailedToStart", + DriverMessage::RetryWorkerLaunch { .. } => "RetryWorkerLaunch", DriverMessage::ProbeIdleWorker { .. } => "ProbeIdleWorker", DriverMessage::ProbeLostWorker { .. } => "ProbeLostWorker", DriverMessage::ExecuteJob { .. } => "ExecuteJob", @@ -149,7 +163,7 @@ impl SpanAssociation for DriverMessage { fn properties(&self) -> impl IntoIterator, Cow<'static, str>)> { let mut p: Vec<(&'static str, String)> = vec![]; match self { - DriverMessage::Activate => {} + DriverMessage::Activate { result: _ } => {} DriverMessage::RegisterWorker { worker_id, host, @@ -166,6 +180,10 @@ impl SpanAssociation for DriverMessage { peer_worker_ids: _, } | DriverMessage::ProbePendingWorker { worker_id } + | DriverMessage::WorkerFailedToStart { + worker_id, + message: _, + } | DriverMessage::ProbeIdleWorker { worker_id, instant: _, @@ -176,6 +194,10 @@ impl SpanAssociation for DriverMessage { } => { p.push((SpanAttribute::CLUSTER_WORKER_ID, worker_id.to_string())); } + DriverMessage::RetryWorkerLaunch { worker_id, launch } => { + p.push((SpanAttribute::CLUSTER_WORKER_ID, worker_id.to_string())); + p.push((SpanAttribute::RETRY_ATTEMPT, launch.attempt.to_string())); + } DriverMessage::ExecuteJob { plan: _, context: _, diff --git a/crates/sail-execution/src/driver/actor/mod.rs b/crates/sail-execution/src/driver/actor/mod.rs index f0499f0348..270e318631 100644 --- a/crates/sail-execution/src/driver/actor/mod.rs +++ b/crates/sail-execution/src/driver/actor/mod.rs @@ -29,6 +29,9 @@ pub struct DriverActor { task_assigner: TaskAssigner, task_runner: Option>, extensions: DriverExtensions, + activated: bool, + /// Whether launch retries have been exhausted for the current queued worker demand. + worker_launch_retries_exhausted: bool, /// The sequence number corresponding to the last task status update from the worker. /// A different sequence number is tracked for each attempt. task_sequences: HashMap, diff --git a/crates/sail-execution/src/driver/actor/options.rs b/crates/sail-execution/src/driver/actor/options.rs index 9588d363c1..d8fd3c2539 100644 --- a/crates/sail-execution/src/driver/actor/options.rs +++ b/crates/sail-execution/src/driver/actor/options.rs @@ -24,6 +24,7 @@ pub struct DriverOptions { pub worker_heartbeat_interval: Duration, pub worker_heartbeat_timeout: Duration, pub worker_launch_timeout: Duration, + pub worker_launch_retry_strategy: RetryStrategy, pub task_launch_timeout: Duration, pub task_stream_buffer: usize, pub task_stream_creation_timeout: Duration, @@ -64,6 +65,7 @@ impl DriverOptions { config.cluster.worker_heartbeat_timeout_secs, ), worker_launch_timeout: Duration::from_secs(config.cluster.worker_launch_timeout_secs), + worker_launch_retry_strategy: (&config.cluster.worker_launch_retry_strategy).into(), rpc_retry_strategy: (&config.cluster.rpc_retry_strategy).into(), task_launch_timeout: Duration::from_secs(config.cluster.task_launch_timeout_secs), task_stream_buffer: config.cluster.task_stream_buffer, diff --git a/crates/sail-execution/src/driver/registry.rs b/crates/sail-execution/src/driver/registry.rs index 06be5f4f9b..fa58d8d962 100644 --- a/crates/sail-execution/src/driver/registry.rs +++ b/crates/sail-execution/src/driver/registry.rs @@ -43,9 +43,11 @@ impl DriverHandle { } pub async fn activate(&self) -> ExecutionResult<()> { - self.send(DriverMessage::Activate) + let (result, receiver) = oneshot::channel(); + self.send(DriverMessage::Activate { result }) .await - .map_err(ExecutionError::from) + .map_err(ExecutionError::from)?; + receiver.await.map_err(ExecutionError::from)? } pub async fn shutdown(&self) -> ExecutionResult<()> { diff --git a/crates/sail-execution/src/driver/task_assigner/core.rs b/crates/sail-execution/src/driver/task_assigner/core.rs index b6e8d83eb1..d4205dfd7d 100644 --- a/crates/sail-execution/src/driver/task_assigner/core.rs +++ b/crates/sail-execution/src/driver/task_assigner/core.rs @@ -12,18 +12,8 @@ use crate::task::scheduling::{ }; impl TaskAssigner { - pub fn request_workers(&mut self) -> usize { - let enqueued_slots = self - .task_queue - .iter() - .map(|region| { - region - .tasks - .iter() - .filter(|(placement, _)| matches!(placement, TaskPlacement::Worker)) - .count() - }) - .sum::(); + pub fn request_workers(&self) -> usize { + let enqueued_slots = self.enqueued_worker_slots(); let vacant_slots = self .workers .values() @@ -31,10 +21,18 @@ impl TaskAssigner { WorkerResource::Active { task_slots, .. } => { task_slots.iter().filter(|x| x.is_vacant()).count() } - WorkerResource::Inactive => 0, + WorkerResource::Pending | WorkerResource::Inactive => 0, }) .sum::(); - let required_slots = enqueued_slots.saturating_sub(vacant_slots); + let pending_workers = self + .workers + .values() + .filter(|worker| matches!(worker, WorkerResource::Pending)) + .count(); + let pending_slots = pending_workers.saturating_mul(self.options.worker_task_slots); + let required_slots = enqueued_slots + .saturating_sub(vacant_slots) + .saturating_sub(pending_slots); let active_workers = self .workers .values() @@ -45,33 +43,97 @@ impl TaskAssigner { } else { self.options .worker_max_count - .saturating_sub(self.requested_worker_count) + .saturating_sub(pending_workers) .saturating_sub(active_workers) }; - let required_workers = required_slots + if self.options.worker_task_slots == 0 { + error!("worker task slots must be greater than zero"); + return 0; + } + required_slots .div_ceil(self.options.worker_task_slots) - .min(allowed_workers); - self.requested_worker_count = self.requested_worker_count.saturating_add(required_workers); - required_workers + .min(allowed_workers) } - pub fn track_worker_failed_to_start(&mut self) { - self.requested_worker_count = self.requested_worker_count.saturating_sub(1); + pub fn has_worker_demand(&self) -> bool { + self.enqueued_worker_slots() > 0 } - pub fn activate_worker(&mut self, worker_id: WorkerId) { - self.requested_worker_count = self.requested_worker_count.saturating_sub(1); - if self.workers.contains_key(&worker_id) { - warn!("worker {worker_id} is already active"); - return; + pub fn has_pending_workers(&self) -> bool { + self.workers + .values() + .any(|worker| matches!(worker, WorkerResource::Pending)) + } + + pub fn request_initial_workers(&self, worker_initial_count: usize) -> usize { + let live_workers = self + .workers + .values() + .filter(|worker| { + matches!( + worker, + WorkerResource::Pending | WorkerResource::Active { .. } + ) + }) + .count(); + let requested_workers = worker_initial_count.saturating_sub(live_workers); + if self.options.worker_max_count == 0 { + requested_workers + } else { + requested_workers.min(self.options.worker_max_count.saturating_sub(live_workers)) } - self.workers.insert( - worker_id, - WorkerResource::Active { - task_slots: vec![TaskSlot::default(); self.options.worker_task_slots], - local_streams: IndexSet::new(), + } + + pub fn track_pending_worker(&mut self, worker_id: WorkerId) { + match self.workers.entry(worker_id) { + indexmap::map::Entry::Vacant(entry) => { + entry.insert(WorkerResource::Pending); + } + indexmap::map::Entry::Occupied(_) => { + warn!("worker {worker_id} is already tracked"); + } + } + } + + pub fn track_worker_failed_to_start(&mut self, worker_id: WorkerId) -> bool { + let Some(worker) = self.workers.get_mut(&worker_id) else { + warn!("worker {worker_id} not found"); + return false; + }; + match worker { + WorkerResource::Pending => { + *worker = WorkerResource::Inactive; + true + } + WorkerResource::Active { .. } => { + warn!("worker {worker_id} is already active"); + false + } + WorkerResource::Inactive => { + warn!("worker {worker_id} is already inactive"); + false + } + } + } + + pub fn activate_worker(&mut self, worker_id: WorkerId) { + let resource = WorkerResource::Active { + task_slots: vec![TaskSlot::default(); self.options.worker_task_slots], + local_streams: IndexSet::new(), + }; + match self.workers.entry(worker_id) { + indexmap::map::Entry::Vacant(entry) => { + warn!("worker {worker_id} was not pending"); + entry.insert(resource); + } + indexmap::map::Entry::Occupied(mut entry) => match entry.get() { + WorkerResource::Pending => { + entry.insert(resource); + } + WorkerResource::Active { .. } => warn!("worker {worker_id} is already active"), + WorkerResource::Inactive => warn!("worker {worker_id} is inactive"), }, - ); + } } pub fn deactivate_worker(&mut self, worker_id: WorkerId) { @@ -80,7 +142,7 @@ impl TaskAssigner { return; }; match worker { - WorkerResource::Active { .. } => { + WorkerResource::Pending | WorkerResource::Active { .. } => { *worker = WorkerResource::Inactive; } WorkerResource::Inactive => { @@ -231,7 +293,7 @@ impl TaskAssigner { task_slots: slots, local_streams: streams, } => slots.iter().all(|s| s.is_vacant()) && streams.is_empty(), - WorkerResource::Inactive => false, + WorkerResource::Pending | WorkerResource::Inactive => false, } } @@ -247,10 +309,23 @@ impl TaskAssigner { .iter() .flat_map(|x| x.list_tasks().cloned().collect::>()) .collect(), - WorkerResource::Inactive => vec![], + WorkerResource::Pending | WorkerResource::Inactive => vec![], } } + fn enqueued_worker_slots(&self) -> usize { + self.task_queue + .iter() + .map(|region| { + region + .tasks + .iter() + .filter(|(placement, _)| matches!(placement, TaskPlacement::Worker)) + .count() + }) + .sum() + } + /// Builds a snapshot of available task slots across the driver and active workers for assignment. fn build_worker_task_slot_assigner(&self) -> TaskSlotAssigner { let slots = self @@ -258,7 +333,7 @@ impl TaskAssigner { .iter() .filter_map(|(id, worker)| { let slots = match worker { - WorkerResource::Inactive => vec![], + WorkerResource::Pending | WorkerResource::Inactive => vec![], WorkerResource::Active { task_slots: slots, .. } => slots @@ -333,3 +408,74 @@ impl TaskSlotAssigner { Ok(assignments) } } + +#[cfg(test)] +mod tests { + use crate::driver::task_assigner::{TaskAssigner, TaskAssignerOptions}; + use crate::id::WorkerId; + use crate::job_graph::TaskPlacement; + use crate::task::scheduling::{TaskRegion, TaskSet}; + + fn task_assigner(worker_task_slots: usize, worker_max_count: usize) -> TaskAssigner { + TaskAssigner::new(TaskAssignerOptions::new( + worker_task_slots, + worker_max_count, + )) + } + + fn worker_region(slots: usize) -> TaskRegion { + TaskRegion { + tasks: (0..slots) + .map(|_| (TaskPlacement::Worker, TaskSet { entries: vec![] })) + .collect(), + } + } + + #[test] + fn pending_worker_capacity_prevents_duplicate_requests() { + let mut assigner = task_assigner(8, 4); + assigner.enqueue_tasks(worker_region(8)); + + assert_eq!(assigner.request_workers(), 1); + assigner.track_pending_worker(WorkerId::from(1)); + assert_eq!(assigner.request_workers(), 0); + + assigner.track_worker_failed_to_start(WorkerId::from(1)); + assert_eq!(assigner.request_workers(), 1); + } + + #[test] + fn worker_requests_respect_pending_and_active_worker_limit() { + let mut assigner = task_assigner(1, 2); + assigner.enqueue_tasks(worker_region(4)); + assigner.track_pending_worker(WorkerId::from(1)); + assigner.track_pending_worker(WorkerId::from(2)); + + assert_eq!(assigner.request_workers(), 0); + + assigner.activate_worker(WorkerId::from(1)); + assigner.track_worker_failed_to_start(WorkerId::from(2)); + assert_eq!(assigner.request_workers(), 1); + } + + #[test] + fn initial_workers_share_the_worker_limit() { + let mut assigner = task_assigner(8, 4); + assigner.track_pending_worker(WorkerId::from(1)); + + assert_eq!(assigner.request_initial_workers(4), 3); + + assigner.track_pending_worker(WorkerId::from(2)); + assigner.track_pending_worker(WorkerId::from(3)); + assigner.track_pending_worker(WorkerId::from(4)); + assert_eq!(assigner.request_initial_workers(4), 0); + } + + #[test] + fn zero_worker_task_slots_does_not_panic() { + let mut assigner = task_assigner(0, 1); + assigner.enqueue_tasks(worker_region(1)); + + assert_eq!(assigner.request_workers(), 0); + } +} diff --git a/crates/sail-execution/src/driver/task_assigner/mod.rs b/crates/sail-execution/src/driver/task_assigner/mod.rs index 8425afe9d8..14366faee7 100644 --- a/crates/sail-execution/src/driver/task_assigner/mod.rs +++ b/crates/sail-execution/src/driver/task_assigner/mod.rs @@ -16,7 +16,6 @@ pub struct TaskAssigner { options: TaskAssignerOptions, driver: DriverResource, workers: IndexMap, - requested_worker_count: usize, /// A lookup table from task attempts to the place they are assigned to. /// This is more convenient than finding the task attempt in the task slots. /// @@ -36,7 +35,6 @@ impl TaskAssigner { options, driver: DriverResource::default(), workers: IndexMap::new(), - requested_worker_count: 0, task_assignments: IndexMap::new(), task_queue: VecDeque::new(), } diff --git a/crates/sail-execution/src/driver/task_assigner/options.rs b/crates/sail-execution/src/driver/task_assigner/options.rs index f5afeb216d..5fb505e704 100644 --- a/crates/sail-execution/src/driver/task_assigner/options.rs +++ b/crates/sail-execution/src/driver/task_assigner/options.rs @@ -6,6 +6,16 @@ pub struct TaskAssignerOptions { pub worker_max_count: usize, } +#[cfg(test)] +impl TaskAssignerOptions { + pub(super) fn new(worker_task_slots: usize, worker_max_count: usize) -> Self { + Self { + worker_task_slots, + worker_max_count, + } + } +} + impl From<&DriverOptions> for TaskAssignerOptions { fn from(options: &DriverOptions) -> Self { Self { diff --git a/crates/sail-execution/src/driver/task_assigner/state.rs b/crates/sail-execution/src/driver/task_assigner/state.rs index 51ca7abaac..6a839387ab 100644 --- a/crates/sail-execution/src/driver/task_assigner/state.rs +++ b/crates/sail-execution/src/driver/task_assigner/state.rs @@ -90,6 +90,7 @@ impl DriverResource { #[derive(Debug)] /// Represents the current state of a worker's resources as seen by the task assigner. pub enum WorkerResource { + Pending, Active { /// The task slots on the worker. task_slots: Vec, @@ -122,7 +123,7 @@ impl WorkerResource { warn!("invalid task slot {slot} on worker"); } } - WorkerResource::Inactive => { + WorkerResource::Pending | WorkerResource::Inactive => { warn!("cannot add tasks to inactive worker"); } } @@ -138,7 +139,7 @@ impl WorkerResource { false } } - WorkerResource::Inactive => { + WorkerResource::Pending | WorkerResource::Inactive => { warn!("cannot remove tasks from inactive worker"); false } @@ -150,7 +151,7 @@ impl WorkerResource { WorkerResource::Active { local_streams, .. } => { local_streams.extend(set.local_streams().cloned()); } - WorkerResource::Inactive => { + WorkerResource::Pending | WorkerResource::Inactive => { warn!("cannot track local streams on inactive worker"); } } @@ -167,7 +168,7 @@ impl WorkerResource { } count != local_streams.len() } - WorkerResource::Inactive => { + WorkerResource::Pending | WorkerResource::Inactive => { warn!("cannot untrack local streams from inactive worker"); false } diff --git a/crates/sail-execution/src/driver/worker_pool/core.rs b/crates/sail-execution/src/driver/worker_pool/core.rs index 58ad06e358..9b82c913b0 100644 --- a/crates/sail-execution/src/driver/worker_pool/core.rs +++ b/crates/sail-execution/src/driver/worker_pool/core.rs @@ -16,7 +16,7 @@ use tokio::time::Instant; use tonic::Code; use crate::driver::worker_pool::state::WorkerState; -use crate::driver::worker_pool::{WorkerDescriptor, WorkerPool, WorkerPoolOptions}; +use crate::driver::worker_pool::{WorkerDescriptor, WorkerLaunch, WorkerPool, WorkerPoolOptions}; use crate::driver::{DriverActor, DriverMessage, TaskStatus}; use crate::error::{ExecutionError, ExecutionResult}; use crate::id::{JobId, TaskKey, TaskKeyDisplay, TaskStreamKey, WorkerId}; @@ -38,14 +38,19 @@ impl WorkerPool { Ok(()) } - pub fn start_worker(&mut self, ctx: &mut ActorContext) { - let Ok(worker_id) = self.worker_id_generator.generate() else { - error!("failed to generate worker ID"); - ctx.send(DriverMessage::Shutdown { result: None }); - return; - }; + pub fn start_worker( + &mut self, + ctx: &mut ActorContext, + launch: WorkerLaunch, + ) -> ExecutionResult { + let worker_id = self.worker_id_generator.generate()?; + info!( + "starting worker {worker_id} (launch attempt {})", + launch.attempt + ); let descriptor = WorkerDescriptor { state: WorkerState::Pending, + launch: Some(launch), messages: vec![], peers: HashSet::new(), }; @@ -93,11 +98,19 @@ impl WorkerPool { let task = self .worker_manager .launch_worker(ctx.children_mut(), worker_id, options); + let driver = ctx.handle().clone(); ctx.spawn(async move { if let Err(e) = task.await { error!("failed to start worker {worker_id}: {e}"); + let _ = driver + .send(DriverMessage::WorkerFailedToStart { + worker_id, + message: e.to_string(), + }) + .await; } }); + Ok(worker_id) } pub fn register_worker( @@ -123,6 +136,7 @@ impl WorkerPool { heartbeat_at: Instant::now(), client: None, }; + worker.launch = None; Self::schedule_lost_worker_probe(ctx, worker_id, worker, &self.options); Self::schedule_idle_worker_probe(ctx, worker_id, worker, &self.options); event_reporter.report(SystemEvent::WorkerUpdated { @@ -211,20 +225,6 @@ impl WorkerPool { } } - /// Returns true if any worker is still launching (pending registration). - /// - /// A task stuck in `Created` should wait for such a worker rather than - /// failing with a scheduling timeout: once the worker registers, - /// `handle_register_worker` runs the pending tasks and can assign it. A - /// worker that never registers is bounded by `worker_launch_timeout` - /// (`fail_worker_if_pending`), after which it leaves the `Pending` state, so - /// this cannot keep a task alive forever. - pub fn has_pending_workers(&self) -> bool { - self.workers - .values() - .any(|worker| matches!(worker.state, WorkerState::Pending)) - } - fn list_running_workers(&self) -> Vec { self.workers .iter() @@ -269,16 +269,19 @@ impl WorkerPool { worker.peers.extend(peer_worker_ids); } - pub fn fail_worker_if_pending(&mut self, worker_id: WorkerId) -> bool { + pub fn fail_worker_if_pending( + &mut self, + worker_id: WorkerId, + message: String, + ) -> Option { let event_reporter = self.event_reporter.clone(); let session_id = self.options.session_id.clone(); let Some(worker) = self.workers.get_mut(&worker_id) else { warn!("worker {worker_id} not found"); - return false; + return None; }; if matches!(&worker.state, WorkerState::Pending) { - warn!("worker {worker_id} registration timeout"); - let message = "worker registration timeout".to_string(); + warn!("worker {worker_id} failed to start: {message}"); worker.state = WorkerState::Failed; worker.messages.push(message); event_reporter.report(SystemEvent::WorkerUpdated { @@ -289,9 +292,9 @@ impl WorkerPool { status: worker.state.status().to_string(), updated_at: Utc::now(), }); - true + worker.launch.take() } else { - false + None } } diff --git a/crates/sail-execution/src/driver/worker_pool/mod.rs b/crates/sail-execution/src/driver/worker_pool/mod.rs index de084b2abe..06d402d0d2 100644 --- a/crates/sail-execution/src/driver/worker_pool/mod.rs +++ b/crates/sail-execution/src/driver/worker_pool/mod.rs @@ -7,6 +7,7 @@ use std::sync::Arc; use indexmap::IndexMap; pub use options::WorkerPoolOptions; use sail_telemetry::events::SystemEventReporter; +pub(crate) use state::{WorkerLaunch, WorkerLaunchReason}; use crate::driver::worker_pool::state::WorkerDescriptor; use crate::id::{IdGenerator, WorkerId}; diff --git a/crates/sail-execution/src/driver/worker_pool/state.rs b/crates/sail-execution/src/driver/worker_pool/state.rs index 53399caff0..b74b55a8d4 100644 --- a/crates/sail-execution/src/driver/worker_pool/state.rs +++ b/crates/sail-execution/src/driver/worker_pool/state.rs @@ -1,5 +1,6 @@ use std::collections::HashSet; +use sail_common::utils::retry::RetrySchedule; use tokio::time::Instant; use crate::id::WorkerId; @@ -7,6 +8,7 @@ use crate::worker::WorkerClientSet; pub struct WorkerDescriptor { pub state: WorkerState, + pub launch: Option, pub messages: Vec, /// A list of peer workers known to the worker. /// The list may or may not cover all the running workers, @@ -16,6 +18,20 @@ pub struct WorkerDescriptor { pub peers: HashSet, } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum WorkerLaunchReason { + Initial, + Demand, +} + +#[derive(Debug, Clone)] +pub struct WorkerLaunch { + pub reason: WorkerLaunchReason, + /// The launch attempt number. The initial attempt is zero. + pub attempt: usize, + pub retries: RetrySchedule, +} + pub enum WorkerState { Pending, Running { diff --git a/crates/sail-session/src/session_manager/actor/core.rs b/crates/sail-session/src/session_manager/actor/core.rs index 6e4153b155..dbeb8e4c7b 100644 --- a/crates/sail-session/src/session_manager/actor/core.rs +++ b/crates/sail-session/src/session_manager/actor/core.rs @@ -88,6 +88,16 @@ impl Actor for SessionManagerActor { user_id, result, } => self.handle_get_or_create_session(ctx, session_id, user_id, result), + SessionManagerMessage::CompleteSessionCreation { + session_id, + user_id, + context, + driver_id, + activation, + result, + } => self.handle_complete_session_creation( + ctx, session_id, user_id, context, driver_id, activation, result, + ), SessionManagerMessage::ProbeIdleSession { session_id, instant, diff --git a/crates/sail-session/src/session_manager/actor/handler.rs b/crates/sail-session/src/session_manager/actor/handler.rs index a2e6d34ab1..18b2c90463 100644 --- a/crates/sail-session/src/session_manager/actor/handler.rs +++ b/crates/sail-session/src/session_manager/actor/handler.rs @@ -39,135 +39,213 @@ impl SessionManagerActor { user_id: String, result: oneshot::Sender>, ) -> ActorAction { - let context = if let Some(session) = self.sessions.get(&session_id) { - if let ServerSessionState::Running { context, .. } = &session.state { + if let Some(session) = self.sessions.get(&session_id) { + let context = if let ServerSessionState::Running { context, .. } = &session.state { Ok(context.clone()) } else { Err(SessionError::invalid(format!( "session {session_id} is not running" ))) + }; + if let Ok(context) = &context { + self.schedule_idle_session_probe(ctx, session_id, context); } - } else { - // TODO: The session ID is used in various storage paths, so it is assumed to be unique - // across all session managers, and it should contain only valid characters for a - // path segment. Right now the session ID is generated as a UUID by the Spark client, - // so this is true in practice, but we may still want some validation here. - let session_id = session_id.clone(); - info!("creating session {session_id}"); - let span = Span::root( - "SessionManagerActor::create_session_context", - SpanContext::random(), - ) - .with_property(|| (SpanAttribute::SESSION_ID, session_id.clone())); - let _guard = span.set_local_parent(); - let driver_id = match self.driver_id_generator.generate() { - Ok(driver_id) => driver_id, - Err(e) => { - let output = Err(SessionError::internal(e.to_string())); - let _ = result.send(output); - return ActorAction::Continue; - } + let _ = result.send(context); + return ActorAction::Continue; + } + + // TODO: The session ID is used in various storage paths, so it is assumed to be unique + // across all session managers, and it should contain only valid characters for a + // path segment. Right now the session ID is generated as a UUID by the Spark client, + // so this is true in practice, but we may still want some validation here. + info!("creating session {session_id}"); + let span = Span::root( + "SessionManagerActor::create_session_context", + SpanContext::random(), + ) + .with_property(|| (SpanAttribute::SESSION_ID, session_id.clone())); + let _guard = span.set_local_parent(); + let driver_id = match self.driver_id_generator.generate() { + Ok(driver_id) => driver_id, + Err(e) => { + let _ = result.send(Err(SessionError::internal(e.to_string()))); + return ActorAction::Continue; + } + }; + let runner = match self.job_runner_factory.create( + ctx.children_mut(), + SessionJobRunnerInfo { + session_id: session_id.clone(), + driver_id, + driver_server_port: self.driver_gateway.as_ref().map(|x| x.port()), + event_reporter: self.event_reporter.clone(), + }, + ) { + Ok(runner) => runner, + Err(e) => { + let _ = result.send(Err(e.into())); + return ActorAction::Continue; + } + }; + let (runner, driver) = runner.into_parts(); + let registered_driver_id = driver.as_ref().map(|_| driver_id); + if let Some(driver) = &driver + && let Err(e) = self.drivers.insert(driver_id, driver.clone()) + { + let session = ServerSession { + state: ServerSessionState::Failed, }; - let runner = self.job_runner_factory.create( - ctx.children_mut(), - SessionJobRunnerInfo { - session_id: session_id.clone(), - driver_id, - driver_server_port: self.driver_gateway.as_ref().map(|x| x.port()), - event_reporter: self.event_reporter.clone(), - }, - ); - match runner { - Ok(runner) => { - let (runner, driver) = runner.into_parts(); - let registered_driver_id = driver.as_ref().map(|_| driver_id); - if let Some(driver) = &driver - && let Err(e) = self.drivers.insert(driver_id, driver.clone()) - { - let session = ServerSession { - state: ServerSessionState::Failed, - }; - let status = session.state.status().to_string(); - self.sessions.insert(session_id.clone(), session); - self.event_reporter.report(SystemEvent::SessionCreated { - session_id: session_id.clone(), - user_id, - status, - created_at: Utc::now(), - }); - let driver = driver.clone(); - ctx.spawn(async move { - if let Err(e) = driver.shutdown().await { - warn!("failed to shut down driver {driver_id}: {e}"); - } - }); - let output = Err(e.into()); - let _ = result.send(output); - return ActorAction::Continue; - } - let info = ServerSessionInfo { - session_id: session_id.clone(), - user_id: user_id.clone(), - session_manager: ctx.handle().clone(), - job_runner: Some(runner), - }; - match self.session_factory.create(info) { - Ok(context) => { - if let Some(driver) = driver { - ctx.spawn(async move { - if let Err(e) = driver.activate().await { - warn!("failed to activate driver {driver_id}: {e}"); - } - }); - } - let session = ServerSession { - state: ServerSessionState::Running { - context: context.clone(), - driver_id: registered_driver_id, - }, - }; - let status = session.state.status().to_string(); - self.sessions.insert(session_id.clone(), session); - self.event_reporter.report(SystemEvent::SessionCreated { - session_id: session_id.clone(), - user_id, - status, - created_at: Utc::now(), - }); - Ok(context) + let status = session.state.status().to_string(); + self.sessions.insert(session_id.clone(), session); + self.event_reporter.report(SystemEvent::SessionCreated { + session_id, + user_id, + status, + created_at: Utc::now(), + }); + let driver = driver.clone(); + ctx.spawn(async move { + if let Err(e) = driver.shutdown().await { + warn!("failed to shut down driver {driver_id}: {e}"); + } + }); + let _ = result.send(Err(e.into())); + return ActorAction::Continue; + } + let info = ServerSessionInfo { + session_id: session_id.clone(), + user_id: user_id.clone(), + session_manager: ctx.handle().clone(), + job_runner: Some(runner), + }; + let context = match self.session_factory.create(info) { + Ok(context) => context, + Err(e) => { + if let Some(driver_id) = registered_driver_id + && let Some(driver) = self.drivers.remove(driver_id) + { + ctx.spawn(async move { + if let Err(e) = driver.shutdown().await { + warn!("failed to shut down driver {driver_id}: {e}"); } - Err(e) => { - if let Some(driver_id) = registered_driver_id - && let Some(driver) = self.drivers.remove(driver_id) - { - ctx.spawn(async move { - if let Err(e) = driver.shutdown().await { - warn!("failed to shut down driver {driver_id}: {e}"); - } - }); - } - let session = ServerSession { - state: ServerSessionState::Failed, - }; - let status = session.state.status().to_string(); - self.sessions.insert(session_id.clone(), session); - self.event_reporter.report(SystemEvent::SessionCreated { - session_id: session_id.clone(), - user_id, - status, - created_at: Utc::now(), - }); - Err(e.into()) + }); + } + let session = ServerSession { + state: ServerSessionState::Failed, + }; + let status = session.state.status().to_string(); + self.sessions.insert(session_id.clone(), session); + self.event_reporter.report(SystemEvent::SessionCreated { + session_id, + user_id, + status, + created_at: Utc::now(), + }); + let _ = result.send(Err(e.into())); + return ActorAction::Continue; + } + }; + self.sessions.insert( + session_id.clone(), + ServerSession { + state: ServerSessionState::Creating { + driver_id: registered_driver_id, + }, + }, + ); + let message = move |activation| SessionManagerMessage::CompleteSessionCreation { + session_id, + user_id, + context, + driver_id: registered_driver_id, + activation, + result, + }; + if let Some(driver) = driver { + let session_manager = ctx.handle().clone(); + ctx.spawn(async move { + let message = message(driver.activate().await); + if session_manager.send(message).await.is_err() { + warn!("failed to complete session creation"); + } + }); + } else { + ctx.send(message(Ok(()))); + } + ActorAction::Continue + } + + pub(super) fn handle_complete_session_creation( + &mut self, + ctx: &mut ActorContext, + session_id: String, + user_id: String, + context: SessionContext, + driver_id: Option, + activation: ExecutionResult<()>, + result: oneshot::Sender>, + ) -> ActorAction { + let Some(session) = self.sessions.get_mut(&session_id) else { + let _ = result.send(Err(SessionError::internal(format!( + "session {session_id} disappeared during creation" + )))); + return ActorAction::Continue; + }; + if !matches!( + session.state, + ServerSessionState::Creating { + driver_id: creating_driver_id + } if creating_driver_id == driver_id + ) { + let _ = result.send(Err(SessionError::internal(format!( + "session {session_id} creation is no longer pending" + )))); + return ActorAction::Continue; + } + let output = match activation { + Ok(()) => { + session.state = ServerSessionState::Running { + context: context.clone(), + driver_id, + }; + Ok(context) + } + Err(e) => { + if let Some(driver_id) = driver_id + && let Some(driver) = self.drivers.remove(driver_id) + { + ctx.spawn(async move { + if let Err(e) = driver.shutdown().await { + warn!("failed to shut down driver {driver_id}: {e}"); } - } + }); } - Err(e) => Err(e.into()), + session.state = ServerSessionState::Failed; + Err(e.into()) } }; - if let Ok(context) = &context - && let Ok(active_at) = context - .extension::() - .and_then(|tracker| tracker.track_activity()) + self.event_reporter.report(SystemEvent::SessionCreated { + session_id: session_id.clone(), + user_id, + status: session.state.status().to_string(), + created_at: Utc::now(), + }); + if let Ok(context) = &output { + self.schedule_idle_session_probe(ctx, session_id, context); + } + let _ = result.send(output); + ActorAction::Continue + } + + fn schedule_idle_session_probe( + &self, + ctx: &mut ActorContext, + session_id: String, + context: &SessionContext, + ) { + if let Ok(active_at) = context + .extension::() + .and_then(|tracker| tracker.track_activity()) { ctx.send_with_delay( SessionManagerMessage::ProbeIdleSession { @@ -177,8 +255,6 @@ impl SessionManagerActor { self.options.session_timeout, ); } - let _ = result.send(context); - ActorAction::Continue } pub(super) fn handle_probe_idle_session( @@ -255,7 +331,8 @@ impl SessionManagerActor { return ActorAction::Continue; }; let driver_id = match &session.state { - ServerSessionState::Running { driver_id, .. } => *driver_id, + ServerSessionState::Creating { driver_id } + | ServerSessionState::Running { driver_id, .. } => *driver_id, ServerSessionState::Deleted | ServerSessionState::Failed => None, }; if let Some(driver_id) = driver_id diff --git a/crates/sail-session/src/session_manager/actor/message.rs b/crates/sail-session/src/session_manager/actor/message.rs index e497177213..002027956c 100644 --- a/crates/sail-session/src/session_manager/actor/message.rs +++ b/crates/sail-session/src/session_manager/actor/message.rs @@ -16,6 +16,14 @@ pub enum SessionManagerMessage { user_id: String, result: oneshot::Sender>, }, + CompleteSessionCreation { + session_id: String, + user_id: String, + context: SessionContext, + driver_id: Option, + activation: ExecutionResult<()>, + result: oneshot::Sender>, + }, ProbeIdleSession { session_id: String, /// The time when the session was known to be active. @@ -41,6 +49,7 @@ impl SpanAssociation for SessionManagerMessage { fn name(&self) -> Cow<'static, str> { let name = match self { SessionManagerMessage::GetOrCreateSession { .. } => "GetOrCreateSession", + SessionManagerMessage::CompleteSessionCreation { .. } => "CompleteSessionCreation", SessionManagerMessage::ProbeIdleSession { .. } => "ProbeIdleSession", SessionManagerMessage::DeleteSession { .. } => "DeleteSession", SessionManagerMessage::SetSessionFailure { .. } => "SetSessionFailure", @@ -58,6 +67,14 @@ impl SpanAssociation for SessionManagerMessage { user_id: _, result: _, } + | SessionManagerMessage::CompleteSessionCreation { + session_id, + user_id: _, + context: _, + driver_id: _, + activation: _, + result: _, + } | SessionManagerMessage::ProbeIdleSession { session_id, instant: _, diff --git a/crates/sail-session/src/session_manager/session.rs b/crates/sail-session/src/session_manager/session.rs index 0dbeee5dae..cde697664c 100644 --- a/crates/sail-session/src/session_manager/session.rs +++ b/crates/sail-session/src/session_manager/session.rs @@ -6,6 +6,9 @@ pub struct ServerSession { } pub enum ServerSessionState { + Creating { + driver_id: Option, + }, Running { context: SessionContext, driver_id: Option, @@ -17,6 +20,7 @@ pub enum ServerSessionState { impl ServerSessionState { pub fn status(&self) -> &'static str { match self { + ServerSessionState::Creating { .. } => "CREATING", ServerSessionState::Running { .. } => "RUNNING", ServerSessionState::Deleted => "DELETED", ServerSessionState::Failed => "FAILED", From 27f5c3a766089fdddea7eb11b47c596eebd68cb0 Mon Sep 17 00:00:00 2001 From: Heran Lin Date: Fri, 4 Sep 2026 18:26:54 +0800 Subject: [PATCH 2/2] Update --- .../sail-execution/src/driver/actor/core.rs | 8 +- .../src/driver/actor/handler.rs | 130 ++++----- .../src/driver/actor/message.rs | 15 +- crates/sail-execution/src/driver/actor/mod.rs | 4 +- .../src/driver/actor/options.rs | 14 +- crates/sail-execution/src/driver/mod.rs | 1 + .../src/driver/task_assigner/core.rs | 131 ++------- .../src/driver/task_assigner/state.rs | 9 +- .../src/driver/worker_pool/core.rs | 22 +- .../src/driver/worker_pool/mod.rs | 1 - .../src/driver/worker_pool/state.rs | 16 -- .../src/driver/worker_scaler/core.rs | 253 ++++++++++++++++++ .../src/driver/worker_scaler/mod.rs | 28 ++ .../src/driver/worker_scaler/options.rs | 25 ++ .../src/driver/worker_scaler/state.rs | 37 +++ crates/sail-execution/src/id.rs | 1 + .../src/session_factory/job_runner.rs | 14 +- .../src/session_manager/actor/core.rs | 3 +- .../src/session_manager/actor/handler.rs | 100 ++++--- .../src/session_manager/actor/message.rs | 2 - .../src/session_manager/session.rs | 4 + 21 files changed, 530 insertions(+), 288 deletions(-) create mode 100644 crates/sail-execution/src/driver/worker_scaler/core.rs create mode 100644 crates/sail-execution/src/driver/worker_scaler/mod.rs create mode 100644 crates/sail-execution/src/driver/worker_scaler/options.rs create mode 100644 crates/sail-execution/src/driver/worker_scaler/state.rs diff --git a/crates/sail-execution/src/driver/actor/core.rs b/crates/sail-execution/src/driver/actor/core.rs index bb8e329015..f2295484d3 100644 --- a/crates/sail-execution/src/driver/actor/core.rs +++ b/crates/sail-execution/src/driver/actor/core.rs @@ -12,6 +12,7 @@ use sail_common::actor::{Actor, ActorAction, ActorContext}; use crate::driver::job_scheduler::{JobScheduler, JobSchedulerOptions}; use crate::driver::task_assigner::{TaskAssigner, TaskAssignerOptions}; use crate::driver::worker_pool::{WorkerPool, WorkerPoolOptions}; +use crate::driver::worker_scaler::{WorkerScaler, WorkerScalerOptions}; use crate::driver::{DriverActor, DriverComponents, DriverMessage, DriverOptions}; use crate::shuffle::{ShuffleBackendKind, celeborn_application_id}; use crate::stream::celeborn::CelebornStreamManager; @@ -43,15 +44,16 @@ impl Actor for DriverActor { ); let job_scheduler = JobScheduler::new(JobSchedulerOptions::from(&options), event_reporter); let task_assigner = TaskAssigner::new(TaskAssignerOptions::from(&options)); + let worker_scaler = WorkerScaler::new(WorkerScalerOptions::from(&options)); Self { options, worker_pool, job_scheduler, task_assigner, + worker_scaler, task_runner: None, extensions: Default::default(), activated: false, - worker_launch_retries_exhausted: false, task_sequences: HashMap::new(), shutdown_notifier: None, } @@ -151,8 +153,8 @@ impl Actor for DriverActor { DriverMessage::WorkerFailedToStart { worker_id, message } => { self.handle_worker_failed_to_start(ctx, worker_id, message) } - DriverMessage::RetryWorkerLaunch { worker_id, launch } => { - self.handle_retry_worker_launch(ctx, worker_id, launch) + DriverMessage::RetryWorkerDemand { request } => { + self.handle_retry_worker_demand(ctx, request) } DriverMessage::ProbeIdleWorker { worker_id, instant } => { self.handle_probe_idle_worker(ctx, worker_id, instant) diff --git a/crates/sail-execution/src/driver/actor/handler.rs b/crates/sail-execution/src/driver/actor/handler.rs index 2d152feb31..49e7eb9f92 100644 --- a/crates/sail-execution/src/driver/actor/handler.rs +++ b/crates/sail-execution/src/driver/actor/handler.rs @@ -16,7 +16,7 @@ use tokio::time::Instant; use crate::driver::actor::DriverActor; use crate::driver::job_scheduler::{JobAction, TaskState}; use crate::driver::output::JobOutputItem; -use crate::driver::worker_pool::{WorkerLaunch, WorkerLaunchReason}; +use crate::driver::worker_scaler::{WorkerLaunchRequest, WorkerRetryRequest}; use crate::driver::{DriverMessage, TaskStatus}; use crate::error::{ExecutionError, ExecutionResult}; use crate::id::{JobId, TaskKey, TaskKeyDisplay, TaskStreamKey, TaskStreamKeyDisplay, WorkerId}; @@ -51,7 +51,9 @@ impl DriverActor { let count = self .task_assigner .request_initial_workers(self.options.worker_initial_count); - self.start_workers(ctx, count, WorkerLaunchReason::Initial) + self.worker_scaler + .request_initial_workers(count) + .and_then(|requests| self.launch_worker_requests(ctx, requests)) .inspect(|_| { self.activated = true; }) @@ -73,10 +75,10 @@ impl DriverActor { info!("worker {worker_id} is available at {host}:{port}"); let out = self.worker_pool.register_worker(ctx, worker_id, host, port); if out.is_ok() { - self.worker_launch_retries_exhausted = false; + self.worker_scaler.worker_registered(worker_id); self.task_assigner.activate_worker(worker_id); self.run_tasks(ctx); - self.scale_up_workers(ctx); + self.reconcile_worker_demands(ctx); } if result.send(out).is_err() { warn!("failed to send worker registration result"); @@ -127,34 +129,17 @@ impl DriverActor { ActorAction::Continue } - pub(super) fn handle_retry_worker_launch( + pub(super) fn handle_retry_worker_demand( &mut self, ctx: &mut ActorContext, - worker_id: WorkerId, - launch: WorkerLaunch, + request: WorkerRetryRequest, ) -> ActorAction { - if !self.task_assigner.track_worker_failed_to_start(worker_id) { - return ActorAction::Continue; - } - let should_retry = match launch.reason { - WorkerLaunchReason::Initial => { - self.task_assigner - .request_initial_workers(self.options.worker_initial_count) - > 0 - } - WorkerLaunchReason::Demand => self.task_assigner.request_workers() > 0, - }; - if should_retry { - match self.worker_pool.start_worker(ctx, launch) { - Ok(worker_id) => self.task_assigner.track_pending_worker(worker_id), - Err(e) => { - error!("failed to retry worker launch: {e}"); - ctx.send(DriverMessage::Shutdown { result: None }); - return ActorAction::Continue; - } - } + if let Some(request) = self.worker_scaler.retry(request) + && let Err(e) = self.launch_worker_request(ctx, request) + { + error!("failed to retry worker launch: {e}"); + ctx.send(DriverMessage::Shutdown { result: None }); } - self.scale_up_workers(ctx); ActorAction::Continue } @@ -213,7 +198,7 @@ impl DriverActor { for job_id in job_ids { self.refresh_job(ctx, job_id); self.run_tasks(ctx); - self.scale_up_workers(ctx); + self.reconcile_worker_demands(ctx); } } ActorAction::Continue @@ -230,7 +215,7 @@ impl DriverActor { if let Ok((job_id, _)) = &out { self.refresh_job(ctx, *job_id); self.run_tasks(ctx); - self.scale_up_workers(ctx); + self.reconcile_worker_demands(ctx); } let _ = result.send(out.map(|(_, stream)| stream)); ActorAction::Continue @@ -242,6 +227,7 @@ impl DriverActor { job_id: JobId, ) -> ActorAction { self.clean_up_job(ctx, job_id); + self.reconcile_worker_demands(ctx); ActorAction::Continue } @@ -278,7 +264,7 @@ impl DriverActor { self.task_assigner.unassign_task(&key); self.refresh_job(ctx, key.job_id); self.run_tasks(ctx); - self.scale_up_workers(ctx); + self.reconcile_worker_demands(ctx); } TaskStatus::Failed => { // Some canceled tasks may report failed status due to closed streams, @@ -288,7 +274,7 @@ impl DriverActor { self.task_assigner.unassign_task(&key); self.refresh_job(ctx, key.job_id); self.run_tasks(ctx); - self.scale_up_workers(ctx); + self.reconcile_worker_demands(ctx); } TaskStatus::Canceled => { // The task attempt state should already be "canceled" but we update it @@ -318,15 +304,15 @@ impl DriverActor { // so reschedule the probe instead of failing. This keeps long, // many-stage jobs alive while the worker pool scales between stages. // It cannot loop forever: each worker launch has a finite retry - // schedule. A failed worker remains pending here while waiting for - // its retry so the replacement capacity is not requested twice. + // schedule. A worker demand remains pending while waiting for its + // retry so the replacement capacity is not requested twice. // // Re-probe at `worker_launch_timeout` (capped by `task_launch_timeout`) // rather than a full `task_launch_timeout`: that is the window a // pending worker takes to register or be failed, so once the last // pending worker resolves the task fails promptly instead of waiting // another full launch window. - if self.task_assigner.has_pending_workers() { + if self.worker_scaler.has_pending_worker_demands() { let delay = self .options .worker_launch_timeout @@ -621,34 +607,44 @@ impl DriverActor { } } - fn scale_up_workers(&mut self, ctx: &mut ActorContext) { - if self.worker_launch_retries_exhausted { - if self.task_assigner.has_worker_demand() { - return; - } - self.worker_launch_retries_exhausted = false; - } - let count = self.task_assigner.request_workers(); - if let Err(e) = self.start_workers(ctx, count, WorkerLaunchReason::Demand) { + fn reconcile_worker_demands(&mut self, ctx: &mut ActorContext) { + let output = self + .worker_scaler + .reconcile(self.task_assigner.count_worker_demands()) + .and_then(|requests| self.launch_worker_requests(ctx, requests)); + if let Err(e) = output { error!("failed to request workers: {e}"); ctx.send(DriverMessage::Shutdown { result: None }); } } - fn start_workers( + fn launch_worker_request( + &mut self, + ctx: &mut ActorContext, + request: WorkerLaunchRequest, + ) -> ExecutionResult<()> { + info!( + "launching worker demand {} attempt {}", + request.demand_id, request.attempt + ); + let worker_id = self.worker_pool.start_worker(ctx)?; + if self.worker_scaler.bind_worker(request, worker_id) { + Ok(()) + } else { + Err(ExecutionError::InternalError(format!( + "failed to bind worker {worker_id} to demand {}", + request.demand_id + ))) + } + } + + fn launch_worker_requests( &mut self, ctx: &mut ActorContext, - count: usize, - reason: WorkerLaunchReason, + requests: Vec, ) -> ExecutionResult<()> { - for _ in 0..count { - let launch = WorkerLaunch { - reason, - attempt: 0, - retries: self.options.worker_launch_retry_strategy.retries(), - }; - let worker_id = self.worker_pool.start_worker(ctx, launch)?; - self.task_assigner.track_pending_worker(worker_id); + for request in requests { + self.launch_worker_request(ctx, request)?; } Ok(()) } @@ -659,26 +655,16 @@ impl DriverActor { worker_id: WorkerId, message: String, ) { - let Some(mut launch) = self.worker_pool.fail_worker_if_pending(worker_id, message) else { + if !self.worker_pool.fail_worker_if_pending(worker_id, message) { return; - }; - if let Some(step) = launch.retries.next() { - launch.attempt = step.retry; - warn!( - "scheduling worker {worker_id} launch retry {} in {:?}", - step.retry, step.delay, - ); - ctx.send_with_delay( - DriverMessage::RetryWorkerLaunch { worker_id, launch }, - step.delay, - ); - } else { - self.task_assigner.track_worker_failed_to_start(worker_id); - self.worker_launch_retries_exhausted = self.task_assigner.has_worker_demand(); + } + if let Some(request) = self.worker_scaler.worker_failed(worker_id) { warn!( - "worker {worker_id} launch retries exhausted after attempt {}", - launch.attempt + "scheduling worker demand {} launch retry {} in {:?}", + request.demand_id, request.attempt, request.delay, ); + ctx.send_with_delay(DriverMessage::RetryWorkerDemand { request }, request.delay); } + self.reconcile_worker_demands(ctx); } } diff --git a/crates/sail-execution/src/driver/actor/message.rs b/crates/sail-execution/src/driver/actor/message.rs index e3032d719a..21743799f0 100644 --- a/crates/sail-execution/src/driver/actor/message.rs +++ b/crates/sail-execution/src/driver/actor/message.rs @@ -14,7 +14,7 @@ use tokio::sync::oneshot; use tokio::time::Instant; use crate::driver::r#gen; -use crate::driver::worker_pool::WorkerLaunch; +use crate::driver::worker_scaler::WorkerRetryRequest; use crate::error::ExecutionResult; use crate::id::{JobId, TaskKey, TaskStreamKey, WorkerId}; use crate::stream::reader::TaskStreamSource; @@ -43,10 +43,8 @@ pub enum DriverMessage { worker_id: WorkerId, message: String, }, - RetryWorkerLaunch { - /// The failed worker whose reserved capacity this retry replaces. - worker_id: WorkerId, - launch: WorkerLaunch, + RetryWorkerDemand { + request: WorkerRetryRequest, }, ProbeIdleWorker { worker_id: WorkerId, @@ -145,7 +143,7 @@ impl SpanAssociation for DriverMessage { DriverMessage::WorkerKnownPeers { .. } => "WorkerKnownPeers", DriverMessage::ProbePendingWorker { .. } => "ProbePendingWorker", DriverMessage::WorkerFailedToStart { .. } => "WorkerFailedToStart", - DriverMessage::RetryWorkerLaunch { .. } => "RetryWorkerLaunch", + DriverMessage::RetryWorkerDemand { .. } => "RetryWorkerDemand", DriverMessage::ProbeIdleWorker { .. } => "ProbeIdleWorker", DriverMessage::ProbeLostWorker { .. } => "ProbeLostWorker", DriverMessage::ExecuteJob { .. } => "ExecuteJob", @@ -194,9 +192,8 @@ impl SpanAssociation for DriverMessage { } => { p.push((SpanAttribute::CLUSTER_WORKER_ID, worker_id.to_string())); } - DriverMessage::RetryWorkerLaunch { worker_id, launch } => { - p.push((SpanAttribute::CLUSTER_WORKER_ID, worker_id.to_string())); - p.push((SpanAttribute::RETRY_ATTEMPT, launch.attempt.to_string())); + DriverMessage::RetryWorkerDemand { request } => { + p.push((SpanAttribute::RETRY_ATTEMPT, request.attempt.to_string())); } DriverMessage::ExecuteJob { plan: _, diff --git a/crates/sail-execution/src/driver/actor/mod.rs b/crates/sail-execution/src/driver/actor/mod.rs index 270e318631..cc8e54ca1d 100644 --- a/crates/sail-execution/src/driver/actor/mod.rs +++ b/crates/sail-execution/src/driver/actor/mod.rs @@ -14,6 +14,7 @@ use tokio::sync::oneshot; use crate::driver::job_scheduler::JobScheduler; use crate::driver::task_assigner::TaskAssigner; use crate::driver::worker_pool::WorkerPool; +use crate::driver::worker_scaler::WorkerScaler; use crate::id::TaskKey; use crate::task_runner::TaskRunnerActor; @@ -27,11 +28,10 @@ pub struct DriverActor { worker_pool: WorkerPool, job_scheduler: JobScheduler, task_assigner: TaskAssigner, + worker_scaler: WorkerScaler, task_runner: Option>, extensions: DriverExtensions, activated: bool, - /// Whether launch retries have been exhausted for the current queued worker demand. - worker_launch_retries_exhausted: bool, /// The sequence number corresponding to the last task status update from the worker. /// A different sequence number is tracked for each attempt. task_sequences: HashMap, diff --git a/crates/sail-execution/src/driver/actor/options.rs b/crates/sail-execution/src/driver/actor/options.rs index d8fd3c2539..678066ddbe 100644 --- a/crates/sail-execution/src/driver/actor/options.rs +++ b/crates/sail-execution/src/driver/actor/options.rs @@ -5,6 +5,7 @@ use sail_common::runtime::RuntimeHandle; use sail_common::utils::retry::RetryStrategy; use sail_telemetry::events::SystemEventReporter; +use crate::error::{ExecutionError, ExecutionResult}; use crate::id::DriverId; use crate::shuffle::ShuffleBackendKind; use crate::worker_manager::WorkerManager; @@ -40,14 +41,19 @@ pub struct DriverComponents { } impl DriverOptions { - pub fn new( + pub fn try_new( config: &AppConfig, runtime: RuntimeHandle, session_id: String, driver_id: DriverId, driver_server_port: u16, - ) -> Self { - Self { + ) -> ExecutionResult { + if config.cluster.worker_task_slots == 0 { + return Err(ExecutionError::InvalidArgument( + "worker task slots must be greater than zero".to_string(), + )); + } + Ok(Self { enable_tls: config.cluster.enable_tls, session_id, driver_id, @@ -75,6 +81,6 @@ impl DriverOptions { task_max_attempts: config.cluster.task_max_attempts, shuffle_backend: (&config.cluster.shuffle_backend).into(), runtime, - } + }) } } diff --git a/crates/sail-execution/src/driver/mod.rs b/crates/sail-execution/src/driver/mod.rs index 8f116ee8e0..49187f8dcd 100644 --- a/crates/sail-execution/src/driver/mod.rs +++ b/crates/sail-execution/src/driver/mod.rs @@ -8,6 +8,7 @@ mod registry; mod server; mod task_assigner; pub(super) mod worker_pool; +mod worker_scaler; #[expect(clippy::allow_attributes)] pub(crate) mod r#gen { diff --git a/crates/sail-execution/src/driver/task_assigner/core.rs b/crates/sail-execution/src/driver/task_assigner/core.rs index d4205dfd7d..e44ab1b377 100644 --- a/crates/sail-execution/src/driver/task_assigner/core.rs +++ b/crates/sail-execution/src/driver/task_assigner/core.rs @@ -12,7 +12,7 @@ use crate::task::scheduling::{ }; impl TaskAssigner { - pub fn request_workers(&self) -> usize { + pub fn count_worker_demands(&self) -> usize { let enqueued_slots = self.enqueued_worker_slots(); let vacant_slots = self .workers @@ -21,18 +21,10 @@ impl TaskAssigner { WorkerResource::Active { task_slots, .. } => { task_slots.iter().filter(|x| x.is_vacant()).count() } - WorkerResource::Pending | WorkerResource::Inactive => 0, + WorkerResource::Inactive => 0, }) .sum::(); - let pending_workers = self - .workers - .values() - .filter(|worker| matches!(worker, WorkerResource::Pending)) - .count(); - let pending_slots = pending_workers.saturating_mul(self.options.worker_task_slots); - let required_slots = enqueued_slots - .saturating_sub(vacant_slots) - .saturating_sub(pending_slots); + let required_slots = enqueued_slots.saturating_sub(vacant_slots); let active_workers = self .workers .values() @@ -41,78 +33,24 @@ impl TaskAssigner { let allowed_workers = if self.options.worker_max_count == 0 { usize::MAX } else { - self.options - .worker_max_count - .saturating_sub(pending_workers) - .saturating_sub(active_workers) + self.options.worker_max_count.saturating_sub(active_workers) }; - if self.options.worker_task_slots == 0 { - error!("worker task slots must be greater than zero"); - return 0; - } required_slots .div_ceil(self.options.worker_task_slots) .min(allowed_workers) } - pub fn has_worker_demand(&self) -> bool { - self.enqueued_worker_slots() > 0 - } - - pub fn has_pending_workers(&self) -> bool { - self.workers - .values() - .any(|worker| matches!(worker, WorkerResource::Pending)) - } - pub fn request_initial_workers(&self, worker_initial_count: usize) -> usize { - let live_workers = self + let active_workers = self .workers .values() - .filter(|worker| { - matches!( - worker, - WorkerResource::Pending | WorkerResource::Active { .. } - ) - }) + .filter(|worker| matches!(worker, WorkerResource::Active { .. })) .count(); - let requested_workers = worker_initial_count.saturating_sub(live_workers); + let requested_workers = worker_initial_count.saturating_sub(active_workers); if self.options.worker_max_count == 0 { requested_workers } else { - requested_workers.min(self.options.worker_max_count.saturating_sub(live_workers)) - } - } - - pub fn track_pending_worker(&mut self, worker_id: WorkerId) { - match self.workers.entry(worker_id) { - indexmap::map::Entry::Vacant(entry) => { - entry.insert(WorkerResource::Pending); - } - indexmap::map::Entry::Occupied(_) => { - warn!("worker {worker_id} is already tracked"); - } - } - } - - pub fn track_worker_failed_to_start(&mut self, worker_id: WorkerId) -> bool { - let Some(worker) = self.workers.get_mut(&worker_id) else { - warn!("worker {worker_id} not found"); - return false; - }; - match worker { - WorkerResource::Pending => { - *worker = WorkerResource::Inactive; - true - } - WorkerResource::Active { .. } => { - warn!("worker {worker_id} is already active"); - false - } - WorkerResource::Inactive => { - warn!("worker {worker_id} is already inactive"); - false - } + requested_workers.min(self.options.worker_max_count.saturating_sub(active_workers)) } } @@ -127,11 +65,11 @@ impl TaskAssigner { entry.insert(resource); } indexmap::map::Entry::Occupied(mut entry) => match entry.get() { - WorkerResource::Pending => { + WorkerResource::Active { .. } => warn!("worker {worker_id} is already active"), + WorkerResource::Inactive => { + warn!("worker {worker_id} was inactive"); entry.insert(resource); } - WorkerResource::Active { .. } => warn!("worker {worker_id} is already active"), - WorkerResource::Inactive => warn!("worker {worker_id} is inactive"), }, } } @@ -142,7 +80,7 @@ impl TaskAssigner { return; }; match worker { - WorkerResource::Pending | WorkerResource::Active { .. } => { + WorkerResource::Active { .. } => { *worker = WorkerResource::Inactive; } WorkerResource::Inactive => { @@ -293,7 +231,7 @@ impl TaskAssigner { task_slots: slots, local_streams: streams, } => slots.iter().all(|s| s.is_vacant()) && streams.is_empty(), - WorkerResource::Pending | WorkerResource::Inactive => false, + WorkerResource::Inactive => false, } } @@ -309,7 +247,7 @@ impl TaskAssigner { .iter() .flat_map(|x| x.list_tasks().cloned().collect::>()) .collect(), - WorkerResource::Pending | WorkerResource::Inactive => vec![], + WorkerResource::Inactive => vec![], } } @@ -333,7 +271,7 @@ impl TaskAssigner { .iter() .filter_map(|(id, worker)| { let slots = match worker { - WorkerResource::Pending | WorkerResource::Inactive => vec![], + WorkerResource::Inactive => vec![], WorkerResource::Active { task_slots: slots, .. } => slots @@ -432,50 +370,29 @@ mod tests { } #[test] - fn pending_worker_capacity_prevents_duplicate_requests() { + fn worker_demand_accounts_for_vacant_slots() { let mut assigner = task_assigner(8, 4); assigner.enqueue_tasks(worker_region(8)); - assert_eq!(assigner.request_workers(), 1); - assigner.track_pending_worker(WorkerId::from(1)); - assert_eq!(assigner.request_workers(), 0); - - assigner.track_worker_failed_to_start(WorkerId::from(1)); - assert_eq!(assigner.request_workers(), 1); + assert_eq!(assigner.count_worker_demands(), 1); + assigner.activate_worker(WorkerId::from(1)); + assert_eq!(assigner.count_worker_demands(), 0); } #[test] - fn worker_requests_respect_pending_and_active_worker_limit() { + fn worker_demand_respects_active_worker_limit() { let mut assigner = task_assigner(1, 2); assigner.enqueue_tasks(worker_region(4)); - assigner.track_pending_worker(WorkerId::from(1)); - assigner.track_pending_worker(WorkerId::from(2)); - - assert_eq!(assigner.request_workers(), 0); - assigner.activate_worker(WorkerId::from(1)); - assigner.track_worker_failed_to_start(WorkerId::from(2)); - assert_eq!(assigner.request_workers(), 1); + + assert_eq!(assigner.count_worker_demands(), 1); } #[test] - fn initial_workers_share_the_worker_limit() { + fn initial_workers_respect_the_worker_limit() { let mut assigner = task_assigner(8, 4); - assigner.track_pending_worker(WorkerId::from(1)); + assigner.activate_worker(WorkerId::from(1)); assert_eq!(assigner.request_initial_workers(4), 3); - - assigner.track_pending_worker(WorkerId::from(2)); - assigner.track_pending_worker(WorkerId::from(3)); - assigner.track_pending_worker(WorkerId::from(4)); - assert_eq!(assigner.request_initial_workers(4), 0); - } - - #[test] - fn zero_worker_task_slots_does_not_panic() { - let mut assigner = task_assigner(0, 1); - assigner.enqueue_tasks(worker_region(1)); - - assert_eq!(assigner.request_workers(), 0); } } diff --git a/crates/sail-execution/src/driver/task_assigner/state.rs b/crates/sail-execution/src/driver/task_assigner/state.rs index 6a839387ab..51ca7abaac 100644 --- a/crates/sail-execution/src/driver/task_assigner/state.rs +++ b/crates/sail-execution/src/driver/task_assigner/state.rs @@ -90,7 +90,6 @@ impl DriverResource { #[derive(Debug)] /// Represents the current state of a worker's resources as seen by the task assigner. pub enum WorkerResource { - Pending, Active { /// The task slots on the worker. task_slots: Vec, @@ -123,7 +122,7 @@ impl WorkerResource { warn!("invalid task slot {slot} on worker"); } } - WorkerResource::Pending | WorkerResource::Inactive => { + WorkerResource::Inactive => { warn!("cannot add tasks to inactive worker"); } } @@ -139,7 +138,7 @@ impl WorkerResource { false } } - WorkerResource::Pending | WorkerResource::Inactive => { + WorkerResource::Inactive => { warn!("cannot remove tasks from inactive worker"); false } @@ -151,7 +150,7 @@ impl WorkerResource { WorkerResource::Active { local_streams, .. } => { local_streams.extend(set.local_streams().cloned()); } - WorkerResource::Pending | WorkerResource::Inactive => { + WorkerResource::Inactive => { warn!("cannot track local streams on inactive worker"); } } @@ -168,7 +167,7 @@ impl WorkerResource { } count != local_streams.len() } - WorkerResource::Pending | WorkerResource::Inactive => { + WorkerResource::Inactive => { warn!("cannot untrack local streams from inactive worker"); false } diff --git a/crates/sail-execution/src/driver/worker_pool/core.rs b/crates/sail-execution/src/driver/worker_pool/core.rs index 9b82c913b0..3428f38965 100644 --- a/crates/sail-execution/src/driver/worker_pool/core.rs +++ b/crates/sail-execution/src/driver/worker_pool/core.rs @@ -16,7 +16,7 @@ use tokio::time::Instant; use tonic::Code; use crate::driver::worker_pool::state::WorkerState; -use crate::driver::worker_pool::{WorkerDescriptor, WorkerLaunch, WorkerPool, WorkerPoolOptions}; +use crate::driver::worker_pool::{WorkerDescriptor, WorkerPool, WorkerPoolOptions}; use crate::driver::{DriverActor, DriverMessage, TaskStatus}; use crate::error::{ExecutionError, ExecutionResult}; use crate::id::{JobId, TaskKey, TaskKeyDisplay, TaskStreamKey, WorkerId}; @@ -41,16 +41,11 @@ impl WorkerPool { pub fn start_worker( &mut self, ctx: &mut ActorContext, - launch: WorkerLaunch, ) -> ExecutionResult { let worker_id = self.worker_id_generator.generate()?; - info!( - "starting worker {worker_id} (launch attempt {})", - launch.attempt - ); + info!("starting worker {worker_id}"); let descriptor = WorkerDescriptor { state: WorkerState::Pending, - launch: Some(launch), messages: vec![], peers: HashSet::new(), }; @@ -136,7 +131,6 @@ impl WorkerPool { heartbeat_at: Instant::now(), client: None, }; - worker.launch = None; Self::schedule_lost_worker_probe(ctx, worker_id, worker, &self.options); Self::schedule_idle_worker_probe(ctx, worker_id, worker, &self.options); event_reporter.report(SystemEvent::WorkerUpdated { @@ -269,16 +263,12 @@ impl WorkerPool { worker.peers.extend(peer_worker_ids); } - pub fn fail_worker_if_pending( - &mut self, - worker_id: WorkerId, - message: String, - ) -> Option { + pub fn fail_worker_if_pending(&mut self, worker_id: WorkerId, message: String) -> bool { let event_reporter = self.event_reporter.clone(); let session_id = self.options.session_id.clone(); let Some(worker) = self.workers.get_mut(&worker_id) else { warn!("worker {worker_id} not found"); - return None; + return false; }; if matches!(&worker.state, WorkerState::Pending) { warn!("worker {worker_id} failed to start: {message}"); @@ -292,9 +282,9 @@ impl WorkerPool { status: worker.state.status().to_string(), updated_at: Utc::now(), }); - worker.launch.take() + true } else { - None + false } } diff --git a/crates/sail-execution/src/driver/worker_pool/mod.rs b/crates/sail-execution/src/driver/worker_pool/mod.rs index 06d402d0d2..de084b2abe 100644 --- a/crates/sail-execution/src/driver/worker_pool/mod.rs +++ b/crates/sail-execution/src/driver/worker_pool/mod.rs @@ -7,7 +7,6 @@ use std::sync::Arc; use indexmap::IndexMap; pub use options::WorkerPoolOptions; use sail_telemetry::events::SystemEventReporter; -pub(crate) use state::{WorkerLaunch, WorkerLaunchReason}; use crate::driver::worker_pool::state::WorkerDescriptor; use crate::id::{IdGenerator, WorkerId}; diff --git a/crates/sail-execution/src/driver/worker_pool/state.rs b/crates/sail-execution/src/driver/worker_pool/state.rs index b74b55a8d4..53399caff0 100644 --- a/crates/sail-execution/src/driver/worker_pool/state.rs +++ b/crates/sail-execution/src/driver/worker_pool/state.rs @@ -1,6 +1,5 @@ use std::collections::HashSet; -use sail_common::utils::retry::RetrySchedule; use tokio::time::Instant; use crate::id::WorkerId; @@ -8,7 +7,6 @@ use crate::worker::WorkerClientSet; pub struct WorkerDescriptor { pub state: WorkerState, - pub launch: Option, pub messages: Vec, /// A list of peer workers known to the worker. /// The list may or may not cover all the running workers, @@ -18,20 +16,6 @@ pub struct WorkerDescriptor { pub peers: HashSet, } -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum WorkerLaunchReason { - Initial, - Demand, -} - -#[derive(Debug, Clone)] -pub struct WorkerLaunch { - pub reason: WorkerLaunchReason, - /// The launch attempt number. The initial attempt is zero. - pub attempt: usize, - pub retries: RetrySchedule, -} - pub enum WorkerState { Pending, Running { diff --git a/crates/sail-execution/src/driver/worker_scaler/core.rs b/crates/sail-execution/src/driver/worker_scaler/core.rs new file mode 100644 index 0000000000..eae1b5f347 --- /dev/null +++ b/crates/sail-execution/src/driver/worker_scaler/core.rs @@ -0,0 +1,253 @@ +use log::warn; + +use crate::driver::worker_scaler::state::{WorkerDemand, WorkerDemandState}; +use crate::driver::worker_scaler::{ + WorkerDemandReason, WorkerLaunchRequest, WorkerRetryRequest, WorkerScaler, +}; +use crate::error::ExecutionResult; +use crate::id::{WorkerDemandId, WorkerId}; + +impl WorkerScaler { + pub fn request_initial_workers( + &mut self, + count: usize, + ) -> ExecutionResult> { + self.create_demands(count, WorkerDemandReason::Initial) + } + + pub fn reconcile(&mut self, target: usize) -> ExecutionResult> { + let initial = self + .demands + .values() + .filter(|demand| matches!(demand.reason, WorkerDemandReason::Initial)) + .count(); + let target = target.saturating_sub(initial); + self.remove_surplus_task_demands(target); + + let current = self + .demands + .values() + .filter(|demand| matches!(demand.reason, WorkerDemandReason::Task)) + .count(); + self.create_demands(target.saturating_sub(current), WorkerDemandReason::Task) + } + + pub fn bind_worker(&mut self, request: WorkerLaunchRequest, worker_id: WorkerId) -> bool { + let Some(demand) = self.demands.get_mut(&request.demand_id) else { + warn!("worker demand {} not found", request.demand_id); + return false; + }; + if !matches!( + demand.state, + WorkerDemandState::Created { attempt } if attempt == request.attempt + ) { + warn!("worker demand {} is not ready to launch", request.demand_id); + return false; + } + demand.state = WorkerDemandState::Launching { + worker_id, + attempt: request.attempt, + }; + self.workers.insert(worker_id, request.demand_id); + true + } + + pub fn worker_registered(&mut self, worker_id: WorkerId) -> bool { + let Some(demand_id) = self.workers.swap_remove(&worker_id) else { + warn!("worker {worker_id} is not associated with a demand"); + return false; + }; + self.demands.swap_remove(&demand_id).is_some() + } + + pub fn worker_failed(&mut self, worker_id: WorkerId) -> Option { + let Some(demand_id) = self.workers.swap_remove(&worker_id) else { + warn!("worker {worker_id} is not associated with a demand"); + return None; + }; + let attempt = match self.demands.get(&demand_id).map(|demand| &demand.state) { + Some(WorkerDemandState::Launching { + worker_id: launching_worker_id, + attempt, + }) if *launching_worker_id == worker_id => *attempt, + Some(_) => { + warn!("worker demand {demand_id} is not launching worker {worker_id}"); + return None; + } + None => { + warn!("worker demand {demand_id} not found"); + return None; + } + }; + self.fail_demand(demand_id, attempt) + } + + pub fn retry(&mut self, request: WorkerRetryRequest) -> Option { + let demand = self.demands.get_mut(&request.demand_id)?; + if !matches!( + demand.state, + WorkerDemandState::WaitingForRetry { attempt } if attempt == request.attempt + ) { + return None; + } + demand.state = WorkerDemandState::Created { + attempt: request.attempt, + }; + Some(WorkerLaunchRequest { + demand_id: request.demand_id, + attempt: request.attempt, + }) + } + + pub fn has_pending_worker_demands(&self) -> bool { + self.demands.values().any(|demand| { + matches!( + demand.state, + WorkerDemandState::Created { .. } + | WorkerDemandState::Launching { .. } + | WorkerDemandState::WaitingForRetry { .. } + ) + }) + } + + fn create_demands( + &mut self, + count: usize, + reason: WorkerDemandReason, + ) -> ExecutionResult> { + let mut requests = Vec::with_capacity(count); + for _ in 0..count { + let demand_id = self.worker_demand_id_generator.generate()?; + let attempt = 0; + self.demands.insert( + demand_id, + WorkerDemand { + reason, + state: WorkerDemandState::Created { attempt }, + retries: self.options.worker_launch_retry_strategy.retries(), + }, + ); + requests.push(WorkerLaunchRequest { demand_id, attempt }); + } + Ok(requests) + } + + fn fail_demand( + &mut self, + demand_id: WorkerDemandId, + attempt: usize, + ) -> Option { + let demand = self.demands.get_mut(&demand_id)?; + if let Some(step) = demand.retries.next() { + demand.state = WorkerDemandState::WaitingForRetry { + attempt: step.retry, + }; + Some(WorkerRetryRequest { + demand_id, + attempt: step.retry, + delay: step.delay, + }) + } else { + warn!("worker demand {demand_id} launch retries exhausted after attempt {attempt}"); + let reason = demand.reason; + if matches!(reason, WorkerDemandReason::Task) { + demand.state = WorkerDemandState::Exhausted; + } + if matches!(reason, WorkerDemandReason::Initial) { + self.demands.swap_remove(&demand_id); + } + None + } + } + + fn remove_surplus_task_demands(&mut self, target: usize) { + let mut surplus = self + .demands + .values() + .filter(|demand| matches!(demand.reason, WorkerDemandReason::Task)) + .count() + .saturating_sub(target); + if surplus == 0 { + return; + } + + let predicates: [fn(&WorkerDemandState) -> bool; 3] = [ + |state: &WorkerDemandState| matches!(state, WorkerDemandState::Exhausted), + |state: &WorkerDemandState| matches!(state, WorkerDemandState::WaitingForRetry { .. }), + |state: &WorkerDemandState| matches!(state, WorkerDemandState::Created { .. }), + ]; + for predicate in predicates { + let removable = self + .demands + .iter() + .filter_map(|(demand_id, demand)| { + (matches!(demand.reason, WorkerDemandReason::Task) && predicate(&demand.state)) + .then_some(*demand_id) + }) + .collect::>(); + for demand_id in removable { + if surplus == 0 { + return; + } + self.demands.swap_remove(&demand_id); + surplus -= 1; + } + } + } +} + +#[cfg(test)] +mod tests { + #![expect(clippy::unwrap_used)] + + use std::time::Duration; + + use sail_common::utils::retry::RetryStrategy; + + use crate::driver::worker_scaler::{WorkerScaler, WorkerScalerOptions}; + use crate::id::WorkerId; + + fn worker_scaler(max_count: usize) -> WorkerScaler { + WorkerScaler::new(WorkerScalerOptions::new(RetryStrategy::Fixed { + max_count, + delay: Duration::ZERO, + })) + } + + #[test] + fn exhausted_demand_is_not_recreated_until_target_decreases() { + let mut scaler = worker_scaler(0); + let request = scaler.reconcile(1).unwrap().pop().unwrap(); + assert!(scaler.bind_worker(request, WorkerId::from(1))); + assert!(scaler.worker_failed(WorkerId::from(1)).is_none()); + + assert!(scaler.reconcile(1).unwrap().is_empty()); + assert!(scaler.reconcile(0).unwrap().is_empty()); + let next = scaler.reconcile(1).unwrap().pop().unwrap(); + assert_ne!(next.demand_id, request.demand_id); + } + + #[test] + fn retry_preserves_worker_demand_id() { + let mut scaler = worker_scaler(1); + let first = scaler.reconcile(1).unwrap().pop().unwrap(); + assert!(scaler.bind_worker(first, WorkerId::from(1))); + + let retry = scaler.worker_failed(WorkerId::from(1)).unwrap(); + assert_eq!(retry.demand_id, first.demand_id); + assert_eq!(retry.attempt, 1); + + let second = scaler.retry(retry).unwrap(); + assert_eq!(second.demand_id, first.demand_id); + assert!(scaler.bind_worker(second, WorkerId::from(2))); + assert!(scaler.worker_registered(WorkerId::from(2))); + } + + #[test] + fn initial_demand_counts_toward_task_target() { + let mut scaler = worker_scaler(0); + assert_eq!(scaler.request_initial_workers(2).unwrap().len(), 2); + assert!(scaler.reconcile(2).unwrap().is_empty()); + assert_eq!(scaler.reconcile(3).unwrap().len(), 1); + } +} diff --git a/crates/sail-execution/src/driver/worker_scaler/mod.rs b/crates/sail-execution/src/driver/worker_scaler/mod.rs new file mode 100644 index 0000000000..af52f733da --- /dev/null +++ b/crates/sail-execution/src/driver/worker_scaler/mod.rs @@ -0,0 +1,28 @@ +mod core; +mod options; +mod state; + +use indexmap::IndexMap; +pub use options::WorkerScalerOptions; +pub(crate) use state::{WorkerDemandReason, WorkerLaunchRequest, WorkerRetryRequest}; + +use crate::driver::worker_scaler::state::WorkerDemand; +use crate::id::{IdGenerator, WorkerDemandId, WorkerId}; + +pub struct WorkerScaler { + options: WorkerScalerOptions, + demands: IndexMap, + workers: IndexMap, + worker_demand_id_generator: IdGenerator, +} + +impl WorkerScaler { + pub fn new(options: WorkerScalerOptions) -> Self { + Self { + options, + demands: IndexMap::new(), + workers: IndexMap::new(), + worker_demand_id_generator: IdGenerator::new(), + } + } +} diff --git a/crates/sail-execution/src/driver/worker_scaler/options.rs b/crates/sail-execution/src/driver/worker_scaler/options.rs new file mode 100644 index 0000000000..921a693817 --- /dev/null +++ b/crates/sail-execution/src/driver/worker_scaler/options.rs @@ -0,0 +1,25 @@ +use sail_common::utils::retry::RetryStrategy; + +use crate::driver::DriverOptions; + +#[readonly::make] +pub struct WorkerScalerOptions { + pub worker_launch_retry_strategy: RetryStrategy, +} + +#[cfg(test)] +impl WorkerScalerOptions { + pub(super) fn new(worker_launch_retry_strategy: RetryStrategy) -> Self { + Self { + worker_launch_retry_strategy, + } + } +} + +impl From<&DriverOptions> for WorkerScalerOptions { + fn from(options: &DriverOptions) -> Self { + Self { + worker_launch_retry_strategy: options.worker_launch_retry_strategy.clone(), + } + } +} diff --git a/crates/sail-execution/src/driver/worker_scaler/state.rs b/crates/sail-execution/src/driver/worker_scaler/state.rs new file mode 100644 index 0000000000..932156b3a0 --- /dev/null +++ b/crates/sail-execution/src/driver/worker_scaler/state.rs @@ -0,0 +1,37 @@ +use std::time::Duration; + +use sail_common::utils::retry::RetrySchedule; + +use crate::id::{WorkerDemandId, WorkerId}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum WorkerDemandReason { + Initial, + Task, +} + +pub struct WorkerDemand { + pub reason: WorkerDemandReason, + pub state: WorkerDemandState, + pub retries: RetrySchedule, +} + +pub enum WorkerDemandState { + Created { attempt: usize }, + Launching { worker_id: WorkerId, attempt: usize }, + WaitingForRetry { attempt: usize }, + Exhausted, +} + +#[derive(Debug, Clone, Copy)] +pub struct WorkerLaunchRequest { + pub demand_id: WorkerDemandId, + pub attempt: usize, +} + +#[derive(Debug, Clone, Copy)] +pub struct WorkerRetryRequest { + pub demand_id: WorkerDemandId, + pub attempt: usize, + pub delay: Duration, +} diff --git a/crates/sail-execution/src/id.rs b/crates/sail-execution/src/id.rs index 7be7412c33..f5ecc175eb 100644 --- a/crates/sail-execution/src/id.rs +++ b/crates/sail-execution/src/id.rs @@ -66,6 +66,7 @@ macro_rules! define_id_type { define_id_type!(JobId, u64); define_id_type!(DriverId, u64); define_id_type!(WorkerId, u64); +define_id_type!(WorkerDemandId, u64); #[derive(Debug)] pub struct IdGenerator { diff --git a/crates/sail-session/src/session_factory/job_runner.rs b/crates/sail-session/src/session_factory/job_runner.rs index 9e4c9417ba..616b0855ba 100644 --- a/crates/sail-session/src/session_factory/job_runner.rs +++ b/crates/sail-session/src/session_factory/job_runner.rs @@ -1,6 +1,5 @@ use std::sync::Arc; -use datafusion::common::{Result, internal_err}; use sail_common::actor::ActorSystem; use sail_common::config::{AppConfig, ExecutionMode}; use sail_common::runtime::RuntimeHandle; @@ -13,6 +12,7 @@ use sail_execution::worker_manager::{ }; use sail_telemetry::events::SystemEventReporter; +use crate::error::{SessionError, SessionResult}; use crate::session_factory::{SessionFactory, WorkerSessionFactory}; pub struct SessionJobRunner { @@ -53,7 +53,7 @@ pub trait SessionJobRunnerFactory: Send { &mut self, system: &mut ActorSystem, info: SessionJobRunnerInfo, - ) -> Result; + ) -> SessionResult; } pub struct ServerSessionJobRunnerFactory { @@ -71,17 +71,17 @@ impl ServerSessionJobRunnerFactory { system: &mut ActorSystem, info: SessionJobRunnerInfo, worker_manager: Box, - ) -> Result { + ) -> SessionResult { let Some(port) = info.driver_server_port else { - return internal_err!("driver gateway is not available"); + return Err(SessionError::internal("driver gateway is not available")); }; - let options = DriverOptions::new( + let options = DriverOptions::try_new( &self.config, self.runtime.clone(), info.session_id, info.driver_id, port, - ); + )?; let components = DriverComponents { worker_manager, event_reporter: info.event_reporter, @@ -97,7 +97,7 @@ impl SessionJobRunnerFactory for ServerSessionJobRunnerFactory { &mut self, system: &mut ActorSystem, info: SessionJobRunnerInfo, - ) -> Result { + ) -> SessionResult { match self.config.mode { ExecutionMode::Local => Ok(SessionJobRunner::local(LocalJobRunner::new( info.session_id, diff --git a/crates/sail-session/src/session_manager/actor/core.rs b/crates/sail-session/src/session_manager/actor/core.rs index dbeb8e4c7b..f3fb136cba 100644 --- a/crates/sail-session/src/session_manager/actor/core.rs +++ b/crates/sail-session/src/session_manager/actor/core.rs @@ -94,9 +94,8 @@ impl Actor for SessionManagerActor { context, driver_id, activation, - result, } => self.handle_complete_session_creation( - ctx, session_id, user_id, context, driver_id, activation, result, + ctx, session_id, user_id, context, driver_id, activation, ), SessionManagerMessage::ProbeIdleSession { session_id, diff --git a/crates/sail-session/src/session_manager/actor/handler.rs b/crates/sail-session/src/session_manager/actor/handler.rs index 18b2c90463..2061e7c874 100644 --- a/crates/sail-session/src/session_manager/actor/handler.rs +++ b/crates/sail-session/src/session_manager/actor/handler.rs @@ -39,18 +39,23 @@ impl SessionManagerActor { user_id: String, result: oneshot::Sender>, ) -> ActorAction { - if let Some(session) = self.sessions.get(&session_id) { - let context = if let ServerSessionState::Running { context, .. } = &session.state { - Ok(context.clone()) + if let Some(session) = self.sessions.get_mut(&session_id) { + let context = match &mut session.state { + ServerSessionState::Running { context, .. } => Some(context.clone()), + ServerSessionState::Creating { waiters, .. } => { + waiters.push(result); + return ActorAction::Continue; + } + ServerSessionState::Deleted | ServerSessionState::Failed => None, + }; + if let Some(context) = context { + self.schedule_idle_session_probe(ctx, session_id, &context); + let _ = result.send(Ok(context)); } else { - Err(SessionError::invalid(format!( + let _ = result.send(Err(SessionError::invalid(format!( "session {session_id} is not running" - ))) - }; - if let Ok(context) = &context { - self.schedule_idle_session_probe(ctx, session_id, context); + )))); } - let _ = result.send(context); return ActorAction::Continue; } @@ -83,7 +88,7 @@ impl SessionManagerActor { ) { Ok(runner) => runner, Err(e) => { - let _ = result.send(Err(e.into())); + let _ = result.send(Err(e)); return ActorAction::Continue; } }; @@ -150,6 +155,7 @@ impl SessionManagerActor { ServerSession { state: ServerSessionState::Creating { driver_id: registered_driver_id, + waiters: vec![result], }, }, ); @@ -159,7 +165,6 @@ impl SessionManagerActor { context, driver_id: registered_driver_id, activation, - result, }; if let Some(driver) = driver { let session_manager = ctx.handle().clone(); @@ -183,32 +188,37 @@ impl SessionManagerActor { context: SessionContext, driver_id: Option, activation: ExecutionResult<()>, - result: oneshot::Sender>, ) -> ActorAction { let Some(session) = self.sessions.get_mut(&session_id) else { - let _ = result.send(Err(SessionError::internal(format!( - "session {session_id} disappeared during creation" - )))); + warn!("session {session_id} disappeared during creation"); return ActorAction::Continue; }; - if !matches!( - session.state, + let waiters = match &mut session.state { ServerSessionState::Creating { - driver_id: creating_driver_id - } if creating_driver_id == driver_id - ) { - let _ = result.send(Err(SessionError::internal(format!( - "session {session_id} creation is no longer pending" - )))); - return ActorAction::Continue; - } - let output = match activation { + driver_id: creating_driver_id, + waiters, + } if *creating_driver_id == driver_id => std::mem::take(waiters), + _ => { + warn!("session {session_id} creation is no longer pending"); + return ActorAction::Continue; + } + }; + match activation { Ok(()) => { session.state = ServerSessionState::Running { context: context.clone(), driver_id, }; - Ok(context) + self.event_reporter.report(SystemEvent::SessionCreated { + session_id: session_id.clone(), + user_id, + status: session.state.status().to_string(), + created_at: Utc::now(), + }); + self.schedule_idle_session_probe(ctx, session_id, &context); + for waiter in waiters { + let _ = waiter.send(Ok(context.clone())); + } } Err(e) => { if let Some(driver_id) = driver_id @@ -221,19 +231,18 @@ impl SessionManagerActor { }); } session.state = ServerSessionState::Failed; - Err(e.into()) + self.event_reporter.report(SystemEvent::SessionCreated { + session_id, + user_id, + status: session.state.status().to_string(), + created_at: Utc::now(), + }); + let message = e.to_string(); + for waiter in waiters { + let _ = waiter.send(Err(SessionError::internal(message.clone()))); + } } - }; - self.event_reporter.report(SystemEvent::SessionCreated { - session_id: session_id.clone(), - user_id, - status: session.state.status().to_string(), - created_at: Utc::now(), - }); - if let Ok(context) = &output { - self.schedule_idle_session_probe(ctx, session_id, context); } - let _ = result.send(output); ActorAction::Continue } @@ -330,10 +339,12 @@ impl SessionManagerActor { warn!("session not found: {session_id}"); return ActorAction::Continue; }; - let driver_id = match &session.state { - ServerSessionState::Creating { driver_id } - | ServerSessionState::Running { driver_id, .. } => *driver_id, - ServerSessionState::Deleted | ServerSessionState::Failed => None, + let (driver_id, waiters) = match &mut session.state { + ServerSessionState::Creating { driver_id, waiters } => { + (*driver_id, std::mem::take(waiters)) + } + ServerSessionState::Running { driver_id, .. } => (*driver_id, vec![]), + ServerSessionState::Deleted | ServerSessionState::Failed => (None, vec![]), }; if let Some(driver_id) = driver_id && let Some(driver) = self.drivers.remove(driver_id) @@ -351,6 +362,11 @@ impl SessionManagerActor { status, updated_at: Utc::now(), }); + for waiter in waiters { + let _ = waiter.send(Err(SessionError::internal( + "session failed during creation", + ))); + } ActorAction::Continue } diff --git a/crates/sail-session/src/session_manager/actor/message.rs b/crates/sail-session/src/session_manager/actor/message.rs index 002027956c..4b819d1a92 100644 --- a/crates/sail-session/src/session_manager/actor/message.rs +++ b/crates/sail-session/src/session_manager/actor/message.rs @@ -22,7 +22,6 @@ pub enum SessionManagerMessage { context: SessionContext, driver_id: Option, activation: ExecutionResult<()>, - result: oneshot::Sender>, }, ProbeIdleSession { session_id: String, @@ -73,7 +72,6 @@ impl SpanAssociation for SessionManagerMessage { context: _, driver_id: _, activation: _, - result: _, } | SessionManagerMessage::ProbeIdleSession { session_id, diff --git a/crates/sail-session/src/session_manager/session.rs b/crates/sail-session/src/session_manager/session.rs index cde697664c..26020cf56d 100644 --- a/crates/sail-session/src/session_manager/session.rs +++ b/crates/sail-session/src/session_manager/session.rs @@ -1,5 +1,8 @@ use datafusion::prelude::SessionContext; use sail_execution::DriverId; +use tokio::sync::oneshot; + +use crate::error::SessionResult; pub struct ServerSession { pub state: ServerSessionState, @@ -8,6 +11,7 @@ pub struct ServerSession { pub enum ServerSessionState { Creating { driver_id: Option, + waiters: Vec>>, }, Running { context: SessionContext,