From 0edc6fc229d67370f4aed1865de8fefc3b9353aa Mon Sep 17 00:00:00 2001 From: Zachery Aaron Shores-Chmielewski Date: Wed, 29 Jul 2026 17:17:57 +0400 Subject: [PATCH] refactor(mvp-system): consolidate runtime-ready ack flow - De-generify run_chat_session_* helpers. - Replace ProvisionedNodeGuard, StageProvisionDispatch, and PipelineSendHandle with a single RuntimeReadyAckLoop driving wait_for_runtime_ready / wait_for_weights_loaded. - Add staging shard offset/header helpers. Signed-off-by: Zachery Aaron Shores-Chmielewski --- crates/mvp-system/src/chat/node_image.rs | 59 +- crates/mvp-system/src/chat/runtime.rs | 680 ++++---- crates/mvp-system/src/node/actor.rs | 159 +- .../src/node/worker_node_runtime.rs | 609 +++---- crates/mvp-system/src/orchestration/app.rs | 1469 ++++++++--------- .../src/orchestration/engine_builder/mod.rs | 5 +- .../mvp-system/src/orchestration/run_fsm.rs | 77 +- crates/mvp-system/src/staging/control.rs | 49 +- .../mvp-system/src/staging/gguf_metadata.rs | 35 +- crates/mvp-system/src/staging/gguf_shard.rs | 139 +- 10 files changed, 1566 insertions(+), 1715 deletions(-) diff --git a/crates/mvp-system/src/chat/node_image.rs b/crates/mvp-system/src/chat/node_image.rs index 9b7e1d2..ee9e90c 100644 --- a/crates/mvp-system/src/chat/node_image.rs +++ b/crates/mvp-system/src/chat/node_image.rs @@ -49,16 +49,6 @@ pub(super) struct NodeImageRequest { pub(super) enabled: bool, } -#[allow(dead_code)] -#[derive(Clone, Debug)] -pub(super) struct PreparedNodeImage { - pub(super) image_ref: String, - pub(super) tag: String, - pub(super) already_available: bool, - pub(super) built: bool, - pub(super) pushed: bool, -} - #[derive(Clone, Debug, PartialEq, Eq)] pub(super) enum NodeImageProgressEventKind { ImageReference { @@ -140,7 +130,7 @@ struct RealImageCommandRunner; pub(super) fn prepare_node_image_with_progress( request: NodeImageRequest, progress: Option<&mut dyn NodeImageProgressSink>, -) -> Result { +) -> Result { let mut progress = progress; let mut runner = RealImageCommandRunner; prepare_node_image_inner(request, &mut progress, &mut runner) @@ -150,16 +140,10 @@ fn prepare_node_image_inner( request: NodeImageRequest, progress: &mut Option<&mut dyn NodeImageProgressSink>, runner: &mut dyn ImageCommandRunner, -) -> Result { +) -> Result { emit_image_reference(progress, "requested", &request.requested_image); if !request.enabled { - return Ok(PreparedNodeImage { - image_ref: request.requested_image, - tag: String::new(), - already_available: false, - built: false, - pushed: false, - }); + return Ok(request.requested_image); } emit_image_reference(progress, "base", &request.base_image); @@ -208,16 +192,9 @@ fn prepare_node_image_inner( docker_image_labels_match(runner, &root, &image_ref, &expected_node_labels)?; let remote_available = remote_required && runner.docker_manifest_exists(&root, &image_ref); if !request.force_refresh && remote_required && remote_available { - let pushed = - ensure_aliases_for_remote(runner, progress, &root, &image_ref, &image, &alias_tags)?; + ensure_aliases_for_remote(runner, progress, &root, &image_ref, &image, &alias_tags)?; prune_old_dirty_images(runner, &root, &image, &tag); - return Ok(PreparedNodeImage { - image_ref, - tag, - already_available: true, - built: false, - pushed, - }); + return Ok(image_ref); } if !request.force_refresh && remote_required && local_image_matches { ensure_aliases_local(runner, progress, &root, &image_ref, &image, &alias_tags)?; @@ -226,24 +203,12 @@ fn prepare_node_image_inner( push_image(runner, progress, &root, &alias)?; } prune_old_dirty_images(runner, &root, &image, &tag); - return Ok(PreparedNodeImage { - image_ref, - tag, - already_available: true, - built: false, - pushed: true, - }); + return Ok(image_ref); } if !request.force_refresh && !remote_required && local_image_matches { ensure_aliases_local(runner, progress, &root, &image_ref, &image, &alias_tags)?; prune_old_dirty_images(runner, &root, &image, &tag); - return Ok(PreparedNodeImage { - image_ref, - tag, - already_available: true, - built: false, - pushed: false, - }); + return Ok(image_ref); } let base_image_matches = docker_image_labels_match(runner, &root, &request.base_image, &expected_base_labels)?; @@ -294,23 +259,15 @@ fn prepare_node_image_inner( )?; ensure_aliases_local(runner, progress, &root, &image_ref, &image, &alias_tags)?; - let mut pushed = false; if remote_required { push_image(runner, progress, &root, &image_ref)?; - pushed = true; for alias in alias_refs(&image, &alias_tags) { push_image(runner, progress, &root, &alias)?; } } prune_old_dirty_images(runner, &root, &image, &tag); - Ok(PreparedNodeImage { - image_ref, - tag, - already_available: false, - built: true, - pushed, - }) + Ok(image_ref) } fn workspace_root() -> Result { diff --git a/crates/mvp-system/src/chat/runtime.rs b/crates/mvp-system/src/chat/runtime.rs index f77463f..111eb60 100644 --- a/crates/mvp-system/src/chat/runtime.rs +++ b/crates/mvp-system/src/chat/runtime.rs @@ -25,7 +25,7 @@ use signal_hook::iterator::Signals; use crate::chat::config as chat_config; use crate::chat::node_image::{ NodeImageProgressEvent, NodeImageProgressEventKind, NodeImageProgressSink, NodeImageProvider, - NodeImageRequest, PreparedNodeImage, prepare_node_image_with_progress, + NodeImageRequest, prepare_node_image_with_progress, }; use crate::node_provisioning::{ProviderKind, provider_kind}; use crate::observability::{benchmark, frame_archive::FrameArchive}; @@ -775,73 +775,21 @@ impl Config { let args = ParsedArgs::parse(provided_args)?; let loaded = load_chat_config(args.config_path.as_deref())?; let toml = loaded.overlay; - let provider = provider_from_sources(args.provider, toml.provider.kind.as_deref())?; + let provider = provider_from_sources(args.provider.clone(), toml.provider.kind.as_deref())?; let node_image = first_non_empty([toml.image.node.clone()]).unwrap_or_default(); if provider != provider_kind::process() && node_image.is_empty() { return Err("node image is required for docker or vastai provider".to_owned()); } - let pipeline_stages = args - .pipeline_stages - .or(toml.runtime.pipeline_stages) - .unwrap_or(1); - if pipeline_stages == 0 { - return Err("--pipeline-stages must be greater than 0".to_owned()); - } - let max_tokens = toml.runtime.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS); - if max_tokens == 0 { - return Err("[runtime].max_tokens must be greater than 0".to_owned()); - } + let pipeline_stages = Self::pipeline_stages(&args, &toml)?; + let max_tokens = Self::max_tokens(&toml)?; let gpu_run = args.gpu || env_flag(MVP_CHAT_GPU_RUN_ENV, false); - let endpoint_addr_mask = match first_non_empty([ - args.endpoint_addr_mask.clone(), - toml.relay.endpoint_addr_mask.clone(), - ]) { - Some(mask) => EndpointAddrMask::parse(&mask)?, - None => EndpointAddrMask::Full, - }; - let relay_mode = first_non_empty([args.relay_mode.clone(), toml.relay.mode.clone()]); - let mut relay_url = first_non_empty([args.relay_url.clone(), toml.relay.url.clone()]); - if endpoint_addr_mask.requires_relay() && relay_url.is_none() { - relay_url = first_non_empty([toml.vastai.relay_url.clone()]); - } - if endpoint_addr_mask.requires_relay() && relay_url.is_none() { - return Err("relay-only endpoint address mask requires [relay].url, --relay-url, or [vastai].relay_url".to_owned()); - } - let relay_mode = relay_mode.or_else(|| relay_url.as_ref().map(|_| "default".to_owned())); - let cached_model_source = match args.cached_model { - Some(source) => Some(source), - None if gpu_run && provider == provider_kind::process() => { - Some(CachedModelSource::Discover) - } - None => None, - }; - let cached_model = cached_model_source + let endpoint_addr_mask = Self::endpoint_addr_mask(&args, &toml)?; + let (relay_mode, relay_url) = Self::relay_settings(&args, &toml, endpoint_addr_mask)?; + let cached_model = Self::cached_model_source(&args, gpu_run, &provider) .map(CachedModelConfig::from_source) .transpose()?; - let model = if provider == provider_kind::vastai() { - match &cached_model { - Some(cached_model) => { - vastai_model_config_for_cached_model(toml.model.clone(), cached_model)? - } - None => toml.model.clone(), - } - } else { - toml.model.clone() - }; - let datastream_frame_log = if args.dump_logs { - Some( - args.dump_log_path - .unwrap_or_else(|| PathBuf::from("mvp-chat.log")), - ) - } else if toml.observability.dump_logs.unwrap_or(false) { - Some( - first_non_empty([toml.observability.dump_log_path.clone()]) - .map(PathBuf::from) - .unwrap_or_else(|| PathBuf::from("mvp-chat.log")), - ) - } else { - None - }; + let model = Self::model_config(&provider, &toml, cached_model.as_ref())?; + let datastream_frame_log = Self::datastream_frame_log(&args, &toml); let vastai = if provider == provider_kind::vastai() { Some(resolve_vastai_config(&toml.vastai, &node_image)?) } else { @@ -871,9 +819,114 @@ impl Config { }) } + fn pipeline_stages(args: &ParsedArgs, toml: &ChatTomlConfig) -> Result { + let pipeline_stages = args + .pipeline_stages + .or(toml.runtime.pipeline_stages) + .unwrap_or(1); + if pipeline_stages == 0 { + return Err("--pipeline-stages must be greater than 0".to_owned()); + } + Ok(pipeline_stages) + } + + fn max_tokens(toml: &ChatTomlConfig) -> Result { + let max_tokens = toml.runtime.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS); + if max_tokens == 0 { + return Err("[runtime].max_tokens must be greater than 0".to_owned()); + } + Ok(max_tokens) + } + + fn endpoint_addr_mask( + args: &ParsedArgs, + toml: &ChatTomlConfig, + ) -> Result { + match first_non_empty([ + args.endpoint_addr_mask.clone(), + toml.relay.endpoint_addr_mask.clone(), + ]) { + Some(mask) => EndpointAddrMask::parse(&mask), + None => Ok(EndpointAddrMask::Full), + } + } + + fn relay_settings( + args: &ParsedArgs, + toml: &ChatTomlConfig, + endpoint_addr_mask: EndpointAddrMask, + ) -> Result<(Option, Option), String> { + let relay_mode = first_non_empty([args.relay_mode.clone(), toml.relay.mode.clone()]); + let mut relay_url = first_non_empty([args.relay_url.clone(), toml.relay.url.clone()]); + if endpoint_addr_mask.requires_relay() && relay_url.is_none() { + relay_url = first_non_empty([toml.vastai.relay_url.clone()]); + } + if endpoint_addr_mask.requires_relay() && relay_url.is_none() { + return Err("relay-only endpoint address mask requires [relay].url, --relay-url, or [vastai].relay_url".to_owned()); + } + let relay_mode = relay_mode.or_else(|| relay_url.as_ref().map(|_| "default".to_owned())); + Ok((relay_mode, relay_url)) + } + + fn cached_model_source( + args: &ParsedArgs, + gpu_run: bool, + provider: &ProviderKind, + ) -> Option { + match &args.cached_model { + Some(source) => Some(source.clone()), + None if gpu_run && provider == &provider_kind::process() => { + Some(CachedModelSource::Discover) + } + None => None, + } + } + + fn model_config( + provider: &ProviderKind, + toml: &ChatTomlConfig, + cached_model: Option<&CachedModelConfig>, + ) -> Result { + if provider != &provider_kind::vastai() { + return Ok(toml.model.clone()); + } + match cached_model { + Some(cached_model) => { + vastai_model_config_for_cached_model(toml.model.clone(), cached_model) + } + None => Ok(toml.model.clone()), + } + } + + fn datastream_frame_log(args: &ParsedArgs, toml: &ChatTomlConfig) -> Option { + if args.dump_logs { + return Some( + args.dump_log_path + .clone() + .unwrap_or_else(|| PathBuf::from("mvp-chat.log")), + ); + } + if !toml.observability.dump_logs.unwrap_or(false) { + return None; + } + Some( + first_non_empty([toml.observability.dump_log_path.clone()]) + .map(PathBuf::from) + .unwrap_or_else(|| PathBuf::from("mvp-chat.log")), + ) + } + // The orchestrator launch spec is still pending. These flags are the current adapter; // adjust this mapping when the approved orchestrator launch contract is finalized. fn orchestrator_cli_args(&self, image_ref: &str) -> Vec { + macro_rules! push_opt { + ($args:ident, $option:expr, $flag:expr, |$value:ident| $arg:expr) => { + if let Some($value) = $option { + $args.extend([$flag.to_owned(), $arg]); + } + }; + } + let mut args = vec![ "--provider".to_owned(), self.provider.as_str().to_owned(), @@ -889,27 +942,36 @@ impl Config { self.pipeline_stages.to_string(), "--dashboard".to_owned(), ]; - if let Some(model_id) = &self.model.id { - args.extend(["--model-id".to_owned(), model_id.clone()]); - } - if let Some(path) = &self.model.gguf_local_path { - args.extend(["--gguf-local-path".to_owned(), path.clone()]); - } - if let Some(repo) = &self.model.gguf_repo { - args.extend(["--gguf-repo".to_owned(), repo.clone()]); - } - if let Some(file) = &self.model.gguf_file { - args.extend(["--gguf-file".to_owned(), file.clone()]); - } - if let Some(revision) = &self.model.gguf_revision { - args.extend(["--gguf-revision".to_owned(), revision.clone()]); - } - if let Some(path) = &self.model.tokenizer_local_path { - args.extend(["--tokenizer-local-path".to_owned(), path.clone()]); - } - if let Some(max_context) = self.model.max_context { - args.extend(["--max-context".to_owned(), max_context.to_string()]); - } + push_opt!(args, &self.model.id, "--model-id", |model_id| model_id + .clone()); + push_opt!( + args, + &self.model.gguf_local_path, + "--gguf-local-path", + |path| path.clone() + ); + push_opt!(args, &self.model.gguf_repo, "--gguf-repo", |repo| repo + .clone()); + push_opt!(args, &self.model.gguf_file, "--gguf-file", |file| file + .clone()); + push_opt!( + args, + &self.model.gguf_revision, + "--gguf-revision", + |revision| { revision.clone() } + ); + push_opt!( + args, + &self.model.tokenizer_local_path, + "--tokenizer-local-path", + |path| path.clone() + ); + push_opt!( + args, + self.model.max_context, + "--max-context", + |max_context| { max_context.to_string() } + ); if self.provider == provider_kind::process() { args.extend([ "--worker-bin".to_owned(), @@ -922,18 +984,14 @@ impl Config { cached_model.host_path.to_string_lossy().to_string(), ]); } - if let Some(path) = &self.datastream_frame_log { - args.extend([ - "--datastream-frame-log".to_owned(), - path.to_string_lossy().to_string(), - ]); - } - if let Some(mode) = &self.relay_mode { - args.extend(["--relay-mode".to_owned(), mode.clone()]); - } - if let Some(url) = &self.relay_url { - args.extend(["--relay-url".to_owned(), url.clone()]); - } + push_opt!( + args, + &self.datastream_frame_log, + "--datastream-frame-log", + |path| { path.to_string_lossy().to_string() } + ); + push_opt!(args, &self.relay_mode, "--relay-mode", |mode| mode.clone()); + push_opt!(args, &self.relay_url, "--relay-url", |url| url.clone()); if self.endpoint_addr_mask != EndpointAddrMask::Full { args.extend([ "--endpoint-addr-mask".to_owned(), @@ -946,39 +1004,41 @@ impl Config { vastai.bootstrap_command.clone(), "--no-vastai-confirm-lease".to_owned(), ]); - if let Some(disk_gb) = vastai.disk_gb { - args.extend(["--vastai-disk-gb".to_owned(), disk_gb.to_string()]); - } - if let Some(gpu_name) = &vastai.gpu_name { - args.extend(["--vastai-gpu-name".to_owned(), gpu_name.clone()]); - } - if let Some(min_gpu_ram_mb) = vastai.min_gpu_ram_mb { - args.extend([ - "--vastai-min-gpu-ram-mb".to_owned(), - min_gpu_ram_mb.to_string(), - ]); - } - if let Some(min_down_mbps) = vastai.min_down_mbps { - args.extend([ - "--vastai-min-down-mbps".to_owned(), - min_down_mbps.to_string(), - ]); - } - if let Some(min_up_mbps) = vastai.min_up_mbps { - args.extend(["--vastai-min-up-mbps".to_owned(), min_up_mbps.to_string()]); - } - if let Some(max_dph_total) = vastai.max_dph_total { - args.extend([ - "--vastai-max-dph-total".to_owned(), - max_dph_total.to_string(), - ]); - } - if let Some(min_reliability) = vastai.min_reliability { - args.extend([ - "--vastai-min-reliability".to_owned(), - min_reliability.to_string(), - ]); - } + push_opt!(args, vastai.disk_gb, "--vastai-disk-gb", |disk_gb| disk_gb + .to_string()); + push_opt!(args, &vastai.gpu_name, "--vastai-gpu-name", |gpu_name| { + gpu_name.clone() + }); + push_opt!( + args, + vastai.min_gpu_ram_mb, + "--vastai-min-gpu-ram-mb", + |min_gpu_ram_mb| min_gpu_ram_mb.to_string() + ); + push_opt!( + args, + vastai.min_down_mbps, + "--vastai-min-down-mbps", + |min_down_mbps| min_down_mbps.to_string() + ); + push_opt!( + args, + vastai.min_up_mbps, + "--vastai-min-up-mbps", + |min_up_mbps| { min_up_mbps.to_string() } + ); + push_opt!( + args, + vastai.max_dph_total, + "--vastai-max-dph-total", + |max_dph_total| { max_dph_total.to_string() } + ); + push_opt!( + args, + vastai.min_reliability, + "--vastai-min-reliability", + |min_reliability| min_reliability.to_string() + ); if let Some(require_verified) = vastai.require_verified { args.push(if require_verified { "--vastai-require-verified".to_owned() @@ -989,12 +1049,14 @@ impl Config { for host_id in &vastai.blacklist_hosts { args.extend(["--vastai-blacklist-host".to_owned(), host_id.to_string()]); } - if let Some(onstart) = &vastai.onstart { - args.extend(["--vastai-onstart".to_owned(), onstart.clone()]); - } - if let Some(ssh_identity) = &vastai.ssh_identity { - args.extend(["--vastai-ssh-identity".to_owned(), ssh_identity.clone()]); - } + push_opt!(args, &vastai.onstart, "--vastai-onstart", |onstart| onstart + .clone()); + push_opt!( + args, + &vastai.ssh_identity, + "--vastai-ssh-identity", + |ssh_identity| { ssh_identity.clone() } + ); } args } @@ -1044,6 +1106,70 @@ impl ParsedArgs { Ok(()) } + fn apply_provider_arg(&mut self, arg: &str) -> Result { + match arg { + "--help" | "-h" | "help" => self.help = true, + "--gpu" => self.gpu = true, + "--vastai" => self.set_provider_selector(provider_kind::vastai())?, + "--process" => self.set_provider_selector(provider_kind::process())?, + "--docker" => self.set_provider_selector(provider_kind::docker())?, + "--yes" | "-y" => self.vastai_yes = true, + "--dump-logs" => self.dump_logs = true, + "--cached-model" => self.cached_model = Some(CachedModelSource::Discover), + "--skip-rebuild" => self.skip_rebuild = true, + _ => return Ok(false), + } + Ok(true) + } + + fn apply_config_arg(&mut self, arg: &str, args: &mut I) -> Result + where + I: Iterator, + { + match arg { + "--config" => self.config_path = Some(PathBuf::from(next_arg(args, "--config")?)), + "--pipeline-stages" | "--pipeline-parallel" => { + if self.pipeline_stages.is_some() { + return Err("pipeline stage count was provided more than once".to_owned()); + } + self.pipeline_stages = Some(parse_pipeline_stages_value(args, arg)?); + } + "--relay-mode" => self.relay_mode = Some(next_arg(args, "--relay-mode")?), + "--relay-url" => self.relay_url = Some(next_arg(args, "--relay-url")?), + "--endpoint-addr-mask" => { + self.endpoint_addr_mask = Some(next_arg(args, "--endpoint-addr-mask")?) + } + "--run-id" => { + let run_id: u64 = parse_next(args, "--run-id")?; + if run_id == 0 { + return Err("--run-id must be greater than 0".to_owned()); + } + self.run_id = Some(run_id); + } + _ => return Ok(false), + } + Ok(true) + } + + fn apply_assignment_arg(&mut self, arg: &str) -> Result { + if let Some(path) = arg.strip_prefix("--dump-logs=") { + if path.is_empty() { + return Err("--dump-logs path must not be empty".to_owned()); + } + self.dump_logs = true; + self.dump_log_path = Some(PathBuf::from(path)); + return Ok(true); + } + if let Some(path) = arg.strip_prefix("--cached-model=") { + if path.is_empty() { + return Err("--cached-model path must not be empty".to_owned()); + } + self.cached_model = Some(CachedModelSource::Path(PathBuf::from(path))); + return Ok(true); + } + Ok(false) + } + fn parse(provided_args: I) -> Result where I: IntoIterator, @@ -1051,61 +1177,13 @@ impl ParsedArgs { let mut parsed = Self::default(); let mut args = provided_args.into_iter().peekable(); while let Some(arg) = args.next() { - match arg.as_str() { - "--help" | "-h" | "help" => parsed.help = true, - "--gpu" => parsed.gpu = true, - "--vastai" => parsed.set_provider_selector(provider_kind::vastai())?, - "--process" => parsed.set_provider_selector(provider_kind::process())?, - "--docker" => parsed.set_provider_selector(provider_kind::docker())?, - "--yes" | "-y" => parsed.vastai_yes = true, - "--config" => { - parsed.config_path = Some(PathBuf::from(next_arg(&mut args, "--config")?)) - } - "--pipeline-stages" | "--pipeline-parallel" => { - if parsed.pipeline_stages.is_some() { - return Err("pipeline stage count was provided more than once".to_owned()); - } - parsed.pipeline_stages = - Some(parse_pipeline_stages_value(&mut args, arg.as_str())?) - } - "--relay-mode" => parsed.relay_mode = Some(next_arg(&mut args, "--relay-mode")?), - "--relay-url" => parsed.relay_url = Some(next_arg(&mut args, "--relay-url")?), - "--endpoint-addr-mask" => { - parsed.endpoint_addr_mask = Some(next_arg(&mut args, "--endpoint-addr-mask")?) - } - "--run-id" => { - let run_id: u64 = parse_next(&mut args, "--run-id")?; - if run_id == 0 { - return Err("--run-id must be greater than 0".to_owned()); - } - parsed.run_id = Some(run_id); - } - "--dump-logs" => { - parsed.dump_logs = true; - } - value if value.starts_with("--dump-logs=") => { - let path = value.strip_prefix("--dump-logs=").expect("prefix checked"); - if path.is_empty() { - return Err("--dump-logs path must not be empty".to_owned()); - } - parsed.dump_logs = true; - parsed.dump_log_path = Some(PathBuf::from(path)); - } - "--cached-model" => { - parsed.cached_model = Some(CachedModelSource::Discover); - } - value if value.starts_with("--cached-model=") => { - let path = value - .strip_prefix("--cached-model=") - .expect("prefix checked"); - if path.is_empty() { - return Err("--cached-model path must not be empty".to_owned()); - } - parsed.cached_model = Some(CachedModelSource::Path(PathBuf::from(path))); - } - "--skip-rebuild" => parsed.skip_rebuild = true, - other => return Err(format!("unsupported mvp-chat argument {other:?}")), + if parsed.apply_provider_arg(&arg)? + || parsed.apply_config_arg(&arg, &mut args)? + || parsed.apply_assignment_arg(&arg)? + { + continue; } + return Err(format!("unsupported mvp-chat argument {arg:?}")); } Ok(parsed) } @@ -1511,7 +1589,7 @@ fn signal_orch_process_group(child: &Child, signal: libc::c_int) -> io::Result<( fn prepare_node_image_progress_adapter( request: NodeImageRequest, progress: Option<&mut dyn NodeImageProgressSink>, -) -> Result { +) -> Result { prepare_node_image_with_progress(request, progress) } @@ -1525,7 +1603,7 @@ fn prepare_runtime(config: &Config) -> Result { #[allow(dead_code)] fn prepare_runtime_with(config: &Config, prepare_node_image_fn: F) -> Result where - F: FnMut(NodeImageRequest) -> Result, + F: FnMut(NodeImageRequest) -> Result, { let mut prepare_node_image_fn = prepare_node_image_fn; prepare_runtime_with_progress( @@ -1541,10 +1619,7 @@ fn prepare_runtime_with_progress( progress: Option<&mut ChatDatastream>, ) -> Result where - F: FnMut( - NodeImageRequest, - Option<&mut dyn NodeImageProgressSink>, - ) -> Result, + F: FnMut(NodeImageRequest, Option<&mut dyn NodeImageProgressSink>) -> Result, { let mut progress = progress; let binary_mode = if config.skip_rebuild { @@ -1758,9 +1833,9 @@ where CHAT_RUNTIME_CHANNEL, "prepare_node_image", "ready", - 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()}), + json!({"provider": config.provider.as_str(), "command_label": "prepare_node_image", "image_ref": &prepared, "elapsed_ms": prepare_node_image_started.elapsed().as_millis()}), ); - Ok(prepared.image_ref) + Ok(prepared) } fn stdin_prompt_events() -> mpsc::Receiver { @@ -1852,32 +1927,23 @@ fn run_chat_loop_with_input_and_progress( } #[cfg(test)] -fn run_chat_session_with_output( - writer: &mut W, - reader: R, +fn run_chat_session_with_output( + writer: &mut impl Write, + reader: impl BufRead, input_rx: mpsc::Receiver, max_tokens: u32, - output: &mut O, -) -> Result<(), String> -where - R: BufRead, - W: Write, - O: Write, -{ + output: &mut impl Write, +) -> Result<(), String> { run_chat_session_with_output_and_progress(writer, reader, input_rx, max_tokens, output, None) } -fn run_chat_session_with_progress( - writer: &mut W, - reader: R, +fn run_chat_session_with_progress( + writer: &mut impl Write, + reader: impl BufRead, input_rx: mpsc::Receiver, max_tokens: u32, progress: Option<&mut ChatDatastream>, -) -> Result<(), String> -where - R: BufRead, - W: Write, -{ +) -> Result<(), String> { let mut output = io::stdout(); run_chat_session_with_output_and_progress( writer, @@ -1905,19 +1971,14 @@ fn prompt_hash_hex(prompt: &str) -> String { blake3::hash(prompt.as_bytes()).to_hex().to_string() } -fn run_chat_session_with_output_and_progress( - writer: &mut W, - mut reader: R, +fn run_chat_session_with_output_and_progress( + writer: &mut impl Write, + mut reader: impl BufRead, input_rx: mpsc::Receiver, max_tokens: u32, - output: &mut O, + output: &mut impl Write, progress: Option<&mut ChatDatastream>, -) -> Result<(), String> -where - R: BufRead, - W: Write, - O: Write, -{ +) -> Result<(), String> { let mut progress = progress; let mut next_request_id = 1_u64; let mut next_prompt_index = 1_u64; @@ -2375,10 +2436,8 @@ fn parse_pipeline_stages_value( mod tests { use super::*; - use std::ffi::{OsStr, OsString}; + use std::ffi::OsString; use std::io::{Cursor, Read}; - #[cfg(target_os = "linux")] - use std::os::unix::process::CommandExt; use std::path::Path; use std::sync::Mutex; use std::sync::atomic::{AtomicU64, Ordering as AtomicOrdering}; @@ -2487,50 +2546,6 @@ mod tests { path } - fn base_config(provider: ProviderKind) -> Config { - Config { - orch_bin: PathBuf::from("/tmp/mvp-orchestrator"), - worker_bin: PathBuf::from("/tmp/mvp-worker-node"), - rpc_addr: DEFAULT_RPC_ADDR.to_owned(), - node_image: "docker.io/acme/node:latest".to_owned(), - provider, - image_tag: None, - cached_model: None, - datastream_frame_log: None, - run_id: 1, - vastai_yes: false, - vastai: None, - model: ChatModelConfig::default(), - pipeline_stages: 1, - max_tokens: DEFAULT_MAX_TOKENS, - skip_rebuild: true, - gpu_run: false, - relay_mode: None, - relay_url: None, - endpoint_addr_mask: EndpointAddrMask::Full, - } - } - - fn valid_vastai() -> ResolvedVastAiConfig { - ResolvedVastAiConfig { - api_key: "secret".to_owned(), - relay_url: "https://relay.example".to_owned(), - image: "docker.io/acme/node:latest".to_owned(), - bootstrap_command: "boot".to_owned(), - disk_gb: None, - gpu_name: None, - min_gpu_ram_mb: None, - min_down_mbps: None, - min_up_mbps: None, - max_dph_total: None, - min_reliability: None, - require_verified: None, - blacklist_hosts: Vec::new(), - onstart: None, - ssh_identity: None, - } - } - fn channel_lines(lines: &[&str]) -> mpsc::Receiver { let (tx, rx) = mpsc::channel(); for line in lines { @@ -2802,131 +2817,6 @@ kind = "mock" }); } - struct MockApproval { - terminal: bool, - answer: Result, - } - - impl VastAiApproval for MockApproval { - fn stdin_is_terminal(&self) -> bool { - self.terminal - } - - fn ask(&mut self) -> Result { - self.answer.clone() - } - } - - fn panic_prepare_node_image(_: NodeImageRequest) -> Result { - 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 prompt_loop_exits_cleanly_and_ignores_empty_prompts() { let mut rpc_writer = Vec::new(); diff --git a/crates/mvp-system/src/node/actor.rs b/crates/mvp-system/src/node/actor.rs index caffb64..9c37dad 100644 --- a/crates/mvp-system/src/node/actor.rs +++ b/crates/mvp-system/src/node/actor.rs @@ -347,7 +347,89 @@ impl NodeAgentActor { } } + fn forward_prompt_or_snapshot(&mut self, ctx: &Ctx, msg: NodeAgentMsg) -> Option { + match msg { + NodeAgentMsg::InferPrompt { + request_id, + prompt, + max_tokens, + reply_to, + } => { + if let Some(report_to) = self.report_to { + let _ = ctx.send( + report_to, + NodeAgentReport::PromptRequested { + request_id, + prompt, + max_tokens, + reply_to, + }, + ); + } + None + } + NodeAgentMsg::EncodePrompt { + request_id, + prompt, + reply_to, + } => { + if let Some(report_to) = self.report_to { + let _ = ctx.send( + report_to, + NodeAgentReport::EncodePromptRequested { + request_id, + prompt, + reply_to, + }, + ); + } + None + } + NodeAgentMsg::DecodeTokens { + request_id, + tokens, + reply_to, + } => { + if let Some(report_to) = self.report_to { + let _ = ctx.send( + report_to, + NodeAgentReport::DecodeTokensRequested { + request_id, + tokens, + reply_to, + }, + ); + } + None + } + NodeAgentMsg::Snapshot { reply_to } => { + let _ = ctx.send( + reply_to, + NodeAgentReport::Snapshot { + commands: self + .core + .commands() + .iter() + .map(StageCommandWire::from) + .collect(), + events: self + .core + .events() + .iter() + .map(StageLifecycleWire::from) + .collect(), + }, + ); + None + } + other => Some(other), + } + } + fn observe(&mut self, ctx: &Ctx, msg: NodeAgentMsg) { + let Some(msg) = self.forward_prompt_or_snapshot(ctx, msg) else { + return; + }; match msg { NodeAgentMsg::ProvisionStage(provision) => { self.core.observe(stage::StageEvent::ProvisionStage { @@ -480,79 +562,10 @@ impl NodeAgentActor { run_id: stage::RunId(run_id), }) } - NodeAgentMsg::InferPrompt { - request_id, - prompt, - max_tokens, - reply_to, - } => { - if let Some(report_to) = self.report_to { - let _ = ctx.send( - report_to, - NodeAgentReport::PromptRequested { - request_id, - prompt, - max_tokens, - reply_to, - }, - ); - } - return; - } - NodeAgentMsg::EncodePrompt { - request_id, - prompt, - reply_to, - } => { - if let Some(report_to) = self.report_to { - let _ = ctx.send( - report_to, - NodeAgentReport::EncodePromptRequested { - request_id, - prompt, - reply_to, - }, - ); - } - return; - } - NodeAgentMsg::DecodeTokens { - request_id, - tokens, - reply_to, - } => { - if let Some(report_to) = self.report_to { - let _ = ctx.send( - report_to, - NodeAgentReport::DecodeTokensRequested { - request_id, - tokens, - reply_to, - }, - ); - } - return; - } - NodeAgentMsg::Snapshot { reply_to } => { - let _ = ctx.send( - reply_to, - NodeAgentReport::Snapshot { - commands: self - .core - .commands() - .iter() - .map(StageCommandWire::from) - .collect(), - events: self - .core - .events() - .iter() - .map(StageLifecycleWire::from) - .collect(), - }, - ); - return; - } + NodeAgentMsg::InferPrompt { .. } + | NodeAgentMsg::EncodePrompt { .. } + | NodeAgentMsg::DecodeTokens { .. } + | NodeAgentMsg::Snapshot { .. } => unreachable!("prompt messages returned early"), } self.drain_outputs(ctx); } diff --git a/crates/mvp-system/src/node/worker_node_runtime.rs b/crates/mvp-system/src/node/worker_node_runtime.rs index 807ea7d..b7cdab7 100644 --- a/crates/mvp-system/src/node/worker_node_runtime.rs +++ b/crates/mvp-system/src/node/worker_node_runtime.rs @@ -1252,303 +1252,330 @@ impl WorkerEdgeRuntime { driver: &mut IrohDriver, ) -> Result<(), String> { loop { - let mut progressed = false; - while self.edge_command_cursor < self.establisher.commands().len() { - let command = self.establisher.commands()[self.edge_command_cursor].clone(); - self.edge_command_cursor += 1; - progressed = true; - match command { - edge::EdgeCommand::LeaseRing { - request_id, - ring_spec, - .. - } => { - let events = arena_manager.lock().request(arena::ArenaRequest::LeaseRing( - arena::LeaseRing { - request_id: arena::LeaseRequestId(request_id.0), - ring_spec: arena::RingSpec { - header_bytes: ring_spec.header_bytes, - data_bytes: ring_spec.data_bytes, - alignment: ring_spec.alignment, - }, - }, - )); - for event in events { - match event { - arena::ArenaEvent::RingLeased { lease } => { - self.establisher.observe(edge::EdgeEvent::RingLeased { - request_id: edge::LeaseRequestId(lease.request_id.0), - ring_id: edge::RingId(lease.ring_id.0), - layout: edge::RingLayout { - start_offset: lease.layout.start_offset, - header_offset: lease.layout.header_offset, - data_offset: lease.layout.data_offset, - end_offset: lease.layout.end_offset, - data_bytes: lease.layout.data_bytes, - alignment: lease.layout.alignment, - }, - }); - } - arena::ArenaEvent::RingLeaseRejected { request_id, reason } => { - let reason = match reason { - arena::RingLeaseRejection::CannotFitWithinCeiling => { - edge::RingLeaseRejection::CannotFit - } - arena::RingLeaseRejection::ArenaShuttingDown => { - edge::RingLeaseRejection::ArenaShuttingDown - } - }; - self.establisher - .observe(edge::EdgeEvent::RingLeaseRejected { - request_id: edge::LeaseRequestId(request_id.0), - reason, - }); - } - arena::ArenaEvent::RingLeaseQueued { .. } - | arena::ArenaEvent::RingReleased { .. } - | arena::ArenaEvent::RingReleaseRejected { .. } - | arena::ArenaEvent::CancelledFreshLeaseReleased { .. } => {} - } - } - } - edge::EdgeCommand::InstallWorkerRing { - edge_id, - ring_id, - direction, - object_spec, - .. - } => { - let lease = arena_manager - .lock() - .lookup_lease(arena::RingId(ring_id.0)) - .ok_or_else(|| format!("ring {} lease missing", ring_id.0))? - .clone(); - let (port, direction_name, wire_spec) = match direction { - edge::RingDirection::Ingress => { - self.inbound_ring_id = Some(ring_id.0); - let spec = self - .inbound_edge - .as_ref() - .map(|edge| edge.object_spec) - .unwrap_or(StageObjectSpecWire { - max_extent: object_spec.max_extent_bytes, - alignment: 4, - }); - ("input", "ingress", spec) - } - edge::RingDirection::Egress => { - self.outbound_ring_id = Some(ring_id.0); - let spec = self - .outbound_edge - .as_ref() - .map(|edge| edge.object_spec) - .unwrap_or(StageObjectSpecWire { - max_extent: object_spec.max_extent_bytes, - alignment: 4, - }); - ("output", "egress", spec) - } - }; - worker.install_ring( - ring_id.0, - edge_id.0, - port, - direction_name, - lease.layout, - wire_spec, - config, - datastream, - &mut || {}, - )?; - self.establisher - .observe(edge::EdgeEvent::RingInstalled { edge_id, ring_id }); - } - edge::EdgeCommand::EstablishSend { - edge_id, - consumer_node_id, - .. - } => { - let outbound = self - .outbound_edge - .as_ref() - .ok_or_else(|| "outbound edge missing".to_owned())?; - let peer = outbound - .consumer_endpoint - .clone() - .ok_or_else(|| "outbound consumer endpoint missing".to_owned())?; - let record = self - .establisher - .local_record(edge_id) - .ok_or_else(|| format!("edge {} record missing", edge_id.0))?; - let ring_id = record - .ring_id - .ok_or_else(|| format!("edge {} ring missing", edge_id.0))?; - let ring_capacity = outbound.ring_spec.data_capacity as usize; - self.driver_model - .observe(driver_model::DriverEvent::EstablishSend( - driver_model::EstablishSend { - edge_id: driver_model::EdgeId(edge_id.0), - peer_node_id: driver_model::NodeId(consumer_node_id.0), - layout: driver_model::RingLayout { - ring_id: driver_model::RingId(ring_id.0), - byte_capacity: ring_capacity, - direction: driver_model::RingDirection::Egress, - }, - }, - )); - self.outbound_sender = Some(driver.spawn_edge_send_pump(peer, edge_id.0)?); - } - edge::EdgeCommand::EstablishRecv { edge_id, .. } => { - let record = self - .establisher - .local_record(edge_id) - .ok_or_else(|| format!("edge {} record missing", edge_id.0))?; - let ring_id = record - .ring_id - .ok_or_else(|| format!("edge {} ring missing", edge_id.0))?; - let ring_capacity = self - .inbound_edge - .as_ref() - .map(|edge| edge.ring_spec.data_capacity as usize) - .unwrap_or(4096); - self.driver_model - .observe(driver_model::DriverEvent::EstablishRecv( - driver_model::EstablishRecv { - edge_id: driver_model::EdgeId(edge_id.0), - layout: driver_model::RingLayout { - ring_id: driver_model::RingId(ring_id.0), - byte_capacity: ring_capacity, - direction: driver_model::RingDirection::Ingress, - }, - }, - )); - } - edge::EdgeCommand::CancelQueuedLease { request_id, .. } => { - let _ = arena_manager - .lock() - .request(arena::ArenaRequest::CancelLease { - request_id: arena::LeaseRequestId(request_id.0), - }); - } - edge::EdgeCommand::StopPump { edge_id, .. } => { - self.driver_model - .observe(driver_model::DriverEvent::StopEdge { - edge_id: driver_model::EdgeId(edge_id.0), - }); - } - edge::EdgeCommand::UninstallWorkerRing { ring_id, .. } => { - let mut pump = || {}; - worker.uninstall_ring(ring_id.0, config, datastream, &mut pump)?; - self.establisher - .observe(edge::EdgeEvent::RingQuiesced { ring_id }); - } - edge::EdgeCommand::ReleaseArenaLease { ring_id, proof } => { - let proof = if proof == edge::QuiescenceProof::verified() { - arena::QuiescenceProof::verified() - } else { - arena::QuiescenceProof::missing() - }; - let _ = arena_manager - .lock() - .request(arena::ArenaRequest::ReleaseRing { - ring_id: arena::RingId(ring_id.0), - proof, - }); - } - } - } - - while self.driver_event_cursor < self.driver_model.events().len() { - let event = self.driver_model.events()[self.driver_event_cursor].clone(); - self.driver_event_cursor += 1; - progressed = true; - match event { - driver_model::DriverEventOut::DriverEdgeReady { edge_id } => { - self.establisher.observe(edge::EdgeEvent::DriverEdgeReady { - edge_id: edge::EdgeId(edge_id.0), - }); - } - driver_model::DriverEventOut::StreamFault { edge_id, reason } => { - let reason = match reason { - driver_model::StreamFaultReason::ReadError => { - edge::StreamFaultReason::ReadError - } - driver_model::StreamFaultReason::WriteError => { - edge::StreamFaultReason::WriteError - } - driver_model::StreamFaultReason::ProtocolError => { - edge::StreamFaultReason::ProtocolError - } - }; - self.establisher.observe(edge::EdgeEvent::StreamFault { - edge_id: edge::EdgeId(edge_id.0), - reason, - }); - } - driver_model::DriverEventOut::PumpStopped { edge_id, ring_id } => { - self.establisher.observe(edge::EdgeEvent::PumpStopped { - edge_id: edge::EdgeId(edge_id.0), - ring_id: edge::RingId(ring_id.0), - }); - } - driver_model::DriverEventOut::StreamClosed { .. } => {} - } - } - - while self.edge_event_cursor < self.establisher.events().len() { - let event = self.establisher.events()[self.edge_event_cursor].clone(); - self.edge_event_cursor += 1; - progressed = true; - match event { - edge::EdgeLifecycleEvent::EdgeReady { edge_id, .. } => { - if self - .inbound_edge - .as_ref() - .is_some_and(|edge| edge.edge_id == edge_id.0) - { - stack - .runtime - .send_to( - node_actor, - NodeAgentMsg::MarkInboundEdgeReady { edge_id: edge_id.0 }, - ) - .map_err(|e| format!("mark inbound ready: {e}"))?; - } - if self - .outbound_edge - .as_ref() - .is_some_and(|edge| edge.edge_id == edge_id.0) - { - stack - .runtime - .send_to( - node_actor, - NodeAgentMsg::MarkOutboundEdgeReady { edge_id: edge_id.0 }, - ) - .map_err(|e| format!("mark outbound ready: {e}"))?; - } - } - edge::EdgeLifecycleEvent::EdgeFaulted { edge_id, reason } => { - stack - .runtime - .send_to( - node_actor, - NodeAgentMsg::WorkerCrashed { - reason: Some(format!("edge {} faulted: {reason:?}", edge_id.0)), - }, - ) - .map_err(|e| format!("mark worker crashed after edge fault: {e}"))?; - return Err(format!("edge {} faulted: {reason:?}", edge_id.0)); - } - edge::EdgeLifecycleEvent::EdgeStopped { .. } => {} - } - } + let progressed = + self.drain_edge_commands(worker, arena_manager, config, datastream, driver)? + || self.drain_driver_events() + || self.drain_edge_events(stack, node_actor)?; if !progressed { break; } } Ok(()) } + + fn drain_edge_commands( + &mut self, + worker: &mut TinygradWorker, + arena_manager: &Arc>, + config: &DeploymentConfig, + datastream: &mut NodeDatastream, + driver: &mut IrohDriver, + ) -> Result { + let mut progressed = false; + while self.edge_command_cursor < self.establisher.commands().len() { + let command = self.establisher.commands()[self.edge_command_cursor].clone(); + self.edge_command_cursor += 1; + progressed = true; + match command { + edge::EdgeCommand::LeaseRing { + request_id, + ring_spec, + .. + } => { + let events = arena_manager.lock().request(arena::ArenaRequest::LeaseRing( + arena::LeaseRing { + request_id: arena::LeaseRequestId(request_id.0), + ring_spec: arena::RingSpec { + header_bytes: ring_spec.header_bytes, + data_bytes: ring_spec.data_bytes, + alignment: ring_spec.alignment, + }, + }, + )); + for event in events { + match event { + arena::ArenaEvent::RingLeased { lease } => { + self.establisher.observe(edge::EdgeEvent::RingLeased { + request_id: edge::LeaseRequestId(lease.request_id.0), + ring_id: edge::RingId(lease.ring_id.0), + layout: edge::RingLayout { + start_offset: lease.layout.start_offset, + header_offset: lease.layout.header_offset, + data_offset: lease.layout.data_offset, + end_offset: lease.layout.end_offset, + data_bytes: lease.layout.data_bytes, + alignment: lease.layout.alignment, + }, + }); + } + arena::ArenaEvent::RingLeaseRejected { request_id, reason } => { + let reason = match reason { + arena::RingLeaseRejection::CannotFitWithinCeiling => { + edge::RingLeaseRejection::CannotFit + } + arena::RingLeaseRejection::ArenaShuttingDown => { + edge::RingLeaseRejection::ArenaShuttingDown + } + }; + self.establisher + .observe(edge::EdgeEvent::RingLeaseRejected { + request_id: edge::LeaseRequestId(request_id.0), + reason, + }); + } + arena::ArenaEvent::RingLeaseQueued { .. } + | arena::ArenaEvent::RingReleased { .. } + | arena::ArenaEvent::RingReleaseRejected { .. } + | arena::ArenaEvent::CancelledFreshLeaseReleased { .. } => {} + } + } + } + edge::EdgeCommand::InstallWorkerRing { + edge_id, + ring_id, + direction, + object_spec, + .. + } => { + let lease = arena_manager + .lock() + .lookup_lease(arena::RingId(ring_id.0)) + .ok_or_else(|| format!("ring {} lease missing", ring_id.0))? + .clone(); + let (port, direction_name, wire_spec) = match direction { + edge::RingDirection::Ingress => { + self.inbound_ring_id = Some(ring_id.0); + let spec = self + .inbound_edge + .as_ref() + .map(|edge| edge.object_spec) + .unwrap_or(StageObjectSpecWire { + max_extent: object_spec.max_extent_bytes, + alignment: 4, + }); + ("input", "ingress", spec) + } + edge::RingDirection::Egress => { + self.outbound_ring_id = Some(ring_id.0); + let spec = self + .outbound_edge + .as_ref() + .map(|edge| edge.object_spec) + .unwrap_or(StageObjectSpecWire { + max_extent: object_spec.max_extent_bytes, + alignment: 4, + }); + ("output", "egress", spec) + } + }; + worker.install_ring( + ring_id.0, + edge_id.0, + port, + direction_name, + lease.layout, + wire_spec, + config, + datastream, + &mut || {}, + )?; + self.establisher + .observe(edge::EdgeEvent::RingInstalled { edge_id, ring_id }); + } + edge::EdgeCommand::EstablishSend { + edge_id, + consumer_node_id, + .. + } => { + let outbound = self + .outbound_edge + .as_ref() + .ok_or_else(|| "outbound edge missing".to_owned())?; + let peer = outbound + .consumer_endpoint + .clone() + .ok_or_else(|| "outbound consumer endpoint missing".to_owned())?; + let record = self + .establisher + .local_record(edge_id) + .ok_or_else(|| format!("edge {} record missing", edge_id.0))?; + let ring_id = record + .ring_id + .ok_or_else(|| format!("edge {} ring missing", edge_id.0))?; + let ring_capacity = outbound.ring_spec.data_capacity as usize; + self.driver_model + .observe(driver_model::DriverEvent::EstablishSend( + driver_model::EstablishSend { + edge_id: driver_model::EdgeId(edge_id.0), + peer_node_id: driver_model::NodeId(consumer_node_id.0), + layout: driver_model::RingLayout { + ring_id: driver_model::RingId(ring_id.0), + byte_capacity: ring_capacity, + direction: driver_model::RingDirection::Egress, + }, + }, + )); + self.outbound_sender = Some(driver.spawn_edge_send_pump(peer, edge_id.0)?); + } + edge::EdgeCommand::EstablishRecv { edge_id, .. } => { + let record = self + .establisher + .local_record(edge_id) + .ok_or_else(|| format!("edge {} record missing", edge_id.0))?; + let ring_id = record + .ring_id + .ok_or_else(|| format!("edge {} ring missing", edge_id.0))?; + let ring_capacity = self + .inbound_edge + .as_ref() + .map(|edge| edge.ring_spec.data_capacity as usize) + .unwrap_or(4096); + self.driver_model + .observe(driver_model::DriverEvent::EstablishRecv( + driver_model::EstablishRecv { + edge_id: driver_model::EdgeId(edge_id.0), + layout: driver_model::RingLayout { + ring_id: driver_model::RingId(ring_id.0), + byte_capacity: ring_capacity, + direction: driver_model::RingDirection::Ingress, + }, + }, + )); + } + edge::EdgeCommand::CancelQueuedLease { request_id, .. } => { + let _ = arena_manager + .lock() + .request(arena::ArenaRequest::CancelLease { + request_id: arena::LeaseRequestId(request_id.0), + }); + } + edge::EdgeCommand::StopPump { edge_id, .. } => { + self.driver_model + .observe(driver_model::DriverEvent::StopEdge { + edge_id: driver_model::EdgeId(edge_id.0), + }); + } + edge::EdgeCommand::UninstallWorkerRing { ring_id, .. } => { + let mut pump = || {}; + worker.uninstall_ring(ring_id.0, config, datastream, &mut pump)?; + self.establisher + .observe(edge::EdgeEvent::RingQuiesced { ring_id }); + } + edge::EdgeCommand::ReleaseArenaLease { ring_id, proof } => { + let proof = if proof == edge::QuiescenceProof::verified() { + arena::QuiescenceProof::verified() + } else { + arena::QuiescenceProof::missing() + }; + let _ = arena_manager + .lock() + .request(arena::ArenaRequest::ReleaseRing { + ring_id: arena::RingId(ring_id.0), + proof, + }); + } + } + } + Ok(progressed) + } + + fn drain_driver_events(&mut self) -> bool { + let mut progressed = false; + while self.driver_event_cursor < self.driver_model.events().len() { + let event = self.driver_model.events()[self.driver_event_cursor].clone(); + self.driver_event_cursor += 1; + progressed = true; + match event { + driver_model::DriverEventOut::DriverEdgeReady { edge_id } => { + self.establisher.observe(edge::EdgeEvent::DriverEdgeReady { + edge_id: edge::EdgeId(edge_id.0), + }); + } + driver_model::DriverEventOut::StreamFault { edge_id, reason } => { + let reason = match reason { + driver_model::StreamFaultReason::ReadError => { + edge::StreamFaultReason::ReadError + } + driver_model::StreamFaultReason::WriteError => { + edge::StreamFaultReason::WriteError + } + driver_model::StreamFaultReason::ProtocolError => { + edge::StreamFaultReason::ProtocolError + } + }; + self.establisher.observe(edge::EdgeEvent::StreamFault { + edge_id: edge::EdgeId(edge_id.0), + reason, + }); + } + driver_model::DriverEventOut::PumpStopped { edge_id, ring_id } => { + self.establisher.observe(edge::EdgeEvent::PumpStopped { + edge_id: edge::EdgeId(edge_id.0), + ring_id: edge::RingId(ring_id.0), + }); + } + driver_model::DriverEventOut::StreamClosed { .. } => {} + } + } + progressed + } + + fn drain_edge_events( + &mut self, + stack: &DistributionRuntimeStack, + node_actor: ActorAddress, + ) -> Result { + let mut progressed = false; + while self.edge_event_cursor < self.establisher.events().len() { + let event = self.establisher.events()[self.edge_event_cursor].clone(); + self.edge_event_cursor += 1; + progressed = true; + match event { + edge::EdgeLifecycleEvent::EdgeReady { edge_id, .. } => { + if self + .inbound_edge + .as_ref() + .is_some_and(|edge| edge.edge_id == edge_id.0) + { + stack + .runtime + .send_to( + node_actor, + NodeAgentMsg::MarkInboundEdgeReady { edge_id: edge_id.0 }, + ) + .map_err(|e| format!("mark inbound ready: {e}"))?; + } + if self + .outbound_edge + .as_ref() + .is_some_and(|edge| edge.edge_id == edge_id.0) + { + stack + .runtime + .send_to( + node_actor, + NodeAgentMsg::MarkOutboundEdgeReady { edge_id: edge_id.0 }, + ) + .map_err(|e| format!("mark outbound ready: {e}"))?; + } + } + edge::EdgeLifecycleEvent::EdgeFaulted { edge_id, reason } => { + stack + .runtime + .send_to( + node_actor, + NodeAgentMsg::WorkerCrashed { + reason: Some(format!("edge {} faulted: {reason:?}", edge_id.0)), + }, + ) + .map_err(|e| format!("mark worker crashed after edge fault: {e}"))?; + return Err(format!("edge {} faulted: {reason:?}", edge_id.0)); + } + edge::EdgeLifecycleEvent::EdgeStopped { .. } => {} + } + } + Ok(progressed) + } } fn duration_ms_u64(duration: Duration) -> u64 { diff --git a/crates/mvp-system/src/orchestration/app.rs b/crates/mvp-system/src/orchestration/app.rs index 2cc2b94..142c46f 100644 --- a/crates/mvp-system/src/orchestration/app.rs +++ b/crates/mvp-system/src/orchestration/app.rs @@ -30,8 +30,6 @@ use crate::observability::telemetry::{ mvp_provision_log_channel, }; use crate::orchestration::distribution_stack::DistributionRuntimeStack; -#[cfg(test)] -use crate::orchestration::provider_adapters::relay::relay_runtime_config_from_env; use crate::orchestration::provider_adapters::relay::{ MVP_IROH_RELAY_URL_ENV, RelayRuntimeConfig, SWACTOR_IROH_RELAY_URL_ENV, relay_mode_env_value, relay_runtime_config_from_settings, @@ -70,8 +68,6 @@ use iroh_driver::{ use parking_lot::Mutex; use serde_json::{Value, json}; use swactor::actor::ActorAddress; -#[cfg(test)] -use tokio::sync::mpsc as tokio_mpsc; const DEFAULT_IMAGE: &str = "swactor-mvp-node:latest"; const MVP_RUNTIME_CONFIG_ENV: &str = "MVP_RUNTIME_CONFIG"; @@ -506,16 +502,21 @@ where provisioner, &config, pipeline_plan.as_ref(), - &mut driver, - &stack, - &obs_rx, - &frame_rx, - &frame_tx, - &orchestrator_reports, - &stop_rx, - dashboard.as_ref(), - &mut orch_datastream, - orch_stdio_rx.as_ref(), + RuntimeReadyAckLoop { + driver: &mut driver, + stack: &stack, + obs_rx: &obs_rx, + frame_rx: &frame_rx, + frame_tx: &frame_tx, + orchestrator_reports: &orchestrator_reports, + stop_rx: &stop_rx, + dashboard: dashboard.as_ref(), + orch_datastream: &mut orch_datastream, + orch_stdio_rx: orch_stdio_rx.as_ref(), + run_id: config.run_id, + orchestrator_node_id: config.node_id, + provider: &config.provider, + }, sink, coordinator_endpoint, pipeline_coordinator_endpoint, @@ -580,26 +581,29 @@ where json!({"mode":"single_active_prompt","poll_interval_ms":PUMP_INTERVAL.as_millis()}), ); let result = serve_prompts( - &mut driver, - &stack, - &obs_rx, - &frame_rx, - &frame_tx, + RuntimeReadyAckLoop { + driver: &mut driver, + stack: &stack, + obs_rx: &obs_rx, + frame_rx: &frame_rx, + frame_tx: &frame_tx, + orchestrator_reports: &orchestrator_reports, + stop_rx: &stop_rx, + dashboard: dashboard.as_ref(), + orch_datastream: &mut orch_datastream, + orch_stdio_rx: orch_stdio_rx.as_ref(), + run_id: config.run_id, + orchestrator_node_id: config.node_id, + provider: &config.provider, + }, &work_rx, &prompt_events, - &stop_rx, - dashboard.as_ref(), - &mut orch_datastream, - orch_stdio_rx.as_ref(), - config.run_id, - config.node_id, ready.first_stage.node_actor, prompt_reply_actor, &tokenizer_events, ready.first_stage.node_actor, ready.final_stage.node_actor, tokenizer_reply_actor, - &config.provider, pipeline_plan.as_ref(), ready.first_stage.endpoint.clone(), ); @@ -659,7 +663,6 @@ struct VastAiRuntimeConfig { provisioning: VastAiProvisioningConfig, bootstrap_command: Option, ssh_identity: Option, - ssh_public_key: Option, ssh_public_fingerprint: Option, } @@ -767,7 +770,6 @@ impl VastAiRuntimeConfig { provisioning, bootstrap_command: builder.vastai_bootstrap_command.clone(), ssh_identity, - ssh_public_key: None, ssh_public_fingerprint: None, }) } @@ -1077,412 +1079,440 @@ impl ConfigBuilder { } fn overlay_toml(mut self, overlay: TomlConfigOverlay) -> Result { - if let Some(profile) = overlay.runtime.profile { + macro_rules! apply { + ($option:expr, |$value:ident| $body:block) => { + if let Some($value) = $option $body + }; + } + + apply!(overlay.runtime.profile, |profile| { self.config_profile = RuntimeConfigProfile::parse(&profile)?; - } - if let Some(run_id) = overlay.runtime.run_id { - self.run_id = run_id; - } - if let Some(node_id) = overlay.runtime.node_id { - self.node_id = node_id; - } - if let Some(stage_index) = overlay.runtime.stage_index { - self.stage_index = stage_index; - } - if let Some(layer_end_exclusive) = overlay.runtime.layer_end_exclusive { - self.layer_end_exclusive = Some(layer_end_exclusive); - } - if let Some(pipeline_stages) = overlay.runtime.pipeline_stages { - self.pipeline_stages = pipeline_stages; - } - if let Some(provider) = overlay.provider.kind { + }); + apply!(overlay.runtime.run_id, |run_id| { self.run_id = run_id }); + apply!(overlay.runtime.node_id, |node_id| { + self.node_id = node_id + }); + apply!(overlay.runtime.stage_index, |stage_index| { + self.stage_index = stage_index + }); + apply!(overlay.runtime.layer_end_exclusive, |layer_end_exclusive| { + self.layer_end_exclusive = Some(layer_end_exclusive) + }); + apply!(overlay.runtime.pipeline_stages, |pipeline_stages| { + self.pipeline_stages = pipeline_stages + }); + apply!(overlay.provider.kind, |provider| { self.provider = Some(provider_kind::parse_deploy(&provider)?); - } - if let Some(image) = overlay.image.node { - self.image = image; - } - if let Some(mode) = overlay.relay.mode { - self.relay_mode = Some(mode); - } - if let Some(url) = overlay.relay.url { - self.relay_url = Some(url); - } - if let Some(rpc_bind) = overlay.prompt.rpc_addr { + }); + apply!(overlay.image.node, |image| { self.image = image }); + apply!(overlay.relay.mode, |mode| { self.relay_mode = Some(mode) }); + apply!(overlay.relay.url, |url| { self.relay_url = Some(url) }); + apply!(overlay.prompt.rpc_addr, |rpc_bind| { self.rpc_bind = rpc_bind; self.rpc_bind_label = "[prompt].rpc_addr"; - } - if let Some(max_tokens) = overlay.prompt.max_tokens { - self.default_max_tokens = max_tokens; - } - if let Some(dashboard) = overlay.prompt.dashboard { - self.dashboard = dashboard; - } - if let Some(model_id) = overlay.model.id { - self.model_id = model_id; - } - if let Some(path) = overlay.model.gguf_local_path { - self.gguf_source = GgufSource::LocalPath(path); - } - if let Some(repo) = overlay.model.gguf_repo { - self.set_gguf_repo(repo); - } - if let Some(file) = overlay.model.gguf_file { - self.set_gguf_file(file); - } - if let Some(revision) = overlay.model.gguf_revision { - self.set_gguf_revision(Some(revision)); - } - if let Some(path) = overlay.model.tokenizer_local_path { - self.tokenizer = TokenizerSource::LocalPath(path); - } - if let Some(max_context) = overlay.model.max_context { - self.max_context = Some(max_context); - } - if let Some(gpus) = overlay.docker.gpus { - self.docker_gpus = gpus; - } - if let Some(path) = overlay.docker.cached_model_host_path { - self.cached_model_host_path = Some(PathBuf::from(path)); - } - if let Some(path) = overlay.observability.datastream_frame_log { - self.datastream_frame_log = Some(PathBuf::from(path)); - } - if let Some(image) = overlay.vastai.image { - self.toml_vastai_image = Some(image); - } - if let Some(api_key) = overlay.vastai.api_key { - self.vastai_api_key = Some(api_key); - } - if let Some(command) = overlay.vastai.bootstrap_command { - self.vastai_bootstrap_command = Some(command); - } - if let Some(disk_gb) = overlay.vastai.disk_gb { - self.vastai_disk_gb = Some(disk_gb); - } - if let Some(ssh_user) = overlay.vastai.ssh_user { - self.vastai_ssh_user = Some(ssh_user); - } - if let Some(confirm_lease) = overlay.vastai.confirm_lease { - self.vastai_confirm_lease = Some(confirm_lease); - } - if let Some(onstart) = overlay.vastai.onstart { - self.vastai_onstart = Some(onstart); - } - if let Some(identity) = overlay.vastai.ssh_identity { - self.vastai_ssh_identity_raw = Some(identity); - } - if let Some(gpu_name) = overlay.vastai.gpu_name { - self.vastai_gpu_name = Some(gpu_name); - } - if let Some(min_gpu_ram_mb) = overlay.vastai.min_gpu_ram_mb { - self.vastai_min_gpu_ram_mb = Some(min_gpu_ram_mb); - } - if let Some(min_down_mbps) = overlay.vastai.min_down_mbps { - self.vastai_min_down_mbps = Some(min_down_mbps); - } - if let Some(max_dph_total) = overlay.vastai.max_dph_total { - self.vastai_max_dph_total = Some(max_dph_total); - } - if let Some(min_up_mbps) = overlay.vastai.min_up_mbps { - self.vastai_min_up_mbps = Some(min_up_mbps); - } - if let Some(min_reliability) = overlay.vastai.min_reliability { - self.vastai_min_reliability = Some(min_reliability); - } - if let Some(require_verified) = overlay.vastai.require_verified { - self.vastai_require_verified = Some(require_verified); - } + }); + apply!(overlay.prompt.max_tokens, |max_tokens| { + self.default_max_tokens = max_tokens + }); + apply!(overlay.prompt.dashboard, |dashboard| { + self.dashboard = dashboard + }); + apply!(overlay.model.id, |model_id| { self.model_id = model_id }); + apply!(overlay.model.gguf_local_path, |path| { + self.gguf_source = GgufSource::LocalPath(path) + }); + apply!(overlay.model.gguf_repo, |repo| { self.set_gguf_repo(repo) }); + apply!(overlay.model.gguf_file, |file| { self.set_gguf_file(file) }); + apply!(overlay.model.gguf_revision, |revision| { + self.set_gguf_revision(Some(revision)) + }); + apply!(overlay.model.tokenizer_local_path, |path| { + self.tokenizer = TokenizerSource::LocalPath(path) + }); + apply!(overlay.model.max_context, |max_context| { + self.max_context = Some(max_context) + }); + apply!(overlay.docker.gpus, |gpus| { self.docker_gpus = gpus }); + apply!(overlay.docker.cached_model_host_path, |path| { + self.cached_model_host_path = Some(PathBuf::from(path)) + }); + apply!(overlay.observability.datastream_frame_log, |path| { + self.datastream_frame_log = Some(PathBuf::from(path)) + }); + apply!(overlay.vastai.image, |image| { + self.toml_vastai_image = Some(image) + }); + apply!(overlay.vastai.api_key, |api_key| { + self.vastai_api_key = Some(api_key) + }); + apply!(overlay.vastai.bootstrap_command, |command| { + self.vastai_bootstrap_command = Some(command) + }); + apply!(overlay.vastai.disk_gb, |disk_gb| { + self.vastai_disk_gb = Some(disk_gb) + }); + apply!(overlay.vastai.ssh_user, |ssh_user| { + self.vastai_ssh_user = Some(ssh_user) + }); + apply!(overlay.vastai.confirm_lease, |confirm_lease| { + self.vastai_confirm_lease = Some(confirm_lease) + }); + apply!(overlay.vastai.onstart, |onstart| { + self.vastai_onstart = Some(onstart) + }); + apply!(overlay.vastai.ssh_identity, |identity| { + self.vastai_ssh_identity_raw = Some(identity) + }); + apply!(overlay.vastai.gpu_name, |gpu_name| { + self.vastai_gpu_name = Some(gpu_name) + }); + apply!(overlay.vastai.min_gpu_ram_mb, |min_gpu_ram_mb| { + self.vastai_min_gpu_ram_mb = Some(min_gpu_ram_mb) + }); + apply!(overlay.vastai.min_down_mbps, |min_down_mbps| { + self.vastai_min_down_mbps = Some(min_down_mbps) + }); + apply!(overlay.vastai.max_dph_total, |max_dph_total| { + self.vastai_max_dph_total = Some(max_dph_total) + }); + apply!(overlay.vastai.min_up_mbps, |min_up_mbps| { + self.vastai_min_up_mbps = Some(min_up_mbps) + }); + apply!(overlay.vastai.min_reliability, |min_reliability| { + self.vastai_min_reliability = Some(min_reliability) + }); + apply!(overlay.vastai.require_verified, |require_verified| { + self.vastai_require_verified = Some(require_verified) + }); for host_id in overlay.vastai.blacklist_hosts { self.push_vastai_blacklist_host(host_id); } - if let Some(poll_interval_secs) = overlay.vastai.poll_interval_secs { - self.vastai_poll_interval_secs = Some(poll_interval_secs); - } + apply!(overlay.vastai.poll_interval_secs, |poll_interval_secs| { + self.vastai_poll_interval_secs = Some(poll_interval_secs) + }); Ok(self) } fn overlay_env(mut self) -> Result { - if let Some(profile) = env_optional(MVP_RUNTIME_CONFIG_ENV) { + macro_rules! apply { + ($option:expr, |$value:ident| $body:block) => { + if let Some($value) = $option $body + }; + } + macro_rules! env_apply { + ($name:expr, |$value:ident| $body:block) => { + apply!(env_optional($name), |$value| $body) + }; + } + macro_rules! env_parse { + ($name:expr, |$value:ident| $body:block) => { + env_apply!($name, |raw| { + let $value = Self::parse_value($name, &raw)?; + $body + }) + }; + } + + env_apply!(MVP_RUNTIME_CONFIG_ENV, |profile| { self.config_profile = RuntimeConfigProfile::parse(&profile)?; - } - if let Some(run_id) = env_optional("MVP_RUN_ID") { - self.run_id = Self::parse_value("MVP_RUN_ID", &run_id)?; - } - if let Some(node_id) = env_optional("MVP_LOGICAL_NODE_ID") { - self.node_id = Self::parse_value("MVP_LOGICAL_NODE_ID", &node_id)?; - } - if let Some(stage_index) = env_optional("MVP_STAGE_INDEX") { - self.stage_index = Self::parse_value("MVP_STAGE_INDEX", &stage_index)?; - } - if let Some(layer_end_exclusive) = env_optional("MVP_LAYER_END_EXCLUSIVE") { - self.layer_end_exclusive = Some(Self::parse_value( - "MVP_LAYER_END_EXCLUSIVE", - &layer_end_exclusive, - )?); - } - if let Some(pipeline_stages) = env_optional("MVP_PIPELINE_STAGES") { - self.pipeline_stages = Self::parse_value("MVP_PIPELINE_STAGES", &pipeline_stages)?; - } - if let Some(provider) = - env_optional("MVP_NODE_PROVIDER").or_else(|| env_optional("MVP_PROVIDER")) - { - self.provider = Some(provider_kind::parse_deploy(&provider)?); - } - if let Some(image) = env_optional("MVP_NODE_IMAGE") { - self.set_process_image(image); - } - if let Some(gpus) = env_optional("MVP_DOCKER_GPUS") { - self.docker_gpus = gpus; - } - if let Some(path) = env_optional(CACHED_MODEL_HOST_ENV) { - self.cached_model_host_path = Some(PathBuf::from(path)); - } - if let Some(path) = env_optional(MVP_WORKER_BIN_ENV) { - self.worker_bin = Some(PathBuf::from(path)); - } - if let Some(rpc_bind) = env_optional("MVP_PROMPT_RPC_BIND") { + }); + env_parse!("MVP_RUN_ID", |run_id| { self.run_id = run_id }); + env_parse!("MVP_LOGICAL_NODE_ID", |node_id| { self.node_id = node_id }); + env_parse!("MVP_STAGE_INDEX", |stage_index| { + self.stage_index = stage_index + }); + env_parse!("MVP_LAYER_END_EXCLUSIVE", |layer_end_exclusive| { + self.layer_end_exclusive = Some(layer_end_exclusive) + }); + env_parse!("MVP_PIPELINE_STAGES", |pipeline_stages| { + self.pipeline_stages = pipeline_stages + }); + apply!( + env_optional("MVP_NODE_PROVIDER").or_else(|| env_optional("MVP_PROVIDER")), + |provider| { + self.provider = Some(provider_kind::parse_deploy(&provider)?); + } + ); + env_apply!("MVP_NODE_IMAGE", |image| { self.set_process_image(image) }); + env_apply!("MVP_DOCKER_GPUS", |gpus| { self.docker_gpus = gpus }); + env_apply!(CACHED_MODEL_HOST_ENV, |path| { + self.cached_model_host_path = Some(PathBuf::from(path)) + }); + env_apply!(MVP_WORKER_BIN_ENV, |path| { + self.worker_bin = Some(PathBuf::from(path)) + }); + env_apply!("MVP_PROMPT_RPC_BIND", |rpc_bind| { self.rpc_bind = rpc_bind; self.rpc_bind_label = "MVP_PROMPT_RPC_BIND"; - } - if let Some(max_tokens) = env_optional("MVP_PROMPT_MAX_TOKENS") { - self.default_max_tokens = Self::parse_value("MVP_PROMPT_MAX_TOKENS", &max_tokens)?; - } - if let Some(dashboard) = env_optional("MVP_DASHBOARD") { + }); + env_parse!("MVP_PROMPT_MAX_TOKENS", |max_tokens| { + self.default_max_tokens = max_tokens + }); + env_apply!("MVP_DASHBOARD", |dashboard| { self.dashboard = Self::parse_bool("MVP_DASHBOARD", &dashboard)?; - } - if let Some(path) = env_optional(DATASTREAM_FRAME_LOG_ENV) { - self.datastream_frame_log = Some(PathBuf::from(path)); - } - if let Some(model_id) = env_optional("MVP_MODEL_ID") { - self.model_id = model_id; - } - if let Some(path) = env_optional("MVP_GGUF_LOCAL_PATH") { - self.gguf_source = GgufSource::LocalPath(path); - } - if let Some(repo) = env_optional("MVP_GGUF_REPO") { - self.set_gguf_repo(repo); - } - if let Some(file) = env_optional("MVP_GGUF_FILE") { - self.set_gguf_file(file); - } - if let Some(revision) = env_optional("MVP_GGUF_REVISION") { - self.set_gguf_revision(Some(revision)); - } - if let Some(path) = env_optional("MVP_TOKENIZER_LOCAL_PATH") { - self.tokenizer = TokenizerSource::LocalPath(path); - } - if let Some(max_context) = env_optional("MVP_MAX_CONTEXT") { - self.max_context = Some(Self::parse_value("MVP_MAX_CONTEXT", &max_context)?); - } - if let Some(mode) = env_optional("MVP_IROH_RELAY_MODE") { - self.relay_mode = Some(mode.to_ascii_lowercase()); - } - if let Some(url) = env_optional(MVP_IROH_RELAY_URL_ENV) - .or_else(|| env_optional(SWACTOR_IROH_RELAY_URL_ENV)) - { - self.relay_url = Some(url); - } - if let Some(mask) = env_optional(MVP_IROH_ENDPOINT_ADDR_MASK_ENV) { - self.endpoint_addr_mask = Some(mask); - } - if let Some(api_key) = env_optional("VAST_API_KEY") - .or_else(|| env_optional("MVP_VASTAI_API_KEY")) - .or_else(|| env_optional("VASTAI_API_KEY")) - { - self.vastai_api_key = Some(api_key); - } - if let Some(command) = env_optional("MVP_VASTAI_BOOTSTRAP_COMMAND") { - self.vastai_bootstrap_command = Some(command); - } - if let Some(identity) = env_optional("MVP_VASTAI_SSH_IDENTITY") { - self.vastai_ssh_identity_raw = Some(identity); - } - if let Some(disk_gb) = env_optional("MVP_VASTAI_DISK_GB") { - self.vastai_disk_gb_raw = Some(disk_gb); - } - if let Some(ssh_user) = env_optional("MVP_VASTAI_SSH_USER") { - self.vastai_ssh_user = Some(ssh_user); - } - if let Some(confirm_lease) = env_optional("MVP_VASTAI_CONFIRM_LEASE") { - self.vastai_confirm_lease_raw = Some(confirm_lease); - } - if let Some(onstart) = env_optional("MVP_VASTAI_ONSTART") { - self.vastai_onstart = Some(onstart); - } - if let Some(gpu_name) = env_optional("MVP_VASTAI_GPU_NAME") { - self.vastai_gpu_name = Some(gpu_name); - } - if let Some(min_gpu_ram_mb) = env_optional("MVP_VASTAI_MIN_GPU_RAM_MB") { - self.vastai_min_gpu_ram_mb_raw = Some(min_gpu_ram_mb); - } - if let Some(min_down_mbps) = env_optional("MVP_VASTAI_MIN_DOWN_MBPS") { - self.vastai_min_down_mbps_raw = Some(min_down_mbps); - } - if let Some(max_dph_total) = env_optional("MVP_VASTAI_MAX_DPH_TOTAL") { - self.vastai_max_dph_total_raw = Some(max_dph_total); - } - if let Some(min_up_mbps) = env_optional("MVP_VASTAI_MIN_UP_MBPS") { - self.vastai_min_up_mbps_raw = Some(min_up_mbps); - } - if let Some(min_reliability) = env_optional("MVP_VASTAI_MIN_RELIABILITY") { - self.vastai_min_reliability_raw = Some(min_reliability); - } - if let Some(require_verified) = env_optional("MVP_VASTAI_REQUIRE_VERIFIED") { - self.vastai_require_verified_raw = Some(require_verified); - } - if let Some(blacklist_hosts) = env_optional("MVP_VASTAI_BLACKLIST_HOSTS") { + }); + env_apply!(DATASTREAM_FRAME_LOG_ENV, |path| { + self.datastream_frame_log = Some(PathBuf::from(path)) + }); + env_apply!("MVP_MODEL_ID", |model_id| { self.model_id = model_id }); + env_apply!("MVP_GGUF_LOCAL_PATH", |path| { + self.gguf_source = GgufSource::LocalPath(path) + }); + env_apply!("MVP_GGUF_REPO", |repo| { self.set_gguf_repo(repo) }); + env_apply!("MVP_GGUF_FILE", |file| { self.set_gguf_file(file) }); + env_apply!("MVP_GGUF_REVISION", |revision| { + self.set_gguf_revision(Some(revision)) + }); + env_apply!("MVP_TOKENIZER_LOCAL_PATH", |path| { + self.tokenizer = TokenizerSource::LocalPath(path) + }); + env_parse!("MVP_MAX_CONTEXT", |max_context| { + self.max_context = Some(max_context) + }); + env_apply!("MVP_IROH_RELAY_MODE", |mode| { + self.relay_mode = Some(mode.to_ascii_lowercase()) + }); + apply!( + env_optional(MVP_IROH_RELAY_URL_ENV) + .or_else(|| env_optional(SWACTOR_IROH_RELAY_URL_ENV)), + |url| { + self.relay_url = Some(url); + } + ); + env_apply!(MVP_IROH_ENDPOINT_ADDR_MASK_ENV, |mask| { + self.endpoint_addr_mask = Some(mask) + }); + apply!( + env_optional("VAST_API_KEY") + .or_else(|| env_optional("MVP_VASTAI_API_KEY")) + .or_else(|| env_optional("VASTAI_API_KEY")), + |api_key| { + self.vastai_api_key = Some(api_key); + } + ); + env_apply!("MVP_VASTAI_BOOTSTRAP_COMMAND", |command| { + self.vastai_bootstrap_command = Some(command) + }); + env_apply!("MVP_VASTAI_SSH_IDENTITY", |identity| { + self.vastai_ssh_identity_raw = Some(identity) + }); + env_apply!("MVP_VASTAI_DISK_GB", |disk_gb| { + self.vastai_disk_gb_raw = Some(disk_gb) + }); + env_apply!("MVP_VASTAI_SSH_USER", |ssh_user| { + self.vastai_ssh_user = Some(ssh_user) + }); + env_apply!("MVP_VASTAI_CONFIRM_LEASE", |confirm_lease| { + self.vastai_confirm_lease_raw = Some(confirm_lease) + }); + env_apply!("MVP_VASTAI_ONSTART", |onstart| { + self.vastai_onstart = Some(onstart) + }); + env_apply!("MVP_VASTAI_GPU_NAME", |gpu_name| { + self.vastai_gpu_name = Some(gpu_name) + }); + env_apply!("MVP_VASTAI_MIN_GPU_RAM_MB", |min_gpu_ram_mb| { + self.vastai_min_gpu_ram_mb_raw = Some(min_gpu_ram_mb) + }); + env_apply!("MVP_VASTAI_MIN_DOWN_MBPS", |min_down_mbps| { + self.vastai_min_down_mbps_raw = Some(min_down_mbps) + }); + env_apply!("MVP_VASTAI_MAX_DPH_TOTAL", |max_dph_total| { + self.vastai_max_dph_total_raw = Some(max_dph_total) + }); + env_apply!("MVP_VASTAI_MIN_UP_MBPS", |min_up_mbps| { + self.vastai_min_up_mbps_raw = Some(min_up_mbps) + }); + env_apply!("MVP_VASTAI_MIN_RELIABILITY", |min_reliability| { + self.vastai_min_reliability_raw = Some(min_reliability) + }); + env_apply!("MVP_VASTAI_REQUIRE_VERIFIED", |require_verified| { + self.vastai_require_verified_raw = Some(require_verified) + }); + env_apply!("MVP_VASTAI_BLACKLIST_HOSTS", |blacklist_hosts| { for host_id in Self::parse_list("MVP_VASTAI_BLACKLIST_HOSTS", &blacklist_hosts)? { self.push_vastai_blacklist_host(host_id); } - } - if let Some(poll_interval_secs) = env_optional("MVP_VASTAI_POLL_INTERVAL_SECS") { - self.vastai_poll_interval_secs_raw = Some(poll_interval_secs); - } + }); + env_apply!("MVP_VASTAI_POLL_INTERVAL_SECS", |poll_interval_secs| { + self.vastai_poll_interval_secs_raw = Some(poll_interval_secs) + }); Ok(self) } + fn apply_core_cli_arg(&mut self, arg: &str, args: &mut I) -> Result + where + I: Iterator, + { + match arg { + "--runtime-config" => { + self.config_profile = + RuntimeConfigProfile::parse(&next_arg(args, "--runtime-config")?)? + } + "--provider" => { + self.provider = Some(provider_kind::parse_deploy(&next_arg(args, "--provider")?)?) + } + "--worker-bin" => { + self.worker_bin = Some(PathBuf::from(next_arg(args, "--worker-bin")?)) + } + "--image" => self.set_process_image(next_arg(args, "--image")?), + "--gpus" => self.docker_gpus = next_arg(args, "--gpus")?, + "--rpc-bind" => { + self.rpc_bind = next_arg(args, "--rpc-bind")?; + self.rpc_bind_label = "--rpc-bind"; + } + "--run-id" => self.run_id = parse_next(args, "--run-id")?, + "--node-id" => self.node_id = parse_next(args, "--node-id")?, + "--stage-index" => self.stage_index = parse_next(args, "--stage-index")?, + "--layer-end-exclusive" => { + self.layer_end_exclusive = Some(parse_next(args, "--layer-end-exclusive")?) + } + "-N" | "--pipeline-stages" => self.pipeline_stages = parse_next(args, arg)?, + "--max-tokens" => self.default_max_tokens = parse_next(args, "--max-tokens")?, + "--dashboard" => self.dashboard = true, + "--no-dashboard" => self.dashboard = false, + "--datastream-frame-log" => { + self.datastream_frame_log = + Some(PathBuf::from(next_arg(args, "--datastream-frame-log")?)); + } + _ => return Ok(false), + } + Ok(true) + } + + fn apply_model_cli_arg(&mut self, arg: &str, args: &mut I) -> Result + where + I: Iterator, + { + match arg { + "--model-id" => self.model_id = next_arg(args, "--model-id")?, + "--gguf-local-path" => { + self.gguf_source = GgufSource::LocalPath(next_arg(args, "--gguf-local-path")?) + } + "--gguf-repo" => self.set_gguf_repo(next_arg(args, "--gguf-repo")?), + "--gguf-file" => self.set_gguf_file(next_arg(args, "--gguf-file")?), + "--gguf-revision" => self.set_gguf_revision(Some(next_arg(args, "--gguf-revision")?)), + "--tokenizer-local-path" => { + self.tokenizer = + TokenizerSource::LocalPath(next_arg(args, "--tokenizer-local-path")?) + } + "--max-context" => self.max_context = Some(parse_next(args, "--max-context")?), + "--cached-model-host-path" => { + self.cached_model_host_path = + Some(PathBuf::from(next_arg(args, "--cached-model-host-path")?)); + } + "--relay-mode" => self.relay_mode = Some(next_arg(args, "--relay-mode")?), + "--relay-url" => self.relay_url = Some(next_arg(args, "--relay-url")?), + "--endpoint-addr-mask" => { + self.endpoint_addr_mask = Some(next_arg(args, "--endpoint-addr-mask")?) + } + _ => return Ok(false), + } + Ok(true) + } + + fn apply_vastai_string_cli_arg(&mut self, arg: &str, args: &mut I) -> Result + where + I: Iterator, + { + match arg { + "--vastai-api-key" => self.vastai_api_key = Some(next_arg(args, "--vastai-api-key")?), + "--vastai-bootstrap-command" => { + self.vastai_bootstrap_command = Some(next_arg(args, "--vastai-bootstrap-command")?) + } + "--vastai-ssh-identity" => { + self.vastai_ssh_identity_raw = Some(next_arg(args, "--vastai-ssh-identity")?); + } + "--vastai-ssh-user" => { + self.vastai_ssh_user = Some(next_arg(args, "--vastai-ssh-user")?) + } + "--vastai-onstart" => self.vastai_onstart = Some(next_arg(args, "--vastai-onstart")?), + "--vastai-gpu-name" => { + self.vastai_gpu_name = Some(next_arg(args, "--vastai-gpu-name")?) + } + _ => return Ok(false), + } + Ok(true) + } + + fn apply_vastai_numeric_cli_arg(&mut self, arg: &str, args: &mut I) -> Result + where + I: Iterator, + { + match arg { + "--vastai-disk-gb" => { + self.vastai_disk_gb = Some(parse_next(args, "--vastai-disk-gb")?); + self.vastai_disk_gb_raw = None; + } + "--vastai-min-gpu-ram-mb" => { + self.vastai_min_gpu_ram_mb = Some(parse_next(args, "--vastai-min-gpu-ram-mb")?); + self.vastai_min_gpu_ram_mb_raw = None; + } + "--vastai-min-down-mbps" => { + self.vastai_min_down_mbps = Some(parse_next(args, "--vastai-min-down-mbps")?); + self.vastai_min_down_mbps_raw = None; + } + "--vastai-max-dph-total" => { + self.vastai_max_dph_total = Some(parse_next(args, "--vastai-max-dph-total")?); + self.vastai_max_dph_total_raw = None; + } + "--vastai-min-up-mbps" => { + self.vastai_min_up_mbps = Some(parse_next(args, "--vastai-min-up-mbps")?); + self.vastai_min_up_mbps_raw = None; + } + "--vastai-min-reliability" => { + self.vastai_min_reliability = Some(parse_next(args, "--vastai-min-reliability")?); + self.vastai_min_reliability_raw = None; + } + "--vastai-blacklist-host" => { + let host_id = parse_next(args, "--vastai-blacklist-host")?; + self.push_vastai_blacklist_host(host_id); + } + "--vastai-poll-interval-secs" => { + self.vastai_poll_interval_secs = + Some(parse_next(args, "--vastai-poll-interval-secs")?); + self.vastai_poll_interval_secs_raw = None; + } + _ => return Ok(false), + } + Ok(true) + } + + fn apply_vastai_bool_cli_arg(&mut self, arg: &str) -> bool { + match arg { + "--vastai-confirm-lease" => { + self.vastai_confirm_lease = Some(true); + self.vastai_confirm_lease_raw = None; + } + "--no-vastai-confirm-lease" => { + self.vastai_confirm_lease = Some(false); + self.vastai_confirm_lease_raw = None; + } + "--vastai-require-verified" => { + self.vastai_require_verified = Some(true); + self.vastai_require_verified_raw = None; + } + "--no-vastai-require-verified" => { + self.vastai_require_verified = Some(false); + self.vastai_require_verified_raw = None; + } + _ => return false, + } + true + } + fn overlay_cli(mut self, args: impl IntoIterator) -> Result { let mut args = args.into_iter(); while let Some(arg) = args.next() { - match arg.as_str() { - "--runtime-config" => { - self.config_profile = - RuntimeConfigProfile::parse(&next_arg(&mut args, "--runtime-config")?)? - } - "--provider" => { - self.provider = Some(provider_kind::parse_deploy(&next_arg( - &mut args, - "--provider", - )?)?) - } - "--worker-bin" => { - self.worker_bin = Some(PathBuf::from(next_arg(&mut args, "--worker-bin")?)) - } - "--image" => self.set_process_image(next_arg(&mut args, "--image")?), - "--gpus" => self.docker_gpus = next_arg(&mut args, "--gpus")?, - "--rpc-bind" => { - self.rpc_bind = next_arg(&mut args, "--rpc-bind")?; - self.rpc_bind_label = "--rpc-bind"; - } - "--run-id" => self.run_id = parse_next(&mut args, "--run-id")?, - "--node-id" => self.node_id = parse_next(&mut args, "--node-id")?, - "--stage-index" => self.stage_index = parse_next(&mut args, "--stage-index")?, - "--layer-end-exclusive" => { - self.layer_end_exclusive = Some(parse_next(&mut args, "--layer-end-exclusive")?) - } - "-N" | "--pipeline-stages" => { - self.pipeline_stages = parse_next(&mut args, arg.as_str())? - } - "--max-tokens" => self.default_max_tokens = parse_next(&mut args, "--max-tokens")?, - "--dashboard" => self.dashboard = true, - "--no-dashboard" => self.dashboard = false, - "--datastream-frame-log" => { - self.datastream_frame_log = Some(PathBuf::from(next_arg( - &mut args, - "--datastream-frame-log", - )?)); - } - "--model-id" => self.model_id = next_arg(&mut args, "--model-id")?, - "--gguf-local-path" => { - self.gguf_source = - GgufSource::LocalPath(next_arg(&mut args, "--gguf-local-path")?) - } - "--gguf-repo" => self.set_gguf_repo(next_arg(&mut args, "--gguf-repo")?), - "--gguf-file" => self.set_gguf_file(next_arg(&mut args, "--gguf-file")?), - "--gguf-revision" => { - self.set_gguf_revision(Some(next_arg(&mut args, "--gguf-revision")?)) - } - "--tokenizer-local-path" => { - self.tokenizer = - TokenizerSource::LocalPath(next_arg(&mut args, "--tokenizer-local-path")?) - } - "--max-context" => self.max_context = Some(parse_next(&mut args, "--max-context")?), - "--cached-model-host-path" => { - self.cached_model_host_path = Some(PathBuf::from(next_arg( - &mut args, - "--cached-model-host-path", - )?)); - } - "--relay-mode" => self.relay_mode = Some(next_arg(&mut args, "--relay-mode")?), - "--relay-url" => self.relay_url = Some(next_arg(&mut args, "--relay-url")?), - "--endpoint-addr-mask" => { - self.endpoint_addr_mask = Some(next_arg(&mut args, "--endpoint-addr-mask")?) - } - "--vastai-api-key" => { - self.vastai_api_key = Some(next_arg(&mut args, "--vastai-api-key")?) - } - "--vastai-bootstrap-command" => { - self.vastai_bootstrap_command = - Some(next_arg(&mut args, "--vastai-bootstrap-command")?) - } - "--vastai-ssh-identity" => { - self.vastai_ssh_identity_raw = - Some(next_arg(&mut args, "--vastai-ssh-identity")?); - } - "--vastai-disk-gb" => { - self.vastai_disk_gb = Some(parse_next(&mut args, "--vastai-disk-gb")?); - self.vastai_disk_gb_raw = None; - } - "--vastai-ssh-user" => { - self.vastai_ssh_user = Some(next_arg(&mut args, "--vastai-ssh-user")?) - } - "--vastai-confirm-lease" => { - self.vastai_confirm_lease = Some(true); - self.vastai_confirm_lease_raw = None; - } - "--no-vastai-confirm-lease" => { - self.vastai_confirm_lease = Some(false); - self.vastai_confirm_lease_raw = None; - } - "--vastai-onstart" => { - self.vastai_onstart = Some(next_arg(&mut args, "--vastai-onstart")?) - } - "--vastai-gpu-name" => { - self.vastai_gpu_name = Some(next_arg(&mut args, "--vastai-gpu-name")?) - } - "--vastai-min-gpu-ram-mb" => { - self.vastai_min_gpu_ram_mb = - Some(parse_next(&mut args, "--vastai-min-gpu-ram-mb")?); - self.vastai_min_gpu_ram_mb_raw = None; - } - "--vastai-min-down-mbps" => { - self.vastai_min_down_mbps = - Some(parse_next(&mut args, "--vastai-min-down-mbps")?); - self.vastai_min_down_mbps_raw = None; - } - "--vastai-max-dph-total" => { - self.vastai_max_dph_total = - Some(parse_next(&mut args, "--vastai-max-dph-total")?); - self.vastai_max_dph_total_raw = None; - } - "--vastai-min-up-mbps" => { - self.vastai_min_up_mbps = Some(parse_next(&mut args, "--vastai-min-up-mbps")?); - self.vastai_min_up_mbps_raw = None; - } - "--vastai-min-reliability" => { - self.vastai_min_reliability = - Some(parse_next(&mut args, "--vastai-min-reliability")?); - self.vastai_min_reliability_raw = None; - } - "--vastai-require-verified" => { - self.vastai_require_verified = Some(true); - self.vastai_require_verified_raw = None; - } - "--no-vastai-require-verified" => { - self.vastai_require_verified = Some(false); - self.vastai_require_verified_raw = None; - } - "--vastai-blacklist-host" => { - let host_id = parse_next(&mut args, "--vastai-blacklist-host")?; - self.push_vastai_blacklist_host(host_id); - } - "--vastai-poll-interval-secs" => { - self.vastai_poll_interval_secs = - Some(parse_next(&mut args, "--vastai-poll-interval-secs")?); - self.vastai_poll_interval_secs_raw = None; - } - other => return Err(format!("unknown argument {other:?}")), + if self.apply_core_cli_arg(&arg, &mut args)? + || self.apply_model_cli_arg(&arg, &mut args)? + || self.apply_vastai_string_cli_arg(&arg, &mut args)? + || self.apply_vastai_numeric_cli_arg(&arg, &mut args)? + || self.apply_vastai_bool_cli_arg(&arg) + { + continue; } + return Err(format!("unknown argument {arg:?}")); } Ok(self) } @@ -1847,8 +1877,7 @@ impl Config { .as_mut() .expect("VastAI config exists when provider is vastai"); vastai.ssh_identity = Some(identity); - vastai.provisioning.ssh_public_key = Some(public_key.clone()); - vastai.ssh_public_key = Some(public_key); + vastai.provisioning.ssh_public_key = Some(public_key); vastai.ssh_public_fingerprint = Some(fingerprint); Ok(()) } @@ -1922,23 +1951,17 @@ impl Config { if local_tinygrad_worker_env(&self.provider).is_some() { keys.push("MVP_TINYGRAD_WORKER"); } - if std::env::var_os("MVP_CPU_LINE_PROFILE").is_some() { - keys.push("MVP_CPU_LINE_PROFILE"); - } - if std::env::var_os("MVP_CPU_LINE_PROFILE_INTERVAL_MS").is_some() { - keys.push("MVP_CPU_LINE_PROFILE_INTERVAL_MS"); - } - if std::env::var_os("MVP_TOKEN_PROGRESS_EVERY").is_some() { - keys.push("MVP_TOKEN_PROGRESS_EVERY"); - } - if std::env::var_os("CUDA_DEVICE_SCHEDULE").is_some() { - keys.push("CUDA_DEVICE_SCHEDULE"); - } - if std::env::var_os("MVP_MODEL_CACHE_DIR").is_some() { - keys.push("MVP_MODEL_CACHE_DIR"); - } - if std::env::var_os("HF_TOKEN").is_some() { - keys.push("HF_TOKEN"); + for key in [ + "MVP_CPU_LINE_PROFILE", + "MVP_CPU_LINE_PROFILE_INTERVAL_MS", + "MVP_TOKEN_PROGRESS_EVERY", + "CUDA_DEVICE_SCHEDULE", + "MVP_MODEL_CACHE_DIR", + "HF_TOKEN", + ] { + if std::env::var_os(key).is_some() { + keys.push(key); + } } match &self.gguf_source { GgufSource::LocalPath(_) => keys.push("MVP_GGUF_LOCAL_PATH"), @@ -1959,19 +1982,6 @@ impl Config { keys } - fn node_spec( - &self, - coordinator: EndpointAddr, - orchestrator_actor: ActorAddress, - ) -> Result { - self.node_spec_for_stage( - coordinator, - orchestrator_actor, - self.node_id, - self.stage_index, - ) - } - fn node_spec_for_stage( &self, coordinator: EndpointAddr, @@ -2159,24 +2169,42 @@ fn enqueue_datastream_subscribe( .map_err(|e| format!("send datastream subscribe: {e}")) } -#[allow(clippy::too_many_arguments)] -fn wait_for_runtime_ready_acks( - driver: &mut IrohDriver, - stack: &DistributionRuntimeStack, - obs_rx: &mpsc::Receiver, - frame_rx: &mpsc::Receiver, - frame_tx: &mpsc::Sender, - orchestrator_reports: &swactor::runtime::Inbox, - stop_rx: &mpsc::Receiver<()>, - dashboard: Option<&DashboardSupport>, - orch_datastream: &mut OrchDatastream, - orch_stdio_rx: Option<&mpsc::Receiver>, +struct RuntimeReadyAckLoop<'a> { + driver: &'a mut IrohDriver, + stack: &'a DistributionRuntimeStack, + obs_rx: &'a mpsc::Receiver, + frame_rx: &'a mpsc::Receiver, + frame_tx: &'a mpsc::Sender, + orchestrator_reports: &'a swactor::runtime::Inbox, + stop_rx: &'a mpsc::Receiver<()>, + dashboard: Option<&'a DashboardSupport>, + orch_datastream: &'a mut OrchDatastream, + orch_stdio_rx: Option<&'a mpsc::Receiver>, run_id: u64, orchestrator_node_id: u64, - provider: &ProviderKind, + provider: &'a ProviderKind, +} + +fn wait_for_runtime_ready_acks( + ctx: RuntimeReadyAckLoop<'_>, targets: &[RuntimeReadyAckTarget], collector_endpoint: &EndpointAddr, ) -> Result<(), String> { + let RuntimeReadyAckLoop { + driver, + stack, + obs_rx, + frame_rx, + frame_tx, + orchestrator_reports, + stop_rx, + dashboard, + orch_datastream, + orch_stdio_rx, + run_id, + orchestrator_node_id, + provider, + } = ctx; let mut pending = targets .iter() .cloned() @@ -2208,21 +2236,13 @@ fn wait_for_runtime_ready_acks( "shutdown requested while waiting for runtime-ready acknowledgements".to_owned(), ); } - while let Ok(observation) = obs_rx.try_recv() { - emit_plugin_observation(orch_datastream, dashboard, provider, &observation); - match observation { - PluginObservation::DatastreamFrame { .. } => {} - PluginObservation::ProviderLine { .. } - | PluginObservation::StdoutLine { .. } - | PluginObservation::StderrLine { .. } => {} - PluginObservation::Failed { reason, .. } => return Err(reason), - PluginObservation::Exited { - status, node_id, .. - } => { - return Err(format!("node {node_id} exited before ready: {status:?}")); - } - } - } + drain_observations_with_exit( + obs_rx, + dashboard, + orch_datastream, + provider, + |node_id, status| format!("node {node_id} exited before ready: {status:?}"), + )?; drain_frames(frame_rx, dashboard, orch_datastream); while let Some(report) = orchestrator_reports.try_recv() { let OrchestratorReport::NodeRuntimeReadyAck { @@ -2305,39 +2325,6 @@ fn wait_for_runtime_ready_acks( Ok(()) } -#[cfg(test)] -struct ProvisionedNodeGuard<'a> { - provisioner: &'a mut dyn ProvisionPlugin, - handle: Option, -} - -#[cfg(test)] -impl<'a> ProvisionedNodeGuard<'a> { - fn new( - provisioner: &'a mut dyn ProvisionPlugin, - handle: crate::provisioning::PluginNodeHandle, - ) -> Self { - Self { - provisioner, - handle: Some(handle), - } - } - - fn stop(&mut self) -> Result<(), String> { - let Some(handle) = self.handle.take() else { - return Ok(()); - }; - self.provisioner.stop_node(&handle) - } -} - -#[cfg(test)] -impl Drop for ProvisionedNodeGuard<'_> { - fn drop(&mut self) { - let _ = self.stop(); - } -} - struct ProvisionedClusterGuard { provisioner: Box, handles: Vec, @@ -2391,6 +2378,11 @@ impl Drop for ProvisionedClusterGuard { } } +type ProviderStartResults = Vec<( + NodeProvisionSpec, + Result, +)>; + fn start_nodes_with_stdio_capture( provisioner: Box, node_specs: Vec, @@ -2400,13 +2392,7 @@ fn start_nodes_with_stdio_capture( orch_datastream: &mut OrchDatastream, run_id: u64, node_id: u64, -) -> ( - Box, - Vec<( - NodeProvisionSpec, - Result, - )>, -) { +) -> (Box, ProviderStartResults) { let (tx, rx) = mpsc::channel(); thread::spawn(move || { let mut provisioner = provisioner; @@ -2447,26 +2433,29 @@ fn start_nodes_with_stdio_capture( } } -#[allow(clippy::too_many_arguments)] fn start_and_provision_workers( mut provisioner: Box, config: &Config, pipeline_plan: Option<&run_plan::RunPlan>, - driver: &mut IrohDriver, - stack: &DistributionRuntimeStack, - obs_rx: &mpsc::Receiver, - frame_rx: &mpsc::Receiver, - frame_tx: &mpsc::Sender, - orchestrator_reports: &swactor::runtime::Inbox, - stop_rx: &mpsc::Receiver<()>, - dashboard: Option<&DashboardSupport>, - orch_datastream: &mut OrchDatastream, - orch_stdio_rx: Option<&mpsc::Receiver>, + ctx: RuntimeReadyAckLoop<'_>, sink: PluginSink, coordinator: EndpointAddr, pipeline_coordinator: EndpointAddr, orchestrator_actor: ActorAddress, ) -> Result<(ProvisionedClusterGuard, PromptRuntimeReady), String> { + let RuntimeReadyAckLoop { + driver, + stack, + obs_rx, + frame_rx, + frame_tx, + orchestrator_reports, + stop_rx, + dashboard, + orch_datastream, + orch_stdio_rx, + .. + } = ctx; let stage_specs = stage_node_specs( config, pipeline_plan, @@ -2628,19 +2617,22 @@ fn start_and_provision_workers( ); let readies = if pipeline_plan.is_some() { match wait_for_runtime_readies( - driver, - stack, - obs_rx, - frame_rx, - frame_tx, - orchestrator_reports, - stop_rx, - dashboard, - orch_datastream, - orch_stdio_rx, - config.run_id, + RuntimeReadyAckLoop { + driver, + stack, + obs_rx, + frame_rx, + frame_tx, + orchestrator_reports, + stop_rx, + dashboard, + orch_datastream, + orch_stdio_rx, + run_id: config.run_id, + orchestrator_node_id: config.node_id, + provider: &config.provider, + }, &expected_node_ids, - &config.provider, ) { Ok(readies) => readies, Err(error) => { @@ -2656,7 +2648,7 @@ fn start_and_provision_workers( } } } else { - let ready = match wait_for_runtime_ready( + let ready = match wait_for_runtime_ready(RuntimeReadyAckLoop { driver, stack, obs_rx, @@ -2667,10 +2659,10 @@ fn start_and_provision_workers( dashboard, orch_datastream, orch_stdio_rx, - config.run_id, - config.node_id, - &config.provider, - ) { + run_id: config.run_id, + orchestrator_node_id: config.node_id, + provider: &config.provider, + }) { Ok(ready) => ready, Err(error) => { orch_datastream.emit_bootstrap( @@ -2704,19 +2696,21 @@ fn start_and_provision_workers( }) .collect::>(); wait_for_runtime_ready_acks( - driver, - stack, - obs_rx, - frame_rx, - frame_tx, - orchestrator_reports, - stop_rx, - dashboard, - orch_datastream, - orch_stdio_rx, - config.run_id, - config.node_id, - &config.provider, + RuntimeReadyAckLoop { + driver, + stack, + obs_rx, + frame_rx, + frame_tx, + orchestrator_reports, + stop_rx, + dashboard, + orch_datastream, + orch_stdio_rx, + run_id: config.run_id, + orchestrator_node_id: config.node_id, + provider: &config.provider, + }, &ack_targets, &pipeline_coordinator, )?; @@ -2749,21 +2743,22 @@ fn start_and_provision_workers( ); let weights_result = if pipeline_plan.is_some() { wait_for_weights_loaded_count( - driver, - stack, - obs_rx, - frame_rx, - frame_tx, - orchestrator_reports, - stop_rx, - dashboard, - orch_datastream, - orch_stdio_rx, - config.run_id, - config.node_id, - &config.provider, + RuntimeReadyAckLoop { + driver, + stack, + obs_rx, + frame_rx, + frame_tx, + orchestrator_reports, + stop_rx, + dashboard, + orch_datastream, + orch_stdio_rx, + run_id: config.run_id, + orchestrator_node_id: config.node_id, + provider: &config.provider, + }, expected_node_ids.len(), - config, pipeline_plan.expect("pipeline mode requires plan"), &readies, &pipeline_coordinator, @@ -2771,19 +2766,21 @@ fn start_and_provision_workers( ) } else { wait_for_weights_loaded( - driver, - stack, - obs_rx, - frame_rx, - frame_tx, - orchestrator_reports, - stop_rx, - dashboard, - orch_datastream, - orch_stdio_rx, - config.run_id, - config.node_id, - &config.provider, + RuntimeReadyAckLoop { + driver, + stack, + obs_rx, + frame_rx, + frame_tx, + orchestrator_reports, + stop_rx, + dashboard, + orch_datastream, + orch_stdio_rx, + run_id: config.run_id, + orchestrator_node_id: config.node_id, + provider: &config.provider, + }, config.stage_index, ) }; @@ -2873,25 +2870,22 @@ fn stage_node_specs( }) .collect() } else { - Ok(vec![config.node_spec(coordinator, orchestrator_actor)?]) + Ok(vec![config.node_spec_for_stage( + coordinator, + orchestrator_actor, + config.node_id, + config.stage_index, + )?]) } } struct ProviderStartOutcome { - results: Vec<( - NodeProvisionSpec, - Result, - )>, + results: ProviderStartResults, successful_handles: Vec, failed_specs: Vec, first_error: Option, } -fn collect_provider_start_outcome( - results: Vec<( - NodeProvisionSpec, - Result, - )>, -) -> ProviderStartOutcome { +fn collect_provider_start_outcome(results: ProviderStartResults) -> ProviderStartOutcome { let mut successful_handles = Vec::new(); let mut failed_specs = Vec::new(); let mut first_error = None; @@ -3133,22 +3127,25 @@ fn stage_ring_spec_wire(spec: run_plan::RingSpec) -> StageRingSpecWire { } } -#[allow(clippy::too_many_arguments)] fn wait_for_runtime_readies( - driver: &mut IrohDriver, - stack: &DistributionRuntimeStack, - obs_rx: &mpsc::Receiver, - frame_rx: &mpsc::Receiver, - frame_tx: &mpsc::Sender, - orchestrator_reports: &swactor::runtime::Inbox, - stop_rx: &mpsc::Receiver<()>, - dashboard: Option<&DashboardSupport>, - orch_datastream: &mut OrchDatastream, - orch_stdio_rx: Option<&mpsc::Receiver>, - run_id: u64, + ctx: RuntimeReadyAckLoop<'_>, expected_node_ids: &[u64], - provider: &ProviderKind, ) -> Result, String> { + let RuntimeReadyAckLoop { + driver, + stack, + obs_rx, + frame_rx, + frame_tx, + orchestrator_reports, + stop_rx, + dashboard, + orch_datastream, + orch_stdio_rx, + run_id, + provider, + .. + } = ctx; let expected = expected_node_ids.iter().copied().collect::>(); let mut pending = BTreeMap::::new(); loop { @@ -3224,28 +3221,29 @@ fn wait_for_runtime_readies( } } -#[allow(clippy::too_many_arguments)] fn wait_for_weights_loaded_count( - driver: &mut IrohDriver, - stack: &DistributionRuntimeStack, - obs_rx: &mpsc::Receiver, - frame_rx: &mpsc::Receiver, - frame_tx: &mpsc::Sender, - orchestrator_reports: &swactor::runtime::Inbox, - stop_rx: &mpsc::Receiver<()>, - dashboard: Option<&DashboardSupport>, - orch_datastream: &mut OrchDatastream, - orch_stdio_rx: Option<&mpsc::Receiver>, - run_id: u64, - node_id: u64, - provider: &ProviderKind, + ctx: RuntimeReadyAckLoop<'_>, expected_count: usize, - _config: &Config, pipeline_plan: &run_plan::RunPlan, readies: &BTreeMap, pipeline_coordinator: &EndpointAddr, stage_shard_plans: &BTreeMap, ) -> Result<(), String> { + let RuntimeReadyAckLoop { + driver, + stack, + obs_rx, + frame_rx, + frame_tx, + orchestrator_reports, + stop_rx, + dashboard, + orch_datastream, + orch_stdio_rx, + run_id, + orchestrator_node_id: node_id, + provider, + } = ctx; let expected_stages = pipeline_plan .stages .iter() @@ -3278,7 +3276,7 @@ fn wait_for_weights_loaded_count( )); } for stage in pending { - let _sent = send_pipeline_stage_provision( + send_pipeline_stage_provision( driver, stack, frame_tx, @@ -3417,7 +3415,7 @@ fn send_pipeline_stage_provision( stage_last_sends: &mut BTreeMap, load_progress: &BTreeMap, attempt: u64, -) -> Result { +) -> Result<(), String> { let stage_node_id = stage.node_id.0; let current_send_count = stage_resend_counts .get(&stage.stage_index) @@ -3489,13 +3487,13 @@ fn send_pipeline_stage_provision( ); return Err(reason); } - let dispatch = stage_provision_dispatch( + let (should_send, dispatch_reason) = 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() { + if !should_send { orch_datastream.emit_bootstrap( dashboard, run_id, @@ -3510,7 +3508,7 @@ fn send_pipeline_stage_provision( "stage_send_count":current_send_count, "loaded_stage_count":loaded_stages.len(), "resend_suppressed":true, - "resend_reason":dispatch.reason(), + "resend_reason":dispatch_reason, "liveness":stage_load_liveness_detail( load_progress.get(&stage_node_id), stage.stage_index, @@ -3529,7 +3527,7 @@ fn send_pipeline_stage_provision( ) }), ); - return Ok(false); + return Ok(()); } let stage_send_count = { let count = stage_resend_counts.entry(stage.stage_index).or_default(); @@ -3550,7 +3548,7 @@ fn send_pipeline_stage_provision( "stage_send_count":stage_send_count, "loaded_stage_count":loaded_stages.len(), "parallel_weight_acquisition":true, - "resend_reason":dispatch.reason(), + "resend_reason":dispatch_reason, }), ); if stage_send_count == 1 || stage_send_count % 15 == 0 { @@ -3596,7 +3594,7 @@ fn send_pipeline_stage_provision( stage_shard_plans, )?; pump(driver, stack, frame_tx); - Ok(true) + Ok(()) } struct FailedProvisionPlugin; @@ -3669,24 +3667,6 @@ impl StageLoadProgress { } } -#[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, @@ -3711,34 +3691,34 @@ fn stage_provision_dispatch( send_count: u64, last_send: Option, now: Instant, -) -> StageProvisionDispatch { +) -> (bool, &'static str) { if send_count == 0 { - return StageProvisionDispatch::Send("initial"); + return (true, "initial"); } let Some(progress) = progress else { - return StageProvisionDispatch::Send("no_progress_after_send"); + return (true, "no_progress_after_send"); }; if progress.failure_reason.is_some() || progress.phase.as_deref() == Some("failed") { - return StageProvisionDispatch::Suppress("worker_load_failed"); + return (false, "worker_load_failed"); } if progress.phase.as_deref() == Some("weights_loaded") { - return StageProvisionDispatch::Suppress("weights_loaded_report_pending"); + return (false, "weights_loaded_report_pending"); } if !stage_load_phase_is_active(progress.phase.as_deref()) { - return StageProvisionDispatch::Send("unknown_or_inactive_progress"); + return (true, "unknown_or_inactive_progress"); } let Some(last_progress) = progress.last_progress else { - return StageProvisionDispatch::Send("active_phase_without_progress_time"); + return (true, "active_phase_without_progress_time"); }; if now.duration_since(last_progress) < STAGE_PROVISION_ACTIVE_RESEND_AFTER { - return StageProvisionDispatch::Suppress("active_progress"); + return (false, "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"); - } + if let Some(last_send) = last_send + && now.duration_since(last_send) < STAGE_PROVISION_ACTIVE_RESEND_AFTER + { + return (false, "recent_stale_progress_resend"); } - StageProvisionDispatch::Send("stale_progress") + (true, "stale_progress") } fn stage_load_liveness_detail( @@ -4407,21 +4387,22 @@ fn handle_prompt_connection( Ok(()) } -fn wait_for_runtime_ready( - driver: &mut IrohDriver, - stack: &DistributionRuntimeStack, - obs_rx: &mpsc::Receiver, - frame_rx: &mpsc::Receiver, - frame_tx: &mpsc::Sender, - orchestrator_reports: &swactor::runtime::Inbox, - stop_rx: &mpsc::Receiver<()>, - dashboard: Option<&DashboardSupport>, - orch_datastream: &mut OrchDatastream, - orch_stdio_rx: Option<&mpsc::Receiver>, - run_id: u64, - node_id: u64, - provider: &ProviderKind, -) -> Result { +fn wait_for_runtime_ready(ctx: RuntimeReadyAckLoop<'_>) -> Result { + let RuntimeReadyAckLoop { + driver, + stack, + obs_rx, + frame_rx, + frame_tx, + orchestrator_reports, + stop_rx, + dashboard, + orch_datastream, + orch_stdio_rx, + run_id, + orchestrator_node_id: node_id, + provider, + } = ctx; let mut pending_ready: Option = None; let mut node_swim_started = false; let mut node_swim_ready = false; @@ -4433,19 +4414,9 @@ fn wait_for_runtime_ready( if stop_requested(stop_rx) { return Err("shutdown requested while waiting for node ready".to_owned()); } - while let Ok(observation) = obs_rx.try_recv() { - emit_plugin_observation(orch_datastream, dashboard, provider, &observation); - match observation { - PluginObservation::DatastreamFrame { .. } => {} - PluginObservation::ProviderLine { .. } - | PluginObservation::StdoutLine { .. } - | PluginObservation::StderrLine { .. } => {} - PluginObservation::Failed { reason, .. } => return Err(reason), - PluginObservation::Exited { status, .. } => { - return Err(format!("node exited before ready: {status:?}")); - } - } - } + drain_observations_with_exit(obs_rx, dashboard, orch_datastream, provider, |_, status| { + format!("node exited before ready: {status:?}") + })?; while let Some(report) = orchestrator_reports.try_recv() { if let OrchestratorReport::NodeRuntimeReady { run_id: report_run_id, @@ -4561,41 +4532,31 @@ fn provision_stage( .map_err(|e| format!("send stage provision: {e}")) } -fn wait_for_weights_loaded( - driver: &mut IrohDriver, - stack: &DistributionRuntimeStack, - obs_rx: &mpsc::Receiver, - frame_rx: &mpsc::Receiver, - frame_tx: &mpsc::Sender, - orchestrator_reports: &swactor::runtime::Inbox, - stop_rx: &mpsc::Receiver<()>, - dashboard: Option<&DashboardSupport>, - orch_datastream: &mut OrchDatastream, - orch_stdio_rx: Option<&mpsc::Receiver>, - run_id: u64, - node_id: u64, - provider: &ProviderKind, - stage_index: u32, -) -> Result<(), String> { +fn wait_for_weights_loaded(ctx: RuntimeReadyAckLoop<'_>, stage_index: u32) -> Result<(), String> { + let RuntimeReadyAckLoop { + driver, + stack, + obs_rx, + frame_rx, + frame_tx, + orchestrator_reports, + stop_rx, + dashboard, + orch_datastream, + orch_stdio_rx, + run_id, + orchestrator_node_id: node_id, + provider, + } = ctx; loop { pump(driver, stack, frame_tx); drain_orch_stdio_capture(orch_stdio_rx, orch_datastream, dashboard, run_id, node_id); if stop_requested(stop_rx) { return Err("shutdown requested while waiting for weights loaded".to_owned()); } - 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 { status, .. } => { - return Err(format!("node exited while loading weights: {status:?}")); - } - PluginObservation::DatastreamFrame { .. } => {} - PluginObservation::ProviderLine { .. } - | PluginObservation::StdoutLine { .. } - | PluginObservation::StderrLine { .. } => {} - } - } + drain_observations_with_exit(obs_rx, dashboard, orch_datastream, provider, |_, status| { + format!("node exited while loading weights: {status:?}") + })?; drain_frames(frame_rx, dashboard, orch_datastream); while let Some(report) = orchestrator_reports.try_recv() { match report { @@ -4634,23 +4595,7 @@ struct PipelineTokenRecord { eos: bool, } -enum PipelineSendHandle { - Driver(EdgeSendHandle), - #[cfg(test)] - Channel(tokio_mpsc::UnboundedSender>), -} - -impl PipelineSendHandle { - fn send(&self, bytes: Vec) -> Result<(), String> { - match self { - Self::Driver(handle) => handle.send(bytes), - #[cfg(test)] - Self::Channel(tx) => tx - .send(bytes) - .map_err(|_| "pipeline token-in sender stopped".to_owned()), - } - } -} +type PipelineSendHandle = EdgeSendHandle; struct PendingEncode { request_id: u64, } @@ -4710,9 +4655,8 @@ impl PipelinePromptRuntime { token_out_edge_id: token_out_edge.edge_id.0, token_spec: token_in_edge.object_spec, token_out_spec: token_out_edge.object_spec, - token_in_sender: PipelineSendHandle::Driver( - driver.spawn_edge_send_pump(first_stage_endpoint, token_in_edge.edge_id.0)?, - ), + token_in_sender: driver + .spawn_edge_send_pump(first_stage_endpoint, token_in_edge.edge_id.0)?, recv_rx, recv_tx, tokenizer_encode_actor, @@ -5181,17 +5125,6 @@ impl PipelinePromptRuntime { } } -#[cfg(test)] -fn encode_token_record( - spec: run_plan::ObjectSpec, - _edge_id: u64, - sequence: u64, - tokens: &[u32], - eos: bool, -) -> Result, String> { - encode_token_record_with_flags(spec, sequence, tokens, eos, false) -} - fn encode_token_record_with_flags( spec: run_plan::ObjectSpec, sequence: u64, @@ -5274,29 +5207,33 @@ fn prompt_runtime_mode(pipeline_plan: Option<&run_plan::RunPlan>) -> PromptRunti } fn serve_prompts( - driver: &mut IrohDriver, - stack: &DistributionRuntimeStack, - obs_rx: &mpsc::Receiver, - frame_rx: &mpsc::Receiver, - frame_tx: &mpsc::Sender, + ctx: RuntimeReadyAckLoop<'_>, work_rx: &mpsc::Receiver, prompt_events: &swactor::runtime::Inbox, - stop_rx: &mpsc::Receiver<()>, - dashboard: Option<&DashboardSupport>, - orch_datastream: &mut OrchDatastream, - orch_stdio_rx: Option<&mpsc::Receiver>, - run_id: u64, - node_id: u64, node_actor: ActorAddress, reply_to: ActorAddress, tokenizer_events: &swactor::runtime::Inbox, tokenizer_encode_actor: ActorAddress, tokenizer_decode_actor: ActorAddress, tokenizer_reply_to: ActorAddress, - provider: &ProviderKind, pipeline_plan: Option<&run_plan::RunPlan>, prompt_endpoint: EndpointAddr, ) -> Result<(), String> { + let RuntimeReadyAckLoop { + driver, + stack, + obs_rx, + frame_rx, + frame_tx, + stop_rx, + dashboard, + orch_datastream, + orch_stdio_rx, + run_id, + orchestrator_node_id: node_id, + provider, + .. + } = ctx; let mut pipeline_runtime = match prompt_runtime_mode(pipeline_plan) { PromptRuntimeMode::PipelineTokenEdges => Some(PipelinePromptRuntime::new( driver, @@ -5324,7 +5261,13 @@ fn serve_prompts( pipeline.drain_tokens(&stack.runtime, dashboard, orch_datastream, run_id, node_id)?; pipeline.emit_wait_progress(dashboard, orch_datastream, run_id, node_id); } - drain_observations(obs_rx, dashboard, orch_datastream, &provider)?; + drain_observations_with_exit( + obs_rx, + dashboard, + orch_datastream, + &provider, + |_, status| format!("node exited: {status:?}"), + )?; drain_frames(frame_rx, dashboard, orch_datastream); drain_orch_stdio_capture(orch_stdio_rx, orch_datastream, dashboard, run_id, node_id); if stop_rx.try_recv().is_ok() { @@ -5569,19 +5512,20 @@ fn spawn_stop_listener() -> mpsc::Receiver<()> { rx } -fn drain_observations( +fn drain_observations_with_exit( obs_rx: &mpsc::Receiver, dashboard: Option<&DashboardSupport>, orch_datastream: &mut OrchDatastream, provider: &ProviderKind, + exit_message: impl Fn(u64, Option) -> String, ) -> Result<(), String> { 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 { status, .. } => { - return Err(format!("node exited: {status:?}")); - } + PluginObservation::Exited { + node_id, status, .. + } => return Err(exit_message(node_id, status)), PluginObservation::DatastreamFrame { .. } => {} PluginObservation::ProviderLine { .. } | PluginObservation::StdoutLine { .. } @@ -6057,11 +6001,6 @@ fn command_output_failure_detail(output: &std::process::Output, secret: Option<& detail } -#[cfg(test)] -fn relay_mode_from_env() -> Result { - relay_runtime_config_from_env(1).map(|relay| relay.mode) -} - fn next_arg(args: &mut impl Iterator, name: &str) -> Result { args.next() .ok_or_else(|| format!("missing value after {name}")) diff --git a/crates/mvp-system/src/orchestration/engine_builder/mod.rs b/crates/mvp-system/src/orchestration/engine_builder/mod.rs index 51f1da3..249ff25 100644 --- a/crates/mvp-system/src/orchestration/engine_builder/mod.rs +++ b/crates/mvp-system/src/orchestration/engine_builder/mod.rs @@ -20,8 +20,7 @@ pub mod runtime_stack; pub mod workload; pub use crate::run_plan::NodeId; -pub use engine::{ClusterBuilder, ClusterHandle}; -pub use error::{EngineBuildError, PlanningError}; +pub use engine::ClusterBuilder; pub use events::EngineEvent; pub use launcher::StaticNodeLauncher; pub use model::{DTypeFamily, ModelArtifact, ModelSpec}; @@ -29,5 +28,3 @@ pub use node_image::{NodeImageSpec, WorkerRuntimeSpec}; pub use planner::FixedLinearPipelinePlanner; pub use pool::{NodeCapability, NodeLease, ResourceFacts, StaticPoolProvider}; pub use roles::RoleKind; -pub use runtime_stack::RuntimeNode; -pub use workload::WorkloadAdapter; diff --git a/crates/mvp-system/src/orchestration/run_fsm.rs b/crates/mvp-system/src/orchestration/run_fsm.rs index 319e020..5b54e27 100644 --- a/crates/mvp-system/src/orchestration/run_fsm.rs +++ b/crates/mvp-system/src/orchestration/run_fsm.rs @@ -230,14 +230,7 @@ impl OrchestratorRun { RunEvent::StageReady { run_id, stage_index, - } => { - if run_id != self.config.run_id || !self.plan_has_stage(stage_index) { - self.fault(RunFaultReason::UnknownStageReady { stage_index }); - return; - } - self.ready_stages.insert(stage_index); - self.maybe_inject_initial(); - } + } => self.stage_ready(run_id, stage_index), RunEvent::TokenInEndpointReady => { self.token_in_ready = true; self.maybe_inject_initial(); @@ -250,45 +243,31 @@ impl OrchestratorRun { sequence, token_id, eos, - } => { - if self.terminal { - return; - } - if sequence != self.expected_token_sequence { - return; - } - self.expected_token_sequence += 1; - if eos { - self.complete(); - } else if (self.injected_sequences.len() as u64) < self.config.max_tokens { - self.inject_decode(sequence + 1, token_id, sequence); - } else { - self.complete(); - } - } + } => self.token_received(sequence, token_id, eos), + RunEvent::StageFault { run_id, .. } + | RunEvent::EndpointFault { run_id, .. } + | RunEvent::OperatorStop { run_id } + | RunEvent::MembershipLost { run_id, .. } + | RunEvent::StageStopped { run_id, .. } + if run_id != self.config.run_id => {} RunEvent::StageFault { - run_id, stage_index, reason, - } if run_id == self.config.run_id => { + .. + } => { self.fault(RunFaultReason::StageFault { stage_index, reason, }); } - RunEvent::EndpointFault { run_id, endpoint } if run_id == self.config.run_id => { + RunEvent::EndpointFault { endpoint, .. } => { self.fault(RunFaultReason::EndpointFault { endpoint }); } - RunEvent::OperatorStop { run_id } if run_id == self.config.run_id => { - self.operator_stop(); - } - RunEvent::MembershipLost { run_id, node_id } if run_id == self.config.run_id => { + RunEvent::OperatorStop { .. } => self.operator_stop(), + RunEvent::MembershipLost { node_id, .. } => { self.fault(RunFaultReason::MembershipLost { node_id }); } - RunEvent::StageStopped { - run_id, - stage_index, - } if run_id == self.config.run_id => { + RunEvent::StageStopped { stage_index, .. } => { self.stopped_stages.insert(stage_index); self.maybe_torn_down(); } @@ -296,11 +275,29 @@ impl OrchestratorRun { self.token_endpoints_stopped = true; self.maybe_torn_down(); } - RunEvent::StageFault { .. } - | RunEvent::EndpointFault { .. } - | RunEvent::OperatorStop { .. } - | RunEvent::MembershipLost { .. } - | RunEvent::StageStopped { .. } => {} + } + } + + fn stage_ready(&mut self, run_id: RunId, stage_index: u32) { + if run_id != self.config.run_id || !self.plan_has_stage(stage_index) { + self.fault(RunFaultReason::UnknownStageReady { stage_index }); + return; + } + self.ready_stages.insert(stage_index); + self.maybe_inject_initial(); + } + + fn token_received(&mut self, sequence: u64, token_id: u32, eos: bool) { + if self.terminal || sequence != self.expected_token_sequence { + return; + } + self.expected_token_sequence += 1; + if eos { + self.complete(); + } else if (self.injected_sequences.len() as u64) < self.config.max_tokens { + self.inject_decode(sequence + 1, token_id, sequence); + } else { + self.complete(); } } diff --git a/crates/mvp-system/src/staging/control.rs b/crates/mvp-system/src/staging/control.rs index dfebb3a..0674179 100644 --- a/crates/mvp-system/src/staging/control.rs +++ b/crates/mvp-system/src/staging/control.rs @@ -302,22 +302,10 @@ impl StageController { StageEvent::WorkerReady => self.worker_ready = true, StageEvent::WeightsReady => self.weights_ready = true, StageEvent::InboundEdgeReady { edge_id } => { - if self - .provision - .as_ref() - .is_some_and(|p| p.inbound.edge_id == edge_id) - { - self.inbound_ready = true; - } + self.edge_ready(edge_id, EdgeDirection::Inbound) } StageEvent::OutboundEdgeReady { edge_id } => { - if self - .provision - .as_ref() - .is_some_and(|p| p.outbound.edge_id == edge_id) - { - self.outbound_ready = true; - } + self.edge_ready(edge_id, EdgeDirection::Outbound) } StageEvent::ObjectLoaded { edge_id, @@ -340,6 +328,21 @@ impl StageController { self.maybe_stage_ready(); } + fn edge_ready(&mut self, edge_id: EdgeId, direction: EdgeDirection) { + let Some(provision) = &self.provision else { + return; + }; + match direction { + EdgeDirection::Inbound if provision.inbound.edge_id == edge_id => { + self.inbound_ready = true + } + EdgeDirection::Outbound if provision.outbound.edge_id == edge_id => { + self.outbound_ready = true + } + _ => {} + } + } + pub fn commands(&self) -> &[StageCommand] { &self.commands } @@ -380,15 +383,17 @@ impl StageController { if self.stage_ready_emitted || self.faulted || self.stopped || self.stopping_run.is_some() { return; } - if self.worker_ready && self.weights_ready && self.inbound_ready && self.outbound_ready { - if let Some(provision) = &self.provision { - self.stage_ready_emitted = true; - self.events.push(StageLifecycleEvent::StageReady { - run_id: provision.run_id, - stage_index: provision.stage_index, - }); - } + if !(self.worker_ready && self.weights_ready && self.inbound_ready && self.outbound_ready) { + return; } + let Some(provision) = &self.provision else { + return; + }; + self.stage_ready_emitted = true; + self.events.push(StageLifecycleEvent::StageReady { + run_id: provision.run_id, + stage_index: provision.stage_index, + }); } fn object_loaded( diff --git a/crates/mvp-system/src/staging/gguf_metadata.rs b/crates/mvp-system/src/staging/gguf_metadata.rs index adc0b95..e26fc64 100644 --- a/crates/mvp-system/src/staging/gguf_metadata.rs +++ b/crates/mvp-system/src/staging/gguf_metadata.rs @@ -174,23 +174,26 @@ enum GgufValueType { impl GgufValueType { fn read(reader: &mut R) -> Result { + const VALUE_TYPES: [GgufValueType; 13] = [ + GgufValueType::Uint8, + GgufValueType::Int8, + GgufValueType::Uint16, + GgufValueType::Int16, + GgufValueType::Uint32, + GgufValueType::Int32, + GgufValueType::Float32, + GgufValueType::Bool, + GgufValueType::String, + GgufValueType::Array, + GgufValueType::Uint64, + GgufValueType::Int64, + GgufValueType::Float64, + ]; let raw = read_u32(reader)?; - match raw { - 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 metadata value type {other}")), - } + VALUE_TYPES + .get(raw as usize) + .copied() + .ok_or_else(|| format!("unsupported GGUF metadata value type {raw}")) } fn is_integer(self) -> bool { diff --git a/crates/mvp-system/src/staging/gguf_shard.rs b/crates/mvp-system/src/staging/gguf_shard.rs index 028b482..1eb534c 100644 --- a/crates/mvp-system/src/staging/gguf_shard.rs +++ b/crates/mvp-system/src/staging/gguf_shard.rs @@ -1,6 +1,6 @@ use std::fs::File; use std::io::{Read, Seek, SeekFrom, Write}; -use std::path::Path; +use std::path::{Path, PathBuf}; use serde::{Deserialize, Serialize}; @@ -537,54 +537,16 @@ where } 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 data_offsets = stage_shard_data_offsets(plan, alignment)?; - 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() - )); + let partial_path = partial_stage_shard_path(output_path); 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)?; + write_stage_shard_header(&mut out, plan, metadata_body, &data_offsets, alignment)?; let mut written_data = 0_u64; let mut tensor_index = 0_usize; for range in &plan.merged_tensor_ranges { @@ -712,6 +674,63 @@ where Ok(()) } +fn stage_shard_data_offsets(plan: &StageShardPlan, alignment: u64) -> Result, String> { + let mut offsets = Vec::with_capacity(plan.tensors.len()); + let mut cursor = 0_u64; + for tensor in &plan.tensors { + cursor = align_to(cursor, alignment)?; + offsets.push(cursor); + cursor = cursor + .checked_add(tensor.byte_len) + .ok_or_else(|| format!("stage shard data size overflow at tensor {}", tensor.name))?; + } + Ok(offsets) +} + +fn partial_stage_shard_path(output_path: &Path) -> PathBuf { + output_path.with_extension(format!( + "{}partial", + output_path + .extension() + .and_then(|value| value.to_str()) + .map(|ext| format!("{ext}.")) + .unwrap_or_default() + )) +} + +fn write_stage_shard_header( + out: &mut File, + plan: &StageShardPlan, + metadata_body: &[u8], + data_offsets: &[u64], + alignment: u64, +) -> Result<(), String> { + 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(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(out, alignment) +} + fn fetch_http_range(url: &str, start: u64, len: u64) -> Result, String> { if len == 0 { return Ok(Vec::new()); @@ -793,22 +812,26 @@ enum GgufValueType { 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}")), - } + const VALUE_TYPES: [GgufValueType; 13] = [ + GgufValueType::Uint8, + GgufValueType::Int8, + GgufValueType::Uint16, + GgufValueType::Int16, + GgufValueType::Uint32, + GgufValueType::Int32, + GgufValueType::Float32, + GgufValueType::Bool, + GgufValueType::String, + GgufValueType::Array, + GgufValueType::Uint64, + GgufValueType::Int64, + GgufValueType::Float64, + ]; + let raw = read_u32(reader)?; + VALUE_TYPES + .get(raw as usize) + .copied() + .ok_or_else(|| format!("unsupported GGUF value type {raw}")) } fn fixed_width(self) -> Option {