diff --git a/crates/mvp-system/src/node_image.rs b/crates/mvp-system/src/node_image.rs index 379fb17..7f646b3 100644 --- a/crates/mvp-system/src/node_image.rs +++ b/crates/mvp-system/src/node_image.rs @@ -7,25 +7,8 @@ use std::sync::mpsc; use std::thread; use std::time::{Duration, Instant}; -const IMAGE_SOURCE_INPUTS: &[&str] = &[ - "Cargo.lock", - "Cargo.toml", - "src", - "crates/datastream/Cargo.toml", - "crates/datastream/src", - "crates/distribution/Cargo.toml", - "crates/distribution/src", - "crates/iroh-driver/Cargo.toml", - "crates/iroh-driver/src", - "crates/mvp-system/Cargo.toml", - "crates/mvp-system/src", - "crates/transport/Cargo.toml", - "crates/transport/src", - "tools/vastai/Cargo.toml", - "tools/vastai/src", +const NODE_IMAGE_CONTENT_INPUTS: &[&str] = &[ "apps/mvp-node/Dockerfile", - "apps/mvp-node/Dockerfile.base", - "apps/mvp-node/mvp_entrypoint.sh", "apps/mvp-node/tinygrad_worker.py", ]; @@ -192,13 +175,31 @@ fn prepare_node_image_inner( )); } - let tag = image_version_tag(&root)?; + run_status( + runner, + progress, + &root, + "cargo", + &[ + "build", + "--quiet", + "-p", + "mvp-system", + "--bin", + "mvp-worker-node", + ], + "build mvp-worker-node", + None, + )?; + + let base_hash = content_hash_for_inputs(&root, BASE_IMAGE_SOURCE_INPUTS)?; + let image_content_hash = node_image_content_hash(&root, &request.node_bin, &base_hash)?; + let tag = image_version_tag(&root, &image_content_hash)?; let image_ref = image.ref_for_tag(&tag); emit_image_reference(progress, "resolved", &image_ref); - let source_hash = source_content_hash(&root)?; let worker_hash = file_content_hash(&root, Path::new("apps/mvp-node/tinygrad_worker.py"))?; - let base_hash = content_hash_for_inputs(&root, BASE_IMAGE_SOURCE_INPUTS)?; - let expected_node_labels = node_image_labels(&tag, &source_hash, &worker_hash, &base_hash); + let expected_node_labels = + node_image_labels(&tag, &image_content_hash, &worker_hash, &base_hash); let expected_base_labels = base_image_labels(&base_hash); let alias_tags = alias_tags(&image, request.extra_tag.as_deref(), &tag)?; for alias in alias_refs(&image, &alias_tags) { @@ -247,23 +248,6 @@ fn prepare_node_image_inner( pushed: false, }); } - - run_status( - runner, - progress, - &root, - "cargo", - &[ - "build", - "--quiet", - "-p", - "mvp-system", - "--bin", - "mvp-worker-node", - ], - "build mvp-worker-node", - None, - )?; let base_image_matches = docker_image_labels_match(runner, &root, &request.base_image, &expected_base_labels)?; if !base_image_matches { @@ -350,12 +334,12 @@ fn workspace_root() -> Result { )) } -fn image_version_tag(root: &Path) -> Result { +fn image_version_tag(root: &Path, image_content_hash: &str) -> Result { if git_worktree_clean(root)? { let sha = git_capture(root, &["rev-parse", "--short=12", "HEAD"])?; Ok(format!("git-{}", sha.trim())) } else { - Ok(format!("dirty-{}", dirty_content_hash(root)?)) + Ok(format!("dirty-{image_content_hash}")) } } @@ -384,12 +368,25 @@ fn git_capture(root: &Path, args: &[&str]) -> Result { } } -fn dirty_content_hash(root: &Path) -> Result { - source_content_hash(root) -} - -fn source_content_hash(root: &Path) -> Result { - content_hash_for_inputs(root, IMAGE_SOURCE_INPUTS) +fn node_image_content_hash( + root: &Path, + node_bin: &Path, + base_hash: &str, +) -> Result { + let mut files = Vec::new(); + for input in NODE_IMAGE_CONTENT_INPUTS { + let path = root.join(input); + collect_hash_inputs(root, &path, &mut files)?; + } + let node_bin = if node_bin.is_absolute() { + node_bin.to_path_buf() + } else { + root.join(node_bin) + }; + files.push(relative_path(root, &node_bin)?); + files.sort(); + files.dedup(); + hash_relative_files_with_salts(root, files, &[("base", base_hash)]) } fn content_hash_for_inputs(root: &Path, inputs: &[&str]) -> Result { @@ -408,7 +405,21 @@ fn file_content_hash(root: &Path, path: &Path) -> Result { } fn hash_relative_files(root: &Path, files: Vec) -> Result { + hash_relative_files_with_salts(root, files, &[]) +} + +fn hash_relative_files_with_salts( + root: &Path, + files: Vec, + salts: &[(&str, &str)], +) -> Result { let mut hasher = blake3::Hasher::new(); + for (key, value) in salts { + hasher.update(key.as_bytes()); + hasher.update(b"\0"); + hasher.update(value.as_bytes()); + hasher.update(b"\0"); + } for relative in files { let full = root.join(&relative); hasher.update(relative.to_string_lossy().as_bytes()); @@ -1348,7 +1359,7 @@ mod tests { #[test] fn prepare_node_image_with_dry_runner_returns_expected_image_and_progress() { - let root = workspace_root().expect("workspace root resolves"); + let node_bin = std::env::current_exe().expect("test binary path resolves"); let mut runner = DryImageCommandRunner::default(); let mut progress = CollectProgress::default(); let mut sink: Option<&mut dyn NodeImageProgressSink> = Some(&mut progress); @@ -1357,7 +1368,7 @@ mod tests { NodeImageRequest { requested_image: "docker.io/acme/mvp-node:latest".to_owned(), base_image: "swactor-mvp-node-base:cuda12.6".to_owned(), - node_bin: root.join("target/debug/mvp-worker-node"), + node_bin, provider: NodeImageProvider::Docker, extra_tag: Some("smoke".to_owned()), push: false, @@ -1391,6 +1402,35 @@ mod tests { ))); } + #[test] + fn node_image_content_hash_tracks_node_payload_not_unrelated_files() { + let root = workspace_root().expect("workspace root resolves"); + let scratch = root + .join("target/node-image-hash-test") + .join(std::process::id().to_string()); + fs::create_dir_all(&scratch).expect("scratch dir is writable"); + let node_bin = scratch.join("mvp-worker-node"); + fs::write(&node_bin, b"worker binary v1").expect("node bin fixture is writable"); + + let initial = + node_image_content_hash(&root, &node_bin, "base-v1").expect("initial hash succeeds"); + fs::write(scratch.join("unrelated.txt"), b"not part of the image") + .expect("unrelated fixture is writable"); + let after_unrelated = + node_image_content_hash(&root, &node_bin, "base-v1").expect("unrelated hash succeeds"); + fs::write(&node_bin, b"worker binary v2").expect("node bin fixture update is writable"); + let after_node_bin = + node_image_content_hash(&root, &node_bin, "base-v1").expect("node bin hash succeeds"); + let after_base = + node_image_content_hash(&root, &node_bin, "base-v2").expect("base hash succeeds"); + + let _ = fs::remove_dir_all(&scratch); + + assert_eq!(initial, after_unrelated); + assert_ne!(initial, after_node_bin); + assert_ne!(after_node_bin, after_base); + } + #[test] fn image_name_splits_tag_after_last_slash() { let image = ImageName::parse("localhost:5000/team/mvp-node:trial").unwrap(); diff --git a/crates/mvp-system/src/tests/vastai_provisioning_guarantees.rs b/crates/mvp-system/src/tests/vastai_provisioning_guarantees.rs index 38b6818..f328123 100644 --- a/crates/mvp-system/src/tests/vastai_provisioning_guarantees.rs +++ b/crates/mvp-system/src/tests/vastai_provisioning_guarantees.rs @@ -115,6 +115,7 @@ impl VastAiBootstrapLauncher for FakeBootstrap { endpoint: VastAiSshEndpoint, _sink: PluginSink, _producer: Option, + _lifecycle: LifecyclePolicy, ) -> Result { if let Some(reason) = self.fail.clone() { return Err(reason); diff --git a/crates/mvp-system/src/vastai_provisioning.rs b/crates/mvp-system/src/vastai_provisioning.rs index c0481cd..2d46f09 100644 --- a/crates/mvp-system/src/vastai_provisioning.rs +++ b/crates/mvp-system/src/vastai_provisioning.rs @@ -1,5 +1,6 @@ use parking_lot::Mutex; -use std::collections::{BTreeMap, BTreeSet}; +use std::collections::{BTreeMap, BTreeSet, HashSet}; +use std::io::{BufRead, BufReader, Read}; use std::path::{Path, PathBuf}; use std::process::{Child, Command, Stdio}; use std::sync::{ @@ -7,14 +8,16 @@ use std::sync::{ atomic::{AtomicBool, Ordering}, mpsc, }; -use std::time::Duration; +use std::thread::JoinHandle; +use std::time::{Duration, Instant}; use datastream::DatastreamProducer; use serde::{Deserialize, Serialize}; use swactor::actor::{ActorAddress, ActorInterface}; use swactor::runtime::{Ctx, Runtime}; use swactor_vastai::{ - LifecyclePolicy, ProvisionRequest, ProvisionedInstance, SelectionPolicy, classify_vastai_error, + CreateInstanceRequest, LifecyclePolicy, Offer, ProvisionRequest, ProvisionedInstance, + SelectionPolicy, classify_vastai_error, create_instance, }; use crate::bootstrap_datastream::{BootstrapDatastreamBridge, node_stream_id}; @@ -60,6 +63,33 @@ pub struct VastAiSshEndpoint { pub user: String, } +pub struct VastAiProviderMonitor { + stopping: Arc, + join: Option>, +} + +impl VastAiProviderMonitor { + fn new(stopping: Arc, join: JoinHandle<()>) -> Self { + Self { + stopping, + join: Some(join), + } + } + + fn stop(&mut self) { + self.stopping.store(true, Ordering::SeqCst); + if let Some(join) = self.join.take() { + let _ = join.join(); + } + } +} + +impl Drop for VastAiProviderMonitor { + fn drop(&mut self) { + self.stop(); + } +} + pub trait VastAiLeaseClient: Send { fn provision_one(&mut self, request: ProvisionRequest) -> Result; fn plan_first_wave_offers( @@ -76,6 +106,16 @@ pub trait VastAiLeaseClient: Send { lifecycle: &LifecyclePolicy, ssh_user: &str, ) -> Result; + fn spawn_provider_monitor( + &mut self, + _contract_id: u64, + _label: String, + _lifecycle: LifecyclePolicy, + _spec: NodeProvisionSpec, + _sink: PluginSink, + ) -> Option { + None + } fn destroy_contract(&mut self, contract_id: u64) -> Result<(), String>; } @@ -83,6 +123,8 @@ pub trait VastAiLeaseClient: Send { pub struct ToolsVastAiLeaseClient { client: swactor_vastai::VastClient, runtime: tokio::runtime::Runtime, + planned_offer_pool: Arc>>, + planned_offer_ids: Arc>>, } impl ToolsVastAiLeaseClient { @@ -91,7 +133,12 @@ impl ToolsVastAiLeaseClient { .enable_all() .build() .map_err(|e| format!("vastai tokio runtime: {e}"))?; - Ok(Self { client, runtime }) + Ok(Self { + client, + runtime, + planned_offer_pool: Arc::new(Mutex::new(Vec::new())), + planned_offer_ids: Arc::new(Mutex::new(HashSet::new())), + }) } pub fn from_api_key(api_key: impl Into) -> Result { @@ -101,25 +148,269 @@ impl ToolsVastAiLeaseClient { pub fn client(&self) -> &swactor_vastai::VastClient { &self.client } + + fn create_request_for_offer( + request: &ProvisionRequest, + offer_id: u64, + ) -> CreateInstanceRequest { + let mut env = request.env.clone(); + if let Some(overlay) = request.per_instance_env.first() { + env.extend( + overlay + .iter() + .map(|(key, value)| (key.clone(), value.clone())), + ); + } + CreateInstanceRequest { + offer_id, + image: request.image.clone(), + disk_gb: request.disk_gb, + label: request.label.clone(), + env, + onstart: request.onstart.clone(), + } + } + + fn candidate_pool(&mut self, request: &ProvisionRequest) -> Result, String> { + let cached = self.planned_offer_pool.lock().clone(); + if request + .preferred_offer_id + .is_some_and(|offer_id| cached.iter().any(|offer| offer.id == offer_id)) + { + return Ok(cached); + } + self.runtime + .block_on(self.client.search_offers(&request.selection, 1)) + } + + fn create_from_offer( + &mut self, + request: &ProvisionRequest, + offer: &Offer, + ) -> Result { + let create = Self::create_request_for_offer(request, offer.id); + let info = self.runtime.block_on(create_instance( + self.client.http(), + self.client.base_url(), + self.client.api_key(), + &create, + ))?; + Ok(ProvisionedInstance { + index: 0, + contract_id: info.contract_id, + offer_id: offer.id, + host_id: offer.host_id, + gpu_name: offer.gpu_name.clone(), + gpu_ram: offer.gpu_ram, + dph_total: offer.dph_total, + }) + } + + fn monitor_provider_status( + &mut self, + contract_id: u64, + label: String, + lifecycle: LifecyclePolicy, + spec: NodeProvisionSpec, + sink: PluginSink, + stopping: Arc, + ) { + let mut last_state: Option = None; + let mut state_since = Instant::now(); + let mut poll = 0_u64; + while !stopping.load(Ordering::SeqCst) { + poll = poll.saturating_add(1); + let status = match self + .runtime + .block_on(self.client.instance_status(contract_id)) + { + Ok(status) => status, + Err(error) + if error.contains("not found while fetching provider status") + || error.contains("parse failed") => + { + let reason = classified_start_error(format!( + "vastai provider monitor node {} contract {contract_id}: {error}", + spec.node_id + )); + sink.observe(PluginObservation::ProviderLine { + run_id: spec.run_id, + node_id: spec.node_id, + line: serde_json::json!({ + "type": "VastAiProviderStatusFailure", + "run_id": spec.run_id, + "node_id": spec.node_id, + "label": &label, + "contract_id": contract_id, + "poll": poll, + "reason": &reason, + }) + .to_string(), + }); + sink.observe(PluginObservation::Failed { + run_id: spec.run_id, + node_id: spec.node_id, + reason, + }); + return; + } + Err(error) => { + sink.observe(PluginObservation::ProviderLine { + run_id: spec.run_id, + node_id: spec.node_id, + line: serde_json::json!({ + "type": "VastAiProviderStatusPollRetry", + "run_id": spec.run_id, + "node_id": spec.node_id, + "label": &label, + "contract_id": contract_id, + "poll": poll, + "reason": error, + }) + .to_string(), + }); + if !sleep_provider_monitor(lifecycle.poll_interval, &stopping) { + return; + } + continue; + } + }; + + let actual = status.actual_status.as_str(); + if last_state.as_deref() != Some(actual) { + state_since = Instant::now(); + last_state = Some(actual.to_owned()); + } + let in_state_ms = state_since.elapsed().as_millis(); + sink.observe(PluginObservation::ProviderLine { + run_id: spec.run_id, + node_id: spec.node_id, + line: serde_json::json!({ + "type": "VastAiProviderStatusObserved", + "run_id": spec.run_id, + "node_id": spec.node_id, + "label": &label, + "contract_id": contract_id, + "poll": poll, + "actual_status": &status.actual_status, + "intended_status": &status.intended_status, + "status_msg": &status.status_msg, + "disk_usage": status.disk_usage, + "in_state_ms": in_state_ms, + }) + .to_string(), + }); + + if let Some(error) = provider_terminal_start_error( + contract_id, + &status.actual_status, + &status.intended_status, + status.status_msg.as_deref(), + ) { + let reason = classified_start_error(format!( + "vastai provider monitor node {}: {error}", + spec.node_id + )); + sink.observe(PluginObservation::ProviderLine { + run_id: spec.run_id, + node_id: spec.node_id, + line: serde_json::json!({ + "type": "VastAiProviderTerminalBeforeRuntimeReady", + "run_id": spec.run_id, + "node_id": spec.node_id, + "label": &label, + "contract_id": contract_id, + "reason": &reason, + }) + .to_string(), + }); + sink.observe(PluginObservation::Failed { + run_id: spec.run_id, + node_id: spec.node_id, + reason, + }); + return; + } + + if !sleep_provider_monitor(lifecycle.poll_interval, &stopping) { + return; + } + } + } } impl Clone for ToolsVastAiLeaseClient { fn clone(&self) -> Self { - Self::new(self.client.clone()).expect("clone VastAI lease client runtime") + Self { + client: self.client.clone(), + runtime: tokio::runtime::Builder::new_current_thread() + .enable_all() + .build() + .expect("clone VastAI lease client runtime"), + planned_offer_pool: Arc::clone(&self.planned_offer_pool), + planned_offer_ids: Arc::clone(&self.planned_offer_ids), + } } } impl VastAiLeaseClient for ToolsVastAiLeaseClient { fn provision_one(&mut self, request: ProvisionRequest) -> Result { - let fleet = self.runtime.block_on(self.client.provision(request))?; - let mut instances = fleet.instances; - if instances.len() != 1 { + if request.count != 1 { return Err(format!( - "vastai provision expected one instance, got {}", - instances.len() + "vastai provision_one expected count=1, got {}", + request.count )); } - Ok(instances.remove(0)) + + let pool = self.candidate_pool(&request)?; + let planned_offer_ids = self.planned_offer_ids.lock().clone(); + let blocked_hosts = request + .selection + .blacklist_hosts + .iter() + .copied() + .collect::>(); + let mut failed_hosts = HashSet::new(); + let mut tried_offer_ids = HashSet::new(); + let mut ordered = Vec::with_capacity(pool.len()); + if let Some(preferred_offer_id) = request.preferred_offer_id + && let Some(offer) = pool.iter().find(|offer| offer.id == preferred_offer_id) + { + ordered.push(offer.clone()); + } + ordered.extend(pool.into_iter()); + + let mut last_error = None; + for offer in ordered { + if !tried_offer_ids.insert(offer.id) { + continue; + } + if request.preferred_offer_id != Some(offer.id) && planned_offer_ids.contains(&offer.id) + { + continue; + } + if offer.host_id.is_some_and(|host_id| { + blocked_hosts.contains(&host_id) || failed_hosts.contains(&host_id) + }) { + continue; + } + match self.create_from_offer(&request, &offer) { + Ok(instance) => return Ok(instance), + Err(error) => { + if let Some(host_id) = offer.host_id { + failed_hosts.insert(host_id); + } + last_error = Some(format!("offer {}: {error}", offer.id)); + } + } + } + + Err(format!( + "vastai provision node exhausted eligible offers{}", + last_error + .map(|error| format!(" after create failure ({error})")) + .unwrap_or_default() + )) } fn plan_first_wave_offers( @@ -127,6 +418,8 @@ impl VastAiLeaseClient for ToolsVastAiLeaseClient { requests: &[ProvisionRequest], ) -> Result>, String> { let Some(first) = requests.first() else { + self.planned_offer_pool.lock().clear(); + self.planned_offer_ids.lock().clear(); return Ok(Vec::new()); }; let pool = self.runtime.block_on( @@ -139,6 +432,9 @@ impl VastAiLeaseClient for ToolsVastAiLeaseClient { &first.selection.blacklist_hosts, &[], ); + let planned_ids = planned.iter().map(|offer| offer.id).collect::>(); + *self.planned_offer_pool.lock() = pool; + *self.planned_offer_ids.lock() = planned_ids; let mut out = planned .into_iter() .map(|offer| Some(offer.id)) @@ -154,23 +450,36 @@ impl VastAiLeaseClient for ToolsVastAiLeaseClient { lifecycle: &LifecyclePolicy, ssh_user: &str, ) -> Result { - self.runtime.block_on(async { - let instances = self.client.list_by_label(label).await?; - if let Some(instance) = instances - .into_iter() - .find(|instance| instance.contract_id == contract_id) - { - let host = if instance.ssh_host.is_empty() { - instance.public_ipaddr - } else { - instance.ssh_host - }; - return endpoint_from_parts(contract_id, host, instance.ssh_port, ssh_user); - } + let endpoint = self.runtime.block_on(self.client.wait_for_ssh_endpoint( + contract_id, + label, + lifecycle, + ))?; + endpoint_from_parts(contract_id, endpoint.ip, endpoint.port, ssh_user) + } - let running = self.client.wait_for_running(contract_id, lifecycle).await?; - endpoint_from_parts(contract_id, running.ip, running.port, ssh_user) - }) + fn spawn_provider_monitor( + &mut self, + contract_id: u64, + label: String, + lifecycle: LifecyclePolicy, + spec: NodeProvisionSpec, + sink: PluginSink, + ) -> Option { + let stopping = Arc::new(AtomicBool::new(false)); + let thread_stopping = Arc::clone(&stopping); + let mut client = self.clone(); + let join = std::thread::spawn(move || { + client.monitor_provider_status( + contract_id, + label, + lifecycle, + spec, + sink, + thread_stopping, + ); + }); + Some(VastAiProviderMonitor::new(stopping, join)) } fn destroy_contract(&mut self, contract_id: u64) -> Result<(), String> { @@ -203,6 +512,44 @@ fn endpoint_from_parts( }) } +fn provider_terminal_start_error( + contract_id: u64, + actual: &str, + intended: &str, + msg: Option<&str>, +) -> Option { + if let Some(message) = msg { + let lower = message.to_ascii_lowercase(); + if lower.contains("error") || lower.contains("failed") { + return Some(format!("instance {contract_id} error: {message}")); + } + } + if intended == "stopped" && actual != "running" { + return Some(format!( + "instance {contract_id} stopped: {}", + msg.unwrap_or_default() + )); + } + match actual { + "exited" | "error" | "stopped" => Some(format!( + "instance {contract_id} reached terminal status: {actual}" + )), + _ => None, + } +} + +fn sleep_provider_monitor(duration: Duration, stopping: &AtomicBool) -> bool { + let deadline = Instant::now() + duration; + while Instant::now() < deadline { + if stopping.load(Ordering::SeqCst) { + return false; + } + let remaining = deadline.saturating_duration_since(Instant::now()); + std::thread::sleep(std::cmp::min(remaining, Duration::from_millis(100))); + } + !stopping.load(Ordering::SeqCst) +} + #[derive(Clone, Debug)] pub struct VastAiProviderPlugin where @@ -421,6 +768,7 @@ pub trait VastAiBootstrapLauncher: Send { endpoint: VastAiSshEndpoint, sink: PluginSink, producer: Option, + lifecycle: LifecyclePolicy, ) -> Result; fn stop_bootstrap(&mut self, handle: &mut Self::Handle, reason: BootstrapStopReason); @@ -492,6 +840,7 @@ impl VastAiBootstrapLauncher for SshCommandBootstrapLauncher { endpoint: VastAiSshEndpoint, sink: PluginSink, producer: Option, + lifecycle: LifecyclePolicy, ) -> Result { if spec.args.is_empty() { return Err(format!( @@ -514,6 +863,7 @@ impl VastAiBootstrapLauncher for SshCommandBootstrapLauncher { self.ssh_identity.clone(), child, stopping, + lifecycle.state_timeout, ); Ok(SshCommandBootstrapHandle { @@ -528,6 +878,70 @@ impl VastAiBootstrapLauncher for SshCommandBootstrapLauncher { } } +const POST_GRACE_BOOTSTRAP_FAILURE_LIMIT: u32 = 2; + +fn classify_ssh_observation(line: &str) -> Option<&'static str> { + let lower = line.to_ascii_lowercase(); + if lower.contains("permission denied (publickey") + || lower.contains("publickey denied") + || lower.contains("public key denied") + || lower.contains("no supported authentication methods") + { + return Some("auth_denied"); + } + if lower.contains("connection refused") + || lower.contains("connect to host") && lower.contains("refused") + { + return Some("refused"); + } + if lower.contains("operation timed out") + || lower.contains("connection timed out") + || lower.contains("connect timed out") + { + return Some("timeout"); + } + None +} + +fn post_grace_terminal_bootstrap_class(class: &str) -> bool { + matches!(class, "auth_denied" | "refused" | "timeout") +} + +fn spawn_classifying_stderr_reader( + stderr: R, + bridge: BootstrapDatastreamBridge, + observed_class: Arc>>, +) -> JoinHandle<()> +where + R: Read + Send + 'static, +{ + std::thread::spawn(move || { + let reader = BufReader::new(stderr); + for next in reader.lines() { + match next { + Ok(line) => { + if let Some(class) = classify_ssh_observation(&line) { + *observed_class.lock() = Some(class); + bridge.observe_provider_line( + serde_json::json!({ + "type": "VastAiBootstrapObservationClass", + "run_id": bridge.spec().run_id, + "node_id": bridge.spec().node_id, + "class": class, + }) + .to_string(), + ); + } + bridge.observe_stderr_line(line); + } + Err(error) => { + bridge.observe_provider_line(format!("read VastAI SSH stderr: {error}")); + break; + } + } + } + }) +} fn spawn_retrying_ssh_bootstrap( spec: NodeProvisionSpec, endpoint: VastAiSshEndpoint, @@ -536,12 +950,15 @@ fn spawn_retrying_ssh_bootstrap( ssh_identity: Option, child_slot: Arc>>, stopping: Arc, + post_grace_failure_after: Duration, ) { std::thread::spawn(move || { let run_id = spec.run_id; let node_id = spec.node_id; let mut attempt = 1u64; let mut backoff = Duration::from_secs(1); + let bootstrap_started = Instant::now(); + let mut post_grace_failures = 0_u32; while !stopping.load(Ordering::SeqCst) { sink.observe(PluginObservation::ProviderLine { @@ -562,7 +979,12 @@ fn spawn_retrying_ssh_bootstrap( producer.clone(), ); bridge.spawn_stdout_reader(stdout); - bridge.spawn_stderr_reader(stderr); + let observed_class = Arc::new(Mutex::new(None)); + let mut stderr_reader = Some(spawn_classifying_stderr_reader( + stderr, + bridge.clone(), + Arc::clone(&observed_class), + )); loop { if stopping.load(Ordering::SeqCst) { @@ -597,17 +1019,50 @@ fn spawn_retrying_ssh_bootstrap( } else { "not ready before runtime ready" }; - let line = format!( - "VastAI SSH bootstrap {readiness} (attempt {attempt}, {status}); retrying" - ); + if let Some(reader) = stderr_reader.take() { + let _ = reader.join(); + } + let observation_class = + (*observed_class.lock()).unwrap_or("process_exit"); sink.observe(PluginObservation::ProviderLine { run_id, node_id, - line, + line: serde_json::json!({ + "type": "VastAiBootstrapAttemptCompleted", + "run_id": run_id, + "node_id": node_id, + "attempt": attempt, + "status": status.to_string(), + "class": observation_class, + "classification": readiness, + }) + .to_string(), }); + if !status.success() + && !post_grace_failure_after.is_zero() + && bootstrap_started.elapsed() >= post_grace_failure_after + && post_grace_terminal_bootstrap_class(observation_class) + { + post_grace_failures = post_grace_failures.saturating_add(1); + if post_grace_failures >= POST_GRACE_BOOTSTRAP_FAILURE_LIMIT { + sink.observe(PluginObservation::Failed { + run_id, + node_id, + reason: format!( + "VastAI SSH bootstrap repeated post-grace {observation_class} failure before runtime ready" + ), + }); + return; + } + } else { + post_grace_failures = 0; + } break; } Some(Err(error)) => { + if let Some(reader) = stderr_reader.take() { + let _ = reader.join(); + } sink.observe(PluginObservation::ProviderLine { run_id, node_id, @@ -740,7 +1195,12 @@ where struct VastAiNode { contract_id: u64, bootstrap: Option, + provider_monitor: Option, host_id: Option, + run_id: u64, + node_id: u64, + label: String, + sink: PluginSink, } impl VastAiProvisioningPlugin @@ -908,6 +1368,18 @@ where }) .to_string(), }); + sink.observe(PluginObservation::ProviderLine { + run_id: spec.run_id, + node_id: spec.node_id, + line: serde_json::json!({ + "type": "VastAiSshEndpointDiscoveryStarted", + "run_id": spec.run_id, + "node_id": spec.node_id, + "contract_id": instance.contract_id, + "label": &label, + }) + .to_string(), + }); let endpoint = match self.client.ssh_endpoint( instance.contract_id, @@ -944,11 +1416,27 @@ where .to_string(), }); + sink.observe(PluginObservation::ProviderLine { + run_id: spec.run_id, + node_id: spec.node_id, + line: serde_json::json!({ + "type": "VastAiBootstrapObservationStarted", + "run_id": spec.run_id, + "node_id": spec.node_id, + "contract_id": instance.contract_id, + "host": &endpoint.host, + "port": endpoint.port, + "user": &endpoint.user, + }) + .to_string(), + }); + let bootstrap = match self.bootstrap.start_bootstrap( spec.clone(), endpoint, - sink, + sink.clone(), self.bootstrap_producer.clone(), + self.config.lifecycle.clone(), ) { Ok(handle) => handle, Err(error) => { @@ -980,7 +1468,18 @@ where VastAiNode { contract_id: instance.contract_id, bootstrap: Some(bootstrap), + provider_monitor: self.client.spawn_provider_monitor( + instance.contract_id, + label.clone(), + self.config.lifecycle.clone(), + spec.clone(), + sink.clone(), + ), host_id, + run_id: spec.run_id, + node_id: spec.node_id, + label, + sink, }, ); Ok(handle) @@ -1059,13 +1558,56 @@ where }) .to_string(), }); + } else { + sink.observe(PluginObservation::ProviderLine { + run_id: spec.run_id, + node_id: spec.node_id, + line: serde_json::json!({ + "type": "VastAiFirstWaveOfferPlanUnavailable", + "run_id": spec.run_id, + "node_id": spec.node_id, + "label": &label, + }) + .to_string(), + }); } let mut client = self.client.clone(); let config = self.config.clone(); let worker_tx = completion_tx.clone(); + let worker_sink = sink.clone(); std::thread::spawn(move || { let started = match client.provision_one(request) { Ok(instance) => { + worker_sink.observe(PluginObservation::ProviderLine { + run_id: spec.run_id, + node_id: spec.node_id, + line: serde_json::json!({ + "type": "VastAiLeaseReady", + "run_id": spec.run_id, + "node_id": spec.node_id, + "label": &label, + "image": &spec.image, + "contract_id": instance.contract_id, + "offer_id": instance.offer_id, + "host_id": instance.host_id, + "gpu_name": &instance.gpu_name, + "gpu_ram": instance.gpu_ram, + "dph_total": instance.dph_total, + }) + .to_string(), + }); + worker_sink.observe(PluginObservation::ProviderLine { + run_id: spec.run_id, + node_id: spec.node_id, + line: serde_json::json!({ + "type": "VastAiSshEndpointDiscoveryStarted", + "run_id": spec.run_id, + "node_id": spec.node_id, + "contract_id": instance.contract_id, + "label": &label, + }) + .to_string(), + }); match client.ssh_endpoint( instance.contract_id, &label, @@ -1112,29 +1654,17 @@ where for (index, spec, started) in completion_rx { match started { Ok(started) => { - sink.observe(PluginObservation::ProviderLine { - run_id: spec.run_id, - node_id: spec.node_id, - line: format!( - "vastai contract {} ready for SSH lookup", - started.instance.contract_id - ), - }); sink.observe(PluginObservation::ProviderLine { run_id: spec.run_id, node_id: spec.node_id, line: serde_json::json!({ - "type": "VastAiLeaseReady", + "type": "VastAiSshEndpointReady", "run_id": spec.run_id, "node_id": spec.node_id, - "label": &started.label, - "image": &spec.image, "contract_id": started.instance.contract_id, - "offer_id": started.instance.offer_id, - "host_id": started.instance.host_id, - "gpu_name": &started.instance.gpu_name, - "gpu_ram": started.instance.gpu_ram, - "dph_total": started.instance.dph_total, + "host": &started.endpoint.host, + "port": started.endpoint.port, + "user": &started.endpoint.user, }) .to_string(), }); @@ -1142,7 +1672,7 @@ where run_id: spec.run_id, node_id: spec.node_id, line: serde_json::json!({ - "type": "VastAiSshEndpointReady", + "type": "VastAiBootstrapObservationStarted", "run_id": spec.run_id, "node_id": spec.node_id, "contract_id": started.instance.contract_id, @@ -1158,6 +1688,7 @@ where started.endpoint, sink.clone(), self.bootstrap_producer.clone(), + self.config.lifecycle.clone(), ) { Ok(handle) => handle, Err(error) => { @@ -1192,7 +1723,18 @@ where VastAiNode { contract_id: started.instance.contract_id, bootstrap: Some(bootstrap), + provider_monitor: self.client.spawn_provider_monitor( + started.instance.contract_id, + started.label.clone(), + self.config.lifecycle.clone(), + spec.clone(), + sink.clone(), + ), host_id, + run_id: spec.run_id, + node_id: spec.node_id, + label: started.label, + sink: sink.clone(), }, ); results[index] = Some((spec, Ok(handle))); @@ -1232,6 +1774,23 @@ where let Some(node) = self.nodes.get_mut(&handle.id) else { return Ok(()); }; + if let Some(monitor) = node.provider_monitor.as_mut() { + monitor.stop(); + } + node.provider_monitor = None; + node.sink.observe(PluginObservation::ProviderLine { + run_id: node.run_id, + node_id: node.node_id, + line: serde_json::json!({ + "type": "VastAiRuntimeReadyAccepted", + "run_id": node.run_id, + "node_id": node.node_id, + "label": &node.label, + "contract_id": node.contract_id, + "classification": "runtime_ready_over_provider_staleness", + }) + .to_string(), + }); if let Some(mut bootstrap) = node.bootstrap.take() { self.bootstrap .stop_bootstrap(&mut bootstrap, BootstrapStopReason::RuntimeReady); @@ -1243,6 +1802,9 @@ where let Some(mut node) = self.nodes.remove(&handle.id) else { return Ok(()); }; + if let Some(mut monitor) = node.provider_monitor.take() { + monitor.stop(); + } if let Some(host_id) = node.host_id { self.leased_host_ids.remove(&host_id); } @@ -1250,7 +1812,22 @@ where self.bootstrap .stop_bootstrap(&mut bootstrap, BootstrapStopReason::NodeStop); } - self.client.destroy_contract(node.contract_id) + let result = self.client.destroy_contract(node.contract_id); + node.sink.observe(PluginObservation::ProviderLine { + run_id: node.run_id, + node_id: node.node_id, + line: serde_json::json!({ + "type": "VastAiContractCleanup", + "run_id": node.run_id, + "node_id": node.node_id, + "label": &node.label, + "contract_id": node.contract_id, + "result": if result.is_ok() { "ok" } else { "failed" }, + "error": result.as_ref().err(), + }) + .to_string(), + }); + result } } @@ -1294,6 +1871,7 @@ mod tests { _endpoint: VastAiSshEndpoint, _sink: PluginSink, _producer: Option, + _lifecycle: LifecyclePolicy, ) -> Result { panic!("build_request tests must not start SSH bootstrap") } @@ -1391,6 +1969,7 @@ mod tests { _endpoint: VastAiSshEndpoint, _sink: PluginSink, _producer: Option, + _lifecycle: LifecyclePolicy, ) -> Result { Ok(spec.node_id) } @@ -1471,6 +2050,21 @@ mod tests { } } + #[test] + fn ssh_bootstrap_observation_classifies_auth_and_transport_failures() { + for (line, expected) in [ + ("Permission denied (publickey).", Some("auth_denied")), + ( + "ssh: connect to host ssh5.vast.ai port 22017: Connection refused", + Some("refused"), + ), + ("ssh: connect timed out", Some("timeout")), + ("debug1: permanently_set_uid", None), + ] { + assert_eq!(classify_ssh_observation(line), expected, "{line}"); + } + } + #[test] fn ssh_bootstrap_args_include_verbose_flag_and_identity_when_configured() { let endpoint = VastAiSshEndpoint { diff --git a/tools/vastai/src/client.rs b/tools/vastai/src/client.rs index 8b78f55..24717b1 100644 --- a/tools/vastai/src/client.rs +++ b/tools/vastai/src/client.rs @@ -1,7 +1,8 @@ use std::time::Duration; use crate::types::{ - LabeledInstance, LifecyclePolicy, Offer, ProvisionRequest, ProvisionedFleet, RunningInstance, + LabeledInstance, LifecyclePolicy, Offer, ProviderInstanceStatus, ProvisionRequest, + ProvisionedFleet, RunningInstance, }; /// Small convenience wrapper around a reqwest client + vast.ai endpoint. @@ -60,6 +61,18 @@ impl VastClient { crate::lease::provision_fleet(&self.http, &self.base_url, &self.api_key, req).await } + pub async fn instance_status( + &self, + contract_id: u64, + ) -> Result { + crate::monitor::fetch_instance_status( + &self.http, + &self.base_url, + &self.api_key, + contract_id, + ) + .await + } pub async fn wait_for_running( &self, contract_id: u64, @@ -75,6 +88,23 @@ impl VastClient { .await } + pub async fn wait_for_ssh_endpoint( + &self, + contract_id: u64, + label: &str, + policy: &LifecyclePolicy, + ) -> Result { + crate::monitor::wait_for_ssh_endpoint_with_policy( + &self.http, + &self.base_url, + &self.api_key, + contract_id, + label, + policy, + ) + .await + } + pub async fn list_by_label(&self, label: &str) -> Result, String> { crate::teardown::list_instances_by_label(&self.http, &self.base_url, &self.api_key, label) .await diff --git a/tools/vastai/src/lease.rs b/tools/vastai/src/lease.rs index c2b5027..4428d9b 100644 --- a/tools/vastai/src/lease.rs +++ b/tools/vastai/src/lease.rs @@ -398,9 +398,9 @@ mod tests { .and(path("/api/v0/instances/101/")) .respond_with(ResponseTemplate::new(200).set_body_json(json!({ "instances": { - "actual_status": "loading", + "actual_status": "error", "intended_status": "running", - "status_msg": "still pulling" + "status_msg": "container failed before runtime readiness" } }))) .mount(&server) diff --git a/tools/vastai/src/lib.rs b/tools/vastai/src/lib.rs index d38ecb5..8d7ab2c 100644 --- a/tools/vastai/src/lib.rs +++ b/tools/vastai/src/lib.rs @@ -20,7 +20,10 @@ pub mod types; pub use client::VastClient; pub use lease::{confirm_lease, provision_fleet}; pub use logs::{fetch_logs, request_logs}; -pub use monitor::{wait_for_running, wait_for_running_with_policy}; +pub use monitor::{ + fetch_instance_status, wait_for_running, wait_for_running_with_policy, + wait_for_ssh_endpoint_with_policy, +}; pub use pricing::CostModel; pub use provision::create_instance; pub use search::{plan_distinct_host_first_wave, select_offer_pool, select_offer_pool_with_policy}; @@ -29,6 +32,6 @@ pub use teardown::{ }; pub use types::{ ContractRef, CreateInstanceRequest, FleetState, InstanceInfo, LabeledInstance, LifecyclePolicy, - Offer, ProvisionRequest, ProvisionedFleet, ProvisionedInstance, RunningInstance, - SelectionPolicy, VastAiFailureClass, classify_vastai_error, + Offer, ProviderInstanceStatus, ProvisionRequest, ProvisionedFleet, ProvisionedInstance, + RunningInstance, SelectionPolicy, VastAiFailureClass, classify_vastai_error, }; diff --git a/tools/vastai/src/monitor.rs b/tools/vastai/src/monitor.rs index 7ea6210..7717118 100644 --- a/tools/vastai/src/monitor.rs +++ b/tools/vastai/src/monitor.rs @@ -1,4 +1,92 @@ -use crate::types::{InstanceResponse, LifecyclePolicy, RunningInstance}; +use crate::types::{ + InstanceResponse, LabeledInstance, LifecyclePolicy, ProviderInstanceStatus, RunningInstance, +}; + +pub async fn fetch_instance_status( + client: &reqwest::Client, + base_url: &str, + api_key: &str, + contract_id: u64, +) -> Result { + let url = format!("{base_url}/api/v0/instances/{contract_id}/"); + let resp = client + .get(&url) + .header("Authorization", format!("Bearer {api_key}")) + .send() + .await + .map_err(|e| format!("fetch_instance_status request failed: {e}"))?; + + if resp.status() == reqwest::StatusCode::NOT_FOUND { + let body = resp.text().await.unwrap_or_default(); + return Err(format!( + "instance {contract_id} not found while fetching provider status: {}", + body.chars().take(80).collect::(), + )); + } + if !resp.status().is_success() { + let status = resp.status(); + let body = resp.text().await.unwrap_or_default(); + return Err(format!( + "fetch_instance_status HTTP {status}: {}", + body.chars().take(80).collect::(), + )); + } + + let wrapper: InstanceResponse = resp + .json() + .await + .map_err(|e| format!("fetch_instance_status parse failed: {e}"))?; + Ok(wrapper.instances.into()) +} + +fn provider_terminal_error( + contract_id: u64, + actual: &str, + intended: &str, + msg: Option<&str>, +) -> Option { + if let Some(m) = msg { + let lower = m.to_ascii_lowercase(); + if lower.contains("error") || lower.contains("failed") { + return Some(format!("instance {contract_id} error: {m}")); + } + } + if intended == "stopped" && actual != "running" { + return Some(format!( + "instance {contract_id} stopped: {}", + msg.unwrap_or_default() + )); + } + match actual { + "exited" | "error" | "stopped" => Some(format!( + "instance {contract_id} reached terminal status: {actual}" + )), + _ => None, + } +} + +fn labeled_endpoint(instance: &LabeledInstance) -> Option { + let host = if instance.ssh_host.trim().is_empty() { + instance.public_ipaddr.trim() + } else { + instance.ssh_host.trim() + }; + if host.is_empty() || instance.ssh_port == 0 { + return None; + } + Some(RunningInstance { + ip: host.to_owned(), + port: instance.ssh_port, + }) +} + +fn missing_instance(error: &str) -> bool { + error.contains("not found while fetching provider status") +} + +fn parse_failed(error: &str) -> bool { + error.contains("parse failed") +} /// Historical env-backed polling wrapper. pub async fn wait_for_running( @@ -20,67 +108,41 @@ pub async fn wait_for_running_with_policy( contract_id: u64, policy: &LifecyclePolicy, ) -> Result { - let url = format!("{base_url}/api/v0/instances/{contract_id}/"); let mut state_since = std::time::Instant::now(); let mut last_state: Option = None; - + let mut running_without_endpoint_since = None; let mut poll = 0_u64; + loop { poll += 1; - let resp = match client - .get(&url) - .header("Authorization", format!("Bearer {api_key}")) - .send() - .await - { - Ok(r) => r, - Err(e) => { - eprintln!(" contract {contract_id} poll {poll}: request error: {e} (retrying)"); + let status = match fetch_instance_status(client, base_url, api_key, contract_id).await { + Ok(status) => status, + Err(error) if missing_instance(&error) => { + return Err(error.replace( + "while fetching provider status", + "while waiting for running", + )); + } + Err(error) if parse_failed(&error) => return Err(error), + Err(error) => { + eprintln!(" contract {contract_id} poll {poll}: {error} (retrying)"); tokio::time::sleep(policy.poll_interval).await; continue; } }; - if resp.status() == reqwest::StatusCode::NOT_FOUND { - let body = resp.text().await.unwrap_or_default(); - return Err(format!( - "instance {contract_id} not found while waiting for running: {}", - body.chars().take(80).collect::(), - )); - } - if !resp.status().is_success() { - let status = resp.status(); - let body = resp.text().await.unwrap_or_default(); - eprintln!( - " contract {contract_id} poll {poll}: HTTP {status} (retrying): {}", - body.chars().take(80).collect::(), - ); - tokio::time::sleep(policy.poll_interval).await; - continue; - } - - let wrapper: InstanceResponse = resp - .json() - .await - .map_err(|e| format!("wait_for_running parse failed: {e}"))?; - let status = wrapper.instances; - - let actual = status.actual_status.as_deref().unwrap_or("unknown"); - let intended = status.intended_status.as_deref().unwrap_or("unknown"); - - let msg = status.status_msg.clone(); - let disk = status.disk_usage; - + let actual = status.actual_status.as_str(); + let intended = status.intended_status.as_str(); if last_state.as_deref() != Some(actual) { state_since = std::time::Instant::now(); } - last_state = Some(actual.to_string()); + last_state = Some(actual.to_owned()); let in_state = state_since.elapsed().as_secs(); - let msg_disp = match msg.as_deref() { + let msg_disp = match status.status_msg.as_deref() { Some(m) if !m.is_empty() => format!(" msg=\"{m}\""), _ => String::new(), }; - let disk_disp = match disk { + let disk_disp = match status.disk_usage { Some(d) if d >= 0.0 => format!(" disk={d:.2}GB"), _ => String::new(), }; @@ -88,41 +150,118 @@ pub async fn wait_for_running_with_policy( " contract {contract_id} poll {poll}: status={actual} in-state={in_state}s{msg_disp}{disk_disp}", ); - if let Some(m) = &msg { - if m.contains("Error") || m.contains("failed") { - return Err(format!("instance {contract_id} error: {m}")); + if let Some(error) = + provider_terminal_error(contract_id, actual, intended, status.status_msg.as_deref()) + { + return Err(error); + } + if actual == "running" { + if let Some(endpoint) = status.ssh_endpoint() { + return Ok(endpoint); } - } - if intended == "stopped" && actual != "running" { - return Err(format!( - "instance {contract_id} stopped: {}", - msg.unwrap_or_default() - )); - } - if !policy.state_timeout.is_zero() && state_since.elapsed() >= policy.state_timeout { - return Err(format!( - "instance {contract_id} stuck in status {actual} for {}s", - policy.state_timeout.as_secs() - )); - } - - match actual { - "running" => { - let ip = status - .public_ipaddr - .unwrap_or_else(|| "unknown".to_string()); - let port = status.ssh_port.unwrap_or(0); - return Ok(RunningInstance { ip, port }); - } - "exited" | "error" => { + let since = running_without_endpoint_since + .get_or_insert_with(std::time::Instant::now) + .elapsed(); + if !policy.state_timeout.is_zero() && since >= policy.state_timeout { return Err(format!( - "instance {contract_id} reached terminal status: {actual}" + "instance {contract_id} running without usable SSH endpoint for {}s", + policy.state_timeout.as_secs() )); } - _ => { - tokio::time::sleep(policy.poll_interval).await; + } else { + running_without_endpoint_since = None; + } + + tokio::time::sleep(policy.poll_interval).await; + } +} + +/// Poll provider data until a usable SSH endpoint exists, without requiring +/// provider `running` status first. +pub async fn wait_for_ssh_endpoint_with_policy( + client: &reqwest::Client, + base_url: &str, + api_key: &str, + contract_id: u64, + label: &str, + policy: &LifecyclePolicy, +) -> Result { + let mut running_without_endpoint_since = None; + let mut poll = 0_u64; + + loop { + poll += 1; + match crate::teardown::list_instances_by_label(client, base_url, api_key, label).await { + Ok(instances) => { + if let Some(instance) = instances + .iter() + .find(|instance| instance.contract_id == contract_id) + { + if let Some(endpoint) = labeled_endpoint(instance) { + eprintln!( + " contract {contract_id} endpoint poll {poll}: endpoint discovered from label status={}", + instance.actual_status + ); + return Ok(endpoint); + } + eprintln!( + " contract {contract_id} endpoint poll {poll}: label status={} endpoint missing", + instance.actual_status + ); + } + } + Err(error) => { + eprintln!( + " contract {contract_id} endpoint poll {poll}: list-by-label error: {error} (retrying)" + ); } } + + let status = match fetch_instance_status(client, base_url, api_key, contract_id).await { + Ok(status) => status, + Err(error) if missing_instance(&error) => return Err(error), + Err(error) if parse_failed(&error) => return Err(error), + Err(error) => { + eprintln!(" contract {contract_id} endpoint poll {poll}: {error} (retrying)"); + tokio::time::sleep(policy.poll_interval).await; + continue; + } + }; + + let actual = status.actual_status.as_str(); + let intended = status.intended_status.as_str(); + let msg = status.status_msg.as_deref(); + if let Some(endpoint) = status.ssh_endpoint() { + eprintln!( + " contract {contract_id} endpoint poll {poll}: endpoint discovered from provider status={actual}", + ); + return Ok(endpoint); + } + if let Some(error) = provider_terminal_error(contract_id, actual, intended, msg) { + return Err(error); + } + if actual == "running" { + let since = running_without_endpoint_since + .get_or_insert_with(std::time::Instant::now) + .elapsed(); + if !policy.state_timeout.is_zero() && since >= policy.state_timeout { + return Err(format!( + "instance {contract_id} running without usable SSH endpoint for {}s", + policy.state_timeout.as_secs() + )); + } + } else { + running_without_endpoint_since = None; + } + + let msg_disp = match msg { + Some(m) if !m.is_empty() => format!(" msg=\"{m}\""), + _ => String::new(), + }; + eprintln!( + " contract {contract_id} endpoint poll {poll}: status={actual} endpoint missing{msg_disp}; waiting" + ); + tokio::time::sleep(policy.poll_interval).await; } } @@ -137,7 +276,7 @@ mod tests { use super::*; #[tokio::test] - async fn stuck_loading_state_returns_error_instead_of_polling_forever() { + async fn loading_state_remains_slow_progress_before_terminal_evidence() { let server = MockServer::start().await; Mock::given(method("GET")) .and(path("/api/v0/instances/123/")) @@ -145,7 +284,7 @@ mod tests { "instances": { "actual_status": "loading", "intended_status": "running", - "status_msg": "afad30e59d72: Already exists" + "status_msg": "pulling image layers" } }))) .mount(&server) @@ -153,26 +292,129 @@ mod tests { let policy = LifecyclePolicy { poll_interval: Duration::from_millis(1), - state_timeout: Duration::from_millis(5), + state_timeout: Duration::from_millis(1), ..LifecyclePolicy::default() }; - let error = wait_for_running_with_policy( - &reqwest::Client::new(), - &server.uri(), - "secret", - 123, - &policy, + let still_pending = tokio::time::timeout( + Duration::from_millis(10), + wait_for_running_with_policy( + &reqwest::Client::new(), + &server.uri(), + "secret", + 123, + &policy, + ), ) - .await - .expect_err("stuck loading should be replaceable"); + .await; assert!( - error.contains("stuck in status loading"), - "error should name stuck provider state: {error}" + still_pending.is_err(), + "loading by itself should remain slow progress instead of replaceable failure" ); } + #[tokio::test] + async fn endpoint_discovery_uses_label_endpoint_before_running_status() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/api/v0/instances/")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "instances": [{ + "id": 789, + "label": "node-789", + "actual_status": "loading", + "ssh_host": "ssh5.vast.ai", + "ssh_port": 22017, + "public_ipaddr": "" + }] + }))) + .mount(&server) + .await; + + let policy = LifecyclePolicy { + poll_interval: Duration::from_secs(60), + state_timeout: Duration::from_secs(300), + ..LifecyclePolicy::default() + }; + + let endpoint = wait_for_ssh_endpoint_with_policy( + &reqwest::Client::new(), + &server.uri(), + "secret", + 789, + "node-789", + &policy, + ) + .await + .expect("known endpoint should start bootstrap observation before running status"); + + assert_eq!( + endpoint, + RunningInstance { + ip: "ssh5.vast.ai".to_owned(), + port: 22017, + } + ); + let requests = server.received_requests().await.expect("requests recorded"); + assert!( + requests + .iter() + .all(|request| request.url.path() != "/api/v0/instances/789/"), + "endpoint discovery should not wait for provider-running status once label data has a usable endpoint" + ); + } + + #[tokio::test] + async fn endpoint_missing_after_running_grace_fails() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/api/v0/instances/")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "instances": [{ + "id": 900, + "label": "node-900", + "actual_status": "running", + "ssh_host": "", + "ssh_port": 0, + "public_ipaddr": "" + }] + }))) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/api/v0/instances/900/")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "instances": { + "actual_status": "running", + "intended_status": "running" + } + }))) + .mount(&server) + .await; + + let policy = LifecyclePolicy { + poll_interval: Duration::from_millis(1), + state_timeout: Duration::from_millis(1), + ..LifecyclePolicy::default() + }; + + let error = wait_for_ssh_endpoint_with_policy( + &reqwest::Client::new(), + &server.uri(), + "secret", + 900, + "node-900", + &policy, + ) + .await + .expect_err("running without endpoint past grace is terminal"); + + assert!( + error.contains("running without usable SSH endpoint"), + "error should identify endpoint classification: {error}" + ); + } #[tokio::test] async fn missing_instance_returns_error_instead_of_polling_forever() { let server = MockServer::start().await; diff --git a/tools/vastai/src/types.rs b/tools/vastai/src/types.rs index 4ebd7a1..902521a 100644 --- a/tools/vastai/src/types.rs +++ b/tools/vastai/src/types.rs @@ -45,6 +45,31 @@ pub struct RunningInstance { pub port: u16, } +/// Provider status snapshot for one Vast.ai contract. +#[derive(Debug, Clone, PartialEq)] +pub struct ProviderInstanceStatus { + pub actual_status: String, + pub intended_status: String, + pub status_msg: Option, + pub public_ipaddr: Option, + pub ssh_port: Option, + pub disk_usage: Option, +} + +impl ProviderInstanceStatus { + pub fn ssh_endpoint(&self) -> Option { + let ip = self + .public_ipaddr + .as_deref() + .filter(|ip| !ip.trim().is_empty())?; + let port = self.ssh_port.filter(|port| *port > 0)?; + Some(RunningInstance { + ip: ip.to_owned(), + port, + }) + } +} + /// SSH endpoint + identity of a held instance, discovered by label. #[derive(Debug, Clone, PartialEq, Eq)] pub struct LabeledInstance { @@ -300,6 +325,23 @@ pub(crate) struct InstanceStatus { pub disk_usage: Option, } +impl From for ProviderInstanceStatus { + fn from(status: InstanceStatus) -> Self { + Self { + actual_status: status + .actual_status + .unwrap_or_else(|| "unknown".to_string()), + intended_status: status + .intended_status + .unwrap_or_else(|| "unknown".to_string()), + status_msg: status.status_msg, + public_ipaddr: status.public_ipaddr, + ssh_port: status.ssh_port, + disk_usage: status.disk_usage, + } + } +} + #[derive(Debug, Deserialize)] pub(crate) struct InstanceListResponse { pub instances: Vec,