diff --git a/crates/application/src/model_server.rs b/crates/application/src/model_server.rs index 679c9af..5903cff 100644 --- a/crates/application/src/model_server.rs +++ b/crates/application/src/model_server.rs @@ -153,6 +153,8 @@ pub struct ReadinessPolicy { pub attempts: usize, /// Delay between attempts. pub backoff: Duration, + /// Maximum time to wait for an auto-started process to expose its endpoint. + pub warmup_deadline: Duration, } impl Default for ReadinessPolicy { @@ -160,6 +162,7 @@ impl Default for ReadinessPolicy { Self { attempts: 20, backoff: Duration::from_millis(250), + warmup_deadline: Duration::from_secs(120), } } } @@ -437,68 +440,9 @@ impl EnsureLocalModelServer { handle: &ManagedProcessHandle, hf_source: Option, ) -> Result { - for attempt in 0..self.readiness.attempts { - match self.probe.probe(&config.endpoint).await { - Err(err) => { - self.stop_started_server(config.id, handle).await; - return self.fail(config.id, err); - } - Ok(ModelServerStatus::ReadyReused | ModelServerStatus::ReadyStarted) => { - self.publish( - config.id, - ModelServerLifecycleStatus::Ready { reused: false }, - ); - return Ok(EnsureLocalModelServerOutput { - ready: ready(config, ModelServerStatus::ReadyStarted), - }); - } - Ok(ModelServerStatus::Unreachable) => { - if attempt + 1 < self.readiness.attempts && !self.readiness.backoff.is_zero() { - tokio::time::sleep(self.readiness.backoff).await; - } - } - } - } - - let Some(source) = hf_source else { - let err = ModelServerError::Timeout; - self.stop_started_server(config.id, handle).await; - return self.fail(config.id, err); - }; - - match self.process.status(handle).await { - Ok(ProcessStatus::Running) => { - self.publish( - config.id, - ModelServerLifecycleStatus::Downloading { - downloaded_bytes: None, - total_bytes: None, - percent: None, - source: Some(source.clone()), - }, - ); - } - Ok(ProcessStatus::Exited { code }) => { - self.active.lock().unwrap().remove(&config.id); - return self.fail(config.id, premature_exit_error(code)); - } - Ok(ProcessStatus::Unknown) => { - self.active.lock().unwrap().remove(&config.id); - return self.fail( - config.id, - ModelServerError::Process("process status unknown".to_owned()), - ); - } - Err(err) => return self.fail(config.id, err), - } - - let deadline = Instant::now() + self.hf_download_deadline; + let mut attempts = 0usize; + let deadline = Instant::now() + self.readiness.warmup_deadline; loop { - if Instant::now() >= deadline { - let err = ModelServerError::Timeout; - self.stop_started_server(config.id, handle).await; - return self.fail(config.id, err); - } match self.probe.probe(&config.endpoint).await { Err(err) => { self.stop_started_server(config.id, handle).await; @@ -515,12 +459,21 @@ impl EnsureLocalModelServer { } Ok(ModelServerStatus::Unreachable) => {} } + match self.process.status(handle).await { Ok(ProcessStatus::Running) => { - if !self.readiness.backoff.is_zero() { - tokio::time::sleep(self.readiness.backoff).await; - } else { - tokio::task::yield_now().await; + if attempts.saturating_add(1) == self.readiness.attempts { + if let Some(source) = hf_source.as_ref() { + self.publish( + config.id, + ModelServerLifecycleStatus::Downloading { + downloaded_bytes: None, + total_bytes: None, + percent: None, + source: Some(source.clone()), + }, + ); + } } } Ok(ProcessStatus::Exited { code }) => { @@ -536,6 +489,18 @@ impl EnsureLocalModelServer { } Err(err) => return self.fail(config.id, err), } + + attempts = attempts.saturating_add(1); + if Instant::now() >= deadline { + let err = ModelServerError::Timeout; + self.stop_started_server(config.id, handle).await; + return self.fail(config.id, err); + } + if !self.readiness.backoff.is_zero() { + tokio::time::sleep(self.readiness.backoff).await; + } else { + tokio::task::yield_now().await; + } } } diff --git a/crates/application/tests/model_server.rs b/crates/application/tests/model_server.rs index 663f629..eada9ee 100644 --- a/crates/application/tests/model_server.rs +++ b/crates/application/tests/model_server.rs @@ -409,7 +409,8 @@ fn ensure( EnsureLocalModelServer::new(registry, probe, process, Arc::new(FakeRuntime), fs, events) .with_readiness_policy(ModelServerReadinessPolicy { attempts: 2, - backoff: Duration::ZERO, + backoff: Duration::from_millis(1), + warmup_deadline: Duration::from_millis(25), }) .with_hf_download_deadline(Duration::from_secs(5)) } @@ -675,6 +676,139 @@ async fn readiness_timeout_kills_started_process() { assert_eq!(process.kills.lock().unwrap().as_slice(), ["h1"]); } +#[tokio::test] +async fn slow_local_warmup_exceeding_short_readiness_window_succeeds_without_stop() { + let registry = Arc::new(FakeRegistry::default()); + registry + .save(config(sid(19), 8099, "/models/qwen.gguf", true)) + .await + .unwrap(); + let fs = Arc::new(FakeFs::default()); + fs.existing + .lock() + .unwrap() + .push("/models/qwen.gguf".to_owned()); + let process = Arc::new(FakeProcess::default()); + let mut statuses = vec![ModelServerStatus::Unreachable]; + statuses.extend(std::iter::repeat(ModelServerStatus::Unreachable).take(22)); + statuses.push(ModelServerStatus::ReadyStarted); + let registry_port: Arc = registry.clone(); + let process_port: Arc = process.clone(); + let usecase = EnsureLocalModelServer::new( + registry_port, + Arc::new(FakeProbe::new(statuses)), + process_port, + Arc::new(FakeRuntime), + fs, + Arc::new(FakeEvents::default()), + ) + .with_readiness_policy(ModelServerReadinessPolicy { + attempts: 20, + backoff: Duration::from_millis(1), + warmup_deadline: Duration::from_millis(100), + }); + + let out = usecase + .execute(EnsureLocalModelServerInput { server_id: sid(19) }) + .await + .unwrap(); + + assert_eq!(out.ready.status, ModelServerStatus::ReadyStarted); + assert_eq!(process.spawns.lock().unwrap().len(), 1); + assert!(process.kills.lock().unwrap().is_empty()); +} + +#[tokio::test] +async fn process_exit_during_local_warmup_fails_fast_and_cleans_active_registry() { + let registry = Arc::new(FakeRegistry::default()); + registry + .save(config(sid(20), 8100, "/models/a.gguf", true)) + .await + .unwrap(); + registry + .save(config(sid(21), 8100, "/models/b.gguf", true)) + .await + .unwrap(); + let fs = Arc::new(FakeFs::default()); + fs.existing + .lock() + .unwrap() + .extend(["/models/a.gguf".to_owned(), "/models/b.gguf".to_owned()]); + let process = Arc::new(FakeProcess::default()); + process + .status_sequence + .lock() + .unwrap() + .push_back(ProcessStatus::Exited { code: Some(42) }); + let usecase = ensure( + Arc::clone(®istry), + Arc::new(FakeProbe::new(vec![ + ModelServerStatus::Unreachable, + ModelServerStatus::Unreachable, + ModelServerStatus::Unreachable, + ModelServerStatus::ReadyStarted, + ])), + Arc::clone(&process), + fs, + Arc::new(FakeEvents::default()), + ); + + let err = usecase + .execute(EnsureLocalModelServerInput { server_id: sid(20) }) + .await + .unwrap_err(); + + assert!(err.to_string().contains("process")); + assert!(err.to_string().contains("42")); + assert!(process.kills.lock().unwrap().is_empty()); + + let out = usecase + .execute(EnsureLocalModelServerInput { server_id: sid(21) }) + .await + .unwrap(); + assert_eq!(out.ready.status, ModelServerStatus::ReadyStarted); + assert_eq!(process.spawns.lock().unwrap().len(), 2); +} + +#[tokio::test] +async fn warmup_deadline_reached_returns_timeout_and_stops_started_process() { + let registry = Arc::new(FakeRegistry::default()); + registry + .save(config(sid(22), 8101, "/models/qwen.gguf", true)) + .await + .unwrap(); + let fs = Arc::new(FakeFs::default()); + fs.existing + .lock() + .unwrap() + .push("/models/qwen.gguf".to_owned()); + let process = Arc::new(FakeProcess::default()); + let registry_port: Arc = registry.clone(); + let process_port: Arc = process.clone(); + let usecase = EnsureLocalModelServer::new( + registry_port, + Arc::new(FakeProbe::new(vec![ModelServerStatus::Unreachable])), + process_port, + Arc::new(FakeRuntime), + fs, + Arc::new(FakeEvents::default()), + ) + .with_readiness_policy(ModelServerReadinessPolicy { + attempts: 20, + backoff: Duration::from_millis(1), + warmup_deadline: Duration::from_millis(3), + }); + + let err = usecase + .execute(EnsureLocalModelServerInput { server_id: sid(22) }) + .await + .unwrap_err(); + + assert!(err.to_string().contains("timeout")); + assert_eq!(process.spawns.lock().unwrap().len(), 1); + assert_eq!(process.kills.lock().unwrap().as_slice(), ["h1"]); +} + #[tokio::test] async fn hf_unreachable_alive_publishes_downloading_without_short_timeout_kill() { let registry = Arc::new(FakeRegistry::default());