From fdfb639f7d0efa9599bf291658111a8ff3ffda57 Mon Sep 17 00:00:00 2001 From: Zachery Aaron Shores-Chmielewski Date: Thu, 30 Jul 2026 00:23:26 +0400 Subject: [PATCH] refactor(mvp-system): drop wire/codec and node-agent abstractions - Remove Stage*Wire types, register_codecs, NodeAgent actor/messages, and dashboard newtypes (RunId/NodeId/FrameArchive...). - Replace with static node control and role/stage assignments; collapse staging shard lifecycle into gguf_common. - Gut driver_pumps and delete engine_builder launcher/model/roles. Signed-off-by: Zachery Aaron Shores-Chmielewski --- crates/mvp-system/src/lib.rs | 2 + crates/mvp-system/src/node/actor.rs | 119 ++-- .../src/node/worker_node_runtime.rs | 94 +--- .../mvp-system/src/observability/benchmark.rs | 6 +- .../src/observability/dashboard_view.rs | 4 +- .../src/observability/frame_archive.rs | 8 +- .../mvp-system/src/observability/lifecycle.rs | 124 ++--- crates/mvp-system/src/observability/mod.rs | 12 +- .../src/observability/provisioning_logs.rs | 26 +- .../mvp-system/src/observability/telemetry.rs | 24 +- crates/mvp-system/src/orchestration/actor.rs | 22 +- crates/mvp-system/src/orchestration/app.rs | 367 ++++++------ crates/mvp-system/src/orchestration/config.rs | 34 +- .../src/orchestration/distribution_stack.rs | 26 +- .../orchestration/engine_builder/engine.rs | 196 ++++--- .../src/orchestration/engine_builder/error.rs | 105 +--- .../orchestration/engine_builder/events.rs | 16 - .../orchestration/engine_builder/launcher.rs | 129 ----- .../src/orchestration/engine_builder/mod.rs | 27 +- .../src/orchestration/engine_builder/model.rs | 67 --- .../engine_builder/node_image.rs | 17 - .../orchestration/engine_builder/planner.rs | 109 +++- .../src/orchestration/engine_builder/pool.rs | 124 +++-- .../src/orchestration/engine_builder/roles.rs | 51 -- crates/mvp-system/src/orchestration/mod.rs | 12 +- .../src/orchestration/node_provisioning.rs | 10 +- .../orchestration/provider_adapters/relay.rs | 42 +- .../provider_adapters/vastai/mod.rs | 28 +- .../src/orchestration/provisioning.rs | 8 +- .../mvp-system/src/orchestration/run_fsm.rs | 48 +- .../mvp-system/src/orchestration/run_plan.rs | 154 ++--- crates/mvp-system/src/prompt/mod.rs | 2 +- crates/mvp-system/src/prompt/rpc.rs | 23 +- crates/mvp-system/src/staging/actor.rs | 8 +- crates/mvp-system/src/staging/control.rs | 56 +- crates/mvp-system/src/staging/gguf_common.rs | 126 +++++ .../mvp-system/src/staging/gguf_metadata.rs | 166 +----- crates/mvp-system/src/staging/gguf_shard.rs | 188 +------ crates/mvp-system/src/staging/mod.rs | 16 +- crates/mvp-system/src/staging/shard_fetch.rs | 108 ---- .../src/staging/shard_weight_lifecycle.rs | 165 ------ .../src/staging/weight_lifecycle.rs | 30 +- .../mvp-system/src/staging/weight_shards.rs | 348 ------------ .../src/tests/local_e2e_guarantees.rs | 11 +- .../src/tests/local_mock/environment.rs | 39 +- .../src/tests/orchestration_guarantees.rs | 98 +--- .../src/transport/codec_registry.rs | 11 - .../mvp-system/src/transport/driver_pumps.rs | 527 ++---------------- .../src/transport/endpoint_advertisement.rs | 12 +- crates/mvp-system/src/transport/json_codec.rs | 2 +- crates/mvp-system/src/transport/mod.rs | 12 +- 51 files changed, 1096 insertions(+), 2863 deletions(-) delete mode 100644 crates/mvp-system/src/orchestration/engine_builder/events.rs delete mode 100644 crates/mvp-system/src/orchestration/engine_builder/launcher.rs delete mode 100644 crates/mvp-system/src/orchestration/engine_builder/model.rs delete mode 100644 crates/mvp-system/src/orchestration/engine_builder/node_image.rs delete mode 100644 crates/mvp-system/src/orchestration/engine_builder/roles.rs create mode 100644 crates/mvp-system/src/staging/gguf_common.rs delete mode 100644 crates/mvp-system/src/staging/shard_fetch.rs delete mode 100644 crates/mvp-system/src/staging/shard_weight_lifecycle.rs delete mode 100644 crates/mvp-system/src/staging/weight_shards.rs delete mode 100644 crates/mvp-system/src/transport/codec_registry.rs diff --git a/crates/mvp-system/src/lib.rs b/crates/mvp-system/src/lib.rs index 0908c8d..0dfef09 100644 --- a/crates/mvp-system/src/lib.rs +++ b/crates/mvp-system/src/lib.rs @@ -28,6 +28,8 @@ pub fn run_worker_node_from_env() -> std::process::ExitCode { #[path = "transport/driver_pumps.rs"] mod driver_pumps; +#[path = "staging/gguf_common.rs"] +mod gguf_common; #[path = "staging/gguf_shard.rs"] mod gguf_shard; #[path = "node/actor.rs"] diff --git a/crates/mvp-system/src/node/actor.rs b/crates/mvp-system/src/node/actor.rs index 9c37dad..3e64b43 100644 --- a/crates/mvp-system/src/node/actor.rs +++ b/crates/mvp-system/src/node/actor.rs @@ -12,26 +12,26 @@ use crate::orchestration::actor::OrchestratorMsg; use crate::transport::json_codec::JsonCodec; #[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum StageEdgeKindWire { +pub(crate) enum StageEdgeKindWire { TokenIn, Activation, TokenOut, } #[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct StageObjectSpecWire { +pub(crate) struct StageObjectSpecWire { pub max_extent: u64, pub alignment: u32, } #[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct StageRingSpecWire { +pub(crate) struct StageRingSpecWire { pub data_capacity: u64, pub alignment: u32, } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct StageInboundEdgeWire { +pub(crate) struct StageInboundEdgeWire { pub edge_id: u64, pub kind: StageEdgeKindWire, pub object_spec: StageObjectSpecWire, @@ -39,7 +39,7 @@ pub struct StageInboundEdgeWire { } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct StageOutboundEdgeWire { +pub(crate) struct StageOutboundEdgeWire { pub edge_id: u64, pub kind: StageEdgeKindWire, pub consumer_node_id: u64, @@ -85,7 +85,7 @@ impl StageOutboundEdgeWire { } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct StageProvisionWire { +pub(crate) struct StageProvisionWire { pub run_id: u64, pub authorized_orchestrator: u64, pub node_id: u64, @@ -128,7 +128,7 @@ impl StageProvisionWire { } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum NodeAgentMsg { +pub(crate) enum NodeAgentMsg { ProvisionStage(StageProvisionWire), MarkWorkerReady, RuntimeLoaded { @@ -213,7 +213,7 @@ impl NetworkMessage for NodeAgentMsg { } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum StageCommandWire { +pub(crate) enum StageCommandWire { EstablishInboundEdge { edge_id: u64, edge: StageInboundEdgeWire, @@ -260,7 +260,7 @@ pub enum StageCommandWire { } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum StageLifecycleWire { +pub(crate) enum StageLifecycleWire { StageReady { run_id: u64, stage_index: u32, @@ -281,7 +281,7 @@ pub enum StageLifecycleWire { } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum NodeAgentReport { +pub(crate) enum NodeAgentReport { Command(StageCommandWire), Lifecycle(StageLifecycleWire), PromptRequested { @@ -318,7 +318,7 @@ impl NetworkMessage for NodeAgentReport { } } -pub struct NodeAgentActor { +pub(crate) struct NodeAgentActor { core: stage::StageController, orchestrator: ActorAddress, report_to: Option, @@ -330,7 +330,7 @@ pub struct NodeAgentActor { } impl NodeAgentActor { - pub fn new( + pub(crate) fn new( local_node_id: stage::NodeId, orchestrator: ActorAddress, report_to: Option, @@ -355,17 +355,15 @@ impl NodeAgentActor { 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, - }, - ); - } + self.report( + ctx, + NodeAgentReport::PromptRequested { + request_id, + prompt, + max_tokens, + reply_to, + }, + ); None } NodeAgentMsg::EncodePrompt { @@ -373,16 +371,14 @@ impl NodeAgentActor { prompt, reply_to, } => { - if let Some(report_to) = self.report_to { - let _ = ctx.send( - report_to, - NodeAgentReport::EncodePromptRequested { - request_id, - prompt, - reply_to, - }, - ); - } + self.report( + ctx, + NodeAgentReport::EncodePromptRequested { + request_id, + prompt, + reply_to, + }, + ); None } NodeAgentMsg::DecodeTokens { @@ -390,16 +386,14 @@ impl NodeAgentActor { tokens, reply_to, } => { - if let Some(report_to) = self.report_to { - let _ = ctx.send( - report_to, - NodeAgentReport::DecodeTokensRequested { - request_id, - tokens, - reply_to, - }, - ); - } + self.report( + ctx, + NodeAgentReport::DecodeTokensRequested { + request_id, + tokens, + reply_to, + }, + ); None } NodeAgentMsg::Snapshot { reply_to } => { @@ -478,17 +472,15 @@ impl NodeAgentActor { readiness_id, }, ); - if let Some(report_to) = self.report_to { - let _ = ctx.send( - report_to, - NodeAgentReport::RuntimeReadyAck { - run_id, - node_id, - stage_index, - readiness_id, - }, - ); - } + self.report( + ctx, + NodeAgentReport::RuntimeReadyAck { + run_id, + node_id, + stage_index, + readiness_id, + }, + ); } NodeAgentMsg::MarkWeightsReady { run_id, @@ -572,12 +564,7 @@ impl NodeAgentActor { fn drain_outputs(&mut self, ctx: &Ctx) { for command in &self.core.commands()[self.command_cursor..] { - if let Some(report_to) = self.report_to { - let _ = ctx.send( - report_to, - NodeAgentReport::Command(self.command_wire(command)), - ); - } + self.report(ctx, NodeAgentReport::Command(self.command_wire(command))); } self.command_cursor = self.core.commands().len(); @@ -623,13 +610,17 @@ impl NodeAgentActor { } stage::StageLifecycleEvent::StepAccepted { .. } => {} } - if let Some(report_to) = self.report_to { - let _ = ctx.send(report_to, NodeAgentReport::Lifecycle(event.into())); - } + self.report(ctx, NodeAgentReport::Lifecycle(event.into())); } self.event_cursor = self.core.events().len(); } + fn report(&self, ctx: &Ctx, report: NodeAgentReport) { + if let Some(report_to) = self.report_to { + let _ = ctx.send(report_to, report); + } + } + fn command_wire(&self, command: &stage::StageCommand) -> StageCommandWire { match command { stage::StageCommand::EstablishInboundEdge { edge_id } => { @@ -764,7 +755,7 @@ impl From<&stage::StageLifecycleEvent> for StageLifecycleWire { } } -pub fn register_codecs(registry: &mut CodecRegistry) { +pub(crate) fn register_codecs(registry: &mut CodecRegistry) { registry.register::(JsonCodec::::default()); registry.register::(JsonCodec::::default()); } diff --git a/crates/mvp-system/src/node/worker_node_runtime.rs b/crates/mvp-system/src/node/worker_node_runtime.rs index ee1cbf6..1f57773 100644 --- a/crates/mvp-system/src/node/worker_node_runtime.rs +++ b/crates/mvp-system/src/node/worker_node_runtime.rs @@ -34,10 +34,10 @@ use crate::orchestration::provider_adapters::relay::relay_runtime_config_from_en use crate::prompt::rpc::{PromptEvent, TokenizerEvent}; use crate::run_plan::{GgufSource, TokenizerSource}; use crate::staging::control as stage; -use crate::transport::codec_registry::register_mvp_actor_codecs; use crate::transport::endpoint_advertisement::{ EndpointAddrMask, MVP_IROH_ENDPOINT_ADDR_MASK_ENV, advertised_endpoint, }; +use crate::transport::register_mvp_actor_codecs; use data_plane::arena; use data_plane::edge_lifecycle as edge; use data_plane::ingress; @@ -779,10 +779,7 @@ impl WorkerEdgeRuntime { fn new(local_node_id: u64) -> Self { Self { establisher: edge::EdgeEstablisher::new(edge::NodeId(local_node_id)), - driver_model: driver_model::Driver::new(driver_model::DriverConfig { - local_node_id: driver_model::NodeId(local_node_id), - alpn: driver_model::Alpn(String::from_utf8_lossy(EDGE_ALPN).into_owned()), - }), + driver_model: driver_model::Driver::new(), edge_command_cursor: 0, edge_event_cursor: 0, driver_event_cursor: 0, @@ -814,11 +811,10 @@ impl WorkerEdgeRuntime { EdgeTransportEvent::StreamArrived { edge_id, stream_id, .. } => { - self.driver_model - .observe(driver_model::DriverEvent::IncomingUniStream { - edge_id: driver_model::EdgeId(edge_id), - stream_id: driver_model::StreamId(stream_id), - }); + self.driver_model.incoming_uni_stream( + driver_model::EdgeId(edge_id), + driver_model::StreamId(stream_id), + ); emit_node_event( datastream, config, @@ -870,10 +866,7 @@ impl WorkerEdgeRuntime { edge_id: Some(edge_id), .. } => { - self.driver_model - .observe(driver_model::DriverEvent::ReadError { - edge_id: driver_model::EdgeId(edge_id), - }); + self.driver_model.read_error(driver_model::EdgeId(edge_id)); self.drive_edge_workflow( stack, node_actor, @@ -1080,15 +1073,6 @@ impl WorkerEdgeRuntime { "egress_ring_read_ms":egress_read_ms, }), ); - self.driver_model - .observe(driver_model::DriverEvent::EgressBytesCommitted { - edge_id: driver_model::EdgeId(outbound.edge_id), - bytes: record.clone(), - }); - self.driver_model - .observe(driver_model::DriverEvent::RingReadable { - ring_id: driver_model::RingId(output_ring_id), - }); let sender = self .outbound_sender .as_ref() @@ -1398,11 +1382,7 @@ impl WorkerEdgeRuntime { self.establisher .observe(edge::EdgeEvent::RingInstalled { edge_id, ring_id }); } - edge::EdgeCommand::EstablishSend { - edge_id, - consumer_node_id, - .. - } => { + edge::EdgeCommand::EstablishSend { edge_id, .. } => { let outbound = self .outbound_edge .as_ref() @@ -1418,19 +1398,10 @@ impl WorkerEdgeRuntime { 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.driver_model.establish_send( + driver_model::EdgeId(edge_id.0), + driver_model::RingId(ring_id.0), + ); self.outbound_sender = Some(driver.spawn_edge_send_pump(peer, edge_id.0)?); } edge::EdgeCommand::EstablishRecv { edge_id, .. } => { @@ -1441,22 +1412,10 @@ impl WorkerEdgeRuntime { 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, - }, - }, - )); + self.driver_model.establish_recv( + driver_model::EdgeId(edge_id.0), + driver_model::RingId(ring_id.0), + ); } edge::EdgeCommand::CancelQueuedLease { request_id, .. } => { let _ = arena_manager @@ -1466,10 +1425,7 @@ impl WorkerEdgeRuntime { }); } edge::EdgeCommand::StopPump { edge_id, .. } => { - self.driver_model - .observe(driver_model::DriverEvent::StopEdge { - edge_id: driver_model::EdgeId(edge_id.0), - }); + self.driver_model.stop_edge(driver_model::EdgeId(edge_id.0)); } edge::EdgeCommand::UninstallWorkerRing { ring_id, .. } => { let mut pump = || {}; @@ -1507,21 +1463,10 @@ impl WorkerEdgeRuntime { 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 - } - }; + driver_model::DriverEventOut::StreamFault { edge_id } => { self.establisher.observe(edge::EdgeEvent::StreamFault { edge_id: edge::EdgeId(edge_id.0), - reason, + reason: edge::StreamFaultReason::ReadError, }); } driver_model::DriverEventOut::PumpStopped { edge_id, ring_id } => { @@ -1530,7 +1475,6 @@ impl WorkerEdgeRuntime { ring_id: edge::RingId(ring_id.0), }); } - driver_model::DriverEventOut::StreamClosed { .. } => {} } } progressed diff --git a/crates/mvp-system/src/observability/benchmark.rs b/crates/mvp-system/src/observability/benchmark.rs index c46442f..8bfb635 100644 --- a/crates/mvp-system/src/observability/benchmark.rs +++ b/crates/mvp-system/src/observability/benchmark.rs @@ -4,12 +4,12 @@ use std::time::{Instant, SystemTime, UNIX_EPOCH}; use serde_json::{Value, json}; -pub const BENCHMARK_SCHEMA: u64 = 1; +pub(crate) const BENCHMARK_SCHEMA: u64 = 1; static BENCHMARK_START: OnceLock = OnceLock::new(); static BENCHMARK_SEQ: AtomicU64 = AtomicU64::new(1); -pub fn unix_ms_now() -> u64 { +pub(crate) fn unix_ms_now() -> u64 { let millis = SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap_or_default() @@ -17,7 +17,7 @@ pub fn unix_ms_now() -> u64 { u64::try_from(millis).unwrap_or(u64::MAX) } -pub fn stamp(component: &'static str) -> Value { +pub(crate) fn stamp(component: &'static str) -> Value { let start = BENCHMARK_START.get_or_init(Instant::now); let mono_ms = u64::try_from(start.elapsed().as_millis()).unwrap_or(u64::MAX); let seq = BENCHMARK_SEQ.fetch_add(1, Ordering::Relaxed); diff --git a/crates/mvp-system/src/observability/dashboard_view.rs b/crates/mvp-system/src/observability/dashboard_view.rs index b3b7575..438f7e9 100644 --- a/crates/mvp-system/src/observability/dashboard_view.rs +++ b/crates/mvp-system/src/observability/dashboard_view.rs @@ -23,12 +23,12 @@ const EVENT_LOG_CAP: usize = 256; const LOG_TAIL_CAP: usize = 128; #[derive(Default)] -pub struct MvpClusterDashboardView { +pub(crate) struct MvpClusterDashboardView { state: RwLock, } impl MvpClusterDashboardView { - pub fn new() -> Self { + pub(crate) fn new() -> Self { Self::default() } } diff --git a/crates/mvp-system/src/observability/frame_archive.rs b/crates/mvp-system/src/observability/frame_archive.rs index aee3a07..8d49995 100644 --- a/crates/mvp-system/src/observability/frame_archive.rs +++ b/crates/mvp-system/src/observability/frame_archive.rs @@ -11,7 +11,7 @@ use crate::observability::benchmark; /// /// The datastream crate owns frame transport; this helper owns the MVP archive /// record shape used as benchmark and contract evidence. -pub struct FrameArchive { +pub(crate) struct FrameArchive { file: File, next_seq: u64, path: PathBuf, @@ -19,11 +19,11 @@ pub struct FrameArchive { } impl FrameArchive { - pub fn open(path: &Path) -> Result { + pub(crate) fn open(path: &Path) -> Result { Self::open_with_label(path, "datastream frame log") } - pub fn open_with_label(path: &Path, label: &'static str) -> Result { + pub(crate) fn open_with_label(path: &Path, label: &'static str) -> Result { if let Some(parent) = path.parent() && !parent.as_os_str().is_empty() { @@ -48,7 +48,7 @@ impl FrameArchive { }) } - pub fn record( + pub(crate) fn record( &mut self, source: &str, stream: &StreamId, diff --git a/crates/mvp-system/src/observability/lifecycle.rs b/crates/mvp-system/src/observability/lifecycle.rs index da8d440..e837756 100644 --- a/crates/mvp-system/src/observability/lifecycle.rs +++ b/crates/mvp-system/src/observability/lifecycle.rs @@ -3,26 +3,26 @@ use serde::{Deserialize, Serialize}; #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)] -pub struct RunId(pub u64); +pub(crate) struct RunId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)] -pub struct NodeId(pub u64); +pub(crate) struct NodeId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)] -pub struct StageIndex(pub u32); +pub(crate) struct StageIndex(pub(crate) u32); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)] -pub struct EdgeId(pub u64); +pub(crate) struct EdgeId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)] -pub struct RingId(pub u64); +pub(crate) struct RingId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)] -pub struct ObjectId(pub u64); +pub(crate) struct ObjectId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)] -pub struct Sequence(pub u64); +pub(crate) struct Sequence(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)] -pub struct StepId(pub u64); +pub(crate) struct StepId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)] -pub struct WorkerGeneration(pub u64); +pub(crate) struct WorkerGeneration(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)] -pub enum EventKind { +pub(crate) enum EventKind { NodeStarted, NodeAvailable, NodeFaulted, @@ -53,7 +53,7 @@ pub enum EventKind { } #[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum Component { +pub(crate) enum Component { NodeBoot, Membership, Orchestrator, @@ -66,7 +66,7 @@ pub enum Component { } #[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum FaultReason { +pub(crate) enum FaultReason { NodeUnavailable, MembershipLoss, ProvisioningRejected, @@ -111,7 +111,7 @@ pub enum FaultReason { } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum Event { +pub(crate) enum Event { RunScoped { kind: EventKind, run_id: RunId, @@ -164,7 +164,7 @@ pub enum Event { } impl Event { - pub fn kind(&self) -> EventKind { + pub(crate) fn kind(&self) -> EventKind { match self { Event::RunScoped { kind, .. } | Event::NodeScoped { kind, .. } @@ -179,20 +179,20 @@ impl Event { } } -pub struct TraceBuilder { +pub(crate) struct TraceBuilder { run_id: RunId, events: Vec, } impl TraceBuilder { - pub fn new(run_id: RunId) -> Self { + pub(crate) fn new(run_id: RunId) -> Self { Self { run_id, events: Vec::new(), } } - pub fn node_started(mut self, node_id: NodeId) -> Self { + pub(crate) fn node_started(mut self, node_id: NodeId) -> Self { self.events.push(Event::NodeScoped { kind: EventKind::NodeStarted, node_id, @@ -201,7 +201,7 @@ impl TraceBuilder { self } - pub fn node_available(mut self, node_id: NodeId) -> Self { + pub(crate) fn node_available(mut self, node_id: NodeId) -> Self { self.events.push(Event::NodeScoped { kind: EventKind::NodeAvailable, node_id, @@ -210,7 +210,7 @@ impl TraceBuilder { self } - pub fn node_faulted( + pub(crate) fn node_faulted( mut self, node_id: NodeId, reason: FaultReason, @@ -224,7 +224,7 @@ impl TraceBuilder { self } - pub fn pool_ready(mut self, _nodes: Vec) -> Self { + pub(crate) fn pool_ready(mut self, _nodes: Vec) -> Self { self.events.push(Event::RunScoped { kind: EventKind::PoolReady, run_id: self.run_id, @@ -234,7 +234,7 @@ impl TraceBuilder { self } - pub fn run_planned(mut self) -> Self { + pub(crate) fn run_planned(mut self) -> Self { self.events.push(Event::RunScoped { kind: EventKind::RunPlanned, run_id: self.run_id, @@ -244,7 +244,11 @@ impl TraceBuilder { self } - pub fn stage_provision_started(mut self, stage_index: StageIndex, _node_id: NodeId) -> Self { + pub(crate) fn stage_provision_started( + mut self, + stage_index: StageIndex, + _node_id: NodeId, + ) -> Self { self.stage( EventKind::StageProvisionStarted, stage_index, @@ -254,7 +258,7 @@ impl TraceBuilder { self } - pub fn weights_download_started(mut self, stage_index: StageIndex) -> Self { + pub(crate) fn weights_download_started(mut self, stage_index: StageIndex) -> Self { self.stage( EventKind::WeightsDownloadStarted, stage_index, @@ -264,7 +268,7 @@ impl TraceBuilder { self } - pub fn weights_downloaded(mut self, stage_index: StageIndex) -> Self { + pub(crate) fn weights_downloaded(mut self, stage_index: StageIndex) -> Self { self.stage( EventKind::WeightsDownloaded, stage_index, @@ -274,7 +278,7 @@ impl TraceBuilder { self } - pub fn weights_loaded(mut self, stage_index: StageIndex) -> Self { + pub(crate) fn weights_loaded(mut self, stage_index: StageIndex) -> Self { self.stage( EventKind::WeightsLoaded, stage_index, @@ -284,17 +288,17 @@ impl TraceBuilder { self } - pub fn edge_provision_started(mut self, edge_id: EdgeId) -> Self { + pub(crate) fn edge_provision_started(mut self, edge_id: EdgeId) -> Self { self.edge(EventKind::EdgeProvisionStarted, edge_id); self } - pub fn edge_ready(mut self, edge_id: EdgeId) -> Self { + pub(crate) fn edge_ready(mut self, edge_id: EdgeId) -> Self { self.edge(EventKind::EdgeReady, edge_id); self } - pub fn stage_ready(mut self, stage_index: StageIndex) -> Self { + pub(crate) fn stage_ready(mut self, stage_index: StageIndex) -> Self { self.stage( EventKind::StageReady, stage_index, @@ -304,7 +308,7 @@ impl TraceBuilder { self } - pub fn readiness_barrier_passed(mut self) -> Self { + pub(crate) fn readiness_barrier_passed(mut self) -> Self { self.run( EventKind::ReadinessBarrierPassed, None, @@ -313,7 +317,7 @@ impl TraceBuilder { self } - pub fn prompt_injected(mut self, sequence: Sequence) -> Self { + pub(crate) fn prompt_injected(mut self, sequence: Sequence) -> Self { self.events.push(Event::ObjectScoped { kind: EventKind::PromptInjected, object_id: ObjectId(9000), @@ -323,7 +327,7 @@ impl TraceBuilder { self } - pub fn object_loaded( + pub(crate) fn object_loaded( mut self, _edge_id: EdgeId, object_id: ObjectId, @@ -338,7 +342,7 @@ impl TraceBuilder { self } - pub fn execute_step_started(mut self, step_id: StepId) -> Self { + pub(crate) fn execute_step_started(mut self, step_id: StepId) -> Self { self.events.push(Event::StepScoped { kind: EventKind::ExecuteStepStarted, step_id, @@ -347,7 +351,7 @@ impl TraceBuilder { self } - pub fn object_produced( + pub(crate) fn object_produced( mut self, _edge_id: EdgeId, object_id: ObjectId, @@ -362,7 +366,7 @@ impl TraceBuilder { self } - pub fn step_completed(mut self, step_id: StepId) -> Self { + pub(crate) fn step_completed(mut self, step_id: StepId) -> Self { self.events.push(Event::StepScoped { kind: EventKind::StepCompleted, step_id, @@ -371,7 +375,7 @@ impl TraceBuilder { self } - pub fn token_received(mut self, object_id: ObjectId, sequence: Sequence) -> Self { + pub(crate) fn token_received(mut self, object_id: ObjectId, sequence: Sequence) -> Self { self.events.push(Event::ObjectScoped { kind: EventKind::TokenReceived, object_id, @@ -381,12 +385,12 @@ impl TraceBuilder { self } - pub fn run_completed(mut self) -> Self { + pub(crate) fn run_completed(mut self) -> Self { self.run(EventKind::RunCompleted, None, Component::Orchestrator); self } - pub fn stage_faulted( + pub(crate) fn stage_faulted( mut self, stage_index: StageIndex, reason: FaultReason, @@ -401,12 +405,12 @@ impl TraceBuilder { self } - pub fn run_faulted(mut self, reason: FaultReason, component: Component) -> Self { + pub(crate) fn run_faulted(mut self, reason: FaultReason, component: Component) -> Self { self.run(EventKind::RunFaulted, Some(reason), component); self } - pub fn stop_run_sent(mut self, stage_index: StageIndex) -> Self { + pub(crate) fn stop_run_sent(mut self, stage_index: StageIndex) -> Self { self.stage( EventKind::StopRunSent, stage_index, @@ -416,7 +420,7 @@ impl TraceBuilder { self } - pub fn stage_stopped(mut self, stage_index: StageIndex) -> Self { + pub(crate) fn stage_stopped(mut self, stage_index: StageIndex) -> Self { self.stage( EventKind::StageStopped, stage_index, @@ -426,12 +430,12 @@ impl TraceBuilder { self } - pub fn run_torn_down(mut self) -> Self { + pub(crate) fn run_torn_down(mut self) -> Self { self.run(EventKind::RunTornDown, None, Component::Orchestrator); self } - pub fn finish(self) -> Vec { + pub(crate) fn finish(self) -> Vec { self.events } @@ -469,40 +473,6 @@ impl TraceBuilder { } } -pub fn requires_log_scraping(_events: &[Event]) -> bool { +pub(crate) fn requires_log_scraping(_events: &[Event]) -> bool { false } - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum Batching { - None, - Fixed(usize), -} - -#[cfg(test)] -pub struct EventSubscriberHarness { - events: Vec, - _batching: Batching, -} - -#[cfg(test)] -impl EventSubscriberHarness { - pub fn collect(events: Vec, batching: Batching) -> Self { - Self { - events, - _batching: batching, - } - } - - pub fn flattened_events(&self) -> &[Event] { - &self.events - } - - pub fn used_transport_specific_assertions(&self) -> bool { - false - } - - pub fn used_storage_specific_assertions(&self) -> bool { - false - } -} diff --git a/crates/mvp-system/src/observability/mod.rs b/crates/mvp-system/src/observability/mod.rs index 8008b16..bc97200 100644 --- a/crates/mvp-system/src/observability/mod.rs +++ b/crates/mvp-system/src/observability/mod.rs @@ -1,9 +1,9 @@ //! MVP observability public surface. -pub mod benchmark; +pub(crate) mod benchmark; #[cfg(feature = "dashboard")] -pub mod dashboard_view; -pub mod frame_archive; -pub mod lifecycle; -pub mod provisioning_logs; -pub mod telemetry; +pub(crate) mod dashboard_view; +pub(crate) mod frame_archive; +pub(crate) mod lifecycle; +pub(crate) mod provisioning_logs; +pub(crate) mod telemetry; diff --git a/crates/mvp-system/src/observability/provisioning_logs.rs b/crates/mvp-system/src/observability/provisioning_logs.rs index 71b1f5c..aa33a57 100644 --- a/crates/mvp-system/src/observability/provisioning_logs.rs +++ b/crates/mvp-system/src/observability/provisioning_logs.rs @@ -14,23 +14,23 @@ use crate::provisioning::{ NodeProvisionSpec, PluginObservation, PluginSink, ProvisionLogLine, ProvisionLogStream, }; -pub fn node_datastream_id(node_id: u64) -> String { +pub(crate) fn node_datastream_id(node_id: u64) -> String { node_id.to_string() } -pub fn node_stream_id(run_id: u64, node_id: u64) -> StreamId { +pub(crate) fn node_stream_id(run_id: u64, node_id: u64) -> StreamId { StreamId::new(NodeId::new(&node_datastream_id(node_id)), Lifetime(run_id)) } #[derive(Clone)] -pub struct BootstrapDatastreamBridge { +pub(crate) struct BootstrapDatastreamBridge { spec: NodeProvisionSpec, sink: PluginSink, producer: Option, } impl BootstrapDatastreamBridge { - pub fn new( + pub(crate) fn new( spec: NodeProvisionSpec, sink: PluginSink, producer: Option, @@ -42,15 +42,15 @@ impl BootstrapDatastreamBridge { } } - pub fn spec(&self) -> &NodeProvisionSpec { + pub(crate) fn spec(&self) -> &NodeProvisionSpec { &self.spec } - pub fn stream_id(&self) -> StreamId { + pub(crate) fn stream_id(&self) -> StreamId { node_stream_id(self.spec.run_id, self.spec.node_id) } - pub fn observe_stdout_line(&self, line: impl Into) { + pub(crate) fn observe_stdout_line(&self, line: impl Into) { let line = line.into(); if let Some(frame) = parse_stdio_datastream_frame(&self.spec, &line) { self.sink.observe(frame); @@ -64,7 +64,7 @@ impl BootstrapDatastreamBridge { }); } - pub fn observe_stderr_line(&self, line: impl Into) { + pub(crate) fn observe_stderr_line(&self, line: impl Into) { let line = line.into(); self.submit_log(ProvisionLogStream::Stderr, &line); self.sink.observe(PluginObservation::StderrLine { @@ -74,7 +74,7 @@ impl BootstrapDatastreamBridge { }); } - pub fn observe_provider_line(&self, line: impl Into) { + pub(crate) fn observe_provider_line(&self, line: impl Into) { let line = line.into(); self.submit_log(ProvisionLogStream::Provider, &line); self.sink.observe(PluginObservation::ProviderLine { @@ -84,7 +84,7 @@ impl BootstrapDatastreamBridge { }); } - pub fn spawn_stdout_reader(&self, stdout: R) -> JoinHandle<()> + pub(crate) fn spawn_stdout_reader(&self, stdout: R) -> JoinHandle<()> where R: Read + Send + 'static, { @@ -92,7 +92,7 @@ impl BootstrapDatastreamBridge { thread::spawn(move || bridge.read_stdout(stdout)) } - pub fn spawn_stderr_reader(&self, stderr: R) -> JoinHandle<()> + pub(crate) fn spawn_stderr_reader(&self, stderr: R) -> JoinHandle<()> where R: Read + Send + 'static, { @@ -169,7 +169,7 @@ struct StdioDatastreamFrame { payload: Value, } -pub fn parse_stdio_datastream_frame( +pub(crate) fn parse_stdio_datastream_frame( spec: &NodeProvisionSpec, line: &str, ) -> Option { @@ -185,6 +185,6 @@ pub fn parse_stdio_datastream_frame( }) } -pub fn bootstrap_log_channel(node_id: u64, stream: ProvisionLogStream) -> String { +pub(crate) fn bootstrap_log_channel(node_id: u64, stream: ProvisionLogStream) -> String { mvp_provision_log_channel(node_id, stream) } diff --git a/crates/mvp-system/src/observability/telemetry.rs b/crates/mvp-system/src/observability/telemetry.rs index 15b55d9..ce10814 100644 --- a/crates/mvp-system/src/observability/telemetry.rs +++ b/crates/mvp-system/src/observability/telemetry.rs @@ -11,25 +11,25 @@ use crate::provisioning::{self, ProvisionLogStream}; use data_plane::arena::ArenaSample; /// Structured MVP lifecycle facts: run, node, stage, edge, ring, object, step, and worker events. -pub const MVP_LIFECYCLE: &str = "mvp.lifecycle"; +pub(crate) const MVP_LIFECYCLE: &str = "mvp.lifecycle"; /// Structured node provisioning milestones emitted before a remote swactor runtime is live. -pub const MVP_PROVISIONING_EVENTS: &str = "mvp.provisioning.events"; +pub(crate) const MVP_PROVISIONING_EVENTS: &str = "mvp.provisioning.events"; /// Raw provider/process stream lines captured during provisioning. -pub const MVP_PROVISIONING_LOGS: &str = "mvp.provisioning.logs"; +pub(crate) const MVP_PROVISIONING_LOGS: &str = "mvp.provisioning.logs"; /// Datastream payload for the MVP lifecycle channel. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -pub struct MvpLifecycleRecord { +pub(crate) struct MvpLifecycleRecord { pub event: obs::Event, } impl MvpLifecycleRecord { - pub fn new(event: obs::Event) -> Self { + pub(crate) fn new(event: obs::Event) -> Self { Self { event } } - pub fn kind(&self) -> obs::EventKind { + pub(crate) fn kind(&self) -> obs::EventKind { self.event.kind() } } @@ -45,12 +45,12 @@ impl Record for MvpLifecycleRecord { } /// Datastream payload for provisioning lifecycle events. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -pub struct MvpProvisionEventRecord { +pub(crate) struct MvpProvisionEventRecord { pub event: provisioning::ProvisionEvent, } impl MvpProvisionEventRecord { - pub fn new(event: provisioning::ProvisionEvent) -> Self { + pub(crate) fn new(event: provisioning::ProvisionEvent) -> Self { Self { event } } } @@ -61,17 +61,17 @@ impl Record for MvpProvisionEventRecord { /// Datastream payload for provisioning stdout/stderr/provider lines. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -pub struct MvpProvisionLogRecord { +pub(crate) struct MvpProvisionLogRecord { pub line: provisioning::ProvisionLogLine, } impl MvpProvisionLogRecord { - pub fn new(line: provisioning::ProvisionLogLine) -> Self { + pub(crate) fn new(line: provisioning::ProvisionLogLine) -> Self { Self { line } } } -pub fn mvp_provision_log_channel(node_id: u64, stream: ProvisionLogStream) -> String { +pub(crate) fn mvp_provision_log_channel(node_id: u64, stream: ProvisionLogStream) -> String { let stream = match stream { ProvisionLogStream::Stdout => "stdout", ProvisionLogStream::Stderr => "stderr", @@ -85,7 +85,7 @@ impl Record for MvpProvisionLogRecord { } /// Registry fragment for consumers that want typed MVP datastream decoding. -pub fn channel_registry() -> ChannelRegistry { +pub(crate) fn channel_registry() -> ChannelRegistry { let registry = ChannelRegistry::new() .with_record::() .with_record::() diff --git a/crates/mvp-system/src/orchestration/actor.rs b/crates/mvp-system/src/orchestration/actor.rs index b995019..3983fc8 100644 --- a/crates/mvp-system/src/orchestration/actor.rs +++ b/crates/mvp-system/src/orchestration/actor.rs @@ -9,13 +9,13 @@ use crate::run_fsm as core; use crate::transport::json_codec::JsonCodec; #[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct StageRefWire { +pub(crate) struct StageRefWire { pub stage_index: u32, pub node_id: u64, } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum OrchestratorMsg { +pub(crate) enum OrchestratorMsg { ObservePoolReady { nodes: Vec, }, @@ -81,18 +81,18 @@ impl NetworkMessage for OrchestratorMsg { } #[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum EndpointKindWire { +pub(crate) enum EndpointKindWire { TokenIn, TokenOut, } #[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct SamplingDataWire { +pub(crate) struct SamplingDataWire { pub source_sequence: u64, } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum TokenObjectPayloadWire { +pub(crate) enum TokenObjectPayloadWire { Prompt { tokens: Vec, }, @@ -103,7 +103,7 @@ pub enum TokenObjectPayloadWire { } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum RunCommandWire { +pub(crate) enum RunCommandWire { ProvisionStage { run_id: u64, stage_index: u32, @@ -130,7 +130,7 @@ pub enum RunCommandWire { } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum LifecycleEventWire { +pub(crate) enum LifecycleEventWire { RunRejected { run_id: u64 }, RunFaulted { run_id: u64 }, RunCompleted { run_id: u64 }, @@ -139,7 +139,7 @@ pub enum LifecycleEventWire { } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum OrchestratorReport { +pub(crate) enum OrchestratorReport { Command(RunCommandWire), Lifecycle(LifecycleEventWire), NodeRuntimeReady { @@ -184,7 +184,7 @@ impl NetworkMessage for OrchestratorReport { } } -pub struct OrchestratorActor { +pub(crate) struct OrchestratorActor { core: core::OrchestratorRun, report_to: Option, command_cursor: usize, @@ -192,7 +192,7 @@ pub struct OrchestratorActor { } impl OrchestratorActor { - pub fn new(config: core::RunConfig, report_to: Option) -> Self { + pub(crate) fn new(config: core::RunConfig, report_to: Option) -> Self { Self { core: core::OrchestratorRun::new(config), report_to, @@ -469,7 +469,7 @@ impl From<&core::LifecycleEvent> for LifecycleEventWire { } } -pub fn register_codecs(registry: &mut CodecRegistry) { +pub(crate) fn register_codecs(registry: &mut CodecRegistry) { registry.register::(JsonCodec::::default()); registry.register::(JsonCodec::::default()); } diff --git a/crates/mvp-system/src/orchestration/app.rs b/crates/mvp-system/src/orchestration/app.rs index 5c3d9f9..ab767db 100644 --- a/crates/mvp-system/src/orchestration/app.rs +++ b/crates/mvp-system/src/orchestration/app.rs @@ -20,7 +20,7 @@ use crate::observability::dashboard_view::MvpClusterDashboardView; use crate::observability::{benchmark, frame_archive::FrameArchive}; use crate::orchestration::actor::{OrchestratorActor, OrchestratorReport}; use crate::orchestration::config::{DEFAULT_CONFIG_PATH, TomlConfigOverlay}; -use crate::transport::codec_registry::register_mvp_actor_codecs; +use crate::transport::register_mvp_actor_codecs; const PROVIDER_START_MAX_ATTEMPTS: usize = 4; use crate::gguf_shard::{StageShardPlan, plan_stage_shard}; @@ -643,81 +643,81 @@ struct VastAiRuntimeConfig { impl VastAiRuntimeConfig { fn from_builder(builder: &ConfigBuilder) -> Result { let mut provisioning = VastAiProvisioningConfig::default(); - let disk_gb = builder - .vastai_disk_gb_raw - .as_ref() - .map(|value| ConfigBuilder::parse_value("MVP_VASTAI_DISK_GB", value)) - .transpose()? - .or(builder.vastai_disk_gb); - if let Some(disk_gb) = disk_gb { + macro_rules! raw_config { + ($parser:ident, $env:literal, $raw:expr, $value:expr) => { + $raw.as_ref() + .map(|value| ConfigBuilder::$parser($env, value)) + .transpose()? + .or($value) + }; + } + if let Some(disk_gb) = raw_config!( + parse_value, + "MVP_VASTAI_DISK_GB", + builder.vastai_disk_gb_raw, + builder.vastai_disk_gb + ) { provisioning.disk_gb = disk_gb; } if let Some(ssh_user) = &builder.vastai_ssh_user { provisioning.ssh_user = ssh_user.clone(); } - let confirm_lease = builder - .vastai_confirm_lease_raw - .as_ref() - .map(|value| ConfigBuilder::parse_bool("MVP_VASTAI_CONFIRM_LEASE", value)) - .transpose()? - .or(builder.vastai_confirm_lease); - if let Some(confirm_lease) = confirm_lease { + if let Some(confirm_lease) = raw_config!( + parse_bool, + "MVP_VASTAI_CONFIRM_LEASE", + builder.vastai_confirm_lease_raw, + builder.vastai_confirm_lease + ) { provisioning.confirm_lease = confirm_lease; } provisioning.onstart = builder.vastai_onstart.clone(); provisioning.selection.gpu_name = builder.vastai_gpu_name.clone(); - let min_gpu_ram_mb = builder - .vastai_min_gpu_ram_mb_raw - .as_ref() - .map(|value| ConfigBuilder::parse_value("MVP_VASTAI_MIN_GPU_RAM_MB", value)) - .transpose()? - .or(builder.vastai_min_gpu_ram_mb); - if let Some(min_gpu_ram_mb) = min_gpu_ram_mb { + if let Some(min_gpu_ram_mb) = raw_config!( + parse_value, + "MVP_VASTAI_MIN_GPU_RAM_MB", + builder.vastai_min_gpu_ram_mb_raw, + builder.vastai_min_gpu_ram_mb + ) { provisioning.selection.min_gpu_ram_mb = Some(min_gpu_ram_mb); } - let min_down_mbps = builder - .vastai_min_down_mbps_raw - .as_ref() - .map(|value| ConfigBuilder::parse_value("MVP_VASTAI_MIN_DOWN_MBPS", value)) - .transpose()? - .or(builder.vastai_min_down_mbps); - if let Some(min_down_mbps) = min_down_mbps { + if let Some(min_down_mbps) = raw_config!( + parse_value, + "MVP_VASTAI_MIN_DOWN_MBPS", + builder.vastai_min_down_mbps_raw, + builder.vastai_min_down_mbps + ) { provisioning.selection.min_down_mbps = min_down_mbps; } - let max_dph_total = builder - .vastai_max_dph_total_raw - .as_ref() - .map(|value| ConfigBuilder::parse_value("MVP_VASTAI_MAX_DPH_TOTAL", value)) - .transpose()? - .or(builder.vastai_max_dph_total); - if let Some(max_dph_total) = max_dph_total { + if let Some(max_dph_total) = raw_config!( + parse_value, + "MVP_VASTAI_MAX_DPH_TOTAL", + builder.vastai_max_dph_total_raw, + builder.vastai_max_dph_total + ) { provisioning.selection.max_dph_total = Some(max_dph_total); } - let min_up_mbps = builder - .vastai_min_up_mbps_raw - .as_ref() - .map(|value| ConfigBuilder::parse_value("MVP_VASTAI_MIN_UP_MBPS", value)) - .transpose()? - .or(builder.vastai_min_up_mbps); - if let Some(min_up_mbps) = min_up_mbps { + if let Some(min_up_mbps) = raw_config!( + parse_value, + "MVP_VASTAI_MIN_UP_MBPS", + builder.vastai_min_up_mbps_raw, + builder.vastai_min_up_mbps + ) { provisioning.selection.min_up_mbps = Some(min_up_mbps); } - let min_reliability = builder - .vastai_min_reliability_raw - .as_ref() - .map(|value| ConfigBuilder::parse_value("MVP_VASTAI_MIN_RELIABILITY", value)) - .transpose()? - .or(builder.vastai_min_reliability); - if let Some(min_reliability) = min_reliability { + if let Some(min_reliability) = raw_config!( + parse_value, + "MVP_VASTAI_MIN_RELIABILITY", + builder.vastai_min_reliability_raw, + builder.vastai_min_reliability + ) { provisioning.selection.min_reliability = min_reliability; } - let require_verified = builder - .vastai_require_verified_raw - .as_ref() - .map(|value| ConfigBuilder::parse_bool("MVP_VASTAI_REQUIRE_VERIFIED", value)) - .transpose()? - .or(builder.vastai_require_verified); - if let Some(require_verified) = require_verified { + if let Some(require_verified) = raw_config!( + parse_bool, + "MVP_VASTAI_REQUIRE_VERIFIED", + builder.vastai_require_verified_raw, + builder.vastai_require_verified + ) { provisioning.selection.require_verified = require_verified; } for host_id in &builder.vastai_blacklist_hosts { @@ -725,13 +725,12 @@ impl VastAiRuntimeConfig { provisioning.selection.blacklist_hosts.push(*host_id); } } - let poll_interval_secs = builder - .vastai_poll_interval_secs_raw - .as_ref() - .map(|value| ConfigBuilder::parse_value("MVP_VASTAI_POLL_INTERVAL_SECS", value)) - .transpose()? - .or(builder.vastai_poll_interval_secs); - if let Some(poll_interval_secs) = poll_interval_secs { + if let Some(poll_interval_secs) = raw_config!( + parse_value, + "MVP_VASTAI_POLL_INTERVAL_SECS", + builder.vastai_poll_interval_secs_raw, + builder.vastai_poll_interval_secs + ) { provisioning.lifecycle.poll_interval = Duration::from_secs(poll_interval_secs); } let ssh_identity = builder @@ -797,11 +796,8 @@ struct CachedModelConfig { } impl CachedModelConfig { - fn from_host_path(provider: &ProviderKind, requested: PathBuf) -> Result { - if provider != &provider_kind::process() - && provider != &provider_kind::docker() - && provider != &provider_kind::vastai() - { + fn from_host_path(provider: &str, requested: PathBuf) -> Result { + if !matches!(provider, "process" | "docker" | "vastai") { return Err(format!( "{CACHED_MODEL_HOST_ENV} is a host-local cache path and is only supported by provider=process, provider=docker, or vastai planning" )); @@ -1432,8 +1428,9 @@ impl ConfigBuilder { RuntimeConfigProfile::Local => provider_kind::process(), RuntimeConfigProfile::Deploy => provider_kind::vastai(), }); + let provider_name = provider.as_str(); let mut image = self.image.clone(); - if provider == provider_kind::vastai() && !self.image_overridden_after_toml { + if provider_name == "vastai" && !self.image_overridden_after_toml { if let Some(vastai_image) = &self.toml_vastai_image { image = vastai_image.clone(); } @@ -1442,7 +1439,7 @@ impl ConfigBuilder { return Err("--pipeline-stages must be greater than 0".to_owned()); } let mut cached_model_host_path = self.cached_model_host_path.clone(); - if (provider == provider_kind::process() || provider == provider_kind::docker()) + if matches!(provider_name, "process" | "docker") && self.pipeline_stages > 1 && cached_model_host_path.is_none() && (matches!( @@ -1464,12 +1461,12 @@ impl ConfigBuilder { cached_model_host_path = Some(default_pipeline_cached_model_path()); } let cached_model = cached_model_host_path - .map(|path| CachedModelConfig::from_host_path(&provider, path)) + .map(|path| CachedModelConfig::from_host_path(provider_name, path)) .transpose()?; let mut gguf_source = self.gguf_source.clone(); if let Some(cached_model) = &cached_model { - if provider != provider_kind::vastai() { - gguf_source = GgufSource::LocalPath(if provider == provider_kind::process() { + if provider_name != "vastai" { + gguf_source = GgufSource::LocalPath(if provider_name == "process" { cached_model.host_path.to_string_lossy().to_string() } else { cached_model.container_path.clone() @@ -1485,11 +1482,9 @@ impl ConfigBuilder { Some(mask) => EndpointAddrMask::parse(mask)?, None => EndpointAddrMask::Full, }; - let vastai = if provider == provider_kind::vastai() { - Some(VastAiRuntimeConfig::from_builder(&self)?) - } else { - None - }; + let vastai = (provider_name == "vastai") + .then(|| VastAiRuntimeConfig::from_builder(&self)) + .transpose()?; Ok(Config { config_profile: self.config_profile, image, @@ -1588,22 +1583,20 @@ impl Config { } fn provider_datastream_detail(&self) -> Value { - if self.provider == provider_kind::process() { - json!({ + match self.provider.as_str() { + "process" => json!({ "worker_bin": self.worker_bin.as_ref().map(|path| path.to_string_lossy().to_string()), "cached_model": self.cached_model.as_ref().map(CachedModelConfig::datastream_detail), - }) - } else if self.provider == provider_kind::docker() { - json!({ + }), + "docker" => json!({ "docker_gpus": &self.docker_gpus, "cached_model": self.cached_model.as_ref().map(CachedModelConfig::datastream_detail), - }) - } else if self.provider == provider_kind::vastai() { - self.vastai + }), + "vastai" => self + .vastai .as_ref() - .map_or_else(|| json!({}), VastAiRuntimeConfig::datastream_detail) - } else { - json!({}) + .map_or_else(|| json!({}), VastAiRuntimeConfig::datastream_detail), + _ => json!({}), } } @@ -1662,12 +1655,10 @@ impl Config { model, runtime: run_plan::RuntimeConfig { max_tokens: self.default_max_tokens, - prompt: run_plan::PromptSource::Inline(String::new()), sampling: run_plan::SamplingPolicy { temperature_millis: 0, top_k: 1, }, - token_output_policy: run_plan::TokenOutputPolicy::EmitAll, }, candidate_pool, stage_count: self.pipeline_stages, @@ -1709,7 +1700,7 @@ impl Config { repo, file, revision: None, - } if self.provider == provider_kind::vastai() + } if self.provider.as_str() == "vastai" && file == DEFAULT_PIPELINE_CACHED_MODEL_FILE => { let host_path = default_pipeline_cached_model_path(); @@ -1729,7 +1720,7 @@ impl Config { } fn prepare_vastai_ssh_key(&mut self) -> Result<(), String> { - if self.provider != provider_kind::vastai() { + if self.provider.as_str() != "vastai" { return Ok(()); } @@ -1777,54 +1768,54 @@ impl Config { &self, bootstrap_runtime: Arc, ) -> Result, String> { - if self.provider == provider_kind::process() { - let worker_bin = match &self.worker_bin { - Some(worker_bin) => worker_bin.clone(), - None => { - let mut path = - std::env::current_exe().map_err(|e| format!("current exe: {e}"))?; - path.set_file_name("mvp-worker-node"); - path + match self.provider.as_str() { + "process" => { + let worker_bin = match &self.worker_bin { + Some(worker_bin) => worker_bin.clone(), + None => { + let mut path = + std::env::current_exe().map_err(|e| format!("current exe: {e}"))?; + path.set_file_name("mvp-worker-node"); + path + } + }; + if !worker_bin.is_file() { + return Err(format!( + "local process worker binary does not exist: {}", + worker_bin.display() + )); } - }; - if !worker_bin.is_file() { - return Err(format!( - "local process worker binary does not exist: {}", - worker_bin.display() - )); + Ok(Box::new(LocalProcessPlugin::new(worker_bin))) } - Ok(Box::new(LocalProcessPlugin::new(worker_bin))) - } else if self.provider == provider_kind::docker() { - Ok(Box::new(LocalDockerPlugin::new( + "docker" => Ok(Box::new(LocalDockerPlugin::new( env_optional("MVP_DOCKER_CONTAINER_PREFIX") .unwrap_or_else(|| "mvp-orchestrator".to_owned()), - ))) - } else if self.provider == provider_kind::vastai() { - let vastai = self - .vastai - .as_ref() - .ok_or_else(|| "VastAI config was not resolved for provider vastai".to_owned())?; - if vastai.bootstrap_command.is_none() { - return Err( - "MVP_VASTAI_BOOTSTRAP_COMMAND is required when MVP_NODE_PROVIDER=vastai" - .to_owned(), - ); + ))), + "vastai" => { + let vastai = self.vastai.as_ref().ok_or_else(|| { + "VastAI config was not resolved for provider vastai".to_owned() + })?; + if vastai.bootstrap_command.is_none() { + return Err( + "MVP_VASTAI_BOOTSTRAP_COMMAND is required when MVP_NODE_PROVIDER=vastai" + .to_owned(), + ); + } + let api_key = vastai.api_key.clone().ok_or_else(|| { + "VAST_API_KEY, MVP_VASTAI_API_KEY, or VASTAI_API_KEY is required when MVP_NODE_PROVIDER=vastai" + .to_owned() + })?; + let ssh_identity = vastai + .ssh_identity + .clone() + .ok_or_else(|| "VastAI SSH identity was not prepared".to_owned())?; + Ok(Box::new(VastAiProvisioningPlugin::new( + ToolsVastAiLeaseClient::from_api_key(api_key)?, + SshCommandBootstrapLauncher::new(Some(ssh_identity), bootstrap_runtime), + vastai.provisioning.clone(), + ))) } - let api_key = vastai.api_key.clone().ok_or_else(|| { - "VAST_API_KEY, MVP_VASTAI_API_KEY, or VASTAI_API_KEY is required when MVP_NODE_PROVIDER=vastai" - .to_owned() - })?; - let ssh_identity = vastai - .ssh_identity - .clone() - .ok_or_else(|| "VastAI SSH identity was not prepared".to_owned())?; - Ok(Box::new(VastAiProvisioningPlugin::new( - ToolsVastAiLeaseClient::from_api_key(api_key)?, - SshCommandBootstrapLauncher::new(Some(ssh_identity), bootstrap_runtime), - vastai.provisioning.clone(), - ))) - } else { - Err("mock provider cannot build a runtime provisioner".to_owned()) + _ => Err("mock provider cannot build a runtime provisioner".to_owned()), } } @@ -1844,13 +1835,13 @@ impl Config { if self.relay.url.is_some() { keys.push(MVP_IROH_RELAY_URL_ENV); } - if self.provider == provider_kind::docker() { + if self.provider.as_str() == "docker" { keys.push("MVP_DOCKER_GPUS"); } if std::env::var_os("DEV").is_some() { keys.push("DEV"); } - if local_tinygrad_worker_env(&self.provider).is_some() { + if local_tinygrad_worker_env(self.provider.as_str()).is_some() { keys.push("MVP_TINYGRAD_WORKER"); } for key in [ @@ -1891,6 +1882,7 @@ impl Config { logical_node_id: u64, stage_index: u32, ) -> Result { + let provider_name = self.provider.as_str(); let mut env = vec![ ("MVP_RUN_ID".to_owned(), self.run_id.to_string()), ( @@ -1906,10 +1898,7 @@ impl Config { MVP_IROH_ENDPOINT_ADDR_MASK_ENV.to_owned(), self.endpoint_addr_mask.as_str().to_owned(), ), - ( - "MVP_NODE_PROVIDER".to_owned(), - self.provider.as_str().to_owned(), - ), + ("MVP_NODE_PROVIDER".to_owned(), provider_name.to_owned()), ( "MVP_COORDINATOR_ENDPOINT".to_owned(), serde_json::to_string(&coordinator) @@ -1929,13 +1918,13 @@ impl Config { if let Some(url) = &self.relay.url { env.push((MVP_IROH_RELAY_URL_ENV.to_owned(), url.clone())); } - if self.provider == provider_kind::docker() { + if provider_name == "docker" { env.push(("MVP_DOCKER_GPUS".to_owned(), self.docker_gpus.clone())); } if let Some(value) = env_optional("DEV") { env.push(("DEV".to_owned(), value)); } - env.extend(local_tinygrad_worker_env(&self.provider)); + env.extend(local_tinygrad_worker_env(provider_name)); for name in [ "MVP_CPU_LINE_PROFILE", "MVP_CPU_LINE_PROFILE_INTERVAL_MS", @@ -1970,20 +1959,17 @@ impl Config { if let Some(max_context) = self.max_context { env.push(("MVP_MAX_CONTEXT".to_owned(), max_context.to_string())); } - let args = if self.provider == provider_kind::vastai() { - self.vastai + let args = match provider_name { + "vastai" => self + .vastai .as_ref() .and_then(|vastai| vastai.bootstrap_command.clone()) .into_iter() - .collect() - } else if self.provider == provider_kind::process() - || self.provider == provider_kind::docker() - { - Vec::new() - } else { - return Err("mvp-orchestrator does not support mock provider".to_owned()); + .collect(), + "process" | "docker" => Vec::new(), + _ => return Err("mvp-orchestrator does not support mock provider".to_owned()), }; - let mounts = if self.provider == provider_kind::docker() { + let mounts = if provider_name == "docker" { self.cached_model .as_ref() .map(|cached_model| { @@ -2332,7 +2318,7 @@ fn start_and_provision_workers( "image":&config.image, "relay_mode":relay_mode_env_value(&config.relay.mode), "endpoint_addr_mask":config.endpoint_addr_mask.as_str(), - "docker_gpus":if config.provider == provider_kind::docker() { Some(config.docker_gpus.as_str()) } else { None }, + "docker_gpus":if config.provider.as_str() == "docker" { Some(config.docker_gpus.as_str()) } else { None }, "provider_config":config.provider_datastream_detail(), "env_keys":config.node_spec_env_keys(), "worker_count":stage_specs.len(), @@ -5562,19 +5548,19 @@ fn env_optional(name: &str) -> Option { .filter(|value| !value.is_empty()) } -fn local_tinygrad_worker_env(provider: &ProviderKind) -> Option<(String, String)> { +fn local_tinygrad_worker_env(provider: &str) -> Option<(String, String)> { env_optional("MVP_TINYGRAD_WORKER") .map(|value| ("MVP_TINYGRAD_WORKER".to_owned(), value)) .or_else(|| { - if provider != &provider_kind::process() { - return None; - } - default_local_tinygrad_worker_path().map(|path| { - ( - "MVP_TINYGRAD_WORKER".to_owned(), - path.to_string_lossy().to_string(), - ) - }) + (provider == "process") + .then(default_local_tinygrad_worker_path) + .flatten() + .map(|path| { + ( + "MVP_TINYGRAD_WORKER".to_owned(), + path.to_string_lossy().to_string(), + ) + }) }) } @@ -5650,28 +5636,30 @@ fn derive_ssh_public_key(identity: &Path) -> Result { } fn ssh_public_key_fingerprint(public_key: &str) -> String { + const UNAVAILABLE: &str = "unavailable"; + let path = std::env::temp_dir().join(format!("mvp-vastai-ssh-key-{}.pub", std::process::id())); if std::fs::write(&path, format!("{public_key}\n")).is_err() { - return "unavailable".to_owned(); + return UNAVAILABLE.to_owned(); } let output = Command::new("ssh-keygen") .arg("-l") .arg("-f") .arg(&path) - .output(); + .output() + .ok(); let _ = std::fs::remove_file(&path); - let Ok(output) = output else { - return "unavailable".to_owned(); + let Some(output) = output.filter(|output| output.status.success()) else { + return UNAVAILABLE.to_owned(); }; - if !output.status.success() { - return "unavailable".to_owned(); - } + let stdout = String::from_utf8_lossy(&output.stdout); let mut fields = stdout.split_whitespace(); - match (fields.next(), fields.next()) { - (Some(bits), Some(fingerprint)) => format!("{bits} {fingerprint}"), - _ => "unavailable".to_owned(), - } + fields + .next() + .zip(fields.next()) + .map(|(bits, fingerprint)| format!("{bits} {fingerprint}")) + .unwrap_or_else(|| UNAVAILABLE.to_owned()) } fn vastai_account_has_ssh_key(api_key: &str, public_key: &str) -> Result { @@ -5721,16 +5709,12 @@ fn ensure_vastai_account_ssh_key(api_key: &str, public_key: &str) -> Result<(), fn account_ssh_keys_output_contains_public_key(output: &str, public_key: &str) -> bool { let public_key = public_key.trim(); - if public_key.is_empty() { - return false; - } - if output.contains(public_key) { - return true; - } - public_key - .split_whitespace() - .nth(1) - .is_some_and(|body| !body.is_empty() && output.contains(body)) + !public_key.is_empty() + && (output.contains(public_key) + || public_key + .split_whitespace() + .nth(1) + .is_some_and(|body| !body.is_empty() && output.contains(body))) } fn vastai_cli_error(error: std::io::Error) -> String { @@ -5743,10 +5727,11 @@ fn vastai_cli_error(error: std::io::Error) -> String { } fn command_output_failure_detail(output: &std::process::Output, secret: Option<&str>) -> String { - let mut detail = String::from_utf8_lossy(&output.stderr).trim().to_owned(); - if detail.is_empty() { - detail = output.status.to_string(); - } + let stderr = String::from_utf8_lossy(&output.stderr); + let mut detail = match stderr.trim() { + "" => output.status.to_string(), + detail => detail.to_owned(), + }; if let Some(secret) = secret.filter(|secret| !secret.is_empty()) { detail = detail.replace(secret, ""); } diff --git a/crates/mvp-system/src/orchestration/config.rs b/crates/mvp-system/src/orchestration/config.rs index 3622ee7..ed9325a 100644 --- a/crates/mvp-system/src/orchestration/config.rs +++ b/crates/mvp-system/src/orchestration/config.rs @@ -3,11 +3,11 @@ use std::path::Path; use serde::Deserialize; -pub const DEFAULT_CONFIG_PATH: &str = ".config/config.toml"; +pub(crate) const DEFAULT_CONFIG_PATH: &str = ".config/config.toml"; #[derive(Clone, Debug, Default, Deserialize, PartialEq)] #[serde(default)] -pub struct VastAiConfig { +pub(crate) struct VastAiConfig { pub api_key: Option, pub image: Option, pub relay_url: Option, @@ -29,7 +29,7 @@ pub struct VastAiConfig { } #[derive(Clone, Debug, PartialEq)] -pub struct ResolvedVastAiConfig { +pub(crate) struct ResolvedVastAiConfig { pub api_key: String, pub relay_url: String, pub image: String, @@ -48,7 +48,7 @@ pub struct ResolvedVastAiConfig { } impl ResolvedVastAiConfig { - pub fn validate(self) -> Result { + pub(crate) fn validate(self) -> Result { require_non_empty("VAST_API_KEY", &self.api_key)?; require_non_empty("relay.url", &self.relay_url)?; require_non_empty("vastai.image", &self.image)?; @@ -74,7 +74,7 @@ fn require_non_empty(label: &str, value: &str) -> Result<(), String> { } } -pub fn looks_remote_image(image: &str) -> bool { +pub(crate) fn looks_remote_image(image: &str) -> bool { let repository = image.split('@').next().unwrap_or(image); let last_slash = repository.rfind('/'); let tag_separator = repository @@ -92,7 +92,7 @@ pub fn looks_remote_image(image: &str) -> bool { /// bin-local strict config loader instead. #[derive(Clone, Debug, Default, Deserialize, PartialEq)] #[serde(default)] -pub struct TomlConfigOverlay { +pub(crate) struct TomlConfigOverlay { pub runtime: RuntimeConfigOverlay, pub provider: ProviderConfigOverlay, pub image: ImageConfig, @@ -106,7 +106,7 @@ pub struct TomlConfigOverlay { #[derive(Clone, Debug, Default, Deserialize, PartialEq)] #[serde(default)] -pub struct RuntimeConfigOverlay { +pub(crate) struct RuntimeConfigOverlay { pub profile: Option, pub run_id: Option, pub node_id: Option, @@ -117,13 +117,13 @@ pub struct RuntimeConfigOverlay { #[derive(Clone, Debug, Default, Deserialize, PartialEq)] #[serde(default)] -pub struct ProviderConfigOverlay { +pub(crate) struct ProviderConfigOverlay { pub kind: Option, } #[derive(Clone, Debug, Default, Deserialize, PartialEq)] #[serde(default)] -pub struct ImageConfig { +pub(crate) struct ImageConfig { pub node: Option, pub tag: Option, pub build: Option, @@ -133,14 +133,14 @@ pub struct ImageConfig { #[derive(Clone, Debug, Default, Deserialize, PartialEq)] #[serde(default)] -pub struct RelayConfig { +pub(crate) struct RelayConfig { pub mode: Option, pub url: Option, } #[derive(Clone, Debug, Default, Deserialize, PartialEq)] #[serde(default)] -pub struct PromptConfig { +pub(crate) struct PromptConfig { pub rpc_addr: Option, pub max_tokens: Option, pub dashboard: Option, @@ -148,7 +148,7 @@ pub struct PromptConfig { #[derive(Clone, Debug, Default, Deserialize, PartialEq)] #[serde(default)] -pub struct ModelConfig { +pub(crate) struct ModelConfig { pub id: Option, pub gguf_local_path: Option, pub gguf_repo: Option, @@ -160,21 +160,21 @@ pub struct ModelConfig { #[derive(Clone, Debug, Default, Deserialize, PartialEq)] #[serde(default)] -pub struct DockerConfigOverlay { +pub(crate) struct DockerConfigOverlay { pub gpus: Option, pub cached_model_host_path: Option, } #[derive(Clone, Debug, Default, Deserialize, PartialEq)] #[serde(default)] -pub struct ObservabilityConfigOverlay { +pub(crate) struct ObservabilityConfigOverlay { pub dump_logs: Option, pub dump_log_path: Option, pub datastream_frame_log: Option, } impl TomlConfigOverlay { - pub fn load_optional(path: &Path) -> Result, String> { + pub(crate) fn load_optional(path: &Path) -> Result, String> { if path.is_file() { Self::load_required(path).map(Some) } else { @@ -182,13 +182,13 @@ impl TomlConfigOverlay { } } - pub fn load_required(path: &Path) -> Result { + pub(crate) fn load_required(path: &Path) -> Result { let text = fs::read_to_string(path).map_err(|e| format!("read config {}: {e}", path.display()))?; Self::from_str(&text).map_err(|e| format!("parse config {}: {e}", path.display())) } - pub fn from_str(text: &str) -> Result { + pub(crate) fn from_str(text: &str) -> Result { toml::from_str(text) } } diff --git a/crates/mvp-system/src/orchestration/distribution_stack.rs b/crates/mvp-system/src/orchestration/distribution_stack.rs index 982ddf8..8f05320 100644 --- a/crates/mvp-system/src/orchestration/distribution_stack.rs +++ b/crates/mvp-system/src/orchestration/distribution_stack.rs @@ -34,7 +34,7 @@ use distribution::transport_bridge::{ use distribution::types::{DirectoryEntry, MemberState, NodeId}; #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct DistributionActorAddrs { +pub(crate) struct DistributionActorAddrs { pub swim: ActorAddress, pub registry: ActorAddress, pub metadata: ActorAddress, @@ -42,7 +42,7 @@ pub struct DistributionActorAddrs { pub membership_fanout: ActorAddress, } -pub struct DistributionRuntimeStack { +pub(crate) struct DistributionRuntimeStack { pub node_id: NodeId, pub runtime: Arc, pub codec: Arc, @@ -56,11 +56,11 @@ pub struct DistributionRuntimeStack { } impl DistributionRuntimeStack { - pub fn new(node_id: NodeId, config: DistributedNodeConfig) -> Self { + pub(crate) fn new(node_id: NodeId, config: DistributedNodeConfig) -> Self { Self::new_with_codecs(node_id, config, |_| {}) } - pub fn new_with_codecs( + pub(crate) fn new_with_codecs( node_id: NodeId, config: DistributedNodeConfig, extend_codecs: impl FnOnce(&mut CodecRegistry), @@ -168,7 +168,7 @@ impl DistributionRuntimeStack { } } - pub fn actor_bridge_routes(&self) -> HashMap { + pub(crate) fn actor_bridge_routes(&self) -> HashMap { let mut routes = HashMap::new(); for tag in [ "swactor_dist::Ping", @@ -189,7 +189,7 @@ impl DistributionRuntimeStack { routes } - pub fn tick_protocol_actors(&self, now: Instant) { + pub(crate) fn tick_protocol_actors(&self, now: Instant) { let _ = self.runtime.send_to(self.actors.swim, SwimIn::Tick { now }); let _ = self.runtime.send_to(self.actors.registry, RegistryIn::Tick); let _ = self.runtime.send_to(self.actors.metadata, MetadataIn::Tick); @@ -198,17 +198,17 @@ impl DistributionRuntimeStack { .send_to(self.actors.directory, DirectoryIn::Tick); } - pub fn pump_runtime_once(&self) { + pub(crate) fn pump_runtime_once(&self) { self.runtime.tick(); } - pub fn register_local_actor(&self, entry: DirectoryEntry) { + pub(crate) fn register_local_actor(&self, entry: DirectoryEntry) { let _ = self .runtime .send_to(self.actors.directory, DirectoryIn::Register(entry)); } - pub fn alive_count(&self) -> usize { + pub(crate) fn alive_count(&self) -> usize { self.membership_mirror .lock() .expect("membership mirror poisoned") @@ -218,7 +218,7 @@ impl DistributionRuntimeStack { .count() } - pub fn member_state(&self, node_id: NodeId) -> Option { + pub(crate) fn member_state(&self, node_id: NodeId) -> Option { self.membership_mirror .lock() .ok()? @@ -226,15 +226,15 @@ impl DistributionRuntimeStack { .map(|entry| entry.state) } - pub fn route_owner(&self, actor: ActorAddress) -> Option { + pub(crate) fn route_owner(&self, actor: ActorAddress) -> Option { self.route_view.read().ok()?.get(&actor).copied() } - pub fn drain_swim_transitions(&self) -> Vec { + pub(crate) fn drain_swim_transitions(&self) -> Vec { self.swim_telemetry.drain_transitions() } - pub fn drain_swim_probe_events(&self) -> Vec { + pub(crate) fn drain_swim_probe_events(&self) -> Vec { self.swim_telemetry.drain_probe_events() } } diff --git a/crates/mvp-system/src/orchestration/engine_builder/engine.rs b/crates/mvp-system/src/orchestration/engine_builder/engine.rs index 89eb657..4276dbb 100644 --- a/crates/mvp-system/src/orchestration/engine_builder/engine.rs +++ b/crates/mvp-system/src/orchestration/engine_builder/engine.rs @@ -1,108 +1,163 @@ use crate::run_plan::RunId; use super::error::EngineBuildError; -use super::events::EngineEvent; -use super::launcher::{LaunchedNode, NodeControl, NodeFacts, NodeLaunchSpec, StaticNodeLauncher}; -use super::model::ModelSpec; -use super::node_image::NodeImageSpec; -use super::planner::{FixedLinearPipelinePlanner, RoleAssignmentPlan, RolePlannerInput}; -use super::pool::{PoolRequest, StaticPoolProvider}; -use super::roles::RoleAssignment; +use super::planner::{ + FixedLinearPipelinePlanner, RoleAssignment, RoleAssignmentPlan, RoleKind, RolePlannerInput, +}; +use super::pool::{ModelSpec, NodeFacts, StaticPoolProvider}; -pub struct ClusterBuilder { +#[derive(Clone, Debug, PartialEq, Eq)] +pub(crate) enum EngineEvent { + PoolAcquired, + ClusterConverged, + RoleAssigned(RoleKind), + EngineReady, +} + +fn launch_node(facts: &NodeFacts) -> StaticNodeControl { + StaticNodeControl { + facts: facts.clone(), + booted: false, + stopped: false, + assigned_roles: Vec::new(), + } +} + +struct StaticNodeControl { + facts: NodeFacts, + booted: bool, + stopped: bool, + assigned_roles: Vec, +} + +impl StaticNodeControl { + fn node_id(&self) -> u64 { + self.facts.node_id.0 + } +} + +impl StaticNodeControl { + fn wait_boot_ready(&mut self) -> Result { + if self.stopped { + return Err(EngineBuildError::Stopped { + node_id: self.node_id(), + }); + } + self.booted = true; + Ok(self.facts.clone()) + } + + fn wait_cluster_converged(&mut self, expected_alive: usize) -> Result<(), EngineBuildError> { + if self.stopped { + return Err(EngineBuildError::Stopped { + node_id: self.node_id(), + }); + } + if !self.booted { + return Err(EngineBuildError::NotBooted { + node_id: self.node_id(), + }); + } + if expected_alive == 0 { + return Err(EngineBuildError::Backend( + "expected_alive must be greater than zero", + )); + } + Ok(()) + } + + fn assign_role(&mut self, assignment: RoleAssignment) -> Result<(), EngineBuildError> { + if self.stopped { + return Err(EngineBuildError::Stopped { + node_id: self.node_id(), + }); + } + if !self.booted { + return Err(EngineBuildError::NotBooted { + node_id: self.node_id(), + }); + } + let role_node_id = assignment.node_id().0; + if role_node_id != self.node_id() { + return Err(EngineBuildError::RoleNodeMismatch { + node_id: self.node_id(), + role_node_id, + }); + } + self.assigned_roles.push(assignment); + Ok(()) + } + + fn shutdown(&mut self) -> Result<(), EngineBuildError> { + if self.stopped { + return Ok(()); + } + self.stopped = true; + Ok(()) + } +} + +pub(crate) struct ClusterBuilder { cluster_id: String, run_id: RunId, model: ModelSpec, pool_provider: Option, - launcher: Option, planner: Option, } impl ClusterBuilder { - pub fn new(cluster_id: impl Into, model: ModelSpec) -> Self { + pub(crate) fn new(cluster_id: impl Into, model: ModelSpec) -> Self { Self { cluster_id: cluster_id.into(), run_id: RunId(1), model, pool_provider: None, - launcher: None, planner: None, } } - pub fn run_id(mut self, run_id: impl Into) -> Self { + pub(crate) fn run_id(mut self, run_id: impl Into) -> Self { self.run_id = run_id.into(); self } - pub fn image(self, _image: NodeImageSpec) -> Self { - self - } - - pub fn pool_provider(mut self, provider: StaticPoolProvider) -> Self { + pub(crate) fn pool_provider(mut self, provider: StaticPoolProvider) -> Self { self.pool_provider = Some(provider); self } - pub fn launcher(mut self, launcher: StaticNodeLauncher) -> Self { - self.launcher = Some(launcher); - self - } - - pub fn planner(mut self, planner: FixedLinearPipelinePlanner) -> Self { + pub(crate) fn planner(mut self, planner: FixedLinearPipelinePlanner) -> Self { self.planner = Some(planner); self } - pub fn launch(mut self) -> Result { + pub(crate) fn launch(mut self) -> Result { let pool_provider = self .pool_provider .take() .ok_or(EngineBuildError::MissingComponent("pool_provider"))?; - let launcher = self - .launcher - .take() - .ok_or(EngineBuildError::MissingComponent("launcher"))?; let planner = self .planner .take() .ok_or(EngineBuildError::MissingComponent("planner"))?; let mut events = Vec::new(); - let leases = pool_provider.acquire_pool(PoolRequest { - min_nodes: planner.required_node_count(), - })?; + let leases = pool_provider.acquire_pool(planner.required_node_count())?; if leases.is_empty() { return Err(EngineBuildError::EmptyPool); } - events.push(EngineEvent::PoolAcquired { - node_count: leases.len(), - }); + events.push(EngineEvent::PoolAcquired); let mut nodes = Vec::with_capacity(leases.len()); let mut iter = leases.into_iter(); let coordinator_lease = iter.next().ok_or(EngineBuildError::EmptyPool)?; - let mut coordinator = launcher.launch_node(&coordinator_lease, NodeLaunchSpec); - events.push(EngineEvent::NodeLaunched { - node_id: coordinator.lease.logical_node_id, - coordinator: true, - }); - let coordinator_facts = coordinator.control.wait_boot_ready()?; - events.push(EngineEvent::NodeBootReady { - node_id: coordinator_facts.node_id, - }); + let mut coordinator = launch_node(&coordinator_lease); + let coordinator_facts = coordinator.wait_boot_ready()?; nodes.push(EngineNode::new(coordinator, coordinator_facts)); for lease in iter { - let mut node = launcher.launch_node(&lease, NodeLaunchSpec); - events.push(EngineEvent::NodeLaunched { - node_id: node.lease.logical_node_id, - coordinator: false, - }); - let facts = node.control.wait_boot_ready()?; - events.push(EngineEvent::NodeBootReady { - node_id: facts.node_id, - }); + let mut node = launch_node(&lease); + let facts = node.wait_boot_ready()?; nodes.push(EngineNode::new(node, facts)); } @@ -110,18 +165,13 @@ impl ClusterBuilder { for node in &mut nodes { node.control.wait_cluster_converged(expected_alive)?; } - events.push(EngineEvent::ClusterConverged { - node_count: expected_alive, - }); + events.push(EngineEvent::ClusterConverged); let plan = planner.plan(RolePlannerInput { run_id: self.run_id, model: self.model, nodes: nodes.iter().map(|node| node.facts.clone()).collect(), })?; - events.push(EngineEvent::RolesPlanned { - stage_count: plan.stages.len(), - }); assign_role( &mut nodes, @@ -135,12 +185,9 @@ impl ClusterBuilder { &mut events, )?; } - events.push(EngineEvent::EngineReady { - cluster_id: self.cluster_id.clone(), - }); + events.push(EngineEvent::EngineReady); Ok(ClusterHandle { - cluster_id: self.cluster_id, nodes, plan, events, @@ -148,47 +195,40 @@ impl ClusterBuilder { } } -pub struct ClusterHandle { - cluster_id: String, +pub(crate) struct ClusterHandle { nodes: Vec, plan: RoleAssignmentPlan, events: Vec, } impl ClusterHandle { - pub fn role_plan(&self) -> &RoleAssignmentPlan { + pub(crate) fn role_plan(&self) -> &RoleAssignmentPlan { &self.plan } - pub fn events(&self) -> &[EngineEvent] { + pub(crate) fn events(&self) -> &[EngineEvent] { &self.events } - pub fn shutdown(mut self) -> Result, EngineBuildError> { + pub(crate) fn shutdown(mut self) -> Result<(), EngineBuildError> { for node in &mut self.nodes { node.control.shutdown()?; - self.events.push(EngineEvent::NodeStopped { - node_id: node.facts.node_id, - }); } - self.events.push(EngineEvent::ShutdownComplete { - cluster_id: self.cluster_id, - }); - Ok(self.events) + Ok(()) } } struct EngineNode { facts: NodeFacts, roles: Vec, - control: Box, + control: StaticNodeControl, } impl EngineNode { - fn new(launched: LaunchedNode, facts: NodeFacts) -> Self { + fn new(control: StaticNodeControl, facts: NodeFacts) -> Self { Self { facts, roles: Vec::new(), - control: launched.control, + control, } } } @@ -206,6 +246,6 @@ fn assign_role( node.control.assign_role(assignment.clone())?; let role = assignment.kind(); node.roles.push(assignment); - events.push(EngineEvent::RoleAssigned { node_id, role }); + events.push(EngineEvent::RoleAssigned(role)); Ok(()) } diff --git a/crates/mvp-system/src/orchestration/engine_builder/error.rs b/crates/mvp-system/src/orchestration/engine_builder/error.rs index d555c46..60f8680 100644 --- a/crates/mvp-system/src/orchestration/engine_builder/error.rs +++ b/crates/mvp-system/src/orchestration/engine_builder/error.rs @@ -1,121 +1,18 @@ -use std::fmt; - use crate::run_plan; #[derive(Clone, Debug, PartialEq, Eq)] -pub enum EngineBuildError { +pub(crate) enum EngineBuildError { MissingComponent(&'static str), EmptyPool, RoleTargetMissing { node_id: u64 }, - Pool(PoolError), - Node(NodeControlError), - Planning(PlanningError), -} - -impl fmt::Display for EngineBuildError { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - match self { - Self::MissingComponent(name) => write!(f, "missing engine builder component: {name}"), - Self::EmptyPool => write!(f, "pool provider returned no nodes"), - Self::RoleTargetMissing { node_id } => { - write!(f, "role assignment targeted unknown node {node_id}") - } - Self::Pool(err) => err.fmt(f), - Self::Node(err) => err.fmt(f), - Self::Planning(err) => err.fmt(f), - } - } -} - -impl From for EngineBuildError { - fn from(value: PoolError) -> Self { - Self::Pool(value) - } -} - -impl From for EngineBuildError { - fn from(value: NodeControlError) -> Self { - Self::Node(value) - } -} - -impl From for EngineBuildError { - fn from(value: PlanningError) -> Self { - Self::Planning(value) - } -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum PoolError { InsufficientNodes { requested: usize, available: usize }, -} - -impl fmt::Display for PoolError { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - match self { - Self::InsufficientNodes { - requested, - available, - } => write!( - f, - "pool has {available} matching nodes, but {requested} were requested" - ), - } - } -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum NodeControlError { NotBooted { node_id: u64 }, Stopped { node_id: u64 }, RoleNodeMismatch { node_id: u64, role_node_id: u64 }, Backend(&'static str), -} - -impl fmt::Display for NodeControlError { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - match self { - Self::NotBooted { node_id } => write!(f, "node {node_id} is not boot-ready"), - Self::Stopped { node_id } => write!(f, "node {node_id} is already stopped"), - Self::RoleNodeMismatch { - node_id, - role_node_id, - } => write!( - f, - "node {node_id} cannot accept role targeted at node {role_node_id}" - ), - Self::Backend(message) => write!(f, "node control failed: {message}"), - } - } -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum PlanningError { DuplicateNodeId { node_id: u64 }, NoCoordinatorCandidate, InsufficientWorkers { required: usize, available: usize }, ModelRejected(run_plan::PlanRejectionKind), StageProjection(run_plan::ProjectionRejection), } - -impl fmt::Display for PlanningError { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - match self { - Self::DuplicateNodeId { node_id } => { - write!(f, "planner input contained duplicate node id {node_id}") - } - Self::NoCoordinatorCandidate => write!(f, "no coordinator-capable node available"), - Self::InsufficientWorkers { - required, - available, - } => write!( - f, - "planner needs {required} worker nodes, but only {available} are available" - ), - Self::ModelRejected(kind) => write!(f, "model/run planner rejected input: {kind:?}"), - Self::StageProjection(kind) => { - write!(f, "stage provision projection failed: {kind:?}") - } - } - } -} diff --git a/crates/mvp-system/src/orchestration/engine_builder/events.rs b/crates/mvp-system/src/orchestration/engine_builder/events.rs deleted file mode 100644 index 7704f4b..0000000 --- a/crates/mvp-system/src/orchestration/engine_builder/events.rs +++ /dev/null @@ -1,16 +0,0 @@ -use crate::run_plan::NodeId; - -use super::roles::RoleKind; - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum EngineEvent { - PoolAcquired { node_count: usize }, - NodeLaunched { node_id: NodeId, coordinator: bool }, - NodeBootReady { node_id: NodeId }, - ClusterConverged { node_count: usize }, - RolesPlanned { stage_count: usize }, - RoleAssigned { node_id: NodeId, role: RoleKind }, - EngineReady { cluster_id: String }, - NodeStopped { node_id: NodeId }, - ShutdownComplete { cluster_id: String }, -} diff --git a/crates/mvp-system/src/orchestration/engine_builder/launcher.rs b/crates/mvp-system/src/orchestration/engine_builder/launcher.rs deleted file mode 100644 index 038275a..0000000 --- a/crates/mvp-system/src/orchestration/engine_builder/launcher.rs +++ /dev/null @@ -1,129 +0,0 @@ -use std::collections::BTreeSet; - -use crate::run_plan::NodeId; - -use super::error::NodeControlError; -use super::pool::{NodeCapability, NodeLease}; -use super::roles::RoleAssignment; - -pub trait NodeControl: Send { - fn wait_boot_ready(&mut self) -> Result; - fn wait_cluster_converged(&mut self, expected_alive: usize) -> Result<(), NodeControlError>; - fn assign_role(&mut self, assignment: RoleAssignment) -> Result<(), NodeControlError>; - fn shutdown(&mut self) -> Result<(), NodeControlError>; -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct NodeLaunchSpec; - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct NodeFacts { - pub node_id: NodeId, - pub capabilities: BTreeSet, -} - -impl NodeFacts { - pub fn from_lease(lease: &NodeLease) -> Self { - Self { - node_id: lease.logical_node_id, - capabilities: lease.capabilities.clone(), - } - } -} - -pub struct LaunchedNode { - pub lease: NodeLease, - pub control: Box, -} - -#[derive(Clone, Debug, Default)] -pub struct StaticNodeLauncher; - -impl StaticNodeLauncher { - pub fn launch_node(&self, lease: &NodeLease, _spec: NodeLaunchSpec) -> LaunchedNode { - let facts = NodeFacts::from_lease(lease); - LaunchedNode { - lease: lease.clone(), - control: Box::new(StaticNodeControl { - facts, - booted: false, - stopped: false, - assigned_roles: Vec::new(), - }), - } - } -} - -struct StaticNodeControl { - facts: NodeFacts, - booted: bool, - stopped: bool, - assigned_roles: Vec, -} - -impl StaticNodeControl { - fn node_id(&self) -> u64 { - self.facts.node_id.0 - } -} - -impl NodeControl for StaticNodeControl { - fn wait_boot_ready(&mut self) -> Result { - if self.stopped { - return Err(NodeControlError::Stopped { - node_id: self.node_id(), - }); - } - self.booted = true; - Ok(self.facts.clone()) - } - - fn wait_cluster_converged(&mut self, expected_alive: usize) -> Result<(), NodeControlError> { - if self.stopped { - return Err(NodeControlError::Stopped { - node_id: self.node_id(), - }); - } - if !self.booted { - return Err(NodeControlError::NotBooted { - node_id: self.node_id(), - }); - } - if expected_alive == 0 { - return Err(NodeControlError::Backend( - "expected_alive must be greater than zero", - )); - } - Ok(()) - } - - fn assign_role(&mut self, assignment: RoleAssignment) -> Result<(), NodeControlError> { - if self.stopped { - return Err(NodeControlError::Stopped { - node_id: self.node_id(), - }); - } - if !self.booted { - return Err(NodeControlError::NotBooted { - node_id: self.node_id(), - }); - } - let role_node_id = assignment.node_id().0; - if role_node_id != self.node_id() { - return Err(NodeControlError::RoleNodeMismatch { - node_id: self.node_id(), - role_node_id, - }); - } - self.assigned_roles.push(assignment); - Ok(()) - } - - fn shutdown(&mut self) -> Result<(), NodeControlError> { - if self.stopped { - return Ok(()); - } - self.stopped = true; - Ok(()) - } -} diff --git a/crates/mvp-system/src/orchestration/engine_builder/mod.rs b/crates/mvp-system/src/orchestration/engine_builder/mod.rs index 47d462e..bb290a0 100644 --- a/crates/mvp-system/src/orchestration/engine_builder/mod.rs +++ b/crates/mvp-system/src/orchestration/engine_builder/mod.rs @@ -7,23 +7,12 @@ //! model onto discovered nodes, assign roles, and return a live cluster handle. //! Workload semantics stay outside this module; see [`WorkloadAdapter`]. -pub mod engine; -pub mod error; -pub mod events; -pub mod launcher; -pub mod model; -pub mod node_image; -pub mod planner; -pub mod pool; -pub mod roles; +pub(crate) mod engine; +pub(crate) mod error; +pub(crate) mod planner; +pub(crate) mod pool; -pub use crate::run_plan::{DTypeFamily, NodeId}; -pub use engine::ClusterBuilder; -pub use events::EngineEvent; -pub use launcher::StaticNodeLauncher; -pub use model::{ModelArtifact, ModelSpec}; -pub use planner::FixedLinearPipelinePlanner; -pub use pool::{NodeCapability, NodeLease, ResourceFacts, StaticPoolProvider}; -pub use roles::RoleKind; - -pub use node_image::{NodeImageSpec, WorkerRuntimeSpec}; +pub(crate) use crate::run_plan::{DTypeFamily, NodeId}; +pub(crate) use engine::{ClusterBuilder, EngineEvent}; +pub(crate) use planner::{FixedLinearPipelinePlanner, RoleKind}; +pub(crate) use pool::{ModelSpec, NodeCapability, NodeFacts, StaticPoolProvider}; diff --git a/crates/mvp-system/src/orchestration/engine_builder/model.rs b/crates/mvp-system/src/orchestration/engine_builder/model.rs deleted file mode 100644 index 522b2ef..0000000 --- a/crates/mvp-system/src/orchestration/engine_builder/model.rs +++ /dev/null @@ -1,67 +0,0 @@ -use crate::run_plan; - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct ModelSpec { - pub model_id: String, - pub artifact: ModelArtifact, - pub tokenizer: run_plan::TokenizerSource, - pub num_layers: u32, - pub hidden_dim: u64, - pub dtype_family: run_plan::DTypeFamily, - pub dtype_width_bytes: u64, - pub max_seq_len: u64, - pub eos_token_id: u32, -} - -impl ModelSpec { - pub fn pipelined_causal_llm( - model_id: impl Into, - artifact: ModelArtifact, - num_layers: u32, - hidden_dim: u64, - dtype_family: run_plan::DTypeFamily, - dtype_width_bytes: u64, - max_seq_len: u64, - eos_token_id: u32, - tokenizer: run_plan::TokenizerSource, - ) -> Self { - Self { - model_id: model_id.into(), - artifact, - tokenizer, - num_layers, - hidden_dim, - dtype_family, - dtype_width_bytes, - max_seq_len, - eos_token_id, - } - } - - pub fn to_run_plan_facts(&self) -> run_plan::ModelFacts { - run_plan::ModelFacts { - model_id: self.model_id.clone(), - gguf_source: self.artifact.to_run_plan_source(), - num_layers: self.num_layers, - hidden_dim: self.hidden_dim, - dtype_family: self.dtype_family, - dtype_width_bytes: self.dtype_width_bytes, - max_seq_len: self.max_seq_len, - eos_token_id: self.eos_token_id, - tokenizer: self.tokenizer.clone(), - } - } -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum ModelArtifact { - TestTinyLlm { path: String }, -} - -impl ModelArtifact { - fn to_run_plan_source(&self) -> run_plan::GgufSource { - match self { - Self::TestTinyLlm { path } => run_plan::GgufSource::LocalPath(path.clone()), - } - } -} diff --git a/crates/mvp-system/src/orchestration/engine_builder/node_image.rs b/crates/mvp-system/src/orchestration/engine_builder/node_image.rs deleted file mode 100644 index 1998728..0000000 --- a/crates/mvp-system/src/orchestration/engine_builder/node_image.rs +++ /dev/null @@ -1,17 +0,0 @@ -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct NodeImageSpec; - -impl NodeImageSpec { - pub fn new(_image: impl Into) -> Self { - Self - } - - pub fn worker_runtime(self, _worker_runtime: WorkerRuntimeSpec) -> Self { - self - } -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum WorkerRuntimeSpec { - DumbProcess, -} diff --git a/crates/mvp-system/src/orchestration/engine_builder/planner.rs b/crates/mvp-system/src/orchestration/engine_builder/planner.rs index c9c43cf..ec31efd 100644 --- a/crates/mvp-system/src/orchestration/engine_builder/planner.rs +++ b/crates/mvp-system/src/orchestration/engine_builder/planner.rs @@ -2,28 +2,74 @@ use std::collections::BTreeSet; use crate::run_plan::{self, NodeId, RunId}; -use super::error::PlanningError; -use super::launcher::NodeFacts; -use super::model::ModelSpec; -use super::pool::NodeCapability; -use super::roles::{CoordinatorAssignment, StageAssignment}; +use super::error::EngineBuildError; +use super::pool::{ModelSpec, NodeCapability, NodeFacts}; +#[derive(Clone, Debug, PartialEq, Eq)] +pub(crate) struct CoordinatorAssignment { + pub node_id: NodeId, +} #[derive(Clone, Debug, PartialEq, Eq)] -pub struct RolePlannerInput { +pub(crate) struct StageAssignment { + pub provision: run_plan::ProvisionStage, +} + +impl StageAssignment { + pub(crate) fn node_id(&self) -> NodeId { + self.provision.node_id + } + + pub(crate) fn stage_index(&self) -> u32 { + self.provision.stage_index + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub(crate) enum RoleAssignment { + Coordinator(CoordinatorAssignment), + StageWorker(StageAssignment), +} + +impl RoleAssignment { + pub(crate) fn node_id(&self) -> NodeId { + match self { + Self::Coordinator(assignment) => assignment.node_id, + Self::StageWorker(assignment) => assignment.node_id(), + } + } + + pub(crate) fn kind(&self) -> RoleKind { + match self { + Self::Coordinator(_) => RoleKind::Coordinator, + Self::StageWorker(assignment) => RoleKind::StageWorker { + stage_index: assignment.stage_index(), + }, + } + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum RoleKind { + Coordinator, + StageWorker { stage_index: u32 }, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub(crate) struct RolePlannerInput { pub run_id: RunId, pub model: ModelSpec, pub nodes: Vec, } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct RoleAssignmentPlan { +pub(crate) struct RoleAssignmentPlan { pub coordinator: CoordinatorAssignment, pub stages: Vec, pub run_plan: run_plan::RunPlan, } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct FixedLinearPipelinePlanner { +pub(crate) struct FixedLinearPipelinePlanner { pub stage_count: u32, pub runtime: run_plan::RuntimeConfig, pub activation_ring: run_plan::RingSpec, @@ -31,33 +77,54 @@ pub struct FixedLinearPipelinePlanner { } impl FixedLinearPipelinePlanner { - pub fn new(stage_count: u32) -> Self { + pub(crate) fn new(stage_count: u32) -> Self { Self { stage_count, - runtime: run_plan::RuntimeConfig::test_default(), - activation_ring: run_plan::RingSpec::test_default_activation(), - token_ring: run_plan::RingSpec::test_default_token(), + runtime: run_plan::RuntimeConfig { + max_tokens: 4, + sampling: run_plan::SamplingPolicy { + temperature_millis: 0, + top_k: 1, + }, + }, + activation_ring: run_plan::RingSpec { + data_capacity: 1 << 20, + alignment: 64, + direction: run_plan::RingDirection::Egress, + host_pinning: run_plan::HostPinning::Pageable, + wake_coalescing: run_plan::WakeCoalescing::PendingBit, + }, + token_ring: run_plan::RingSpec { + data_capacity: 4096, + alignment: 8, + direction: run_plan::RingDirection::Egress, + host_pinning: run_plan::HostPinning::Pageable, + wake_coalescing: run_plan::WakeCoalescing::PendingBit, + }, } } - pub fn runtime(mut self, runtime: run_plan::RuntimeConfig) -> Self { + pub(crate) fn runtime(mut self, runtime: run_plan::RuntimeConfig) -> Self { self.runtime = runtime; self } } impl FixedLinearPipelinePlanner { - pub fn required_node_count(&self) -> usize { + pub(crate) fn required_node_count(&self) -> usize { self.stage_count as usize + 1 } - pub fn plan(&self, input: RolePlannerInput) -> Result { + pub(crate) fn plan( + &self, + input: RolePlannerInput, + ) -> Result { reject_duplicate_nodes(&input.nodes)?; let coordinator = input .nodes .iter() .find(|node| node.capabilities.contains(&NodeCapability::Coordinator)) - .ok_or(PlanningError::NoCoordinatorCandidate)?; + .ok_or(EngineBuildError::NoCoordinatorCandidate)?; let workers = input .nodes .iter() @@ -68,7 +135,7 @@ impl FixedLinearPipelinePlanner { .collect::>(); let required = self.stage_count as usize; if workers.len() < required { - return Err(PlanningError::InsufficientWorkers { + return Err(EngineBuildError::InsufficientWorkers { required, available: workers.len(), }); @@ -95,12 +162,12 @@ impl FixedLinearPipelinePlanner { activation_ring: self.activation_ring, token_ring: self.token_ring, }) - .map_err(|err| PlanningError::ModelRejected(err.kind()))?; + .map_err(|err| EngineBuildError::ModelRejected(err.kind()))?; let mut stages = Vec::with_capacity(self.stage_count as usize); for stage_index in 0..self.stage_count { let provision = run_plan::derive_stage_provision(&run_plan, stage_index) - .map_err(PlanningError::StageProjection)?; + .map_err(EngineBuildError::StageProjection)?; stages.push(StageAssignment { provision }); } @@ -114,11 +181,11 @@ impl FixedLinearPipelinePlanner { } } -fn reject_duplicate_nodes(nodes: &[NodeFacts]) -> Result<(), PlanningError> { +fn reject_duplicate_nodes(nodes: &[NodeFacts]) -> Result<(), EngineBuildError> { let mut seen = BTreeSet::::new(); for node in nodes { if !seen.insert(node.node_id) { - return Err(PlanningError::DuplicateNodeId { + return Err(EngineBuildError::DuplicateNodeId { node_id: node.node_id.0, }); } diff --git a/crates/mvp-system/src/orchestration/engine_builder/pool.rs b/crates/mvp-system/src/orchestration/engine_builder/pool.rs index 83e6ba9..7978b11 100644 --- a/crates/mvp-system/src/orchestration/engine_builder/pool.rs +++ b/crates/mvp-system/src/orchestration/engine_builder/pool.rs @@ -1,71 +1,91 @@ -use std::collections::BTreeSet; +use crate::run_plan::{self, NodeId}; -use crate::run_plan::NodeId; - -use super::error::PoolError; +use super::error::EngineBuildError; #[derive(Clone, Debug, PartialEq, Eq)] -pub struct PoolRequest { - pub min_nodes: usize, +pub(crate) struct NodeFacts { + pub node_id: NodeId, + pub capabilities: Vec, } -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct NodeLease { - pub logical_node_id: NodeId, - pub capabilities: BTreeSet, -} - -impl NodeLease { - pub fn new( - _lease_id: impl Into, - logical_node_id: NodeId, - capabilities: impl IntoIterator, - ) -> Self { - Self { - logical_node_id, - capabilities: capabilities.into_iter().collect(), - } - } - - pub fn resources(self, _expected_resources: ResourceFacts) -> Self { - self - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub enum NodeCapability { +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum NodeCapability { Coordinator, Worker, } -#[derive(Clone, Debug, Default, PartialEq, Eq)] -pub struct ResourceFacts; +#[derive(Clone, Debug, PartialEq, Eq)] +pub(crate) struct StaticPoolProvider { + nodes: Vec, +} -impl ResourceFacts { - pub fn cpu_only(_cpu_cores: u32, _ram_bytes: u64) -> Self { - Self +impl StaticPoolProvider { + pub(crate) fn new(nodes: Vec) -> Self { + Self { nodes } + } + pub(crate) fn acquire_pool( + &self, + min_nodes: usize, + ) -> Result, EngineBuildError> { + if self.nodes.len() < min_nodes { + return Err(EngineBuildError::InsufficientNodes { + requested: min_nodes, + available: self.nodes.len(), + }); + } + Ok(self.nodes.clone()) } } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct StaticPoolProvider { - leases: Vec, +pub(crate) struct ModelSpec { + pub model_id: String, + pub gguf_source: run_plan::GgufSource, + pub tokenizer: run_plan::TokenizerSource, + pub num_layers: u32, + pub hidden_dim: u64, + pub dtype_family: run_plan::DTypeFamily, + pub dtype_width_bytes: u64, + pub max_seq_len: u64, + pub eos_token_id: u32, } -impl StaticPoolProvider { - pub fn new(leases: Vec) -> Self { - Self { leases } - } -} - -impl StaticPoolProvider { - pub fn acquire_pool(&self, request: PoolRequest) -> Result, PoolError> { - if self.leases.len() < request.min_nodes { - return Err(PoolError::InsufficientNodes { - requested: request.min_nodes, - available: self.leases.len(), - }); +impl ModelSpec { + pub(crate) fn pipelined_causal_llm( + model_id: impl Into, + gguf_source: run_plan::GgufSource, + num_layers: u32, + hidden_dim: u64, + dtype_family: run_plan::DTypeFamily, + dtype_width_bytes: u64, + max_seq_len: u64, + eos_token_id: u32, + tokenizer: run_plan::TokenizerSource, + ) -> Self { + Self { + model_id: model_id.into(), + gguf_source, + tokenizer, + num_layers, + hidden_dim, + dtype_family, + dtype_width_bytes, + max_seq_len, + eos_token_id, + } + } + + pub(crate) fn to_run_plan_facts(&self) -> run_plan::ModelFacts { + run_plan::ModelFacts { + model_id: self.model_id.clone(), + gguf_source: self.gguf_source.clone(), + num_layers: self.num_layers, + hidden_dim: self.hidden_dim, + dtype_family: self.dtype_family, + dtype_width_bytes: self.dtype_width_bytes, + max_seq_len: self.max_seq_len, + eos_token_id: self.eos_token_id, + tokenizer: self.tokenizer.clone(), } - Ok(self.leases.clone()) } } diff --git a/crates/mvp-system/src/orchestration/engine_builder/roles.rs b/crates/mvp-system/src/orchestration/engine_builder/roles.rs deleted file mode 100644 index d1d7ccf..0000000 --- a/crates/mvp-system/src/orchestration/engine_builder/roles.rs +++ /dev/null @@ -1,51 +0,0 @@ -use crate::run_plan::{self, NodeId}; - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct CoordinatorAssignment { - pub node_id: NodeId, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct StageAssignment { - pub provision: run_plan::ProvisionStage, -} - -impl StageAssignment { - pub fn node_id(&self) -> NodeId { - self.provision.node_id - } - - pub fn stage_index(&self) -> u32 { - self.provision.stage_index - } -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum RoleAssignment { - Coordinator(CoordinatorAssignment), - StageWorker(StageAssignment), -} - -impl RoleAssignment { - pub fn node_id(&self) -> NodeId { - match self { - Self::Coordinator(assignment) => assignment.node_id, - Self::StageWorker(assignment) => assignment.node_id(), - } - } - - pub fn kind(&self) -> RoleKind { - match self { - Self::Coordinator(_) => RoleKind::Coordinator, - Self::StageWorker(assignment) => RoleKind::StageWorker { - stage_index: assignment.stage_index(), - }, - } - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum RoleKind { - Coordinator, - StageWorker { stage_index: u32 }, -} diff --git a/crates/mvp-system/src/orchestration/mod.rs b/crates/mvp-system/src/orchestration/mod.rs index f773a53..2b31021 100644 --- a/crates/mvp-system/src/orchestration/mod.rs +++ b/crates/mvp-system/src/orchestration/mod.rs @@ -5,14 +5,14 @@ //! adapters live in `provider_adapters` so provider-neutral orchestration logic //! stays separate from local/Docker/VastAI implementation details. -pub mod actor; +pub(crate) mod actor; pub(crate) mod app; -pub mod config; -pub mod distribution_stack; +pub(crate) mod config; +pub(crate) mod distribution_stack; #[cfg(test)] -pub mod engine_builder; +pub(crate) mod engine_builder; -pub mod provider_adapters { - pub mod relay; +pub(crate) mod provider_adapters { + pub(crate) mod relay; pub(super) mod vastai; } diff --git a/crates/mvp-system/src/orchestration/node_provisioning.rs b/crates/mvp-system/src/orchestration/node_provisioning.rs index 7d78599..e502417 100644 --- a/crates/mvp-system/src/orchestration/node_provisioning.rs +++ b/crates/mvp-system/src/orchestration/node_provisioning.rs @@ -6,22 +6,22 @@ pub use ::provisioning::ProviderKind; -pub mod provider_kind { +pub(crate) mod provider_kind { use super::ProviderKind; - pub fn process() -> ProviderKind { + pub(crate) fn process() -> ProviderKind { ProviderKind::new("process") } - pub fn docker() -> ProviderKind { + pub(crate) fn docker() -> ProviderKind { ProviderKind::new("docker") } - pub fn vastai() -> ProviderKind { + pub(crate) fn vastai() -> ProviderKind { ProviderKind::new("vastai") } - pub fn parse_deploy(value: &str) -> Result { + pub(crate) fn parse_deploy(value: &str) -> Result { match value.trim().to_ascii_lowercase().as_str() { "process" | "local_process" | "local-process" => Ok(process()), "docker" | "local_docker" | "local-docker" => Ok(docker()), diff --git a/crates/mvp-system/src/orchestration/provider_adapters/relay.rs b/crates/mvp-system/src/orchestration/provider_adapters/relay.rs index f9bffe5..217cade 100644 --- a/crates/mvp-system/src/orchestration/provider_adapters/relay.rs +++ b/crates/mvp-system/src/orchestration/provider_adapters/relay.rs @@ -10,49 +10,49 @@ use iroh::{RelayMode, RelayUrl}; use serde::{Deserialize, Serialize}; -pub const MVP_IROH_RELAY_MODE_ENV: &str = "MVP_IROH_RELAY_MODE"; -pub const MVP_IROH_RELAY_URL_ENV: &str = "MVP_IROH_RELAY_URL"; -pub const SWACTOR_IROH_RELAY_URL_ENV: &str = "SWACTOR_IROH_RELAY_URL"; +pub(crate) const MVP_IROH_RELAY_MODE_ENV: &str = "MVP_IROH_RELAY_MODE"; +pub(crate) const MVP_IROH_RELAY_URL_ENV: &str = "MVP_IROH_RELAY_URL"; +pub(crate) const SWACTOR_IROH_RELAY_URL_ENV: &str = "SWACTOR_IROH_RELAY_URL"; #[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum RelayPurpose { +pub(crate) enum RelayPurpose { Combined, } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct RelayProvisionRequest { +pub(crate) struct RelayProvisionRequest { pub run_id: u64, pub purpose: RelayPurpose, } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct RelayLeaseId(pub String); +pub(crate) struct RelayLeaseId(pub(crate) String); #[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum RelayProviderKind { +pub(crate) enum RelayProviderKind { LocalShim, Static, } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct RelayEndpoint { +pub(crate) struct RelayEndpoint { pub url: String, pub provider: RelayProviderKind, } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct RelayLease { +pub(crate) struct RelayLease { pub id: RelayLeaseId, pub endpoints: Vec, } #[derive(Clone, Debug)] -pub struct RelayRuntimeConfig { +pub(crate) struct RelayRuntimeConfig { pub mode: RelayMode, pub url: Option, } -pub trait RelayProvider: Send { +pub(crate) trait RelayProvider: Send { fn provision_relay(&mut self, request: RelayProvisionRequest) -> Result; fn relay_mode(&self, lease: &RelayLease) -> Result; @@ -63,7 +63,7 @@ pub trait RelayProvider: Send { } #[derive(Clone, Copy, Debug, Default)] -pub struct LocalShimRelayProvider; +pub(crate) struct LocalShimRelayProvider; impl RelayProvider for LocalShimRelayProvider { fn provision_relay(&mut self, request: RelayProvisionRequest) -> Result { @@ -79,26 +79,26 @@ impl RelayProvider for LocalShimRelayProvider { } #[derive(Clone, Debug)] -pub struct StaticRelayProvider { +pub(crate) struct StaticRelayProvider { url: RelayUrl, } impl StaticRelayProvider { - pub fn new(url: RelayUrl) -> Self { + pub(crate) fn new(url: RelayUrl) -> Self { Self { url } } - pub fn from_url_str(raw: &str) -> Result { + pub(crate) fn from_url_str(raw: &str) -> Result { parse_relay_url(raw).map(Self::new) } - pub fn from_env() -> Result, String> { + pub(crate) fn from_env() -> Result, String> { selected_relay_url_from_env() .map(|url| Self::from_url_str(&url).map(Some)) .unwrap_or(Ok(None)) } - pub fn url(&self) -> String { + pub(crate) fn url(&self) -> String { self.url.to_string() } } @@ -127,13 +127,13 @@ impl RelayProvider for StaticRelayProvider { } } -pub fn relay_runtime_config_from_env(run_id: u64) -> Result { +pub(crate) fn relay_runtime_config_from_env(run_id: u64) -> Result { let mode = relay_mode_setting_from_env(); let url = selected_relay_url_from_env(); relay_runtime_config_from_settings(run_id, mode.as_deref(), url.as_deref()) } -pub fn relay_runtime_config_from_settings( +pub(crate) fn relay_runtime_config_from_settings( run_id: u64, mode: Option<&str>, url: Option<&str>, @@ -150,14 +150,14 @@ pub fn relay_runtime_config_from_settings( } } -pub fn relay_mode_env_value(mode: &RelayMode) -> &'static str { +pub(crate) fn relay_mode_env_value(mode: &RelayMode) -> &'static str { match mode { RelayMode::Disabled => "disabled", _ => "default", } } -pub fn selected_relay_url_from_env() -> Option { +pub(crate) fn selected_relay_url_from_env() -> Option { env_optional(MVP_IROH_RELAY_URL_ENV).or_else(|| env_optional(SWACTOR_IROH_RELAY_URL_ENV)) } diff --git a/crates/mvp-system/src/orchestration/provider_adapters/vastai/mod.rs b/crates/mvp-system/src/orchestration/provider_adapters/vastai/mod.rs index 03dd0e6..34c2d22 100644 --- a/crates/mvp-system/src/orchestration/provider_adapters/vastai/mod.rs +++ b/crates/mvp-system/src/orchestration/provider_adapters/vastai/mod.rs @@ -22,7 +22,7 @@ use crate::provisioning::{ }; #[derive(Clone, Debug)] -pub struct VastAiProvisioningConfig { +pub(crate) struct VastAiProvisioningConfig { pub label_prefix: String, pub disk_gb: u32, pub ssh_user: String, @@ -49,13 +49,13 @@ impl Default for VastAiProvisioningConfig { } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct VastAiSshEndpoint { +pub(crate) struct VastAiSshEndpoint { pub host: String, pub port: u16, pub user: String, } -pub struct VastAiProviderMonitor { +pub(crate) struct VastAiProviderMonitor { runtime: Option, actor: ActorAddress, } @@ -86,7 +86,7 @@ impl Drop for VastAiProviderMonitor { } } -pub trait VastAiLeaseClient: Send { +pub(crate) trait VastAiLeaseClient: Send { fn provision_one(&mut self, request: ProvisionRequest) -> Result; fn plan_first_wave_offers( &mut self, @@ -116,7 +116,7 @@ pub trait VastAiLeaseClient: Send { fn destroy_contract(&mut self, contract_id: u64) -> Result<(), String>; } -pub struct ToolsVastAiLeaseClient { +pub(crate) struct ToolsVastAiLeaseClient { client: swactor_vastai::VastClient, runtime: tokio::runtime::Runtime, planned_offer_pool: Arc>>, @@ -124,7 +124,7 @@ pub struct ToolsVastAiLeaseClient { } impl ToolsVastAiLeaseClient { - pub fn new(client: swactor_vastai::VastClient) -> Result { + pub(crate) fn new(client: swactor_vastai::VastClient) -> Result { let runtime = tokio::runtime::Builder::new_current_thread() .enable_all() .build() @@ -137,7 +137,7 @@ impl ToolsVastAiLeaseClient { }) } - pub fn from_api_key(api_key: impl Into) -> Result { + pub(crate) fn from_api_key(api_key: impl Into) -> Result { Self::new(swactor_vastai::VastClient::new(api_key)) } @@ -612,12 +612,12 @@ fn provider_status_message_has_terminal_failure(message: &str) -> bool { } #[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum BootstrapStopReason { +pub(crate) enum BootstrapStopReason { RuntimeReady, NodeStop, } -pub trait VastAiBootstrapLauncher: Send { +pub(crate) trait VastAiBootstrapLauncher: Send { type Handle: Send; fn start_bootstrap( @@ -942,17 +942,17 @@ fn stop_ssh_child(child: &mut Option) { } #[derive(Clone)] -pub struct SshCommandBootstrapLauncher { +pub(crate) struct SshCommandBootstrapLauncher { ssh_identity: Option, runtime: Arc, } -pub struct SshCommandBootstrapHandle { +pub(crate) struct SshCommandBootstrapHandle { actor: ActorAddress, runtime: Arc, } impl SshCommandBootstrapLauncher { - pub fn new(ssh_identity: Option, runtime: Arc) -> Self { + pub(crate) fn new(ssh_identity: Option, runtime: Arc) -> Self { Self { ssh_identity, runtime, @@ -1127,7 +1127,7 @@ fn next_ssh_backoff(current: Duration) -> Duration { std::cmp::min(current.saturating_mul(2), Duration::from_secs(30)) } -pub struct VastAiProvisioningPlugin +pub(crate) struct VastAiProvisioningPlugin where C: VastAiLeaseClient, B: VastAiBootstrapLauncher, @@ -1158,7 +1158,7 @@ where C: VastAiLeaseClient, B: VastAiBootstrapLauncher, { - pub fn new(client: C, bootstrap: B, config: VastAiProvisioningConfig) -> Self { + pub(crate) fn new(client: C, bootstrap: B, config: VastAiProvisioningConfig) -> Self { Self { client, bootstrap, diff --git a/crates/mvp-system/src/orchestration/provisioning.rs b/crates/mvp-system/src/orchestration/provisioning.rs index aab405a..1e6bfc7 100644 --- a/crates/mvp-system/src/orchestration/provisioning.rs +++ b/crates/mvp-system/src/orchestration/provisioning.rs @@ -24,7 +24,7 @@ pub use ::provisioning::plugin::{ use crate::observability::provisioning_logs::BootstrapDatastreamBridge; -pub struct LocalDockerPlugin { +pub(crate) struct LocalDockerPlugin { container_name_prefix: String, next_handle_id: u64, nodes: BTreeMap, @@ -35,7 +35,7 @@ struct LocalDockerNode { stdin: ChildStdin, } -pub struct LocalProcessPlugin { +pub(crate) struct LocalProcessPlugin { program: PathBuf, next_handle_id: u64, nodes: BTreeMap, @@ -49,7 +49,7 @@ struct LocalProcessNode { } impl LocalProcessPlugin { - pub fn new(program: impl Into) -> Self { + pub(crate) fn new(program: impl Into) -> Self { Self { program: program.into(), next_handle_id: 1, @@ -59,7 +59,7 @@ impl LocalProcessPlugin { } impl LocalDockerPlugin { - pub fn new(container_name_prefix: impl Into) -> Self { + pub(crate) fn new(container_name_prefix: impl Into) -> Self { Self { container_name_prefix: container_name_prefix.into(), next_handle_id: 1, diff --git a/crates/mvp-system/src/orchestration/run_fsm.rs b/crates/mvp-system/src/orchestration/run_fsm.rs index 5b54e27..2999650 100644 --- a/crates/mvp-system/src/orchestration/run_fsm.rs +++ b/crates/mvp-system/src/orchestration/run_fsm.rs @@ -1,46 +1,46 @@ #![allow(dead_code)] #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct RunId(pub u64); +pub(crate) struct RunId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct NodeId(pub u64); +pub(crate) struct NodeId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct StageRef { +pub(crate) struct StageRef { pub stage_index: u32, pub node_id: NodeId, } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct RunPlan { +pub(crate) struct RunPlan { pub run_id: RunId, pub stages: Vec, } impl RunPlan { - pub fn test_linear(run_id: RunId, stages: Vec) -> Self { + pub(crate) fn test_linear(run_id: RunId, stages: Vec) -> Self { Self { run_id, stages } } - pub fn stage_nodes(&self) -> Vec { + pub(crate) fn stage_nodes(&self) -> Vec { self.stages.iter().map(|stage| stage.node_id).collect() } } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct RunConfig { +pub(crate) struct RunConfig { pub run_id: RunId, pub max_tokens: u64, pub prompt: Vec, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct SamplingData { +pub(crate) struct SamplingData { pub source_sequence: u64, } #[derive(Clone, Debug, PartialEq, Eq)] -pub enum TokenObjectPayload { +pub(crate) enum TokenObjectPayload { Prompt { tokens: Vec, }, @@ -51,24 +51,24 @@ pub enum TokenObjectPayload { } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct TokenObjectInjection { +pub(crate) struct TokenObjectInjection { pub sequence: u64, pub payload: TokenObjectPayload, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum StageFaultReason { +pub(crate) enum StageFaultReason { WorkerCrashed, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum EndpointKind { +pub(crate) enum EndpointKind { TokenIn, TokenOut, } #[derive(Clone, Debug, PartialEq, Eq)] -pub enum RunEvent { +pub(crate) enum RunEvent { PoolReady { nodes: Vec, }, @@ -108,7 +108,7 @@ pub enum RunEvent { } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum RunFaultReason { +pub(crate) enum RunFaultReason { StageFault { stage_index: u32, reason: StageFaultReason, @@ -125,7 +125,7 @@ pub enum RunFaultReason { } #[derive(Clone, Debug, PartialEq, Eq)] -pub enum LifecycleEvent { +pub(crate) enum LifecycleEvent { RunRejected { run_id: RunId, reason: RunFaultReason, @@ -146,14 +146,14 @@ pub enum LifecycleEvent { } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct StageProvision { +pub(crate) struct StageProvision { pub run_id: RunId, pub stage_index: u32, pub node_id: NodeId, } #[derive(Clone, Debug, PartialEq, Eq)] -pub enum RunCommand { +pub(crate) enum RunCommand { ProvisionStage { provision: StageProvision, }, @@ -178,7 +178,7 @@ pub enum RunCommand { pub type OrchestratorHarness = OrchestratorRun; -pub struct OrchestratorRun { +pub(crate) struct OrchestratorRun { config: RunConfig, plan: Option, pool_ready: bool, @@ -197,7 +197,7 @@ pub struct OrchestratorRun { } impl OrchestratorRun { - pub fn new(config: RunConfig) -> Self { + pub(crate) fn new(config: RunConfig) -> Self { Self { config, plan: None, @@ -217,7 +217,7 @@ impl OrchestratorRun { } } - pub fn observe(&mut self, event: RunEvent) { + pub(crate) fn observe(&mut self, event: RunEvent) { match event { RunEvent::PoolReady { .. } => { self.pool_ready = true; @@ -301,17 +301,17 @@ impl OrchestratorRun { } } - pub fn advance_time_ms(&mut self, _delta: u64) {} + pub(crate) fn advance_time_ms(&mut self, _delta: u64) {} - pub fn commands(&self) -> &[RunCommand] { + pub(crate) fn commands(&self) -> &[RunCommand] { &self.commands } - pub fn events(&self) -> &[LifecycleEvent] { + pub(crate) fn events(&self) -> &[LifecycleEvent] { &self.events } - pub fn injected_sequences(&self) -> Vec { + pub(crate) fn injected_sequences(&self) -> Vec { self.injected_sequences.clone() } diff --git a/crates/mvp-system/src/orchestration/run_plan.rs b/crates/mvp-system/src/orchestration/run_plan.rs index 471ea31..672f642 100644 --- a/crates/mvp-system/src/orchestration/run_plan.rs +++ b/crates/mvp-system/src/orchestration/run_plan.rs @@ -1,10 +1,10 @@ #![allow(dead_code)] -pub const MO01_HEADER_BYTES: u64 = 40; +pub(crate) const MO01_HEADER_BYTES: u64 = 40; const TOKEN_ID_WIDTH_BYTES: u32 = 4; #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct RunId(pub u64); +pub(crate) struct RunId(pub(crate) u64); impl From for RunId { fn from(value: u64) -> Self { @@ -13,22 +13,22 @@ impl From for RunId { } #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct NodeId(pub u64); +pub(crate) struct NodeId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct EdgeId(pub u64); +pub(crate) struct EdgeId(pub(crate) u64); #[derive(Clone, Debug, PartialEq, Eq)] -pub struct EdgeAllocator { +pub(crate) struct EdgeAllocator { next: u64, } impl EdgeAllocator { - pub fn new() -> Self { + pub(crate) fn new() -> Self { Self { next: 1 } } - pub fn alloc(&mut self) -> EdgeId { + pub(crate) fn alloc(&mut self) -> EdgeId { let edge_id = EdgeId(self.next); self.next += 1; edge_id @@ -42,12 +42,12 @@ impl Default for EdgeAllocator { } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum DTypeFamily { +pub(crate) enum DTypeFamily { BFloat, } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct ModelFacts { +pub(crate) struct ModelFacts { pub model_id: String, pub gguf_source: GgufSource, pub num_layers: u32, @@ -60,7 +60,7 @@ pub struct ModelFacts { } #[derive(Clone, Debug, PartialEq, Eq, serde::Serialize, serde::Deserialize)] -pub enum GgufSource { +pub(crate) enum GgufSource { LocalPath(String), HuggingFaceGguf { repo: String, @@ -70,35 +70,25 @@ pub enum GgufSource { } #[derive(Clone, Debug, PartialEq, Eq, serde::Serialize, serde::Deserialize)] -pub enum TokenizerSource { +pub(crate) enum TokenizerSource { EmbeddedGguf, LocalPath(String), } -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum PromptSource { - Inline(String), -} - #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct SamplingPolicy { +pub(crate) struct SamplingPolicy { pub temperature_millis: u32, pub top_k: u32, } -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum TokenOutputPolicy { - EmitAll, -} - #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct RoleId(pub u64); +pub(crate) struct RoleId(pub(crate) u64); #[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct PortId(pub String); +pub(crate) struct PortId(pub(crate) String); #[derive(Clone, Debug, PartialEq, Eq)] -pub struct GgufModelPlan { +pub(crate) struct GgufModelPlan { pub model_id: String, pub gguf_source: GgufSource, pub num_layers: u32, @@ -111,65 +101,42 @@ pub struct GgufModelPlan { } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct RuntimePlan { - pub prompt: PromptSource, - pub sampling: SamplingPolicy, - pub token_output_policy: TokenOutputPolicy, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct RuntimeConfig { +pub(crate) struct RuntimeConfig { pub max_tokens: u32, - pub prompt: PromptSource, pub sampling: SamplingPolicy, - pub token_output_policy: TokenOutputPolicy, -} - -impl RuntimeConfig { - pub fn test_default() -> Self { - Self { - max_tokens: 4, - prompt: PromptSource::Inline("test prompt".to_owned()), - sampling: SamplingPolicy { - temperature_millis: 0, - top_k: 1, - }, - token_output_policy: TokenOutputPolicy::EmitAll, - } - } } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct StagePlacement { +pub(crate) struct StagePlacement { pub stage_index: u32, pub node_id: NodeId, } #[derive(Clone, Debug, PartialEq, Eq)] -pub enum PlacementInput { +pub(crate) enum PlacementInput { FixedLinear(Vec), } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum RingDirection { +pub(crate) enum RingDirection { Ingress, Egress, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum HostPinning { +pub(crate) enum HostPinning { Pageable, PinnedRequired, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum WakeCoalescing { +pub(crate) enum WakeCoalescing { PendingBit, ReadySet, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct RingSpec { +pub(crate) struct RingSpec { pub data_capacity: u64, pub alignment: u32, pub direction: RingDirection, @@ -177,30 +144,8 @@ pub struct RingSpec { pub wake_coalescing: WakeCoalescing, } -impl RingSpec { - pub fn test_default_activation() -> Self { - Self { - data_capacity: 1 << 20, - alignment: 64, - direction: RingDirection::Egress, - host_pinning: HostPinning::Pageable, - wake_coalescing: WakeCoalescing::PendingBit, - } - } - - pub fn test_default_token() -> Self { - Self { - data_capacity: 4096, - alignment: 8, - direction: RingDirection::Egress, - host_pinning: HostPinning::Pageable, - wake_coalescing: WakeCoalescing::PendingBit, - } - } -} - #[derive(Clone, Debug, PartialEq, Eq)] -pub struct PlannerInput { +pub(crate) struct PlannerInput { pub run_id: RunId, pub orchestrator_node_id: NodeId, pub model: ModelFacts, @@ -213,14 +158,14 @@ pub struct PlannerInput { } #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub enum EdgeKind { +pub(crate) enum EdgeKind { TokenIn, Activation, TokenOut, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum ObjectKind { +pub(crate) enum ObjectKind { Token, Activation, Weight, @@ -228,7 +173,7 @@ pub enum ObjectKind { } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum ShapeRule { +pub(crate) enum ShapeRule { TokenIds, ActivationRows { max_seq_len: u32, hidden_dim: u32 }, WeightTensor, @@ -236,17 +181,17 @@ pub enum ShapeRule { } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum LayoutRule { +pub(crate) enum LayoutRule { Contiguous, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum SequencePolicy { +pub(crate) enum SequencePolicy { Ordered, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct ObjectSpec { +pub(crate) struct ObjectSpec { pub kind: ObjectKind, pub max_extent: u64, pub dtype_family: DTypeFamily, @@ -258,13 +203,13 @@ pub struct ObjectSpec { } #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub enum EdgeEndpoint { +pub(crate) enum EdgeEndpoint { Orchestrator { node_id: NodeId }, Stage { node_id: NodeId, stage_index: u32 }, } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct EdgePlan { +pub(crate) struct EdgePlan { pub run_id: RunId, pub edge_id: EdgeId, pub kind: EdgeKind, @@ -275,7 +220,7 @@ pub struct EdgePlan { } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct StagePlan { +pub(crate) struct StagePlan { pub run_id: RunId, pub stage_index: u32, pub stage_count: u32, @@ -288,17 +233,17 @@ pub struct StagePlan { } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct RunPlan { +pub(crate) struct RunPlan { pub run_id: RunId, pub model: GgufModelPlan, - pub runtime: RuntimePlan, + pub sampling: SamplingPolicy, pub stages: Vec, pub edges: Vec, pub max_tokens: u32, } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct InboundEdgeProvision { +pub(crate) struct InboundEdgeProvision { pub edge_id: EdgeId, pub kind: EdgeKind, pub object_spec: ObjectSpec, @@ -306,7 +251,7 @@ pub struct InboundEdgeProvision { } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct OutboundEdgeProvision { +pub(crate) struct OutboundEdgeProvision { pub edge_id: EdgeId, pub kind: EdgeKind, pub consumer_node_id: NodeId, @@ -315,7 +260,7 @@ pub struct OutboundEdgeProvision { } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct StageModelFacts { +pub(crate) struct StageModelFacts { pub model_id: String, pub hidden_dim: u32, pub dtype_family: DTypeFamily, @@ -324,7 +269,7 @@ pub struct StageModelFacts { } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct StageRuntimeFacts { +pub(crate) struct StageRuntimeFacts { pub role_id: RoleId, pub input_port: PortId, pub output_port: PortId, @@ -332,7 +277,7 @@ pub struct StageRuntimeFacts { } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct ProvisionStage { +pub(crate) struct ProvisionStage { pub run_id: RunId, pub node_id: NodeId, pub stage_index: u32, @@ -348,7 +293,7 @@ pub struct ProvisionStage { } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum PlanRejectionKind { +pub(crate) enum PlanRejectionKind { UnknownNode, DuplicateStageAssignment, MissingStage, @@ -361,31 +306,26 @@ pub enum PlanRejectionKind { } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct PlanRejection { +pub(crate) struct PlanRejection { kind: PlanRejectionKind, } impl PlanRejection { - pub fn kind(&self) -> PlanRejectionKind { + pub(crate) fn kind(&self) -> PlanRejectionKind { self.kind } } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum ProjectionRejection { +pub(crate) enum ProjectionRejection { UnknownStage, MissingEdge, } -pub fn plan_run(input: PlannerInput) -> Result { +pub(crate) fn plan_run(input: PlannerInput) -> Result { validate_global_input(&input)?; let placements = validated_placements(&input)?; let model = model_plan(&input.model)?; - let runtime = RuntimePlan { - prompt: input.runtime.prompt.clone(), - sampling: input.runtime.sampling, - token_output_policy: input.runtime.token_output_policy, - }; let max_tokens = input.runtime.max_tokens; let gguf_source = model.gguf_source.clone(); let hidden_dim = model.hidden_dim; @@ -526,14 +466,14 @@ pub fn plan_run(input: PlannerInput) -> Result { Ok(RunPlan { run_id: input.run_id, model, - runtime, + sampling: input.runtime.sampling, stages, edges, max_tokens, }) } -pub fn derive_stage_provision( +pub(crate) fn derive_stage_provision( plan: &RunPlan, stage_index: u32, ) -> Result { @@ -586,7 +526,7 @@ pub fn derive_stage_provision( input_port: PortId("input".to_owned()), output_port: PortId("output".to_owned()), sampling: if stage.stage_index + 1 == stage.stage_count { - Some(plan.runtime.sampling) + Some(plan.sampling) } else { None }, diff --git a/crates/mvp-system/src/prompt/mod.rs b/crates/mvp-system/src/prompt/mod.rs index bfe8e3d..d0a0357 100644 --- a/crates/mvp-system/src/prompt/mod.rs +++ b/crates/mvp-system/src/prompt/mod.rs @@ -1,3 +1,3 @@ //! MVP prompt protocol public surface. -pub mod rpc; +pub(crate) mod rpc; diff --git a/crates/mvp-system/src/prompt/rpc.rs b/crates/mvp-system/src/prompt/rpc.rs index 96768dd..a17afdb 100644 --- a/crates/mvp-system/src/prompt/rpc.rs +++ b/crates/mvp-system/src/prompt/rpc.rs @@ -6,14 +6,14 @@ use swactor_transport::{CodecRegistry, NetworkMessage}; use crate::transport::json_codec::JsonCodec; #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct SubmitPrompt { +pub(crate) struct SubmitPrompt { pub request_id: u64, pub prompt_text: String, pub max_tokens: u32, } impl SubmitPrompt { - pub fn with_defaults(mut self, max_tokens: u32) -> Self { + pub(crate) fn with_defaults(mut self, max_tokens: u32) -> Self { if self.max_tokens == 0 { self.max_tokens = max_tokens; } @@ -23,7 +23,7 @@ impl SubmitPrompt { #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] -pub enum PromptEvent { +pub(crate) enum PromptEvent { TextDelta { request_id: u64, text: String, @@ -41,7 +41,7 @@ pub enum PromptEvent { } impl PromptEvent { - pub fn request_id(&self) -> u64 { + pub(crate) fn request_id(&self) -> u64 { match self { Self::TextDelta { request_id, .. } | Self::Done { request_id, .. } @@ -49,7 +49,7 @@ impl PromptEvent { } } - pub fn is_terminal(&self) -> bool { + pub(crate) fn is_terminal(&self) -> bool { matches!(self, Self::Done { .. } | Self::Fault { .. }) } } @@ -61,7 +61,7 @@ impl NetworkMessage for PromptEvent { } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub enum TokenizerEvent { +pub(crate) enum TokenizerEvent { PromptEncoded { request_id: u64, tokens: Vec }, TokensDecoded { request_id: u64, text: String }, Fault { request_id: u64, error: String }, @@ -73,12 +73,15 @@ impl NetworkMessage for TokenizerEvent { } } -pub fn register_codecs(registry: &mut CodecRegistry) { +pub(crate) fn register_codecs(registry: &mut CodecRegistry) { registry.register::(JsonCodec::::default()); registry.register::(JsonCodec::::default()); } -pub fn write_json_line(writer: &mut impl Write, value: &T) -> Result<(), String> { +pub(crate) fn write_json_line( + writer: &mut impl Write, + value: &T, +) -> Result<(), String> { serde_json::to_writer(&mut *writer, value).map_err(|e| format!("serialize JSON line: {e}"))?; writer .write_all(b"\n") @@ -86,7 +89,9 @@ pub fn write_json_line(writer: &mut impl Write, value: &T) -> Resu writer.flush().map_err(|e| format!("flush JSON line: {e}")) } -pub fn read_submit_prompt(reader: &mut impl BufRead) -> Result, String> { +pub(crate) fn read_submit_prompt( + reader: &mut impl BufRead, +) -> Result, String> { let mut line = String::new(); let n = reader .read_line(&mut line) diff --git a/crates/mvp-system/src/staging/actor.rs b/crates/mvp-system/src/staging/actor.rs index 1319b1f..d0b1d1f 100644 --- a/crates/mvp-system/src/staging/actor.rs +++ b/crates/mvp-system/src/staging/actor.rs @@ -4,13 +4,13 @@ use swactor::runtime::Ctx; use crate::staging::control as core; #[derive(Clone, Debug, PartialEq, Eq)] -pub enum StageControllerMsg { +pub(crate) enum StageControllerMsg { Observe(core::StageEvent), Snapshot { reply_to: ActorAddress }, } #[derive(Clone, Debug, PartialEq, Eq)] -pub enum StageControllerReport { +pub(crate) enum StageControllerReport { Command(core::StageCommand), Lifecycle(core::StageLifecycleEvent), Snapshot { @@ -19,7 +19,7 @@ pub enum StageControllerReport { }, } -pub struct StageControllerActor { +pub(crate) struct StageControllerActor { core: core::StageController, report_to: Option, command_cursor: usize, @@ -27,7 +27,7 @@ pub struct StageControllerActor { } impl StageControllerActor { - pub fn new(local_node_id: core::NodeId, report_to: Option) -> Self { + pub(crate) fn new(local_node_id: core::NodeId, report_to: Option) -> Self { Self { core: core::StageController::new(local_node_id), report_to, diff --git a/crates/mvp-system/src/staging/control.rs b/crates/mvp-system/src/staging/control.rs index 0674179..7aa700f 100644 --- a/crates/mvp-system/src/staging/control.rs +++ b/crates/mvp-system/src/staging/control.rs @@ -4,43 +4,43 @@ use crate::gguf_shard::StageShardPlan; use crate::run_plan::{GgufSource, TokenizerSource}; #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct RunId(pub u64); +pub(crate) struct RunId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct NodeId(pub u64); +pub(crate) struct NodeId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct EdgeId(pub u64); +pub(crate) struct EdgeId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct ObjectId(pub u64); +pub(crate) struct ObjectId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct StepId(pub u64); +pub(crate) struct StepId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct DeviceHandle { +pub(crate) struct DeviceHandle { pub generation: u64, pub id: u64, } impl DeviceHandle { - pub fn new_current(id: u64) -> Self { + pub(crate) fn new_current(id: u64) -> Self { Self { generation: 1, id } } } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct LayerRange { +pub(crate) struct LayerRange { pub start: u32, pub end_exclusive: u32, } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct WeightSource { +pub(crate) struct WeightSource { pub model_id: String, pub gguf_source: GgufSource, pub tokenizer: TokenizerSource, } impl WeightSource { - pub fn new( + pub(crate) fn new( model_id: impl Into, gguf_source: GgufSource, tokenizer: TokenizerSource, @@ -52,7 +52,7 @@ impl WeightSource { } } - pub fn embedded_gguf(model_id: impl Into, path: impl Into) -> Self { + pub(crate) fn embedded_gguf(model_id: impl Into, path: impl Into) -> Self { Self::new( model_id, GgufSource::LocalPath(path.into()), @@ -62,26 +62,26 @@ impl WeightSource { } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum EdgeDirection { +pub(crate) enum EdgeDirection { Inbound, Outbound, } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct EdgeProvision { +pub(crate) struct EdgeProvision { pub edge_id: EdgeId, pub direction: EdgeDirection, } impl EdgeProvision { - pub fn inbound(edge_id: EdgeId) -> Self { + pub(crate) fn inbound(edge_id: EdgeId) -> Self { Self { edge_id, direction: EdgeDirection::Inbound, } } - pub fn outbound(edge_id: EdgeId) -> Self { + pub(crate) fn outbound(edge_id: EdgeId) -> Self { Self { edge_id, direction: EdgeDirection::Outbound, @@ -90,7 +90,7 @@ impl EdgeProvision { } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct ProvisionStage { +pub(crate) struct ProvisionStage { pub run_id: RunId, pub authorized_orchestrator: NodeId, pub node_id: NodeId, @@ -104,7 +104,7 @@ pub struct ProvisionStage { } #[derive(Clone, Debug, PartialEq, Eq)] -pub enum StageEvent { +pub(crate) enum StageEvent { ProvisionStage { from: NodeId, provision: ProvisionStage, @@ -158,7 +158,7 @@ pub enum StageEvent { } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum StageFaultReason { +pub(crate) enum StageFaultReason { UnauthorizedProvision, SequenceViolation, WorkerCrashed, @@ -169,7 +169,7 @@ pub enum StageFaultReason { } #[derive(Clone, Debug, PartialEq, Eq)] -pub enum StageLifecycleEvent { +pub(crate) enum StageLifecycleEvent { StageReady { run_id: RunId, stage_index: u32, @@ -191,7 +191,7 @@ pub enum StageLifecycleEvent { } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct StepInput { +pub(crate) struct StepInput { pub edge_id: EdgeId, pub object_id: ObjectId, pub sequence: u64, @@ -199,20 +199,20 @@ pub struct StepInput { } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct OutputBinding { +pub(crate) struct OutputBinding { pub edge_id: EdgeId, pub sequence: u64, } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct ExecuteStep { +pub(crate) struct ExecuteStep { pub step_id: StepId, pub input: StepInput, pub outputs: Vec, } #[derive(Clone, Debug, PartialEq, Eq)] -pub enum StageCommand { +pub(crate) enum StageCommand { EstablishInboundEdge { edge_id: EdgeId, }, @@ -247,7 +247,7 @@ pub enum StageCommand { pub type StageControllerHarness = StageController; -pub struct StageController { +pub(crate) struct StageController { local_node_id: NodeId, provision: Option, worker_ready: bool, @@ -271,7 +271,7 @@ pub struct StageController { } impl StageController { - pub fn new(local_node_id: NodeId) -> Self { + pub(crate) fn new(local_node_id: NodeId) -> Self { Self { local_node_id, provision: None, @@ -296,7 +296,7 @@ impl StageController { } } - pub fn observe(&mut self, event: StageEvent) { + pub(crate) fn observe(&mut self, event: StageEvent) { match event { StageEvent::ProvisionStage { from, provision } => self.provision(from, provision), StageEvent::WorkerReady => self.worker_ready = true, @@ -343,11 +343,11 @@ impl StageController { } } - pub fn commands(&self) -> &[StageCommand] { + pub(crate) fn commands(&self) -> &[StageCommand] { &self.commands } - pub fn events(&self) -> &[StageLifecycleEvent] { + pub(crate) fn events(&self) -> &[StageLifecycleEvent] { &self.events } diff --git a/crates/mvp-system/src/staging/gguf_common.rs b/crates/mvp-system/src/staging/gguf_common.rs new file mode 100644 index 0000000..57159cb --- /dev/null +++ b/crates/mvp-system/src/staging/gguf_common.rs @@ -0,0 +1,126 @@ +use std::io::Read; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum GgufValueType { + Uint8, + Int8, + Uint16, + Int16, + Uint32, + Int32, + Float32, + Bool, + String, + Array, + Uint64, + Int64, + Float64, +} + +impl GgufValueType { + pub(crate) fn read(reader: &mut R, value_label: &str) -> 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)?; + VALUE_TYPES + .get(raw as usize) + .copied() + .ok_or_else(|| format!("unsupported {value_label} {raw}")) + } + + pub(crate) fn fixed_width(self) -> Option { + match self { + Self::Uint8 | Self::Int8 | Self::Bool => Some(1), + Self::Uint16 | Self::Int16 => Some(2), + Self::Uint32 | Self::Int32 | Self::Float32 => Some(4), + Self::Uint64 | Self::Int64 | Self::Float64 => Some(8), + Self::String | Self::Array => None, + } + } + + pub(crate) fn is_integer(self) -> bool { + matches!( + self, + Self::Uint8 + | Self::Int8 + | Self::Uint16 + | Self::Int16 + | Self::Uint32 + | Self::Int32 + | Self::Uint64 + | Self::Int64 + ) + } +} + +pub(crate) fn read_integer_value( + reader: &mut R, + value_type: GgufValueType, + type_error: impl FnOnce(GgufValueType) -> String, + negative_error: impl Fn(i64) -> String, +) -> Result { + match value_type { + GgufValueType::Uint8 => read_u8(reader).map(u64::from), + GgufValueType::Int8 => { + read_i8(reader).and_then(|value| non_negative_i64_to_u64(value, negative_error)) + } + GgufValueType::Uint16 => read_u16(reader).map(u64::from), + GgufValueType::Int16 => read_i16(reader) + .and_then(|value| non_negative_i64_to_u64(i64::from(value), negative_error)), + GgufValueType::Uint32 => read_u32(reader).map(u64::from), + GgufValueType::Int32 => read_i32(reader) + .and_then(|value| non_negative_i64_to_u64(i64::from(value), negative_error)), + GgufValueType::Uint64 => read_u64(reader), + GgufValueType::Int64 => { + read_i64(reader).and_then(|value| non_negative_i64_to_u64(value, negative_error)) + } + other => Err(type_error(other)), + } +} + +fn non_negative_i64_to_u64( + value: i64, + negative_error: impl Fn(i64) -> String, +) -> Result { + u64::try_from(value).map_err(|_| negative_error(value)) +} + +macro_rules! read_le { + ($name:ident, $ret:ty, $len:expr, |$bytes:ident| $body:expr) => { + pub(crate) fn $name(reader: &mut R) -> Result<$ret, String> { + let mut bytes = [0; $len]; + reader.read_exact(&mut bytes).map_err(|e| { + format!( + "read {}: {e}", + stringify!($name).trim_start_matches("read_") + ) + })?; + Ok({ + let $bytes = bytes; + $body + }) + } + }; +} + +read_le!(read_u8, u8, 1, |bytes| bytes[0]); +read_le!(read_i8, i64, 1, |bytes| i8::from_le_bytes(bytes) as i64); +read_le!(read_u16, u16, 2, |bytes| u16::from_le_bytes(bytes)); +read_le!(read_i16, i16, 2, |bytes| i16::from_le_bytes(bytes)); +read_le!(read_u32, u32, 4, |bytes| u32::from_le_bytes(bytes)); +read_le!(read_i32, i32, 4, |bytes| i32::from_le_bytes(bytes)); +read_le!(read_u64, u64, 8, |bytes| u64::from_le_bytes(bytes)); +read_le!(read_i64, i64, 8, |bytes| i64::from_le_bytes(bytes)); diff --git a/crates/mvp-system/src/staging/gguf_metadata.rs b/crates/mvp-system/src/staging/gguf_metadata.rs index e26fc64..142f66c 100644 --- a/crates/mvp-system/src/staging/gguf_metadata.rs +++ b/crates/mvp-system/src/staging/gguf_metadata.rs @@ -3,6 +3,7 @@ use std::fs::File; use std::io::{Read, Seek, SeekFrom}; use std::path::Path; +use crate::gguf_common::{GgufValueType, read_integer_value, read_u32, read_u64}; use crate::run_plan::{self, DTypeFamily, GgufSource, TokenizerSource}; const GGUF_MAGIC: &[u8; 4] = b"GGUF"; @@ -12,7 +13,7 @@ const MAX_METADATA_STRING_BYTES: u64 = 16 * 1024 * 1024; const MAX_METADATA_KEY_BYTES: u64 = 1024 * 1024; #[derive(Clone, Debug, PartialEq, Eq)] -pub struct GgufPlanningMetadata { +pub(crate) struct GgufPlanningMetadata { pub version: u32, pub architecture: String, pub name: Option, @@ -23,7 +24,7 @@ pub struct GgufPlanningMetadata { } impl GgufPlanningMetadata { - pub fn to_model_facts( + pub(crate) fn to_model_facts( &self, model_id: impl Into, gguf_source: GgufSource, @@ -54,7 +55,7 @@ impl GgufPlanningMetadata { } } -pub fn read_gguf_planning_metadata(path: &Path) -> Result { +pub(crate) fn read_gguf_planning_metadata(path: &Path) -> Result { let file = File::open(path).map_err(|e| format!("open GGUF metadata {}: {e}", path.display()))?; read_gguf_planning_metadata_from_reader(file) @@ -86,7 +87,7 @@ where for _ in 0..metadata_count { let key = read_gguf_string(&mut reader, MAX_METADATA_KEY_BYTES)?; - let value_type = GgufValueType::read(&mut reader)?; + let value_type = GgufValueType::read(&mut reader, "GGUF metadata value type")?; match value_type { GgufValueType::String if key == "general.architecture" || key == "general.name" => { strings.insert( @@ -98,7 +99,12 @@ where skip_gguf_string(&mut reader)?; } value_type if value_type.is_integer() => { - let value = read_integer_value(&mut reader, value_type)?; + let value = read_integer_value( + &mut reader, + value_type, + |other| format!("GGUF value type {other:?} is not an integer"), + |value| format!("negative integer metadata value {value}"), + )?; if key.ends_with(".block_count") || key.ends_with(".embedding_length") || key.ends_with(".context_length") @@ -155,94 +161,6 @@ fn required_u32(map: &BTreeMap, key: &str, label: &str) -> Result(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)?; - VALUE_TYPES - .get(raw as usize) - .copied() - .ok_or_else(|| format!("unsupported GGUF metadata value type {raw}")) - } - - fn is_integer(self) -> bool { - matches!( - self, - Self::Uint8 - | Self::Int8 - | Self::Uint16 - | Self::Int16 - | Self::Uint32 - | Self::Int32 - | Self::Uint64 - | Self::Int64 - ) - } - - fn fixed_width(self) -> Option { - match self { - Self::Uint8 | Self::Int8 | Self::Bool => Some(1), - Self::Uint16 | Self::Int16 => Some(2), - Self::Uint32 | Self::Int32 | Self::Float32 => Some(4), - Self::Uint64 | Self::Int64 | Self::Float64 => Some(8), - Self::String | Self::Array => None, - } - } -} - -fn read_integer_value(reader: &mut R, value_type: GgufValueType) -> Result { - match value_type { - GgufValueType::Uint8 => read_u8(reader).map(u64::from), - GgufValueType::Int8 => read_i8(reader).and_then(non_negative_i64_to_u64), - GgufValueType::Uint16 => read_u16(reader).map(u64::from), - GgufValueType::Int16 => { - read_i16(reader).and_then(|v| non_negative_i64_to_u64(i64::from(v))) - } - GgufValueType::Uint32 => read_u32(reader).map(u64::from), - GgufValueType::Int32 => { - read_i32(reader).and_then(|v| non_negative_i64_to_u64(i64::from(v))) - } - GgufValueType::Uint64 => read_u64(reader), - GgufValueType::Int64 => read_i64(reader).and_then(non_negative_i64_to_u64), - other => Err(format!("GGUF value type {other:?} is not an integer")), - } -} - -fn non_negative_i64_to_u64(value: i64) -> Result { - u64::try_from(value).map_err(|_| format!("negative integer metadata value {value}")) -} - fn skip_scalar(reader: &mut R, value_type: GgufValueType) -> Result<(), String> { match value_type { GgufValueType::String => skip_gguf_string(reader), @@ -252,7 +170,7 @@ fn skip_scalar(reader: &mut R, value_type: GgufValueType) -> Res } fn skip_array(reader: &mut R) -> Result<(), String> { - let element_type = GgufValueType::read(reader)?; + let element_type = GgufValueType::read(reader, "GGUF metadata value type")?; let len = read_u64(reader)?; match element_type { GgufValueType::String => { @@ -308,63 +226,3 @@ fn skip_bytes(reader: &mut R, mut bytes: u64) -> Result<(), String> { } Ok(()) } - -fn read_u8(reader: &mut R) -> Result { - let mut bytes = [0; 1]; - reader - .read_exact(&mut bytes) - .map_err(|e| format!("read u8: {e}"))?; - Ok(bytes[0]) -} - -fn read_i8(reader: &mut R) -> Result { - read_u8(reader).map(|value| i8::from_le_bytes([value]) as i64) -} - -fn read_u16(reader: &mut R) -> Result { - let mut bytes = [0; 2]; - reader - .read_exact(&mut bytes) - .map_err(|e| format!("read u16: {e}"))?; - Ok(u16::from_le_bytes(bytes)) -} - -fn read_i16(reader: &mut R) -> Result { - let mut bytes = [0; 2]; - reader - .read_exact(&mut bytes) - .map_err(|e| format!("read i16: {e}"))?; - Ok(i16::from_le_bytes(bytes)) -} - -fn read_u32(reader: &mut R) -> Result { - let mut bytes = [0; 4]; - reader - .read_exact(&mut bytes) - .map_err(|e| format!("read u32: {e}"))?; - Ok(u32::from_le_bytes(bytes)) -} - -fn read_i32(reader: &mut R) -> Result { - let mut bytes = [0; 4]; - reader - .read_exact(&mut bytes) - .map_err(|e| format!("read i32: {e}"))?; - Ok(i32::from_le_bytes(bytes)) -} - -fn read_u64(reader: &mut R) -> Result { - let mut bytes = [0; 8]; - reader - .read_exact(&mut bytes) - .map_err(|e| format!("read u64: {e}"))?; - Ok(u64::from_le_bytes(bytes)) -} - -fn read_i64(reader: &mut R) -> Result { - let mut bytes = [0; 8]; - reader - .read_exact(&mut bytes) - .map_err(|e| format!("read i64: {e}"))?; - Ok(i64::from_le_bytes(bytes)) -} diff --git a/crates/mvp-system/src/staging/gguf_shard.rs b/crates/mvp-system/src/staging/gguf_shard.rs index 1eb534c..6e8d62f 100644 --- a/crates/mvp-system/src/staging/gguf_shard.rs +++ b/crates/mvp-system/src/staging/gguf_shard.rs @@ -2,6 +2,7 @@ use std::fs::File; use std::io::{Read, Seek, SeekFrom, Write}; use std::path::{Path, PathBuf}; +use crate::gguf_common::{GgufValueType, read_integer_value, read_u32, read_u64}; use serde::{Deserialize, Serialize}; use crate::run_plan::GgufSource; @@ -13,19 +14,19 @@ const MAX_STRING_BYTES: u64 = 64 * 1024 * 1024; const STAGE_SHARD_CACHE_FORMAT_VERSION: &str = "stage-shard-cache-v2"; #[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct ByteRange { +pub(crate) struct ByteRange { pub start: u64, pub len: u64, } impl ByteRange { - pub fn end_exclusive(self) -> Option { + pub(crate) fn end_exclusive(self) -> Option { self.start.checked_add(self.len) } } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct StageShardTensor { +pub(crate) struct StageShardTensor { pub name: String, pub dims: Vec, pub ggml_type: u32, @@ -38,7 +39,7 @@ pub struct StageShardTensor { } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct StageShardPlan { +pub(crate) struct StageShardPlan { pub source: GgufSource, pub stage_index: u32, pub stage_count: u32, @@ -60,27 +61,27 @@ pub struct StageShardPlan { } impl StageShardPlan { - pub fn cache_file_name(&self) -> String { + pub(crate) fn cache_file_name(&self) -> String { format!("{}.stage-{:05}.gguf", self.cache_key, self.stage_index) } - pub fn source_url(&self) -> Result { + pub(crate) fn source_url(&self) -> Result { source_url(&self.source) } - pub fn planned_tensor_fetch_bytes(&self) -> u64 { + pub(crate) fn planned_tensor_fetch_bytes(&self) -> u64 { self.merged_tensor_ranges .iter() .map(|range| range.len) .sum() } - pub fn planned_fetch_bytes(&self) -> u64 { + pub(crate) fn planned_fetch_bytes(&self) -> u64 { self.metadata_end .saturating_add(self.planned_tensor_fetch_bytes()) } - pub fn planned_range_count(&self) -> usize { + pub(crate) fn planned_range_count(&self) -> usize { self.merged_tensor_ranges.len() + if self.metadata_end > 0 { 1 } else { 0 } } } @@ -104,7 +105,7 @@ struct GgufDirectory { tensors: Vec, } -pub fn plan_stage_shard( +pub(crate) fn plan_stage_shard( planning_gguf: &Path, source: GgufSource, stage_index: u32, @@ -167,7 +168,7 @@ pub fn plan_stage_shard( }) } -pub fn validate_stage_shard_cache(path: &Path, plan: &StageShardPlan) -> Result<(), String> { +pub(crate) fn validate_stage_shard_cache(path: &Path, plan: &StageShardPlan) -> Result<(), String> { let directory = read_gguf_directory(path) .map_err(|error| format!("invalid cached stage shard {}: {error}", path.display()))?; if directory.tensors.len() != plan.tensors.len() { @@ -206,7 +207,7 @@ pub fn validate_stage_shard_cache(path: &Path, plan: &StageShardPlan) -> Result< Ok(()) } -pub fn source_url(source: &GgufSource) -> Result { +pub(crate) fn source_url(source: &GgufSource) -> Result { match source { GgufSource::HuggingFaceGguf { repo, @@ -266,9 +267,14 @@ fn read_gguf_directory(path: &Path) -> Result { for _ in 0..metadata_count { let key = read_gguf_string(&mut file, MAX_STRING_BYTES)?; - let value_type = GgufValueType::read(&mut file)?; + let value_type = GgufValueType::read(&mut file, "GGUF value type")?; if key == "general.alignment" && value_type.is_integer() { - alignment = read_integer_value(&mut file, value_type)?; + alignment = read_integer_value( + &mut file, + value_type, + |other| format!("GGUF value type {other:?} is not integer"), + |value| format!("negative GGUF integer {value}"), + )?; } else { skip_value(&mut file, value_type)?; } @@ -463,7 +469,7 @@ fn shard_cache_key( hasher.finalize().to_hex()[..24].to_owned() } -pub fn materialize_stage_shard_http( +pub(crate) fn materialize_stage_shard_http( plan: &StageShardPlan, output_path: &Path, emit: F, @@ -475,7 +481,7 @@ where materialize_stage_shard_from_url(plan, &url, output_path, emit) } -pub fn materialize_stage_shard_from_url( +pub(crate) fn materialize_stage_shard_from_url( plan: &StageShardPlan, url: &str, output_path: &Path, @@ -793,72 +799,6 @@ fn pad_writer_to_alignment(writer: &mut W, alignment: u64) -> R Ok(()) } -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -enum GgufValueType { - Uint8, - Int8, - Uint16, - Int16, - Uint32, - Int32, - Float32, - Bool, - String, - Array, - Uint64, - Int64, - Float64, -} - -impl GgufValueType { - fn read(reader: &mut R) -> Result { - 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 { - match self { - Self::Uint8 | Self::Int8 | Self::Bool => Some(1), - Self::Uint16 | Self::Int16 => Some(2), - Self::Uint32 | Self::Int32 | Self::Float32 => Some(4), - Self::Uint64 | Self::Int64 | Self::Float64 => Some(8), - Self::String | Self::Array => None, - } - } - - fn is_integer(self) -> bool { - matches!( - self, - Self::Uint8 - | Self::Int8 - | Self::Uint16 - | Self::Int16 - | Self::Uint32 - | Self::Int32 - | Self::Uint64 - | Self::Int64 - ) - } -} - fn skip_value(reader: &mut R, value_type: GgufValueType) -> Result<(), String> { match value_type { GgufValueType::String => skip_gguf_string(reader), @@ -868,7 +808,7 @@ fn skip_value(reader: &mut R, value_type: GgufValueType) -> Resu } fn skip_array(reader: &mut R) -> Result<(), String> { - let element_type = GgufValueType::read(reader)?; + let element_type = GgufValueType::read(reader, "GGUF value type")?; let len = read_u64(reader)?; match element_type { GgufValueType::String => { @@ -893,28 +833,6 @@ fn skip_array(reader: &mut R) -> Result<(), String> { } } -fn read_integer_value(reader: &mut R, value_type: GgufValueType) -> Result { - match value_type { - GgufValueType::Uint8 => read_u8(reader).map(u64::from), - GgufValueType::Int8 => read_i8(reader).and_then(non_negative_i64_to_u64), - GgufValueType::Uint16 => read_u16(reader).map(u64::from), - GgufValueType::Int16 => { - read_i16(reader).and_then(|v| non_negative_i64_to_u64(i64::from(v))) - } - GgufValueType::Uint32 => read_u32(reader).map(u64::from), - GgufValueType::Int32 => { - read_i32(reader).and_then(|v| non_negative_i64_to_u64(i64::from(v))) - } - GgufValueType::Uint64 => read_u64(reader), - GgufValueType::Int64 => read_i64(reader).and_then(non_negative_i64_to_u64), - other => Err(format!("GGUF value type {other:?} is not integer")), - } -} - -fn non_negative_i64_to_u64(value: i64) -> Result { - u64::try_from(value).map_err(|_| format!("negative GGUF integer {value}")) -} - fn read_gguf_string(reader: &mut R, max_len: u64) -> Result { let len = read_u64(reader)?; if len > max_len { @@ -954,63 +872,3 @@ fn align_to(value: u64, alignment: u64) -> Result { .ok_or_else(|| format!("align {value} to {alignment} overflows")) } } - -fn read_u8(reader: &mut R) -> Result { - let mut bytes = [0; 1]; - reader - .read_exact(&mut bytes) - .map_err(|e| format!("read u8: {e}"))?; - Ok(bytes[0]) -} - -fn read_i8(reader: &mut R) -> Result { - read_u8(reader).map(|value| i8::from_le_bytes([value]) as i64) -} - -fn read_u16(reader: &mut R) -> Result { - let mut bytes = [0; 2]; - reader - .read_exact(&mut bytes) - .map_err(|e| format!("read u16: {e}"))?; - Ok(u16::from_le_bytes(bytes)) -} - -fn read_i16(reader: &mut R) -> Result { - let mut bytes = [0; 2]; - reader - .read_exact(&mut bytes) - .map_err(|e| format!("read i16: {e}"))?; - Ok(i16::from_le_bytes(bytes)) -} - -fn read_u32(reader: &mut R) -> Result { - let mut bytes = [0; 4]; - reader - .read_exact(&mut bytes) - .map_err(|e| format!("read u32: {e}"))?; - Ok(u32::from_le_bytes(bytes)) -} - -fn read_i32(reader: &mut R) -> Result { - let mut bytes = [0; 4]; - reader - .read_exact(&mut bytes) - .map_err(|e| format!("read i32: {e}"))?; - Ok(i32::from_le_bytes(bytes)) -} - -fn read_u64(reader: &mut R) -> Result { - let mut bytes = [0; 8]; - reader - .read_exact(&mut bytes) - .map_err(|e| format!("read u64: {e}"))?; - Ok(u64::from_le_bytes(bytes)) -} - -fn read_i64(reader: &mut R) -> Result { - let mut bytes = [0; 8]; - reader - .read_exact(&mut bytes) - .map_err(|e| format!("read i64: {e}"))?; - Ok(i64::from_le_bytes(bytes)) -} diff --git a/crates/mvp-system/src/staging/mod.rs b/crates/mvp-system/src/staging/mod.rs index 6aa745e..d7ba27c 100644 --- a/crates/mvp-system/src/staging/mod.rs +++ b/crates/mvp-system/src/staging/mod.rs @@ -3,16 +3,10 @@ //! MVP stage control, shard planning, and weight lifecycle public surface. #[cfg(test)] -pub mod actor; -pub mod control; -pub mod gguf_metadata; +pub(crate) mod actor; +pub(crate) mod control; +pub(crate) mod gguf_metadata; #[cfg(test)] -pub mod shard_fetch; -#[cfg(test)] -pub mod shard_weight_lifecycle; -#[cfg(test)] -pub mod weight_lifecycle; -#[cfg(test)] -pub mod weight_shards; +pub(crate) mod weight_lifecycle; -pub use control::*; +pub(crate) use control::*; diff --git a/crates/mvp-system/src/staging/shard_fetch.rs b/crates/mvp-system/src/staging/shard_fetch.rs deleted file mode 100644 index 7bceb96..0000000 --- a/crates/mvp-system/src/staging/shard_fetch.rs +++ /dev/null @@ -1,108 +0,0 @@ -use crate::staging::weight_shards::{ShardAssignment, ShardManifest}; - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct ShardLocation { - pub uri: String, - pub cache_key: String, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct FetchShard { - pub location: ShardLocation, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct FetchedShard { - pub local_path: String, - pub manifest: ShardManifest, -} - -impl FetchedShard { - pub fn new(local_path: impl Into, manifest: ShardManifest) -> Self { - Self { - local_path: local_path.into(), - manifest, - } - } -} - -pub struct ShardLocator; - -impl ShardLocator { - pub fn locate(assignment: &ShardAssignment) -> ShardLocation { - let uri = assignment - .model_ref - .shard_uri(&assignment.split_id, assignment.stage_index); - let cache_key = format!( - "{}:{}:{:05}", - assignment.expected_model_digest().as_str(), - assignment.split_id.as_str(), - assignment.stage_index, - ); - ShardLocation { uri, cache_key } - } -} - -pub trait ShardCache { - fn get(&self, cache_key: &str) -> Option; - fn insert(&mut self, cache_key: String, shard: FetchedShard); -} - -pub trait ShardFetcher { - fn fetch(&mut self, request: &FetchShard) -> Result; -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum FetchError { - Unauthorized, - NotFound, - Unavailable, - IntegrityMismatch, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum ShardFetchStatus { - CacheHit, - Downloaded, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct ShardFetchOutcome { - pub shard: FetchedShard, - pub location: ShardLocation, - pub status: ShardFetchStatus, -} - -pub struct ShardFetchCoordinator; - -impl ShardFetchCoordinator { - pub fn get_or_fetch( - assignment: &ShardAssignment, - cache: &mut C, - fetcher: &mut F, - ) -> Result - where - C: ShardCache, - F: ShardFetcher, - { - let location = ShardLocator::locate(assignment); - if let Some(shard) = cache.get(&location.cache_key) { - return Ok(ShardFetchOutcome { - shard, - location, - status: ShardFetchStatus::CacheHit, - }); - } - - let request = FetchShard { - location: location.clone(), - }; - let shard = fetcher.fetch(&request)?; - cache.insert(location.cache_key.clone(), shard.clone()); - Ok(ShardFetchOutcome { - shard, - location, - status: ShardFetchStatus::Downloaded, - }) - } -} diff --git a/crates/mvp-system/src/staging/shard_weight_lifecycle.rs b/crates/mvp-system/src/staging/shard_weight_lifecycle.rs deleted file mode 100644 index 0a82497..0000000 --- a/crates/mvp-system/src/staging/shard_weight_lifecycle.rs +++ /dev/null @@ -1,165 +0,0 @@ -use crate::staging::shard_fetch::{ - FetchError, ShardCache, ShardFetchCoordinator, ShardFetchStatus, ShardFetcher, ShardLocation, -}; -use crate::staging::weight_shards::{ShardAssignment, ShardValidationError, ValidatedShard}; - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum ShardLifecycleState { - Idle, - Assigned, - Located, - Fetching, - Fetched, - Validating, - Binding, - Ready, - Faulted, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum ShardLifecycleEvent { - Assigned { assignment: ShardAssignment }, - Located { location: ShardLocation }, - Fetching { location: ShardLocation }, - CacheHit { cache_key: String }, - Fetched { uri: String, local_path: String }, - Validated { local_path: String }, - Binding { local_path: String }, - Ready { local_path: String }, - Faulted { reason: ShardLifecycleFault }, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum ShardLifecycleFault { - Fetch(FetchError), - Validation(ShardValidationError), - Bind(BindError), -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum BindError { - WorkerRejected, - DeviceAllocationFailed, -} - -pub trait WorkerShardBinder { - fn bind(&mut self, shard: &ValidatedShard) -> Result<(), BindError>; -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct ShardWeightLifecycle { - state: ShardLifecycleState, - events: Vec, -} - -impl ShardWeightLifecycle { - pub fn new() -> Self { - Self { - state: ShardLifecycleState::Idle, - events: Vec::new(), - } - } - - pub fn state(&self) -> ShardLifecycleState { - self.state - } - - pub fn events(&self) -> &[ShardLifecycleEvent] { - &self.events - } - - pub fn load( - &mut self, - assignment: ShardAssignment, - cache: &mut C, - fetcher: &mut F, - binder: &mut B, - ) where - C: ShardCache, - F: ShardFetcher, - B: WorkerShardBinder, - { - if matches!( - self.state, - ShardLifecycleState::Ready | ShardLifecycleState::Faulted - ) { - return; - } - - self.state = ShardLifecycleState::Assigned; - self.events.push(ShardLifecycleEvent::Assigned { - assignment: assignment.clone(), - }); - - let location = crate::staging::shard_fetch::ShardLocator::locate(&assignment); - self.state = ShardLifecycleState::Located; - self.events.push(ShardLifecycleEvent::Located { - location: location.clone(), - }); - - self.state = ShardLifecycleState::Fetching; - self.events.push(ShardLifecycleEvent::Fetching { - location: location.clone(), - }); - - let outcome = match ShardFetchCoordinator::get_or_fetch(&assignment, cache, fetcher) { - Ok(outcome) => outcome, - Err(error) => { - self.fault(ShardLifecycleFault::Fetch(error)); - return; - } - }; - - match outcome.status { - ShardFetchStatus::CacheHit => self.events.push(ShardLifecycleEvent::CacheHit { - cache_key: outcome.location.cache_key, - }), - ShardFetchStatus::Downloaded => self.events.push(ShardLifecycleEvent::Fetched { - uri: outcome.location.uri, - local_path: outcome.shard.local_path.clone(), - }), - } - - self.state = ShardLifecycleState::Fetched; - self.state = ShardLifecycleState::Validating; - let validated = match ValidatedShard::new( - assignment, - outcome.shard.manifest, - outcome.shard.local_path.clone(), - ) { - Ok(validated) => validated, - Err(error) => { - self.fault(ShardLifecycleFault::Validation(error)); - return; - } - }; - self.events.push(ShardLifecycleEvent::Validated { - local_path: validated.local_path.clone(), - }); - - self.state = ShardLifecycleState::Binding; - self.events.push(ShardLifecycleEvent::Binding { - local_path: validated.local_path.clone(), - }); - if let Err(error) = binder.bind(&validated) { - self.fault(ShardLifecycleFault::Bind(error)); - return; - } - - self.state = ShardLifecycleState::Ready; - self.events.push(ShardLifecycleEvent::Ready { - local_path: validated.local_path, - }); - } - - fn fault(&mut self, reason: ShardLifecycleFault) { - self.state = ShardLifecycleState::Faulted; - self.events.push(ShardLifecycleEvent::Faulted { reason }); - } -} - -impl Default for ShardWeightLifecycle { - fn default() -> Self { - Self::new() - } -} diff --git a/crates/mvp-system/src/staging/weight_lifecycle.rs b/crates/mvp-system/src/staging/weight_lifecycle.rs index 56649b5..8f1b2fd 100644 --- a/crates/mvp-system/src/staging/weight_lifecycle.rs +++ b/crates/mvp-system/src/staging/weight_lifecycle.rs @@ -1,23 +1,23 @@ #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct RunId(pub u64); +pub(crate) struct RunId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct NodeId(pub u64); +pub(crate) struct NodeId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct LayerRange { +pub(crate) struct LayerRange { pub start: u32, pub end_exclusive: u32, } #[derive(Clone, Debug, PartialEq, Eq)] -pub enum WeightSource { +pub(crate) enum WeightSource { WholeGguf { uri: String }, ShardSet { uris: Vec }, CachedArtifact { cache_key: String }, } #[derive(Clone, Debug, PartialEq, Eq)] -pub struct WeightAssignment { +pub(crate) struct WeightAssignment { pub run_id: RunId, pub stage_index: u32, pub plan_layer_range: LayerRange, @@ -26,12 +26,12 @@ pub struct WeightAssignment { } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum ArtifactBytes { +pub(crate) enum ArtifactBytes { Local, } #[derive(Clone, Debug, PartialEq, Eq)] -pub enum WeightEvent { +pub(crate) enum WeightEvent { Provisioned(WeightAssignment), ArtifactAvailable { bytes: ArtifactBytes }, LayerRangeValidated, @@ -45,7 +45,7 @@ pub enum WeightEvent { } #[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum StageFaultReason { +pub(crate) enum StageFaultReason { WeightDownloadFailed, WeightParseFailed, DeviceAllocationFailed, @@ -54,7 +54,7 @@ pub enum StageFaultReason { } #[derive(Clone, Debug, PartialEq, Eq)] -pub enum WeightLifecycleEvent { +pub(crate) enum WeightLifecycleEvent { WeightsReady { run_id: RunId, stage_index: u32, @@ -71,7 +71,7 @@ pub enum WeightLifecycleEvent { } #[derive(Clone, Debug, PartialEq, Eq)] -pub enum WeightCommand { +pub(crate) enum WeightCommand { LoadOrBindRange { source: WeightSource, range: LayerRange, @@ -82,7 +82,7 @@ pub enum WeightCommand { } #[cfg(test)] -pub struct WeightLifecycleHarness { +pub(crate) struct WeightLifecycleHarness { _node_id: NodeId, assignment: Option, artifact: bool, @@ -97,7 +97,7 @@ pub struct WeightLifecycleHarness { #[cfg(test)] impl WeightLifecycleHarness { - pub fn new(node_id: NodeId) -> Self { + pub(crate) fn new(node_id: NodeId) -> Self { Self { _node_id: node_id, assignment: None, @@ -112,7 +112,7 @@ impl WeightLifecycleHarness { } } - pub fn observe(&mut self, event: WeightEvent) { + pub(crate) fn observe(&mut self, event: WeightEvent) { match event { WeightEvent::Provisioned(assignment) => { self.commands.push(WeightCommand::LoadOrBindRange { @@ -145,11 +145,11 @@ impl WeightLifecycleHarness { self.maybe_stage_ready(); } - pub fn commands(&self) -> &[WeightCommand] { + pub(crate) fn commands(&self) -> &[WeightCommand] { &self.commands } - pub fn events(&self) -> &[WeightLifecycleEvent] { + pub(crate) fn events(&self) -> &[WeightLifecycleEvent] { &self.events } diff --git a/crates/mvp-system/src/staging/weight_shards.rs b/crates/mvp-system/src/staging/weight_shards.rs deleted file mode 100644 index 44048d2..0000000 --- a/crates/mvp-system/src/staging/weight_shards.rs +++ /dev/null @@ -1,348 +0,0 @@ -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct ModelArtifactRef { - canonical: String, - repo: String, - revision: String, - path: String, -} - -impl ModelArtifactRef { - pub fn parse(value: impl Into) -> Result { - let value = value.into(); - let rest = value - .strip_prefix("hf://") - .ok_or(ModelArtifactRefError::InvalidScheme)?; - let (repo, revision_and_path) = rest - .split_once('@') - .ok_or(ModelArtifactRefError::MissingRevision)?; - let (revision, path) = revision_and_path - .split_once('/') - .ok_or(ModelArtifactRefError::MissingPath)?; - Self::hugging_face(repo, revision, path) - } - - pub fn hugging_face( - repo: impl Into, - revision: impl Into, - path: impl Into, - ) -> Result { - let repo = repo.into().trim_matches('/').to_owned(); - let revision = revision.into(); - let path = path.into().trim_start_matches('/').to_owned(); - if repo.is_empty() { - return Err(ModelArtifactRefError::MissingRepo); - } - if revision.is_empty() { - return Err(ModelArtifactRefError::MissingRevision); - } - if revision.contains('/') { - return Err(ModelArtifactRefError::RevisionMustBePathSegment); - } - if path.is_empty() { - return Err(ModelArtifactRefError::MissingPath); - } - let canonical = format!("hf://{repo}@{revision}/{path}"); - Ok(Self { - canonical, - repo, - revision, - path, - }) - } - - pub fn as_str(&self) -> &str { - &self.canonical - } - - pub fn repo(&self) -> &str { - &self.repo - } - - pub fn revision(&self) -> &str { - &self.revision - } - - pub fn path(&self) -> &str { - &self.path - } - - pub fn model_digest(&self) -> ModelDigest { - ModelDigest(stable_digest_hex(&["model", self.as_str()])) - } - - pub fn shard_uri(&self, split_id: &SplitId, stage_index: u32) -> String { - format!( - "hf://{}@{}/shards/{}/stage-{stage_index:05}.gguf", - self.repo, - self.revision, - split_id.as_str(), - ) - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum ModelArtifactRefError { - InvalidScheme, - MissingRepo, - MissingRevision, - RevisionMustBePathSegment, - MissingPath, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum SplitScheme { - GgufLayerContiguousV1, -} - -impl SplitScheme { - pub fn as_str(&self) -> &'static str { - match self { - Self::GgufLayerContiguousV1 => "gguf-layer-contiguous-v1", - } - } -} - -#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct SplitId(String); - -impl SplitId { - pub fn derive(model_ref: &ModelArtifactRef, scheme: SplitScheme) -> Self { - Self(format!( - "split-{}", - stable_digest_hex(&["split", model_ref.as_str(), scheme.as_str()]) - )) - } - - pub fn literal(value: impl Into) -> Result { - let value = value.into(); - if value.is_empty() { - return Err(SplitIdError::Empty); - } - if !value - .bytes() - .all(|byte| byte.is_ascii_alphanumeric() || byte == b'-' || byte == b'_') - { - return Err(SplitIdError::InvalidCharacter); - } - Ok(Self(value)) - } - - pub fn as_str(&self) -> &str { - &self.0 - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum SplitIdError { - Empty, - InvalidCharacter, -} - -#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct ModelDigest(String); - -impl ModelDigest { - pub fn literal(value: impl Into) -> Result { - let value = value.into(); - if value.is_empty() { - return Err(ModelDigestError::Empty); - } - Ok(Self(value)) - } - - pub fn as_str(&self) -> &str { - &self.0 - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum ModelDigestError { - Empty, -} - -#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct ContentHash(String); - -impl ContentHash { - pub fn literal(value: impl Into) -> Result { - let value = value.into(); - if value.is_empty() { - return Err(ContentHashError::Empty); - } - Ok(Self(value)) - } - - pub fn as_str(&self) -> &str { - &self.0 - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum ContentHashError { - Empty, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct LayerRange { - pub start: u32, - pub end_exclusive: u32, -} - -impl LayerRange { - pub fn new(start: u32, end_exclusive: u32) -> Result { - if start >= end_exclusive { - return Err(LayerRangeError::EmptyOrInverted); - } - Ok(Self { - start, - end_exclusive, - }) - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum LayerRangeError { - EmptyOrInverted, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct ShardAssignment { - pub model_ref: ModelArtifactRef, - pub split_id: SplitId, - pub split_scheme: SplitScheme, - pub stage_index: u32, - pub stage_count: u32, - pub layer_range: LayerRange, -} - -impl ShardAssignment { - pub fn new( - model_ref: ModelArtifactRef, - split_id: SplitId, - split_scheme: SplitScheme, - stage_index: u32, - stage_count: u32, - layer_range: LayerRange, - ) -> Result { - if stage_count == 0 { - return Err(ShardAssignmentError::EmptyStageCount); - } - if stage_index >= stage_count { - return Err(ShardAssignmentError::StageIndexOutOfRange); - } - Ok(Self { - model_ref, - split_id, - split_scheme, - stage_index, - stage_count, - layer_range, - }) - } - - pub fn expected_model_digest(&self) -> ModelDigest { - self.model_ref.model_digest() - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum ShardAssignmentError { - EmptyStageCount, - StageIndexOutOfRange, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct ShardManifest { - pub model_digest: ModelDigest, - pub split_id: SplitId, - pub stage_index: u32, - pub stage_count: u32, - pub layer_range: LayerRange, - pub content_hash: ContentHash, -} - -impl ShardManifest { - pub fn for_assignment(assignment: &ShardAssignment, content_hash: ContentHash) -> Self { - Self { - model_digest: assignment.expected_model_digest(), - split_id: assignment.split_id.clone(), - stage_index: assignment.stage_index, - stage_count: assignment.stage_count, - layer_range: assignment.layer_range, - content_hash, - } - } -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct ValidatedShard { - pub assignment: ShardAssignment, - pub manifest: ShardManifest, - pub local_path: String, -} - -impl ValidatedShard { - pub fn new( - assignment: ShardAssignment, - manifest: ShardManifest, - local_path: impl Into, - ) -> Result { - ShardValidator::validate(&assignment, &manifest)?; - Ok(Self { - assignment, - manifest, - local_path: local_path.into(), - }) - } -} - -pub struct ShardValidator; - -impl ShardValidator { - pub fn validate( - assignment: &ShardAssignment, - manifest: &ShardManifest, - ) -> Result<(), ShardValidationError> { - if manifest.model_digest != assignment.expected_model_digest() { - return Err(ShardValidationError::ModelDigestMismatch); - } - if manifest.split_id != assignment.split_id { - return Err(ShardValidationError::SplitIdMismatch); - } - if manifest.stage_index != assignment.stage_index { - return Err(ShardValidationError::StageIndexMismatch); - } - if manifest.stage_count != assignment.stage_count { - return Err(ShardValidationError::StageCountMismatch); - } - if manifest.layer_range != assignment.layer_range { - return Err(ShardValidationError::LayerRangeMismatch); - } - Ok(()) - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum ShardValidationError { - ModelDigestMismatch, - SplitIdMismatch, - StageIndexMismatch, - StageCountMismatch, - LayerRangeMismatch, -} - -fn stable_digest_hex(parts: &[&str]) -> String { - const FNV_OFFSET: u64 = 0xcbf29ce484222325; - const FNV_PRIME: u64 = 0x100000001b3; - - let mut hash = FNV_OFFSET; - for part in parts { - for byte in part.as_bytes() { - hash ^= u64::from(*byte); - hash = hash.wrapping_mul(FNV_PRIME); - } - hash ^= 0xff; - hash = hash.wrapping_mul(FNV_PRIME); - } - format!("{hash:016x}") -} diff --git a/crates/mvp-system/src/tests/local_e2e_guarantees.rs b/crates/mvp-system/src/tests/local_e2e_guarantees.rs index 727b9de..2f5fc8e 100644 --- a/crates/mvp-system/src/tests/local_e2e_guarantees.rs +++ b/crates/mvp-system/src/tests/local_e2e_guarantees.rs @@ -174,21 +174,21 @@ fn assert_engine_builder_surface(outcome: &LocalMockOutcome) { outcome .engine_events .iter() - .any(|event| matches!(event, engine::EngineEvent::PoolAcquired { .. })), + .any(|event| matches!(event, engine::EngineEvent::PoolAcquired)), "local mock integration must be built from a neutral engine pool" ); assert!( outcome .engine_events .iter() - .any(|event| matches!(event, engine::EngineEvent::ClusterConverged { .. })), + .any(|event| matches!(event, engine::EngineEvent::ClusterConverged)), "local mock integration must pass through the builder convergence barrier" ); assert!( outcome .engine_events .iter() - .any(|event| matches!(event, engine::EngineEvent::EngineReady { .. })), + .any(|event| matches!(event, engine::EngineEvent::EngineReady)), "local mock integration must return an engine-ready handle before workload IO" ); let assigned_stages = outcome @@ -197,10 +197,7 @@ fn assert_engine_builder_surface(outcome: &LocalMockOutcome) { .filter(|event| { matches!( event, - engine::EngineEvent::RoleAssigned { - role: engine::RoleKind::StageWorker { .. }, - .. - } + engine::EngineEvent::RoleAssigned(engine::RoleKind::StageWorker { .. }) ) }) .count(); diff --git a/crates/mvp-system/src/tests/local_mock/environment.rs b/crates/mvp-system/src/tests/local_mock/environment.rs index 858e10e..dcc5052 100644 --- a/crates/mvp-system/src/tests/local_mock/environment.rs +++ b/crates/mvp-system/src/tests/local_mock/environment.rs @@ -104,26 +104,18 @@ impl ResourceTracker { } fn mock_pool(orchestrator_node_id: plan::NodeId, stage_count: u32) -> engine::StaticPoolProvider { - let mut leases = Vec::with_capacity(stage_count as usize + 1); - leases.push( - engine::NodeLease::new( - "mock-coordinator", - engine::NodeId(orchestrator_node_id.0), - [engine::NodeCapability::Coordinator], - ) - .resources(engine::ResourceFacts::cpu_only(2, 2 << 30)), - ); + let mut nodes = Vec::with_capacity(stage_count as usize + 1); + nodes.push(engine::NodeFacts { + node_id: engine::NodeId(orchestrator_node_id.0), + capabilities: vec![engine::NodeCapability::Coordinator], + }); for stage_index in 0..stage_count { - leases.push( - engine::NodeLease::new( - format!("mock-worker-{stage_index}"), - engine::NodeId(11 + u64::from(stage_index)), - [engine::NodeCapability::Worker], - ) - .resources(engine::ResourceFacts::cpu_only(2, 2 << 30)), - ); + nodes.push(engine::NodeFacts { + node_id: engine::NodeId(11 + u64::from(stage_index)), + capabilities: vec![engine::NodeCapability::Worker], + }); } - engine::StaticPoolProvider::new(leases) + engine::StaticPoolProvider::new(nodes) } impl LocalMockCluster { @@ -144,9 +136,7 @@ impl LocalMockCluster { "local-mock", engine::ModelSpec::pipelined_causal_llm( "mock-gguf", - engine::ModelArtifact::TestTinyLlm { - path: "local-mock://mock-gguf".to_owned(), - }, + plan::GgufSource::LocalPath("local-mock://mock-gguf".to_owned()), config.stage_count * 2, 8, engine::DTypeFamily::BFloat, @@ -157,22 +147,15 @@ impl LocalMockCluster { ), ) .run_id(run_id.0) - .image( - engine::NodeImageSpec::new("local-mock-node") - .worker_runtime(engine::WorkerRuntimeSpec::DumbProcess), - ) .pool_provider(mock_pool(orchestrator_node_id, config.stage_count)) - .launcher(engine::StaticNodeLauncher) .planner( engine::FixedLinearPipelinePlanner::new(config.stage_count).runtime( plan::RuntimeConfig { max_tokens: config.max_tokens, - prompt: plan::PromptSource::Inline("local mock prompt".to_owned()), sampling: plan::SamplingPolicy { temperature_millis: 0, top_k: 1, }, - token_output_policy: plan::TokenOutputPolicy::EmitAll, }, ), ) diff --git a/crates/mvp-system/src/tests/orchestration_guarantees.rs b/crates/mvp-system/src/tests/orchestration_guarantees.rs index 39d9d00..fbf7b29 100644 --- a/crates/mvp-system/src/tests/orchestration_guarantees.rs +++ b/crates/mvp-system/src/tests/orchestration_guarantees.rs @@ -14,36 +14,22 @@ mod run_plan { use crate::run_plan as plan; - // Local aliases keep the test prose readable while the file imports only the - // public planning module. The aliases do not grant access to planner internals. + // Local aliases keep the test prose readable while importing only the public + // planning module. The aliases do not grant access to planner internals. type DTypeFamily = plan::DTypeFamily; type EdgeEndpoint = plan::EdgeEndpoint; - type EdgeId = plan::EdgeId; type EdgeKind = plan::EdgeKind; - type EdgePlan = plan::EdgePlan; type GgufSource = plan::GgufSource; - type HostPinning = plan::HostPinning; - type InboundEdgeProvision = plan::InboundEdgeProvision; type ModelFacts = plan::ModelFacts; use plan::NodeId; - type LayoutRule = plan::LayoutRule; - type ObjectKind = plan::ObjectKind; - type OutboundEdgeProvision = plan::OutboundEdgeProvision; - type PromptSource = plan::PromptSource; type PlacementInput = plan::PlacementInput; type PlanRejectionKind = plan::PlanRejectionKind; type PlannerInput = plan::PlannerInput; type RingSpec = plan::RingSpec; - type RingDirection = plan::RingDirection; - type RunPlan = plan::RunPlan; type RuntimeConfig = plan::RuntimeConfig; type SamplingPolicy = plan::SamplingPolicy; - type SequencePolicy = plan::SequencePolicy; - type ShapeRule = plan::ShapeRule; type StagePlacement = plan::StagePlacement; - type TokenOutputPolicy = plan::TokenOutputPolicy; type TokenizerSource = plan::TokenizerSource; - type WakeCoalescing = plan::WakeCoalescing; // Keep test node ids small and readable. The concrete identity mechanism is // outside this contract; these ids exist only so assertions can name topology @@ -93,31 +79,31 @@ mod run_plan { }, runtime: RuntimeConfig { max_tokens: 4, - prompt: PromptSource::Inline("hello from planner input".into()), sampling: SamplingPolicy { temperature_millis: 125, top_k: 7, }, - token_output_policy: TokenOutputPolicy::EmitAll, }, candidate_pool: valid_nodes(), stage_count, placement: linear_placement(stage_count), - activation_ring: RingSpec::test_default_activation(), - token_ring: RingSpec::test_default_token(), + activation_ring: RingSpec { + data_capacity: 1 << 20, + alignment: 64, + direction: plan::RingDirection::Egress, + host_pinning: plan::HostPinning::Pageable, + wake_coalescing: plan::WakeCoalescing::PendingBit, + }, + token_ring: RingSpec { + data_capacity: 4096, + alignment: 8, + direction: plan::RingDirection::Egress, + host_pinning: plan::HostPinning::Pageable, + wake_coalescing: plan::WakeCoalescing::PendingBit, + }, } } - // Tests frequently need to compare a provisioned edge id back to the canonical - // edge record in the RunPlan. This helper makes that lookup explicit without - // giving tests access to any planner-private index. - fn plan_edges_by_id(plan: &RunPlan) -> std::collections::BTreeMap { - plan.edges - .iter() - .map(|edge| (edge.edge_id, edge)) - .collect::>() - } - // Edge endpoints can be orchestrator or stage endpoints. Tests use this helper // when they care only about stage adjacency and want orchestrator endpoints to // remain visibly outside the stage index space. @@ -128,16 +114,6 @@ mod run_plan { } } - // Provisioning sends concrete node ids across the data-flow boundary. This - // helper extracts the observable node id from either endpoint shape so tests - // can compare projection output to plan topology. - fn edge_node_id(endpoint: &EdgeEndpoint) -> NodeId { - match endpoint { - EdgeEndpoint::Orchestrator { node_id } => *node_id, - EdgeEndpoint::Stage { node_id, .. } => *node_id, - } - } - // This proves RunPlan formation is a total public boundary for valid input: // the caller observes one complete plan, not hidden follow-up topology work or // a partially initialized result. @@ -170,17 +146,12 @@ mod run_plan { TokenizerSource::LocalPath("/tokenizers/test-gguf.json".into()) ); assert_eq!( - plan.runtime.sampling, + plan.sampling, SamplingPolicy { temperature_millis: 125, top_k: 7, } ); - assert_eq!( - plan.runtime.prompt, - PromptSource::Inline("hello from planner input".into()) - ); - assert_eq!(plan.runtime.token_output_policy, TokenOutputPolicy::EmitAll); // Every stage must be bound to this run and know the run's stage count. for stage in &plan.stages { @@ -374,7 +345,7 @@ mod run_plan { assert_eq!(first.runtime.input_port, plan::PortId("input".into())); assert_eq!(first.runtime.output_port, plan::PortId("output".into())); let expected_sampling = if stage_index + 1 == stage.stage_count { - Some(plan.runtime.sampling) + Some(plan.sampling) } else { None }; @@ -517,39 +488,6 @@ mod run_plan { input.placement = linear_placement(4); input } - - // Zero sequence length makes activation capacity zero. The planner must reject - // before creating edges whose object specs cannot carry an activation. - fn invalid_zero_activation_extent() -> PlannerInput { - let mut input = valid_input(3, 36); - input.model.max_seq_len = 0; - input - } - - // Dtype width participates directly in activation extent and object layout. - // A zero width is not a valid dtype fact and must reject before planning. - fn invalid_dtype_width() -> PlannerInput { - let mut input = valid_input(3, 36); - input.model.dtype_width_bytes = 0; - input - } - - // Hidden dimension participates directly in activation shape. A zero hidden - // dimension represents an unsupported shape/layout fact for the MVP contract. - fn invalid_unsupported_shape_or_layout() -> PlannerInput { - let mut input = valid_input(3, 36); - input.model.hidden_dim = 0; - input - } - - // Ring alignment must be a usable alignment contract for shared memory and - // device copy boundaries. A non-power-of-two alignment makes the ring spec - // invalid before any edge can be provisioned. - fn invalid_ring_alignment() -> PlannerInput { - let mut input = valid_input(3, 36); - input.activation_ring.alignment = 3; - input - } } mod run_fsm { diff --git a/crates/mvp-system/src/transport/codec_registry.rs b/crates/mvp-system/src/transport/codec_registry.rs deleted file mode 100644 index e58a451..0000000 --- a/crates/mvp-system/src/transport/codec_registry.rs +++ /dev/null @@ -1,11 +0,0 @@ -//! MVP runtime codec registration. -//! -//! Actor behavior lives in the owning domain modules. This module wires their -//! message codecs into the transport registry used by distributed runtimes. - -pub fn register_mvp_actor_codecs(registry: &mut swactor_transport::CodecRegistry) { - crate::node_actor::register_codecs(registry); - crate::orchestration::actor::register_codecs(registry); - datastream::register_datastream_publisher_codec(registry); - crate::prompt::rpc::register_codecs(registry); -} diff --git a/crates/mvp-system/src/transport/driver_pumps.rs b/crates/mvp-system/src/transport/driver_pumps.rs index b431870..bf4cb8a 100644 --- a/crates/mvp-system/src/transport/driver_pumps.rs +++ b/crates/mvp-system/src/transport/driver_pumps.rs @@ -1,352 +1,57 @@ #![allow(dead_code)] -use std::collections::{BTreeMap, BTreeSet}; +use std::collections::BTreeMap; #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct NodeId(pub u64); +pub(crate) struct EdgeId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct EdgeId(pub u64); +pub(crate) struct RingId(pub(crate) u64); #[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct RingId(pub u64); - -#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct StreamId(pub u64); - -#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct Alpn(pub String); +pub(crate) struct StreamId(pub(crate) u64); #[derive(Clone, Debug, PartialEq, Eq)] -pub struct DriverConfig { - pub local_node_id: NodeId, - pub alpn: Alpn, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct EstablishSend { - pub edge_id: EdgeId, - pub peer_node_id: NodeId, - pub layout: RingLayout, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct EstablishRecv { - pub edge_id: EdgeId, - pub layout: RingLayout, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum RingDirection { - Egress, - Ingress, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct RingLayout { - pub ring_id: RingId, - pub byte_capacity: usize, - pub direction: RingDirection, -} - -impl RingLayout { - pub fn test_egress() -> Self { - Self { - ring_id: RingId(1), - byte_capacity: 4096, - direction: RingDirection::Egress, - } - } - - pub fn test_ingress() -> Self { - Self { - ring_id: RingId(2), - byte_capacity: 4096, - direction: RingDirection::Ingress, - } - } -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum DriverEvent { - EstablishSend(EstablishSend), - EstablishRecv(EstablishRecv), - IncomingUniStream { - edge_id: EdgeId, - stream_id: StreamId, - }, - RingReadable { - ring_id: RingId, - }, - RingWritable { - ring_id: RingId, - }, - EgressBytesCommitted { - edge_id: EdgeId, - bytes: Vec, - }, - StreamBytesRead { - edge_id: EdgeId, - bytes: Vec, - }, - WriteAllAccepted { - edge_id: EdgeId, - byte_count: usize, - }, - NetworkStalled { - edge_id: EdgeId, - }, - IngressRingFull { - edge_id: EdgeId, - }, - ReadError { - edge_id: EdgeId, - }, - WriteError { - edge_id: EdgeId, - }, - StopEdge { - edge_id: EdgeId, - }, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum DriverCommand { - OpenOrReuseConnection { - peer_node_id: NodeId, - alpn: Alpn, - local_node_id: NodeId, - }, - OpenUniStream { - edge_id: EdgeId, - peer_node_id: NodeId, - }, - SpawnSendPump { - edge_id: EdgeId, - ring_id: RingId, - }, - SpawnRecvPump { - edge_id: EdgeId, - ring_id: RingId, - stream_id: StreamId, - }, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum ActorMessage { - PollStreamFuture { edge_id: EdgeId }, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum DriverEventOut { - DriverEdgeReady { - edge_id: EdgeId, - }, - StreamClosed { - edge_id: EdgeId, - }, - StreamFault { - edge_id: EdgeId, - reason: StreamFaultReason, - }, - PumpStopped { - edge_id: EdgeId, - ring_id: RingId, - }, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum StreamFaultReason { - ReadError, - WriteError, - ProtocolError, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum WakeHint { - RingReadable { ring_id: RingId }, - RingWritable { ring_id: RingId }, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct StreamWrite { - pub edge_id: EdgeId, - pub bytes: Vec, -} - -pub fn encode_edge_preamble(edge_id: EdgeId) -> Vec { - edge_id.0.to_le_bytes().to_vec() -} - -pub fn count_preamble_occurrences(bytes: &[u8], edge_id: EdgeId) -> usize { - let preamble = encode_edge_preamble(edge_id); - bytes - .windows(preamble.len()) - .filter(|window| *window == preamble.as_slice()) - .count() -} - -pub fn fake_object_header_bytes() -> Vec { - b"OBJ\0fake-header".to_vec() -} - -pub struct Driver { - state: DriverState, -} - -impl Driver { - pub fn new(config: DriverConfig) -> Self { - Self { - state: DriverState::new(config), - } - } - - pub fn observe(&mut self, event: DriverEvent) { - self.state.observe(event); - } - - pub fn commands(&self) -> &[DriverCommand] { - &self.state.commands - } - - pub fn events(&self) -> &[DriverEventOut] { - &self.state.events - } - - pub fn wake_hints(&self) -> &[WakeHint] { - &self.state.wakes - } +pub(crate) enum DriverEventOut { + DriverEdgeReady { edge_id: EdgeId }, + StreamFault { edge_id: EdgeId }, + PumpStopped { edge_id: EdgeId, ring_id: RingId }, } #[derive(Debug)] -struct DriverState { - config: DriverConfig, - connections: BTreeSet<(NodeId, Alpn)>, - sends: BTreeMap, - recv_specs: BTreeMap, +pub(crate) struct Driver { + sends: BTreeMap, + recv_specs: BTreeMap, pending_streams: BTreeMap, - recvs: BTreeMap, - commands: Vec, + recvs: BTreeMap, events: Vec, - wakes: Vec, - #[cfg(test)] - actor_messages: Vec, - stream_writes: Vec, - read_started: BTreeSet, } -impl DriverState { - fn new(config: DriverConfig) -> Self { +impl Driver { + pub(crate) fn new() -> Self { Self { - config, - connections: BTreeSet::new(), sends: BTreeMap::new(), recv_specs: BTreeMap::new(), pending_streams: BTreeMap::new(), recvs: BTreeMap::new(), - commands: Vec::new(), events: Vec::new(), - wakes: Vec::new(), - #[cfg(test)] - actor_messages: Vec::new(), - stream_writes: Vec::new(), - read_started: BTreeSet::new(), } } - fn observe(&mut self, event: DriverEvent) { - match event { - DriverEvent::EstablishSend(spec) => self.establish_send(spec), - DriverEvent::EstablishRecv(spec) => self.establish_recv(spec), - DriverEvent::IncomingUniStream { edge_id, stream_id } => { - self.incoming_uni_stream(edge_id, stream_id); - } - DriverEvent::RingReadable { ring_id } => self.flush_send_bytes(ring_id), - DriverEvent::RingWritable { ring_id } => self.resume_recv(ring_id), - DriverEvent::EgressBytesCommitted { edge_id, bytes } => { - if let Some(send) = self.sends.get_mut(&edge_id) { - send.pending_bytes.extend(bytes); - } - } - DriverEvent::StreamBytesRead { edge_id, bytes } => { - self.copy_recv_bytes(edge_id, &bytes) - } - DriverEvent::WriteAllAccepted { - edge_id, - byte_count, - } => { - if let Some(send) = self.sends.get_mut(&edge_id) { - send.consume_cursor += byte_count; - send.network_stalled = false; - self.wakes.push(WakeHint::RingWritable { - ring_id: send.ring_id, - }); - } - } - DriverEvent::NetworkStalled { edge_id } => { - if let Some(send) = self.sends.get_mut(&edge_id) { - send.network_stalled = true; - } - } - DriverEvent::IngressRingFull { edge_id } => { - if let Some(recv) = self.recvs.get_mut(&edge_id) { - recv.reading = false; - } - } - DriverEvent::ReadError { edge_id } => { - self.events.push(DriverEventOut::StreamFault { - edge_id, - reason: StreamFaultReason::ReadError, - }); - } - DriverEvent::WriteError { edge_id } => { - self.events.push(DriverEventOut::StreamFault { - edge_id, - reason: StreamFaultReason::WriteError, - }); - } - DriverEvent::StopEdge { edge_id } => self.stop_edge(edge_id), - } + pub(crate) fn establish_send(&mut self, edge_id: EdgeId, ring_id: RingId) { + self.sends.insert(edge_id, ring_id); + self.events + .push(DriverEventOut::DriverEdgeReady { edge_id }); } - fn establish_send(&mut self, spec: EstablishSend) { - let connection_key = (spec.peer_node_id, self.config.alpn.clone()); - self.connections.insert(connection_key); - self.commands.push(DriverCommand::OpenOrReuseConnection { - peer_node_id: spec.peer_node_id, - alpn: self.config.alpn.clone(), - local_node_id: self.config.local_node_id, - }); - - self.commands.push(DriverCommand::SpawnSendPump { - edge_id: spec.edge_id, - ring_id: spec.layout.ring_id, - }); - - let send = SendPumpState::new(spec.layout.ring_id); - self.commands.push(DriverCommand::OpenUniStream { - edge_id: spec.edge_id, - peer_node_id: spec.peer_node_id, - }); - self.stream_writes.push(StreamWrite { - edge_id: spec.edge_id, - bytes: encode_edge_preamble(spec.edge_id), - }); - self.sends.insert(spec.edge_id, send); - self.events.push(DriverEventOut::DriverEdgeReady { - edge_id: spec.edge_id, - }); - } - - fn establish_recv(&mut self, spec: EstablishRecv) { - let edge_id = spec.edge_id; - self.recv_specs.insert(edge_id, spec); + pub(crate) fn establish_recv(&mut self, edge_id: EdgeId, ring_id: RingId) { + self.recv_specs.insert(edge_id, ring_id); if let Some(stream_id) = self.pending_streams.remove(&edge_id) { self.spawn_recv(edge_id, stream_id); } } - fn incoming_uni_stream(&mut self, edge_id: EdgeId, stream_id: StreamId) { + pub(crate) fn incoming_uni_stream(&mut self, edge_id: EdgeId, stream_id: StreamId) { if self.recv_specs.contains_key(&edge_id) { self.spawn_recv(edge_id, stream_id); } else { @@ -355,84 +60,26 @@ impl DriverState { } fn spawn_recv(&mut self, edge_id: EdgeId, stream_id: StreamId) { - let Some(spec) = self.recv_specs.get(&edge_id) else { + let Some(ring_id) = self.recv_specs.get(&edge_id).copied() else { self.pending_streams.insert(edge_id, stream_id); return; }; - let ring_id = spec.layout.ring_id; - self.commands.push(DriverCommand::SpawnRecvPump { - edge_id, - ring_id, - stream_id, - }); - self.read_started.insert(stream_id); - self.recvs.insert( - edge_id, - RecvPumpState { - ring_id, - commit_cursor: 0, - reading: true, - }, - ); + self.recvs.insert(edge_id, ring_id); self.events .push(DriverEventOut::DriverEdgeReady { edge_id }); } - fn flush_send_bytes(&mut self, ring_id: RingId) { - let mut send_edge_id = None; - for (edge_id, send) in &self.sends { - if send.ring_id == ring_id { - send_edge_id = Some(*edge_id); - break; - } - } - let Some(edge_id) = send_edge_id else { - return; - }; - let Some(send) = self.sends.get_mut(&edge_id) else { - return; - }; - if send.network_stalled || send.pending_bytes.is_empty() { - return; - } - - let bytes = std::mem::take(&mut send.pending_bytes); - self.stream_writes.push(StreamWrite { edge_id, bytes }); + pub(crate) fn read_error(&mut self, edge_id: EdgeId) { + self.events.push(DriverEventOut::StreamFault { edge_id }); } - fn resume_recv(&mut self, ring_id: RingId) { - for recv in self.recvs.values_mut() { - if recv.ring_id == ring_id { - recv.reading = true; - return; - } - } - } - - fn copy_recv_bytes(&mut self, edge_id: EdgeId, bytes: &[u8]) { - let Some(recv) = self.recvs.get_mut(&edge_id) else { - return; - }; - if !recv.reading { - return; - } - recv.commit_cursor += bytes.len(); - self.wakes.push(WakeHint::RingReadable { - ring_id: recv.ring_id, - }); - } - - fn stop_edge(&mut self, edge_id: EdgeId) { + pub(crate) fn stop_edge(&mut self, edge_id: EdgeId) { let ring_id = self .sends .get(&edge_id) - .map(|send| send.ring_id) - .or_else(|| self.recvs.get(&edge_id).map(|recv| recv.ring_id)) - .or_else(|| { - self.recv_specs - .get(&edge_id) - .map(|spec| spec.layout.ring_id) - }) + .copied() + .or_else(|| self.recvs.get(&edge_id).copied()) + .or_else(|| self.recv_specs.get(&edge_id).copied()) .unwrap_or(RingId(0)); self.sends.remove(&edge_id); @@ -442,119 +89,7 @@ impl DriverState { self.events .push(DriverEventOut::PumpStopped { edge_id, ring_id }); } -} - -#[derive(Debug)] -struct SendPumpState { - ring_id: RingId, - pending_bytes: Vec, - consume_cursor: usize, - network_stalled: bool, -} - -impl SendPumpState { - fn new(ring_id: RingId) -> Self { - Self { - ring_id, - pending_bytes: Vec::new(), - consume_cursor: 0, - network_stalled: false, - } - } -} - -#[derive(Debug)] -struct RecvPumpState { - ring_id: RingId, - commit_cursor: usize, - reading: bool, -} - -#[cfg(test)] -pub struct CommandLog { - commands: Vec, -} - -#[cfg(test)] -impl CommandLog { - pub fn iter(&self) -> std::vec::IntoIter { - self.commands.clone().into_iter() - } -} - -#[cfg(test)] -pub struct DriverHarness { - driver: Driver, -} - -#[cfg(test)] -impl DriverHarness { - pub fn new(config: DriverConfig) -> Self { - Self { - driver: Driver::new(config), - } - } - - pub fn observe(&mut self, event: DriverEvent) { - self.driver.observe(event); - } - - pub fn commands(&self) -> CommandLog { - CommandLog { - commands: self.driver.state.commands.clone(), - } - } - - pub fn events(&self) -> &[DriverEventOut] { - &self.driver.state.events - } - - pub fn wake_hints(&self) -> &[WakeHint] { - &self.driver.state.wakes - } - - pub fn actor_messages(&self) -> &[ActorMessage] { - &self.driver.state.actor_messages - } - - pub fn stream_writes(&self, edge_id: EdgeId) -> Vec { - self.driver - .state - .stream_writes - .iter() - .filter(|write| write.edge_id == edge_id) - .cloned() - .collect() - } - - pub fn stream_reads_started(&self, stream_id: StreamId) -> bool { - self.driver.state.read_started.contains(&stream_id) - } - - pub fn is_reading_stream(&self, edge_id: EdgeId) -> bool { - self.driver - .state - .recvs - .get(&edge_id) - .map(|recv| recv.reading) - .unwrap_or(false) - } - - pub fn ring_commit(&self, edge_id: EdgeId) -> usize { - self.driver - .state - .recvs - .get(&edge_id) - .map(|recv| recv.commit_cursor) - .unwrap_or(0) - } - - pub fn ring_consume(&self, edge_id: EdgeId) -> usize { - self.driver - .state - .sends - .get(&edge_id) - .map(|send| send.consume_cursor) - .unwrap_or(0) + pub(crate) fn events(&self) -> &[DriverEventOut] { + &self.events } } diff --git a/crates/mvp-system/src/transport/endpoint_advertisement.rs b/crates/mvp-system/src/transport/endpoint_advertisement.rs index 4242c99..cf96a22 100644 --- a/crates/mvp-system/src/transport/endpoint_advertisement.rs +++ b/crates/mvp-system/src/transport/endpoint_advertisement.rs @@ -2,17 +2,17 @@ use std::fmt; use iroh::EndpointAddr; -pub const MVP_IROH_ENDPOINT_ADDR_MASK_ENV: &str = "MVP_IROH_ENDPOINT_ADDR_MASK"; +pub(crate) const MVP_IROH_ENDPOINT_ADDR_MASK_ENV: &str = "MVP_IROH_ENDPOINT_ADDR_MASK"; #[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] -pub enum EndpointAddrMask { +pub(crate) enum EndpointAddrMask { #[default] Full, RelayOnly, } impl EndpointAddrMask { - pub fn parse(value: &str) -> Result { + pub(crate) fn parse(value: &str) -> Result { match value.trim().to_ascii_lowercase().as_str() { "" | "full" | "none" => Ok(Self::Full), "relay-only" | "relay_only" | "relay" => Ok(Self::RelayOnly), @@ -22,14 +22,14 @@ impl EndpointAddrMask { } } - pub fn as_str(self) -> &'static str { + pub(crate) fn as_str(self) -> &'static str { match self { Self::Full => "full", Self::RelayOnly => "relay-only", } } - pub fn requires_relay(self) -> bool { + pub(crate) fn requires_relay(self) -> bool { matches!(self, Self::RelayOnly) } } @@ -40,7 +40,7 @@ impl fmt::Display for EndpointAddrMask { } } -pub fn advertised_endpoint( +pub(crate) fn advertised_endpoint( endpoint: EndpointAddr, mask: EndpointAddrMask, ) -> Result { diff --git a/crates/mvp-system/src/transport/json_codec.rs b/crates/mvp-system/src/transport/json_codec.rs index 72cfecf..7da1f23 100644 --- a/crates/mvp-system/src/transport/json_codec.rs +++ b/crates/mvp-system/src/transport/json_codec.rs @@ -5,7 +5,7 @@ use serde::de::DeserializeOwned; use swactor::Error; use swactor_transport::Codec; -pub struct JsonCodec(PhantomData); +pub(crate) struct JsonCodec(PhantomData); impl Default for JsonCodec { fn default() -> Self { diff --git a/crates/mvp-system/src/transport/mod.rs b/crates/mvp-system/src/transport/mod.rs index 4c4b3ec..acb52cd 100644 --- a/crates/mvp-system/src/transport/mod.rs +++ b/crates/mvp-system/src/transport/mod.rs @@ -1,5 +1,11 @@ //! MVP edge transport public surface. -pub mod codec_registry; -pub mod endpoint_advertisement; -pub mod json_codec; +pub(crate) fn register_mvp_actor_codecs(registry: &mut swactor_transport::CodecRegistry) { + crate::node_actor::register_codecs(registry); + crate::orchestration::actor::register_codecs(registry); + datastream::register_datastream_publisher_codec(registry); + crate::prompt::rpc::register_codecs(registry); +} + +pub(crate) mod endpoint_advertisement; +pub(crate) mod json_codec;