diff --git a/Cargo.lock b/Cargo.lock index d2c222d..f1d7b40 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,6 +2,12 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "adler2" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa" + [[package]] name = "aead" version = "0.5.2" @@ -671,6 +677,15 @@ dependencies = [ "libc", ] +[[package]] +name = "crc32fast" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9481c1c90cbf2ac953f07c8d4a58aa3945c425b7185c9154d67a65e4230da511" +dependencies = [ + "cfg-if", +] + [[package]] name = "criterion" version = "0.5.1" @@ -1265,6 +1280,16 @@ version = "0.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" +[[package]] +name = "flate2" +version = "1.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "843fba2746e448b37e26a819579957415c8cef339bf08564fe8b7ddbd959573c" +dependencies = [ + "crc32fast", + "miniz_oxide", +] + [[package]] name = "fnv" version = "1.0.7" @@ -1757,7 +1782,7 @@ dependencies = [ "tokio", "tokio-rustls", "tower-service", - "webpki-roots", + "webpki-roots 1.0.8", ] [[package]] @@ -2062,7 +2087,7 @@ dependencies = [ "tracing", "url", "wasm-bindgen-futures", - "webpki-roots", + "webpki-roots 1.0.8", ] [[package]] @@ -2207,7 +2232,7 @@ dependencies = [ "tracing-subscriber", "url", "vergen-gitcl", - "webpki-roots", + "webpki-roots 1.0.8", "ws_stream_wasm", ] @@ -2443,6 +2468,16 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" +[[package]] +name = "miniz_oxide" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1fa76a2c86f704bdb222d66965fb3d63269ce38518b83cb0575fca855ebb6316" +dependencies = [ + "adler2", + "simd-adler32", +] + [[package]] name = "mio" version = "1.2.1" @@ -2491,6 +2526,7 @@ dependencies = [ "swactor-vastai", "tokio", "toml 0.8.23", + "ureq", ] [[package]] @@ -3662,7 +3698,7 @@ dependencies = [ "wasm-bindgen", "wasm-bindgen-futures", "web-sys", - "webpki-roots", + "webpki-roots 1.0.8", ] [[package]] @@ -4184,6 +4220,12 @@ version = "3.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "28d567dcbaf0049cb8ac2608a76cd95ff9e4412e1899d389ee400918ca7537f5" +[[package]] +name = "simd-adler32" +version = "0.3.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3a219298ac11a56ea9a6d2120044824d6f01aeb034955e7af7bc16858527deea" + [[package]] name = "simd_cesu8" version = "1.1.1" @@ -4639,7 +4681,7 @@ dependencies = [ "time", "tokio", "tokio-rustls", - "webpki-roots", + "webpki-roots 1.0.8", "x509-parser", ] @@ -4956,6 +4998,22 @@ version = "0.9.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8ecb6da28b8a351d773b68d5825ac39017e680750f980f3a1a85cd8dd28a47c1" +[[package]] +name = "ureq" +version = "2.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "02d1a66277ed75f640d608235660df48c8e3c19f3b4edb6a263315626cc3c01d" +dependencies = [ + "base64", + "flate2", + "log", + "once_cell", + "rustls", + "rustls-pki-types", + "url", + "webpki-roots 0.26.11", +] + [[package]] name = "url" version = "2.5.8" @@ -5223,6 +5281,15 @@ dependencies = [ "rustls-pki-types", ] +[[package]] +name = "webpki-roots" +version = "0.26.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "521bc38abb08001b01866da9f51eb7c5d647a19260e00054a8c7fd5f9e57f7a9" +dependencies = [ + "webpki-roots 1.0.8", +] + [[package]] name = "webpki-roots" version = "1.0.8" diff --git a/crates/mvp-system/Cargo.toml b/crates/mvp-system/Cargo.toml index 07d634e..74417fe 100644 --- a/crates/mvp-system/Cargo.toml +++ b/crates/mvp-system/Cargo.toml @@ -24,6 +24,7 @@ swactor-vastai = { path = "../../tools/vastai" } parking_lot = "0.12" blake3 = "1" toml = "0.8" +ureq = "2" [target.'cfg(target_os = "linux")'.dependencies] libc = "0.2" diff --git a/crates/mvp-system/src/actors/node_agent.rs b/crates/mvp-system/src/actors/node_agent.rs index cf3abab..a846627 100644 --- a/crates/mvp-system/src/actors/node_agent.rs +++ b/crates/mvp-system/src/actors/node_agent.rs @@ -4,7 +4,7 @@ use swactor::actor::{ActorAddress, ActorInterface}; use swactor::runtime::Ctx; use swactor_transport::{CodecRegistry, NetworkMessage}; -use crate::{run_plan, stage_controller as stage}; +use crate::{gguf_shard::StageShardPlan, run_plan, stage_controller as stage}; use super::codec::JsonCodec; use super::orchestrator::OrchestratorMsg; @@ -98,6 +98,7 @@ pub struct StageProvisionWire { pub model_id: String, pub gguf_source: run_plan::GgufSource, pub tokenizer: run_plan::TokenizerSource, + pub stage_shard_plan: Option, } impl StageProvisionWire { @@ -119,6 +120,7 @@ impl StageProvisionWire { self.gguf_source.clone(), self.tokenizer.clone(), ), + shard_plan: self.stage_shard_plan.clone(), } } } @@ -228,6 +230,7 @@ pub enum StageCommandWire { tokenizer: run_plan::TokenizerSource, layer_start: u32, layer_end_exclusive: u32, + stage_shard_plan: Option, }, RewireEdge { edge_id: u64, @@ -663,12 +666,17 @@ impl From<&stage::StageCommand> for StageCommandWire { layer_start: layer_range.start, layer_end_exclusive: layer_range.end_exclusive, }, - stage::StageCommand::LoadWeights { source, range } => Self::LoadWeights { + stage::StageCommand::LoadWeights { + source, + range, + shard_plan, + } => Self::LoadWeights { model_id: source.model_id.clone(), gguf_source: source.gguf_source.clone(), tokenizer: source.tokenizer.clone(), layer_start: range.start, layer_end_exclusive: range.end_exclusive, + stage_shard_plan: shard_plan.clone(), }, stage::StageCommand::RewireEdge { edge_id } => Self::RewireEdge { edge_id: edge_id.0 }, stage::StageCommand::ExecuteStep(step) => Self::ExecuteStep { diff --git a/crates/mvp-system/src/bin/mvp_chat.rs b/crates/mvp-system/src/bin/mvp_chat.rs index eaacbcd..5d3cbca 100644 --- a/crates/mvp-system/src/bin/mvp_chat.rs +++ b/crates/mvp-system/src/bin/mvp_chat.rs @@ -27,7 +27,8 @@ use mvp_system::config as chat_config; use mvp_system::config::ResolvedVastAiConfig; use mvp_system::endpoint_advertisement::EndpointAddrMask; use mvp_system::node_image::{ - NodeImageProvider, NodeImageRequest, PreparedNodeImage, prepare_node_image, + NodeImageProgressEvent, NodeImageProgressEventKind, NodeImageProgressSink, NodeImageProvider, + NodeImageRequest, PreparedNodeImage, prepare_node_image_with_progress, }; use mvp_system::node_provisioning::ProviderKind; use mvp_system::prompt_rpc::{PromptEvent, SubmitPrompt, write_json_line}; @@ -154,34 +155,38 @@ where ); progress.emit_benchmark_envelope(&config); confirm_vastai_if_needed(&config)?; + let prepare_runtime_started = Instant::now(); progress.emit( CHAT_RUNTIME_CHANNEL, "prepare_runtime", "started", json!({"provider": config.provider.as_str()}), ); - let image_ref = - match prepare_runtime_with_progress(&config, prepare_node_image, Some(&mut progress)) { - Ok(image_ref) => { - progress.emit( + let image_ref = match prepare_runtime_with_progress( + &config, + prepare_node_image_progress_adapter, + Some(&mut progress), + ) { + Ok(image_ref) => { + progress.emit( CHAT_RUNTIME_CHANNEL, "prepare_runtime", "ready", - json!({"image_ref": image_ref}), + json!({"image_ref": image_ref, "elapsed_ms": prepare_runtime_started.elapsed().as_millis()}), ); - image_ref - } - Err(error) => { - progress.emit( + image_ref + } + Err(error) => { + progress.emit( CHAT_RUNTIME_CHANNEL, "prepare_runtime", "failed", - json!({"error": error}), + json!({"error": error, "elapsed_ms": prepare_runtime_started.elapsed().as_millis()}), ); - progress.archive_pending()?; - return Err(error); - } - }; + progress.archive_pending()?; + return Err(error); + } + }; progress.emit( CHAT_COMPONENT_CHANNEL, "orchestrator_process_spawn", @@ -510,6 +515,66 @@ impl ChatDatastream { } } +impl NodeImageProgressSink for ChatDatastream { + fn emit(&mut self, event: NodeImageProgressEvent) { + let mut detail = serde_json::Map::new(); + if let Some(command_label) = event.command_label { + detail.insert("command_label".to_owned(), json!(command_label)); + } + if let Some(image_ref) = event.image_ref { + detail.insert("image_ref".to_owned(), json!(image_ref)); + } + if let Some(elapsed_ms) = event.elapsed_ms { + detail.insert("elapsed_ms".to_owned(), json!(elapsed_ms)); + } + + let (phase, status) = match event.kind { + NodeImageProgressEventKind::ImageReference { role, image_ref } => { + detail.insert("event".to_owned(), json!("image_ref")); + detail.insert("role".to_owned(), json!(role)); + detail.insert("image_ref".to_owned(), json!(image_ref)); + ("prepare_node_image", "image_ref") + } + NodeImageProgressEventKind::CommandStarted { program, args } => { + detail.insert("event".to_owned(), json!("command_start")); + detail.insert("program".to_owned(), json!(program)); + detail.insert("args".to_owned(), json!(args)); + ("node_image_command", "started") + } + NodeImageProgressEventKind::CommandStdout { line } => { + detail.insert("event".to_owned(), json!("stdout")); + detail.insert("stream".to_owned(), json!("stdout")); + detail.insert("line".to_owned(), json!(line)); + ("node_image_command", "stdout") + } + NodeImageProgressEventKind::CommandStderr { line } => { + detail.insert("event".to_owned(), json!("stderr")); + detail.insert("stream".to_owned(), json!("stderr")); + detail.insert("line".to_owned(), json!(line)); + ("node_image_command", "stderr") + } + NodeImageProgressEventKind::CommandExited { + status: command_status, + code, + success, + } => { + detail.insert("event".to_owned(), json!("command_exit")); + detail.insert("command_status".to_owned(), json!(command_status)); + detail.insert("exit_code".to_owned(), json!(code)); + detail.insert("success".to_owned(), json!(success)); + if let Some(elapsed_ms) = detail.get("elapsed_ms").cloned() { + detail.insert("duration_ms".to_owned(), elapsed_ms); + } + ( + "node_image_command", + if success { "exited" } else { "failed" }, + ) + } + }; + self.emit(CHAT_RUNTIME_CHANNEL, phase, status, Value::Object(detail)); + } +} + struct ChatFrameArchive { file: File, next_seq: u64, @@ -1366,26 +1431,44 @@ fn signal_orch_process_group(child: &Child, signal: libc::c_int) -> io::Result<( } } -type PrepareNodeImageFn = fn(NodeImageRequest) -> Result; +fn prepare_node_image_progress_adapter( + request: NodeImageRequest, + progress: Option<&mut dyn NodeImageProgressSink>, +) -> Result { + prepare_node_image_with_progress(request, progress) +} #[allow(dead_code)] fn prepare_runtime(config: &Config) -> Result { - prepare_runtime_with(config, prepare_node_image) + prepare_runtime_with(config, |request| { + prepare_node_image_with_progress(request, None) + }) } #[allow(dead_code)] -fn prepare_runtime_with( - config: &Config, - prepare_node_image_fn: PrepareNodeImageFn, -) -> Result { - prepare_runtime_with_progress(config, prepare_node_image_fn, None) +fn prepare_runtime_with(config: &Config, prepare_node_image_fn: F) -> Result +where + F: FnMut(NodeImageRequest) -> Result, +{ + let mut prepare_node_image_fn = prepare_node_image_fn; + prepare_runtime_with_progress( + config, + move |request, _progress| prepare_node_image_fn(request), + None, + ) } -fn prepare_runtime_with_progress( +fn prepare_runtime_with_progress( config: &Config, - prepare_node_image_fn: PrepareNodeImageFn, + mut prepare_node_image_fn: F, progress: Option<&mut ChatDatastream>, -) -> Result { +) -> Result +where + F: FnMut( + NodeImageRequest, + Option<&mut dyn NodeImageProgressSink>, + ) -> Result, +{ let mut progress = progress; let binary_mode = if config.skip_rebuild { "existing_artifact" @@ -1408,12 +1491,13 @@ fn prepare_runtime_with_progress( json!({"mode": config.orchestrator_launch_mode()}), ); } else { + let ensure_orch_started = Instant::now(); emit_chat_progress( &mut progress, CHAT_RUNTIME_CHANNEL, "ensure_orch_binary", "started", - json!({"mode": binary_mode}), + json!({"mode": binary_mode, "command_label": "ensure_orch_binary"}), ); match ensure_orch_binary(config) { Ok(()) => emit_chat_progress( @@ -1421,7 +1505,7 @@ fn prepare_runtime_with_progress( CHAT_RUNTIME_CHANNEL, "ensure_orch_binary", "ready", - json!({"mode": binary_mode}), + json!({"mode": binary_mode, "command_label": "ensure_orch_binary", "elapsed_ms": ensure_orch_started.elapsed().as_millis()}), ), Err(error) => { emit_chat_progress( @@ -1429,7 +1513,7 @@ fn prepare_runtime_with_progress( CHAT_RUNTIME_CHANNEL, "ensure_orch_binary", "failed", - json!({"mode": binary_mode, "error": error.as_str()}), + json!({"mode": binary_mode, "command_label": "ensure_orch_binary", "elapsed_ms": ensure_orch_started.elapsed().as_millis(), "error": error.as_str()}), ); return Err(error); } @@ -1468,7 +1552,7 @@ fn prepare_runtime_with_progress( CHAT_RUNTIME_CHANNEL, "prepare_node_image", "skipped", - json!({"provider": config.provider.as_str(), "reason": "process_provider"}), + json!({"provider": config.provider.as_str(), "command_label": "prepare_node_image", "reason": "process_provider"}), ); return Ok(config.node_image.clone()); } @@ -1480,23 +1564,24 @@ fn prepare_runtime_with_progress( CHAT_RUNTIME_CHANNEL, "ensure_worker_binary", "skipped", - json!({"mode": binary_mode, "reason": "vastai_remote_image"}), + json!({"mode": binary_mode, "command_label": "ensure_worker_binary", "reason": "vastai_remote_image"}), ); emit_chat_progress( &mut progress, CHAT_RUNTIME_CHANNEL, "prepare_node_image", "skipped", - json!({"provider": config.provider.as_str(), "reason": "skip_rebuild"}), + json!({"provider": config.provider.as_str(), "command_label": "prepare_node_image", "reason": "skip_rebuild"}), ); return Ok(config.node_image.clone()); } + let ensure_worker_started = Instant::now(); emit_chat_progress( &mut progress, CHAT_RUNTIME_CHANNEL, "ensure_worker_binary", "started", - json!({"mode": binary_mode}), + json!({"mode": binary_mode, "command_label": "ensure_worker_binary"}), ); match ensure_worker_binary(config) { Ok(()) => emit_chat_progress( @@ -1504,7 +1589,7 @@ fn prepare_runtime_with_progress( CHAT_RUNTIME_CHANNEL, "ensure_worker_binary", "ready", - json!({"mode": binary_mode}), + json!({"mode": binary_mode, "command_label": "ensure_worker_binary", "elapsed_ms": ensure_worker_started.elapsed().as_millis()}), ), Err(error) => { emit_chat_progress( @@ -1512,7 +1597,7 @@ fn prepare_runtime_with_progress( CHAT_RUNTIME_CHANNEL, "ensure_worker_binary", "failed", - json!({"mode": binary_mode, "error": error.as_str()}), + json!({"mode": binary_mode, "command_label": "ensure_worker_binary", "elapsed_ms": ensure_worker_started.elapsed().as_millis(), "error": error.as_str()}), ); return Err(error); } @@ -1522,17 +1607,18 @@ fn prepare_runtime_with_progress( CHAT_RUNTIME_CHANNEL, "prepare_node_image", "skipped", - json!({"provider": config.provider.as_str(), "reason": "skip_rebuild"}), + json!({"provider": config.provider.as_str(), "command_label": "prepare_node_image", "reason": "skip_rebuild"}), ); return Ok(config.node_image.clone()); } + let prepare_node_image_started = Instant::now(); emit_chat_progress( &mut progress, CHAT_RUNTIME_CHANNEL, "prepare_node_image", "started", - json!({"provider": config.provider.as_str()}), + json!({"provider": config.provider.as_str(), "command_label": "prepare_node_image", "image_tag": config.image_tag.as_deref()}), ); let node_bin = match node_bin_for_current_profile() { Ok(path) => path, @@ -1542,7 +1628,7 @@ fn prepare_runtime_with_progress( CHAT_RUNTIME_CHANNEL, "prepare_node_image", "failed", - json!({"provider": config.provider.as_str(), "error": error.as_str()}), + json!({"provider": config.provider.as_str(), "command_label": "prepare_node_image", "elapsed_ms": prepare_node_image_started.elapsed().as_millis(), "error": error.as_str()}), ); return Err(error); } @@ -1555,31 +1641,39 @@ fn prepare_runtime_with_progress( CHAT_RUNTIME_CHANNEL, "prepare_node_image", "failed", - json!({"provider": config.provider.as_str(), "error": error.as_str()}), + json!({"provider": config.provider.as_str(), "command_label": "prepare_node_image", "elapsed_ms": prepare_node_image_started.elapsed().as_millis(), "error": error.as_str()}), ); return Err(error); } }; - let prepared = match prepare_node_image_fn(NodeImageRequest { - requested_image: config.node_image.clone(), - base_image: BASE_NODE_IMAGE.to_owned(), - node_bin, - provider, - extra_tag: config.image_tag.clone(), - push: false, - force_refresh: false, - enabled: true, - }) { - Ok(prepared) => prepared, - Err(error) => { - emit_chat_progress( - &mut progress, - CHAT_RUNTIME_CHANNEL, - "prepare_node_image", - "failed", - json!({"provider": config.provider.as_str(), "error": error.as_str()}), - ); - return Err(error); + let prepared = { + let command_progress = progress + .as_deref_mut() + .map(|sink| sink as &mut dyn NodeImageProgressSink); + match prepare_node_image_fn( + NodeImageRequest { + requested_image: config.node_image.clone(), + base_image: BASE_NODE_IMAGE.to_owned(), + node_bin, + provider, + extra_tag: config.image_tag.clone(), + push: false, + force_refresh: false, + enabled: true, + }, + command_progress, + ) { + Ok(prepared) => prepared, + Err(error) => { + emit_chat_progress( + &mut progress, + CHAT_RUNTIME_CHANNEL, + "prepare_node_image", + "failed", + json!({"provider": config.provider.as_str(), "command_label": "prepare_node_image", "elapsed_ms": prepare_node_image_started.elapsed().as_millis(), "error": error.as_str()}), + ); + return Err(error); + } } }; emit_chat_progress( @@ -1587,7 +1681,7 @@ fn prepare_runtime_with_progress( CHAT_RUNTIME_CHANNEL, "prepare_node_image", "ready", - json!({"provider": config.provider.as_str(), "image_ref": prepared.image_ref}), + json!({"provider": config.provider.as_str(), "command_label": "prepare_node_image", "image_ref": prepared.image_ref, "elapsed_ms": prepare_node_image_started.elapsed().as_millis()}), ); Ok(prepared.image_ref) } @@ -3088,6 +3182,112 @@ relay_url = "https://relay.example" panic!("image preparer must not be called when --skip-rebuild is set") } + fn panic_prepare_node_image_with_progress( + _request: NodeImageRequest, + _progress: Option<&mut dyn NodeImageProgressSink>, + ) -> Result { + panic!("image preparer must not be called when --skip-rebuild is set") + } + + fn runtime_events(path: &Path) -> Vec { + fs::read_to_string(path) + .expect("read progress archive") + .lines() + .filter_map(|line| { + let outer: serde_json::Value = serde_json::from_str(line).ok()?; + if outer.get("channel").and_then(serde_json::Value::as_str) + != Some(CHAT_RUNTIME_CHANNEL) + { + return None; + } + outer + .get("payload")? + .get("value")? + .as_str() + .and_then(|text| serde_json::from_str::(text).ok()) + }) + .collect() + } + + fn emit_fake_node_image_progress( + progress: Option<&mut dyn NodeImageProgressSink>, + success: bool, + ) { + let Some(sink) = progress else { + return; + }; + sink.emit(NodeImageProgressEvent { + command_label: None, + image_ref: Some("docker.io/acme/node:prepared".to_owned()), + elapsed_ms: None, + kind: NodeImageProgressEventKind::ImageReference { + role: "resolved".to_owned(), + image_ref: "docker.io/acme/node:prepared".to_owned(), + }, + }); + sink.emit(NodeImageProgressEvent { + command_label: Some("build mvp node image".to_owned()), + image_ref: Some("docker.io/acme/node:prepared".to_owned()), + elapsed_ms: Some(0), + kind: NodeImageProgressEventKind::CommandStarted { + program: "fake-docker".to_owned(), + args: vec!["build".to_owned()], + }, + }); + sink.emit(NodeImageProgressEvent { + command_label: Some("build mvp node image".to_owned()), + image_ref: Some("docker.io/acme/node:prepared".to_owned()), + elapsed_ms: Some(1), + kind: NodeImageProgressEventKind::CommandStdout { + line: "building layer".to_owned(), + }, + }); + sink.emit(NodeImageProgressEvent { + command_label: Some("build mvp node image".to_owned()), + image_ref: Some("docker.io/acme/node:prepared".to_owned()), + elapsed_ms: Some(2), + kind: NodeImageProgressEventKind::CommandStderr { + line: "pushing metadata".to_owned(), + }, + }); + sink.emit(NodeImageProgressEvent { + command_label: Some("build mvp node image".to_owned()), + image_ref: Some("docker.io/acme/node:prepared".to_owned()), + elapsed_ms: Some(3), + kind: NodeImageProgressEventKind::CommandExited { + status: if success { + "exit status: 0".to_owned() + } else { + "exit status: 42".to_owned() + }, + code: Some(if success { 0 } else { 42 }), + success, + }, + }); + } + + fn fake_prepare_node_image_with_progress( + _request: NodeImageRequest, + progress: Option<&mut dyn NodeImageProgressSink>, + ) -> Result { + emit_fake_node_image_progress(progress, true); + Ok(PreparedNodeImage { + image_ref: "docker.io/acme/node:prepared".to_owned(), + tag: "prepared".to_owned(), + already_available: false, + built: true, + pushed: false, + }) + } + + fn failing_prepare_node_image_with_progress( + _request: NodeImageRequest, + progress: Option<&mut dyn NodeImageProgressSink>, + ) -> Result { + emit_fake_node_image_progress(progress, false); + Err("build mvp node image failed with exit status: 42".to_owned()) + } + #[test] fn skip_rebuild_requires_existing_artifacts_and_skips_image_preparation() { let temp = TempDir::new("skip-rebuild"); @@ -3110,6 +3310,188 @@ relay_url = "https://relay.example" assert_eq!(image_ref, "docker.io/acme/node:latest"); } + #[test] + fn prepare_runtime_progress_records_local_prep_details() { + let temp = TempDir::new("prep-progress"); + let archive_path = temp.path().join("frames.ndjson"); + let orch_bin = temp.path().join("mvp-orchestrator"); + let worker_bin = temp.path().join("mvp-worker-node"); + fs::write(&orch_bin, b"orch").expect("write orchestrator artifact"); + fs::write(&worker_bin, b"worker").expect("write worker artifact"); + let mut config = base_config(ProviderKind::Docker); + config.skip_rebuild = true; + config.orch_bin = orch_bin; + config.worker_bin = worker_bin; + config.node_image = "docker.io/acme/node:latest".to_owned(); + let mut progress = + ChatDatastream::new(91, Some(archive_path.clone())).expect("datastream constructs"); + + prepare_runtime_with_progress( + &config, + panic_prepare_node_image_with_progress, + Some(&mut progress), + ) + .expect("skip rebuild uses existing artifacts"); + progress.archive_pending().expect("archive prep frames"); + + let events = fs::read_to_string(&archive_path).expect("read archive"); + let inner_events = events + .lines() + .filter_map(|line| { + let outer: serde_json::Value = serde_json::from_str(line).ok()?; + if outer.get("channel").and_then(serde_json::Value::as_str) + != Some(CHAT_RUNTIME_CHANNEL) + { + return None; + } + let inner = outer + .get("payload")? + .get("value")? + .as_str() + .and_then(|text| serde_json::from_str::(text).ok())?; + Some(inner) + }) + .collect::>(); + let ensure_ready = inner_events + .iter() + .find(|event| { + event.get("phase").and_then(serde_json::Value::as_str) == Some("ensure_orch_binary") + && event.get("status").and_then(serde_json::Value::as_str) == Some("ready") + }) + .expect("ensure_orch_binary ready event"); + assert_eq!( + ensure_ready + .pointer("/detail/command_label") + .and_then(serde_json::Value::as_str), + Some("ensure_orch_binary") + ); + assert!( + ensure_ready + .pointer("/detail/elapsed_ms") + .and_then(serde_json::Value::as_u64) + .is_some() + ); + let image_skip = inner_events + .iter() + .find(|event| { + event.get("phase").and_then(serde_json::Value::as_str) == Some("prepare_node_image") + && event.get("status").and_then(serde_json::Value::as_str) == Some("skipped") + }) + .expect("prepare_node_image skipped event"); + assert_eq!( + image_skip + .pointer("/detail/reason") + .and_then(serde_json::Value::as_str), + Some("skip_rebuild") + ); + } + + #[test] + fn prepare_runtime_streams_node_image_command_progress() { + let temp = TempDir::new("node-image-command-progress"); + let archive_path = temp.path().join("frames.ndjson"); + let mut config = base_config(ProviderKind::Docker); + config.skip_rebuild = false; + config.gpu_run = true; + config.node_image = "docker.io/acme/node:latest".to_owned(); + let mut progress = + ChatDatastream::new(92, Some(archive_path.clone())).expect("datastream constructs"); + + let image_ref = prepare_runtime_with_progress( + &config, + fake_prepare_node_image_with_progress, + Some(&mut progress), + ) + .expect("fake image preparation succeeds"); + progress.archive_pending().expect("archive prep frames"); + assert_eq!(image_ref, "docker.io/acme/node:prepared"); + + let events = runtime_events(&archive_path); + assert!(events.iter().any(|event| { + event.get("phase").and_then(serde_json::Value::as_str) == Some("prepare_node_image") + && event.get("status").and_then(serde_json::Value::as_str) == Some("image_ref") + && event + .pointer("/detail/image_ref") + .and_then(serde_json::Value::as_str) + == Some("docker.io/acme/node:prepared") + })); + assert!(events.iter().any(|event| { + event.get("phase").and_then(serde_json::Value::as_str) == Some("node_image_command") + && event.get("status").and_then(serde_json::Value::as_str) == Some("started") + && event + .pointer("/detail/command_label") + .and_then(serde_json::Value::as_str) + == Some("build mvp node image") + && event + .pointer("/detail/program") + .and_then(serde_json::Value::as_str) + == Some("fake-docker") + })); + assert!(events.iter().any(|event| { + event.get("status").and_then(serde_json::Value::as_str) == Some("stdout") + && event + .pointer("/detail/line") + .and_then(serde_json::Value::as_str) + == Some("building layer") + })); + assert!(events.iter().any(|event| { + event.get("status").and_then(serde_json::Value::as_str) == Some("stderr") + && event + .pointer("/detail/line") + .and_then(serde_json::Value::as_str) + == Some("pushing metadata") + })); + assert!(events.iter().any(|event| { + event.get("status").and_then(serde_json::Value::as_str) == Some("exited") + && event + .pointer("/detail/command_status") + .and_then(serde_json::Value::as_str) + == Some("exit status: 0") + && event.pointer("/detail/duration_ms").is_some() + })); + } + + #[test] + fn prepare_runtime_command_failure_preserves_label_and_status() { + let temp = TempDir::new("node-image-command-failure"); + let archive_path = temp.path().join("frames.ndjson"); + let mut config = base_config(ProviderKind::Docker); + config.skip_rebuild = false; + config.gpu_run = true; + let mut progress = + ChatDatastream::new(93, Some(archive_path.clone())).expect("datastream constructs"); + + let error = prepare_runtime_with_progress( + &config, + failing_prepare_node_image_with_progress, + Some(&mut progress), + ) + .expect_err("fake image preparation failure propagates"); + progress.archive_pending().expect("archive prep frames"); + assert!(error.contains("build mvp node image"), "{error}"); + + let events = runtime_events(&archive_path); + let failure = events + .iter() + .find(|event| { + event.get("phase").and_then(serde_json::Value::as_str) == Some("node_image_command") + && event.get("status").and_then(serde_json::Value::as_str) == Some("failed") + }) + .expect("failed command progress event"); + assert_eq!( + failure + .pointer("/detail/command_label") + .and_then(serde_json::Value::as_str), + Some("build mvp node image") + ); + assert_eq!( + failure + .pointer("/detail/command_status") + .and_then(serde_json::Value::as_str), + Some("exit status: 42") + ); + } + #[test] fn vastai_skip_rebuild_uses_remote_image_without_worker_artifact() { let temp = TempDir::new("vastai-skip-rebuild"); diff --git a/crates/mvp-system/src/bin/worker_node.rs b/crates/mvp-system/src/bin/worker_node.rs index a3e5585..222b618 100644 --- a/crates/mvp-system/src/bin/worker_node.rs +++ b/crates/mvp-system/src/bin/worker_node.rs @@ -2,10 +2,10 @@ use std::collections::BTreeMap; use std::fs::{self, File, OpenOptions}; use std::io::{BufRead, BufReader, Read, Write}; use std::path::PathBuf; -use std::process::{Child, ChildStdin, ChildStdout, Command, ExitCode, Stdio}; +use std::process::{Child, ChildStdin, Command, ExitCode, Stdio}; use std::sync::{ Arc, - mpsc::{self, Receiver}, + mpsc::{self, Receiver, RecvTimeoutError, Sender}, }; use std::thread; use std::time::{Duration, Instant}; @@ -38,6 +38,7 @@ use mvp_system::edge_establisher as edge; use mvp_system::endpoint_advertisement::{ EndpointAddrMask, MVP_IROH_ENDPOINT_ADDR_MASK_ENV, advertised_endpoint, }; +use mvp_system::gguf_shard::{StageShardPlan, materialize_stage_shard_http}; use mvp_system::gpu_worker_ingress_parser as ingress; use mvp_system::prompt_rpc::{PromptEvent, TokenizerEvent}; use mvp_system::relay_provisioning::relay_runtime_config_from_env; @@ -64,6 +65,8 @@ const NODE_STAGE_CHANNEL: &str = "mvp.node.stage"; const NODE_WORKER_CHANNEL: &str = "mvp.node.worker"; const NODE_PROMPT_CHANNEL: &str = "mvp.node.prompt"; const NODE_SHUTDOWN_CHANNEL: &str = "mvp.node.shutdown"; +const NODE_SAMPLER_CHANNEL: &str = "mvp.node.sampler"; +const WORKER_COMMAND_WAIT_TELEMETRY_INTERVAL: Duration = Duration::from_secs(1); fn node_event_payload( config: &DeploymentConfig, @@ -417,12 +420,136 @@ fn drain_debug_join_commands( } } +#[derive(Clone, Copy)] +struct SamplerHealthContext { + run_id: u64, + node_id: u64, + stage_index: u32, +} + +impl SamplerHealthContext { + fn from_config(config: &DeploymentConfig) -> Self { + Self { + run_id: config.run_id, + node_id: config.logical_node_id, + stage_index: config.stage_index, + } + } +} + +fn sampler_health_payload( + context: SamplerHealthContext, + sampler: &str, + sample_channel: &str, + status: &str, + detail: Value, +) -> Value { + json!({ + "type":"SamplerHealth", + "schema":"mvp.node.sampler.health.v1", + "run_id":context.run_id, + "node_id":context.node_id, + "stage_index":context.stage_index, + "phase":"host_sampler_health", + "status":status, + "sampler":sampler, + "sample_channel":sample_channel, + "detail":detail, + "benchmark":benchmark_observability::stamp("mvp-worker-node"), + }) +} + +fn submit_sampler_health( + producer: &DatastreamProducer, + channel: ChannelId, + context: SamplerHealthContext, + sampler: &str, + sample_channel: &str, + status: &str, + detail: Value, +) { + producer.submit_text( + channel, + sampler_health_payload(context, sampler, sample_channel, status, detail).to_string(), + ); +} + +fn submit_sampler_started( + producer: &DatastreamProducer, + health_channel: ChannelId, + context: SamplerHealthContext, + sampler: &str, + sample_channel: &str, + interval: Duration, +) { + submit_sampler_health( + producer, + health_channel, + context, + sampler, + sample_channel, + "started", + json!({"state":"started","sample_interval_ms":duration_ms_u64(interval)}), + ); + submit_sampler_health( + producer, + health_channel, + context, + sampler, + sample_channel, + "waiting", + json!({"state":"no_sample_yet","sample_interval_ms":duration_ms_u64(interval)}), + ); +} + +fn submit_sampler_sample_health( + producer: &DatastreamProducer, + health_channel: ChannelId, + context: SamplerHealthContext, + sampler: &str, + sample_channel: &str, + seq: u64, + error: Option<&str>, +) { + match error { + Some(error) => submit_sampler_health( + producer, + health_channel, + context, + sampler, + sample_channel, + "failed", + json!({"state":"error","sample_seq":seq,"error":error}), + ), + None => submit_sampler_health( + producer, + health_channel, + context, + sampler, + sample_channel, + "ready", + json!({"state":"sample_observed","sample_seq":seq}), + ), + } +} + fn spawn_host_gpu_sampler( handle: tokio::runtime::Handle, producer: DatastreamProducer, channel: ChannelId, + health_channel: ChannelId, + health_context: SamplerHealthContext, ) { handle.spawn(async move { + let sample_channel = datastream::hardware::gpu::HOST_GPU_CHANNEL; + submit_sampler_started( + &producer, + health_channel, + health_context, + "gpu", + sample_channel, + datastream::hardware::gpu::GPU_SAMPLE_INTERVAL, + ); let mut seq = 0_u64; let mut interval = tokio::time::interval(datastream::hardware::gpu::GPU_SAMPLE_INTERVAL); @@ -442,6 +569,15 @@ fn spawn_host_gpu_sampler( ), }; + submit_sampler_sample_health( + &producer, + health_channel, + health_context, + "gpu", + sample_channel, + sample_seq, + sample.error.as_deref(), + ); seq = seq.saturating_add(1); producer.submit_record(channel, &sample); } @@ -452,9 +588,20 @@ fn spawn_host_cpu_sampler( handle: tokio::runtime::Handle, producer: DatastreamProducer, channel: ChannelId, + health_channel: ChannelId, + health_context: SamplerHealthContext, watched_pids: Vec, ) { handle.spawn(async move { + let sample_channel = datastream::hardware::cpu::HOST_CPU_CHANNEL; + submit_sampler_started( + &producer, + health_channel, + health_context, + "cpu", + sample_channel, + datastream::hardware::cpu::CPU_SAMPLE_INTERVAL, + ); let mut seq = 0_u64; let mut sampler = datastream::hardware::cpu::CpuSampler::new(watched_pids); let mut interval = tokio::time::interval(datastream::hardware::cpu::CPU_SAMPLE_INTERVAL); @@ -463,6 +610,15 @@ fn spawn_host_cpu_sampler( interval.tick().await; let sample = sampler.sample(seq); + submit_sampler_sample_health( + &producer, + health_channel, + health_context, + "cpu", + sample_channel, + seq, + sample.error.as_deref(), + ); seq = seq.saturating_add(1); producer.submit_record(channel, &sample); } @@ -472,8 +628,19 @@ fn spawn_host_net_sampler( handle: tokio::runtime::Handle, producer: DatastreamProducer, channel: ChannelId, + health_channel: ChannelId, + health_context: SamplerHealthContext, ) { handle.spawn(async move { + let sample_channel = datastream::hardware::net::HOST_NET_CHANNEL; + submit_sampler_started( + &producer, + health_channel, + health_context, + "net", + sample_channel, + datastream::hardware::net::HOST_NET_SAMPLE_INTERVAL, + ); let mut seq = 0_u64; let mut interval = tokio::time::interval(datastream::hardware::net::HOST_NET_SAMPLE_INTERVAL); @@ -494,6 +661,15 @@ fn spawn_host_net_sampler( ), }; + submit_sampler_sample_health( + &producer, + health_channel, + health_context, + "net", + sample_channel, + sample_seq, + sample.error.as_deref(), + ); seq = seq.saturating_add(1); producer.submit_record(channel, &sample); } @@ -1391,6 +1567,9 @@ fn main() -> ExitCode { args.remove(0); return debug_join_client_main(args); } + if args.first().map(String::as_str) == Some("stage-shard-fetcher") { + return stage_shard_fetcher_main(); + } match run() { Ok(()) => ExitCode::SUCCESS, Err(error) => { @@ -1400,6 +1579,36 @@ fn main() -> ExitCode { } } +#[derive(serde::Deserialize, serde::Serialize)] +struct StageShardFetchRequest { + plan: StageShardPlan, + output_path: PathBuf, +} + +fn stage_shard_fetcher_main() -> ExitCode { + match run_stage_shard_fetcher() { + Ok(()) => ExitCode::SUCCESS, + Err(error) => { + let event = json!({"type":"StageShardFetchFailed","error":error}); + println!("{event}"); + ExitCode::from(1) + } + } +} + +fn run_stage_shard_fetcher() -> Result<(), String> { + let mut input = String::new(); + std::io::stdin() + .read_to_string(&mut input) + .map_err(|e| format!("read stage shard fetch request: {e}"))?; + let request: StageShardFetchRequest = serde_json::from_str(&input) + .map_err(|e| format!("parse stage shard fetch request: {e}"))?; + materialize_stage_shard_http(&request.plan, &request.output_path, |event| { + println!("{event}"); + let _ = std::io::stdout().flush(); + }) +} + fn run() -> Result<(), String> { let config = DeploymentConfig::from_env()?; emit_stdio_node_event( @@ -1588,15 +1797,21 @@ fn run() -> Result<(), String> { } }; stack.register_local_actor(driver.register_actor(datastream_publisher, 1)); + let sampler_health_channel = datastream.channel_by_name(NODE_SAMPLER_CHANNEL); + let sampler_health_context = SamplerHealthContext::from_config(&config); spawn_host_gpu_sampler( tokio.handle().clone(), datastream.producer.clone(), datastream.channels.host_gpu, + sampler_health_channel, + sampler_health_context, ); spawn_host_net_sampler( tokio.handle().clone(), datastream.producer.clone(), datastream.channels.host_net, + sampler_health_channel, + sampler_health_context, ); spawn_arena_sampler( tokio.handle().clone(), @@ -1752,6 +1967,8 @@ fn run() -> Result<(), String> { tokio.handle().clone(), datastream.producer.clone(), datastream.channels.host_cpu, + sampler_health_channel, + sampler_health_context, vec![std::process::id(), worker.pid()], ); emit_stdio_node_event( @@ -2134,6 +2351,7 @@ impl NodeDatastream { NODE_WORKER_CHANNEL, NODE_PROMPT_CHANNEL, NODE_SHUTDOWN_CHANNEL, + NODE_SAMPLER_CHANNEL, "mvp.worker.initialize", "mvp.worker.role", "mvp.worker.weights", @@ -2806,6 +3024,189 @@ fn handle_decode_tokens_request( .map_err(|e| format!("send tokenizer decode response: {e}")) } +#[derive(Clone, Copy)] +enum StageShardProcessStream { + Stdout, + Stderr, +} + +struct StageShardProcessLine { + stream: StageShardProcessStream, + line: String, +} + +fn stage_shard_cache_path(plan: &StageShardPlan) -> PathBuf { + let root = std::env::var("MVP_MODEL_CACHE_DIR") + .ok() + .filter(|value| !value.trim().is_empty()) + .map(PathBuf::from) + .unwrap_or_else(|| PathBuf::from("/var/cache/mvp-models")); + root.join(plan.cache_file_name()) +} + +fn spawn_stage_shard_reader( + stream: StageShardProcessStream, + reader: R, + tx: mpsc::Sender, +) { + thread::spawn(move || { + let mut reader = BufReader::new(reader); + let mut line = String::new(); + loop { + line.clear(); + match reader.read_line(&mut line) { + Ok(0) => break, + Ok(_) => { + let _ = tx.send(StageShardProcessLine { + stream, + line: line.trim_end_matches(['\r', '\n']).to_owned(), + }); + } + Err(error) => { + let _ = tx.send(StageShardProcessLine { + stream, + line: format!("reader error: {error}"), + }); + break; + } + } + } + }); +} + +fn publish_stage_shard_fetch_event( + datastream: &mut NodeDatastream, + config: &DeploymentConfig, + event: &Value, +) -> Result<(), String> { + let channel = datastream.channel_by_name("mvp.worker.weights"); + let payload = node_event_payload(config, "stage_shard_fetch", "event", event.clone()); + datastream.submit_text(channel, payload.to_string()); + emit_stdio_datastream_frame("mvp.worker.weights", &payload) + .map_err(|e| format!("emit stage shard fetch datastream frame: {e}"))?; + datastream.tick(); + Ok(()) +} + +fn materialize_stage_shard_with_process( + plan: &StageShardPlan, + config: &DeploymentConfig, + datastream: &mut NodeDatastream, + driver: &mut IrohDriver, + stack: &DistributionRuntimeStack, +) -> Result { + let output_path = stage_shard_cache_path(plan); + if output_path.is_file() { + let event = json!({ + "type":"StageShardCacheReady", + "stage_index":plan.stage_index, + "path":output_path, + "cache_hit":true, + }); + publish_stage_shard_fetch_event(datastream, config, &event)?; + return Ok(output_path); + } + + let request = StageShardFetchRequest { + plan: plan.clone(), + output_path: output_path.clone(), + }; + let request_json = serde_json::to_vec(&request) + .map_err(|e| format!("serialize stage shard fetch request: {e}"))?; + let exe = std::env::current_exe().map_err(|e| format!("locate worker node executable: {e}"))?; + let mut child = Command::new(exe) + .arg("stage-shard-fetcher") + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .map_err(|e| format!("spawn stage shard fetcher: {e}"))?; + if let Some(mut stdin) = child.stdin.take() { + stdin + .write_all(&request_json) + .map_err(|e| format!("write stage shard fetch request: {e}"))?; + } + let (tx, rx) = mpsc::channel::(); + if let Some(stdout) = child.stdout.take() { + spawn_stage_shard_reader(StageShardProcessStream::Stdout, stdout, tx.clone()); + } + if let Some(stderr) = child.stderr.take() { + spawn_stage_shard_reader(StageShardProcessStream::Stderr, stderr, tx); + } + + let mut ready_path = None; + + loop { + while let Ok(line) = rx.try_recv() { + if line.line.is_empty() { + continue; + } + let event = match line.stream { + StageShardProcessStream::Stdout => { + match serde_json::from_str::(&line.line) { + Ok(value) => value, + Err(error) => { + json!({"type":"StageShardFetchOutputParseFailed","line":line.line,"error":error.to_string()}) + } + } + } + StageShardProcessStream::Stderr => { + json!({"type":"StageShardFetchStderr","line":line.line}) + } + }; + if event.get("type").and_then(Value::as_str) == Some("StageShardReady") { + ready_path = event + .get("path") + .and_then(Value::as_str) + .map(PathBuf::from) + .or_else(|| Some(output_path.clone())); + } + publish_stage_shard_fetch_event(datastream, config, &event)?; + } + if let Some(status) = child + .try_wait() + .map_err(|e| format!("poll stage shard fetcher: {e}"))? + { + while let Ok(line) = rx.try_recv() { + if line.line.is_empty() { + continue; + } + let event = match line.stream { + StageShardProcessStream::Stdout => serde_json::from_str::(&line.line) + .unwrap_or_else(|error| { + json!({"type":"StageShardFetchOutputParseFailed","line":line.line,"error":error.to_string()}) + }), + StageShardProcessStream::Stderr => { + json!({"type":"StageShardFetchStderr","line":line.line}) + } + }; + if event.get("type").and_then(Value::as_str) == Some("StageShardReady") { + ready_path = event + .get("path") + .and_then(Value::as_str) + .map(PathBuf::from) + .or_else(|| Some(output_path.clone())); + } + publish_stage_shard_fetch_event(datastream, config, &event)?; + } + if status.success() { + let path = ready_path.unwrap_or_else(|| output_path.clone()); + if path.is_file() { + return Ok(path); + } + return Err(format!( + "stage shard fetcher exited successfully but {} is missing", + path.display() + )); + } + return Err(format!("stage shard fetcher exited with {status}")); + } + pump_network(driver, stack); + datastream.tick(); + thread::sleep(PUMP_INTERVAL); + } +} + fn handle_stage_command( command: StageCommandWire, config: &DeploymentConfig, @@ -2877,6 +3278,7 @@ fn handle_stage_command( tokenizer, layer_start, layer_end_exclusive, + stage_shard_plan, } => { let gguf_source_kind = match &gguf_source { GgufSource::LocalPath(_) => "local_path", @@ -2886,18 +3288,34 @@ fn handle_stage_command( TokenizerSource::EmbeddedGguf => "gguf", TokenizerSource::LocalPath(_) => "local_path", }; + let (resolved_gguf_source, using_stage_shard) = + if let Some(stage_plan) = stage_shard_plan { + let local_path = materialize_stage_shard_with_process( + &stage_plan, + config, + datastream, + driver, + stack, + )?; + ( + GgufSource::LocalPath(local_path.to_string_lossy().into_owned()), + true, + ) + } else { + (gguf_source, false) + }; emit_node_event( datastream, config, NODE_STAGE_CHANNEL, "load_weights", "started", - json!({"model_id":&model_id,"gguf_source":gguf_source_kind,"tokenizer":tokenizer_kind,"layer_range":{"start":layer_start,"end_exclusive":layer_end_exclusive}}), + json!({"model_id":&model_id,"gguf_source":gguf_source_kind,"tokenizer":tokenizer_kind,"stage_shard":using_stage_shard,"layer_range":{"start":layer_start,"end_exclusive":layer_end_exclusive}}), ); let mut pump = || pump_network(driver, stack); match worker.load_weights( model_id.clone(), - gguf_source, + resolved_gguf_source, tokenizer, layer_start, layer_end_exclusive, @@ -3223,10 +3641,242 @@ impl DeploymentConfig { } } +#[derive(Debug)] +enum HelperStdoutEvent { + Line(String), + Closed, + ReadError(String), +} + +#[derive(Clone, Copy)] +struct HelperCommandWaitConfig { + poll_interval: Duration, + telemetry_interval: Duration, +} + +impl HelperCommandWaitConfig { + fn production() -> Self { + Self { + poll_interval: PUMP_INTERVAL, + telemetry_interval: WORKER_COMMAND_WAIT_TELEMETRY_INTERVAL, + } + } +} + +fn spawn_helper_stdout_reader(reader: R, tx: Sender) { + thread::spawn(move || { + let mut reader = BufReader::new(reader); + loop { + let mut line = String::new(); + match reader.read_line(&mut line) { + Ok(0) => { + let _ = tx.send(HelperStdoutEvent::Closed); + break; + } + Ok(_) => { + if tx.send(HelperStdoutEvent::Line(line)).is_err() { + break; + } + } + Err(error) => { + let _ = tx.send(HelperStdoutEvent::ReadError(error.to_string())); + break; + } + } + } + }); +} + +fn drain_worker_stderr( + stderr_rx: &Receiver, + config: &DeploymentConfig, + datastream: &mut NodeDatastream, +) { + let mut emitted = false; + while let Ok(line) = stderr_rx.try_recv() { + let payload = node_event_payload(config, "worker_stderr", "observed", json!({"line":line})); + datastream.submit_text(datastream.channels.worker_stderr, payload.to_string()); + emitted = true; + } + if emitted { + datastream.tick(); + } +} + +#[allow(clippy::too_many_arguments)] +fn wait_for_helper_event( + stdout_rx: &Receiver, + stderr_rx: Option<&Receiver>, + expected: &str, + command_type: &str, + config: &DeploymentConfig, + datastream: &mut NodeDatastream, + channel: ChannelId, + channel_name: &str, + wait_config: HelperCommandWaitConfig, + pump: &mut dyn FnMut(), +) -> Result { + emit_node_event( + datastream, + config, + NODE_WORKER_CHANNEL, + "worker_stdout_read", + "started", + json!({"command_type":command_type,"expected_event_type":expected,"channel":channel_name}), + ); + let wait_started = Instant::now(); + let mut wait_cycles = 0_u64; + let mut next_telemetry_at = wait_started; + loop { + match stdout_rx.recv_timeout(wait_config.poll_interval) { + Ok(HelperStdoutEvent::Line(line)) => { + if let Some(stderr_rx) = stderr_rx { + drain_worker_stderr(stderr_rx, config, datastream); + } + let line_bytes = line.len(); + emit_node_event( + datastream, + config, + NODE_WORKER_CHANNEL, + "worker_stdout_read", + "ready", + json!({"command_type":command_type,"expected_event_type":expected,"channel":channel_name,"line_bytes":line_bytes}), + ); + emit_node_event( + datastream, + config, + NODE_WORKER_CHANNEL, + "worker_stdout_parse", + "started", + json!({"command_type":command_type,"expected_event_type":expected,"channel":channel_name,"line_bytes":line_bytes}), + ); + let value: Value = match serde_json::from_str(&line) { + Ok(value) => value, + Err(error) => { + emit_node_event( + datastream, + config, + NODE_WORKER_CHANNEL, + "worker_stdout_parse", + "failed", + json!({"command_type":command_type,"expected_event_type":expected,"channel":channel_name,"line_bytes":line_bytes,"error":error.to_string()}), + ); + return Err(format!("parse helper stdout {line:?}: {error}")); + } + }; + let worker_event_type = value + .get("type") + .and_then(Value::as_str) + .unwrap_or("unknown"); + emit_node_event( + datastream, + config, + NODE_WORKER_CHANNEL, + "worker_stdout_parse", + "ready", + json!({"command_type":command_type,"expected_event_type":expected,"channel":channel_name,"line_bytes":line_bytes,"worker_event_type":worker_event_type}), + ); + datastream.submit_text(channel, value.to_string()); + emit_stdio_datastream_frame(channel_name, &value) + .map_err(|e| format!("emit worker stdio datastream frame: {e}"))?; + datastream.tick(); + pump(); + if worker_event_type == "WorkerFatal" { + return Err(format!("worker fatal: {value}")); + } + if value.get("type").and_then(Value::as_str) == Some(expected) { + return Ok(value); + } + emit_node_event( + datastream, + config, + NODE_WORKER_CHANNEL, + "worker_event", + "observed", + json!({"command_type":command_type,"command_waiting_for":expected,"worker_event_type":worker_event_type,"event":value}), + ); + } + Ok(HelperStdoutEvent::Closed) => { + if let Some(stderr_rx) = stderr_rx { + drain_worker_stderr(stderr_rx, config, datastream); + } + emit_node_event( + datastream, + config, + NODE_WORKER_CHANNEL, + "worker_stdout_read", + "failed", + json!({"command_type":command_type,"expected_event_type":expected,"channel":channel_name,"line_bytes":0,"error":"stdout closed"}), + ); + return Err(format!( + "tinygrad helper stdout closed while waiting for {expected}" + )); + } + Ok(HelperStdoutEvent::ReadError(error)) => { + if let Some(stderr_rx) = stderr_rx { + drain_worker_stderr(stderr_rx, config, datastream); + } + emit_node_event( + datastream, + config, + NODE_WORKER_CHANNEL, + "worker_stdout_read", + "failed", + json!({"command_type":command_type,"expected_event_type":expected,"channel":channel_name,"error":error}), + ); + return Err(format!("read helper stdout: {error}")); + } + Err(RecvTimeoutError::Timeout) => { + wait_cycles = wait_cycles.saturating_add(1); + pump(); + if let Some(stderr_rx) = stderr_rx { + drain_worker_stderr(stderr_rx, config, datastream); + } + let now = Instant::now(); + if now >= next_telemetry_at { + emit_node_event( + datastream, + config, + NODE_WORKER_CHANNEL, + "worker_command_wait", + "waiting", + json!({ + "command_type":command_type, + "expected_event_type":expected, + "channel":channel_name, + "state":"busy_waiting_for_helper_stdout", + "elapsed_ms":duration_ms_u64(now.saturating_duration_since(wait_started)), + "wait_cycles":wait_cycles, + "poll_interval_ms":duration_ms_u64(wait_config.poll_interval), + }), + ); + next_telemetry_at = now + .checked_add(wait_config.telemetry_interval) + .unwrap_or(now); + } + datastream.tick(); + } + Err(RecvTimeoutError::Disconnected) => { + emit_node_event( + datastream, + config, + NODE_WORKER_CHANNEL, + "worker_stdout_read", + "failed", + json!({"command_type":command_type,"expected_event_type":expected,"channel":channel_name,"line_bytes":0,"error":"stdout reader disconnected"}), + ); + return Err(format!( + "tinygrad helper stdout reader disconnected while waiting for {expected}" + )); + } + } + } +} + struct TinygradWorker { child: Child, stdin: ChildStdin, - stdout: BufReader, + stdout_rx: Receiver, stderr_rx: Receiver, } @@ -3257,6 +3907,8 @@ impl TinygradWorker { .stderr .take() .ok_or_else(|| "tinygrad helper stderr missing".to_owned())?; + let (stdout_tx, stdout_rx) = mpsc::channel(); + spawn_helper_stdout_reader(stdout, stdout_tx); let (stderr_tx, stderr_rx) = mpsc::channel(); thread::spawn(move || { for line in BufReader::new(stderr).lines().map_while(Result::ok) { @@ -3268,7 +3920,7 @@ impl TinygradWorker { Ok(Self { child, stdin, - stdout: BufReader::new(stdout), + stdout_rx, stderr_rx, }) } @@ -3605,16 +4257,7 @@ impl TinygradWorker { } fn drain_stderr(&mut self, config: &DeploymentConfig, datastream: &mut NodeDatastream) { - let mut emitted = false; - while let Ok(line) = self.stderr_rx.try_recv() { - let payload = - node_event_payload(config, "worker_stderr", "observed", json!({"line":line})); - datastream.submit_text(datastream.channels.worker_stderr, payload.to_string()); - emitted = true; - } - if emitted { - datastream.tick(); - } + drain_worker_stderr(&self.stderr_rx, config, datastream); } fn command( @@ -3672,118 +4315,39 @@ impl TinygradWorker { "ready", json!({"command_type":command_type.as_str(),"expected_event_type":expected,"command_bytes":command_bytes}), ); - self.expect_event(expected, config, datastream, channel, channel_name, pump) + self.expect_event( + expected, + command_type.as_str(), + config, + datastream, + channel, + channel_name, + pump, + ) } fn expect_event( &mut self, expected: &str, + command_type: &str, config: &DeploymentConfig, datastream: &mut NodeDatastream, channel: ChannelId, channel_name: &str, pump: &mut dyn FnMut(), ) -> Result { - loop { - let mut line = String::new(); - emit_node_event( - datastream, - config, - NODE_WORKER_CHANNEL, - "worker_stdout_read", - "started", - json!({"expected_event_type":expected,"channel":channel_name}), - ); - let n = match self.stdout.read_line(&mut line) { - Ok(n) => n, - Err(error) => { - emit_node_event( - datastream, - config, - NODE_WORKER_CHANNEL, - "worker_stdout_read", - "failed", - json!({"expected_event_type":expected,"channel":channel_name,"error":error.to_string()}), - ); - return Err(format!("read helper stdout: {error}")); - } - }; - self.drain_stderr(config, datastream); - if n == 0 { - emit_node_event( - datastream, - config, - NODE_WORKER_CHANNEL, - "worker_stdout_read", - "failed", - json!({"expected_event_type":expected,"channel":channel_name,"line_bytes":0,"error":"stdout closed"}), - ); - return Err(format!( - "tinygrad helper stdout closed while waiting for {expected}" - )); - } - emit_node_event( - datastream, - config, - NODE_WORKER_CHANNEL, - "worker_stdout_read", - "ready", - json!({"expected_event_type":expected,"channel":channel_name,"line_bytes":n}), - ); - emit_node_event( - datastream, - config, - NODE_WORKER_CHANNEL, - "worker_stdout_parse", - "started", - json!({"expected_event_type":expected,"channel":channel_name,"line_bytes":n}), - ); - let value: Value = match serde_json::from_str(&line) { - Ok(value) => value, - Err(error) => { - emit_node_event( - datastream, - config, - NODE_WORKER_CHANNEL, - "worker_stdout_parse", - "failed", - json!({"expected_event_type":expected,"channel":channel_name,"line_bytes":n,"error":error.to_string()}), - ); - return Err(format!("parse helper stdout {line:?}: {error}")); - } - }; - let worker_event_type = value - .get("type") - .and_then(Value::as_str) - .unwrap_or("unknown"); - emit_node_event( - datastream, - config, - NODE_WORKER_CHANNEL, - "worker_stdout_parse", - "ready", - json!({"expected_event_type":expected,"channel":channel_name,"line_bytes":n,"worker_event_type":worker_event_type}), - ); - datastream.submit_text(channel, value.to_string()); - emit_stdio_datastream_frame(channel_name, &value) - .map_err(|e| format!("emit worker stdio datastream frame: {e}"))?; - datastream.tick(); - pump(); - if worker_event_type == "WorkerFatal" { - return Err(format!("worker fatal: {value}")); - } - if value.get("type").and_then(Value::as_str) == Some(expected) { - return Ok(value); - } - emit_node_event( - datastream, - config, - NODE_WORKER_CHANNEL, - "worker_event", - "observed", - json!({"command_waiting_for":expected,"worker_event_type":worker_event_type,"event":value}), - ); - } + wait_for_helper_event( + &self.stdout_rx, + Some(&self.stderr_rx), + expected, + command_type, + config, + datastream, + channel, + channel_name, + HelperCommandWaitConfig::production(), + pump, + ) } } @@ -4146,4 +4710,217 @@ mod tests { assert!(pending.backoff <= RUNTIME_READY_RETRY_MAX); } } + + fn fast_helper_wait_config() -> HelperCommandWaitConfig { + HelperCommandWaitConfig { + poll_interval: Duration::from_millis(5), + telemetry_interval: Duration::from_millis(10), + } + } + + fn drain_json_frames( + datastream: &NodeDatastream, + subscription: &DatastreamSubscription, + ) -> Vec<(String, Value)> { + datastream.endpoint.tick(); + subscription + .drain_available() + .into_iter() + .filter_map(|event| { + let DatastreamEvent::Frame(frame) = event else { + return None; + }; + let channel = datastream.by_id.get(&frame.channel.channel)?.clone(); + let value = serde_json::from_slice(&frame.payload).ok()?; + Some((channel, value)) + }) + .collect() + } + + #[test] + fn helper_wait_pumps_and_emits_busy_telemetry_during_quiet_stdout() { + let config = test_config(None); + let mut datastream = NodeDatastream::new(&config); + let channel = datastream.channel_by_name("mvp.worker.weights"); + let subscription = datastream.endpoint.subscribe_all("helper-wait-test"); + let (tx, rx) = mpsc::channel(); + thread::spawn(move || { + thread::sleep(Duration::from_millis(35)); + tx.send(HelperStdoutEvent::Line( + json!({"type":"WeightsLoaded","model_id":"fake"}).to_string() + "\n", + )) + .expect("send fake helper event"); + }); + let mut pump_count = 0_u32; + + let event = wait_for_helper_event( + &rx, + None, + "WeightsLoaded", + "LoadWeights", + &config, + &mut datastream, + channel, + "mvp.worker.weights", + fast_helper_wait_config(), + &mut || { + pump_count = pump_count.saturating_add(1); + }, + ) + .expect("quiet helper eventually returns expected event"); + + assert_eq!( + event.get("type").and_then(Value::as_str), + Some("WeightsLoaded") + ); + assert!(pump_count >= 2, "pump_count={pump_count}"); + let frames = drain_json_frames(&datastream, &subscription); + assert!(frames.iter().any(|(channel, value)| { + channel == "mvp.worker.weights" + && value.get("type").and_then(Value::as_str) == Some("WeightsLoaded") + })); + assert!(frames.iter().any(|(channel, value)| { + channel == NODE_WORKER_CHANNEL + && value.get("phase").and_then(Value::as_str) == Some("worker_command_wait") + && value.get("status").and_then(Value::as_str) == Some("waiting") + && value + .get("detail") + .and_then(|detail| detail.get("state")) + .and_then(Value::as_str) + == Some("busy_waiting_for_helper_stdout") + })); + } + + #[test] + fn helper_wait_errors_on_worker_fatal_event() { + let config = test_config(None); + let mut datastream = NodeDatastream::new(&config); + let channel = datastream.channel_by_name("mvp.worker.weights"); + let (tx, rx) = mpsc::channel(); + tx.send(HelperStdoutEvent::Line( + json!({"type":"WorkerFatal","error":"boom"}).to_string() + "\n", + )) + .expect("send fatal event"); + let mut pump_count = 0_u32; + + let error = wait_for_helper_event( + &rx, + None, + "WeightsLoaded", + "LoadWeights", + &config, + &mut datastream, + channel, + "mvp.worker.weights", + fast_helper_wait_config(), + &mut || { + pump_count = pump_count.saturating_add(1); + }, + ) + .expect_err("worker fatal must fail command wait"); + + assert!(error.contains("worker fatal")); + assert_eq!(pump_count, 1); + } + + #[test] + fn helper_wait_errors_on_closed_stdout_event() { + let config = test_config(None); + let mut datastream = NodeDatastream::new(&config); + let channel = datastream.channel_by_name("mvp.worker.weights"); + let (tx, rx) = mpsc::channel(); + tx.send(HelperStdoutEvent::Closed) + .expect("send closed stdout event"); + let mut pump_count = 0_u32; + + let error = wait_for_helper_event( + &rx, + None, + "WeightsLoaded", + "LoadWeights", + &config, + &mut datastream, + channel, + "mvp.worker.weights", + fast_helper_wait_config(), + &mut || { + pump_count = pump_count.saturating_add(1); + }, + ) + .expect_err("closed stdout must fail command wait"); + + assert!(error.contains("stdout closed")); + assert_eq!(pump_count, 0); + } + + #[test] + fn sampler_health_start_and_no_sample_records_are_json_events() { + let config = test_config(None); + let mut datastream = NodeDatastream::new(&config); + let health_channel = datastream.channel_by_name(NODE_SAMPLER_CHANNEL); + let subscription = datastream.endpoint.subscribe_all("sampler-health-test"); + let context = SamplerHealthContext::from_config(&config); + + for (sampler, sample_channel) in [ + ("gpu", datastream::hardware::gpu::HOST_GPU_CHANNEL), + ("cpu", datastream::hardware::cpu::HOST_CPU_CHANNEL), + ("net", datastream::hardware::net::HOST_NET_CHANNEL), + ] { + submit_sampler_started( + &datastream.producer, + health_channel, + context, + sampler, + sample_channel, + Duration::from_secs(1), + ); + } + submit_sampler_sample_health( + &datastream.producer, + health_channel, + context, + "gpu", + datastream::hardware::gpu::HOST_GPU_CHANNEL, + 0, + Some("nvidia-smi unavailable"), + ); + + let frames = drain_json_frames(&datastream, &subscription); + let sampler_events = frames + .iter() + .filter(|(channel, _)| channel == NODE_SAMPLER_CHANNEL) + .map(|(_, value)| value) + .collect::>(); + assert_eq!(sampler_events.len(), 7); + for sampler in ["gpu", "cpu", "net"] { + assert!(sampler_events.iter().any(|value| { + value.get("type").and_then(Value::as_str) == Some("SamplerHealth") + && value.get("schema").and_then(Value::as_str) + == Some("mvp.node.sampler.health.v1") + && value.get("run_id").and_then(Value::as_u64) == Some(7) + && value.get("node_id").and_then(Value::as_u64) == Some(11) + && value.get("stage_index").and_then(Value::as_u64) == Some(3) + && value.get("sampler").and_then(Value::as_str) == Some(sampler) + && value.get("status").and_then(Value::as_str) == Some("started") + })); + assert!(sampler_events.iter().any(|value| { + value.get("sampler").and_then(Value::as_str) == Some(sampler) + && value.get("status").and_then(Value::as_str) == Some("waiting") + && value + .get("detail") + .and_then(|detail| detail.get("state")) + .and_then(Value::as_str) + == Some("no_sample_yet") + })); + } + assert!(sampler_events.iter().any(|value| { + value.get("sampler").and_then(Value::as_str) == Some("gpu") + && value.get("status").and_then(Value::as_str) == Some("failed") + && value + .get("detail") + .and_then(|detail| detail.get("error")) + .and_then(Value::as_str) + == Some("nvidia-smi unavailable") + })); + } } diff --git a/crates/mvp-system/src/gguf_shard.rs b/crates/mvp-system/src/gguf_shard.rs new file mode 100644 index 0000000..b0c175f --- /dev/null +++ b/crates/mvp-system/src/gguf_shard.rs @@ -0,0 +1,1320 @@ +use std::fs::File; +use std::io::{Read, Seek, SeekFrom, Write}; +use std::path::Path; + +use serde::{Deserialize, Serialize}; + +use crate::run_plan::GgufSource; + +const GGUF_MAGIC: &[u8; 4] = b"GGUF"; +const SUPPORTED_GGUF_VERSION: u32 = 3; +const DEFAULT_ALIGNMENT: u64 = 32; +const MAX_STRING_BYTES: u64 = 64 * 1024 * 1024; + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] +pub struct ByteRange { + pub start: u64, + pub len: u64, +} + +impl ByteRange { + pub fn end_exclusive(self) -> Option { + self.start.checked_add(self.len) + } +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] +pub struct StageShardTensor { + pub name: String, + pub dims: Vec, + pub ggml_type: u32, + /// Tensor byte offset relative to the source GGUF data section. + pub source_offset: u64, + /// Tensor storage bytes in the source GGUF. Includes any source-side tensor padding + /// before the next tensor, which is safe to copy and keeps range math independent of + /// GGML quantization block-size tables. + pub byte_len: u64, +} + +#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] +pub struct StageShardPlan { + pub source: GgufSource, + pub stage_index: u32, + pub stage_count: u32, + pub layer_start: u32, + pub layer_end_exclusive: u32, + pub metadata_count: u64, + /// Offset immediately after the metadata KV section in the source file. + pub metadata_end: u64, + /// Offset of the source data section. + pub data_start: u64, + pub alignment: u32, + pub source_total_bytes: u64, + pub tensors: Vec, + /// Coalesced source-file byte ranges this stage must fetch from the origin. + /// Metadata/header bytes are intentionally separate: every stage fetches + /// `0..metadata_end` so it can write a valid stage-local GGUF header. + pub merged_tensor_ranges: Vec, + pub cache_key: String, +} + +impl StageShardPlan { + pub fn cache_file_name(&self) -> String { + format!("{}.stage-{:05}.gguf", self.cache_key, self.stage_index) + } + + pub fn source_url(&self) -> Result { + source_url(&self.source) + } + + pub fn planned_tensor_fetch_bytes(&self) -> u64 { + self.merged_tensor_ranges + .iter() + .map(|range| range.len) + .sum() + } + + pub fn planned_fetch_bytes(&self) -> u64 { + self.metadata_end + .saturating_add(self.planned_tensor_fetch_bytes()) + } + + pub fn planned_range_count(&self) -> usize { + self.merged_tensor_ranges.len() + if self.metadata_end > 0 { 1 } else { 0 } + } +} + +#[derive(Clone, Debug)] +struct GgufTensorEntry { + name: String, + dims: Vec, + ggml_type: u32, + source_offset: u64, + byte_len: u64, +} + +#[derive(Clone, Debug)] +struct GgufDirectory { + metadata_count: u64, + metadata_end: u64, + data_start: u64, + alignment: u32, + total_bytes: u64, + tensors: Vec, +} + +pub fn plan_stage_shard( + planning_gguf: &Path, + source: GgufSource, + stage_index: u32, + stage_count: u32, + layer_start: u32, + layer_end_exclusive: u32, +) -> Result { + if layer_start >= layer_end_exclusive { + return Err(format!( + "stage {stage_index} has empty layer range {layer_start}..{layer_end_exclusive}" + )); + } + if stage_count == 0 || stage_index >= stage_count { + return Err(format!( + "invalid stage index/count: stage {stage_index}, count {stage_count}" + )); + } + + let directory = read_gguf_directory(planning_gguf)?; + let tensors = select_stage_tensors( + &directory.tensors, + stage_index, + stage_count, + layer_start, + layer_end_exclusive, + )?; + let merged_tensor_ranges = merge_tensor_ranges(directory.data_start, &tensors)?; + let cache_key = shard_cache_key( + &source, + stage_index, + stage_count, + layer_start, + layer_end_exclusive, + &tensors, + ); + + Ok(StageShardPlan { + source, + stage_index, + stage_count, + layer_start, + layer_end_exclusive, + metadata_count: directory.metadata_count, + metadata_end: directory.metadata_end, + data_start: directory.data_start, + alignment: directory.alignment, + source_total_bytes: directory.total_bytes, + tensors: tensors + .into_iter() + .map(|tensor| StageShardTensor { + name: tensor.name, + dims: tensor.dims, + ggml_type: tensor.ggml_type, + source_offset: tensor.source_offset, + byte_len: tensor.byte_len, + }) + .collect(), + merged_tensor_ranges, + cache_key, + }) +} + +pub fn source_url(source: &GgufSource) -> Result { + match source { + GgufSource::HuggingFaceGguf { + repo, + file, + revision, + } => Ok(format!( + "https://huggingface.co/{repo}/resolve/{}/{}", + revision.as_deref().unwrap_or("main"), + encode_hf_path(file) + )), + GgufSource::LocalPath(path) => Err(format!( + "stage shard range fetching requires a remote Hugging Face source; got local path {path:?}" + )), + } +} + +fn encode_hf_path(path: &str) -> String { + path.split('/') + .map(percent_encode_path_segment) + .collect::>() + .join("/") +} + +fn percent_encode_path_segment(segment: &str) -> String { + let mut out = String::new(); + for byte in segment.bytes() { + if byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.' | b'~') { + out.push(char::from(byte)); + } else { + out.push_str(&format!("%{byte:02X}")); + } + } + out +} + +fn read_gguf_directory(path: &Path) -> Result { + let mut file = File::open(path).map_err(|e| format!("open GGUF {}: {e}", path.display()))?; + let total_bytes = file + .metadata() + .map_err(|e| format!("stat GGUF {}: {e}", path.display()))? + .len(); + let mut magic = [0; 4]; + file.read_exact(&mut magic) + .map_err(|e| format!("read GGUF magic: {e}"))?; + if &magic != GGUF_MAGIC { + return Err("invalid GGUF magic".to_owned()); + } + let version = read_u32(&mut file)?; + if version != SUPPORTED_GGUF_VERSION { + return Err(format!( + "unsupported GGUF version {version}; expected {SUPPORTED_GGUF_VERSION}" + )); + } + let tensor_count = read_u64(&mut file)?; + let metadata_count = read_u64(&mut file)?; + let mut alignment = DEFAULT_ALIGNMENT; + + for _ in 0..metadata_count { + let key = read_gguf_string(&mut file, MAX_STRING_BYTES)?; + let value_type = GgufValueType::read(&mut file)?; + if key == "general.alignment" && value_type.is_integer() { + alignment = read_integer_value(&mut file, value_type)?; + } else { + skip_value(&mut file, value_type)?; + } + } + let metadata_end = file + .stream_position() + .map_err(|e| format!("locate GGUF metadata end: {e}"))?; + + let mut tensor_infos = Vec::new(); + for _ in 0..tensor_count { + let name = read_gguf_string(&mut file, MAX_STRING_BYTES)?; + let dims_len = read_u32(&mut file)?; + let mut dims = Vec::with_capacity(dims_len as usize); + for _ in 0..dims_len { + dims.push(read_u64(&mut file)?); + } + let ggml_type = read_u32(&mut file)?; + let source_offset = read_u64(&mut file)?; + tensor_infos.push((name, dims, ggml_type, source_offset)); + } + let tensor_table_end = file + .stream_position() + .map_err(|e| format!("locate GGUF tensor table end: {e}"))?; + let data_start = align_to(tensor_table_end, alignment)?; + if data_start > total_bytes { + return Err(format!( + "GGUF data section starts at {data_start}, beyond file size {total_bytes}" + )); + } + + let mut order = tensor_infos + .iter() + .enumerate() + .map(|(index, (_, _, _, offset))| (*offset, index)) + .collect::>(); + order.sort_by_key(|(offset, _)| *offset); + let mut byte_lens = vec![0_u64; tensor_infos.len()]; + for (position, (offset, tensor_index)) in order.iter().copied().enumerate() { + let absolute = data_start + .checked_add(offset) + .ok_or_else(|| format!("tensor offset overflow at {offset}"))?; + if absolute > total_bytes { + return Err(format!( + "tensor offset {offset} points beyond GGUF data size in {}", + path.display() + )); + } + let next_absolute = if let Some((next_offset, _)) = order.get(position + 1) { + data_start + .checked_add(*next_offset) + .ok_or_else(|| format!("next tensor offset overflow at {next_offset}"))? + } else { + total_bytes + }; + if next_absolute < absolute { + return Err("GGUF tensor offsets are not monotonic".to_owned()); + } + byte_lens[tensor_index] = next_absolute - absolute; + } + + let tensors = tensor_infos + .into_iter() + .enumerate() + .map( + |(index, (name, dims, ggml_type, source_offset))| GgufTensorEntry { + name, + dims, + ggml_type, + source_offset, + byte_len: byte_lens[index], + }, + ) + .collect(); + + Ok(GgufDirectory { + metadata_count, + metadata_end, + data_start, + alignment: u32::try_from(alignment) + .map_err(|_| format!("GGUF alignment {alignment} exceeds u32"))?, + total_bytes, + tensors, + }) +} + +fn select_stage_tensors( + tensors: &[GgufTensorEntry], + stage_index: u32, + stage_count: u32, + layer_start: u32, + layer_end_exclusive: u32, +) -> Result, String> { + let first_stage = layer_start == 0; + let final_stage = stage_index + 1 == stage_count; + let has_output_weight = tensors.iter().any(|tensor| tensor.name == "output.weight"); + let mut selected = Vec::new(); + for tensor in tensors { + if tensor.name == "token_embd.weight" && (first_stage || final_stage && !has_output_weight) + { + selected.push(tensor.clone()); + continue; + } + if final_stage && matches!(tensor.name.as_str(), "output.weight" | "output_norm.weight") { + selected.push(tensor.clone()); + continue; + } + if let Some(layer) = tensor_layer_index(&tensor.name) + && layer_start <= layer + && layer < layer_end_exclusive + { + selected.push(tensor.clone()); + } + } + if selected.is_empty() { + return Err(format!( + "stage {stage_index} selected no tensors for layer range {layer_start}..{layer_end_exclusive}" + )); + } + selected.sort_by_key(|tensor| tensor.source_offset); + Ok(selected) +} + +fn tensor_layer_index(name: &str) -> Option { + let rest = name.strip_prefix("blk.")?; + let (raw, _) = rest.split_once('.')?; + raw.parse().ok() +} + +fn merge_tensor_ranges( + data_start: u64, + tensors: &[GgufTensorEntry], +) -> Result, String> { + let mut ranges = tensors + .iter() + .map(|tensor| { + let start = data_start + .checked_add(tensor.source_offset) + .ok_or_else(|| format!("range start overflow for tensor {}", tensor.name))?; + Ok(ByteRange { + start, + len: tensor.byte_len, + }) + }) + .collect::, String>>()?; + ranges.sort_by_key(|range| range.start); + let mut merged: Vec = Vec::new(); + for range in ranges { + if range.len == 0 { + continue; + } + let range_end = range + .end_exclusive() + .ok_or_else(|| format!("range end overflow at {}", range.start))?; + if let Some(last) = merged.last_mut() { + let last_end = last + .end_exclusive() + .ok_or_else(|| format!("range end overflow at {}", last.start))?; + if range.start <= last_end { + last.len = range_end.saturating_sub(last.start).max(last.len); + continue; + } + } + merged.push(range); + } + Ok(merged) +} + +fn shard_cache_key( + source: &GgufSource, + stage_index: u32, + stage_count: u32, + layer_start: u32, + layer_end_exclusive: u32, + tensors: &[GgufTensorEntry], +) -> String { + let mut hasher = blake3::Hasher::new(); + hasher.update(format!("source:{source:?}\n").as_bytes()); + hasher.update( + format!("stage:{stage_index}/{stage_count}:{layer_start}-{layer_end_exclusive}\n") + .as_bytes(), + ); + for tensor in tensors { + hasher.update( + format!( + "{}:{}:{}:{:?}\n", + tensor.name, tensor.source_offset, tensor.byte_len, tensor.dims + ) + .as_bytes(), + ); + } + hasher.finalize().to_hex()[..24].to_owned() +} + +pub fn materialize_stage_shard_http( + plan: &StageShardPlan, + output_path: &Path, + emit: F, +) -> Result<(), String> +where + F: FnMut(serde_json::Value), +{ + let url = plan.source_url()?; + materialize_stage_shard_from_url(plan, &url, output_path, emit) +} + +pub fn materialize_stage_shard_from_url( + plan: &StageShardPlan, + url: &str, + output_path: &Path, + mut emit: F, +) -> Result<(), String> +where + F: FnMut(serde_json::Value), +{ + let bytes_total = plan.planned_fetch_bytes(); + let range_count = plan.planned_range_count(); + emit(serde_json::json!({ + "type":"StageShardFetchStarted", + "stage_index":plan.stage_index, + "url":url, + "tensor_count":plan.tensors.len(), + "range_count":range_count, + "tensor_range_count":plan.merged_tensor_ranges.len(), + "source_total_bytes":plan.source_total_bytes, + "bytes_done":0_u64, + "bytes_total":bytes_total, + "output_path":output_path, + })); + let mut bytes_done = 0_u64; + let mut next_range_index = 0_usize; + if plan.metadata_end > 0 { + emit(serde_json::json!({ + "type":"StageShardRangeFetchStarted", + "stage_index":plan.stage_index, + "range_index":next_range_index, + "range_count":range_count, + "range_kind":"metadata", + "source_start":0_u64, + "bytes":plan.metadata_end, + "bytes_done":bytes_done, + "bytes_total":bytes_total, + })); + } + let metadata_prefix = fetch_http_range(&url, 0, plan.metadata_end)?; + bytes_done = bytes_done.saturating_add(plan.metadata_end); + if plan.metadata_end > 0 { + emit(serde_json::json!({ + "type":"StageShardRangeFetchReady", + "stage_index":plan.stage_index, + "range_index":next_range_index, + "range_count":range_count, + "range_kind":"metadata", + "source_start":0_u64, + "bytes":plan.metadata_end, + "bytes_done":bytes_done, + "bytes_total":bytes_total, + })); + next_range_index += 1; + } + if metadata_prefix.len() < 24 { + return Err(format!( + "GGUF metadata prefix too short: {} bytes", + metadata_prefix.len() + )); + } + let metadata_body = &metadata_prefix[24..]; + let alignment = u64::from(plan.alignment.max(1)); + let mut data_offsets = Vec::with_capacity(plan.tensors.len()); + let mut data_cursor = 0_u64; + for tensor in &plan.tensors { + data_cursor = align_to(data_cursor, alignment)?; + data_offsets.push(data_cursor); + data_cursor = data_cursor + .checked_add(tensor.byte_len) + .ok_or_else(|| format!("stage shard data size overflow at tensor {}", tensor.name))?; + } + + let partial_path = output_path.with_extension(format!( + "{}partial", + output_path + .extension() + .and_then(|value| value.to_str()) + .map(|ext| format!("{ext}.")) + .unwrap_or_default() + )); + if let Some(parent) = output_path.parent() { + std::fs::create_dir_all(parent) + .map_err(|e| format!("create stage shard cache dir {}: {e}", parent.display()))?; + } + let mut out = File::create(&partial_path) + .map_err(|e| format!("create stage shard {}: {e}", partial_path.display()))?; + out.write_all(GGUF_MAGIC) + .map_err(|e| format!("write stage shard magic: {e}"))?; + out.write_all(&SUPPORTED_GGUF_VERSION.to_le_bytes()) + .map_err(|e| format!("write stage shard version: {e}"))?; + out.write_all(&(plan.tensors.len() as u64).to_le_bytes()) + .map_err(|e| format!("write stage shard tensor count: {e}"))?; + out.write_all(&plan.metadata_count.to_le_bytes()) + .map_err(|e| format!("write stage shard metadata count: {e}"))?; + out.write_all(metadata_body) + .map_err(|e| format!("write stage shard metadata: {e}"))?; + for (tensor, data_offset) in plan.tensors.iter().zip(data_offsets.iter().copied()) { + write_gguf_string(&mut out, &tensor.name)?; + out.write_all(&(tensor.dims.len() as u32).to_le_bytes()) + .map_err(|e| format!("write tensor dim count for {}: {e}", tensor.name))?; + for dim in &tensor.dims { + out.write_all(&dim.to_le_bytes()) + .map_err(|e| format!("write tensor dim for {}: {e}", tensor.name))?; + } + out.write_all(&tensor.ggml_type.to_le_bytes()) + .map_err(|e| format!("write tensor type for {}: {e}", tensor.name))?; + out.write_all(&data_offset.to_le_bytes()) + .map_err(|e| format!("write tensor offset for {}: {e}", tensor.name))?; + } + pad_writer_to_alignment(&mut out, alignment)?; + let mut written_data = 0_u64; + let mut tensor_index = 0_usize; + for range in &plan.merged_tensor_ranges { + let range_index = next_range_index; + emit(serde_json::json!({ + "type":"StageShardRangeFetchStarted", + "stage_index":plan.stage_index, + "range_index":range_index, + "range_count":range_count, + "range_kind":"tensor_data", + "source_start":range.start, + "bytes":range.len, + "bytes_done":bytes_done, + "bytes_total":bytes_total, + })); + let range_bytes = fetch_http_range(&url, range.start, range.len)?; + bytes_done = bytes_done.saturating_add(range.len); + emit(serde_json::json!({ + "type":"StageShardRangeFetchReady", + "stage_index":plan.stage_index, + "range_index":range_index, + "range_count":range_count, + "range_kind":"tensor_data", + "source_start":range.start, + "bytes":range.len, + "bytes_done":bytes_done, + "bytes_total":bytes_total, + })); + let range_end = range + .end_exclusive() + .ok_or_else(|| format!("stage shard range end overflow at {}", range.start))?; + while let Some(tensor) = plan.tensors.get(tensor_index) { + let source_start = plan + .data_start + .checked_add(tensor.source_offset) + .ok_or_else(|| format!("source range overflow for {}", tensor.name))?; + if source_start >= range_end { + break; + } + let source_end = source_start + .checked_add(tensor.byte_len) + .ok_or_else(|| format!("source range overflow for {}", tensor.name))?; + if source_start < range.start || source_end > range_end { + return Err(format!( + "tensor {} source range {source_start}..{source_end} is not covered by planned range {}..{range_end}", + tensor.name, range.start + )); + } + let target_offset = data_offsets + .get(tensor_index) + .copied() + .ok_or_else(|| format!("missing target offset for {}", tensor.name))?; + while written_data < target_offset { + out.write_all(&[0]) + .map_err(|e| format!("pad tensor data before {}: {e}", tensor.name))?; + written_data += 1; + } + emit(serde_json::json!({ + "type":"StageShardTensorFetchStarted", + "stage_index":plan.stage_index, + "range_index":range_index, + "range_count":range_count, + "tensor_index":tensor_index, + "tensor_count":plan.tensors.len(), + "tensor":tensor.name, + "source_start":source_start, + "bytes":tensor.byte_len, + "bytes_done":bytes_done, + "bytes_total":bytes_total, + })); + let offset = usize::try_from(source_start - range.start) + .map_err(|_| format!("tensor {} range offset exceeds usize", tensor.name))?; + let len = usize::try_from(tensor.byte_len) + .map_err(|_| format!("tensor {} byte length exceeds usize", tensor.name))?; + let end = offset + .checked_add(len) + .ok_or_else(|| format!("tensor {} range slice overflows", tensor.name))?; + out.write_all(&range_bytes[offset..end]) + .map_err(|e| format!("write tensor {} bytes: {e}", tensor.name))?; + written_data = written_data + .checked_add(tensor.byte_len) + .ok_or_else(|| format!("written data overflow after {}", tensor.name))?; + emit(serde_json::json!({ + "type":"StageShardTensorFetchReady", + "stage_index":plan.stage_index, + "range_index":range_index, + "range_count":range_count, + "tensor_index":tensor_index, + "tensor_count":plan.tensors.len(), + "tensor":tensor.name, + "bytes":tensor.byte_len, + "bytes_done":bytes_done, + "bytes_total":bytes_total, + })); + tensor_index += 1; + } + next_range_index += 1; + } + if tensor_index != plan.tensors.len() { + return Err(format!( + "planned ranges covered {tensor_index} of {} stage tensors", + plan.tensors.len() + )); + } + out.flush() + .map_err(|e| format!("flush stage shard {}: {e}", partial_path.display()))?; + drop(out); + std::fs::rename(&partial_path, output_path).map_err(|e| { + format!( + "commit stage shard {} -> {}: {e}", + partial_path.display(), + output_path.display() + ) + })?; + emit(serde_json::json!({ + "type":"StageShardReady", + "stage_index":plan.stage_index, + "path":output_path, + "range_count":range_count, + "tensor_count":plan.tensors.len(), + "bytes":std::fs::metadata(output_path).map(|metadata| metadata.len()).unwrap_or(0), + "bytes_done":bytes_done, + "bytes_total":bytes_total, + })); + Ok(()) +} + +fn fetch_http_range(url: &str, start: u64, len: u64) -> Result, String> { + if len == 0 { + return Ok(Vec::new()); + } + let end = start + .checked_add(len - 1) + .ok_or_else(|| format!("HTTP range overflow at {start}+{len}"))?; + let range = format!("bytes={start}-{end}"); + let mut request = ureq::get(url) + .set("Range", &range) + .set("User-Agent", "swactor-mvp-node/0.1"); + if let Ok(token) = std::env::var("HF_TOKEN") + && !token.trim().is_empty() + { + request = request.set("Authorization", &format!("Bearer {}", token.trim())); + } + let response = request + .call() + .map_err(|error| format!("GET {url} {range}: {error}"))?; + if response.status() != 206 { + return Err(format!( + "GET {url} {range} returned HTTP {}; refusing full-body fallback", + response.status() + )); + } + let mut reader = response.into_reader(); + let capacity = usize::try_from(len).unwrap_or(usize::MAX.min(64 * 1024 * 1024)); + let mut bytes = Vec::with_capacity(capacity.min(64 * 1024 * 1024)); + reader + .read_to_end(&mut bytes) + .map_err(|e| format!("read {url} {range}: {e}"))?; + if bytes.len() as u64 != len { + return Err(format!( + "GET {url} {range} returned {} bytes, expected {len}", + bytes.len() + )); + } + Ok(bytes) +} + +fn write_gguf_string(writer: &mut W, value: &str) -> Result<(), String> { + writer + .write_all(&(value.len() as u64).to_le_bytes()) + .map_err(|e| format!("write GGUF string len: {e}"))?; + writer + .write_all(value.as_bytes()) + .map_err(|e| format!("write GGUF string bytes: {e}")) +} + +fn pad_writer_to_alignment(writer: &mut W, alignment: u64) -> Result<(), String> { + let pos = writer + .stream_position() + .map_err(|e| format!("locate writer for alignment: {e}"))?; + let aligned = align_to(pos, alignment)?; + for _ in pos..aligned { + writer + .write_all(&[0]) + .map_err(|e| format!("write alignment padding: {e}"))?; + } + Ok(()) +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum GgufValueType { + Uint8, + Int8, + Uint16, + Int16, + Uint32, + Int32, + Float32, + Bool, + String, + Array, + Uint64, + Int64, + Float64, +} + +impl GgufValueType { + fn read(reader: &mut R) -> Result { + match read_u32(reader)? { + 0 => Ok(Self::Uint8), + 1 => Ok(Self::Int8), + 2 => Ok(Self::Uint16), + 3 => Ok(Self::Int16), + 4 => Ok(Self::Uint32), + 5 => Ok(Self::Int32), + 6 => Ok(Self::Float32), + 7 => Ok(Self::Bool), + 8 => Ok(Self::String), + 9 => Ok(Self::Array), + 10 => Ok(Self::Uint64), + 11 => Ok(Self::Int64), + 12 => Ok(Self::Float64), + other => Err(format!("unsupported GGUF value type {other}")), + } + } + + fn fixed_width(self) -> Option { + match self { + Self::Uint8 | Self::Int8 | Self::Bool => Some(1), + Self::Uint16 | Self::Int16 => Some(2), + Self::Uint32 | Self::Int32 | Self::Float32 => Some(4), + Self::Uint64 | Self::Int64 | Self::Float64 => Some(8), + Self::String | Self::Array => None, + } + } + + fn is_integer(self) -> bool { + matches!( + self, + Self::Uint8 + | Self::Int8 + | Self::Uint16 + | Self::Int16 + | Self::Uint32 + | Self::Int32 + | Self::Uint64 + | Self::Int64 + ) + } +} + +fn skip_value(reader: &mut R, value_type: GgufValueType) -> Result<(), String> { + match value_type { + GgufValueType::String => skip_gguf_string(reader), + GgufValueType::Array => skip_array(reader), + scalar => skip_bytes(reader, scalar.fixed_width().expect("scalar width")), + } +} + +fn skip_array(reader: &mut R) -> Result<(), String> { + let element_type = GgufValueType::read(reader)?; + let len = read_u64(reader)?; + match element_type { + GgufValueType::String => { + for _ in 0..len { + skip_gguf_string(reader)?; + } + Ok(()) + } + GgufValueType::Array => { + for _ in 0..len { + skip_array(reader)?; + } + Ok(()) + } + scalar => { + let width = scalar.fixed_width().expect("scalar array width"); + let bytes = width + .checked_mul(len) + .ok_or_else(|| "GGUF array byte count overflow".to_owned())?; + skip_bytes(reader, bytes) + } + } +} + +fn read_integer_value(reader: &mut R, value_type: GgufValueType) -> Result { + match value_type { + GgufValueType::Uint8 => read_u8(reader).map(u64::from), + GgufValueType::Int8 => read_i8(reader).and_then(non_negative_i64_to_u64), + GgufValueType::Uint16 => read_u16(reader).map(u64::from), + GgufValueType::Int16 => { + read_i16(reader).and_then(|v| non_negative_i64_to_u64(i64::from(v))) + } + GgufValueType::Uint32 => read_u32(reader).map(u64::from), + GgufValueType::Int32 => { + read_i32(reader).and_then(|v| non_negative_i64_to_u64(i64::from(v))) + } + GgufValueType::Uint64 => read_u64(reader), + GgufValueType::Int64 => read_i64(reader).and_then(non_negative_i64_to_u64), + other => Err(format!("GGUF value type {other:?} is not integer")), + } +} + +fn non_negative_i64_to_u64(value: i64) -> Result { + u64::try_from(value).map_err(|_| format!("negative GGUF integer {value}")) +} + +fn read_gguf_string(reader: &mut R, max_len: u64) -> Result { + let len = read_u64(reader)?; + if len > max_len { + return Err(format!("GGUF string length {len} exceeds {max_len}")); + } + let len = usize::try_from(len).map_err(|_| "GGUF string length exceeds usize".to_owned())?; + let mut bytes = vec![0_u8; len]; + reader + .read_exact(&mut bytes) + .map_err(|e| format!("read GGUF string: {e}"))?; + String::from_utf8(bytes).map_err(|e| format!("GGUF string is not UTF-8: {e}")) +} + +fn skip_gguf_string(reader: &mut R) -> Result<(), String> { + let len = read_u64(reader)?; + skip_bytes(reader, len) +} + +fn skip_bytes(reader: &mut R, bytes: u64) -> Result<(), String> { + let offset = i64::try_from(bytes).map_err(|_| format!("cannot seek over {bytes} bytes"))?; + reader + .seek(SeekFrom::Current(offset)) + .map_err(|e| format!("skip bytes: {e}"))?; + Ok(()) +} + +fn align_to(value: u64, alignment: u64) -> Result { + if alignment == 0 { + return Err("GGUF alignment must be non-zero".to_owned()); + } + let remainder = value % alignment; + if remainder == 0 { + Ok(value) + } else { + value + .checked_add(alignment - remainder) + .ok_or_else(|| format!("align {value} to {alignment} overflows")) + } +} + +fn read_u8(reader: &mut R) -> Result { + let mut bytes = [0; 1]; + reader + .read_exact(&mut bytes) + .map_err(|e| format!("read u8: {e}"))?; + Ok(bytes[0]) +} + +fn read_i8(reader: &mut R) -> Result { + read_u8(reader).map(|value| i8::from_le_bytes([value]) as i64) +} + +fn read_u16(reader: &mut R) -> Result { + let mut bytes = [0; 2]; + reader + .read_exact(&mut bytes) + .map_err(|e| format!("read u16: {e}"))?; + Ok(u16::from_le_bytes(bytes)) +} + +fn read_i16(reader: &mut R) -> Result { + let mut bytes = [0; 2]; + reader + .read_exact(&mut bytes) + .map_err(|e| format!("read i16: {e}"))?; + Ok(i16::from_le_bytes(bytes)) +} + +fn read_u32(reader: &mut R) -> Result { + let mut bytes = [0; 4]; + reader + .read_exact(&mut bytes) + .map_err(|e| format!("read u32: {e}"))?; + Ok(u32::from_le_bytes(bytes)) +} + +fn read_i32(reader: &mut R) -> Result { + let mut bytes = [0; 4]; + reader + .read_exact(&mut bytes) + .map_err(|e| format!("read i32: {e}"))?; + Ok(i32::from_le_bytes(bytes)) +} + +fn read_u64(reader: &mut R) -> Result { + let mut bytes = [0; 8]; + reader + .read_exact(&mut bytes) + .map_err(|e| format!("read u64: {e}"))?; + Ok(u64::from_le_bytes(bytes)) +} + +fn read_i64(reader: &mut R) -> Result { + let mut bytes = [0; 8]; + reader + .read_exact(&mut bytes) + .map_err(|e| format!("read i64: {e}"))?; + Ok(i64::from_le_bytes(bytes)) +} + +#[cfg(test)] +mod tests { + use super::*; + use parking_lot::Mutex; + use std::io::{Read, Write}; + use std::net::{TcpListener, TcpStream}; + use std::sync::Arc; + use std::thread; + + #[test] + fn stage_plans_select_only_assigned_layer_tensors_and_boundaries() { + let fixture = SyntheticGguf::new(8); + let source = GgufSource::HuggingFaceGguf { + repo: "org/repo".to_owned(), + file: "model file.gguf".to_owned(), + revision: Some("abc123".to_owned()), + }; + + let stage0 = plan_stage_shard(&fixture.path, source.clone(), 0, 4, 0, 2).unwrap(); + let stage1 = plan_stage_shard(&fixture.path, source.clone(), 1, 4, 2, 4).unwrap(); + let stage3 = plan_stage_shard(&fixture.path, source.clone(), 3, 4, 6, 8).unwrap(); + + assert_eq!( + names(&stage0), + vec![ + "token_embd.weight", + "blk.0.attn_q.weight", + "blk.0.ffn_up.weight", + "blk.1.attn_q.weight", + "blk.1.ffn_up.weight", + ] + ); + assert_eq!( + names(&stage1), + vec![ + "blk.2.attn_q.weight", + "blk.2.ffn_up.weight", + "blk.3.attn_q.weight", + "blk.3.ffn_up.weight", + ] + ); + assert_eq!( + names(&stage3), + vec![ + "blk.6.attn_q.weight", + "blk.6.ffn_up.weight", + "blk.7.attn_q.weight", + "blk.7.ffn_up.weight", + "output_norm.weight", + "output.weight", + ] + ); + assert_eq!( + stage0.source_url().unwrap(), + "https://huggingface.co/org/repo/resolve/abc123/model%20file.gguf" + ); + assert!( + stage1 + .merged_tensor_ranges + .iter() + .all(|range| range.start >= stage1.data_start) + ); + assert!( + stage1 + .merged_tensor_ranges + .iter() + .map(|range| range.len) + .sum::() + < fixture.bytes_len + ); + } + + #[test] + fn tied_output_final_stage_includes_token_embedding_when_output_weight_is_missing() { + let fixture = SyntheticGguf::without_output_weight(4); + let source = GgufSource::HuggingFaceGguf { + repo: "org/repo".to_owned(), + file: "model.gguf".to_owned(), + revision: None, + }; + + let final_stage = plan_stage_shard(&fixture.path, source, 1, 2, 2, 4).unwrap(); + + assert!(names(&final_stage).contains(&"token_embd.weight")); + assert!(names(&final_stage).contains(&"output_norm.weight")); + assert!(!names(&final_stage).contains(&"output.weight")); + } + + #[test] + fn materialized_stage_shard_fetches_only_http_ranges_and_loads_as_gguf() { + let fixture = SyntheticGguf::new(4); + let bytes = std::fs::read(&fixture.path).unwrap(); + let server = RangeServer::start(bytes); + let source = GgufSource::HuggingFaceGguf { + repo: "org/repo".to_owned(), + file: "model.gguf".to_owned(), + revision: None, + }; + let plan = plan_stage_shard(&fixture.path, source, 1, 2, 2, 4).unwrap(); + let output_path = fixture.path.with_file_name("stage-1.gguf"); + let mut events = Vec::new(); + + materialize_stage_shard_from_url(&plan, &server.url, &output_path, |event| { + events.push(event); + }) + .unwrap(); + + let materialized = read_gguf_directory(&output_path).unwrap(); + assert_eq!(materialized.tensors.len(), plan.tensors.len()); + assert_eq!( + materialized + .tensors + .iter() + .map(|tensor| tensor.name.as_str()) + .collect::>(), + names(&plan) + ); + let ready = events.iter().find(|event| { + event.get("type").and_then(serde_json::Value::as_str) == Some("StageShardReady") + }); + assert_eq!( + ready + .and_then(|event| event.get("bytes_done")) + .and_then(serde_json::Value::as_u64), + Some(plan.planned_fetch_bytes()) + ); + assert_eq!( + ready + .and_then(|event| event.get("bytes_total")) + .and_then(serde_json::Value::as_u64), + Some(plan.planned_fetch_bytes()) + ); + let ranges = server.ranges(); + let expected_ranges = std::iter::once((0, plan.metadata_end.saturating_sub(1))) + .chain(plan.merged_tensor_ranges.iter().map(|range| { + ( + range.start, + range + .end_exclusive() + .expect("planned range end") + .saturating_sub(1), + ) + })) + .collect::>(); + assert_eq!(ranges, expected_ranges); + assert!( + ranges.len() < plan.tensors.len() + 1, + "coalesced tensor ranges should replace one HTTP request per tensor" + ); + let range_ready = events + .iter() + .filter(|event| { + event.get("type").and_then(serde_json::Value::as_str) + == Some("StageShardRangeFetchReady") + }) + .collect::>(); + assert_eq!(range_ready.len(), plan.planned_range_count()); + assert_eq!( + range_ready + .last() + .and_then(|event| event.get("bytes_done")) + .and_then(serde_json::Value::as_u64), + Some(plan.planned_fetch_bytes()) + ); + assert_eq!( + range_ready + .last() + .and_then(|event| event.get("range_index")) + .and_then(serde_json::Value::as_u64), + Some((plan.planned_range_count() - 1) as u64) + ); + let total_requested = ranges + .iter() + .map(|(start, end)| end.saturating_sub(*start) + 1) + .sum::(); + assert_eq!(total_requested, plan.planned_fetch_bytes()); + assert!(total_requested < fixture.bytes_len); + } + + fn names(plan: &StageShardPlan) -> Vec<&str> { + plan.tensors + .iter() + .map(|tensor| tensor.name.as_str()) + .collect() + } + + struct RangeServer { + url: String, + ranges: Arc>>, + } + + impl RangeServer { + fn start(bytes: Vec) -> Self { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let url = format!("http://{}", listener.local_addr().unwrap()); + let bytes = Arc::new(bytes); + let ranges = Arc::new(Mutex::new(Vec::new())); + let server_ranges = Arc::clone(&ranges); + thread::spawn(move || { + for stream in listener.incoming().flatten() { + handle_range_request(stream, Arc::clone(&bytes), Arc::clone(&server_ranges)); + } + }); + Self { url, ranges } + } + + fn ranges(&self) -> Vec<(u64, u64)> { + self.ranges.lock().clone() + } + } + + fn handle_range_request( + mut stream: TcpStream, + bytes: Arc>, + ranges: Arc>>, + ) { + let mut request = Vec::new(); + let mut buffer = [0_u8; 1024]; + while !request.windows(4).any(|window| window == b"\r\n\r\n") { + let Ok(read) = stream.read(&mut buffer) else { + return; + }; + if read == 0 { + return; + } + request.extend_from_slice(&buffer[..read]); + } + let request = String::from_utf8_lossy(&request); + let Some(range_header) = request + .lines() + .find(|line| line.to_ascii_lowercase().starts_with("range: bytes=")) + else { + let body = b"missing range"; + let _ = write!( + stream, + "HTTP/1.1 400 Bad Request\r\nContent-Length: {}\r\n\r\n", + body.len() + ); + let _ = stream.write_all(body); + return; + }; + let range = range_header + .split_once("bytes=") + .map(|(_, range)| range.trim()) + .unwrap(); + let (start, end) = range.split_once('-').unwrap(); + let start = start.parse::().unwrap(); + let end = end.parse::().unwrap(); + let content = &bytes[start as usize..=end as usize]; + ranges.lock().push((start, end)); + let _ = write!( + stream, + "HTTP/1.1 206 Partial Content\r\nContent-Length: {}\r\nContent-Range: bytes {}-{}/{}\r\nAccept-Ranges: bytes\r\n\r\n", + content.len(), + start, + end, + bytes.len() + ); + let _ = stream.write_all(content); + } + + struct SyntheticGguf { + path: std::path::PathBuf, + bytes_len: u64, + _dir: std::path::PathBuf, + } + + impl SyntheticGguf { + fn new(layers: u32) -> Self { + Self::build(layers, true) + } + + fn without_output_weight(layers: u32) -> Self { + Self::build(layers, false) + } + + fn build(layers: u32, output_weight: bool) -> Self { + let dir = std::env::temp_dir().join(format!( + "mvp-gguf-shard-test-{}-{:?}", + std::process::id(), + std::thread::current().id() + )); + let _ = std::fs::remove_dir_all(&dir); + std::fs::create_dir_all(&dir).unwrap(); + let path = dir.join(if output_weight { + "model.gguf" + } else { + "tied.gguf" + }); + let bytes = synthetic_gguf(layers, output_weight); + std::fs::File::create(&path) + .unwrap() + .write_all(&bytes) + .unwrap(); + Self { + path, + bytes_len: bytes.len() as u64, + _dir: dir, + } + } + } + + fn synthetic_gguf(layers: u32, output_weight: bool) -> Vec { + let mut tensors = vec!["token_embd.weight".to_owned()]; + for layer in 0..layers { + tensors.push(format!("blk.{layer}.attn_q.weight")); + tensors.push(format!("blk.{layer}.ffn_up.weight")); + } + tensors.push("output_norm.weight".to_owned()); + if output_weight { + tensors.push("output.weight".to_owned()); + } + + let mut metadata = Vec::new(); + push_string_kv(&mut metadata, "general.architecture", "llama"); + push_u32_kv(&mut metadata, "general.alignment", 32); + push_u32_kv(&mut metadata, "llama.block_count", layers); + push_u32_kv(&mut metadata, "llama.embedding_length", 8); + push_u32_kv(&mut metadata, "llama.context_length", 16); + push_u32_kv(&mut metadata, "tokenizer.ggml.eos_token_id", 2); + + let mut tensor_infos = Vec::new(); + let mut data = Vec::new(); + let mut offset = 0_u64; + for (index, name) in tensors.iter().enumerate() { + push_gguf_string(&mut tensor_infos, name); + tensor_infos.extend_from_slice(&2_u32.to_le_bytes()); + tensor_infos.extend_from_slice(&2_u64.to_le_bytes()); + tensor_infos.extend_from_slice(&2_u64.to_le_bytes()); + tensor_infos.extend_from_slice(&0_u32.to_le_bytes()); + tensor_infos.extend_from_slice(&offset.to_le_bytes()); + let len = 16 + index as u64; + data.extend(std::iter::repeat_n(index as u8, len as usize)); + offset += len; + } + + let mut bytes = Vec::new(); + bytes.extend_from_slice(b"GGUF"); + bytes.extend_from_slice(&3_u32.to_le_bytes()); + bytes.extend_from_slice(&(tensors.len() as u64).to_le_bytes()); + bytes.extend_from_slice(&6_u64.to_le_bytes()); + bytes.extend_from_slice(&metadata); + bytes.extend_from_slice(&tensor_infos); + while bytes.len() % 32 != 0 { + bytes.push(0); + } + bytes.extend_from_slice(&data); + bytes + } + + fn push_string_kv(bytes: &mut Vec, key: &str, value: &str) { + push_gguf_string(bytes, key); + bytes.extend_from_slice(&8_u32.to_le_bytes()); + push_gguf_string(bytes, value); + } + + fn push_u32_kv(bytes: &mut Vec, key: &str, value: u32) { + push_gguf_string(bytes, key); + bytes.extend_from_slice(&4_u32.to_le_bytes()); + bytes.extend_from_slice(&value.to_le_bytes()); + } + + fn push_gguf_string(bytes: &mut Vec, value: &str) { + bytes.extend_from_slice(&(value.len() as u64).to_le_bytes()); + bytes.extend_from_slice(value.as_bytes()); + } +} diff --git a/crates/mvp-system/src/lib.rs b/crates/mvp-system/src/lib.rs index 94269c7..1cfb88f 100644 --- a/crates/mvp-system/src/lib.rs +++ b/crates/mvp-system/src/lib.rs @@ -15,6 +15,7 @@ pub mod edge_establisher; pub mod endpoint_advertisement; pub mod engine_builder; pub mod gguf_metadata; +pub mod gguf_shard; pub mod gpu_worker_ctl; pub mod gpu_worker_egress_producer; pub mod gpu_worker_ingress_parser; diff --git a/crates/mvp-system/src/node_image.rs b/crates/mvp-system/src/node_image.rs index d708d3c..379fb17 100644 --- a/crates/mvp-system/src/node_image.rs +++ b/crates/mvp-system/src/node_image.rs @@ -1,8 +1,11 @@ use std::collections::{BTreeMap, BTreeSet}; use std::fs::{self, File}; -use std::io::Read; +use std::io::{BufRead, BufReader, Read}; use std::path::{Path, PathBuf}; use std::process::{Command, Stdio}; +use std::sync::mpsc; +use std::thread; +use std::time::{Duration, Instant}; const IMAGE_SOURCE_INPUTS: &[&str] = &[ "Cargo.lock", @@ -72,7 +75,103 @@ pub struct PreparedNodeImage { pub pushed: bool, } +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum NodeImageProgressEventKind { + ImageReference { + role: String, + image_ref: String, + }, + CommandStarted { + program: String, + args: Vec, + }, + CommandStdout { + line: String, + }, + CommandStderr { + line: String, + }, + CommandExited { + status: String, + code: Option, + success: bool, + }, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct NodeImageProgressEvent { + pub command_label: Option, + pub image_ref: Option, + pub elapsed_ms: Option, + pub kind: NodeImageProgressEventKind, +} + +pub trait NodeImageProgressSink { + fn emit(&mut self, event: NodeImageProgressEvent); +} + +impl NodeImageProgressSink for F +where + F: FnMut(NodeImageProgressEvent), +{ + fn emit(&mut self, event: NodeImageProgressEvent) { + self(event); + } +} + +trait ImageCommandRunner { + fn run_status( + &mut self, + root: &Path, + program: &str, + args: &[String], + label: &str, + image_ref: Option<&str>, + progress: &mut Option<&mut dyn NodeImageProgressSink>, + ) -> Result<(), String>; + + fn docker_image_exists(&mut self, root: &Path, image_ref: &str) -> bool; + + fn docker_image_labels( + &mut self, + root: &Path, + image_ref: &str, + ) -> Result>, String>; + + fn docker_manifest_exists(&mut self, root: &Path, image_ref: &str) -> bool; + + fn docker_image_has_container(&mut self, root: &Path, image_ref: &str) -> bool; + + fn docker_image_tags( + &mut self, + root: &Path, + repository: &str, + ) -> Result, String>; + + fn docker_image_remove(&mut self, root: &Path, image_ref: &str) -> Result<(), String>; +} + +struct RealImageCommandRunner; + pub fn prepare_node_image(request: NodeImageRequest) -> Result { + prepare_node_image_with_progress(request, None) +} + +pub fn prepare_node_image_with_progress( + request: NodeImageRequest, + progress: Option<&mut dyn NodeImageProgressSink>, +) -> Result { + let mut progress = progress; + let mut runner = RealImageCommandRunner; + prepare_node_image_inner(request, &mut progress, &mut runner) +} + +fn prepare_node_image_inner( + request: NodeImageRequest, + progress: &mut Option<&mut dyn NodeImageProgressSink>, + runner: &mut dyn ImageCommandRunner, +) -> Result { + emit_image_reference(progress, "requested", &request.requested_image); if !request.enabled { return Ok(PreparedNodeImage { image_ref: request.requested_image, @@ -83,6 +182,7 @@ pub fn prepare_node_image(request: NodeImageRequest) -> Result Result Result Result Result Result Result Result(base_hash: &'a str) -> Vec<(&'static str, &'a str)> { } fn ensure_aliases_local( + runner: &mut dyn ImageCommandRunner, + progress: &mut Option<&mut dyn NodeImageProgressSink>, root: &Path, source_ref: &str, image: &ImageName, @@ -434,10 +556,13 @@ fn ensure_aliases_local( for alias in alias_refs(image, alias_tags) { if alias != source_ref { run_status( + runner, + progress, root, "docker", &["tag", source_ref, &alias], "tag mvp node image", + Some(&alias), )?; } } @@ -445,6 +570,8 @@ fn ensure_aliases_local( } fn ensure_aliases_for_remote( + runner: &mut dyn ImageCommandRunner, + progress: &mut Option<&mut dyn NodeImageProgressSink>, root: &Path, source_ref: &str, image: &ImageName, @@ -453,12 +580,20 @@ fn ensure_aliases_for_remote( if alias_tags.is_empty() { return Ok(false); } - if !docker_image_exists(root, source_ref) { - run_status(root, "docker", &["pull", source_ref], "pull mvp node image")?; + if !runner.docker_image_exists(root, source_ref) { + run_status( + runner, + progress, + root, + "docker", + &["pull", source_ref], + "pull mvp node image", + Some(source_ref), + )?; } - ensure_aliases_local(root, source_ref, image, alias_tags)?; + ensure_aliases_local(runner, progress, root, source_ref, image, alias_tags)?; for alias in alias_refs(image, alias_tags) { - push_image(root, &alias)?; + push_image(runner, progress, root, &alias)?; } Ok(true) } @@ -470,8 +605,21 @@ fn alias_refs(image: &ImageName, alias_tags: &BTreeSet) -> Vec { .collect() } -fn push_image(root: &Path, image_ref: &str) -> Result<(), String> { - run_status(root, "docker", &["push", image_ref], "push mvp node image") +fn push_image( + runner: &mut dyn ImageCommandRunner, + progress: &mut Option<&mut dyn NodeImageProgressSink>, + root: &Path, + image_ref: &str, +) -> Result<(), String> { + run_status( + runner, + progress, + root, + "docker", + &["push", image_ref], + "push mvp node image", + Some(image_ref), + ) } fn docker_image_exists(root: &Path, image_ref: &str) -> bool { @@ -487,11 +635,12 @@ fn docker_image_exists(root: &Path, image_ref: &str) -> bool { } fn docker_image_labels_match( + runner: &mut dyn ImageCommandRunner, root: &Path, image_ref: &str, expected: &[(&str, &str)], ) -> Result { - let Some(labels) = docker_image_labels(root, image_ref)? else { + let Some(labels) = runner.docker_image_labels(root, image_ref)? else { return Ok(false); }; Ok(expected @@ -536,44 +685,27 @@ fn docker_manifest_exists(root: &Path, image_ref: &str) -> bool { .unwrap_or(false) } -fn prune_old_dirty_images(root: &Path, image: &ImageName, keep_tag: &str) { +fn prune_old_dirty_images( + runner: &mut dyn ImageCommandRunner, + root: &Path, + image: &ImageName, + keep_tag: &str, +) { if !dirty_image_prune_enabled() { return; } - let output = match Command::new("docker") - .current_dir(root) - .args([ - "image", - "ls", - "--format", - "{{.Repository}}\t{{.Tag}}", - &image.repository, - ]) - .stdin(Stdio::null()) - .output() - { - Ok(output) => output, + let tags = match runner.docker_image_tags(root, &image.repository) { + Ok(tags) => tags, Err(error) => { eprintln!("mvp-node-image: prune old dirty images skipped: {error}"); return; } }; - if !output.status.success() { - eprintln!( - "mvp-node-image: prune old dirty images skipped: docker image ls failed with {}", - output.status - ); - return; - } let keep_old = dirty_image_prune_keep(); let mut retained_old = 0_usize; - let stdout = String::from_utf8_lossy(&output.stdout); - for line in stdout.lines() { - let Some((repository, tag)) = line.split_once('\t') else { - continue; - }; + for (repository, tag) in tags { if repository != image.repository || !tag.starts_with("dirty-") || tag == keep_tag @@ -582,18 +714,18 @@ fn prune_old_dirty_images(root: &Path, image: &ImageName, keep_tag: &str) { continue; } - let image_ref = image.ref_for_tag(tag); - let Ok(Some(labels)) = docker_image_labels(root, &image_ref) else { + let image_ref = image.ref_for_tag(&tag); + let Ok(Some(labels)) = runner.docker_image_labels(root, &image_ref) else { continue; }; - if labels.get(NODE_IMAGE_TAG_LABEL).map(String::as_str) != Some(tag) + if labels.get(NODE_IMAGE_TAG_LABEL).map(String::as_str) != Some(tag.as_str()) || !labels.contains_key(NODE_IMAGE_SOURCE_HASH_LABEL) || !labels.contains_key(NODE_IMAGE_WORKER_HASH_LABEL) || !labels.contains_key(NODE_IMAGE_BASE_HASH_LABEL) { continue; } - if docker_image_has_container(root, &image_ref) { + if runner.docker_image_has_container(root, &image_ref) { eprintln!( "mvp-node-image: prune old dirty image {image_ref} skipped: container exists" ); @@ -605,23 +737,8 @@ fn prune_old_dirty_images(root: &Path, image: &ImageName, keep_tag: &str) { } eprintln!("mvp-node-image: prune old dirty image {image_ref}"); - match Command::new("docker") - .current_dir(root) - .args(["image", "rm", &image_ref]) - .stdin(Stdio::null()) - .stdout(Stdio::null()) - .stderr(Stdio::null()) - .status() - { - Ok(status) if status.success() => {} - Ok(status) => { - eprintln!( - "mvp-node-image: prune old dirty image {image_ref} skipped: docker image rm failed with {status}" - ); - } - Err(error) => { - eprintln!("mvp-node-image: prune old dirty image {image_ref} skipped: {error}"); - } + if let Err(error) = runner.docker_image_remove(root, &image_ref) { + eprintln!("mvp-node-image: prune old dirty image {image_ref} skipped: {error}"); } } } @@ -659,42 +776,314 @@ fn docker_image_has_container(root: &Path, image_ref: &str) -> bool { .unwrap_or(true) } -fn run_status(root: &Path, program: &str, args: &[&str], label: &str) -> Result<(), String> { - eprintln!("mvp-node-image: {label}"); - let status = Command::new(program) - .current_dir(root) - .args(args) - .stdin(Stdio::null()) - .stdout(Stdio::inherit()) - .stderr(Stdio::inherit()) - .status() - .map_err(|e| format!("run {label}: {e}"))?; - if status.success() { - Ok(()) - } else { - Err(format!("{label} failed with {status}")) - } +fn run_status( + runner: &mut dyn ImageCommandRunner, + progress: &mut Option<&mut dyn NodeImageProgressSink>, + root: &Path, + program: &str, + args: &[&str], + label: &str, + image_ref: Option<&str>, +) -> Result<(), String> { + let args = args.iter().map(|arg| (*arg).to_owned()).collect::>(); + runner.run_status(root, program, &args, label, image_ref, progress) } fn run_status_vec( + runner: &mut dyn ImageCommandRunner, + progress: &mut Option<&mut dyn NodeImageProgressSink>, root: &Path, program: &str, args: Vec, label: &str, + image_ref: Option<&str>, ) -> Result<(), String> { - eprintln!("mvp-node-image: {label}"); - let status = Command::new(program) - .current_dir(root) - .args(&args) - .stdin(Stdio::null()) - .stdout(Stdio::inherit()) - .stderr(Stdio::inherit()) - .status() - .map_err(|e| format!("run {label}: {e}"))?; - if status.success() { - Ok(()) - } else { - Err(format!("{label} failed with {status}")) + runner.run_status(root, program, &args, label, image_ref, progress) +} + +fn emit_image_reference( + progress: &mut Option<&mut dyn NodeImageProgressSink>, + role: &str, + image_ref: &str, +) { + emit_progress( + progress, + NodeImageProgressEvent { + command_label: None, + image_ref: Some(image_ref.to_owned()), + elapsed_ms: None, + kind: NodeImageProgressEventKind::ImageReference { + role: role.to_owned(), + image_ref: image_ref.to_owned(), + }, + }, + ); +} + +fn emit_progress( + progress: &mut Option<&mut dyn NodeImageProgressSink>, + event: NodeImageProgressEvent, +) { + if let Some(sink) = progress.as_deref_mut() { + sink.emit(event); + } +} + +fn emit_command_progress( + progress: &mut Option<&mut dyn NodeImageProgressSink>, + label: &str, + image_ref: Option<&str>, + elapsed_ms: u128, + kind: NodeImageProgressEventKind, +) { + emit_progress( + progress, + NodeImageProgressEvent { + command_label: Some(label.to_owned()), + image_ref: image_ref.map(str::to_owned), + elapsed_ms: Some(elapsed_ms), + kind, + }, + ); +} + +enum CommandOutputLine { + Stdout(String), + Stderr(String), +} + +fn spawn_line_reader( + reader: R, + to_line: fn(String) -> CommandOutputLine, + tx: mpsc::Sender, +) -> thread::JoinHandle<()> +where + R: Read + Send + 'static, +{ + thread::spawn(move || { + for line in BufReader::new(reader).lines().map_while(Result::ok) { + if tx.send(to_line(line)).is_err() { + break; + } + } + }) +} + +fn drain_command_lines( + rx: &mpsc::Receiver, + progress: &mut Option<&mut dyn NodeImageProgressSink>, + label: &str, + image_ref: Option<&str>, + started: Instant, +) { + while let Ok(line) = rx.try_recv() { + let kind = match line { + CommandOutputLine::Stdout(line) => NodeImageProgressEventKind::CommandStdout { line }, + CommandOutputLine::Stderr(line) => NodeImageProgressEventKind::CommandStderr { line }, + }; + emit_command_progress( + progress, + label, + image_ref, + started.elapsed().as_millis(), + kind, + ); + } +} + +impl ImageCommandRunner for RealImageCommandRunner { + fn run_status( + &mut self, + root: &Path, + program: &str, + args: &[String], + label: &str, + image_ref: Option<&str>, + progress: &mut Option<&mut dyn NodeImageProgressSink>, + ) -> Result<(), String> { + eprintln!("mvp-node-image: {label}"); + if progress.is_none() { + let status = Command::new(program) + .current_dir(root) + .args(args) + .stdin(Stdio::null()) + .stdout(Stdio::inherit()) + .stderr(Stdio::inherit()) + .status() + .map_err(|e| format!("run {label}: {e}"))?; + return if status.success() { + Ok(()) + } else { + Err(format!("{label} failed with {status}")) + }; + } + + let started = Instant::now(); + emit_command_progress( + progress, + label, + image_ref, + 0, + NodeImageProgressEventKind::CommandStarted { + program: program.to_owned(), + args: args.to_vec(), + }, + ); + let mut child = match Command::new(program) + .current_dir(root) + .args(args) + .stdin(Stdio::null()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + { + Ok(child) => child, + Err(error) => { + emit_command_progress( + progress, + label, + image_ref, + started.elapsed().as_millis(), + NodeImageProgressEventKind::CommandExited { + status: format!("spawn error: {error}"), + code: None, + success: false, + }, + ); + return Err(format!("run {label}: {error}")); + } + }; + + let (tx, rx) = mpsc::channel(); + let mut readers = Vec::new(); + if let Some(stdout) = child.stdout.take() { + readers.push(spawn_line_reader( + stdout, + CommandOutputLine::Stdout, + tx.clone(), + )); + } + if let Some(stderr) = child.stderr.take() { + readers.push(spawn_line_reader( + stderr, + CommandOutputLine::Stderr, + tx.clone(), + )); + } + drop(tx); + + let status = loop { + match child.try_wait() { + Ok(Some(status)) => break status, + Ok(None) => { + drain_command_lines(&rx, progress, label, image_ref, started); + thread::sleep(Duration::from_millis(10)); + } + Err(error) => { + drain_command_lines(&rx, progress, label, image_ref, started); + emit_command_progress( + progress, + label, + image_ref, + started.elapsed().as_millis(), + NodeImageProgressEventKind::CommandExited { + status: format!("wait error: {error}"), + code: None, + success: false, + }, + ); + return Err(format!("run {label}: {error}")); + } + } + }; + for reader in readers { + let _ = reader.join(); + } + drain_command_lines(&rx, progress, label, image_ref, started); + let status_text = status.to_string(); + let success = status.success(); + emit_command_progress( + progress, + label, + image_ref, + started.elapsed().as_millis(), + NodeImageProgressEventKind::CommandExited { + status: status_text.clone(), + code: status.code(), + success, + }, + ); + if success { + Ok(()) + } else { + Err(format!("{label} failed with {status_text}")) + } + } + + fn docker_image_exists(&mut self, root: &Path, image_ref: &str) -> bool { + docker_image_exists(root, image_ref) + } + + fn docker_image_labels( + &mut self, + root: &Path, + image_ref: &str, + ) -> Result>, String> { + docker_image_labels(root, image_ref) + } + + fn docker_manifest_exists(&mut self, root: &Path, image_ref: &str) -> bool { + docker_manifest_exists(root, image_ref) + } + + fn docker_image_has_container(&mut self, root: &Path, image_ref: &str) -> bool { + docker_image_has_container(root, image_ref) + } + + fn docker_image_tags( + &mut self, + root: &Path, + repository: &str, + ) -> Result, String> { + let output = Command::new("docker") + .current_dir(root) + .args([ + "image", + "ls", + "--format", + "{{.Repository}}\t{{.Tag}}", + repository, + ]) + .stdin(Stdio::null()) + .output() + .map_err(|error| format!("docker image ls failed: {error}"))?; + if !output.status.success() { + return Err(format!("docker image ls failed with {}", output.status)); + } + let stdout = String::from_utf8_lossy(&output.stdout); + Ok(stdout + .lines() + .filter_map(|line| { + let (repository, tag) = line.split_once('\t')?; + Some((repository.to_owned(), tag.to_owned())) + }) + .collect()) + } + + fn docker_image_remove(&mut self, root: &Path, image_ref: &str) -> Result<(), String> { + let status = Command::new("docker") + .current_dir(root) + .args(["image", "rm", image_ref]) + .stdin(Stdio::null()) + .stdout(Stdio::null()) + .stderr(Stdio::null()) + .status() + .map_err(|error| format!("docker image rm failed: {error}"))?; + if status.success() { + Ok(()) + } else { + Err(format!("docker image rm failed with {status}")) + } } } @@ -743,6 +1132,7 @@ impl ImageName { } Ok(Self { repository, + requested_tag, }) } @@ -756,6 +1146,251 @@ impl ImageName { mod tests { use super::*; + #[derive(Default)] + struct CollectProgress { + events: Vec, + } + + impl NodeImageProgressSink for CollectProgress { + fn emit(&mut self, event: NodeImageProgressEvent) { + self.events.push(event); + } + } + + #[derive(Default)] + struct DryImageCommandRunner { + commands: Vec<(String, Vec, String, Option)>, + labels: BTreeMap>, + existing_images: BTreeSet, + manifests: BTreeSet, + image_tags: Vec<(String, String)>, + containers: BTreeSet, + removed_images: Vec, + } + + impl ImageCommandRunner for DryImageCommandRunner { + fn run_status( + &mut self, + _root: &Path, + program: &str, + args: &[String], + label: &str, + image_ref: Option<&str>, + progress: &mut Option<&mut dyn NodeImageProgressSink>, + ) -> Result<(), String> { + self.commands.push(( + program.to_owned(), + args.to_vec(), + label.to_owned(), + image_ref.map(str::to_owned), + )); + emit_command_progress( + progress, + label, + image_ref, + 0, + NodeImageProgressEventKind::CommandStarted { + program: program.to_owned(), + args: args.to_vec(), + }, + ); + emit_command_progress( + progress, + label, + image_ref, + 1, + NodeImageProgressEventKind::CommandStdout { + line: format!("{label} stdout"), + }, + ); + emit_command_progress( + progress, + label, + image_ref, + 2, + NodeImageProgressEventKind::CommandStderr { + line: format!("{label} stderr"), + }, + ); + emit_command_progress( + progress, + label, + image_ref, + 3, + NodeImageProgressEventKind::CommandExited { + status: "exit status: 0".to_owned(), + code: Some(0), + success: true, + }, + ); + Ok(()) + } + + fn docker_image_exists(&mut self, _root: &Path, image_ref: &str) -> bool { + self.existing_images.contains(image_ref) + } + + fn docker_image_labels( + &mut self, + _root: &Path, + image_ref: &str, + ) -> Result>, String> { + Ok(self.labels.get(image_ref).cloned()) + } + + fn docker_manifest_exists(&mut self, _root: &Path, image_ref: &str) -> bool { + self.manifests.contains(image_ref) + } + + fn docker_image_has_container(&mut self, _root: &Path, image_ref: &str) -> bool { + self.containers.contains(image_ref) + } + + fn docker_image_tags( + &mut self, + _root: &Path, + _repository: &str, + ) -> Result, String> { + Ok(self.image_tags.clone()) + } + + fn docker_image_remove(&mut self, _root: &Path, image_ref: &str) -> Result<(), String> { + self.removed_images.push(image_ref.to_owned()); + Ok(()) + } + } + + #[test] + fn fake_child_command_emits_structured_progress_and_duration() { + let mut runner = RealImageCommandRunner; + let mut progress = CollectProgress::default(); + let mut sink: Option<&mut dyn NodeImageProgressSink> = Some(&mut progress); + + runner + .run_status( + Path::new("."), + "sh", + &[ + "-c".to_owned(), + "printf 'stdout line\\n'; printf 'stderr line\\n' >&2".to_owned(), + ], + "fake child progress", + Some("docker.io/acme/node:test"), + &mut sink, + ) + .expect("fake child exits successfully"); + + assert!(progress.events.iter().any(|event| matches!( + &event.kind, + NodeImageProgressEventKind::CommandStarted { program, .. } if program == "sh" + ))); + assert!(progress.events.iter().any(|event| matches!( + &event.kind, + NodeImageProgressEventKind::CommandStdout { line } if line == "stdout line" + ))); + assert!(progress.events.iter().any(|event| matches!( + &event.kind, + NodeImageProgressEventKind::CommandStderr { line } if line == "stderr line" + ))); + let exit = progress + .events + .iter() + .find(|event| matches!(event.kind, NodeImageProgressEventKind::CommandExited { .. })) + .expect("exit progress event"); + assert_eq!(exit.command_label.as_deref(), Some("fake child progress")); + assert_eq!(exit.image_ref.as_deref(), Some("docker.io/acme/node:test")); + assert!(exit.elapsed_ms.is_some()); + assert!(matches!( + exit.kind, + NodeImageProgressEventKind::CommandExited { + success: true, + code: Some(0), + .. + } + )); + } + + #[test] + fn fake_child_failure_preserves_command_label_and_status() { + let mut runner = RealImageCommandRunner; + let mut progress = CollectProgress::default(); + let mut sink: Option<&mut dyn NodeImageProgressSink> = Some(&mut progress); + + let error = runner + .run_status( + Path::new("."), + "sh", + &["-c".to_owned(), "exit 7".to_owned()], + "failing fake child", + None, + &mut sink, + ) + .expect_err("fake child failure propagates"); + assert!(error.contains("failing fake child"), "{error}"); + + let exit = progress + .events + .iter() + .find_map(|event| match &event.kind { + NodeImageProgressEventKind::CommandExited { + status, + code, + success, + } => Some((event, status, code, success)), + _ => None, + }) + .expect("exit progress event"); + assert_eq!(exit.0.command_label.as_deref(), Some("failing fake child")); + assert!(exit.1.contains("exit status"), "{}", exit.1); + assert_eq!(*exit.2, Some(7)); + assert!(!*exit.3); + } + + #[test] + fn prepare_node_image_with_dry_runner_returns_expected_image_and_progress() { + let root = workspace_root().expect("workspace root resolves"); + let mut runner = DryImageCommandRunner::default(); + let mut progress = CollectProgress::default(); + let mut sink: Option<&mut dyn NodeImageProgressSink> = Some(&mut progress); + + let prepared = prepare_node_image_inner( + 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"), + provider: NodeImageProvider::Docker, + extra_tag: Some("smoke".to_owned()), + push: false, + force_refresh: false, + enabled: true, + }, + &mut sink, + &mut runner, + ) + .expect("dry image preparation succeeds"); + + assert_eq!( + prepared.image_ref, + format!("docker.io/acme/mvp-node:{}", prepared.tag) + ); + assert!(prepared.built); + assert!(!prepared.pushed); + assert!(runner.commands.iter().any(|(_, _, label, image_ref)| { + label == "build mvp node image" + && image_ref.as_deref() == Some(prepared.image_ref.as_str()) + })); + assert!(progress.events.iter().any(|event| matches!( + &event.kind, + NodeImageProgressEventKind::ImageReference { role, image_ref } + if role == "resolved" && image_ref == &prepared.image_ref + ))); + assert!(progress.events.iter().any(|event| matches!( + &event.kind, + NodeImageProgressEventKind::CommandStdout { line } + if line == "build mvp node image stdout" + ))); + } + #[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/orchestrator_app.rs b/crates/mvp-system/src/orchestrator_app.rs index f0d67fc..af5d49d 100644 --- a/crates/mvp-system/src/orchestrator_app.rs +++ b/crates/mvp-system/src/orchestrator_app.rs @@ -24,6 +24,7 @@ use crate::distribution_stack::DistributionRuntimeStack; use crate::endpoint_advertisement::{ EndpointAddrMask, MVP_IROH_ENDPOINT_ADDR_MASK_ENV, advertised_endpoint, }; +use crate::gguf_shard::{StageShardPlan, plan_stage_shard}; use crate::gpu_worker_ingress_parser as ingress; use crate::node_provisioning::ProviderKind; use crate::orchestrator_run_fsm::{RunConfig, RunId}; @@ -84,6 +85,7 @@ const DEFAULT_MAX_TOKENS: u32 = 64; const PUMP_INTERVAL: Duration = Duration::from_millis(10); const RUNTIME_READY_ACK_RETRY_INTERVAL: Duration = Duration::from_millis(250); const RUNTIME_READY_ACK_TIMEOUT: Duration = Duration::from_secs(60); +const STAGE_PROVISION_ACTIVE_RESEND_AFTER: Duration = Duration::from_secs(60); const RUNTIME_READY_TIMEOUT: Duration = Duration::from_secs(60); const MVP_ORCH_BOOTSTRAP: &str = "mvp.orch.bootstrap"; const MVP_ORCH_PROMPT: &str = "mvp.orch.prompt"; @@ -2286,9 +2288,9 @@ impl Drop for ProvisionedClusterGuard { } } -fn start_node_with_stdio_capture( +fn start_nodes_with_stdio_capture( provisioner: Box, - node_spec: NodeProvisionSpec, + node_specs: Vec, sink: PluginSink, orch_stdio_rx: Option<&mpsc::Receiver>, dashboard: Option<&DashboardSupport>, @@ -2297,13 +2299,16 @@ fn start_node_with_stdio_capture( node_id: u64, ) -> ( Box, - Result, + Vec<( + NodeProvisionSpec, + Result, + )>, ) { let (tx, rx) = mpsc::channel(); thread::spawn(move || { let mut provisioner = provisioner; - let result = provisioner.start_node(node_spec, sink); - let _ = tx.send((provisioner, result)); + let results = provisioner.start_nodes(node_specs, sink); + let _ = tx.send((provisioner, results)); }); loop { @@ -2321,7 +2326,18 @@ fn start_node_with_stdio_capture( Err(mpsc::RecvTimeoutError::Disconnected) => { return ( Box::new(FailedProvisionPlugin), - Err("provider start worker disconnected".to_owned()), + vec![( + NodeProvisionSpec { + run_id, + node_id, + stage_index: None, + image: String::new(), + env: Vec::new(), + args: Vec::new(), + mounts: Vec::new(), + }, + Err("provider start worker disconnected".to_owned()), + )], ); } } @@ -2377,8 +2393,20 @@ fn start_and_provision_workers( }), ); - let mut handles = Vec::with_capacity(stage_specs.len()); - for node_spec in stage_specs { + let stage_shard_plans = if let Some(plan) = pipeline_plan { + let plans = pipeline_stage_shard_plans(config, plan)?; + emit_stage_shard_plan_summaries( + dashboard, + orch_datastream, + config.run_id, + config.node_id, + &plans, + ); + plans + } else { + BTreeMap::new() + }; + for node_spec in &stage_specs { orch_datastream.emit_event( dashboard, ProvisionEvent { @@ -2406,19 +2434,36 @@ fn start_and_provision_workers( "stage_index":node_spec.stage_index, }), ); - let (returned_provisioner, handle_result) = start_node_with_stdio_capture( - provisioner, - node_spec.clone(), - sink.clone(), - orch_stdio_rx, - dashboard, - orch_datastream, - config.run_id, - node_spec.node_id, - ); - provisioner = returned_provisioner; + } + let (returned_provisioner, start_results) = start_nodes_with_stdio_capture( + provisioner, + stage_specs, + sink.clone(), + orch_stdio_rx, + dashboard, + orch_datastream, + config.run_id, + config.node_id, + ); + provisioner = returned_provisioner; + let mut handles = Vec::with_capacity(start_results.len()); + for (node_spec, handle_result) in start_results { match handle_result { - Ok(handle) => handles.push(handle), + Ok(handle) => { + orch_datastream.emit_bootstrap( + dashboard, + config.run_id, + config.node_id, + "provider_start", + "ready", + json!({ + "provider":config.provider.as_str(), + "node_id":node_spec.node_id, + "stage_index":node_spec.stage_index, + }), + ); + handles.push(handle); + } Err(error) => { orch_datastream.emit_bootstrap( dashboard, @@ -2592,9 +2637,11 @@ fn start_and_provision_workers( config.node_id, config.provider, expected_node_ids.len(), + config, pipeline_plan.expect("pipeline mode requires plan"), &readies, &pipeline_coordinator, + &stage_shard_plans, ) } else { wait_for_weights_loaded( @@ -2746,6 +2793,7 @@ fn stage_provision_wire_from_plan( stage_index: u32, readies: &BTreeMap, coordinator: &EndpointAddr, + stage_shard_plans: &BTreeMap, ) -> Result { let provision = run_plan::derive_stage_provision(plan, stage_index) .map_err(|e| format!("derive stage {stage_index} provision: {e:?}"))?; @@ -2781,6 +2829,7 @@ fn stage_provision_wire_from_plan( model_id: provision.model.model_id, gguf_source: provision.gguf_source, tokenizer: provision.tokenizer, + stage_shard_plan: stage_shard_plans.get(&stage_index).cloned(), }) } @@ -2791,14 +2840,86 @@ fn provision_stage_from_plan( stage_index: u32, readies: &BTreeMap, coordinator: &EndpointAddr, + stage_shard_plans: &BTreeMap, ) -> Result<(), String> { - let provision = stage_provision_wire_from_plan(plan, stage_index, readies, coordinator)?; + let provision = + stage_provision_wire_from_plan(plan, stage_index, readies, coordinator, stage_shard_plans)?; stack .runtime .send_to(node_actor, NodeAgentMsg::ProvisionStage(provision)) .map_err(|e| format!("send stage {stage_index} provision: {e}")) } +fn pipeline_stage_shard_plans( + config: &Config, + plan: &run_plan::RunPlan, +) -> Result, String> { + if !matches!(config.gguf_source, GgufSource::HuggingFaceGguf { .. }) { + return Ok(BTreeMap::new()); + } + let planning_gguf = config.local_planning_gguf_path()?; + let mut out = BTreeMap::new(); + for stage in &plan.stages { + let shard_plan = plan_stage_shard( + &planning_gguf, + stage.gguf_source.clone(), + stage.stage_index, + stage.stage_count, + stage.layer_start, + stage.layer_end_exclusive, + ) + .map_err(|error| { + format!( + "plan stage {} HF shard ranges from {}: {error}", + stage.stage_index, + planning_gguf.display() + ) + })?; + out.insert(stage.stage_index, shard_plan); + } + Ok(out) +} + +fn stage_shard_plan_summary_detail(plan: &StageShardPlan) -> Value { + let planned_fetch_bytes = plan.planned_fetch_bytes(); + json!({ + "stage_index":plan.stage_index, + "stage_count":plan.stage_count, + "layer_start":plan.layer_start, + "layer_end_exclusive":plan.layer_end_exclusive, + "planned_fetch_bytes":planned_fetch_bytes, + "source_total_bytes":plan.source_total_bytes, + "tensor_count":plan.tensors.len(), + "range_count":plan.planned_range_count(), + "tensor_range_count":plan.merged_tensor_ranges.len(), + "metadata_bytes":plan.metadata_end, + "planned_fraction":if plan.source_total_bytes == 0 { + Value::Null + } else { + json!(planned_fetch_bytes as f64 / plan.source_total_bytes as f64) + }, + }) +} + +fn emit_stage_shard_plan_summaries( + dashboard: Option<&DashboardSupport>, + orch_datastream: &mut OrchDatastream, + run_id: u64, + node_id: u64, + stage_shard_plans: &BTreeMap, +) { + for plan in stage_shard_plans.values() { + orch_datastream.emit_bootstrap( + dashboard, + run_id, + node_id, + "stage_shard_plan", + "ready", + stage_shard_plan_summary_detail(plan), + ); + } +} + fn stage_consumer_endpoint( consumer_node_id: u64, plan: &run_plan::RunPlan, @@ -2972,9 +3093,11 @@ fn wait_for_weights_loaded_count( node_id: u64, provider: ProviderKind, expected_count: usize, + _config: &Config, pipeline_plan: &run_plan::RunPlan, readies: &BTreeMap, pipeline_coordinator: &EndpointAddr, + stage_shard_plans: &BTreeMap, ) -> Result<(), String> { let expected_stages = pipeline_plan .stages @@ -2982,10 +3105,11 @@ fn wait_for_weights_loaded_count( .map(|stage| stage.stage_index) .collect::>(); let mut loaded_stages = BTreeSet::::new(); - let mut last_resend = Instant::now(); - let mut active_stage = None::; + let mut last_resend = Instant::now() - Duration::from_secs(15); let mut resend_attempt = 0_u64; let mut stage_resend_counts = BTreeMap::::new(); + let mut stage_last_sends = BTreeMap::::new(); + let mut load_progress = BTreeMap::::new(); loop { pump(driver, stack, frame_tx); emit_swim_transitions(orch_datastream, dashboard, run_id, node_id, stack); @@ -2997,140 +3121,83 @@ fn wait_for_weights_loaded_count( if loaded_stages.len() >= expected_count { return Ok(()); } - if active_stage.map_or(true, |stage_index| loaded_stages.contains(&stage_index)) { - active_stage = - next_pipeline_weight_load_stage(pipeline_plan, &loaded_stages, active_stage) - .map(|stage| stage.stage_index); - last_resend = Instant::now() - Duration::from_secs(1); - } - if last_resend.elapsed() >= Duration::from_secs(1) { + if last_resend.elapsed() >= Duration::from_secs(15) { resend_attempt += 1; - let Some(stage) = - next_pipeline_weight_load_stage(pipeline_plan, &loaded_stages, active_stage) - else { + let pending = pending_pipeline_weight_load_stages(pipeline_plan, &loaded_stages); + if pending.is_empty() { return Err(format!( "missing unloaded pipeline weight stage; loaded {} of {expected_count}", loaded_stages.len() )); - }; - let stage_node_id = stage.node_id.0; - let stage_send_count = { - let count = stage_resend_counts.entry(stage.stage_index).or_default(); - *count += 1; - *count - }; - let emit_wait_headline = stage_send_count == 1 || stage_send_count % 15 == 0; - let ready = readies.get(&stage_node_id).ok_or_else(|| { - format!("missing runtime-ready node for stage {}", stage.stage_index) - })?; - orch_datastream.emit_bootstrap( - dashboard, - run_id, - node_id, - "stage_provision_send", - "sent", - json!({ - "attempt":resend_attempt, - "stage_count":pipeline_plan.stages.len(), - "stage_index":stage.stage_index, - "stage_send_count":stage_send_count, - "loaded_stage_count":loaded_stages.len() - }), - ); - let route_owner = stack.route_owner(ready.node_actor); - let datastream_route_owner = stack.route_owner(ready.datastream_publisher); - let member_state = stack.member_state(ready.swim_node_id); - let route_matches_ready = route_owner == Some(ready.swim_node_id); - orch_datastream.emit_bootstrap_to_channel( - dashboard, - MVP_STAGE_ROUTE, - run_id, - node_id, - "stage_route_check", - "observed", - json!({ - "attempt":resend_attempt, - "stage_index":stage.stage_index, - "stage_node_id":stage_node_id, - "node_actor":ready.node_actor, - "datastream_publisher":ready.datastream_publisher, - "swim_node_id":format!("{:?}", ready.swim_node_id), - "member_state":member_state.map(|state| format!("{:?}", state)), - "route_owner":route_owner.map(|owner| format!("{:?}", owner)), - "datastream_route_owner":datastream_route_owner.map(|owner| format!("{:?}", owner)), - "route_matches_ready":route_matches_ready, - }), - ); - if member_state == Some(MemberState::Dead) { - let reason = format!( - "stage {} node {} is dead while loading pipeline weights", - stage.stage_index, stage_node_id - ); - orch_datastream.emit_bootstrap( + } + for stage in pending { + let _sent = send_pipeline_stage_provision( + driver, + stack, + frame_tx, dashboard, + orch_datastream, run_id, node_id, - "stage_provision_wait", - "failed", - json!({ - "attempt":resend_attempt, - "stage_count":pipeline_plan.stages.len(), - "stage_index":stage.stage_index, - "stage_node_id":stage_node_id, - "stage_send_count":stage_send_count, - "loaded_stage_count":loaded_stages.len(), - "member_state":"Dead", - "route_owner":route_owner.map(|owner| format!("{:?}", owner)), - "datastream_route_owner":datastream_route_owner.map(|owner| format!("{:?}", owner)), - "reason":reason, - }), - ); - return Err(reason); + pipeline_plan, + stage, + readies, + pipeline_coordinator, + &stage_shard_plans, + &loaded_stages, + &mut stage_resend_counts, + &mut stage_last_sends, + &load_progress, + resend_attempt, + )?; } - if emit_wait_headline { - orch_datastream.emit_bootstrap( - dashboard, - run_id, - node_id, - "stage_provision_wait", - "observed", - json!({ - "attempt":resend_attempt, - "stage_count":pipeline_plan.stages.len(), - "stage_index":stage.stage_index, - "stage_node_id":stage_node_id, - "stage_send_count":stage_send_count, - "loaded_stage_count":loaded_stages.len(), - "message":format!( - "loaded {} of {}; waiting on stage {}", - loaded_stages.len(), - pipeline_plan.stages.len(), - stage.stage_index - ) - }), - ); - } - provision_stage_from_plan( - stack, - ready.node_actor, - pipeline_plan, - stage.stage_index, - readies, - pipeline_coordinator, - )?; - pump(driver, stack, frame_tx); last_resend = Instant::now(); } while let Ok(observation) = obs_rx.try_recv() { emit_plugin_observation(orch_datastream, dashboard, provider, &observation); match observation { - PluginObservation::Failed { reason, .. } => return Err(reason), - PluginObservation::Exited { - node_id, status, .. + PluginObservation::Failed { + reason, + node_id: failed_node_id, + .. } => { - return Err(format!( - "node {node_id} exited while loading weights: {status:?}" - )); + orch_datastream.emit_bootstrap( + dashboard, + run_id, + node_id, + "stage_provision_wait", + "failed", + json!({ + "classification":"worker_process_failed", + "stage_node_id":failed_node_id, + "last_load_progress":load_progress.get(&failed_node_id).map(StageLoadProgress::to_json), + "reason":reason, + }), + ); + return Err(reason); + } + PluginObservation::Exited { + node_id: exited_node_id, + status, + .. + } => { + let reason = + format!("node {exited_node_id} exited while loading weights: {status:?}"); + orch_datastream.emit_bootstrap( + dashboard, + run_id, + node_id, + "stage_provision_wait", + "failed", + json!({ + "classification":"worker_process_exited", + "stage_node_id":exited_node_id, + "last_load_progress":load_progress.get(&exited_node_id).map(StageLoadProgress::to_json), + "status":status, + "reason":reason, + }), + ); + return Err(reason); } PluginObservation::DatastreamFrame { .. } => {} PluginObservation::ProviderLine { .. } @@ -3138,7 +3205,7 @@ fn wait_for_weights_loaded_count( | PluginObservation::StderrLine { .. } => {} } } - drain_frames(frame_rx, dashboard, orch_datastream); + drain_frames_with_load_progress(frame_rx, dashboard, orch_datastream, &mut load_progress); while let Some(report) = orchestrator_reports.try_recv() { match report { OrchestratorReport::WeightsReady { @@ -3147,10 +3214,6 @@ fn wait_for_weights_loaded_count( stage_index, } if report_run_id == run_id && expected_stages.contains(&stage_index) => { loaded_stages.insert(stage_index); - if active_stage == Some(stage_index) { - active_stage = None; - last_resend = Instant::now() - Duration::from_secs(1); - } if loaded_stages.len() >= expected_count { return Ok(()); } @@ -3170,27 +3233,218 @@ fn wait_for_weights_loaded_count( } } -fn next_pipeline_weight_load_stage<'a>( +fn pending_pipeline_weight_load_stages<'a>( pipeline_plan: &'a run_plan::RunPlan, loaded_stages: &BTreeSet, - active_stage: Option, -) -> Option<&'a run_plan::StagePlan> { - if let Some(stage_index) = active_stage { - if !loaded_stages.contains(&stage_index) { - if let Some(stage) = pipeline_plan - .stages - .iter() - .find(|stage| stage.stage_index == stage_index) - { - return Some(stage); - } - } - } - pipeline_plan +) -> Vec<&'a run_plan::StagePlan> { + let mut pending = pipeline_plan .stages .iter() .filter(|stage| !loaded_stages.contains(&stage.stage_index)) - .max_by_key(|stage| stage.stage_index) + .collect::>(); + pending.sort_by_key(|stage| stage.stage_index); + pending +} + +#[allow(clippy::too_many_arguments)] +fn send_pipeline_stage_provision( + driver: &mut IrohDriver, + stack: &DistributionRuntimeStack, + frame_tx: &mpsc::Sender, + dashboard: Option<&DashboardSupport>, + orch_datastream: &mut OrchDatastream, + run_id: u64, + node_id: u64, + pipeline_plan: &run_plan::RunPlan, + stage: &run_plan::StagePlan, + readies: &BTreeMap, + pipeline_coordinator: &EndpointAddr, + stage_shard_plans: &BTreeMap, + loaded_stages: &BTreeSet, + stage_resend_counts: &mut BTreeMap, + stage_last_sends: &mut BTreeMap, + load_progress: &BTreeMap, + attempt: u64, +) -> Result { + let stage_node_id = stage.node_id.0; + let current_send_count = stage_resend_counts + .get(&stage.stage_index) + .copied() + .unwrap_or_default(); + let now = Instant::now(); + let ready = readies + .get(&stage_node_id) + .ok_or_else(|| format!("missing runtime-ready node for stage {}", stage.stage_index))?; + let route_owner = stack.route_owner(ready.node_actor); + let datastream_route_owner = stack.route_owner(ready.datastream_publisher); + let member_state = stack.member_state(ready.swim_node_id); + let route_matches_ready = route_owner == Some(ready.swim_node_id); + orch_datastream.emit_bootstrap_to_channel( + dashboard, + MVP_STAGE_ROUTE, + run_id, + node_id, + "stage_route_check", + "observed", + json!({ + "attempt":attempt, + "stage_index":stage.stage_index, + "stage_node_id":stage_node_id, + "node_actor":ready.node_actor, + "datastream_publisher":ready.datastream_publisher, + "swim_node_id":format!("{:?}", ready.swim_node_id), + "member_state":member_state.map(|state| format!("{:?}", state)), + "route_owner":route_owner.map(|owner| format!("{:?}", owner)), + "datastream_route_owner":datastream_route_owner.map(|owner| format!("{:?}", owner)), + "route_matches_ready":route_matches_ready, + }), + ); + if member_state == Some(MemberState::Dead) { + let reason = format!( + "stage {} node {} is dead while loading pipeline weights", + stage.stage_index, stage_node_id + ); + let liveness = stage_load_liveness_detail( + load_progress.get(&stage_node_id), + stage.stage_index, + stage_node_id, + member_state, + route_owner, + datastream_route_owner, + route_matches_ready, + "heartbeat_missed", + ); + orch_datastream.emit_bootstrap( + dashboard, + run_id, + node_id, + "stage_provision_wait", + "failed", + json!({ + "attempt":attempt, + "stage_count":pipeline_plan.stages.len(), + "stage_index":stage.stage_index, + "stage_node_id":stage_node_id, + "stage_send_count":current_send_count, + "loaded_stage_count":loaded_stages.len(), + "member_state":"Dead", + "route_owner":route_owner.map(|owner| format!("{:?}", owner)), + "datastream_route_owner":datastream_route_owner.map(|owner| format!("{:?}", owner)), + "classification":"heartbeat_missed", + "liveness":liveness, + "reason":reason, + }), + ); + return Err(reason); + } + let dispatch = stage_provision_dispatch( + load_progress.get(&stage_node_id), + current_send_count, + stage_last_sends.get(&stage.stage_index).copied(), + now, + ); + if !dispatch.should_send() { + orch_datastream.emit_bootstrap( + dashboard, + run_id, + node_id, + "stage_provision_wait", + "observed", + json!({ + "attempt":attempt, + "stage_count":pipeline_plan.stages.len(), + "stage_index":stage.stage_index, + "stage_node_id":stage_node_id, + "stage_send_count":current_send_count, + "loaded_stage_count":loaded_stages.len(), + "resend_suppressed":true, + "resend_reason":dispatch.reason(), + "liveness":stage_load_liveness_detail( + load_progress.get(&stage_node_id), + stage.stage_index, + stage_node_id, + member_state, + route_owner, + datastream_route_owner, + route_matches_ready, + "waiting", + ), + "message":format!( + "loaded {} of {}; waiting on stage {}", + loaded_stages.len(), + pipeline_plan.stages.len(), + stage.stage_index + ) + }), + ); + return Ok(false); + } + let stage_send_count = { + let count = stage_resend_counts.entry(stage.stage_index).or_default(); + *count += 1; + *count + }; + stage_last_sends.insert(stage.stage_index, now); + orch_datastream.emit_bootstrap( + dashboard, + run_id, + node_id, + "stage_provision_send", + "sent", + json!({ + "attempt":attempt, + "stage_count":pipeline_plan.stages.len(), + "stage_index":stage.stage_index, + "stage_send_count":stage_send_count, + "loaded_stage_count":loaded_stages.len(), + "parallel_weight_acquisition":true, + "resend_reason":dispatch.reason(), + }), + ); + if stage_send_count == 1 || stage_send_count % 15 == 0 { + orch_datastream.emit_bootstrap( + dashboard, + run_id, + node_id, + "stage_provision_wait", + "observed", + json!({ + "attempt":attempt, + "stage_count":pipeline_plan.stages.len(), + "stage_index":stage.stage_index, + "stage_node_id":stage_node_id, + "stage_send_count":stage_send_count, + "loaded_stage_count":loaded_stages.len(), + "liveness":stage_load_liveness_detail( + load_progress.get(&stage_node_id), + stage.stage_index, + stage_node_id, + member_state, + route_owner, + datastream_route_owner, + route_matches_ready, + "waiting", + ), + "message":format!( + "loaded {} of {}; waiting on stage {}", + loaded_stages.len(), + pipeline_plan.stages.len(), + stage.stage_index + ) + }), + ); + } + provision_stage_from_plan( + stack, + ready.node_actor, + pipeline_plan, + stage.stage_index, + readies, + pipeline_coordinator, + stage_shard_plans, + )?; + pump(driver, stack, frame_tx); + Ok(true) } struct FailedProvisionPlugin; @@ -3233,6 +3487,269 @@ struct CollectedDatastreamFrame { frame: Frame, } +#[derive(Clone, Debug, Default)] +struct StageLoadProgress { + node_id: u64, + stage_index: Option, + phase: Option, + bytes_done: Option, + bytes_total: Option, + last_progress: Option, + last_worker_event: Option, + failure_reason: Option, + host_gpu_samples: u64, +} + +impl StageLoadProgress { + fn to_json(&self) -> Value { + json!({ + "node_id": self.node_id, + "stage_index": self.stage_index, + "phase": self.phase.as_deref().unwrap_or("unknown"), + "bytes_done": self.bytes_done, + "bytes_total": self.bytes_total, + "last_progress_age_ms": self.last_progress.map(|at| at.elapsed().as_millis()), + "last_worker_event": self.last_worker_event, + "failure_reason": self.failure_reason, + "host_gpu_samples": self.host_gpu_samples, + "host_gpu_missing": self.host_gpu_samples == 0, + }) + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum StageProvisionDispatch { + Send(&'static str), + Suppress(&'static str), +} + +impl StageProvisionDispatch { + fn reason(self) -> &'static str { + match self { + Self::Send(reason) | Self::Suppress(reason) => reason, + } + } + + fn should_send(self) -> bool { + matches!(self, Self::Send(_)) + } +} + +fn stage_load_phase_is_active(phase: Option<&str>) -> bool { + matches!( + phase, + Some( + "loading_weights" + | "prefetching_model" + | "prefetching_stage_shard" + | "fetching_stage_shard" + | "stage_shard_cache_ready" + | "stage_shard_ready" + | "cache_ready" + | "constructing_stage" + | "stage_constructed" + | "building_tokenizer" + | "tokenizer_ready" + ) + ) +} + +fn stage_provision_dispatch( + progress: Option<&StageLoadProgress>, + send_count: u64, + last_send: Option, + now: Instant, +) -> StageProvisionDispatch { + if send_count == 0 { + return StageProvisionDispatch::Send("initial"); + } + let Some(progress) = progress else { + return StageProvisionDispatch::Send("no_progress_after_send"); + }; + if progress.failure_reason.is_some() || progress.phase.as_deref() == Some("failed") { + return StageProvisionDispatch::Suppress("worker_load_failed"); + } + if progress.phase.as_deref() == Some("weights_loaded") { + return StageProvisionDispatch::Suppress("weights_loaded_report_pending"); + } + if !stage_load_phase_is_active(progress.phase.as_deref()) { + return StageProvisionDispatch::Send("unknown_or_inactive_progress"); + } + let Some(last_progress) = progress.last_progress else { + return StageProvisionDispatch::Send("active_phase_without_progress_time"); + }; + if now.duration_since(last_progress) < STAGE_PROVISION_ACTIVE_RESEND_AFTER { + return StageProvisionDispatch::Suppress("active_progress"); + } + if let Some(last_send) = last_send { + if now.duration_since(last_send) < STAGE_PROVISION_ACTIVE_RESEND_AFTER { + return StageProvisionDispatch::Suppress("recent_stale_progress_resend"); + } + } + StageProvisionDispatch::Send("stale_progress") +} + +fn stage_load_liveness_detail( + progress: Option<&StageLoadProgress>, + stage_index: u32, + stage_node_id: u64, + member_state: Option, + route_owner: Option, + datastream_route_owner: Option, + route_matches_ready: bool, + classification: &str, +) -> Value { + json!({ + "classification": classification, + "stage_index": stage_index, + "stage_node_id": stage_node_id, + "member_state": member_state.map(|state| format!("{:?}", state)), + "route_owner": route_owner.map(|owner| format!("{:?}", owner)), + "datastream_route_owner": datastream_route_owner.map(|owner| format!("{:?}", owner)), + "route_matches_ready": route_matches_ready, + "load_progress": progress.map(StageLoadProgress::to_json), + "host_gpu_missing": progress.is_none_or(|progress| progress.host_gpu_samples == 0), + }) +} + +fn update_load_progress_from_frame( + progress: &mut BTreeMap, + collected: &CollectedDatastreamFrame, + now: Instant, +) { + let stream_node_id = collected.stream.node.as_str().parse::().ok(); + if collected.channel_name == "host.gpu" { + if let Some(node_id) = stream_node_id { + let entry = progress + .entry(node_id) + .or_insert_with(|| StageLoadProgress { + node_id, + ..StageLoadProgress::default() + }); + entry.host_gpu_samples = entry.host_gpu_samples.saturating_add(1); + } + return; + } + + let Ok(value) = serde_json::from_slice::(&collected.frame.payload) else { + return; + }; + if value.get("type").and_then(Value::as_str) == Some("NodeEvent") { + update_load_progress_from_node_event(progress, &value, now); + return; + } + if collected.channel_name == "mvp.worker.weights" { + let Some(node_id) = stream_node_id else { + return; + }; + update_load_progress_from_worker_event(progress, node_id, None, &value, now); + } +} + +fn update_load_progress_from_node_event( + progress: &mut BTreeMap, + value: &Value, + now: Instant, +) { + let Some(node_id) = numeric_json_field(value, "node_id") else { + return; + }; + let stage_index = + numeric_json_field(value, "stage_index").and_then(|stage| u32::try_from(stage).ok()); + let phase = value.get("phase").and_then(Value::as_str); + let status = value.get("status").and_then(Value::as_str); + let detail = value.get("detail").unwrap_or(&Value::Null); + if phase == Some("load_weights") { + let load_phase = match status { + Some("started") => Some("loading_weights"), + Some("ready") => Some("weights_loaded"), + Some("failed") => Some("failed"), + _ => None, + }; + if let Some(load_phase) = load_phase { + let entry = progress + .entry(node_id) + .or_insert_with(|| StageLoadProgress { + node_id, + ..StageLoadProgress::default() + }); + entry.stage_index = stage_index.or(entry.stage_index); + entry.phase = Some(load_phase.to_owned()); + entry.last_progress = Some(now); + if status == Some("failed") { + entry.failure_reason = detail + .get("error") + .and_then(Value::as_str) + .map(str::to_owned); + } + } + } + if let Some(worker_event) = detail.get("event") { + update_load_progress_from_worker_event(progress, node_id, stage_index, worker_event, now); + } +} + +fn update_load_progress_from_worker_event( + progress: &mut BTreeMap, + node_id: u64, + stage_index: Option, + event: &Value, + now: Instant, +) { + let Some(event_type) = event.get("type").and_then(Value::as_str) else { + return; + }; + let Some(phase) = load_phase_for_worker_event(event_type) else { + return; + }; + let entry = progress + .entry(node_id) + .or_insert_with(|| StageLoadProgress { + node_id, + ..StageLoadProgress::default() + }); + entry.stage_index = stage_index.or(entry.stage_index); + entry.phase = Some(phase.to_owned()); + entry.last_worker_event = Some(event_type.to_owned()); + entry.last_progress = Some(now); + if let Some(bytes_done) = + numeric_json_field(event, "bytes_done").or_else(|| numeric_json_field(event, "bytes")) + { + entry.bytes_done = Some(bytes_done); + } + if let Some(bytes_total) = numeric_json_field(event, "bytes_total") { + entry.bytes_total = Some(bytes_total); + } +} + +fn numeric_json_field(value: &Value, field: &str) -> Option { + value + .get(field) + .and_then(|value| value.as_u64().or_else(|| value.as_str()?.parse().ok())) +} + +fn load_phase_for_worker_event(event_type: &str) -> Option<&'static str> { + match event_type { + "GgufDownloadStarted" | "GgufDownloadProgress" => Some("prefetching_model"), + "GgufCacheReady" => Some("cache_ready"), + "StageShardFetchStarted" + | "StageShardRangeFetchStarted" + | "StageShardRangeFetchReady" + | "StageShardTensorFetchStarted" + | "StageShardTensorFetchReady" => Some("fetching_stage_shard"), + "StageShardCacheReady" => Some("stage_shard_cache_ready"), + "StageShardReady" => Some("stage_shard_ready"), + "StageShardFetchFailed" => Some("failed"), + "PipelineStageFromGgufStarted" => Some("constructing_stage"), + "PipelineStageFromGgufReady" => Some("stage_constructed"), + "TokenizerBuildStarted" => Some("building_tokenizer"), + "TokenizerBuildReady" => Some("tokenizer_ready"), + "WeightsLoaded" => Some("weights_loaded"), + "WorkerFatal" => Some("failed"), + _ => None, + } +} + fn drain_datastream_connections( driver: &mut IrohDriver, frame_tx: &mpsc::Sender, @@ -3916,6 +4433,7 @@ fn provision_stage( model_id: config.model_id.clone(), gguf_source: config.gguf_source.clone(), tokenizer: config.tokenizer.clone(), + stage_shard_plan: None, } }), ) @@ -4950,21 +5468,41 @@ fn drain_frames( orch_datastream: &mut OrchDatastream, ) { while let Ok(collected) = frame_rx.try_recv() { - ingest_dashboard_frame( - dashboard, - &collected.stream, - &collected.channel_name, - &collected.frame, - ); - orch_datastream.archive_frame( - "node", - &collected.stream, - &collected.channel_name, - &collected.frame, - ); + archive_collected_frame(collected, dashboard, orch_datastream); } } +fn drain_frames_with_load_progress( + frame_rx: &mpsc::Receiver, + dashboard: Option<&DashboardSupport>, + orch_datastream: &mut OrchDatastream, + progress: &mut BTreeMap, +) { + while let Ok(collected) = frame_rx.try_recv() { + update_load_progress_from_frame(progress, &collected, Instant::now()); + archive_collected_frame(collected, dashboard, orch_datastream); + } +} + +fn archive_collected_frame( + collected: CollectedDatastreamFrame, + dashboard: Option<&DashboardSupport>, + orch_datastream: &mut OrchDatastream, +) { + ingest_dashboard_frame( + dashboard, + &collected.stream, + &collected.channel_name, + &collected.frame, + ); + orch_datastream.archive_frame( + "node", + &collected.stream, + &collected.channel_name, + &collected.frame, + ); +} + fn ingest_dashboard_frame( dashboard: Option<&DashboardSupport>, stream: &StreamId, @@ -5506,7 +6044,7 @@ mod tests { .collect() } #[test] - fn pipeline_weight_load_scheduler_keeps_one_active_stage() { + fn pipeline_weight_load_scheduler_keeps_all_unloaded_stages_pending() { let model = TempModelFile::with_metadata( "seven-stage-scheduler.gguf", TestGgufMetadata { @@ -5518,24 +6056,30 @@ mod tests { let plan = config.build_run_plan().expect("seven-stage plan builds"); let mut loaded = BTreeSet::new(); - let first = - next_pipeline_weight_load_stage(&plan, &loaded, None).expect("first stage selected"); - assert_eq!(first.stage_index, 6); + assert_eq!( + pending_pipeline_weight_load_stages(&plan, &loaded) + .iter() + .map(|stage| stage.stage_index) + .collect::>(), + vec![0, 1, 2, 3, 4, 5, 6] + ); - let resent = next_pipeline_weight_load_stage(&plan, &loaded, Some(first.stage_index)) - .expect("active stage is resent before it loads"); - assert_eq!(resent.stage_index, 6); + loaded.insert(0); + loaded.insert(3); + assert_eq!( + pending_pipeline_weight_load_stages(&plan, &loaded) + .iter() + .map(|stage| stage.stage_index) + .collect::>(), + vec![1, 2, 4, 5, 6] + ); - for expected_stage in (0..7).rev() { - let active = next_pipeline_weight_load_stage(&plan, &loaded, None) - .expect("next unloaded stage selected"); - assert_eq!(active.stage_index, expected_stage); - loaded.insert(expected_stage); + for stage_index in 0..7 { + loaded.insert(stage_index); } - assert!( - next_pipeline_weight_load_stage(&plan, &loaded, None).is_none(), - "all stages loaded should leave no active load" + pending_pipeline_weight_load_stages(&plan, &loaded).is_empty(), + "all stages loaded should leave no pending load" ); } @@ -7631,12 +8175,15 @@ bootstrap_command = "/run" }) .collect::>(); + let stage_shard_plans = BTreeMap::new(); + for stage in &plan.stages { let wire = stage_provision_wire_from_plan( &plan, stage.stage_index, &readies, &coordinator, + &stage_shard_plans, ) .expect("stage provision wire builds from plan"); let inbound = wire.inbound_edge.as_ref().expect("planned inbound edge"); @@ -7759,6 +8306,375 @@ bootstrap_command = "/run" } } + #[test] + fn load_liveness_tracks_worker_progress_and_missing_gpu_samples() { + let mut progress = BTreeMap::new(); + let frame = CollectedDatastreamFrame { + stream: StreamId::new(NodeId::new("2"), Lifetime(17)), + channel_name: "mvp.worker.weights".to_owned(), + frame: Frame::new( + ChannelId(1), + datastream::Position(1), + serde_json::to_vec(&json!({ + "type": "GgufDownloadProgress", + "bytes_done": 4_218_u64, + "bytes_total": 4_683_u64, + })) + .expect("serialize progress frame"), + ), + }; + + update_load_progress_from_frame(&mut progress, &frame, Instant::now()); + let detail = stage_load_liveness_detail( + progress.get(&2), + 0, + 2, + Some(MemberState::Dead), + Some(DistNodeId([3; 32])), + None, + false, + "heartbeat_missed", + ); + + assert_eq!( + detail + .pointer("/load_progress/phase") + .and_then(Value::as_str), + Some("prefetching_model") + ); + assert_eq!( + detail + .pointer("/load_progress/bytes_done") + .and_then(Value::as_u64), + Some(4_218) + ); + assert_eq!( + detail + .pointer("/load_progress/bytes_total") + .and_then(Value::as_u64), + Some(4_683) + ); + assert_eq!( + detail.pointer("/classification").and_then(Value::as_str), + Some("heartbeat_missed") + ); + assert_eq!( + detail.pointer("/host_gpu_missing").and_then(Value::as_bool), + Some(true) + ); + + let gpu_frame = CollectedDatastreamFrame { + stream: StreamId::new(NodeId::new("2"), Lifetime(17)), + channel_name: "host.gpu".to_owned(), + frame: Frame::new(ChannelId(2), datastream::Position(2), b"{}".to_vec()), + }; + update_load_progress_from_frame(&mut progress, &gpu_frame, Instant::now()); + let detail = stage_load_liveness_detail( + progress.get(&2), + 0, + 2, + Some(MemberState::Alive), + Some(DistNodeId([3; 32])), + None, + true, + "waiting", + ); + assert_eq!( + detail.pointer("/host_gpu_missing").and_then(Value::as_bool), + Some(false) + ); + } + + #[test] + fn stage_shard_events_update_load_progress_and_liveness_phase() { + let mut progress = BTreeMap::new(); + let now = Instant::now(); + let frame = CollectedDatastreamFrame { + stream: StreamId::new(NodeId::new("7"), Lifetime(17)), + channel_name: "mvp.worker.weights".to_owned(), + frame: Frame::new( + ChannelId(1), + datastream::Position(1), + serde_json::to_vec(&json!({ + "type":"NodeEvent", + "phase":"stage_shard_fetch", + "status":"event", + "run_id":17, + "node_id":7, + "stage_index":3, + "detail":{ + "event":{ + "type":"StageShardRangeFetchReady", + "stage_index":3, + "range_index":1, + "range_count":4, + "tensor_index":2, + "tensor_count":9, + "bytes_done":384_u64, + "bytes_total":1024_u64 + } + } + })) + .expect("serialize stage shard progress frame"), + ), + }; + + update_load_progress_from_frame(&mut progress, &frame, now); + let entry = progress.get(&7).expect("stage shard progress tracked"); + assert_eq!(entry.stage_index, Some(3)); + assert_eq!(entry.phase.as_deref(), Some("fetching_stage_shard")); + assert_eq!( + entry.last_worker_event.as_deref(), + Some("StageShardRangeFetchReady") + ); + assert_eq!(entry.bytes_done, Some(384)); + assert_eq!(entry.bytes_total, Some(1024)); + assert_eq!(entry.last_progress, Some(now)); + + let ready_frame = CollectedDatastreamFrame { + stream: StreamId::new(NodeId::new("7"), Lifetime(17)), + channel_name: "mvp.worker.weights".to_owned(), + frame: Frame::new( + ChannelId(1), + datastream::Position(2), + serde_json::to_vec(&json!({ + "type":"NodeEvent", + "phase":"stage_shard_fetch", + "status":"event", + "run_id":17, + "node_id":7, + "stage_index":3, + "detail":{ + "event":{ + "type":"StageShardReady", + "stage_index":3, + "bytes_done":1024_u64, + "bytes_total":1024_u64 + } + } + })) + .expect("serialize stage shard ready frame"), + ), + }; + update_load_progress_from_frame(&mut progress, &ready_frame, Instant::now()); + let detail = stage_load_liveness_detail( + progress.get(&7), + 3, + 7, + Some(MemberState::Dead), + Some(DistNodeId([4; 32])), + None, + false, + "heartbeat_missed", + ); + assert_eq!( + detail + .pointer("/load_progress/phase") + .and_then(Value::as_str), + Some("stage_shard_ready") + ); + assert_eq!( + detail + .pointer("/load_progress/last_worker_event") + .and_then(Value::as_str), + Some("StageShardReady") + ); + assert!( + detail + .pointer("/load_progress/last_progress_age_ms") + .and_then(Value::as_u64) + .is_some() + ); + + let cache_frame = CollectedDatastreamFrame { + stream: StreamId::new(NodeId::new("8"), Lifetime(17)), + channel_name: "mvp.worker.weights".to_owned(), + frame: Frame::new( + ChannelId(1), + datastream::Position(3), + serde_json::to_vec(&json!({ + "type":"NodeEvent", + "phase":"stage_shard_fetch", + "status":"event", + "run_id":17, + "node_id":8, + "stage_index":4, + "detail":{"event":{"type":"StageShardCacheReady","stage_index":4,"cache_hit":true}} + })) + .expect("serialize stage shard cache frame"), + ), + }; + update_load_progress_from_frame(&mut progress, &cache_frame, Instant::now()); + assert_eq!( + progress.get(&8).and_then(|entry| entry.phase.as_deref()), + Some("stage_shard_cache_ready") + ); + } + + #[test] + fn stage_shard_plan_summary_events_include_fetch_facts() { + static NEXT_TEMP_FILE: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0); + let suffix = NEXT_TEMP_FILE.fetch_add(1, std::sync::atomic::Ordering::Relaxed); + let path = std::env::temp_dir().join(format!( + "mvp-stage-shard-plan-summary-{}-{suffix}.jsonl", + std::process::id() + )); + let _ = std::fs::remove_file(&path); + let plan0 = test_stage_shard_plan(0, 2_000, 100, vec![(800, 400), (1_600, 200)]); + let plan1 = test_stage_shard_plan(1, 2_000, 120, vec![(1_000, 300)]); + let plans = BTreeMap::from([(0, plan0.clone()), (1, plan1.clone())]); + + let mut datastream = OrchDatastream::new(41, Some(&path)).expect("datastream opens"); + emit_stage_shard_plan_summaries(None, &mut datastream, 41, 9, &plans); + drop(datastream); + + let contents = std::fs::read_to_string(&path).expect("read summary archive"); + let _ = std::fs::remove_file(&path); + let summaries = contents + .lines() + .map(|line| serde_json::from_str::(line).expect("archive line is json")) + .filter(|record| { + record.get("channel").and_then(Value::as_str) == Some(MVP_ORCH_BOOTSTRAP) + }) + .map(|record| { + let payload = record + .pointer("/payload/value") + .and_then(Value::as_str) + .expect("bootstrap payload is text"); + serde_json::from_str::(payload).expect("bootstrap payload is json") + }) + .filter(|event| event.get("phase").and_then(Value::as_str) == Some("stage_shard_plan")) + .collect::>(); + assert_eq!(summaries.len(), 2); + + let first = summaries + .iter() + .find(|event| event.pointer("/detail/stage_index").and_then(Value::as_u64) == Some(0)) + .expect("stage 0 summary emitted"); + assert_eq!( + first + .pointer("/detail/planned_fetch_bytes") + .and_then(Value::as_u64), + Some(plan0.planned_fetch_bytes()) + ); + assert_eq!( + first + .pointer("/detail/source_total_bytes") + .and_then(Value::as_u64), + Some(plan0.source_total_bytes) + ); + assert_eq!( + first + .pointer("/detail/tensor_count") + .and_then(Value::as_u64), + Some(plan0.tensors.len() as u64) + ); + assert_eq!( + first.pointer("/detail/range_count").and_then(Value::as_u64), + Some(plan0.planned_range_count() as u64) + ); + assert!( + first + .pointer("/detail/planned_fraction") + .and_then(Value::as_f64) + .expect("planned fraction is numeric") + < 1.0 + ); + assert!( + plan0.planned_fetch_bytes() < plan0.source_total_bytes, + "nontrivial stage shard fetches less than the source GGUF" + ); + } + + #[test] + fn stage_provision_dispatch_suppresses_active_progress_and_recovers_when_stale() { + let now = Instant::now(); + assert_eq!( + stage_provision_dispatch(None, 0, None, now), + StageProvisionDispatch::Send("initial") + ); + assert_eq!( + stage_provision_dispatch(None, 1, Some(now - Duration::from_secs(15)), now), + StageProvisionDispatch::Send("no_progress_after_send") + ); + + let active = StageLoadProgress { + node_id: 2, + stage_index: Some(0), + phase: Some("fetching_stage_shard".to_owned()), + last_progress: Some(now - Duration::from_secs(5)), + ..StageLoadProgress::default() + }; + assert_eq!( + stage_provision_dispatch(Some(&active), 1, Some(now - Duration::from_secs(15)), now), + StageProvisionDispatch::Suppress("active_progress") + ); + + let stale = StageLoadProgress { + last_progress: Some(now - STAGE_PROVISION_ACTIVE_RESEND_AFTER - Duration::from_secs(1)), + ..active.clone() + }; + assert_eq!( + stage_provision_dispatch(Some(&stale), 1, Some(now - Duration::from_secs(15)), now), + StageProvisionDispatch::Suppress("recent_stale_progress_resend") + ); + assert_eq!( + stage_provision_dispatch( + Some(&stale), + 1, + Some(now - STAGE_PROVISION_ACTIVE_RESEND_AFTER - Duration::from_secs(1)), + now, + ), + StageProvisionDispatch::Send("stale_progress") + ); + + let failed = StageLoadProgress { + phase: Some("failed".to_owned()), + failure_reason: Some("cache write failed".to_owned()), + ..active + }; + assert_eq!( + stage_provision_dispatch(Some(&failed), 1, Some(now), now), + StageProvisionDispatch::Suppress("worker_load_failed") + ); + } + + fn test_stage_shard_plan( + stage_index: u32, + source_total_bytes: u64, + metadata_end: u64, + ranges: Vec<(u64, u64)>, + ) -> StageShardPlan { + StageShardPlan { + source: GgufSource::HuggingFaceGguf { + repo: "org/repo".to_owned(), + file: "model.gguf".to_owned(), + revision: None, + }, + stage_index, + stage_count: 2, + layer_start: stage_index * 2, + layer_end_exclusive: stage_index * 2 + 2, + metadata_count: 1, + metadata_end, + data_start: 512, + alignment: 32, + source_total_bytes, + tensors: vec![crate::gguf_shard::StageShardTensor { + name: format!("blk.{stage_index}.attn_q.weight"), + dims: vec![2, 2], + ggml_type: 0, + source_offset: 0, + byte_len: 16, + }], + merged_tensor_ranges: ranges + .into_iter() + .map(|(start, len)| crate::gguf_shard::ByteRange { start, len }) + .collect(), + cache_key: format!("test-stage-{stage_index}"), + } + } + #[derive(Default)] struct FakeProvisionPlugin { stopped: Vec, diff --git a/crates/mvp-system/src/provisioning.rs b/crates/mvp-system/src/provisioning.rs index 2d67b06..deef89d 100644 --- a/crates/mvp-system/src/provisioning.rs +++ b/crates/mvp-system/src/provisioning.rs @@ -138,6 +138,20 @@ pub trait ProvisionPlugin: Send { sink: PluginSink, ) -> Result; + fn start_nodes( + &mut self, + specs: Vec, + sink: PluginSink, + ) -> Vec<(NodeProvisionSpec, Result)> { + specs + .into_iter() + .map(|spec| { + let result = self.start_node(spec.clone(), sink.clone()); + (spec, result) + }) + .collect() + } + fn complete_bootstrap(&mut self, handle: &PluginNodeHandle) -> Result<(), String>; fn stop_node(&mut self, handle: &PluginNodeHandle) -> Result<(), String>; diff --git a/crates/mvp-system/src/stage_controller.rs b/crates/mvp-system/src/stage_controller.rs index 84a1ba4..ba7bdb4 100644 --- a/crates/mvp-system/src/stage_controller.rs +++ b/crates/mvp-system/src/stage_controller.rs @@ -1,3 +1,4 @@ +use crate::gguf_shard::StageShardPlan; use crate::run_plan::{GgufSource, TokenizerSource}; #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] @@ -97,6 +98,7 @@ pub struct ProvisionStage { pub inbound: EdgeProvision, pub outbound: EdgeProvision, pub weight_source: WeightSource, + pub shard_plan: Option, } #[derive(Clone, Debug, PartialEq, Eq)] @@ -223,6 +225,7 @@ pub enum StageCommand { LoadWeights { source: WeightSource, range: LayerRange, + shard_plan: Option, }, RewireEdge { edge_id: EdgeId, @@ -366,6 +369,7 @@ impl StageController { self.commands.push(StageCommand::LoadWeights { source: provision.weight_source.clone(), range: provision.layer_range, + shard_plan: provision.shard_plan.clone(), }); self.provision = Some(provision); } diff --git a/crates/mvp-system/src/tests/local_mock/environment.rs b/crates/mvp-system/src/tests/local_mock/environment.rs index 3a1e468..b3cb048 100644 --- a/crates/mvp-system/src/tests/local_mock/environment.rs +++ b/crates/mvp-system/src/tests/local_mock/environment.rs @@ -802,6 +802,7 @@ impl LocalMockCluster { provision.gguf_source, provision.tokenizer, ), + shard_plan: None, } } diff --git a/crates/mvp-system/src/tests/stage_controller_guarantees.rs b/crates/mvp-system/src/tests/stage_controller_guarantees.rs index bf2ea25..54ad9dd 100644 --- a/crates/mvp-system/src/tests/stage_controller_guarantees.rs +++ b/crates/mvp-system/src/tests/stage_controller_guarantees.rs @@ -27,6 +27,7 @@ fn valid_provision() -> stage::ProvisionStage { inbound: stage::EdgeProvision::inbound(stage::EdgeId(7001)), outbound: stage::EdgeProvision::outbound(stage::EdgeId(7002)), weight_source: stage::WeightSource::embedded_gguf("model", "model.gguf"), + shard_plan: None, } } diff --git a/crates/mvp-system/src/tests/vastai_provisioning_guarantees.rs b/crates/mvp-system/src/tests/vastai_provisioning_guarantees.rs index 9ce0f6e..38b6818 100644 --- a/crates/mvp-system/src/tests/vastai_provisioning_guarantees.rs +++ b/crates/mvp-system/src/tests/vastai_provisioning_guarantees.rs @@ -1,5 +1,6 @@ use std::collections::VecDeque; use std::sync::Arc; +use std::time::Duration; use mvp_system::node_provisioning as provision; use mvp_system::node_provisioning::ProviderPlugin; @@ -10,7 +11,7 @@ use mvp_system::vastai_provisioning::{ BootstrapStopReason, VastAiBootstrapLauncher, VastAiLeaseClient, VastAiProviderPlugin, VastAiProvisioningConfig, VastAiProvisioningPlugin, VastAiSshEndpoint, }; -use parking_lot::Mutex; +use parking_lot::{Condvar, Mutex}; use swactor_vastai::{LifecyclePolicy, ProvisionRequest, ProvisionedInstance, SelectionPolicy}; #[derive(Default)] @@ -28,7 +29,7 @@ fn sink() -> PluginSink { PluginSink::new(Arc::new(RecordingSink::default())) } -#[derive(Default)] +#[derive(Clone, Default)] struct FakeLeaseClient { requests: Vec, endpoint_lookups: Vec<(u64, String, String)>, @@ -37,6 +38,7 @@ struct FakeLeaseClient { destroy_result: Option>, next_contract_id: u64, host_ids: VecDeque>, + first_wave_plan: Vec>, } impl FakeLeaseClient { @@ -63,6 +65,15 @@ impl VastAiLeaseClient for FakeLeaseClient { }) } + fn plan_first_wave_offers( + &mut self, + requests: &[ProvisionRequest], + ) -> Result>, String> { + let mut plan = self.first_wave_plan.clone(); + plan.resize(requests.len(), None); + Ok(plan) + } + fn ssh_endpoint( &mut self, contract_id: u64, @@ -118,6 +129,209 @@ impl VastAiBootstrapLauncher for FakeBootstrap { } } +#[derive(Clone)] +struct ParallelLeaseClient { + state: Arc>, + gate: Arc, +} + +struct ParallelLeaseGate { + target: usize, + started: Mutex, + all_started: Condvar, +} + +#[derive(Default)] +struct ParallelLeaseState { + requests: Vec, + endpoint_lookups: Vec, + destroyed: Vec, + first_endpoint_request_count: Option, + next_contract_id: u64, + host_ids: VecDeque>, + first_wave_plan_requests: usize, + first_wave_plan: Vec>, +} + +impl ParallelLeaseClient { + fn new(target: usize) -> Self { + Self { + state: Arc::new(Mutex::new(ParallelLeaseState { + next_contract_id: 100, + host_ids: (0..target) + .map(|index| Some(10_000 + u64::try_from(index).unwrap())) + .collect(), + first_wave_plan: (0..target) + .map(|index| Some(9_000 + u64::try_from(index).unwrap())) + .collect(), + ..ParallelLeaseState::default() + })), + gate: Arc::new(ParallelLeaseGate { + target, + started: Mutex::new(0), + all_started: Condvar::new(), + }), + } + } +} + +impl VastAiLeaseClient for ParallelLeaseClient { + fn plan_first_wave_offers( + &mut self, + requests: &[ProvisionRequest], + ) -> Result>, String> { + let mut state = self.state.lock(); + state.first_wave_plan_requests = requests.len(); + let mut plan = state.first_wave_plan.clone(); + plan.resize(requests.len(), None); + Ok(plan) + } + + fn provision_one(&mut self, request: ProvisionRequest) -> Result { + { + self.state.lock().requests.push(request); + } + let mut started = self.gate.started.lock(); + *started += 1; + if *started < self.gate.target { + let wait = self + .gate + .all_started + .wait_for(&mut started, Duration::from_secs(2)); + assert!( + !wait.timed_out(), + "all concurrent Vast.ai lease requests should start before any waits for SSH" + ); + } else { + self.gate.all_started.notify_all(); + } + drop(started); + + let mut state = self.state.lock(); + let contract_id = state.next_contract_id; + state.next_contract_id = state.next_contract_id.wrapping_add(1).max(1); + let host_id = state.host_ids.pop_front().unwrap_or(Some(77)); + Ok(ProvisionedInstance { + index: 0, + contract_id, + offer_id: 55, + host_id, + gpu_name: "RTX 4090".to_owned(), + gpu_ram: Some(24_000.0), + dph_total: 0.42, + }) + } + + fn ssh_endpoint( + &mut self, + contract_id: u64, + _label: &str, + _lifecycle: &LifecyclePolicy, + ssh_user: &str, + ) -> Result { + let mut state = self.state.lock(); + let request_count = state.requests.len(); + state + .first_endpoint_request_count + .get_or_insert(request_count); + state.endpoint_lookups.push(contract_id); + Ok(VastAiSshEndpoint { + host: "ssh5.vast.ai".to_owned(), + port: 22017, + user: ssh_user.to_owned(), + }) + } + + fn destroy_contract(&mut self, contract_id: u64) -> Result<(), String> { + self.state.lock().destroyed.push(contract_id); + Ok(()) + } +} + +#[derive(Clone)] +struct OutOfOrderLeaseClient { + state: Arc>, +} + +#[derive(Default)] +struct OutOfOrderLeaseState { + requests: Vec, + endpoint_lookups: Vec, + destroyed: Vec, + slow_node_ids: Vec, + endpoint_fail_node_ids: Vec, +} + +impl OutOfOrderLeaseClient { + fn new(slow_node_ids: Vec, endpoint_fail_node_ids: Vec) -> Self { + Self { + state: Arc::new(Mutex::new(OutOfOrderLeaseState { + slow_node_ids, + endpoint_fail_node_ids, + ..OutOfOrderLeaseState::default() + })), + } + } + + fn node_id_from_label(label: Option<&str>) -> u64 { + label + .and_then(|label| label.rsplit('-').next()) + .and_then(|node| node.parse::().ok()) + .expect("test request labels include node id suffix") + } +} + +impl VastAiLeaseClient for OutOfOrderLeaseClient { + fn provision_one(&mut self, request: ProvisionRequest) -> Result { + let node_id = Self::node_id_from_label(request.label.as_deref()); + let should_sleep = { + let mut state = self.state.lock(); + state.requests.push(node_id); + state.slow_node_ids.contains(&node_id) + }; + if should_sleep { + std::thread::sleep(Duration::from_millis(150)); + } + Ok(ProvisionedInstance { + index: 0, + contract_id: 1_000 + node_id, + offer_id: 55 + node_id, + host_id: Some(10_000 + node_id), + gpu_name: "RTX 4090".to_owned(), + gpu_ram: Some(24_000.0), + dph_total: 0.42, + }) + } + + fn ssh_endpoint( + &mut self, + contract_id: u64, + _label: &str, + _lifecycle: &LifecyclePolicy, + ssh_user: &str, + ) -> Result { + let node_id = contract_id - 1_000; + let should_fail = { + let mut state = self.state.lock(); + state.endpoint_lookups.push(node_id); + state.endpoint_fail_node_ids.contains(&node_id) + }; + if should_fail { + return Err("connection refused".to_owned()); + } + Ok(VastAiSshEndpoint { + host: "ssh5.vast.ai".to_owned(), + port: 22017, + user: ssh_user.to_owned(), + }) + } + + fn destroy_contract(&mut self, contract_id: u64) -> Result<(), String> { + self.state.lock().destroyed.push(contract_id); + Ok(()) + } +} + fn spec() -> NodeProvisionSpec { NodeProvisionSpec { run_id: 9, @@ -272,6 +486,171 @@ fn pipeline_starts_blacklist_hosts_already_leased_in_run() { plugin.stop_node(&first_handle).unwrap(); } +#[test] +fn failed_vastai_host_is_blacklisted_for_later_requests() { + let mut client = FakeLeaseClient::default().with_contract(100); + client.host_ids.extend([Some(77), Some(88)]); + client + .endpoint_results + .push_back(Err("connection refused".to_owned())); + let mut plugin = VastAiProvisioningPlugin::new(client, FakeBootstrap::default(), config()); + let first_error = plugin.start_node(spec(), sink()).unwrap_err(); + assert!(first_error.contains("connection refused")); + + let mut second = spec(); + second.node_id = 12; + second.stage_index = Some(3); + let second_handle = plugin.start_node(second, sink()).unwrap(); + + assert!( + plugin.client().requests[1] + .selection + .blacklist_hosts + .contains(&77), + "host that failed before runtime-ready must be excluded from later Vast.ai requests" + ); + plugin.stop_node(&second_handle).unwrap(); +} + +#[test] +fn vastai_start_nodes_starts_lease_requests_concurrently() { + let client = ParallelLeaseClient::new(4); + let state = client.state.clone(); + let mut plugin = VastAiProvisioningPlugin::new(client, FakeBootstrap::default(), config()); + let specs = (0..4) + .map(|index| { + let mut spec = spec(); + spec.node_id = 11 + index; + spec.stage_index = Some(u32::try_from(index).unwrap()); + spec + }) + .collect::>(); + + let results = plugin.start_nodes(specs, sink()); + + assert!(results.iter().all(|(_, result)| result.is_ok())); + assert_eq!(plugin.active_contract_count(), 4); + assert_eq!(plugin.bootstrap().starts.len(), 4); + let state = state.lock(); + assert_eq!(state.requests.len(), 4); + assert_eq!(state.endpoint_lookups.len(), 4); + assert_eq!( + state.first_endpoint_request_count, + Some(4), + "SSH lookup must not begin before every lease request has started" + ); +} + +#[test] +fn vastai_start_nodes_assigns_shared_first_wave_offer_plan() { + let client = ParallelLeaseClient::new(3); + let state = client.state.clone(); + let mut plugin = VastAiProvisioningPlugin::new(client, FakeBootstrap::default(), config()); + let specs = (0..3) + .map(|index| { + let mut spec = spec(); + spec.node_id = 11 + index; + spec.stage_index = Some(u32::try_from(index).unwrap()); + spec + }) + .collect::>(); + + let results = plugin.start_nodes(specs, sink()); + + assert!(results.iter().all(|(_, result)| result.is_ok())); + let state = state.lock(); + assert_eq!(state.first_wave_plan_requests, 3); + let mut assigned = state + .requests + .iter() + .map(|request| { + ( + request.label.clone().expect("request label"), + request.preferred_offer_id, + ) + }) + .collect::>(); + assigned.sort_by(|left, right| left.0.cmp(&right.0)); + assert_eq!( + assigned + .into_iter() + .map(|(_, preferred)| preferred) + .collect::>(), + vec![Some(9_000), Some(9_001), Some(9_002)], + "per-node requests should carry the coordinated first-wave offer plan" + ); +} + +#[test] +fn vastai_start_nodes_bootstraps_fast_completion_before_earlier_slow_node() { + let client = OutOfOrderLeaseClient::new(vec![11], Vec::new()); + let state = client.state.clone(); + let mut plugin = VastAiProvisioningPlugin::new(client, FakeBootstrap::default(), config()); + let mut slow = spec(); + slow.node_id = 11; + slow.stage_index = Some(0); + let mut fast = spec(); + fast.node_id = 12; + fast.stage_index = Some(1); + + let results = plugin.start_nodes(vec![slow, fast], sink()); + + assert!(results.iter().all(|(_, result)| result.is_ok())); + assert_eq!( + plugin + .bootstrap() + .starts + .iter() + .map(|(spec, _)| spec.node_id) + .collect::>(), + vec![12, 11], + "later fast completion should bootstrap before earlier slow completion" + ); + assert_eq!(state.lock().endpoint_lookups.len(), 2); +} + +#[test] +fn vastai_start_nodes_one_failure_does_not_block_completed_node_bootstrap() { + let client = OutOfOrderLeaseClient::new(vec![11], vec![11]); + let state = client.state.clone(); + let mut plugin = VastAiProvisioningPlugin::new(client, FakeBootstrap::default(), config()); + let mut slow_failure = spec(); + slow_failure.node_id = 11; + slow_failure.stage_index = Some(0); + let mut fast_success = spec(); + fast_success.node_id = 12; + fast_success.stage_index = Some(1); + + let results = plugin.start_nodes(vec![slow_failure, fast_success], sink()); + + assert!( + results[0] + .1 + .as_ref() + .unwrap_err() + .contains("connection refused") + ); + assert!( + results[0] + .1 + .as_ref() + .unwrap_err() + .contains("class=connection_refused") + ); + assert!(results[1].1.is_ok()); + assert_eq!( + plugin + .bootstrap() + .starts + .iter() + .map(|(spec, _)| spec.node_id) + .collect::>(), + vec![12], + "successful completed node should bootstrap even though another node fails" + ); + assert_eq!(state.lock().destroyed.as_slice(), &[1_011]); +} + #[test] fn stop_destroys_known_vastai_contract_exactly_once() { let mut plugin = VastAiProvisioningPlugin::new( @@ -293,7 +672,7 @@ fn stop_destroys_known_vastai_contract_exactly_once() { } #[test] -fn vastai_complete_bootstrap_keeps_log_tail_until_node_stop() { +fn vastai_complete_bootstrap_stops_optional_log_tail_before_node_stop() { let mut plugin = VastAiProvisioningPlugin::new( FakeLeaseClient::default().with_contract(100), FakeBootstrap::default(), @@ -303,9 +682,9 @@ fn vastai_complete_bootstrap_keeps_log_tail_until_node_stop() { plugin.complete_bootstrap(&handle).unwrap(); - assert!( - plugin.bootstrap().stops.is_empty(), - "runtime-ready completion should keep the SSH log tail alive" + assert_eq!( + plugin.bootstrap().stops, + vec![(1, BootstrapStopReason::RuntimeReady)] ); assert_eq!(plugin.client().destroyed, Vec::::new()); assert_eq!(plugin.active_contract_count(), 1); @@ -315,7 +694,7 @@ fn vastai_complete_bootstrap_keeps_log_tail_until_node_stop() { assert_eq!(plugin.client().destroyed, vec![100]); assert_eq!( plugin.bootstrap().stops, - vec![(1, BootstrapStopReason::NodeStop)] + vec![(1, BootstrapStopReason::RuntimeReady)] ); assert_eq!(plugin.active_contract_count(), 0); } diff --git a/crates/mvp-system/src/tests/worker_edge_adapter_guarantees.rs b/crates/mvp-system/src/tests/worker_edge_adapter_guarantees.rs index c832612..47ac844 100644 --- a/crates/mvp-system/src/tests/worker_edge_adapter_guarantees.rs +++ b/crates/mvp-system/src/tests/worker_edge_adapter_guarantees.rs @@ -58,6 +58,7 @@ fn provision_wire() -> StageProvisionWire { model_id: "smollm2-135m-q4".to_owned(), gguf_source: GgufSource::LocalPath("/models/smollm.gguf".to_owned()), tokenizer: TokenizerSource::EmbeddedGguf, + stage_shard_plan: None, } } @@ -139,6 +140,7 @@ fn provision_wire_from_plan(plan: &run_plan::RunPlan, stage_index: u32) -> Stage model_id: provision.model.model_id, gguf_source: provision.gguf_source, tokenizer: provision.tokenizer, + stage_shard_plan: None, } } diff --git a/crates/mvp-system/src/vastai_provisioning.rs b/crates/mvp-system/src/vastai_provisioning.rs index a5279f9..c0481cd 100644 --- a/crates/mvp-system/src/vastai_provisioning.rs +++ b/crates/mvp-system/src/vastai_provisioning.rs @@ -5,6 +5,7 @@ use std::process::{Child, Command, Stdio}; use std::sync::{ Arc, atomic::{AtomicBool, Ordering}, + mpsc, }; use std::time::Duration; @@ -12,7 +13,9 @@ use datastream::DatastreamProducer; use serde::{Deserialize, Serialize}; use swactor::actor::{ActorAddress, ActorInterface}; use swactor::runtime::{Ctx, Runtime}; -use swactor_vastai::{LifecyclePolicy, ProvisionRequest, ProvisionedInstance, SelectionPolicy}; +use swactor_vastai::{ + LifecyclePolicy, ProvisionRequest, ProvisionedInstance, SelectionPolicy, classify_vastai_error, +}; use crate::bootstrap_datastream::{BootstrapDatastreamBridge, node_stream_id}; use crate::node_provisioning::{ @@ -59,6 +62,12 @@ pub struct VastAiSshEndpoint { pub trait VastAiLeaseClient: Send { fn provision_one(&mut self, request: ProvisionRequest) -> Result; + fn plan_first_wave_offers( + &mut self, + requests: &[ProvisionRequest], + ) -> Result>, String> { + Ok(vec![None; requests.len()]) + } fn ssh_endpoint( &mut self, @@ -94,6 +103,12 @@ impl ToolsVastAiLeaseClient { } } +impl Clone for ToolsVastAiLeaseClient { + fn clone(&self) -> Self { + Self::new(self.client.clone()).expect("clone VastAI lease client runtime") + } +} + impl VastAiLeaseClient for ToolsVastAiLeaseClient { fn provision_one(&mut self, request: ProvisionRequest) -> Result { let fleet = self.runtime.block_on(self.client.provision(request))?; @@ -107,6 +122,31 @@ impl VastAiLeaseClient for ToolsVastAiLeaseClient { Ok(instances.remove(0)) } + fn plan_first_wave_offers( + &mut self, + requests: &[ProvisionRequest], + ) -> Result>, String> { + let Some(first) = requests.first() else { + return Ok(Vec::new()); + }; + let pool = self.runtime.block_on( + self.client + .search_offers(&first.selection, requests.len() as u32), + )?; + let planned = swactor_vastai::plan_distinct_host_first_wave( + &pool, + requests.len() as u32, + &first.selection.blacklist_hosts, + &[], + ); + let mut out = planned + .into_iter() + .map(|offer| Some(offer.id)) + .collect::>(); + out.resize(requests.len(), None); + Ok(out) + } + fn ssh_endpoint( &mut self, contract_id: u64, @@ -240,6 +280,7 @@ where disk_gb: spec.shape.disk_gb, env, per_instance_env: vec![BTreeMap::new()], + preferred_offer_id: None, onstart: self.config.onstart.clone(), selection: self.selection_for(spec), lifecycle: self.config.lifecycle.clone(), @@ -602,13 +643,27 @@ fn spawn_retrying_ssh_bootstrap( backoff.as_secs() ), }); - std::thread::sleep(backoff); + if !sleep_ssh_backoff(backoff, &stopping) { + return; + } backoff = next_ssh_backoff(backoff); attempt += 1; } }); } +fn sleep_ssh_backoff(backoff: Duration, stopping: &AtomicBool) -> bool { + let deadline = std::time::Instant::now() + backoff; + while std::time::Instant::now() < deadline { + if stopping.load(Ordering::SeqCst) { + return false; + } + let remaining = deadline.saturating_duration_since(std::time::Instant::now()); + std::thread::sleep(std::cmp::min(remaining, Duration::from_millis(50))); + } + !stopping.load(Ordering::SeqCst) +} + fn spawn_ssh_bootstrap_attempt( spec: &NodeProvisionSpec, endpoint: &VastAiSshEndpoint, @@ -677,6 +732,7 @@ where config: VastAiProvisioningConfig, bootstrap_producer: Option, leased_host_ids: BTreeSet, + failed_host_ids: BTreeSet, next_handle_id: u64, nodes: BTreeMap>, } @@ -700,6 +756,7 @@ where bootstrap_producer: None, next_handle_id: 1, leased_host_ids: BTreeSet::new(), + failed_host_ids: BTreeSet::new(), nodes: BTreeMap::new(), } } @@ -751,7 +808,11 @@ where env.insert("SSH_PUBLIC_KEY".to_owned(), key.to_owned()); } let mut selection = self.config.selection.clone(); - for host_id in &self.leased_host_ids { + for host_id in self + .leased_host_ids + .iter() + .chain(self.failed_host_ids.iter()) + { if !selection.blacklist_hosts.contains(host_id) { selection.blacklist_hosts.push(*host_id); } @@ -764,6 +825,7 @@ where disk_gb: self.config.disk_gb, env, per_instance_env: vec![BTreeMap::new()], + preferred_offer_id: None, onstart: self.config.onstart.clone(), selection, lifecycle: self.config.lifecycle.clone(), @@ -779,9 +841,25 @@ where } } +fn classified_start_error(reason: String) -> String { + let class = classify_vastai_error(&reason).as_str(); + format!("{reason} [class={class}]") +} + +struct VastAiStartedLease { + label: String, + instance: ProvisionedInstance, + endpoint: VastAiSshEndpoint, +} + +struct VastAiBatchStartError { + reason: String, + failed_host_id: Option, +} + impl ProvisionPlugin for VastAiProvisioningPlugin where - C: VastAiLeaseClient + 'static, + C: VastAiLeaseClient + Clone + 'static, B: VastAiBootstrapLauncher + 'static, { fn start_node( @@ -801,10 +879,9 @@ where }); let request = self.build_request(&spec, label.clone()); - let instance = self - .client - .provision_one(request) - .map_err(|e| format!("vastai provision node {}: {e}", spec.node_id))?; + let instance = self.client.provision_one(request).map_err(|e| { + classified_start_error(format!("vastai provision node {}: {e}", spec.node_id)) + })?; sink.observe(PluginObservation::ProviderLine { run_id: spec.run_id, node_id: spec.node_id, @@ -840,9 +917,15 @@ where ) { Ok(endpoint) => endpoint, Err(error) => { + if let Some(host_id) = instance.host_id { + self.failed_host_ids.insert(host_id); + } return Err(self.cleanup_contract_after_start_error( instance.contract_id, - format!("vastai SSH endpoint node {}: {error}", spec.node_id), + classified_start_error(format!( + "vastai SSH endpoint node {}: {error}", + spec.node_id + )), )); } }; @@ -869,9 +952,15 @@ where ) { Ok(handle) => handle, Err(error) => { + if let Some(host_id) = instance.host_id { + self.failed_host_ids.insert(host_id); + } return Err(self.cleanup_contract_after_start_error( instance.contract_id, - format!("vastai bootstrap node {}: {error}", spec.node_id), + classified_start_error(format!( + "vastai bootstrap node {}: {error}", + spec.node_id + )), )); } }; @@ -897,7 +986,256 @@ where Ok(handle) } - fn complete_bootstrap(&mut self, _handle: &PluginNodeHandle) -> Result<(), String> { + fn start_nodes( + &mut self, + specs: Vec, + sink: PluginSink, + ) -> Vec<(NodeProvisionSpec, Result)> { + if specs.len() <= 1 { + return specs + .into_iter() + .map(|spec| { + let result = self.start_node(spec.clone(), sink.clone()); + (spec, result) + }) + .collect(); + } + + let mut results = (0..specs.len()).map(|_| None).collect::>(); + let mut start_inputs = Vec::new(); + for (index, spec) in specs.into_iter().enumerate() { + if !spec.mounts.is_empty() { + results[index] = Some(( + spec, + Err("vastai provider does not support host file mounts".to_owned()), + )); + continue; + } + let stream_id = node_stream_id(spec.run_id, spec.node_id); + let label = self.label_for(&spec); + sink.observe(PluginObservation::ProviderLine { + run_id: spec.run_id, + node_id: spec.node_id, + line: format!("vastai provisioning label={label} stream={stream_id}"), + }); + let request = self.build_request(&spec, label.clone()); + start_inputs.push((index, spec, label, request)); + } + + let request_plan = start_inputs + .iter() + .map(|(_, _, _, request)| request.clone()) + .collect::>(); + let offer_plan = match self.client.plan_first_wave_offers(&request_plan) { + Ok(plan) => plan, + Err(error) => { + for (_, spec, _, _) in &start_inputs { + sink.observe(PluginObservation::ProviderLine { + run_id: spec.run_id, + node_id: spec.node_id, + line: format!( + "vastai first-wave offer planning failed; falling back to per-node selection: {error}" + ), + }); + } + vec![None; start_inputs.len()] + } + }; + + let (completion_tx, completion_rx) = mpsc::channel(); + for (plan_index, (index, spec, label, mut request)) in start_inputs.into_iter().enumerate() + { + request.preferred_offer_id = offer_plan.get(plan_index).copied().flatten(); + if let Some(offer_id) = request.preferred_offer_id { + sink.observe(PluginObservation::ProviderLine { + run_id: spec.run_id, + node_id: spec.node_id, + line: serde_json::json!({ + "type": "VastAiFirstWaveOfferPlanned", + "run_id": spec.run_id, + "node_id": spec.node_id, + "label": &label, + "offer_id": offer_id, + }) + .to_string(), + }); + } + let mut client = self.client.clone(); + let config = self.config.clone(); + let worker_tx = completion_tx.clone(); + std::thread::spawn(move || { + let started = match client.provision_one(request) { + Ok(instance) => { + match client.ssh_endpoint( + instance.contract_id, + &label, + &config.lifecycle, + &config.ssh_user, + ) { + Ok(endpoint) => Ok(VastAiStartedLease { + label, + instance, + endpoint, + }), + Err(error) => { + let failed_host_id = instance.host_id; + let reason = match client.destroy_contract(instance.contract_id) { + Ok(()) => classified_start_error(format!( + "vastai SSH endpoint node {}: {error}", + spec.node_id + )), + Err(cleanup) => classified_start_error(format!( + "vastai SSH endpoint node {}: {error}; cleanup destroy {} failed: {cleanup}", + spec.node_id, instance.contract_id + )), + }; + Err(VastAiBatchStartError { + reason, + failed_host_id, + }) + } + } + } + Err(error) => Err(VastAiBatchStartError { + reason: classified_start_error(format!( + "vastai provision node {}: {error}", + spec.node_id + )), + failed_host_id: None, + }), + }; + let _ = worker_tx.send((index, spec, started)); + }); + } + drop(completion_tx); + + 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", + "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, + }) + .to_string(), + }); + sink.observe(PluginObservation::ProviderLine { + run_id: spec.run_id, + node_id: spec.node_id, + line: serde_json::json!({ + "type": "VastAiSshEndpointReady", + "run_id": spec.run_id, + "node_id": spec.node_id, + "contract_id": started.instance.contract_id, + "host": &started.endpoint.host, + "port": started.endpoint.port, + "user": &started.endpoint.user, + }) + .to_string(), + }); + + let bootstrap = match self.bootstrap.start_bootstrap( + spec.clone(), + started.endpoint, + sink.clone(), + self.bootstrap_producer.clone(), + ) { + Ok(handle) => handle, + Err(error) => { + if let Some(host_id) = started.instance.host_id { + self.failed_host_ids.insert(host_id); + } + let node_id = spec.node_id; + results[index] = Some(( + spec, + Err(self.cleanup_contract_after_start_error( + started.instance.contract_id, + classified_start_error(format!( + "vastai bootstrap node {node_id}: {error}" + )), + )), + )); + continue; + } + }; + + let host_id = started.instance.host_id; + if let Some(host_id) = host_id { + self.leased_host_ids.insert(host_id); + } + let handle = PluginNodeHandle { + id: self.next_handle_id, + provider_process_id: None, + }; + self.next_handle_id = self.next_handle_id.wrapping_add(1).max(1); + self.nodes.insert( + handle.id, + VastAiNode { + contract_id: started.instance.contract_id, + bootstrap: Some(bootstrap), + host_id, + }, + ); + results[index] = Some((spec, Ok(handle))); + } + Err(error) => { + if let Some(host_id) = error.failed_host_id { + self.failed_host_ids.insert(host_id); + } + results[index] = Some((spec, Err(error.reason))); + } + } + } + + results + .into_iter() + .enumerate() + .map(|(index, result)| { + result.unwrap_or_else(|| { + ( + NodeProvisionSpec { + run_id: 0, + node_id: u64::try_from(index).unwrap_or(u64::MAX), + stage_index: None, + image: String::new(), + env: Vec::new(), + args: Vec::new(), + mounts: Vec::new(), + }, + Err("vastai provision worker panicked".to_owned()), + ) + }) + }) + .collect() + } + + fn complete_bootstrap(&mut self, handle: &PluginNodeHandle) -> Result<(), String> { + let Some(node) = self.nodes.get_mut(&handle.id) else { + return Ok(()); + }; + if let Some(mut bootstrap) = node.bootstrap.take() { + self.bootstrap + .stop_bootstrap(&mut bootstrap, BootstrapStopReason::RuntimeReady); + } Ok(()) } @@ -1083,7 +1421,7 @@ mod tests { } #[test] - fn vastai_complete_bootstrap_keeps_log_tail_until_node_stop() { + fn vastai_complete_bootstrap_stops_optional_log_tail_before_node_stop() { let destroyed_contracts = Arc::new(Mutex::new(Vec::new())); let stop_reasons = Arc::new(Mutex::new(Vec::new())); let sink = PluginSink::new(Arc::new(ObservationSink::default())); @@ -1104,13 +1442,16 @@ mod tests { plugin .complete_bootstrap(&handle) .expect("runtime-ready bootstrap completion succeeds"); - assert!( - stop_reasons.lock().is_empty(), - "VastAI bootstrap SSH tail must remain alive for post-ready worker logs" + assert_eq!( + *stop_reasons.lock(), + vec![BootstrapStopReason::RuntimeReady] ); plugin.stop_node(&handle).expect("VastAI node stops"); - assert_eq!(*stop_reasons.lock(), vec![BootstrapStopReason::NodeStop]); + assert_eq!( + *stop_reasons.lock(), + vec![BootstrapStopReason::RuntimeReady] + ); assert_eq!(*destroyed_contracts.lock(), vec![42]); } @@ -1189,6 +1530,18 @@ mod tests { assert!(!args.iter().any(|arg| arg == "IdentitiesOnly=yes")); } + #[test] + fn ssh_bootstrap_backoff_sleep_observes_stop_without_waiting_full_backoff() { + let stopping = AtomicBool::new(true); + let started = std::time::Instant::now(); + + assert!(!sleep_ssh_backoff(Duration::from_secs(5), &stopping)); + assert!( + started.elapsed() < Duration::from_millis(250), + "stopped bootstrap backoff should not wait for the full retry delay" + ); + } + #[test] fn ssh_bootstrap_actor_stop_kills_child() { let runtime = Arc::new(swactor::runtime::Runtime::new( diff --git a/tools/vastai/src/lease.rs b/tools/vastai/src/lease.rs index 8f14dde..c2b5027 100644 --- a/tools/vastai/src/lease.rs +++ b/tools/vastai/src/lease.rs @@ -4,9 +4,9 @@ use std::time::Duration; use crate::config::{ENV_ASSUME_YES, truthy_env}; use crate::monitor::wait_for_running_with_policy; -use crate::pricing::{CostModel, plan_picks}; +use crate::pricing::CostModel; use crate::provision::create_instance; -use crate::search::select_offer_pool_with_policy; +use crate::search::{plan_distinct_host_first_wave, select_offer_pool_with_policy}; use crate::teardown::{destroy_instance_with_retry, rollback}; use crate::types::{ CreateInstanceRequest, InstanceInfo, Offer, ProvisionRequest, ProvisionedFleet, @@ -15,7 +15,7 @@ use crate::types::{ /// Print the planned lease + hourly cost and, on TTY, require y/N confirmation. pub fn confirm_lease(pool: &[Offer], num_instances: u32, cost: &CostModel) -> Result<(), String> { - let picks = plan_picks(pool, num_instances); + let picks = plan_distinct_host_first_wave(pool, num_instances, &[], &[]); let total_dph: f64 = picks.iter().map(|o| o.dph_total).sum(); let total_eff: f64 = picks.iter().map(|o| cost.effective_price(o)).sum(); @@ -80,9 +80,23 @@ fn next_eligible_offer<'a>( pool: &'a [Offer], tried_offer_ids: &[u64], used_host_ids: &HashSet, + failed_host_ids: &HashSet, + preferred_offer_id: Option, ) -> Option<&'a Offer> { + if let Some(offer_id) = preferred_offer_id { + if let Some(offer) = pool.iter().find(|o| o.id == offer_id) { + if !tried_offer_ids.contains(&offer.id) + && offer.host_id.is_none_or(|h| !failed_host_ids.contains(&h)) + { + return Some(offer); + } + } + } + pool.iter().find(|o| { - !tried_offer_ids.contains(&o.id) && o.host_id.map_or(true, |h| !used_host_ids.contains(&h)) + !tried_offer_ids.contains(&o.id) + && o.host_id + .is_none_or(|h| !used_host_ids.contains(&h) && !failed_host_ids.contains(&h)) }) } @@ -103,14 +117,22 @@ async fn provision_one( index: u32, tried_offer_ids: &mut Vec, used_host_ids: &mut HashSet, + failed_host_ids: &mut HashSet, + preferred_offer_id: Option, ) -> Result { let mut attempt = 1_u64; loop { - let offer = match next_eligible_offer(pool, tried_offer_ids, used_host_ids) { + let offer = match next_eligible_offer( + pool, + tried_offer_ids, + used_host_ids, + failed_host_ids, + preferred_offer_id.filter(|offer_id| !tried_offer_ids.contains(offer_id)), + ) { Some(o) => o.clone(), None => { return Err(format!( - "pool exhausted for index {index} (no untried offer on an unused host)" + "pool exhausted for index {index} (no untried offer outside failed hosts)" )); } }; @@ -201,10 +223,12 @@ pub async fn provision_fleet( confirm_lease(&pool, req.count, &CostModel::from_policy(&req.selection))?; } + let first_wave = + plan_distinct_host_first_wave(&pool, req.count, &req.selection.blacklist_hosts, &[]); let mut tried_offer_ids = Vec::new(); let mut created: Vec = Vec::with_capacity(req.count as usize); let mut used_host_ids = HashSet::new(); - + let mut failed_host_ids = HashSet::new(); for index in 0..req.count { match provision_one( client, @@ -215,6 +239,10 @@ pub async fn provision_fleet( index, &mut tried_offer_ids, &mut used_host_ids, + &mut failed_host_ids, + req.preferred_offer_id + .filter(|_| req.count == 1) + .or_else(|| first_wave.get(index as usize).map(|offer| offer.id)), ) .await { @@ -237,6 +265,9 @@ pub async fn provision_fleet( { Ok(_) => break, Err(e) => { + if let Some(host_id) = created[idx].host_id { + failed_host_ids.insert(host_id); + } eprintln!( "lease_chain: index {index} contract {cid} did not reach running: {e}" ); @@ -257,6 +288,8 @@ pub async fn provision_fleet( index, &mut tried_offer_ids, &mut used_host_ids, + &mut failed_host_ids, + None, ) .await { @@ -286,3 +319,136 @@ pub async fn provision_fleet( instances: created, }) } + +#[cfg(test)] +mod tests { + use std::collections::BTreeMap; + use std::time::Duration; + + use serde_json::json; + use wiremock::matchers::{method, path}; + use wiremock::{Mock, MockServer, ResponseTemplate}; + + use super::*; + use crate::types::{LifecyclePolicy, SelectionPolicy}; + + fn offer(id: u64, host_id: u64) -> serde_json::Value { + json!({ + "id": id, + "gpu_name": "RTX 4090", + "dph_total": id as f64 / 100.0, + "gpu_ram": 24_000.0, + "compute_cap": 890, + "geolocation": "US", + "internet_down_cost_per_tb": 0.0, + "internet_up_cost_per_tb": 0.0, + "host_id": host_id, + "verification": "verified" + }) + } + + fn request(count: u32) -> ProvisionRequest { + ProvisionRequest { + count, + image: "registry.example/mvp-worker:latest".to_owned(), + label: Some("lease-test".to_owned()), + disk_gb: 80, + env: BTreeMap::new(), + per_instance_env: Vec::new(), + preferred_offer_id: None, + onstart: None, + selection: SelectionPolicy { + drop_cheap_frac: 0.0, + ..SelectionPolicy::default() + }, + lifecycle: LifecyclePolicy { + lease_pace: Duration::ZERO, + poll_interval: Duration::from_millis(1), + state_timeout: Duration::from_millis(5), + }, + confirm_lease: false, + } + } + + #[tokio::test] + async fn replacement_excludes_failed_host_from_shared_offer_pool() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/api/v0/bundles/")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "offers": [ + offer(1, 10), + offer(2, 20), + offer(3, 10), + offer(4, 30) + ] + }))) + .mount(&server) + .await; + for (offer_id, contract_id) in [(1, 101), (2, 102), (4, 104)] { + Mock::given(method("PUT")) + .and(path(format!("/api/v0/asks/{offer_id}/"))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "new_contract": contract_id + }))) + .mount(&server) + .await; + } + Mock::given(method("GET")) + .and(path("/api/v0/instances/101/")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "instances": { + "actual_status": "loading", + "intended_status": "running", + "status_msg": "still pulling" + } + }))) + .mount(&server) + .await; + for contract_id in [102, 104] { + Mock::given(method("GET")) + .and(path(format!("/api/v0/instances/{contract_id}/"))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "instances": { + "actual_status": "running", + "intended_status": "running", + "public_ipaddr": "127.0.0.1", + "ssh_port": 22 + } + }))) + .mount(&server) + .await; + } + Mock::given(method("DELETE")) + .and(path("/api/v0/instances/101/")) + .respond_with(ResponseTemplate::new(200)) + .mount(&server) + .await; + + let fleet = provision_fleet(&reqwest::Client::new(), &server.uri(), "secret", request(2)) + .await + .expect("replacement should use non-failed host"); + + assert_eq!( + fleet + .instances + .iter() + .map(|instance| (instance.index, instance.offer_id, instance.host_id)) + .collect::>(), + vec![(0, 4, Some(30)), (1, 2, Some(20))] + ); + let requests = server.received_requests().await.expect("recorded requests"); + assert!( + requests + .iter() + .any(|request| request.url.path() == "/api/v0/asks/4/"), + "replacement should rent an offer from a non-failed host" + ); + assert!( + !requests + .iter() + .any(|request| request.url.path() == "/api/v0/asks/3/"), + "replacement must skip untried offers on the failed host" + ); + } +} diff --git a/tools/vastai/src/lib.rs b/tools/vastai/src/lib.rs index 7c7ec4b..d38ecb5 100644 --- a/tools/vastai/src/lib.rs +++ b/tools/vastai/src/lib.rs @@ -23,12 +23,12 @@ pub use logs::{fetch_logs, request_logs}; pub use monitor::{wait_for_running, wait_for_running_with_policy}; pub use pricing::CostModel; pub use provision::create_instance; -pub use search::{select_offer_pool, select_offer_pool_with_policy}; +pub use search::{plan_distinct_host_first_wave, select_offer_pool, select_offer_pool_with_policy}; pub use teardown::{ destroy_all_instances, destroy_instance, destroy_instance_with_retry, list_instances_by_label, }; pub use types::{ ContractRef, CreateInstanceRequest, FleetState, InstanceInfo, LabeledInstance, LifecyclePolicy, Offer, ProvisionRequest, ProvisionedFleet, ProvisionedInstance, RunningInstance, - SelectionPolicy, + SelectionPolicy, VastAiFailureClass, classify_vastai_error, }; diff --git a/tools/vastai/src/pricing.rs b/tools/vastai/src/pricing.rs index ebbabd6..405f65b 100644 --- a/tools/vastai/src/pricing.rs +++ b/tools/vastai/src/pricing.rs @@ -48,20 +48,3 @@ pub(crate) fn rank_survivors(offers: Vec, cost: &CostModel, drop_frac: f6 survivors.sort_by(&by_price); survivors } - -pub(crate) fn plan_picks(pool: &[Offer], num_instances: u32) -> Vec<&Offer> { - let mut picks = Vec::with_capacity(num_instances as usize); - let mut used = std::collections::HashSet::new(); - for o in pool { - if picks.len() == num_instances as usize { - break; - } - if let Some(h) = o.host_id { - if !used.insert(h) { - continue; - } - } - picks.push(o); - } - picks -} diff --git a/tools/vastai/src/search.rs b/tools/vastai/src/search.rs index a62f502..55ea64e 100644 --- a/tools/vastai/src/search.rs +++ b/tools/vastai/src/search.rs @@ -1,6 +1,60 @@ use crate::filters::reachable_offers; use crate::pricing::{CostModel, rank_survivors}; use crate::types::{Offer, SearchResponse, SelectionPolicy}; +use std::collections::HashSet; + +/// Choose the first batch of offers from one ranked pool, preferring different +/// hosts whenever the filtered pool can satisfy that. +pub fn plan_distinct_host_first_wave( + pool: &[Offer], + target_count: u32, + blacklisted_hosts: &[u64], + failed_hosts: &[u64], +) -> Vec { + let target = target_count as usize; + let blocked = blacklisted_hosts + .iter() + .chain(failed_hosts.iter()) + .copied() + .collect::>(); + let mut selected = Vec::with_capacity(target); + let mut selected_ids = HashSet::new(); + let mut selected_hosts = HashSet::new(); + + for offer in pool.iter().filter(|offer| { + offer + .host_id + .is_none_or(|host_id| !blocked.contains(&host_id)) + }) { + if selected.len() == target { + break; + } + if let Some(host_id) = offer.host_id { + if !selected_hosts.insert(host_id) { + continue; + } + } + selected_ids.insert(offer.id); + selected.push(offer.clone()); + } + + if selected.len() < target { + for offer in pool.iter().filter(|offer| { + offer + .host_id + .is_none_or(|host_id| !blocked.contains(&host_id)) + }) { + if selected.len() == target { + break; + } + if selected_ids.insert(offer.id) { + selected.push(offer.clone()); + } + } + } + + selected +} /// Historical env-backed offer search wrapper. pub async fn select_offer_pool( @@ -103,3 +157,70 @@ pub async fn select_offer_pool_with_policy( ); Ok(pool) } + +#[cfg(test)] +mod tests { + use super::*; + + fn offer(id: u64, host_id: Option) -> Offer { + Offer { + id, + gpu_name: "RTX 4090".to_owned(), + dph_total: id as f64 / 100.0, + gpu_ram: Some(24_000.0), + compute_cap: 890, + geolocation: Some("US".to_owned()), + inet_down_cost_per_tb: 0.0, + inet_up_cost_per_tb: 0.0, + host_id, + verification: Some("verified".to_owned()), + } + } + + #[test] + fn first_wave_plan_prefers_distinct_hosts_and_preserves_blacklists() { + let pool = vec![ + offer(1, Some(10)), + offer(2, Some(10)), + offer(3, Some(20)), + offer(4, Some(30)), + offer(5, Some(40)), + ]; + + let plan = plan_distinct_host_first_wave(&pool, 3, &[30], &[]); + + assert_eq!( + plan.iter().map(|offer| offer.id).collect::>(), + vec![1, 3, 5] + ); + assert_eq!( + plan.iter() + .filter_map(|offer| offer.host_id) + .collect::>() + .len(), + 3 + ); + assert!( + plan.iter().all(|offer| offer.host_id != Some(30)), + "operator blacklist must remain authoritative" + ); + } + + #[test] + fn first_wave_plan_excludes_failed_hosts_from_replacements() { + let pool = vec![ + offer(1, Some(10)), + offer(2, Some(20)), + offer(3, Some(30)), + offer(4, Some(20)), + ]; + + let plan = plan_distinct_host_first_wave(&pool, 2, &[], &[20]); + + assert_eq!( + plan.iter().map(|offer| offer.id).collect::>(), + vec![1, 3] + ); + assert!(plan.iter().all(|offer| offer.host_id != Some(20))); + } +} diff --git a/tools/vastai/src/types.rs b/tools/vastai/src/types.rs index 41190ba..4ebd7a1 100644 --- a/tools/vastai/src/types.rs +++ b/tools/vastai/src/types.rs @@ -133,6 +133,9 @@ pub struct ProvisionRequest { pub env: BTreeMap, /// Per-index env overlays, merged after `env`. pub per_instance_env: Vec>, + /// Preferred offer for single-instance requests after an app-level shared + /// first-wave planner has already coordinated distinct hosts. + pub preferred_offer_id: Option, pub onstart: Option, pub selection: SelectionPolicy, pub lifecycle: LifecyclePolicy, @@ -158,6 +161,99 @@ pub struct ProvisionedFleet { pub instances: Vec, } +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum VastAiFailureClass { + ConnectionRefused, + PublicKeyDenied, + EndpointMissing, + ProviderLoadingTimeout, + VanishedOffer, + Other, +} + +impl VastAiFailureClass { + pub fn as_str(self) -> &'static str { + match self { + Self::ConnectionRefused => "connection_refused", + Self::PublicKeyDenied => "publickey_denied", + Self::EndpointMissing => "endpoint_missing", + Self::ProviderLoadingTimeout => "provider_loading_timeout", + Self::VanishedOffer => "vanished_offer", + Self::Other => "other", + } + } +} + +pub fn classify_vastai_error(raw: &str) -> VastAiFailureClass { + let lower = raw.to_ascii_lowercase(); + if lower.contains("connection refused") || lower.contains("os error 111") { + return VastAiFailureClass::ConnectionRefused; + } + if lower.contains("permission denied (publickey") + || lower.contains("publickey denied") + || lower.contains("public key denied") + || lower.contains("no supported authentication methods") + { + return VastAiFailureClass::PublicKeyDenied; + } + if lower.contains("no ssh host") + || lower.contains("no ssh port") + || lower.contains("no ssh endpoint") + || lower.contains("endpoint missing") + || lower.contains("ssh missing") + { + return VastAiFailureClass::EndpointMissing; + } + if lower.contains("stuck in status loading") + || lower.contains("provider loading timeout") + || (lower.contains("status loading") && lower.contains("timeout")) + { + return VastAiFailureClass::ProviderLoadingTimeout; + } + if lower.contains("no_such_ask") + || lower.contains("no such ask") + || lower.contains("vanished offer") + || (lower.contains("404") && lower.contains("/asks/")) + { + return VastAiFailureClass::VanishedOffer; + } + VastAiFailureClass::Other +} + +#[cfg(test)] +mod failure_class_tests { + use super::*; + + #[test] + fn representative_vastai_errors_classify_to_stable_failure_classes() { + for (raw, class) in [ + ( + "ssh: connect to host ssh5.vast.ai port 22017: Connection refused", + VastAiFailureClass::ConnectionRefused, + ), + ( + "Permission denied (publickey).", + VastAiFailureClass::PublicKeyDenied, + ), + ( + "vastai contract 123 has no SSH host", + VastAiFailureClass::EndpointMissing, + ), + ( + "instance 123 stuck in status loading for 300s", + VastAiFailureClass::ProviderLoadingTimeout, + ), + ( + "create_instance HTTP 400: {\"error\":\"no_such_ask\"}", + VastAiFailureClass::VanishedOffer, + ), + ] { + assert_eq!(classify_vastai_error(raw), class, "{raw}"); + assert_ne!(class.as_str(), "other"); + } + } +} + /// Generic held-fleet handle file. #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] pub struct FleetState { diff --git a/xtask/src/main.rs b/xtask/src/main.rs index c4d29c3..dfbf0f1 100644 --- a/xtask/src/main.rs +++ b/xtask/src/main.rs @@ -755,6 +755,8 @@ fn run_mvp_chat_check(args: Vec) -> ExitCode { }; let run_id = mvp_chat_check_run_id(); println!("mvp-chat-check: scenario {}", invocation.name()); + println!("mvp-chat-check: artifacts {}", paths.root.display()); + println!("mvp-chat-check: datastream {}", paths.dump_log.display()); let output = match run_mvp_chat_check_process(&workspace, &paths, run_id, &invocation) { Ok(output) => output,