stash mvp system unit test buildout

This commit is contained in:
Zachery Aaron Shores-Chmielewski 2026-06-23 17:51:34 +04:00
parent 4820d751a6
commit 3e206c931a
44 changed files with 8543 additions and 290 deletions

1
Cargo.lock generated
View file

@ -2925,6 +2925,7 @@ name = "mvp-system"
version = "0.1.0"
dependencies = [
"libc",
"serde_json",
]
[[package]]

View file

@ -59,9 +59,6 @@ proptest = "1"
proptest-state-machine = "0.3"
mvp-system = { path = "crates/mvp-system" }
[[test]]
name = "arena_manager_guarantees"
path = "tests/mvp_system/arena_manager_guarantees.rs"
[[bench]]
name = "runtime_benchmarks"

View file

@ -4,5 +4,9 @@ version = "0.1.0"
edition = "2024"
publish = false
[dependencies]
serde_json = "1"
[target.'cfg(target_os = "linux")'.dependencies]
libc = "0.2"

View file

@ -0,0 +1,548 @@
use std::collections::{BTreeMap, VecDeque};
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct LeaseRequestId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct RingId(pub u64);
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ArenaConfig {
pub node_id: NodeId,
pub reservation_ceiling: u64,
pub base_alignment: u64,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RingSpec {
pub header_bytes: u64,
pub data_bytes: u64,
pub alignment: u64,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct LeaseRing {
pub request_id: LeaseRequestId,
pub ring_spec: RingSpec,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum LayoutPointer {
NoProcessPointer,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RingLayout {
pub start_offset: u64,
pub header_offset: u64,
pub data_offset: u64,
pub end_offset: u64,
pub data_bytes: u64,
pub alignment: u64,
pub pointer: LayoutPointer,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RingLease {
pub request_id: LeaseRequestId,
pub ring_id: RingId,
pub node_id: NodeId,
pub layout: RingLayout,
pub requested_alignment: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct QuiescenceProof {
verified: bool,
}
impl QuiescenceProof {
pub fn verified() -> Self {
Self { verified: true }
}
pub fn missing() -> Self {
Self { verified: false }
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ArenaRequest {
LeaseRing(LeaseRing),
CancelLease {
request_id: LeaseRequestId,
},
ReleaseRing {
ring_id: RingId,
proof: QuiescenceProof,
},
Shutdown,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ArenaFault {
InvalidReservationCeiling,
InvalidBaseAlignment,
BackingUnavailable,
UnsupportedPlatform,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RingLeaseRejection {
CannotFitWithinCeiling,
ArenaShuttingDown,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RingReleaseRejection {
MissingQuiescenceProof,
UnknownRingId,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ArenaEvent {
RingLeased {
lease: RingLease,
},
RingLeaseQueued {
request_id: LeaseRequestId,
},
RingLeaseRejected {
request_id: LeaseRequestId,
reason: RingLeaseRejection,
},
RingReleased {
ring_id: RingId,
start_offset: u64,
end_offset: u64,
},
RingReleaseRejected {
ring_id: RingId,
reason: RingReleaseRejection,
},
CancelledFreshLeaseReleased {
request_id: LeaseRequestId,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ArenaCommand {
InstallWorkerOrPumpState {
request_id: LeaseRequestId,
ring_id: RingId,
layout: RingLayout,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum ArenaState {
Ready,
ShuttingDown,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct FreeRange {
start: u64,
end: u64,
}
#[derive(Clone, Debug, PartialEq, Eq)]
struct QueuedLease {
request: LeaseRing,
cancelled: bool,
}
pub struct ArenaManager {
config: ArenaConfig,
_backing: ArenaBacking,
state: ArenaState,
next_ring_id: u64,
free_ranges: Vec<FreeRange>,
pending: VecDeque<QueuedLease>,
live_order: Vec<RingLease>,
live_index: BTreeMap<RingId, usize>,
}
impl ArenaManager {
pub fn boot(config: ArenaConfig) -> Result<Self, ArenaFault> {
let backing = ArenaBacking::create(&config)?;
Ok(Self {
free_ranges: vec![FreeRange {
start: 0,
end: config.reservation_ceiling,
}],
config,
_backing: backing,
state: ArenaState::Ready,
next_ring_id: 1,
pending: VecDeque::new(),
live_order: Vec::new(),
live_index: BTreeMap::new(),
})
}
fn request(&mut self, request: ArenaRequest) -> Vec<ArenaEvent> {
match request {
ArenaRequest::LeaseRing(request) => self.lease_ring(request),
ArenaRequest::CancelLease { request_id } => {
self.cancel_lease(request_id);
Vec::new()
}
ArenaRequest::ReleaseRing { ring_id, proof } => self.release_ring(ring_id, proof),
ArenaRequest::Shutdown => self.shutdown(),
}
}
fn live_leases(&self) -> &[RingLease] {
&self.live_order
}
fn lookup_lease(&self, ring_id: RingId) -> Option<&RingLease> {
self.live_index
.get(&ring_id)
.and_then(|index| self.live_order.get(*index))
}
fn lease_ring(&mut self, request: LeaseRing) -> Vec<ArenaEvent> {
if self.state == ArenaState::ShuttingDown {
return vec![ArenaEvent::RingLeaseRejected {
request_id: request.request_id,
reason: RingLeaseRejection::ArenaShuttingDown,
}];
}
if self.request_layout_at(&request.ring_spec, 0).is_none() {
return vec![ArenaEvent::RingLeaseRejected {
request_id: request.request_id,
reason: RingLeaseRejection::CannotFitWithinCeiling,
}];
}
if self.pending.is_empty() {
if let Some(lease) = self.try_allocate(&request) {
return vec![ArenaEvent::RingLeased { lease }];
}
}
let request_id = request.request_id;
self.pending.push_back(QueuedLease {
request,
cancelled: false,
});
vec![ArenaEvent::RingLeaseQueued { request_id }]
}
fn cancel_lease(&mut self, request_id: LeaseRequestId) {
for pending in &mut self.pending {
if pending.request.request_id == request_id {
pending.cancelled = true;
}
}
}
fn release_ring(&mut self, ring_id: RingId, proof: QuiescenceProof) -> Vec<ArenaEvent> {
let Some(index) = self.live_index.get(&ring_id).copied() else {
return vec![ArenaEvent::RingReleaseRejected {
ring_id,
reason: RingReleaseRejection::UnknownRingId,
}];
};
if !proof.verified {
return vec![ArenaEvent::RingReleaseRejected {
ring_id,
reason: RingReleaseRejection::MissingQuiescenceProof,
}];
}
let lease = self.live_order.remove(index);
self.rebuild_live_index();
self.insert_free_range(FreeRange {
start: lease.layout.start_offset,
end: lease.layout.end_offset,
});
let mut events = vec![ArenaEvent::RingReleased {
ring_id,
start_offset: lease.layout.start_offset,
end_offset: lease.layout.end_offset,
}];
if self.state == ArenaState::Ready {
self.retry_pending_leases(&mut events);
}
events
}
fn shutdown(&mut self) -> Vec<ArenaEvent> {
self.state = ArenaState::ShuttingDown;
let mut events = Vec::new();
while let Some(pending) = self.pending.pop_front() {
if pending.cancelled {
events.push(ArenaEvent::CancelledFreshLeaseReleased {
request_id: pending.request.request_id,
});
} else {
events.push(ArenaEvent::RingLeaseRejected {
request_id: pending.request.request_id,
reason: RingLeaseRejection::ArenaShuttingDown,
});
}
}
events
}
fn retry_pending_leases(&mut self, events: &mut Vec<ArenaEvent>) {
while let Some(pending) = self.pending.pop_front() {
if pending.cancelled {
events.push(ArenaEvent::CancelledFreshLeaseReleased {
request_id: pending.request.request_id,
});
continue;
}
if let Some(lease) = self.try_allocate(&pending.request) {
events.push(ArenaEvent::RingLeased { lease });
continue;
}
self.pending.push_front(pending);
break;
}
}
fn try_allocate(&mut self, request: &LeaseRing) -> Option<RingLease> {
for index in 0..self.free_ranges.len() {
let range = self.free_ranges[index];
let Some(layout) = self.request_layout_at(&request.ring_spec, range.start) else {
continue;
};
if layout.end_offset > range.end {
continue;
}
self.free_ranges.remove(index);
let mut insert_index = index;
if range.start < layout.start_offset {
self.free_ranges.insert(
insert_index,
FreeRange {
start: range.start,
end: layout.start_offset,
},
);
insert_index += 1;
}
if layout.end_offset < range.end {
self.free_ranges.insert(
insert_index,
FreeRange {
start: layout.end_offset,
end: range.end,
},
);
}
let ring_id = RingId(self.next_ring_id);
self.next_ring_id = self.next_ring_id.checked_add(1)?;
let lease = RingLease {
request_id: request.request_id,
ring_id,
node_id: self.config.node_id,
layout,
requested_alignment: request.ring_spec.alignment,
};
self.live_index.insert(ring_id, self.live_order.len());
self.live_order.push(lease.clone());
return Some(lease);
}
None
}
fn request_layout_at(&self, spec: &RingSpec, range_start: u64) -> Option<RingLayout> {
let alignment = lcm(self.config.base_alignment, spec.alignment)?;
let start_offset = align_up(range_start, alignment)?;
let header_offset = start_offset;
let after_header = header_offset.checked_add(spec.header_bytes)?;
let data_offset = align_up(after_header, alignment)?;
let end_offset = data_offset.checked_add(spec.data_bytes)?;
if end_offset > self.config.reservation_ceiling {
return None;
}
Some(RingLayout {
start_offset,
header_offset,
data_offset,
end_offset,
data_bytes: spec.data_bytes,
alignment,
pointer: LayoutPointer::NoProcessPointer,
})
}
fn insert_free_range(&mut self, range: FreeRange) {
self.free_ranges.push(range);
self.free_ranges.sort_by_key(|range| range.start);
let mut coalesced: Vec<FreeRange> = Vec::with_capacity(self.free_ranges.len());
for range in self.free_ranges.drain(..) {
if range.start == range.end {
continue;
}
if let Some(last) = coalesced.last_mut() {
if range.start <= last.end {
last.end = last.end.max(range.end);
continue;
}
}
coalesced.push(range);
}
self.free_ranges = coalesced;
}
fn rebuild_live_index(&mut self) {
self.live_index.clear();
for (index, lease) in self.live_order.iter().enumerate() {
self.live_index.insert(lease.ring_id, index);
}
}
}
pub struct ArenaManagerHarness {
manager: ArenaManager,
events: Vec<ArenaEvent>,
commands: Vec<ArenaCommand>,
}
impl ArenaManagerHarness {
pub fn boot(config: ArenaConfig) -> Result<Self, ArenaFault> {
Ok(Self {
manager: ArenaManager::boot(config)?,
events: Vec::new(),
commands: Vec::new(),
})
}
pub fn request(&mut self, request: ArenaRequest) {
self.events.extend(self.manager.request(request));
}
pub fn events(&self) -> &[ArenaEvent] {
&self.events
}
pub fn commands(&self) -> &[ArenaCommand] {
&self.commands
}
pub fn live_leases(&self) -> &[RingLease] {
self.manager.live_leases()
}
pub fn lookup_lease(&self, ring_id: RingId) -> Option<&RingLease> {
self.manager.lookup_lease(ring_id)
}
}
fn align_up(value: u64, alignment: u64) -> Option<u64> {
if alignment == 0 {
return None;
}
let remainder = value % alignment;
if remainder == 0 {
Some(value)
} else {
value.checked_add(alignment - remainder)
}
}
fn lcm(left: u64, right: u64) -> Option<u64> {
if left == 0 || right == 0 {
return None;
}
let gcd = gcd(left, right);
(left / gcd).checked_mul(right)
}
fn gcd(mut left: u64, mut right: u64) -> u64 {
while right != 0 {
let remainder = left % right;
left = right;
right = remainder;
}
left
}
#[cfg(target_os = "linux")]
struct ArenaBacking {
len: usize,
ptr: *mut libc::c_void,
_fd: std::os::fd::OwnedFd,
}
#[cfg(target_os = "linux")]
impl ArenaBacking {
fn create(config: &ArenaConfig) -> Result<Self, ArenaFault> {
if config.reservation_ceiling == 0 {
return Err(ArenaFault::InvalidReservationCeiling);
}
if config.base_alignment == 0 {
return Err(ArenaFault::InvalidBaseAlignment);
}
let len = usize::try_from(config.reservation_ceiling)
.map_err(|_| ArenaFault::InvalidReservationCeiling)?;
let fd = unsafe {
let name = b"mvp-system-arena\0";
libc::memfd_create(name.as_ptr().cast(), libc::MFD_CLOEXEC)
};
if fd < 0 {
return Err(ArenaFault::BackingUnavailable);
}
let fd = unsafe { <std::os::fd::OwnedFd as std::os::fd::FromRawFd>::from_raw_fd(fd) };
let truncate_result =
unsafe { libc::ftruncate(std::os::fd::AsRawFd::as_raw_fd(&fd), len as libc::off_t) };
if truncate_result != 0 {
return Err(ArenaFault::BackingUnavailable);
}
let ptr = unsafe {
libc::mmap(
std::ptr::null_mut(),
len,
libc::PROT_READ | libc::PROT_WRITE,
libc::MAP_SHARED,
std::os::fd::AsRawFd::as_raw_fd(&fd),
0,
)
};
if ptr == libc::MAP_FAILED {
return Err(ArenaFault::BackingUnavailable);
}
Ok(Self { len, ptr, _fd: fd })
}
}
#[cfg(target_os = "linux")]
impl Drop for ArenaBacking {
fn drop(&mut self) {
unsafe {
libc::munmap(self.ptr, self.len);
}
}
}
#[cfg(not(target_os = "linux"))]
struct ArenaBacking;
#[cfg(not(target_os = "linux"))]
impl ArenaBacking {
fn create(_config: &ArenaConfig) -> Result<Self, ArenaFault> {
Err(ArenaFault::UnsupportedPlatform)
}
}

View file

@ -0,0 +1,633 @@
use std::collections::{BTreeMap, BTreeSet};
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct WorkerGeneration(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct StepId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct CopyEvent(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct DeviceHandle {
pub generation: WorkerGeneration,
pub id: u64,
}
impl DeviceHandle {
pub const fn new(generation: WorkerGeneration, id: u64) -> Self {
Self { generation, id }
}
}
pub type DeviceAllocation = DeviceHandle;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DType {
U32,
F16,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Shape {
Vector,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ObjectSpec {
pub max_extent: u64,
pub alignment: u64,
pub dtype: DType,
pub shape: Shape,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct TensorViewSpec {
pub dtype: DType,
pub shape: Shape,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct TensorView {
pub handle: DeviceHandle,
pub dtype: DType,
pub shape: Shape,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct HostRange {
pub offset: u64,
pub len: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct DeviceRange {
pub handle: DeviceHandle,
pub offset: u64,
pub len: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CopyMode {
Sync,
Async,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DeviceError {
AllocationFailed,
InvalidExtent,
InvalidAlignment,
InvalidRange,
UnknownDeviceHandle,
OldGenerationHandle,
DeviceCopyFailed,
InvalidViewDType,
InvalidTensorView,
AllocationStillInUse,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DeviceEvent {
CopyCompleted {
copy: CopyEvent,
},
ComputeStarted {
handle: DeviceHandle,
step_id: StepId,
},
ComputeCompleted {
handle: DeviceHandle,
step_id: StepId,
},
WorkerRestarted {
generation: WorkerGeneration,
},
}
pub trait DeviceBridgeBackend {
fn alloc_device(
&mut self,
allocation: DeviceAllocation,
spec: ObjectSpec,
extent: u64,
) -> Result<(), DeviceError>;
fn free_device(&mut self, allocation: DeviceAllocation) -> Result<(), DeviceError>;
fn host_to_device(
&mut self,
host: HostRange,
device: DeviceRange,
copy: CopyEvent,
mode: CopyMode,
) -> Result<(), DeviceError>;
fn device_to_host(
&mut self,
device: DeviceRange,
host: HostRange,
copy: CopyEvent,
mode: CopyMode,
) -> Result<(), DeviceError>;
fn wrap_for_tinygrad(
&mut self,
allocation: DeviceAllocation,
view: TensorViewSpec,
) -> Result<TensorView, DeviceError>;
}
pub struct DeviceBridge<B> {
current_generation: WorkerGeneration,
next_handle_id: u64,
next_copy_id: u64,
allocations: BTreeMap<DeviceHandle, AllocationRecord>,
copies: BTreeMap<CopyEvent, CopyRecord>,
backend: B,
}
impl<B: DeviceBridgeBackend> DeviceBridge<B> {
pub fn new(current_generation: WorkerGeneration, backend: B) -> Self {
Self {
current_generation,
next_handle_id: 1,
next_copy_id: 1,
allocations: BTreeMap::new(),
copies: BTreeMap::new(),
backend,
}
}
pub fn alloc_device(
&mut self,
spec: ObjectSpec,
extent: u64,
) -> Result<DeviceAllocation, DeviceError> {
validate_object_spec(spec, extent)?;
let allocation = DeviceHandle::new(self.current_generation, self.next_handle_id);
self.next_handle_id += 1;
self.backend.alloc_device(allocation, spec, extent)?;
self.allocations.insert(
allocation,
AllocationRecord {
spec,
extent,
active_copies: BTreeSet::new(),
active_steps: BTreeSet::new(),
},
);
Ok(allocation)
}
pub fn free_device(&mut self, allocation: DeviceAllocation) -> Result<(), DeviceError> {
self.validate_generation(allocation)?;
let record = self
.allocations
.get(&allocation)
.ok_or(DeviceError::UnknownDeviceHandle)?;
if !record.active_copies.is_empty() || !record.active_steps.is_empty() {
return Err(DeviceError::AllocationStillInUse);
}
self.backend.free_device(allocation)?;
self.allocations.remove(&allocation);
Ok(())
}
pub fn host_to_device(
&mut self,
host: HostRange,
device: DeviceRange,
mode: CopyMode,
) -> Result<CopyEvent, DeviceError> {
self.validate_device_range(device)?;
validate_equal_copy_len(host.len, device.len)?;
let copy = self.next_copy_event();
self.backend.host_to_device(host, device, copy, mode)?;
self.record_copy(copy, device.handle, CopyDirection::HostToDevice, mode);
Ok(copy)
}
pub fn device_to_host(
&mut self,
device: DeviceRange,
host: HostRange,
mode: CopyMode,
) -> Result<CopyEvent, DeviceError> {
self.validate_device_range(device)?;
validate_equal_copy_len(device.len, host.len)?;
let copy = self.next_copy_event();
self.backend.device_to_host(device, host, copy, mode)?;
self.record_copy(copy, device.handle, CopyDirection::DeviceToHost, mode);
Ok(copy)
}
pub fn wrap_for_tinygrad(
&mut self,
allocation: DeviceAllocation,
view: TensorViewSpec,
) -> Result<TensorView, DeviceError> {
let record = self.allocation_record(allocation)?;
if view.dtype != record.spec.dtype {
return Err(DeviceError::InvalidViewDType);
}
if view.shape != record.spec.shape {
return Err(DeviceError::InvalidTensorView);
}
self.backend.wrap_for_tinygrad(allocation, view)
}
pub fn observe(&mut self, event: DeviceEvent) {
match event {
DeviceEvent::CopyCompleted { copy } => self.complete_copy(copy),
DeviceEvent::ComputeStarted { handle, step_id } => {
if let Some(record) = self.current_allocation_record_mut(handle) {
record.active_steps.insert(step_id);
}
}
DeviceEvent::ComputeCompleted { handle, step_id } => {
if let Some(record) = self.current_allocation_record_mut(handle) {
record.active_steps.remove(&step_id);
}
}
DeviceEvent::WorkerRestarted { generation } => {
self.current_generation = generation;
self.next_handle_id = 1;
self.next_copy_id = 1;
self.allocations.clear();
self.copies.clear();
}
}
}
pub fn safe_to_release_host(&self, copy: CopyEvent) -> bool {
self.copies
.get(&copy)
.map(|record| record.direction == CopyDirection::HostToDevice && record.completed)
.unwrap_or(false)
}
pub fn host_bytes_valid(&self, copy: CopyEvent) -> bool {
self.copies
.get(&copy)
.map(|record| record.direction == CopyDirection::DeviceToHost && record.completed)
.unwrap_or(false)
}
pub fn copy_event_complete(&self, copy: CopyEvent) -> bool {
self.copies
.get(&copy)
.map(|record| record.completed)
.unwrap_or(false)
}
pub fn backend(&self) -> &B {
&self.backend
}
pub fn backend_mut(&mut self) -> &mut B {
&mut self.backend
}
fn allocation_record(
&self,
allocation: DeviceAllocation,
) -> Result<&AllocationRecord, DeviceError> {
self.validate_generation(allocation)?;
self.allocations
.get(&allocation)
.ok_or(DeviceError::UnknownDeviceHandle)
}
fn current_allocation_record_mut(
&mut self,
allocation: DeviceAllocation,
) -> Option<&mut AllocationRecord> {
if allocation.generation != self.current_generation {
return None;
}
self.allocations.get_mut(&allocation)
}
fn validate_generation(&self, allocation: DeviceAllocation) -> Result<(), DeviceError> {
if allocation.generation != self.current_generation {
return Err(DeviceError::OldGenerationHandle);
}
Ok(())
}
fn validate_device_range(&self, device: DeviceRange) -> Result<(), DeviceError> {
let record = self.allocation_record(device.handle)?;
let end = device
.offset
.checked_add(device.len)
.ok_or(DeviceError::InvalidRange)?;
if end > record.extent {
return Err(DeviceError::InvalidRange);
}
Ok(())
}
fn next_copy_event(&mut self) -> CopyEvent {
let copy = CopyEvent(self.next_copy_id);
self.next_copy_id += 1;
copy
}
fn record_copy(
&mut self,
copy: CopyEvent,
allocation: DeviceAllocation,
direction: CopyDirection,
mode: CopyMode,
) {
let completed = mode == CopyMode::Sync;
self.copies.insert(
copy,
CopyRecord {
allocation,
direction,
completed,
},
);
if !completed {
if let Some(record) = self.allocations.get_mut(&allocation) {
record.active_copies.insert(copy);
}
}
}
fn complete_copy(&mut self, copy: CopyEvent) {
let Some(record) = self.copies.get_mut(&copy) else {
return;
};
record.completed = true;
let allocation = record.allocation;
if let Some(allocation_record) = self.allocations.get_mut(&allocation) {
allocation_record.active_copies.remove(&copy);
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
struct AllocationRecord {
spec: ObjectSpec,
extent: u64,
active_copies: BTreeSet<CopyEvent>,
active_steps: BTreeSet<StepId>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum CopyDirection {
HostToDevice,
DeviceToHost,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct CopyRecord {
allocation: DeviceAllocation,
direction: CopyDirection,
completed: bool,
}
fn validate_object_spec(spec: ObjectSpec, extent: u64) -> Result<(), DeviceError> {
if spec.alignment == 0 {
return Err(DeviceError::InvalidAlignment);
}
if extent > spec.max_extent || extent % spec.alignment != 0 {
return Err(DeviceError::InvalidExtent);
}
Ok(())
}
fn validate_equal_copy_len(left: u64, right: u64) -> Result<(), DeviceError> {
if left != right {
return Err(DeviceError::InvalidRange);
}
Ok(())
}
#[cfg(test)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum BackendFailure {
AllocationFailed,
CopyFailed,
InvalidView,
}
#[cfg(test)]
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum BackendCall {
Alloc {
allocation: DeviceAllocation,
spec: ObjectSpec,
extent: u64,
},
Free {
freed: DeviceAllocation,
},
HostToDevice {
host: HostRange,
device: DeviceRange,
copy: CopyEvent,
mode: CopyMode,
},
DeviceToHost {
device: DeviceRange,
host: HostRange,
copy: CopyEvent,
mode: CopyMode,
},
WrapForTinygrad {
allocation: DeviceAllocation,
view: TensorViewSpec,
},
}
#[cfg(test)]
#[derive(Default)]
pub struct MockDeviceBackend {
calls: Vec<BackendCall>,
next_failure: Option<BackendFailure>,
}
#[cfg(test)]
impl MockDeviceBackend {
pub fn inject_failure(&mut self, failure: BackendFailure) {
self.next_failure = Some(failure);
}
pub fn calls(&self) -> &[BackendCall] {
&self.calls
}
fn take_failure(&mut self, failure: BackendFailure) -> bool {
if self.next_failure == Some(failure) {
self.next_failure = None;
return true;
}
false
}
}
#[cfg(test)]
impl DeviceBridgeBackend for MockDeviceBackend {
fn alloc_device(
&mut self,
allocation: DeviceAllocation,
spec: ObjectSpec,
extent: u64,
) -> Result<(), DeviceError> {
self.calls.push(BackendCall::Alloc {
allocation,
spec,
extent,
});
if self.take_failure(BackendFailure::AllocationFailed) {
return Err(DeviceError::AllocationFailed);
}
Ok(())
}
fn free_device(&mut self, allocation: DeviceAllocation) -> Result<(), DeviceError> {
self.calls.push(BackendCall::Free { freed: allocation });
Ok(())
}
fn host_to_device(
&mut self,
host: HostRange,
device: DeviceRange,
copy: CopyEvent,
mode: CopyMode,
) -> Result<(), DeviceError> {
self.calls.push(BackendCall::HostToDevice {
host,
device,
copy,
mode,
});
if self.take_failure(BackendFailure::CopyFailed) {
return Err(DeviceError::DeviceCopyFailed);
}
Ok(())
}
fn device_to_host(
&mut self,
device: DeviceRange,
host: HostRange,
copy: CopyEvent,
mode: CopyMode,
) -> Result<(), DeviceError> {
self.calls.push(BackendCall::DeviceToHost {
device,
host,
copy,
mode,
});
if self.take_failure(BackendFailure::CopyFailed) {
return Err(DeviceError::DeviceCopyFailed);
}
Ok(())
}
fn wrap_for_tinygrad(
&mut self,
allocation: DeviceAllocation,
view: TensorViewSpec,
) -> Result<TensorView, DeviceError> {
self.calls
.push(BackendCall::WrapForTinygrad { allocation, view });
if self.take_failure(BackendFailure::InvalidView) {
return Err(DeviceError::InvalidTensorView);
}
Ok(TensorView {
handle: allocation,
dtype: view.dtype,
shape: view.shape,
})
}
}
#[cfg(test)]
pub struct DeviceBridgeHarness {
bridge: DeviceBridge<MockDeviceBackend>,
}
#[cfg(test)]
impl DeviceBridgeHarness {
pub fn new(generation: WorkerGeneration) -> Self {
Self {
bridge: DeviceBridge::new(generation, MockDeviceBackend::default()),
}
}
pub fn alloc_device(
&mut self,
spec: ObjectSpec,
extent: u64,
) -> Result<DeviceHandle, DeviceError> {
self.bridge.alloc_device(spec, extent)
}
pub fn free_device(&mut self, allocation: DeviceAllocation) -> Result<(), DeviceError> {
self.bridge.free_device(allocation)
}
pub fn host_to_device(
&mut self,
host: HostRange,
device: DeviceRange,
mode: CopyMode,
) -> Result<CopyEvent, DeviceError> {
self.bridge.host_to_device(host, device, mode)
}
pub fn device_to_host(
&mut self,
device: DeviceRange,
host: HostRange,
mode: CopyMode,
) -> Result<CopyEvent, DeviceError> {
self.bridge.device_to_host(device, host, mode)
}
pub fn wrap_for_tinygrad(
&mut self,
allocation: DeviceAllocation,
view: TensorViewSpec,
) -> Result<TensorView, DeviceError> {
self.bridge.wrap_for_tinygrad(allocation, view)
}
pub fn observe(&mut self, event: DeviceEvent) {
self.bridge.observe(event);
}
pub fn safe_to_release_host(&self, copy: CopyEvent) -> bool {
self.bridge.safe_to_release_host(copy)
}
pub fn host_bytes_valid(&self, copy: CopyEvent) -> bool {
self.bridge.host_bytes_valid(copy)
}
pub fn copy_event_complete(&self, copy: CopyEvent) -> bool {
self.bridge.copy_event_complete(copy)
}
pub fn inject_backend_failure(&mut self, failure: BackendFailure) {
self.bridge.backend_mut().inject_failure(failure);
}
pub fn backend_calls(&self) -> &[BackendCall] {
self.bridge.backend().calls()
}
}

View file

@ -0,0 +1,544 @@
use std::collections::{BTreeMap, BTreeSet};
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct EdgeId(pub 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);
#[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 {
edge_id: EdgeId,
},
RingWritable {
edge_id: EdgeId,
},
EgressBytesCommitted {
edge_id: EdgeId,
bytes: Vec<u8>,
},
StreamBytesRead {
edge_id: EdgeId,
bytes: Vec<u8>,
},
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,
},
ObjectHeaderParsed {
edge_id: EdgeId,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum StreamFaultReason {
ReadError,
WriteError,
ProtocolError,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WakeHint {
RingReadable { edge_id: EdgeId },
RingWritable { edge_id: EdgeId },
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct StreamWrite {
pub edge_id: EdgeId,
pub bytes: Vec<u8>,
}
pub fn encode_edge_preamble(edge_id: EdgeId) -> Vec<u8> {
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<u8> {
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
}
}
#[derive(Debug)]
struct DriverState {
config: DriverConfig,
connections: BTreeSet<(NodeId, Alpn)>,
sends: BTreeMap<EdgeId, SendPumpState>,
recv_specs: BTreeMap<EdgeId, EstablishRecv>,
pending_streams: BTreeMap<EdgeId, StreamId>,
recvs: BTreeMap<EdgeId, RecvPumpState>,
commands: Vec<DriverCommand>,
events: Vec<DriverEventOut>,
wakes: Vec<WakeHint>,
#[cfg(test)]
actor_messages: Vec<ActorMessage>,
stream_writes: Vec<StreamWrite>,
read_started: BTreeSet<StreamId>,
}
impl DriverState {
fn new(config: DriverConfig) -> 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 { edge_id } => self.flush_send_bytes(edge_id),
DriverEvent::RingWritable { edge_id } => self.resume_recv(edge_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 { edge_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),
}
}
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);
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) {
if self.recv_specs.contains_key(&edge_id) {
self.spawn_recv(edge_id, stream_id);
} else {
self.pending_streams.insert(edge_id, stream_id);
}
}
fn spawn_recv(&mut self, edge_id: EdgeId, stream_id: StreamId) {
let Some(spec) = self.recv_specs.get(&edge_id) 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.events
.push(DriverEventOut::DriverEdgeReady { edge_id });
}
fn flush_send_bytes(&mut self, edge_id: EdgeId) {
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 });
}
fn resume_recv(&mut self, edge_id: EdgeId) {
if let Some(recv) = self.recvs.get_mut(&edge_id) {
recv.reading = true;
}
}
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 { edge_id });
}
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)
})
.unwrap_or(RingId(0));
self.sends.remove(&edge_id);
self.recvs.remove(&edge_id);
self.recv_specs.remove(&edge_id);
self.pending_streams.remove(&edge_id);
self.events
.push(DriverEventOut::PumpStopped { edge_id, ring_id });
}
}
#[derive(Debug)]
struct SendPumpState {
ring_id: RingId,
pending_bytes: Vec<u8>,
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<DriverCommand>,
}
#[cfg(test)]
impl CommandLog {
pub fn iter(&self) -> std::vec::IntoIter<DriverCommand> {
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<StreamWrite> {
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)
}
}

View file

@ -0,0 +1,751 @@
use std::collections::BTreeMap;
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct RunId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct EdgeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct LeaseRequestId(pub 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 ActorAddress(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RingDirection {
Egress,
Ingress,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ObjectKind {
Activation,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DType {
F16,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ObjectSpec {
pub kind: ObjectKind,
pub dtype: DType,
pub max_extent_bytes: u64,
}
impl ObjectSpec {
pub const fn test_activation() -> Self {
Self {
kind: ObjectKind::Activation,
dtype: DType::F16,
max_extent_bytes: 4096,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct RingSpec {
pub header_bytes: u64,
pub data_bytes: u64,
pub alignment: u64,
}
impl RingSpec {
pub const fn test_activation() -> Self {
Self {
header_bytes: 128,
data_bytes: 4096,
alignment: 64,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct RingLayout {
pub start_offset: u64,
pub header_offset: u64,
pub data_offset: u64,
pub end_offset: u64,
pub data_bytes: u64,
pub alignment: u64,
}
impl RingLayout {
pub const fn test_layout(start_offset: u64) -> Self {
Self {
start_offset,
header_offset: start_offset,
data_offset: start_offset + 128,
end_offset: start_offset + 128 + 4096,
data_bytes: 4096,
alignment: 64,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct QuiescenceProof {
verified: bool,
}
impl QuiescenceProof {
pub const fn verified() -> Self {
Self { verified: true }
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ProvisionTx {
pub run_id: RunId,
pub edge_id: EdgeId,
pub local_node_id: NodeId,
pub consumer_node_id: NodeId,
pub object_spec: ObjectSpec,
pub ring_spec: RingSpec,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ProvisionRx {
pub run_id: RunId,
pub edge_id: EdgeId,
pub local_node_id: NodeId,
pub object_spec: ObjectSpec,
pub ring_spec: RingSpec,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RingLeaseRejection {
CannotFit,
ArenaShuttingDown,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum StreamFaultReason {
ReadError,
WriteError,
ProtocolError,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RingFaultReason {
WorkerRejectedRing,
WorkerCrashed,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum EdgeEvent {
ProvisionTx(ProvisionTx),
ProvisionRx(ProvisionRx),
RingLeased {
request_id: LeaseRequestId,
ring_id: RingId,
layout: RingLayout,
},
RingLeaseRejected {
request_id: LeaseRequestId,
reason: RingLeaseRejection,
},
RingInstalled {
edge_id: EdgeId,
ring_id: RingId,
},
DriverEdgeReady {
edge_id: EdgeId,
},
StreamFault {
edge_id: EdgeId,
reason: StreamFaultReason,
},
RingFault {
edge_id: EdgeId,
ring_id: RingId,
reason: RingFaultReason,
},
StopEdge {
edge_id: EdgeId,
},
PumpStopped {
edge_id: EdgeId,
ring_id: RingId,
},
RingQuiesced {
ring_id: RingId,
},
QuiescenceProven {
ring_id: RingId,
},
Stopped {
edge_id: EdgeId,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum EdgeCommand {
LeaseRing {
request_id: LeaseRequestId,
edge_id: EdgeId,
direction: RingDirection,
ring_spec: RingSpec,
},
CancelQueuedLease {
request_id: LeaseRequestId,
edge_id: EdgeId,
},
InstallWorkerRing {
edge_id: EdgeId,
ring_id: RingId,
direction: RingDirection,
layout: RingLayout,
object_spec: ObjectSpec,
ring_spec: RingSpec,
},
UninstallWorkerRing {
edge_id: EdgeId,
ring_id: RingId,
},
EstablishSend {
edge_id: EdgeId,
consumer_node_id: NodeId,
layout: RingLayout,
},
EstablishRecv {
edge_id: EdgeId,
layout: RingLayout,
},
StopPump {
edge_id: EdgeId,
ring_id: RingId,
},
ReleaseArenaLease {
ring_id: RingId,
proof: QuiescenceProof,
},
CopyHotPathBytes {
edge_id: EdgeId,
bytes: Vec<u8>,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum EdgeFaultReason {
RingLeaseRejected(RingLeaseRejection),
RingFault(RingFaultReason),
StreamFault(StreamFaultReason),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum EdgeLifecycleEvent {
EdgeReady {
edge_id: EdgeId,
local_edge_actor: ActorAddress,
},
EdgeFaulted {
edge_id: EdgeId,
reason: EdgeFaultReason,
},
EdgeStopped {
edge_id: EdgeId,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct LocalEdgeRecord {
pub edge_id: EdgeId,
pub direction: RingDirection,
pub state: EdgeProvisionState,
pub lease_request_id: Option<LeaseRequestId>,
pub ring_id: Option<RingId>,
pub peer_node_id: Option<NodeId>,
pub local_edge_actor: ActorAddress,
pub remote_actor_address: Option<ActorAddress>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum EdgeProvisionState {
WaitingForLease,
WaitingForWorkerRing,
WaitingForDriver,
Ready,
Stopping,
Stopped,
Failed,
}
pub struct EdgeEstablisher {
state: EdgeEstablisherState,
}
impl EdgeEstablisher {
pub fn new(local_node_id: NodeId) -> Self {
Self {
state: EdgeEstablisherState::new(local_node_id),
}
}
pub fn observe(&mut self, event: EdgeEvent) {
self.state.observe(event);
}
pub fn commands(&self) -> &[EdgeCommand] {
&self.state.commands
}
pub fn events(&self) -> &[EdgeLifecycleEvent] {
&self.state.events
}
pub fn local_record(&self, edge_id: EdgeId) -> Option<LocalEdgeRecord> {
self.state.records.get(&edge_id).map(EdgeRecord::snapshot)
}
}
struct EdgeEstablisherState {
local_node_id: NodeId,
next_request_id: u64,
next_actor_id: u64,
records: BTreeMap<EdgeId, EdgeRecord>,
commands: Vec<EdgeCommand>,
events: Vec<EdgeLifecycleEvent>,
}
impl EdgeEstablisherState {
fn new(local_node_id: NodeId) -> Self {
Self {
local_node_id,
next_request_id: 1,
next_actor_id: 1,
records: BTreeMap::new(),
commands: Vec::new(),
events: Vec::new(),
}
}
fn observe(&mut self, event: EdgeEvent) {
match event {
EdgeEvent::ProvisionTx(provision) => self.provision_tx(provision),
EdgeEvent::ProvisionRx(provision) => self.provision_rx(provision),
EdgeEvent::RingLeased {
request_id,
ring_id,
layout,
} => self.ring_leased(request_id, ring_id, layout),
EdgeEvent::RingLeaseRejected { request_id, reason } => {
self.ring_lease_rejected(request_id, reason);
}
EdgeEvent::RingInstalled { edge_id, ring_id } => self.ring_installed(edge_id, ring_id),
EdgeEvent::DriverEdgeReady { edge_id } => self.driver_edge_ready(edge_id),
EdgeEvent::StreamFault { edge_id, reason } => {
self.fault_edge(edge_id, EdgeFaultReason::StreamFault(reason));
}
EdgeEvent::RingFault {
edge_id,
ring_id,
reason,
} => self.ring_fault(edge_id, ring_id, reason),
EdgeEvent::StopEdge { edge_id } => self.stop_edge(edge_id),
EdgeEvent::PumpStopped { edge_id, ring_id } => self.pump_stopped(edge_id, ring_id),
EdgeEvent::RingQuiesced { ring_id } | EdgeEvent::QuiescenceProven { ring_id } => {
self.release_after_quiescence(ring_id);
}
EdgeEvent::Stopped { edge_id } => self.mark_stopped(edge_id),
}
}
fn provision_tx(&mut self, provision: ProvisionTx) {
if provision.local_node_id != self.local_node_id {
return;
}
let request_id = self.next_request_id();
let actor = self.next_actor_address();
let edge_id = provision.edge_id;
let ring_spec = provision.ring_spec;
self.records
.insert(edge_id, EdgeRecord::new_tx(provision, request_id, actor));
self.commands.push(EdgeCommand::LeaseRing {
request_id,
edge_id,
direction: RingDirection::Egress,
ring_spec,
});
}
fn provision_rx(&mut self, provision: ProvisionRx) {
if provision.local_node_id != self.local_node_id {
return;
}
let request_id = self.next_request_id();
let actor = self.next_actor_address();
let edge_id = provision.edge_id;
let ring_spec = provision.ring_spec;
self.records
.insert(edge_id, EdgeRecord::new_rx(provision, request_id, actor));
self.commands.push(EdgeCommand::LeaseRing {
request_id,
edge_id,
direction: RingDirection::Ingress,
ring_spec,
});
}
fn ring_leased(&mut self, request_id: LeaseRequestId, ring_id: RingId, layout: RingLayout) {
let Some(edge_id) = self.edge_for_request(request_id) else {
return;
};
let Some(record) = self.records.get_mut(&edge_id) else {
return;
};
if record.state != EdgeProvisionState::WaitingForLease {
self.commands.push(EdgeCommand::ReleaseArenaLease {
ring_id,
proof: QuiescenceProof::verified(),
});
return;
}
record.ring_id = Some(ring_id);
record.layout = Some(layout);
record.state = EdgeProvisionState::WaitingForWorkerRing;
self.commands.push(EdgeCommand::InstallWorkerRing {
edge_id,
ring_id,
direction: record.direction,
layout,
object_spec: record.object_spec,
ring_spec: record.ring_spec,
});
}
fn ring_lease_rejected(&mut self, request_id: LeaseRequestId, reason: RingLeaseRejection) {
let Some(edge_id) = self.edge_for_waiting_request(request_id) else {
return;
};
let Some(record) = self.records.get_mut(&edge_id) else {
return;
};
record.state = EdgeProvisionState::Failed;
self.events.push(EdgeLifecycleEvent::EdgeFaulted {
edge_id,
reason: EdgeFaultReason::RingLeaseRejected(reason),
});
}
fn ring_installed(&mut self, edge_id: EdgeId, ring_id: RingId) {
let Some(record) = self.records.get_mut(&edge_id) else {
return;
};
if record.state != EdgeProvisionState::WaitingForWorkerRing
|| record.ring_id != Some(ring_id)
{
return;
}
let Some(layout) = record.layout else {
return;
};
record.worker_installed = true;
record.driver_established = true;
record.state = EdgeProvisionState::WaitingForDriver;
match record.direction {
RingDirection::Egress => {
if let Some(consumer_node_id) = record.peer_node_id {
self.commands.push(EdgeCommand::EstablishSend {
edge_id,
consumer_node_id,
layout,
});
}
}
RingDirection::Ingress => {
self.commands
.push(EdgeCommand::EstablishRecv { edge_id, layout });
}
}
}
fn driver_edge_ready(&mut self, edge_id: EdgeId) {
let Some(record) = self.records.get_mut(&edge_id) else {
return;
};
if record.state != EdgeProvisionState::WaitingForDriver {
return;
}
record.state = EdgeProvisionState::Ready;
self.events.push(EdgeLifecycleEvent::EdgeReady {
edge_id,
local_edge_actor: record.local_edge_actor,
});
}
fn ring_fault(&mut self, edge_id: EdgeId, ring_id: RingId, reason: RingFaultReason) {
let Some(record) = self.records.get(&edge_id) else {
return;
};
if record.ring_id != Some(ring_id) {
return;
}
self.fault_edge(edge_id, EdgeFaultReason::RingFault(reason));
}
fn fault_edge(&mut self, edge_id: EdgeId, reason: EdgeFaultReason) {
let Some(record) = self.records.get(&edge_id) else {
return;
};
if matches!(
record.state,
EdgeProvisionState::Stopping | EdgeProvisionState::Stopped | EdgeProvisionState::Failed
) {
return;
}
self.events
.push(EdgeLifecycleEvent::EdgeFaulted { edge_id, reason });
self.start_stopping(edge_id, true);
}
fn stop_edge(&mut self, edge_id: EdgeId) {
self.start_stopping(edge_id, true);
}
fn start_stopping(&mut self, edge_id: EdgeId, cancel_lease: bool) {
let Some(record) = self.records.get_mut(&edge_id) else {
return;
};
if matches!(
record.state,
EdgeProvisionState::Stopping | EdgeProvisionState::Stopped
) {
return;
}
let request_id = record.lease_request_id;
let ring_id = record.ring_id;
let driver_established = record.driver_established;
let worker_installed = record.worker_installed;
record.state = EdgeProvisionState::Stopping;
if cancel_lease {
if let Some(request_id) = request_id {
self.commands.push(EdgeCommand::CancelQueuedLease {
request_id,
edge_id,
});
}
}
let Some(ring_id) = ring_id else {
self.mark_stopped_with_event(edge_id);
return;
};
if driver_established {
self.commands
.push(EdgeCommand::StopPump { edge_id, ring_id });
}
if worker_installed {
self.commands
.push(EdgeCommand::UninstallWorkerRing { edge_id, ring_id });
}
if !driver_established && !worker_installed {
self.release_ring(edge_id, ring_id);
}
}
fn pump_stopped(&mut self, edge_id: EdgeId, ring_id: RingId) {
let Some(record) = self.records.get_mut(&edge_id) else {
return;
};
if record.ring_id == Some(ring_id) {
record.pump_stopped = true;
}
}
fn release_after_quiescence(&mut self, ring_id: RingId) {
let Some(edge_id) = self.edge_for_ring(ring_id) else {
return;
};
let Some(record) = self.records.get(&edge_id) else {
return;
};
if record.state != EdgeProvisionState::Stopping {
return;
}
self.release_ring(edge_id, ring_id);
}
fn release_ring(&mut self, edge_id: EdgeId, ring_id: RingId) {
self.commands.push(EdgeCommand::ReleaseArenaLease {
ring_id,
proof: QuiescenceProof::verified(),
});
self.mark_stopped_with_event(edge_id);
}
fn mark_stopped_with_event(&mut self, edge_id: EdgeId) {
let Some(record) = self.records.get_mut(&edge_id) else {
return;
};
if record.state == EdgeProvisionState::Stopped {
return;
}
record.state = EdgeProvisionState::Stopped;
self.events
.push(EdgeLifecycleEvent::EdgeStopped { edge_id });
}
fn mark_stopped(&mut self, edge_id: EdgeId) {
let Some(record) = self.records.get_mut(&edge_id) else {
return;
};
record.state = EdgeProvisionState::Stopped;
}
fn edge_for_request(&self, request_id: LeaseRequestId) -> Option<EdgeId> {
self.records.iter().find_map(|(edge_id, record)| {
(record.lease_request_id == Some(request_id)).then_some(*edge_id)
})
}
fn edge_for_waiting_request(&self, request_id: LeaseRequestId) -> Option<EdgeId> {
self.records.iter().find_map(|(edge_id, record)| {
(record.state == EdgeProvisionState::WaitingForLease
&& record.lease_request_id == Some(request_id))
.then_some(*edge_id)
})
}
fn edge_for_ring(&self, ring_id: RingId) -> Option<EdgeId> {
self.records
.iter()
.find_map(|(edge_id, record)| (record.ring_id == Some(ring_id)).then_some(*edge_id))
}
fn next_request_id(&mut self) -> LeaseRequestId {
let request_id = LeaseRequestId(self.next_request_id);
self.next_request_id += 1;
request_id
}
fn next_actor_address(&mut self) -> ActorAddress {
let actor = ActorAddress(self.next_actor_id);
self.next_actor_id += 1;
actor
}
}
struct EdgeRecord {
run_id: RunId,
edge_id: EdgeId,
direction: RingDirection,
state: EdgeProvisionState,
lease_request_id: Option<LeaseRequestId>,
ring_id: Option<RingId>,
layout: Option<RingLayout>,
peer_node_id: Option<NodeId>,
local_edge_actor: ActorAddress,
remote_actor_address: Option<ActorAddress>,
object_spec: ObjectSpec,
ring_spec: RingSpec,
worker_installed: bool,
driver_established: bool,
pump_stopped: bool,
}
impl EdgeRecord {
fn new_tx(provision: ProvisionTx, request_id: LeaseRequestId, actor: ActorAddress) -> Self {
Self {
run_id: provision.run_id,
edge_id: provision.edge_id,
direction: RingDirection::Egress,
state: EdgeProvisionState::WaitingForLease,
lease_request_id: Some(request_id),
ring_id: None,
layout: None,
peer_node_id: Some(provision.consumer_node_id),
local_edge_actor: actor,
remote_actor_address: None,
object_spec: provision.object_spec,
ring_spec: provision.ring_spec,
worker_installed: false,
driver_established: false,
pump_stopped: false,
}
}
fn new_rx(provision: ProvisionRx, request_id: LeaseRequestId, actor: ActorAddress) -> Self {
Self {
run_id: provision.run_id,
edge_id: provision.edge_id,
direction: RingDirection::Ingress,
state: EdgeProvisionState::WaitingForLease,
lease_request_id: Some(request_id),
ring_id: None,
layout: None,
peer_node_id: None,
local_edge_actor: actor,
remote_actor_address: None,
object_spec: provision.object_spec,
ring_spec: provision.ring_spec,
worker_installed: false,
driver_established: false,
pump_stopped: false,
}
}
fn snapshot(&self) -> LocalEdgeRecord {
let _ = self.run_id;
LocalEdgeRecord {
edge_id: self.edge_id,
direction: self.direction,
state: self.state,
lease_request_id: self.lease_request_id,
ring_id: self.ring_id,
peer_node_id: self.peer_node_id,
local_edge_actor: self.local_edge_actor,
remote_actor_address: self.remote_actor_address,
}
}
}
#[cfg(test)]
pub struct EdgeEstablisherHarness {
establisher: EdgeEstablisher,
}
#[cfg(test)]
impl EdgeEstablisherHarness {
pub fn new(local_node_id: NodeId) -> Self {
Self {
establisher: EdgeEstablisher::new(local_node_id),
}
}
pub fn observe(&mut self, event: EdgeEvent) {
self.establisher.observe(event);
}
pub fn commands(&self) -> &[EdgeCommand] {
self.establisher.commands()
}
pub fn events(&self) -> &[EdgeLifecycleEvent] {
self.establisher.events()
}
pub fn local_record(&self, edge_id: EdgeId) -> Option<LocalEdgeRecord> {
self.establisher.local_record(edge_id)
}
}

View file

@ -0,0 +1,378 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ProcessId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct WorkerGeneration(pub 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 StepId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ObjectId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct DeviceHandle {
pub generation: WorkerGeneration,
pub id: u64,
}
impl DeviceHandle {
pub fn new(generation: WorkerGeneration, id: u64) -> Self {
Self { generation, id }
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ArenaEnv {
pub arena_fd: i32,
pub arena_bytes: u64,
}
impl ArenaEnv {
pub fn test_default() -> Self {
Self {
arena_fd: 3,
arena_bytes: 4096,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WorkerConfig {
pub node_id: NodeId,
pub arena_env: ArenaEnv,
pub initialization_timeout_ms: u64,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ActorCommand {
InstallRing {
generation: WorkerGeneration,
ring_id: RingId,
},
ExecuteStep {
generation: WorkerGeneration,
step_id: StepId,
input: DeviceHandle,
},
ReleaseDeviceObject {
generation: WorkerGeneration,
handle: DeviceHandle,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WorkerCommand {
InitializeWorker {
generation: WorkerGeneration,
},
InstallRing {
ring_id: RingId,
},
ExecuteStep {
step_id: StepId,
input: DeviceHandle,
},
ReleaseDeviceObject {
handle: DeviceHandle,
},
ShutdownWorker,
PayloadBytes {
bytes: Vec<u8>,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WorkerEvent {
RingInstalled { ring_id: RingId },
ObjectLoaded { object_id: ObjectId, sequence: u64 },
ObjectProduced { object_id: ObjectId, sequence: u64 },
StepCompleted { step_id: StepId },
RingReadable { ring_id: RingId },
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ExitStatus {
Code(i32),
Signal(i32),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WorkerCtlEvent {
StartWorker,
ProcessStarted { pid: ProcessId },
WorkerReady { generation: WorkerGeneration },
ActorCommand(ActorCommand),
StdoutEvent(WorkerEvent),
ProcessExited { status: ExitStatus },
RestartRequested,
ShutdownRequested,
WorkerStopped { generation: WorkerGeneration },
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum WorkerFailure {
InitializationTimeout,
ProcessExited,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CommandRejection {
NotRunning,
OldGenerationHandle,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WorkerCtlOut {
WorkerRunning {
generation: WorkerGeneration,
},
WorkerFailed {
generation: WorkerGeneration,
reason: WorkerFailure,
},
CommandRejected {
reason: CommandRejection,
},
RingFaulted {
ring_id: RingId,
},
WorkerStopped {
generation: WorkerGeneration,
},
RingQuiesced {
ring_id: RingId,
},
TerminalStopped {
generation: WorkerGeneration,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WorkerCtlCommand {
SpawnProcessActor { node_id: NodeId },
StopDriverPump { ring_id: RingId },
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum RoutedEvent {
ToEdgeEstablisher(WorkerEvent),
ToRxOrRole(WorkerEvent),
ToTxOrRole(WorkerEvent),
ToStageController(WorkerEvent),
ToDriverOrWorkerSide(WorkerEvent),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum CtlState {
Idle,
Starting,
Running,
Crashed,
ShuttingDown,
Stopped,
}
#[cfg(test)]
pub struct GpuWorkerCtlHarness {
config: WorkerConfig,
state: CtlState,
current_generation: WorkerGeneration,
now_ms: u64,
start_time_ms: Option<u64>,
commands: Vec<WorkerCtlCommand>,
serialized: Vec<WorkerCommand>,
events: Vec<WorkerCtlOut>,
routed: Vec<RoutedEvent>,
installed_rings: std::collections::BTreeSet<RingId>,
}
#[cfg(test)]
impl GpuWorkerCtlHarness {
pub fn new(config: WorkerConfig) -> Self {
Self {
config,
state: CtlState::Idle,
current_generation: WorkerGeneration(1),
now_ms: 0,
start_time_ms: None,
commands: Vec::new(),
serialized: Vec::new(),
events: Vec::new(),
routed: Vec::new(),
installed_rings: std::collections::BTreeSet::new(),
}
}
pub fn observe(&mut self, event: WorkerCtlEvent) {
match event {
WorkerCtlEvent::StartWorker => self.start(),
WorkerCtlEvent::ProcessStarted { .. } => {
self.state = CtlState::Starting;
self.serialized.push(WorkerCommand::InitializeWorker {
generation: self.current_generation,
});
}
WorkerCtlEvent::WorkerReady { generation } => {
self.current_generation = generation;
self.state = CtlState::Running;
self.events.push(WorkerCtlOut::WorkerRunning { generation });
}
WorkerCtlEvent::ActorCommand(command) => self.actor_command(command),
WorkerCtlEvent::StdoutEvent(event) => self.route(event),
WorkerCtlEvent::ProcessExited { status } => self.process_exited(status),
WorkerCtlEvent::RestartRequested => self.restart(),
WorkerCtlEvent::ShutdownRequested => self.shutdown(),
WorkerCtlEvent::WorkerStopped { generation } => {
self.events.push(WorkerCtlOut::WorkerStopped { generation });
}
}
}
pub fn advance_time_ms(&mut self, delta: u64) {
self.now_ms = self.now_ms.saturating_add(delta);
if self.state == CtlState::Starting {
if let Some(start_time) = self.start_time_ms {
if self.now_ms.saturating_sub(start_time) > self.config.initialization_timeout_ms {
self.events.push(WorkerCtlOut::WorkerFailed {
generation: self.current_generation,
reason: WorkerFailure::InitializationTimeout,
});
self.state = CtlState::Crashed;
}
}
}
}
pub fn commands(&self) -> &[WorkerCtlCommand] {
&self.commands
}
pub fn serialized_worker_commands(&self) -> &[WorkerCommand] {
&self.serialized
}
pub fn events(&self) -> &[WorkerCtlOut] {
&self.events
}
pub fn routed(&self) -> &[RoutedEvent] {
&self.routed
}
pub fn current_generation(&self) -> WorkerGeneration {
self.current_generation
}
fn start(&mut self) {
self.state = CtlState::Starting;
self.start_time_ms = Some(self.now_ms);
self.commands.push(WorkerCtlCommand::SpawnProcessActor {
node_id: self.config.node_id,
});
}
fn actor_command(&mut self, command: ActorCommand) {
if self.state != CtlState::Running {
self.events.push(WorkerCtlOut::CommandRejected {
reason: CommandRejection::NotRunning,
});
return;
}
let command_generation = match command {
ActorCommand::InstallRing { generation, .. }
| ActorCommand::ExecuteStep { generation, .. }
| ActorCommand::ReleaseDeviceObject { generation, .. } => generation,
};
if command_generation != self.current_generation {
self.events.push(WorkerCtlOut::CommandRejected {
reason: CommandRejection::OldGenerationHandle,
});
return;
}
match command {
ActorCommand::InstallRing { ring_id, .. } => {
self.installed_rings.insert(ring_id);
self.serialized.push(WorkerCommand::InstallRing { ring_id });
}
ActorCommand::ExecuteStep { step_id, input, .. } => {
if input.generation != self.current_generation {
self.events.push(WorkerCtlOut::CommandRejected {
reason: CommandRejection::OldGenerationHandle,
});
} else {
self.serialized
.push(WorkerCommand::ExecuteStep { step_id, input });
}
}
ActorCommand::ReleaseDeviceObject { handle, .. } => {
if handle.generation != self.current_generation {
self.events.push(WorkerCtlOut::CommandRejected {
reason: CommandRejection::OldGenerationHandle,
});
} else {
self.serialized
.push(WorkerCommand::ReleaseDeviceObject { handle });
}
}
}
}
fn route(&mut self, event: WorkerEvent) {
match event.clone() {
WorkerEvent::RingInstalled { .. } => {
self.routed.push(RoutedEvent::ToEdgeEstablisher(event))
}
WorkerEvent::ObjectLoaded { .. } => self.routed.push(RoutedEvent::ToRxOrRole(event)),
WorkerEvent::ObjectProduced { .. } => self.routed.push(RoutedEvent::ToTxOrRole(event)),
WorkerEvent::StepCompleted { .. } => {
self.routed.push(RoutedEvent::ToStageController(event))
}
WorkerEvent::RingReadable { .. } => {
self.routed.push(RoutedEvent::ToDriverOrWorkerSide(event))
}
}
}
fn process_exited(&mut self, status: ExitStatus) {
match (self.state, status) {
(CtlState::ShuttingDown, ExitStatus::Code(0)) => {
for ring_id in &self.installed_rings {
self.events
.push(WorkerCtlOut::RingQuiesced { ring_id: *ring_id });
}
self.events.push(WorkerCtlOut::TerminalStopped {
generation: self.current_generation,
});
self.state = CtlState::Stopped;
}
(_, _) => {
for ring_id in &self.installed_rings {
self.events
.push(WorkerCtlOut::RingFaulted { ring_id: *ring_id });
self.commands
.push(WorkerCtlCommand::StopDriverPump { ring_id: *ring_id });
}
self.events.push(WorkerCtlOut::WorkerFailed {
generation: self.current_generation,
reason: WorkerFailure::ProcessExited,
});
self.state = CtlState::Crashed;
}
}
}
fn restart(&mut self) {
self.state = CtlState::Starting;
self.current_generation = WorkerGeneration(self.current_generation.0 + 1);
self.start_time_ms = Some(self.now_ms);
}
fn shutdown(&mut self) {
if self.state == CtlState::Running {
self.serialized.push(WorkerCommand::ShutdownWorker);
self.state = CtlState::ShuttingDown;
}
}
}

View file

@ -0,0 +1,391 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct WorkerGeneration(pub 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 EdgeId(pub u64);
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct PortId(pub String);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ObjectId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct StepId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct DeviceHandle {
pub generation: WorkerGeneration,
pub id: u64,
}
impl DeviceHandle {
pub fn new(generation: WorkerGeneration, id: u64) -> Self {
Self { generation, id }
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RingDirection {
Ingress,
Egress,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ObjectLayout {
Token,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ObjectSpec {
pub max_extent: u64,
pub alignment: u64,
pub layout: ObjectLayout,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct InstallRing {
pub ring_id: RingId,
pub edge_id: EdgeId,
pub port_id: PortId,
pub direction: RingDirection,
pub object_spec: ObjectSpec,
pub generation: WorkerGeneration,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
pub struct ObjectFlags {
pub end_of_sequence: bool,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct OutputBinding {
pub ring_id: RingId,
pub object_id: ObjectId,
pub sequence: u64,
pub extent: u64,
pub flags: ObjectFlags,
pub device_source: DeviceHandle,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WorkerEgressEvent {
InstallRing(InstallRing),
ExecuteStep {
step_id: StepId,
outputs: Vec<OutputBinding>,
},
HeaderReady {
object_id: ObjectId,
},
EgressRingFull {
ring_id: RingId,
},
RingWritable {
ring_id: RingId,
},
DeviceToHostCopyCompleted {
object_id: ObjectId,
byte_count: u64,
},
DeviceCopyFailed {
object_id: ObjectId,
},
RoleStateUpdated {
step_id: StepId,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum StepFailureReason {
InvalidOutputRing,
OutputExtentViolation,
DeviceCopyFailed,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WorkerEgressOut {
ObjectProduced {
ring_id: RingId,
object_id: ObjectId,
sequence: u64,
},
StepCompleted {
step_id: StepId,
},
StepFailed {
step_id: StepId,
reason: StepFailureReason,
},
RingFault {
ring_id: RingId,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WakeHint {
RingReadable { ring_id: RingId },
RingWritable { ring_id: RingId },
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ObjectHeader {
pub object_id: ObjectId,
pub sequence: u64,
pub extent: u64,
}
impl ObjectHeader {
pub fn decode(bytes: &[u8]) -> Result<Self, HeaderDecodeError> {
if bytes.len() < HEADER_LEN || &bytes[0..4] != b"MO01" || bytes[4] != 1 {
return Err(HeaderDecodeError);
}
Ok(Self {
object_id: ObjectId(u64::from_le_bytes(bytes[8..16].try_into().unwrap())),
sequence: u64::from_le_bytes(bytes[16..24].try_into().unwrap()),
extent: u64::from_le_bytes(bytes[24..32].try_into().unwrap()),
})
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct HeaderDecodeError;
const HEADER_LEN: usize = 48;
#[derive(Clone, Debug, PartialEq, Eq)]
struct PendingStep {
step_id: StepId,
outputs: Vec<OutputBinding>,
role_state_updated: bool,
}
#[cfg(test)]
pub struct EgressProducerHarness {
generation: WorkerGeneration,
rings: std::collections::BTreeMap<RingId, InstallRing>,
pending_outputs: Vec<OutputBinding>,
steps: Vec<PendingStep>,
committed: std::collections::BTreeMap<RingId, Vec<u8>>,
payload_committed: std::collections::BTreeMap<RingId, u64>,
full_rings: std::collections::BTreeSet<RingId>,
produced: std::collections::BTreeSet<ObjectId>,
wake_hints: Vec<WakeHint>,
events: Vec<WorkerEgressOut>,
}
#[cfg(test)]
impl EgressProducerHarness {
pub fn new(generation: WorkerGeneration) -> Self {
Self {
generation,
rings: std::collections::BTreeMap::new(),
pending_outputs: Vec::new(),
steps: Vec::new(),
committed: std::collections::BTreeMap::new(),
payload_committed: std::collections::BTreeMap::new(),
full_rings: std::collections::BTreeSet::new(),
produced: std::collections::BTreeSet::new(),
wake_hints: Vec::new(),
events: Vec::new(),
}
}
pub fn observe(&mut self, event: WorkerEgressEvent) {
match event {
WorkerEgressEvent::InstallRing(install) => {
if install.direction == RingDirection::Egress
&& install.generation == self.generation
{
self.rings.insert(install.ring_id, install);
}
}
WorkerEgressEvent::ExecuteStep { step_id, outputs } => {
self.execute_step(step_id, outputs)
}
WorkerEgressEvent::HeaderReady { object_id } => self.header_ready(object_id),
WorkerEgressEvent::EgressRingFull { ring_id } => {
self.full_rings.insert(ring_id);
}
WorkerEgressEvent::RingWritable { ring_id } => {
self.full_rings.remove(&ring_id);
self.wake_hints.push(WakeHint::RingWritable { ring_id });
}
WorkerEgressEvent::DeviceToHostCopyCompleted {
object_id,
byte_count,
} => self.copy_completed(object_id, byte_count),
WorkerEgressEvent::DeviceCopyFailed { object_id } => self.copy_failed(object_id),
WorkerEgressEvent::RoleStateUpdated { step_id } => {
if let Some(step) = self.steps.iter_mut().find(|step| step.step_id == step_id) {
step.role_state_updated = true;
}
self.maybe_step_completed(step_id);
}
}
}
pub fn pending_outputs(&self) -> &[OutputBinding] {
&self.pending_outputs
}
pub fn committed_bytes(&self, ring_id: RingId) -> &[u8] {
self.committed
.get(&ring_id)
.map(Vec::as_slice)
.unwrap_or(&[])
}
pub fn committed_payload_bytes(&self, ring_id: RingId) -> u64 {
self.payload_committed.get(&ring_id).copied().unwrap_or(0)
}
pub fn wake_hints(&self) -> &[WakeHint] {
&self.wake_hints
}
pub fn events(&self) -> &[WorkerEgressOut] {
&self.events
}
pub fn complete_output(&mut self, object_id: ObjectId) {
self.header_ready(object_id);
let Some(output) = self
.pending_outputs
.iter()
.find(|output| output.object_id == object_id)
.copied()
else {
return;
};
self.copy_completed(object_id, output.extent);
}
fn execute_step(&mut self, step_id: StepId, outputs: Vec<OutputBinding>) {
for output in &outputs {
let Some(ring) = self.rings.get(&output.ring_id) else {
self.events.push(WorkerEgressOut::StepFailed {
step_id,
reason: StepFailureReason::InvalidOutputRing,
});
return;
};
if output.extent > ring.object_spec.max_extent
|| (ring.object_spec.alignment != 0
&& output.extent % ring.object_spec.alignment != 0)
{
self.events.push(WorkerEgressOut::StepFailed {
step_id,
reason: StepFailureReason::OutputExtentViolation,
});
return;
}
}
self.pending_outputs.extend(outputs.iter().copied());
self.steps.push(PendingStep {
step_id,
outputs,
role_state_updated: false,
});
}
fn header_ready(&mut self, object_id: ObjectId) {
let Some(output) = self
.pending_outputs
.iter()
.find(|output| output.object_id == object_id)
.copied()
else {
return;
};
let bytes = encode_header(output);
self.committed
.entry(output.ring_id)
.or_default()
.extend(bytes);
self.wake_hints.push(WakeHint::RingReadable {
ring_id: output.ring_id,
});
}
fn copy_completed(&mut self, object_id: ObjectId, byte_count: u64) {
let Some(output) = self
.pending_outputs
.iter()
.find(|output| output.object_id == object_id)
.copied()
else {
return;
};
if self.full_rings.contains(&output.ring_id) || byte_count != output.extent {
return;
}
self.committed
.entry(output.ring_id)
.or_default()
.extend(std::iter::repeat(0).take(byte_count as usize));
*self.payload_committed.entry(output.ring_id).or_insert(0) += byte_count;
if self.produced.insert(object_id) {
self.events.push(WorkerEgressOut::ObjectProduced {
ring_id: output.ring_id,
object_id,
sequence: output.sequence,
});
}
let step_ids = self
.steps
.iter()
.filter(|step| {
step.outputs
.iter()
.any(|output| output.object_id == object_id)
})
.map(|step| step.step_id)
.collect::<Vec<_>>();
for step_id in step_ids {
self.maybe_step_completed(step_id);
}
}
fn copy_failed(&mut self, object_id: ObjectId) {
let step_id = self
.steps
.iter()
.find(|step| {
step.outputs
.iter()
.any(|output| output.object_id == object_id)
})
.map(|step| step.step_id)
.unwrap_or(StepId(0));
self.events.push(WorkerEgressOut::StepFailed {
step_id,
reason: StepFailureReason::DeviceCopyFailed,
});
}
fn maybe_step_completed(&mut self, step_id: StepId) {
let Some(step) = self.steps.iter().find(|step| step.step_id == step_id) else {
return;
};
if !step.role_state_updated {
return;
}
if step.outputs.iter().all(|output| self.produced.contains(&output.object_id))
&& !self.events.iter().any(|event| matches!(event, WorkerEgressOut::StepCompleted { step_id: seen } if *seen == step_id))
{
self.events.push(WorkerEgressOut::StepCompleted { step_id });
}
}
}
fn encode_header(output: OutputBinding) -> Vec<u8> {
let mut out = Vec::with_capacity(HEADER_LEN);
out.extend_from_slice(b"MO01");
out.push(1);
out.push(HEADER_LEN as u8);
out.extend_from_slice(&[0u8; 2]);
out.extend_from_slice(&output.object_id.0.to_le_bytes());
out.extend_from_slice(&output.sequence.to_le_bytes());
out.extend_from_slice(&output.extent.to_le_bytes());
out.extend_from_slice(&[0u8; 16]);
out
}

View file

@ -0,0 +1,425 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct WorkerGeneration(pub 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 EdgeId(pub u64);
#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct PortId(pub String);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ObjectId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct DeviceHandle {
pub generation: WorkerGeneration,
pub id: u64,
}
impl DeviceHandle {
pub fn new(generation: WorkerGeneration, id: u64) -> Self {
Self { generation, id }
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RingDirection {
Ingress,
Egress,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ObjectLayout {
Token,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ObjectSpec {
pub max_extent: u64,
pub alignment: u64,
pub layout: ObjectLayout,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct InstallRing {
pub ring_id: RingId,
pub edge_id: EdgeId,
pub port_id: PortId,
pub direction: RingDirection,
pub object_spec: ObjectSpec,
pub generation: WorkerGeneration,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WorkerIngressEvent {
InstallRing(InstallRing),
RingReadable {
ring_id: RingId,
},
Eof {
ring_id: RingId,
},
DeviceCopyCompleted {
object_id: ObjectId,
byte_count: u64,
},
DeviceHandleCreated {
object_id: ObjectId,
handle: DeviceHandle,
},
RingFault {
ring_id: RingId,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ObjectFailureReason {
UnsupportedMagic,
UnsupportedVersion,
MalformedHeaderLength,
ExtentExceedsMax,
ExtentAlignmentViolation,
SequenceViolation,
EofBeforeFullPayload,
DeviceCopyFailed,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WorkerIngressOut {
ObjectLoaded {
ring_id: RingId,
edge_id: EdgeId,
port_id: PortId,
object_id: ObjectId,
sequence: u64,
extent: u64,
handle: DeviceHandle,
},
ObjectFailed {
ring_id: RingId,
object_id: Option<ObjectId>,
reason: ObjectFailureReason,
},
RingFault {
ring_id: RingId,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct DeviceCopyLog {
pub object_id: ObjectId,
pub byte_count: u64,
}
#[derive(Clone, Debug)]
pub struct ObjectRecordBuilder {
spec: ObjectSpec,
object_id: ObjectId,
sequence: u64,
extent: u64,
payload: Vec<u8>,
magic: [u8; 4],
version: u8,
header_len: u8,
}
impl ObjectRecordBuilder {
pub fn new(spec: ObjectSpec) -> Self {
Self {
spec,
object_id: ObjectId(9000),
sequence: 0,
extent: 0,
payload: Vec::new(),
magic: *b"MO01",
version: 1,
header_len: HEADER_LEN as u8,
}
}
pub fn object_id(mut self, object_id: ObjectId) -> Self {
self.object_id = object_id;
self
}
pub fn sequence(mut self, sequence: u64) -> Self {
self.sequence = sequence;
self
}
pub fn extent(mut self, extent: u64) -> Self {
self.extent = extent;
self
}
pub fn payload(mut self, payload: Vec<u8>) -> Self {
self.payload = payload;
self
}
pub fn partial_payload(mut self, payload: Vec<u8>) -> Self {
self.payload = payload;
self
}
pub fn unsupported_magic(mut self) -> Self {
self.magic = *b"BAD!";
self
}
pub fn unsupported_version(mut self) -> Self {
self.version = 99;
self
}
pub fn malformed_header_length(mut self) -> Self {
self.header_len = 1;
self
}
pub fn encode(mut self) -> Vec<u8> {
if self.extent == 0 && !self.payload.is_empty() {
self.extent = self.payload.len() as u64;
}
let mut out = Vec::with_capacity(HEADER_LEN + self.payload.len());
out.extend_from_slice(&self.magic);
out.push(self.version);
out.push(self.header_len);
out.extend_from_slice(&[0u8; 2]);
out.extend_from_slice(&self.object_id.0.to_le_bytes());
out.extend_from_slice(&self.sequence.to_le_bytes());
out.extend_from_slice(&self.extent.to_le_bytes());
out.extend_from_slice(&self.spec.max_extent.to_le_bytes());
out.extend_from_slice(&self.spec.alignment.to_le_bytes());
out.extend_from_slice(&self.payload);
out
}
}
const HEADER_LEN: usize = 48;
#[derive(Clone, Debug, PartialEq, Eq)]
struct ParsedRecord {
object_id: ObjectId,
sequence: u64,
extent: u64,
total_len: usize,
}
#[derive(Clone, Debug, PartialEq, Eq)]
struct PendingObject {
record: ParsedRecord,
copy_done: bool,
handle: Option<DeviceHandle>,
}
#[cfg(test)]
pub struct IngressParserHarness {
generation: WorkerGeneration,
install: Option<InstallRing>,
buffers: std::collections::BTreeMap<RingId, Vec<u8>>,
consume: std::collections::BTreeMap<RingId, u64>,
cursor_reload: std::collections::BTreeMap<RingId, u64>,
faulted_rings: std::collections::BTreeSet<RingId>,
expected_sequence: u64,
pending: std::collections::BTreeMap<ObjectId, PendingObject>,
copy_log: Vec<DeviceCopyLog>,
events: Vec<WorkerIngressOut>,
}
#[cfg(test)]
impl IngressParserHarness {
pub fn new(generation: WorkerGeneration) -> Self {
Self {
generation,
install: None,
buffers: std::collections::BTreeMap::new(),
consume: std::collections::BTreeMap::new(),
cursor_reload: std::collections::BTreeMap::new(),
faulted_rings: std::collections::BTreeSet::new(),
expected_sequence: 0,
pending: std::collections::BTreeMap::new(),
copy_log: Vec::new(),
events: Vec::new(),
}
}
pub fn observe(&mut self, event: WorkerIngressEvent) {
match event {
WorkerIngressEvent::InstallRing(install) => {
if install.direction == RingDirection::Ingress
&& install.generation == self.generation
{
self.consume.entry(install.ring_id).or_insert(0);
self.install = Some(install);
}
}
WorkerIngressEvent::RingReadable { ring_id } => self.parse_ring(ring_id, false),
WorkerIngressEvent::Eof { ring_id } => self.parse_ring(ring_id, true),
WorkerIngressEvent::DeviceCopyCompleted {
object_id,
byte_count,
} => {
self.copy_log.push(DeviceCopyLog {
object_id,
byte_count,
});
if let Some(pending) = self.pending.get_mut(&object_id) {
pending.copy_done = true;
if byte_count == pending.record.extent {
if let Some(install) = &self.install {
*self.consume.entry(install.ring_id).or_insert(0) +=
pending.record.total_len as u64;
}
}
}
self.maybe_loaded(object_id);
}
WorkerIngressEvent::DeviceHandleCreated { object_id, handle } => {
if let Some(pending) = self.pending.get_mut(&object_id) {
pending.handle = Some(handle);
}
self.maybe_loaded(object_id);
}
WorkerIngressEvent::RingFault { ring_id } => {
self.faulted_rings.insert(ring_id);
self.events.push(WorkerIngressOut::RingFault { ring_id });
}
}
}
pub fn write_committed_bytes(&mut self, ring_id: RingId, bytes: Vec<u8>) {
self.buffers.entry(ring_id).or_default().extend(bytes);
}
pub fn write_uncommitted_bytes(&mut self, _ring_id: RingId, _bytes: Vec<u8>) {}
pub fn consume_cursor(&self, ring_id: RingId) -> u64 {
self.consume.get(&ring_id).copied().unwrap_or(0)
}
pub fn cursor_reload_count(&self, ring_id: RingId) -> u64 {
self.cursor_reload.get(&ring_id).copied().unwrap_or(0)
}
pub fn device_copy_log(&self) -> &[DeviceCopyLog] {
&self.copy_log
}
pub fn events(&self) -> &[WorkerIngressOut] {
&self.events
}
fn parse_ring(&mut self, ring_id: RingId, eof: bool) {
let Some(install) = &self.install else {
return;
};
if install.ring_id != ring_id || self.faulted_rings.contains(&ring_id) {
return;
}
*self.cursor_reload.entry(ring_id).or_insert(0) += 1;
let buffer = self.buffers.get(&ring_id).cloned().unwrap_or_default();
if buffer.is_empty() {
return;
}
match decode_record(&buffer, install.object_spec, eof) {
Ok(record) => {
if record.sequence != self.expected_sequence {
self.events.push(WorkerIngressOut::ObjectFailed {
ring_id,
object_id: Some(record.object_id),
reason: ObjectFailureReason::SequenceViolation,
});
return;
}
self.expected_sequence += 1;
self.pending.insert(
record.object_id,
PendingObject {
record: record.clone(),
copy_done: false,
handle: None,
},
);
self.copy_log.push(DeviceCopyLog {
object_id: record.object_id,
byte_count: record.extent,
});
}
Err(reason) => self.events.push(WorkerIngressOut::ObjectFailed {
ring_id,
object_id: None,
reason,
}),
}
}
fn maybe_loaded(&mut self, object_id: ObjectId) {
let Some(pending) = self.pending.get(&object_id).cloned() else {
return;
};
let Some(handle) = pending.handle else {
return;
};
if !pending.copy_done || handle.generation != self.generation {
return;
}
let install = self.install.as_ref().unwrap();
if !self.events.iter().any(|event| matches!(event, WorkerIngressOut::ObjectLoaded { object_id: seen, .. } if *seen == object_id)) {
self.events.push(WorkerIngressOut::ObjectLoaded {
ring_id: install.ring_id,
edge_id: install.edge_id,
port_id: install.port_id.clone(),
object_id,
sequence: pending.record.sequence,
extent: pending.record.extent,
handle,
});
}
}
}
fn decode_record(
bytes: &[u8],
spec: ObjectSpec,
eof: bool,
) -> Result<ParsedRecord, ObjectFailureReason> {
if bytes.len() < HEADER_LEN {
return if eof {
Err(ObjectFailureReason::EofBeforeFullPayload)
} else {
Err(ObjectFailureReason::MalformedHeaderLength)
};
}
if &bytes[0..4] != b"MO01" {
return Err(ObjectFailureReason::UnsupportedMagic);
}
if bytes[4] != 1 {
return Err(ObjectFailureReason::UnsupportedVersion);
}
if bytes[5] as usize != HEADER_LEN {
return Err(ObjectFailureReason::MalformedHeaderLength);
}
let object_id = ObjectId(u64::from_le_bytes(bytes[8..16].try_into().unwrap()));
let sequence = u64::from_le_bytes(bytes[16..24].try_into().unwrap());
let extent = u64::from_le_bytes(bytes[24..32].try_into().unwrap());
if extent > spec.max_extent {
return Err(ObjectFailureReason::ExtentExceedsMax);
}
if spec.alignment != 0 && extent % spec.alignment != 0 {
return Err(ObjectFailureReason::ExtentAlignmentViolation);
}
let total_len = HEADER_LEN + extent as usize;
if bytes.len() < total_len {
return if eof {
Err(ObjectFailureReason::EofBeforeFullPayload)
} else {
Err(ObjectFailureReason::MalformedHeaderLength)
};
}
Ok(ParsedRecord {
object_id,
sequence,
extent,
total_len,
})
}

View file

@ -0,0 +1,290 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ArenaFd(pub i32);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct HelperAbiVersion(pub u32);
#[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 StepId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct DeviceHandle(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct WorkerGeneration(pub u64);
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct AdapterConfig {
pub arena_fd: ArenaFd,
pub arena_bytes: u64,
pub helper_abi_version: HelperAbiVersion,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WorkerCommand {
InitializeWorker {
helper_abi_version: HelperAbiVersion,
},
InstallRing {
ring_id: RingId,
},
ExecuteStep {
step_id: StepId,
},
ReleaseDeviceObject {
handle: DeviceHandle,
},
ShutdownWorker,
RingReadable {
ring_id: RingId,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct JsonLine(serde_json::Map<String, serde_json::Value>);
impl JsonLine {
pub fn parse(line: &str) -> Result<Self, serde_json::Error> {
let value: serde_json::Value = serde_json::from_str(line.trim_end())?;
match value {
serde_json::Value::Object(map) => Ok(Self(map)),
other => Ok(Self(serde_json::Map::from_iter([("value".into(), other)]))),
}
}
pub fn is_object(&self) -> bool {
true
}
pub fn contains_key(&self, key: &str) -> bool {
self.0.contains_key(key)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ProcessFaultReason {
InvalidJson,
UnknownEventShape,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum WorkerFatalReason {
UnsupportedHelperAbi,
BackendUnavailable,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum AdapterEvent {
WorkerReady {
generation: WorkerGeneration,
},
WorkerFatal {
reason: WorkerFatalReason,
},
ProcessFault {
reason: ProcessFaultReason,
line: String,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WorkerAction {
ReadArenaEnvironment {
arena_fd: ArenaFd,
arena_bytes: u64,
},
MapArena {
arena_fd: ArenaFd,
arena_bytes: u64,
},
InitializeRingHelper {
helper_abi_version: HelperAbiVersion,
},
InitializeBackend,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum InitializationFailure {
BackendUnavailable,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ExitStatus {
Code(i32),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CommandRejectionReason {
PayloadBytesForbidden,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct CommandRejection {
pub reason: CommandRejectionReason,
}
#[cfg(test)]
pub struct ProcessAdapterHarness {
config: AdapterConfig,
stdin_lines: Vec<String>,
events: Vec<AdapterEvent>,
worker_actions: Vec<WorkerAction>,
command_rejections: Vec<CommandRejection>,
exit_status: Option<ExitStatus>,
}
#[cfg(test)]
impl ProcessAdapterHarness {
pub fn new(config: AdapterConfig) -> Self {
Self {
config,
stdin_lines: Vec::new(),
events: Vec::new(),
worker_actions: Vec::new(),
command_rejections: Vec::new(),
exit_status: None,
}
}
pub fn start_worker_process(&mut self) {
self.worker_actions
.push(WorkerAction::ReadArenaEnvironment {
arena_fd: self.config.arena_fd,
arena_bytes: self.config.arena_bytes,
});
}
pub fn send_command(&mut self, command: WorkerCommand) {
if let WorkerCommand::InitializeWorker { helper_abi_version } = command {
if helper_abi_version != self.config.helper_abi_version {
self.events.push(AdapterEvent::WorkerFatal {
reason: WorkerFatalReason::UnsupportedHelperAbi,
});
return;
}
self.worker_actions.push(WorkerAction::MapArena {
arena_fd: self.config.arena_fd,
arena_bytes: self.config.arena_bytes,
});
self.worker_actions
.push(WorkerAction::InitializeRingHelper { helper_abi_version });
self.worker_actions.push(WorkerAction::InitializeBackend);
}
self.stdin_lines.push(format_json_command(&command));
}
pub fn send_raw_json_command(&mut self, raw: &str) {
match JsonLine::parse(raw) {
Ok(parsed)
if parsed.contains_key("payload")
|| parsed.contains_key("bytes")
|| parsed.contains_key("data") =>
{
self.command_rejections.push(CommandRejection {
reason: CommandRejectionReason::PayloadBytesForbidden,
});
}
Ok(_) => self.stdin_lines.push(format!("{}\n", raw.trim_end())),
Err(_) => self.command_rejections.push(CommandRejection {
reason: CommandRejectionReason::PayloadBytesForbidden,
}),
}
}
pub fn receive_stdout_line(&mut self, line: &str) {
let parsed = match JsonLine::parse(line) {
Ok(parsed) => parsed,
Err(_) => {
self.events.push(AdapterEvent::ProcessFault {
reason: ProcessFaultReason::InvalidJson,
line: line.into(),
});
return;
}
};
let Some(serde_json::Value::String(kind)) = parsed.0.get("type") else {
self.events.push(AdapterEvent::ProcessFault {
reason: ProcessFaultReason::UnknownEventShape,
line: line.into(),
});
return;
};
match kind.as_str() {
"WorkerReady" => {
let generation = parsed
.0
.get("generation")
.and_then(|value| value.as_u64())
.unwrap_or(1);
self.events.push(AdapterEvent::WorkerReady {
generation: WorkerGeneration(generation),
});
}
_ => self.events.push(AdapterEvent::ProcessFault {
reason: ProcessFaultReason::UnknownEventShape,
line: line.into(),
}),
}
}
pub fn receive_stderr_line(&mut self, _line: &str) {}
pub fn inject_initialization_failure(&mut self, failure: InitializationFailure) {
match failure {
InitializationFailure::BackendUnavailable => {
self.events.push(AdapterEvent::WorkerFatal {
reason: WorkerFatalReason::BackendUnavailable,
});
self.exit_status = Some(ExitStatus::Code(1));
}
}
}
pub fn stdin_lines(&self) -> &[String] {
&self.stdin_lines
}
pub fn events(&self) -> &[AdapterEvent] {
&self.events
}
pub fn worker_actions(&self) -> &[WorkerAction] {
&self.worker_actions
}
pub fn exit_status(&self) -> Option<ExitStatus> {
self.exit_status
}
pub fn command_rejections(&self) -> &[CommandRejection] {
&self.command_rejections
}
}
#[cfg(test)]
fn format_json_command(command: &WorkerCommand) -> String {
let value = match command {
WorkerCommand::InitializeWorker { helper_abi_version } => serde_json::json!({
"type": "InitializeWorker",
"helper_abi_version": helper_abi_version.0,
}),
WorkerCommand::InstallRing { ring_id } => serde_json::json!({
"type": "InstallRing",
"ring_id": ring_id.0,
}),
WorkerCommand::ExecuteStep { step_id } => serde_json::json!({
"type": "ExecuteStep",
"step_id": step_id.0,
}),
WorkerCommand::ReleaseDeviceObject { handle } => serde_json::json!({
"type": "ReleaseDeviceObject",
"handle": handle.0,
}),
WorkerCommand::ShutdownWorker => serde_json::json!({
"type": "ShutdownWorker",
}),
WorkerCommand::RingReadable { ring_id } => serde_json::json!({
"type": "RingReadable",
"ring_id": ring_id.0,
}),
};
format!("{}\n", value)
}

View file

@ -1 +1,29 @@
#[cfg(test)]
extern crate self as mvp_system;
pub mod arena_manager;
pub mod device_bridge;
pub mod driver_pumps;
pub mod edge_establisher;
pub mod gpu_worker_ctl;
pub mod gpu_worker_egress_producer;
pub mod gpu_worker_ingress_parser;
pub mod gpu_worker_process_adapter;
pub mod membership_pool_readiness;
pub mod node_boot_lifecycle;
pub mod observability_surface;
pub mod orchestrator_run_fsm;
pub mod orchestrator_token_endpoint;
pub mod resource_inventory;
pub mod run_plan;
pub mod shared_ring_helper_abi;
pub mod stage_controller;
pub mod tx_rx_edge_actor;
pub mod weight_lifecycle;
#[cfg(test)]
mod tests;
#[cfg(test)]
#[path = "tests/driver_pumps_guarantees.rs"]
mod driver_pumps_guarantees;

View file

@ -0,0 +1,248 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct RunId(pub u64);
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PoolId(pub String);
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ReadinessConfig {
pub pool_id: PoolId,
pub convergence_window_ms: u64,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RunRequest {
pub run_id: RunId,
}
impl RunRequest {
pub fn new(run_id: RunId) -> Self {
Self { run_id }
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Observation {
NodeKnown {
node_id: NodeId,
},
SwimLive {
node_id: NodeId,
},
NodeAvailable {
node_id: NodeId,
},
DataPlaneIdentityReady {
node_id: NodeId,
},
SwimSuspect {
node_id: NodeId,
},
NodeFaulted {
node_id: NodeId,
},
DataPlaneIdentityMissing {
node_id: NodeId,
},
SwimLost {
node_id: NodeId,
},
RunProvisioned {
run_id: RunId,
required_nodes: Vec<NodeId>,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RunFaultReason {
RequiredNodeLost { node_id: NodeId },
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ReadinessEvent {
PoolReady {
pool: Vec<NodeId>,
},
RunFaulted {
run_id: RunId,
reason: RunFaultReason,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ReadinessCommand {
EmitPoolReady { pool: Vec<NodeId> },
StartPlanning { run_id: RunId, pool: Vec<NodeId> },
WaitForStability { pool_id: PoolId },
AbortPendingRun { run_id: RunId },
CommitRunPlan { run_id: RunId },
RecomputePlacement { run_id: RunId },
AssignStage { node_id: NodeId },
AssignEdge { node_id: NodeId },
AssignLayerRange { node_id: NodeId },
AssignObjectSpec { node_id: NodeId },
}
#[derive(Clone, Default)]
struct NodeFacts {
known: bool,
live: bool,
available: bool,
identity_ready: bool,
poisoned: bool,
}
#[cfg(test)]
pub struct ReadinessGateHarness {
config: ReadinessConfig,
candidates: Vec<NodeId>,
facts: std::collections::BTreeMap<NodeId, NodeFacts>,
events: Vec<ReadinessEvent>,
commands: Vec<ReadinessCommand>,
now_ms: u64,
stable_since_ms: Option<u64>,
emitted_ready: bool,
pending_run: Option<RunRequest>,
provisioned_run: Option<(RunId, Vec<NodeId>)>,
}
#[cfg(test)]
impl ReadinessGateHarness {
pub fn new(config: ReadinessConfig, candidates: Vec<NodeId>) -> Self {
let facts = candidates
.iter()
.copied()
.map(|node_id| (node_id, NodeFacts::default()))
.collect();
Self {
config,
candidates,
facts,
events: Vec::new(),
commands: Vec::new(),
now_ms: 0,
stable_since_ms: None,
emitted_ready: false,
pending_run: None,
provisioned_run: None,
}
}
pub fn observe(&mut self, observation: Observation) {
match observation {
Observation::NodeKnown { node_id } => self.fact(node_id).known = true,
Observation::SwimLive { node_id } => self.fact(node_id).live = true,
Observation::NodeAvailable { node_id } => self.fact(node_id).available = true,
Observation::DataPlaneIdentityReady { node_id } => {
self.fact(node_id).identity_ready = true
}
Observation::SwimSuspect { node_id }
| Observation::NodeFaulted { node_id }
| Observation::DataPlaneIdentityMissing { node_id } => {
self.fact(node_id).poisoned = true
}
Observation::SwimLost { node_id } => {
self.fact(node_id).live = false;
self.fact(node_id).poisoned = true;
if let Some((run_id, required)) = &self.provisioned_run {
if required.contains(&node_id)
&& !self.events.iter().any(|event| {
matches!(event, ReadinessEvent::RunFaulted { run_id: seen, .. } if seen == run_id)
})
{
self.events.push(ReadinessEvent::RunFaulted {
run_id: *run_id,
reason: RunFaultReason::RequiredNodeLost { node_id },
});
}
} else if let Some(pending) = &self.pending_run {
self.commands.push(ReadinessCommand::AbortPendingRun {
run_id: pending.run_id,
});
}
}
Observation::RunProvisioned {
run_id,
required_nodes,
} => {
self.provisioned_run = Some((run_id, required_nodes));
}
}
self.reset_stability_if_needed();
self.maybe_ready();
}
pub fn advance_time_ms(&mut self, delta: u64) {
self.now_ms = self.now_ms.saturating_add(delta);
self.maybe_ready();
}
pub fn request_run_planning(&mut self, request: RunRequest) {
self.pending_run = Some(request);
self.maybe_ready();
}
pub fn events(&self) -> &[ReadinessEvent] {
&self.events
}
pub fn commands(&self) -> &[ReadinessCommand] {
&self.commands
}
fn fact(&mut self, node_id: NodeId) -> &mut NodeFacts {
self.facts.entry(node_id).or_default()
}
fn complete_now(&self) -> bool {
self.candidates.iter().all(|node_id| {
self.facts.get(node_id).is_some_and(|facts| {
facts.known
&& facts.live
&& facts.available
&& facts.identity_ready
&& !facts.poisoned
})
})
}
fn reset_stability_if_needed(&mut self) {
if self.complete_now() {
if self.stable_since_ms.is_none() {
self.stable_since_ms = Some(self.now_ms);
self.commands.push(ReadinessCommand::WaitForStability {
pool_id: self.config.pool_id.clone(),
});
}
} else {
self.stable_since_ms = None;
}
}
fn maybe_ready(&mut self) {
if self.emitted_ready || !self.complete_now() {
return;
}
let Some(stable_since) = self.stable_since_ms else {
return;
};
if self.now_ms.saturating_sub(stable_since) < self.config.convergence_window_ms {
return;
}
self.emitted_ready = true;
let pool = self.candidates.clone();
self.events
.push(ReadinessEvent::PoolReady { pool: pool.clone() });
self.commands
.push(ReadinessCommand::EmitPoolReady { pool: pool.clone() });
if let Some(request) = &self.pending_run {
self.commands.push(ReadinessCommand::StartPlanning {
run_id: request.run_id,
pool,
});
}
}
}

View file

@ -0,0 +1,226 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u64);
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PoolId(pub String);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum WorkerStartPolicy {
StartBeforeAvailable,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct BootConfig {
pub node_id: NodeId,
pub intended_pool_id: PoolId,
pub worker_policy: WorkerStartPolicy,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum ResourceFact {
RustProcessAlive,
RuntimeAcceptingControl,
StableNodeIdKnown,
ArenaMapped,
GpuWorkerReady,
TransportEndpointBound,
SwimJoiningPool,
ProvisioningReceiverOpen,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ResourceOutcome {
RustProcessAlive,
RuntimeAcceptingControl,
StableNodeIdKnown(NodeId),
ArenaMapped,
GpuWorkerReady,
TransportEndpointBound,
SwimJoiningPool,
ProvisioningReceiverOpen,
ArenaFault,
GpuWorkerFault,
TransportEndpointFault,
InvalidNodeIdentity,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum BootFaultKind {
ArenaConstructionFailed,
WorkerStartupFailed,
TransportEndpointFailed,
InvalidNodeIdentity,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum LifecycleEvent {
NodeAvailable {
node_id: NodeId,
},
NodeFaulted {
node_id: NodeId,
kind: BootFaultKind,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum BootCommand {
AdvertiseLifecycle { node_id: NodeId },
JoinMembership { pool_id: PoolId },
OpenProvisioningInbox { node_id: NodeId },
LoadWeights { node_id: NodeId },
ConfigureRole { node_id: NodeId },
EstablishEdge { node_id: NodeId },
AssignStage { node_id: NodeId },
AssignLayerRange { node_id: NodeId },
AssignEdge { node_id: NodeId },
AssignObjectSpec { node_id: NodeId },
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ProvisionRequest {
pub node_id: NodeId,
}
impl ProvisionRequest {
pub fn for_node(node_id: NodeId) -> Self {
Self { node_id }
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ProvisioningAdmissionRejection {
NodeNotAvailable,
WrongNode,
}
#[cfg(test)]
pub struct BootHarness {
config: BootConfig,
facts: std::collections::BTreeSet<ResourceFact>,
events: Vec<LifecycleEvent>,
commands: Vec<BootCommand>,
available: bool,
faulted: bool,
}
#[cfg(test)]
impl BootHarness {
pub fn launch(config: BootConfig) -> Self {
let commands = vec![
BootCommand::AdvertiseLifecycle {
node_id: config.node_id,
},
BootCommand::JoinMembership {
pool_id: config.intended_pool_id.clone(),
},
BootCommand::OpenProvisioningInbox {
node_id: config.node_id,
},
];
Self {
config,
facts: std::collections::BTreeSet::new(),
events: Vec::new(),
commands,
available: false,
faulted: false,
}
}
pub fn observe(&mut self, outcome: ResourceOutcome) {
if self.available || self.faulted {
return;
}
match outcome {
ResourceOutcome::RustProcessAlive => self.insert_fact(ResourceFact::RustProcessAlive),
ResourceOutcome::RuntimeAcceptingControl => {
self.insert_fact(ResourceFact::RuntimeAcceptingControl)
}
ResourceOutcome::StableNodeIdKnown(node_id) if node_id == self.config.node_id => {
self.insert_fact(ResourceFact::StableNodeIdKnown)
}
ResourceOutcome::StableNodeIdKnown(_) | ResourceOutcome::InvalidNodeIdentity => {
self.fault(BootFaultKind::InvalidNodeIdentity)
}
ResourceOutcome::ArenaMapped => self.insert_fact(ResourceFact::ArenaMapped),
ResourceOutcome::GpuWorkerReady => self.insert_fact(ResourceFact::GpuWorkerReady),
ResourceOutcome::TransportEndpointBound => {
self.insert_fact(ResourceFact::TransportEndpointBound)
}
ResourceOutcome::SwimJoiningPool => self.insert_fact(ResourceFact::SwimJoiningPool),
ResourceOutcome::ProvisioningReceiverOpen => {
self.insert_fact(ResourceFact::ProvisioningReceiverOpen)
}
ResourceOutcome::ArenaFault => self.fault(BootFaultKind::ArenaConstructionFailed),
ResourceOutcome::GpuWorkerFault => self.fault(BootFaultKind::WorkerStartupFailed),
ResourceOutcome::TransportEndpointFault => {
self.fault(BootFaultKind::TransportEndpointFailed)
}
}
self.maybe_available();
}
pub fn events(&self) -> &[LifecycleEvent] {
&self.events
}
pub fn commands(&self) -> &[BootCommand] {
&self.commands
}
pub fn is_candidate_eligible(&self, node_id: NodeId) -> bool {
node_id == self.config.node_id && self.available && !self.faulted
}
pub fn try_accept_provisioning(
&self,
request: ProvisionRequest,
) -> Result<(), ProvisioningAdmissionRejection> {
if request.node_id != self.config.node_id {
return Err(ProvisioningAdmissionRejection::WrongNode);
}
if self.available && !self.faulted {
Ok(())
} else {
Err(ProvisioningAdmissionRejection::NodeNotAvailable)
}
}
fn insert_fact(&mut self, fact: ResourceFact) {
self.facts.insert(fact);
}
fn maybe_available(&mut self) {
if self.available || self.faulted {
return;
}
let required = [
ResourceFact::RustProcessAlive,
ResourceFact::RuntimeAcceptingControl,
ResourceFact::StableNodeIdKnown,
ResourceFact::ArenaMapped,
ResourceFact::GpuWorkerReady,
ResourceFact::TransportEndpointBound,
ResourceFact::SwimJoiningPool,
ResourceFact::ProvisioningReceiverOpen,
];
if required.iter().all(|fact| self.facts.contains(fact)) {
self.available = true;
self.events.push(LifecycleEvent::NodeAvailable {
node_id: self.config.node_id,
});
}
}
fn fault(&mut self, kind: BootFaultKind) {
if self.faulted || self.available {
return;
}
self.faulted = true;
self.events.push(LifecycleEvent::NodeFaulted {
node_id: self.config.node_id,
kind,
});
}
}

View file

@ -0,0 +1,443 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct RunId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct StageIndex(pub u32);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct EdgeId(pub 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 ObjectId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct Sequence(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct StepId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct WorkerGeneration(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum EventKind {
NodeStarted,
NodeAvailable,
PoolReady,
RunPlanned,
StageProvisionStarted,
WeightsDownloadStarted,
WeightsDownloaded,
WeightsLoaded,
EdgeProvisionStarted,
EdgeReady,
StageReady,
ReadinessBarrierPassed,
PromptInjected,
ObjectLoaded,
ExecuteStepStarted,
ObjectProduced,
StepCompleted,
TokenReceived,
RunCompleted,
RunFaulted,
StopRunSent,
StageStopped,
RunTornDown,
StageFaulted,
RingReadable,
WorkerReady,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Component {
NodeBoot,
Membership,
Orchestrator,
StageController,
WeightLifecycle,
EdgeEstablisher,
GpuWorkerCtl,
SharedRingHelper,
TokenEndpoint,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum FaultReason {
WorkerCrashed,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum Event {
RunScoped {
kind: EventKind,
run_id: RunId,
reason: Option<FaultReason>,
component: Component,
},
NodeScoped {
kind: EventKind,
node_id: NodeId,
component: Component,
},
StageScoped {
kind: EventKind,
run_id: RunId,
stage_index: StageIndex,
reason: Option<FaultReason>,
component: Component,
},
EdgeScoped {
kind: EventKind,
edge_id: EdgeId,
component: Component,
},
RingScoped {
kind: EventKind,
ring_id: RingId,
component: Component,
},
ObjectScoped {
kind: EventKind,
object_id: ObjectId,
sequence: Sequence,
component: Component,
},
StepScoped {
kind: EventKind,
step_id: StepId,
component: Component,
},
WorkerScoped {
kind: EventKind,
worker_generation: WorkerGeneration,
component: Component,
},
}
impl Event {
pub fn kind(&self) -> EventKind {
match self {
Event::RunScoped { kind, .. }
| Event::NodeScoped { kind, .. }
| Event::StageScoped { kind, .. }
| Event::EdgeScoped { kind, .. }
| Event::RingScoped { kind, .. }
| Event::ObjectScoped { kind, .. }
| Event::StepScoped { kind, .. }
| Event::WorkerScoped { kind, .. } => *kind,
}
}
}
pub struct TraceBuilder {
run_id: RunId,
events: Vec<Event>,
}
impl TraceBuilder {
pub fn new(run_id: RunId) -> Self {
Self {
run_id,
events: Vec::new(),
}
}
pub fn node_started(mut self, node_id: NodeId) -> Self {
self.events.push(Event::NodeScoped {
kind: EventKind::NodeStarted,
node_id,
component: Component::NodeBoot,
});
self
}
pub fn node_available(mut self, node_id: NodeId) -> Self {
self.events.push(Event::NodeScoped {
kind: EventKind::NodeAvailable,
node_id,
component: Component::NodeBoot,
});
self
}
pub fn pool_ready(mut self, _nodes: Vec<NodeId>) -> Self {
self.events.push(Event::RunScoped {
kind: EventKind::PoolReady,
run_id: self.run_id,
reason: None,
component: Component::Membership,
});
self
}
pub fn run_planned(mut self) -> Self {
self.events.push(Event::RunScoped {
kind: EventKind::RunPlanned,
run_id: self.run_id,
reason: None,
component: Component::Orchestrator,
});
self
}
pub fn stage_provision_started(mut self, stage_index: StageIndex, _node_id: NodeId) -> Self {
self.stage(
EventKind::StageProvisionStarted,
stage_index,
None,
Component::Orchestrator,
);
self
}
pub fn weights_download_started(mut self, stage_index: StageIndex) -> Self {
self.stage(
EventKind::WeightsDownloadStarted,
stage_index,
None,
Component::WeightLifecycle,
);
self
}
pub fn weights_downloaded(mut self, stage_index: StageIndex) -> Self {
self.stage(
EventKind::WeightsDownloaded,
stage_index,
None,
Component::WeightLifecycle,
);
self
}
pub fn weights_loaded(mut self, stage_index: StageIndex) -> Self {
self.stage(
EventKind::WeightsLoaded,
stage_index,
None,
Component::WeightLifecycle,
);
self
}
pub 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 {
self.edge(EventKind::EdgeReady, edge_id);
self
}
pub fn stage_ready(mut self, stage_index: StageIndex) -> Self {
self.stage(
EventKind::StageReady,
stage_index,
None,
Component::StageController,
);
self
}
pub fn readiness_barrier_passed(mut self) -> Self {
self.run(
EventKind::ReadinessBarrierPassed,
None,
Component::Orchestrator,
);
self
}
pub fn prompt_injected(mut self, sequence: Sequence) -> Self {
self.events.push(Event::ObjectScoped {
kind: EventKind::PromptInjected,
object_id: ObjectId(9000),
sequence,
component: Component::TokenEndpoint,
});
self
}
pub fn object_loaded(
mut self,
_edge_id: EdgeId,
object_id: ObjectId,
sequence: Sequence,
) -> Self {
self.events.push(Event::ObjectScoped {
kind: EventKind::ObjectLoaded,
object_id,
sequence,
component: Component::GpuWorkerCtl,
});
self
}
pub fn execute_step_started(mut self, step_id: StepId) -> Self {
self.events.push(Event::StepScoped {
kind: EventKind::ExecuteStepStarted,
step_id,
component: Component::StageController,
});
self
}
pub fn object_produced(
mut self,
_edge_id: EdgeId,
object_id: ObjectId,
sequence: Sequence,
) -> Self {
self.events.push(Event::ObjectScoped {
kind: EventKind::ObjectProduced,
object_id,
sequence,
component: Component::GpuWorkerCtl,
});
self
}
pub fn step_completed(mut self, step_id: StepId) -> Self {
self.events.push(Event::StepScoped {
kind: EventKind::StepCompleted,
step_id,
component: Component::StageController,
});
self
}
pub fn token_received(mut self, object_id: ObjectId, sequence: Sequence) -> Self {
self.events.push(Event::ObjectScoped {
kind: EventKind::TokenReceived,
object_id,
sequence,
component: Component::TokenEndpoint,
});
self
}
pub fn run_completed(mut self) -> Self {
self.run(EventKind::RunCompleted, None, Component::Orchestrator);
self
}
pub fn stage_faulted(
mut self,
stage_index: StageIndex,
reason: FaultReason,
component: Component,
) -> Self {
self.stage(
EventKind::StageFaulted,
stage_index,
Some(reason),
component,
);
self
}
pub 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 {
self.stage(
EventKind::StopRunSent,
stage_index,
None,
Component::Orchestrator,
);
self
}
pub fn stage_stopped(mut self, stage_index: StageIndex) -> Self {
self.stage(
EventKind::StageStopped,
stage_index,
None,
Component::StageController,
);
self
}
pub fn run_torn_down(mut self) -> Self {
self.run(EventKind::RunTornDown, None, Component::Orchestrator);
self
}
pub fn finish(self) -> Vec<Event> {
self.events
}
fn run(&mut self, kind: EventKind, reason: Option<FaultReason>, component: Component) {
self.events.push(Event::RunScoped {
kind,
run_id: self.run_id,
reason,
component,
});
}
fn stage(
&mut self,
kind: EventKind,
stage_index: StageIndex,
reason: Option<FaultReason>,
component: Component,
) {
self.events.push(Event::StageScoped {
kind,
run_id: self.run_id,
stage_index,
reason,
component,
});
}
fn edge(&mut self, kind: EventKind, edge_id: EdgeId) {
self.events.push(Event::EdgeScoped {
kind,
edge_id,
component: Component::EdgeEstablisher,
});
}
}
pub 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<Event>,
_batching: Batching,
}
#[cfg(test)]
impl EventSubscriberHarness {
pub fn collect(events: Vec<Event>, 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
}
}

View file

@ -0,0 +1,424 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct RunId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct StageRef {
pub stage_index: u32,
pub node_id: NodeId,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RunPlan {
pub run_id: RunId,
pub stages: Vec<StageRef>,
}
impl RunPlan {
pub fn test_linear(run_id: RunId, stages: Vec<StageRef>) -> Self {
Self { run_id, stages }
}
pub fn stage_nodes(&self) -> Vec<NodeId> {
self.stages.iter().map(|stage| stage.node_id).collect()
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RunConfig {
pub run_id: RunId,
pub max_tokens: u64,
pub prompt: Vec<u32>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum StageFaultReason {
WorkerCrashed,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum EndpointKind {
TokenIn,
TokenOut,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum TimeoutKind {
Execution,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum RunEvent {
PoolReady {
nodes: Vec<NodeId>,
},
PlanAvailable(RunPlan),
StageReady {
run_id: RunId,
stage_index: u32,
},
TokenInEndpointReady,
TokenOutEndpointReady,
TokenReceived {
sequence: u64,
token_id: u32,
eos: bool,
},
StageFault {
run_id: RunId,
stage_index: u32,
reason: StageFaultReason,
},
EndpointFault {
run_id: RunId,
endpoint: EndpointKind,
},
Timeout {
run_id: RunId,
kind: TimeoutKind,
},
StageStopped {
run_id: RunId,
stage_index: u32,
},
TokenEndpointsStopped,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RunFaultReason {
StageFault {
stage_index: u32,
reason: StageFaultReason,
},
EndpointFault {
endpoint: EndpointKind,
},
Timeout {
kind: TimeoutKind,
},
UnknownStageReady {
stage_index: u32,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum LifecycleEvent {
RunRejected {
run_id: RunId,
reason: RunFaultReason,
},
RunFaulted {
run_id: RunId,
reason: RunFaultReason,
},
RunCompleted {
run_id: RunId,
},
RunTornDown {
run_id: RunId,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct StageProvision {
pub run_id: RunId,
pub stage_index: u32,
pub node_id: NodeId,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum RunCommand {
ProvisionStage {
provision: StageProvision,
},
CreateTokenInEndpoint {
run_id: RunId,
},
CreateTokenOutEndpoint {
run_id: RunId,
},
InjectPrompt {
run_id: RunId,
sequence: u64,
prompt: Vec<u32>,
},
BroadcastStart {
run_id: RunId,
},
StopRun {
run_id: RunId,
stage_index: u32,
},
TearDownTokenEndpoints {
run_id: RunId,
},
}
#[cfg(test)]
pub struct OrchestratorHarness {
config: RunConfig,
plan: Option<RunPlan>,
pool_ready: bool,
provisioned: bool,
token_in_ready: bool,
token_out_ready: bool,
ready_stages: std::collections::BTreeSet<u32>,
injected_sequences: Vec<u64>,
expected_token_sequence: u64,
events: Vec<LifecycleEvent>,
commands: Vec<RunCommand>,
terminal: bool,
teardown_started: bool,
stopped_stages: std::collections::BTreeSet<u32>,
token_endpoints_stopped: bool,
}
#[cfg(test)]
impl OrchestratorHarness {
pub fn new(config: RunConfig) -> Self {
Self {
config,
plan: None,
pool_ready: false,
provisioned: false,
token_in_ready: false,
token_out_ready: false,
ready_stages: std::collections::BTreeSet::new(),
injected_sequences: Vec::new(),
expected_token_sequence: 0,
events: Vec::new(),
commands: Vec::new(),
terminal: false,
teardown_started: false,
stopped_stages: std::collections::BTreeSet::new(),
token_endpoints_stopped: false,
}
}
pub fn observe(&mut self, event: RunEvent) {
match event {
RunEvent::PoolReady { .. } => {
self.pool_ready = true;
self.maybe_provision();
}
RunEvent::PlanAvailable(plan) => {
self.plan = Some(plan);
self.maybe_provision();
}
RunEvent::StageReady {
run_id,
stage_index,
} => {
if run_id != self.config.run_id || !self.plan_has_stage(stage_index) {
self.fault(RunFaultReason::UnknownStageReady { stage_index });
return;
}
self.ready_stages.insert(stage_index);
self.maybe_inject_initial();
}
RunEvent::TokenInEndpointReady => {
self.token_in_ready = true;
self.maybe_inject_initial();
}
RunEvent::TokenOutEndpointReady => {
self.token_out_ready = true;
self.maybe_inject_initial();
}
RunEvent::TokenReceived { sequence, eos, .. } => {
if self.terminal {
return;
}
if sequence != self.expected_token_sequence {
return;
}
self.expected_token_sequence += 1;
if eos {
self.complete();
} else if (self.injected_sequences.len() as u64) < self.config.max_tokens {
self.inject(sequence + 1);
} else {
self.complete();
}
}
RunEvent::StageFault {
run_id,
stage_index,
reason,
} if run_id == self.config.run_id => {
self.fault(RunFaultReason::StageFault {
stage_index,
reason,
});
}
RunEvent::EndpointFault { run_id, endpoint } if run_id == self.config.run_id => {
self.fault(RunFaultReason::EndpointFault { endpoint });
}
RunEvent::Timeout { run_id, kind } if run_id == self.config.run_id => {
self.fault(RunFaultReason::Timeout { kind });
}
RunEvent::StageStopped {
run_id,
stage_index,
} if run_id == self.config.run_id => {
self.stopped_stages.insert(stage_index);
self.maybe_torn_down();
}
RunEvent::TokenEndpointsStopped => {
self.token_endpoints_stopped = true;
self.maybe_torn_down();
}
RunEvent::StageFault { .. }
| RunEvent::EndpointFault { .. }
| RunEvent::Timeout { .. }
| RunEvent::StageStopped { .. } => {}
}
}
pub fn advance_time_ms(&mut self, _delta: u64) {}
pub fn commands(&self) -> &[RunCommand] {
&self.commands
}
pub fn events(&self) -> &[LifecycleEvent] {
&self.events
}
pub fn injected_sequences(&self) -> Vec<u64> {
self.injected_sequences.clone()
}
fn maybe_provision(&mut self) {
if !self.pool_ready || self.provisioned {
return;
}
let Some(plan) = &self.plan else {
return;
};
self.provisioned = true;
for stage in &plan.stages {
self.commands.push(RunCommand::ProvisionStage {
provision: StageProvision {
run_id: plan.run_id,
stage_index: stage.stage_index,
node_id: stage.node_id,
},
});
}
self.commands.push(RunCommand::CreateTokenInEndpoint {
run_id: self.config.run_id,
});
self.commands.push(RunCommand::CreateTokenOutEndpoint {
run_id: self.config.run_id,
});
}
fn maybe_inject_initial(&mut self) {
if self.terminal || !self.provisioned || !self.injected_sequences.is_empty() {
return;
}
if self.token_in_ready && self.token_out_ready && self.all_stages_ready() {
self.inject(0);
}
}
fn inject(&mut self, sequence: u64) {
if self.terminal {
return;
}
self.injected_sequences.push(sequence);
self.commands.push(RunCommand::InjectPrompt {
run_id: self.config.run_id,
sequence,
prompt: if sequence == 0 {
self.config.prompt.clone()
} else {
Vec::new()
},
});
}
fn complete(&mut self) {
if self.terminal {
return;
}
self.terminal = true;
self.events.push(LifecycleEvent::RunCompleted {
run_id: self.config.run_id,
});
self.start_teardown();
}
fn fault(&mut self, reason: RunFaultReason) {
if self.terminal {
return;
}
self.terminal = true;
self.events.push(LifecycleEvent::RunFaulted {
run_id: self.config.run_id,
reason,
});
self.start_teardown();
}
fn start_teardown(&mut self) {
if self.teardown_started {
return;
}
self.teardown_started = true;
if let Some(plan) = &self.plan {
for stage in &plan.stages {
self.commands.push(RunCommand::StopRun {
run_id: self.config.run_id,
stage_index: stage.stage_index,
});
}
}
self.commands.push(RunCommand::TearDownTokenEndpoints {
run_id: self.config.run_id,
});
}
fn maybe_torn_down(&mut self) {
if !self.teardown_started || !self.token_endpoints_stopped {
return;
}
if !self.all_stages_stopped() {
return;
}
if !self
.events
.iter()
.any(|event| matches!(event, LifecycleEvent::RunTornDown { .. }))
{
self.events.push(LifecycleEvent::RunTornDown {
run_id: self.config.run_id,
});
}
}
fn plan_has_stage(&self, stage_index: u32) -> bool {
self.plan.as_ref().is_some_and(|plan| {
plan.stages
.iter()
.any(|stage| stage.stage_index == stage_index)
})
}
fn all_stages_ready(&self) -> bool {
self.plan.as_ref().is_some_and(|plan| {
plan.stages
.iter()
.all(|stage| self.ready_stages.contains(&stage.stage_index))
})
}
fn all_stages_stopped(&self) -> bool {
self.plan.as_ref().is_some_and(|plan| {
plan.stages
.iter()
.all(|stage| self.stopped_stages.contains(&stage.stage_index))
})
}
}

View file

@ -0,0 +1,325 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct RunId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct EdgeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ObjectId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ObjectSpec {
pub max_extent: u64,
pub alignment: u64,
}
impl ObjectSpec {
pub fn test_tokens() -> Self {
Self {
max_extent: 64,
alignment: 4,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum EndpointDirection {
TokenIn,
TokenOut,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum EdgeSemantics {
SingleProducerSingleConsumer,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct EdgePlan {
pub edge_id: EdgeId,
pub producer_node_id: NodeId,
pub consumer_node_id: NodeId,
pub direction: EndpointDirection,
}
impl EdgePlan {
pub fn token_in(
edge_id: EdgeId,
orchestrator_node_id: NodeId,
first_stage_node: NodeId,
) -> Self {
Self {
edge_id,
producer_node_id: orchestrator_node_id,
consumer_node_id: first_stage_node,
direction: EndpointDirection::TokenIn,
}
}
pub fn token_out(
edge_id: EdgeId,
last_stage_node: NodeId,
orchestrator_node_id: NodeId,
) -> Self {
Self {
edge_id,
producer_node_id: last_stage_node,
consumer_node_id: orchestrator_node_id,
direction: EndpointDirection::TokenOut,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct TokenEndpointPlan {
pub run_id: RunId,
pub orchestrator_node_id: NodeId,
pub token_in_edge: EdgePlan,
pub token_out_edge: EdgePlan,
pub token_spec: ObjectSpec,
pub max_tokens: u64,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct LocalEndpoint {
pub edge_id: EdgeId,
pub orchestrator_node_id: NodeId,
pub edge_semantics: EdgeSemantics,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct TokenObjectWrite {
pub run_id: RunId,
pub edge_id: EdgeId,
pub sequence: u64,
pub tokens: Vec<u32>,
pub object_spec: ObjectSpec,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum EndpointCommand {
CreateTokenInProducer {
edge_id: EdgeId,
orchestrator_node_id: NodeId,
},
CreateTokenOutConsumer {
edge_id: EdgeId,
orchestrator_node_id: NodeId,
},
WriteTokenObject(TokenObjectWrite),
BroadcastStart {
run_id: RunId,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum EndpointEvent {
TokenInReady {
edge_id: EdgeId,
},
TokenOutReady {
edge_id: EdgeId,
},
AllStagesReady,
ReadinessBarrierPassed,
TokenObjectReceived {
edge_id: EdgeId,
object_id: ObjectId,
sequence: u64,
token_id: u32,
eos: bool,
},
EndpointFault {
edge_id: EdgeId,
direction: EndpointDirection,
},
MalformedTokenObject {
edge_id: EdgeId,
object_id: ObjectId,
},
TeardownFailed {
edge_id: EdgeId,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum RunFaultReason {
TokenSequenceViolation,
TokenEndpointFault,
MalformedTokenObject,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum EndpointLifecycleEvent {
RunFaulted {
run_id: RunId,
reason: RunFaultReason,
},
TeardownFailed {
run_id: RunId,
edge_id: EdgeId,
},
}
#[cfg(test)]
pub struct TokenEndpointHarness {
plan: TokenEndpointPlan,
commands: Vec<EndpointCommand>,
events: Vec<EndpointLifecycleEvent>,
local_endpoints: Vec<LocalEndpoint>,
prompt: Option<Vec<u32>>,
token_in_ready: bool,
token_out_ready: bool,
stages_ready: bool,
barrier_passed: bool,
injected_sequences: Vec<u64>,
expected_token_sequence: u64,
terminal: bool,
}
#[cfg(test)]
impl TokenEndpointHarness {
pub fn new(plan: TokenEndpointPlan) -> Self {
let commands = vec![
EndpointCommand::CreateTokenInProducer {
edge_id: plan.token_in_edge.edge_id,
orchestrator_node_id: plan.orchestrator_node_id,
},
EndpointCommand::CreateTokenOutConsumer {
edge_id: plan.token_out_edge.edge_id,
orchestrator_node_id: plan.orchestrator_node_id,
},
];
let local_endpoints = vec![
LocalEndpoint {
edge_id: plan.token_in_edge.edge_id,
orchestrator_node_id: plan.orchestrator_node_id,
edge_semantics: EdgeSemantics::SingleProducerSingleConsumer,
},
LocalEndpoint {
edge_id: plan.token_out_edge.edge_id,
orchestrator_node_id: plan.orchestrator_node_id,
edge_semantics: EdgeSemantics::SingleProducerSingleConsumer,
},
];
Self {
plan,
commands,
events: Vec::new(),
local_endpoints,
prompt: None,
token_in_ready: false,
token_out_ready: false,
stages_ready: false,
barrier_passed: false,
injected_sequences: Vec::new(),
expected_token_sequence: 0,
terminal: false,
}
}
pub fn commands(&self) -> &[EndpointCommand] {
&self.commands
}
pub fn events(&self) -> &[EndpointLifecycleEvent] {
&self.events
}
pub fn local_endpoints(&self) -> &[LocalEndpoint] {
&self.local_endpoints
}
pub fn injected_sequences(&self) -> Vec<u64> {
self.injected_sequences.clone()
}
pub fn request_prompt_injection(&mut self, prompt: Vec<u32>) {
self.prompt = Some(prompt);
self.maybe_inject_initial();
}
pub fn advance_time_ms(&mut self, _delta: u64) {}
pub fn observe(&mut self, event: EndpointEvent) {
match event {
EndpointEvent::TokenInReady { edge_id }
if edge_id == self.plan.token_in_edge.edge_id =>
{
self.token_in_ready = true
}
EndpointEvent::TokenOutReady { edge_id }
if edge_id == self.plan.token_out_edge.edge_id =>
{
self.token_out_ready = true
}
EndpointEvent::AllStagesReady => self.stages_ready = true,
EndpointEvent::ReadinessBarrierPassed => self.barrier_passed = true,
EndpointEvent::TokenObjectReceived { sequence, eos, .. } => {
if self.terminal {
return;
}
if sequence != self.expected_token_sequence {
self.fault(RunFaultReason::TokenSequenceViolation);
return;
}
self.expected_token_sequence += 1;
if eos || self.injected_sequences.len() as u64 >= self.plan.max_tokens {
self.terminal = true;
} else {
self.inject_sequence(sequence + 1, vec![0]);
}
}
EndpointEvent::EndpointFault { .. } => self.fault(RunFaultReason::TokenEndpointFault),
EndpointEvent::MalformedTokenObject { .. } => {
self.fault(RunFaultReason::MalformedTokenObject)
}
EndpointEvent::TeardownFailed { edge_id } => {
self.events.push(EndpointLifecycleEvent::TeardownFailed {
run_id: self.plan.run_id,
edge_id,
});
}
EndpointEvent::TokenInReady { .. } | EndpointEvent::TokenOutReady { .. } => {}
}
self.maybe_inject_initial();
}
fn maybe_inject_initial(&mut self) {
if self.injected_sequences.is_empty()
&& self.prompt.is_some()
&& self.token_in_ready
&& self.token_out_ready
&& self.stages_ready
&& self.barrier_passed
{
let prompt = self.prompt.clone().unwrap_or_default();
self.inject_sequence(0, prompt);
}
}
fn inject_sequence(&mut self, sequence: u64, tokens: Vec<u32>) {
if self.terminal || self.injected_sequences.len() as u64 >= self.plan.max_tokens {
return;
}
self.injected_sequences.push(sequence);
self.commands
.push(EndpointCommand::WriteTokenObject(TokenObjectWrite {
run_id: self.plan.run_id,
edge_id: self.plan.token_in_edge.edge_id,
sequence,
tokens,
object_spec: self.plan.token_spec,
}));
}
fn fault(&mut self, reason: RunFaultReason) {
if self.terminal {
return;
}
self.terminal = true;
self.events.push(EndpointLifecycleEvent::RunFaulted {
run_id: self.plan.run_id,
reason,
});
}
}

View file

@ -0,0 +1,199 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct RunId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum GpuClass {
TestSmall,
TestLarge,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum BootHealth {
Ready,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct InventoryEntry {
pub node_id: NodeId,
pub gpu_class: GpuClass,
pub ready: bool,
}
impl InventoryEntry {
pub fn ready(node_id: NodeId, gpu_class: GpuClass) -> Self {
Self {
node_id,
gpu_class,
ready: true,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct StagePlacement {
pub stage_index: u32,
pub node_id: NodeId,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum PlacementInput {
FixedLinear(Vec<StagePlacement>),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PlanningRequest {
pub run_id: RunId,
pub stage_count: u32,
pub entries: Vec<InventoryEntry>,
pub placement: PlacementInput,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct StagePlan {
pub stage_index: u32,
pub node_id: NodeId,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RunPlan {
pub run_id: RunId,
pub stages: Vec<StagePlan>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum InventoryRejectionKind {
UnknownNode { node_id: NodeId },
DuplicateStage { stage_index: u32 },
MissingStage { stage_index: u32 },
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct InventoryRejection {
pub kind: InventoryRejectionKind,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum NodeReport {
BootHealth {
node_id: NodeId,
health: BootHealth,
},
CapacityChanged {
node_id: NodeId,
gpu_class: GpuClass,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum InventoryCommand {
AcceptPlacementNegotiation { node_id: NodeId },
RewriteStagePlacement { node_id: NodeId },
ProvisionStage { stage_index: u32, node_id: NodeId },
}
pub fn plan_from_inventory(request: PlanningRequest) -> Result<RunPlan, InventoryRejection> {
let known_nodes = request
.entries
.iter()
.map(|entry| entry.node_id)
.collect::<std::collections::BTreeSet<_>>();
let PlacementInput::FixedLinear(stages) = request.placement;
let mut by_stage = std::collections::BTreeMap::new();
for placement in stages {
if !known_nodes.contains(&placement.node_id) {
return Err(InventoryRejection {
kind: InventoryRejectionKind::UnknownNode {
node_id: placement.node_id,
},
});
}
if by_stage
.insert(placement.stage_index, placement.node_id)
.is_some()
{
return Err(InventoryRejection {
kind: InventoryRejectionKind::DuplicateStage {
stage_index: placement.stage_index,
},
});
}
}
let mut stages = Vec::with_capacity(request.stage_count as usize);
for stage_index in 0..request.stage_count {
let node_id = by_stage.remove(&stage_index).ok_or(InventoryRejection {
kind: InventoryRejectionKind::MissingStage { stage_index },
})?;
stages.push(StagePlan {
stage_index,
node_id,
});
}
Ok(RunPlan {
run_id: request.run_id,
stages,
})
}
#[cfg(test)]
pub struct InventoryHarness {
entries: Vec<InventoryEntry>,
commands: Vec<InventoryCommand>,
committed_plan: Option<RunPlan>,
}
#[cfg(test)]
impl InventoryHarness {
pub fn new(entries: Vec<InventoryEntry>) -> Self {
Self {
entries,
commands: Vec::new(),
committed_plan: None,
}
}
pub fn entries(&self) -> &[InventoryEntry] {
&self.entries
}
pub fn commands(&self) -> &[InventoryCommand] {
&self.commands
}
pub fn committed_plan(&self) -> Option<&RunPlan> {
self.committed_plan.as_ref()
}
pub fn observe_node_report(&mut self, report: NodeReport) {
match report {
NodeReport::BootHealth { node_id, health } => {
if let Some(entry) = self
.entries
.iter_mut()
.find(|entry| entry.node_id == node_id)
{
entry.ready = matches!(health, BootHealth::Ready);
}
}
NodeReport::CapacityChanged { .. } => {}
}
}
pub fn commit_plan(&mut self, plan: RunPlan) {
self.committed_plan = Some(plan);
}
pub fn provision_stages(&mut self) {
if let Some(plan) = &self.committed_plan {
for stage in &plan.stages {
self.commands.push(InventoryCommand::ProvisionStage {
stage_index: stage.stage_index,
node_id: stage.node_id,
});
}
}
}
}

View file

@ -0,0 +1,422 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct RunId(pub u64);
impl From<u64> for RunId {
fn from(value: u64) -> Self {
Self(value)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct EdgeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DTypeFamily {
BFloat,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ModelFacts {
pub model_id: String,
pub num_layers: u32,
pub hidden_dim: u64,
pub dtype_family: DTypeFamily,
pub dtype_width_bytes: u64,
pub max_seq_len: u64,
pub eos_token_id: u32,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RuntimeConfig {
pub max_tokens: u32,
}
impl RuntimeConfig {
pub fn test_default() -> Self {
Self { max_tokens: 4 }
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct StagePlacement {
pub stage_index: u32,
pub node_id: NodeId,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum PlacementInput {
FixedLinear(Vec<StagePlacement>),
FixedLinearWithEdgeOverride {
stages: Vec<StagePlacement>,
forced_activation_edges: Vec<(u32, u32)>,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct RingSpec {
pub data_capacity: u64,
pub alignment: u64,
}
impl RingSpec {
pub fn test_default_activation() -> Self {
Self {
data_capacity: 1 << 20,
alignment: 64,
}
}
pub fn test_default_token() -> Self {
Self {
data_capacity: 4096,
alignment: 8,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PlannerInput {
pub run_id: RunId,
pub orchestrator_node_id: NodeId,
pub model: ModelFacts,
pub runtime: RuntimeConfig,
pub candidate_pool: Vec<NodeId>,
pub stage_count: u32,
pub placement: PlacementInput,
pub activation_ring: RingSpec,
pub token_ring: RingSpec,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum EdgeKind {
TokenIn,
Activation,
TokenOut,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ObjectKind {
Token,
Activation,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ObjectSpec {
pub kind: ObjectKind,
pub max_extent: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum EdgeEndpoint {
Orchestrator { node_id: NodeId },
Stage { node_id: NodeId, stage_index: u32 },
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct EdgePlan {
pub run_id: RunId,
pub edge_id: EdgeId,
pub kind: EdgeKind,
pub producer: EdgeEndpoint,
pub consumer: EdgeEndpoint,
pub object_spec: ObjectSpec,
pub ring_spec: RingSpec,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct StagePlan {
pub run_id: RunId,
pub stage_index: u32,
pub stage_count: u32,
pub node_id: NodeId,
pub layer_start: u32,
pub layer_end_exclusive: u32,
pub inbound_edge: EdgeId,
pub outbound_edge: EdgeId,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RunPlan {
pub run_id: RunId,
pub stages: Vec<StagePlan>,
pub edges: Vec<EdgePlan>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct InboundEdgeProvision {
pub edge_id: EdgeId,
pub kind: EdgeKind,
pub object_spec: ObjectSpec,
pub ring_spec: RingSpec,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct OutboundEdgeProvision {
pub edge_id: EdgeId,
pub kind: EdgeKind,
pub consumer_node_id: NodeId,
pub object_spec: ObjectSpec,
pub ring_spec: RingSpec,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ProvisionStage {
pub run_id: RunId,
pub node_id: NodeId,
pub stage_index: u32,
pub stage_count: u32,
pub layer_start: u32,
pub layer_end_exclusive: u32,
pub inbound: InboundEdgeProvision,
pub outbound: OutboundEdgeProvision,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum PlanRejectionKind {
UnknownNode,
DuplicateStageAssignment,
MissingStage,
InvalidStageCount,
EdgeEndpointMismatch,
ModelStageLayoutMismatch,
InvalidObjectSpec,
UnsupportedShapeOrLayout,
InvalidRingSpec,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PlanRejection {
kind: PlanRejectionKind,
}
impl PlanRejection {
pub fn kind(&self) -> PlanRejectionKind {
self.kind
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ProjectionRejection {
UnknownStage,
MissingEdge,
}
pub fn plan_run(input: PlannerInput) -> Result<RunPlan, PlanRejection> {
validate_global_input(&input)?;
let placements = validated_placements(&input)?;
let activation_extent = input
.model
.max_seq_len
.checked_mul(input.model.hidden_dim)
.and_then(|value| value.checked_mul(input.model.dtype_width_bytes))
.ok_or_else(|| reject(PlanRejectionKind::InvalidObjectSpec))?;
if activation_extent == 0 {
return Err(reject(PlanRejectionKind::InvalidObjectSpec));
}
let token_spec = ObjectSpec {
kind: ObjectKind::Token,
max_extent: 64,
};
let activation_spec = ObjectSpec {
kind: ObjectKind::Activation,
max_extent: activation_extent,
};
let mut edges = Vec::with_capacity(input.stage_count as usize + 1);
edges.push(EdgePlan {
run_id: input.run_id,
edge_id: EdgeId(7000),
kind: EdgeKind::TokenIn,
producer: EdgeEndpoint::Orchestrator {
node_id: input.orchestrator_node_id,
},
consumer: EdgeEndpoint::Stage {
node_id: placements[0].node_id,
stage_index: 0,
},
object_spec: token_spec,
ring_spec: input.token_ring,
});
for stage_index in 0..input.stage_count.saturating_sub(1) {
edges.push(EdgePlan {
run_id: input.run_id,
edge_id: EdgeId(7001 + u64::from(stage_index)),
kind: EdgeKind::Activation,
producer: EdgeEndpoint::Stage {
node_id: placements[stage_index as usize].node_id,
stage_index,
},
consumer: EdgeEndpoint::Stage {
node_id: placements[stage_index as usize + 1].node_id,
stage_index: stage_index + 1,
},
object_spec: activation_spec,
ring_spec: input.activation_ring,
});
}
edges.push(EdgePlan {
run_id: input.run_id,
edge_id: EdgeId(7000 + u64::from(input.stage_count)),
kind: EdgeKind::TokenOut,
producer: EdgeEndpoint::Stage {
node_id: placements[input.stage_count as usize - 1].node_id,
stage_index: input.stage_count - 1,
},
consumer: EdgeEndpoint::Orchestrator {
node_id: input.orchestrator_node_id,
},
object_spec: token_spec,
ring_spec: input.token_ring,
});
let mut stages = Vec::with_capacity(input.stage_count as usize);
for placement in &placements {
let stage_index = placement.stage_index;
let (start, end) = layer_range(input.model.num_layers, input.stage_count, stage_index);
stages.push(StagePlan {
run_id: input.run_id,
stage_index,
stage_count: input.stage_count,
node_id: placement.node_id,
layer_start: start,
layer_end_exclusive: end,
inbound_edge: EdgeId(7000 + u64::from(stage_index)),
outbound_edge: EdgeId(7001 + u64::from(stage_index)),
});
}
Ok(RunPlan {
run_id: input.run_id,
stages,
edges,
})
}
pub fn derive_stage_provision(
plan: &RunPlan,
stage_index: u32,
) -> Result<ProvisionStage, ProjectionRejection> {
let stage = plan
.stages
.iter()
.find(|stage| stage.stage_index == stage_index)
.ok_or(ProjectionRejection::UnknownStage)?;
let inbound = plan
.edges
.iter()
.find(|edge| edge.edge_id == stage.inbound_edge)
.ok_or(ProjectionRejection::MissingEdge)?;
let outbound = plan
.edges
.iter()
.find(|edge| edge.edge_id == stage.outbound_edge)
.ok_or(ProjectionRejection::MissingEdge)?;
Ok(ProvisionStage {
run_id: plan.run_id,
node_id: stage.node_id,
stage_index,
stage_count: stage.stage_count,
layer_start: stage.layer_start,
layer_end_exclusive: stage.layer_end_exclusive,
inbound: InboundEdgeProvision {
edge_id: inbound.edge_id,
kind: inbound.kind,
object_spec: inbound.object_spec,
ring_spec: inbound.ring_spec,
},
outbound: OutboundEdgeProvision {
edge_id: outbound.edge_id,
kind: outbound.kind,
consumer_node_id: endpoint_node_id(&outbound.consumer),
object_spec: outbound.object_spec,
ring_spec: outbound.ring_spec,
},
})
}
fn validate_global_input(input: &PlannerInput) -> Result<(), PlanRejection> {
if input.stage_count == 0 {
return Err(reject(PlanRejectionKind::InvalidStageCount));
}
if input.model.num_layers < input.stage_count {
return Err(reject(PlanRejectionKind::ModelStageLayoutMismatch));
}
if input.model.max_seq_len == 0 || input.model.dtype_width_bytes == 0 {
return Err(reject(PlanRejectionKind::InvalidObjectSpec));
}
if input.model.hidden_dim == 0 {
return Err(reject(PlanRejectionKind::UnsupportedShapeOrLayout));
}
if !valid_ring(input.activation_ring) || !valid_ring(input.token_ring) {
return Err(reject(PlanRejectionKind::InvalidRingSpec));
}
Ok(())
}
fn validated_placements(input: &PlannerInput) -> Result<Vec<StagePlacement>, PlanRejection> {
let (stages, forced_edges) = match &input.placement {
PlacementInput::FixedLinear(stages) => (stages.as_slice(), &[][..]),
PlacementInput::FixedLinearWithEdgeOverride {
stages,
forced_activation_edges,
} => (stages.as_slice(), forced_activation_edges.as_slice()),
};
if forced_edges
.iter()
.any(|(producer, consumer)| *consumer != producer.saturating_add(1))
{
return Err(reject(PlanRejectionKind::EdgeEndpointMismatch));
}
let candidate_nodes = input
.candidate_pool
.iter()
.copied()
.collect::<std::collections::BTreeSet<_>>();
let mut by_stage = std::collections::BTreeMap::new();
for placement in stages {
if !candidate_nodes.contains(&placement.node_id) {
return Err(reject(PlanRejectionKind::UnknownNode));
}
if by_stage.insert(placement.stage_index, *placement).is_some() {
return Err(reject(PlanRejectionKind::DuplicateStageAssignment));
}
}
let mut dense = Vec::with_capacity(input.stage_count as usize);
for stage_index in 0..input.stage_count {
let placement = by_stage
.remove(&stage_index)
.ok_or_else(|| reject(PlanRejectionKind::MissingStage))?;
dense.push(placement);
}
Ok(dense)
}
fn layer_range(num_layers: u32, stage_count: u32, stage_index: u32) -> (u32, u32) {
let start = (u64::from(num_layers) * u64::from(stage_index) / u64::from(stage_count)) as u32;
let end = (u64::from(num_layers) * u64::from(stage_index + 1) / u64::from(stage_count)) as u32;
(start, end)
}
fn endpoint_node_id(endpoint: &EdgeEndpoint) -> NodeId {
match endpoint {
EdgeEndpoint::Orchestrator { node_id } | EdgeEndpoint::Stage { node_id, .. } => *node_id,
}
}
fn valid_ring(spec: RingSpec) -> bool {
spec.data_capacity > 0 && spec.alignment > 0 && spec.alignment.is_power_of_two()
}
fn reject(kind: PlanRejectionKind) -> PlanRejection {
PlanRejection { kind }
}

View file

@ -0,0 +1,293 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct RingId(pub u64);
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct EndpointId(pub String);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ArenaBase(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ProcessLocalPointer(pub u64);
impl ProcessLocalPointer {
pub fn from_base_plus_offset(base: ArenaBase, offset: u64) -> Self {
Self(base.0 + offset)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RingConfig {
pub node_id: NodeId,
pub ring_id: RingId,
pub capacity: u64,
pub producer: EndpointId,
pub consumer: EndpointId,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RingIdentity {
pub ring_id: RingId,
pub capacity: u64,
pub producer_count: u32,
pub consumer_count: u32,
generation: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct CursorSnapshot {
pub commit: u64,
pub consume: u64,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WakeHint {
RingReadable { ring_id: RingId },
RingWritable { ring_id: RingId },
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ReserveError {
InsufficientSpace,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ReadError {
BeyondCommittedBytes,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Reservation {
start: u64,
len: u64,
generation: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct StaleWake {
ring_id: RingId,
generation: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct WakeDelivery {
pub was_accepted: bool,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct SchedulerState {
pub readable_rings: std::collections::BTreeSet<RingId>,
pub writable_rings: std::collections::BTreeSet<RingId>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct RingLayout {
pub data_offset: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct ProcessLocalView {
pub layout: RingLayout,
pub data_pointer: ProcessLocalPointer,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum PythonOperation {
ReserveViaHelper { ring_id: RingId },
CommitViaHelper { ring_id: RingId },
ReadableViaHelper { ring_id: RingId },
ConsumeViaHelper { ring_id: RingId },
MapPointerViaHelper { ring_id: RingId },
DirectAtomicAccess { ring_id: RingId },
DirectWrapArithmetic { ring_id: RingId },
}
#[cfg(test)]
pub struct RingHelperHarness {
identity: RingIdentity,
buffer: Vec<u8>,
commit: u64,
consume: u64,
retired: bool,
wake_hints: Vec<WakeHint>,
wake_log: Vec<WakeDelivery>,
scheduler_state: SchedulerState,
}
#[cfg(test)]
impl RingHelperHarness {
pub fn create(config: RingConfig) -> Self {
Self::with_generation(config.ring_id, config.capacity, 1)
}
pub fn create_replacement(ring_id: RingId) -> Self {
Self::with_generation(ring_id, 8, 2)
}
fn with_generation(ring_id: RingId, capacity: u64, generation: u64) -> Self {
Self {
identity: RingIdentity {
ring_id,
capacity,
producer_count: 1,
consumer_count: 1,
generation,
},
buffer: vec![0; capacity as usize],
commit: 0,
consume: 0,
retired: false,
wake_hints: Vec::new(),
wake_log: Vec::new(),
scheduler_state: SchedulerState {
readable_rings: std::collections::BTreeSet::new(),
writable_rings: std::collections::BTreeSet::new(),
},
}
}
pub fn identity(&self) -> RingIdentity {
self.identity.clone()
}
pub fn retire_and_capture_stale_wake(&mut self) -> StaleWake {
self.retired = true;
StaleWake {
ring_id: self.identity.ring_id,
generation: self.identity.generation,
}
}
pub fn deliver_wake(&mut self, wake: StaleWake) {
self.wake_log.push(WakeDelivery {
was_accepted: wake.ring_id == self.identity.ring_id
&& wake.generation == self.identity.generation
&& !self.retired,
});
}
pub fn wake_log(&self) -> &[WakeDelivery] {
&self.wake_log
}
pub fn producer_reserve(&self, len: u64) -> Result<Reservation, ReserveError> {
let used = self.commit.saturating_sub(self.consume);
if len <= self.identity.capacity.saturating_sub(used) {
Ok(Reservation {
start: self.commit,
len,
generation: self.identity.generation,
})
} else {
Err(ReserveError::InsufficientSpace)
}
}
pub fn producer_write(&mut self, reservation: &Reservation, bytes: &[u8]) {
assert_eq!(reservation.generation, self.identity.generation);
assert_eq!(reservation.len as usize, bytes.len());
for (i, byte) in bytes.iter().copied().enumerate() {
let idx = (reservation.start + i as u64) % self.identity.capacity;
self.buffer[idx as usize] = byte;
}
}
pub fn producer_commit(&mut self, reservation: Reservation) {
assert_eq!(reservation.start, self.commit);
self.commit = self.commit.saturating_add(reservation.len);
self.wake_hints.push(WakeHint::RingReadable {
ring_id: self.identity.ring_id,
});
self.scheduler_state
.readable_rings
.insert(self.identity.ring_id);
}
pub fn consumer_readable(&self) -> u64 {
self.commit.saturating_sub(self.consume)
}
pub fn consumer_try_read(&self, len: u64) -> Result<Vec<u8>, ReadError> {
if len > self.consumer_readable() {
return Err(ReadError::BeyondCommittedBytes);
}
Ok(self.consumer_read(len))
}
pub fn consumer_read(&self, len: u64) -> Vec<u8> {
assert!(len <= self.consumer_readable());
(0..len)
.map(|i| self.buffer[((self.consume + i) % self.identity.capacity) as usize])
.collect()
}
pub fn consumer_consume(&mut self, len: u64) {
assert!(len <= self.consumer_readable());
self.consume = self.consume.saturating_add(len);
self.wake_hints.push(WakeHint::RingWritable {
ring_id: self.identity.ring_id,
});
self.scheduler_state
.writable_rings
.insert(self.identity.ring_id);
}
pub fn cursor_snapshot(&self) -> CursorSnapshot {
CursorSnapshot {
commit: self.commit,
consume: self.consume,
}
}
pub fn wake_hints(&self) -> &[WakeHint] {
&self.wake_hints
}
pub fn coalesce_duplicate_wakes(&mut self) {
for wake in &self.wake_hints {
match wake {
WakeHint::RingReadable { ring_id } => {
self.scheduler_state.readable_rings.insert(*ring_id);
}
WakeHint::RingWritable { ring_id } => {
self.scheduler_state.writable_rings.insert(*ring_id);
}
}
}
}
pub fn scheduler_state(&self) -> &SchedulerState {
&self.scheduler_state
}
pub fn map_process_local_view(&self, base: ArenaBase) -> ProcessLocalView {
let layout = RingLayout { data_offset: 128 };
ProcessLocalView {
layout,
data_pointer: ProcessLocalPointer::from_base_plus_offset(base, layout.data_offset),
}
}
pub fn python_visible_operations(&self) -> Vec<PythonOperation> {
vec![
PythonOperation::ReserveViaHelper {
ring_id: self.identity.ring_id,
},
PythonOperation::CommitViaHelper {
ring_id: self.identity.ring_id,
},
PythonOperation::ReadableViaHelper {
ring_id: self.identity.ring_id,
},
PythonOperation::ConsumeViaHelper {
ring_id: self.identity.ring_id,
},
PythonOperation::MapPointerViaHelper {
ring_id: self.identity.ring_id,
},
]
}
}

View file

@ -0,0 +1,399 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct RunId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct EdgeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ObjectId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct StepId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct DeviceHandle {
pub generation: u64,
pub id: u64,
}
impl DeviceHandle {
pub fn new_current(id: u64) -> Self {
Self { generation: 1, id }
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct LayerRange {
pub start: u32,
pub end_exclusive: u32,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WeightSource {
TestArtifact(String),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum EdgeDirection {
Inbound,
Outbound,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct EdgeProvision {
pub edge_id: EdgeId,
pub direction: EdgeDirection,
}
impl EdgeProvision {
pub fn inbound(edge_id: EdgeId) -> Self {
Self {
edge_id,
direction: EdgeDirection::Inbound,
}
}
pub fn outbound(edge_id: EdgeId) -> Self {
Self {
edge_id,
direction: EdgeDirection::Outbound,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ProvisionStage {
pub run_id: RunId,
pub authorized_orchestrator: NodeId,
pub node_id: NodeId,
pub stage_index: u32,
pub stage_count: u32,
pub layer_range: LayerRange,
pub inbound: EdgeProvision,
pub outbound: EdgeProvision,
pub weight_source: WeightSource,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum StageEvent {
ProvisionStage {
from: NodeId,
provision: ProvisionStage,
},
WorkerReady,
WeightsReady,
InboundEdgeReady {
edge_id: EdgeId,
},
OutboundEdgeReady {
edge_id: EdgeId,
},
ObjectLoaded {
edge_id: EdgeId,
object_id: ObjectId,
sequence: u64,
handle: DeviceHandle,
},
StepCompleted {
step_id: StepId,
},
WorkerCrashed,
StopRun {
run_id: RunId,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum StageFaultReason {
UnauthorizedProvision,
SequenceViolation,
WorkerCrashed,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum StageLifecycleEvent {
StageReady {
run_id: RunId,
stage_index: u32,
},
StepAccepted {
run_id: RunId,
stage_index: u32,
sequence: u64,
},
StageFault {
run_id: RunId,
stage_index: u32,
reason: StageFaultReason,
},
StageStopped {
run_id: RunId,
stage_index: u32,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct StepInput {
pub edge_id: EdgeId,
pub object_id: ObjectId,
pub sequence: u64,
pub handle: DeviceHandle,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct OutputBinding {
pub edge_id: EdgeId,
pub sequence: u64,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ExecuteStep {
pub step_id: StepId,
pub input: StepInput,
pub outputs: Vec<OutputBinding>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum StageCommand {
EstablishInboundEdge {
edge_id: EdgeId,
},
EstablishOutboundEdge {
edge_id: EdgeId,
},
ConfigureWorkerRole {
run_id: RunId,
stage_index: u32,
layer_range: LayerRange,
},
LoadWeights {
source: WeightSource,
range: LayerRange,
},
RewireEdge {
edge_id: EdgeId,
},
ExecuteStep(ExecuteStep),
ReleaseInputHandle {
object_id: ObjectId,
handle: DeviceHandle,
},
StopLocalEdges {
run_id: RunId,
},
ReleaseRunDeviceObjects {
run_id: RunId,
},
}
#[cfg(test)]
pub struct StageControllerHarness {
local_node_id: NodeId,
provision: Option<ProvisionStage>,
worker_ready: bool,
weights_ready: bool,
inbound_ready: bool,
outbound_ready: bool,
stage_ready_emitted: bool,
busy: bool,
expected_sequence: u64,
active_input: Option<StepInput>,
commands: Vec<StageCommand>,
events: Vec<StageLifecycleEvent>,
faulted: bool,
stopped: bool,
}
#[cfg(test)]
impl StageControllerHarness {
pub fn new(local_node_id: NodeId) -> Self {
Self {
local_node_id,
provision: None,
worker_ready: false,
weights_ready: false,
inbound_ready: false,
outbound_ready: false,
stage_ready_emitted: false,
busy: false,
expected_sequence: 0,
active_input: None,
commands: Vec::new(),
events: Vec::new(),
faulted: false,
stopped: false,
}
}
pub fn observe(&mut self, event: StageEvent) {
match event {
StageEvent::ProvisionStage { from, provision } => self.provision(from, provision),
StageEvent::WorkerReady => self.worker_ready = true,
StageEvent::WeightsReady => self.weights_ready = true,
StageEvent::InboundEdgeReady { edge_id } => {
if self
.provision
.as_ref()
.is_some_and(|p| p.inbound.edge_id == edge_id)
{
self.inbound_ready = true;
}
}
StageEvent::OutboundEdgeReady { edge_id } => {
if self
.provision
.as_ref()
.is_some_and(|p| p.outbound.edge_id == edge_id)
{
self.outbound_ready = true;
}
}
StageEvent::ObjectLoaded {
edge_id,
object_id,
sequence,
handle,
} => self.object_loaded(edge_id, object_id, sequence, handle),
StageEvent::StepCompleted { .. } => self.step_completed(),
StageEvent::WorkerCrashed => self.fault(StageFaultReason::WorkerCrashed),
StageEvent::StopRun { run_id } => self.stop(run_id),
}
self.maybe_stage_ready();
}
pub fn commands(&self) -> &[StageCommand] {
&self.commands
}
pub fn events(&self) -> &[StageLifecycleEvent] {
&self.events
}
fn provision(&mut self, from: NodeId, provision: ProvisionStage) {
if from != provision.authorized_orchestrator || provision.node_id != self.local_node_id {
self.provision = Some(provision.clone());
self.fault(StageFaultReason::UnauthorizedProvision);
return;
}
self.commands.push(StageCommand::EstablishInboundEdge {
edge_id: provision.inbound.edge_id,
});
self.commands.push(StageCommand::EstablishOutboundEdge {
edge_id: provision.outbound.edge_id,
});
self.commands.push(StageCommand::ConfigureWorkerRole {
run_id: provision.run_id,
stage_index: provision.stage_index,
layer_range: provision.layer_range,
});
self.commands.push(StageCommand::LoadWeights {
source: provision.weight_source.clone(),
range: provision.layer_range,
});
self.provision = Some(provision);
}
fn maybe_stage_ready(&mut self) {
if self.stage_ready_emitted || self.faulted || self.stopped {
return;
}
if self.worker_ready && self.weights_ready && self.inbound_ready && self.outbound_ready {
if let Some(provision) = &self.provision {
self.stage_ready_emitted = true;
self.events.push(StageLifecycleEvent::StageReady {
run_id: provision.run_id,
stage_index: provision.stage_index,
});
}
}
}
fn object_loaded(
&mut self,
edge_id: EdgeId,
object_id: ObjectId,
sequence: u64,
handle: DeviceHandle,
) {
if self.faulted || self.stopped || !self.stage_ready_emitted || self.busy {
return;
}
let Some(provision) = &self.provision else {
return;
};
if edge_id != provision.inbound.edge_id {
return;
}
if sequence != self.expected_sequence {
self.fault(StageFaultReason::SequenceViolation);
return;
}
let input = StepInput {
edge_id,
object_id,
sequence,
handle,
};
let step = ExecuteStep {
step_id: StepId(sequence),
input: input.clone(),
outputs: vec![OutputBinding {
edge_id: provision.outbound.edge_id,
sequence,
}],
};
self.busy = true;
self.active_input = Some(input);
self.commands.push(StageCommand::ExecuteStep(step));
self.events.push(StageLifecycleEvent::StepAccepted {
run_id: provision.run_id,
stage_index: provision.stage_index,
sequence,
});
}
fn step_completed(&mut self) {
if self.faulted || self.stopped || !self.busy {
return;
}
if let Some(input) = self.active_input.take() {
self.commands.push(StageCommand::ReleaseInputHandle {
object_id: input.object_id,
handle: input.handle,
});
self.expected_sequence = input.sequence + 1;
}
self.busy = false;
}
fn fault(&mut self, reason: StageFaultReason) {
if self.faulted || self.stopped {
return;
}
self.faulted = true;
let (run_id, stage_index) = self
.provision
.as_ref()
.map(|p| (p.run_id, p.stage_index))
.unwrap_or((RunId(0), 0));
self.events.push(StageLifecycleEvent::StageFault {
run_id,
stage_index,
reason,
});
}
fn stop(&mut self, run_id: RunId) {
if self.stopped {
return;
}
self.stopped = true;
self.commands.push(StageCommand::StopLocalEdges { run_id });
self.commands
.push(StageCommand::ReleaseRunDeviceObjects { run_id });
let stage_index = self.provision.as_ref().map_or(0, |p| p.stage_index);
self.events.push(StageLifecycleEvent::StageStopped {
run_id,
stage_index,
});
}
}

View file

@ -74,7 +74,10 @@ fn arena_boot_creates_stable_offset_only_layout_domain() {
// Layout facts are arena offsets and stay under the reservation ceiling.
assert!(lease.layout.start_offset < arena_config().reservation_ceiling);
assert!(lease.layout.end_offset <= arena_config().reservation_ceiling);
assert!(matches!(lease.layout.pointer, arena::LayoutPointer::NoProcessPointer));
assert!(matches!(
lease.layout.pointer,
arena::LayoutPointer::NoProcessPointer
));
// Boot failure is typed and emits no usable arena.
let failed = arena::ArenaManagerHarness::boot(arena::ArenaConfig {
@ -232,9 +235,12 @@ fn release_requires_quiescence_and_reuse_happens_only_after_release() {
proof: arena::QuiescenceProof::verified(),
});
harness.request(arena::ArenaRequest::LeaseRing(lease_request(3, 1024, 64)));
assert!(harness.events().iter().any(|event| {
matches!(event, arena::ArenaEvent::RingLeased { .. })
}));
assert!(
harness
.events()
.iter()
.any(|event| { matches!(event, arena::ArenaEvent::RingLeased { .. }) })
);
}
// This proves shutdown rejects new leases, preserves live lease records, and

View file

@ -77,11 +77,17 @@ fn allocation_creates_current_generation_handle_or_typed_failure() {
fn host_to_device_copies_exact_range_and_defers_release_until_safe() {
// Allocate a device object and request an asynchronous copy.
let (mut harness, handle) = allocated_bridge();
let copy = harness.host_to_device(
device::HostRange { offset: 4, len: 8 },
device::DeviceRange { handle, offset: 0, len: 8 },
device::CopyMode::Async,
).expect("copy request must be accepted");
let copy = harness
.host_to_device(
device::HostRange { offset: 4, len: 8 },
device::DeviceRange {
handle,
offset: 0,
len: 8,
},
device::CopyMode::Async,
)
.expect("copy request must be accepted");
// The backend sees the exact ranges.
assert!(harness.backend_calls().iter().any(|call| {
@ -89,7 +95,11 @@ fn host_to_device_copies_exact_range_and_defers_release_until_safe() {
call,
device::BackendCall::HostToDevice {
host: device::HostRange { offset: 4, len: 8 },
device: device::DeviceRange { offset: 0, len: 8, .. },
device: device::DeviceRange {
offset: 0,
len: 8,
..
},
..
}
)
@ -110,18 +120,28 @@ fn host_to_device_copies_exact_range_and_defers_release_until_safe() {
fn device_to_host_copies_exact_range_and_defers_host_validity_until_safe() {
// Allocate a device object and request an asynchronous device-to-host copy.
let (mut harness, handle) = allocated_bridge();
let copy = harness.device_to_host(
device::DeviceRange { handle, offset: 0, len: 8 },
device::HostRange { offset: 12, len: 8 },
device::CopyMode::Async,
).expect("copy request must be accepted");
let copy = harness
.device_to_host(
device::DeviceRange {
handle,
offset: 0,
len: 8,
},
device::HostRange { offset: 12, len: 8 },
device::CopyMode::Async,
)
.expect("copy request must be accepted");
// The backend sees the exact ranges.
assert!(harness.backend_calls().iter().any(|call| {
matches!(
call,
device::BackendCall::DeviceToHost {
device: device::DeviceRange { offset: 0, len: 8, .. },
device: device::DeviceRange {
offset: 0,
len: 8,
..
},
host: device::HostRange { offset: 12, len: 8 },
..
}
@ -176,11 +196,17 @@ fn lifetime_blocks_free_until_dependencies_clear_and_rejects_old_generation() {
handle,
step_id: device::StepId(77),
});
let copy = harness.host_to_device(
device::HostRange { offset: 0, len: 8 },
device::DeviceRange { handle, offset: 0, len: 8 },
device::CopyMode::Async,
).expect("copy request must be accepted");
let copy = harness
.host_to_device(
device::HostRange { offset: 0, len: 8 },
device::DeviceRange {
handle,
offset: 0,
len: 8,
},
device::CopyMode::Async,
)
.expect("copy request must be accepted");
// Free is blocked while compute/copy depends on the allocation.
assert_eq!(
@ -195,9 +221,11 @@ fn lifetime_blocks_free_until_dependencies_clear_and_rejects_old_generation() {
});
harness.observe(device::DeviceEvent::CopyCompleted { copy });
assert_eq!(harness.free_device(handle), Ok(()));
assert!(harness.backend_calls().iter().any(|call| {
matches!(call, device::BackendCall::Free { freed } if *freed == handle)
}));
assert!(
harness.backend_calls().iter().any(|call| {
matches!(call, device::BackendCall::Free { freed } if *freed == handle)
})
);
// Restart invalidates prior handles.
harness.observe(device::DeviceEvent::WorkerRestarted {
@ -219,7 +247,11 @@ fn backend_copy_and_view_failures_are_typed() {
let copy_failure = harness
.host_to_device(
device::HostRange { offset: 0, len: 8 },
device::DeviceRange { handle, offset: 0, len: 8 },
device::DeviceRange {
handle,
offset: 0,
len: 8,
},
device::CopyMode::Sync,
)
.expect_err("copy failure must surface");

View file

@ -77,13 +77,22 @@ fn driver_owns_endpoint_connection_demux_and_pump_tasks() {
// Pump tasks are driver-owned.
assert!(harness.commands().iter().any(|command| {
matches!(command, driver::DriverCommand::SpawnSendPump { edge_id: driver::EdgeId(7001), .. })
matches!(
command,
driver::DriverCommand::SpawnSendPump {
edge_id: driver::EdgeId(7001),
..
}
)
}));
// Actor commands must not expose stream polling.
assert!(!harness.actor_messages().iter().any(|message| {
matches!(message, driver::ActorMessage::PollStreamFuture { .. })
}));
assert!(
!harness
.actor_messages()
.iter()
.any(|message| { matches!(message, driver::ActorMessage::PollStreamFuture { .. }) })
);
}
// This proves each edge uses one persistent uni-stream, writes an edge-id
@ -131,15 +140,24 @@ fn recv_rendezvous_starts_pump_only_after_spec_and_stream_exist() {
// Spec first, stream second.
let mut spec_first = new_driver();
spec_first.observe(driver::DriverEvent::EstablishRecv(recv_spec()));
assert!(!spec_first.commands().iter().any(|command| {
matches!(command, driver::DriverCommand::SpawnRecvPump { .. })
}));
assert!(
!spec_first
.commands()
.iter()
.any(|command| { matches!(command, driver::DriverCommand::SpawnRecvPump { .. }) })
);
spec_first.observe(driver::DriverEvent::IncomingUniStream {
edge_id: driver::EdgeId(7001),
stream_id: driver::StreamId(1),
});
assert!(spec_first.commands().iter().any(|command| {
matches!(command, driver::DriverCommand::SpawnRecvPump { edge_id: driver::EdgeId(7001), .. })
matches!(
command,
driver::DriverCommand::SpawnRecvPump {
edge_id: driver::EdgeId(7001),
..
}
)
}));
// Stream first, spec second.
@ -172,12 +190,20 @@ fn recv_pump_copies_bytes_without_parsing_and_respects_backpressure() {
edge_id: driver::EdgeId(7001),
bytes: driver::fake_object_header_bytes(),
});
assert!(!harness.events().iter().any(|event| {
matches!(event, driver::DriverEventOut::ObjectHeaderParsed { .. })
}));
assert!(
!harness
.events()
.iter()
.any(|event| { matches!(event, driver::DriverEventOut::ObjectHeaderParsed { .. }) })
);
assert!(harness.ring_commit(driver::EdgeId(7001)) > 0);
assert!(harness.wake_hints().iter().any(|wake| {
matches!(wake, driver::WakeHint::RingReadable { edge_id: driver::EdgeId(7001) })
matches!(
wake,
driver::WakeHint::RingReadable {
edge_id: driver::EdgeId(7001)
}
)
}));
// With no ring space, the pump stops reading and waits for RingWritable.
@ -217,7 +243,12 @@ fn send_pump_advances_consume_only_after_write_acceptance() {
});
assert_eq!(harness.ring_consume(driver::EdgeId(7001)), 7);
assert!(harness.wake_hints().iter().any(|wake| {
matches!(wake, driver::WakeHint::RingWritable { edge_id: driver::EdgeId(7001) })
matches!(
wake,
driver::WakeHint::RingWritable {
edge_id: driver::EdgeId(7001)
}
)
}));
}
@ -236,7 +267,13 @@ fn driver_faults_and_stop_emit_stream_fault_or_pump_stopped() {
edge_id: driver::EdgeId(7001),
});
assert!(recv.events().iter().any(|event| {
matches!(event, driver::DriverEventOut::StreamFault { edge_id: driver::EdgeId(7001), .. })
matches!(
event,
driver::DriverEventOut::StreamFault {
edge_id: driver::EdgeId(7001),
..
}
)
}));
// Write error faults the send edge.
@ -246,7 +283,13 @@ fn driver_faults_and_stop_emit_stream_fault_or_pump_stopped() {
edge_id: driver::EdgeId(7001),
});
assert!(send.events().iter().any(|event| {
matches!(event, driver::DriverEventOut::StreamFault { edge_id: driver::EdgeId(7001), .. })
matches!(
event,
driver::DriverEventOut::StreamFault {
edge_id: driver::EdgeId(7001),
..
}
)
}));
// StopEdge stops the corresponding pump.
@ -254,6 +297,12 @@ fn driver_faults_and_stop_emit_stream_fault_or_pump_stopped() {
edge_id: driver::EdgeId(7001),
});
assert!(send.events().iter().any(|event| {
matches!(event, driver::DriverEventOut::PumpStopped { edge_id: driver::EdgeId(7001), .. })
matches!(
event,
driver::DriverEventOut::PumpStopped {
edge_id: driver::EdgeId(7001),
..
}
)
}));
}

View file

@ -89,7 +89,10 @@ fn provisioning_creates_local_edge_records_without_remote_actor_addresses() {
assert!(rx.local_record(edge::EdgeId(7001)).is_some());
// Data-flow provisioning must not require remote actor addresses.
for record in [tx.local_record(edge::EdgeId(7001)).unwrap(), rx.local_record(edge::EdgeId(7001)).unwrap()] {
for record in [
tx.local_record(edge::EdgeId(7001)).unwrap(),
rx.local_record(edge::EdgeId(7001)).unwrap(),
] {
assert!(record.remote_actor_address.is_none());
}
}
@ -109,9 +112,12 @@ fn lease_flow_matches_records_and_suppresses_stale_or_cancelled_leases() {
});
// Mismatched lease must not install worker or pump state.
assert!(!harness.commands().iter().any(|command| {
matches!(command, edge::EdgeCommand::InstallWorkerRing { .. })
}));
assert!(
!harness
.commands()
.iter()
.any(|command| { matches!(command, edge::EdgeCommand::InstallWorkerRing { .. }) })
);
// A matching rejection faults the record.
harness.observe(edge::EdgeEvent::RingLeaseRejected {
@ -164,9 +170,12 @@ fn worker_ring_install_precedes_driver_establishment_and_uses_provision_specs()
});
// Before ring installation, no driver command may be issued.
assert!(!harness.commands().iter().any(|command| {
matches!(command, edge::EdgeCommand::EstablishSend { .. })
}));
assert!(
!harness
.commands()
.iter()
.any(|command| { matches!(command, edge::EdgeCommand::EstablishSend { .. }) })
);
// Matching ring installation advances establishment.
harness.observe(edge::EdgeEvent::RingInstalled {
@ -188,7 +197,13 @@ fn worker_ring_install_precedes_driver_establishment_and_uses_provision_specs()
// Now the driver may be established.
assert!(harness.commands().iter().any(|command| {
matches!(command, edge::EdgeCommand::EstablishSend { edge_id: edge::EdgeId(7001), .. })
matches!(
command,
edge::EdgeCommand::EstablishSend {
edge_id: edge::EdgeId(7001),
..
}
)
}));
}
@ -212,9 +227,11 @@ fn driver_ready_marks_local_edge_actor_ready() {
}));
// The edge is not ready until DriverEdgeReady arrives.
assert!(!tx.events().iter().any(|event| {
matches!(event, edge::EdgeLifecycleEvent::EdgeReady { .. })
}));
assert!(
!tx.events()
.iter()
.any(|event| { matches!(event, edge::EdgeLifecycleEvent::EdgeReady { .. }) })
);
tx.observe(edge::EdgeEvent::DriverEdgeReady {
edge_id: edge::EdgeId(7001),
});
@ -230,9 +247,11 @@ fn driver_ready_marks_local_edge_actor_ready() {
// After readiness, stream and pump behavior belongs to the driver; the
// establisher should not emit hot-path byte commands.
assert!(!tx.commands().iter().any(|command| {
matches!(command, edge::EdgeCommand::CopyHotPathBytes { .. })
}));
assert!(
!tx.commands()
.iter()
.any(|command| { matches!(command, edge::EdgeCommand::CopyHotPathBytes { .. }) })
);
}
// This proves StopEdge cancels queued leases, stops pumps, uninstalls worker
@ -251,20 +270,32 @@ fn stop_edge_tears_down_local_state_and_terminal_stopped_ignores_late_events() {
});
// Stop commands must cover queued lease, pump, and worker ring cleanup.
assert!(harness.commands().iter().any(|command| {
matches!(command, edge::EdgeCommand::CancelQueuedLease { .. })
}));
assert!(harness.commands().iter().any(|command| {
matches!(command, edge::EdgeCommand::StopPump { .. })
}));
assert!(harness.commands().iter().any(|command| {
matches!(command, edge::EdgeCommand::UninstallWorkerRing { .. })
}));
assert!(
harness
.commands()
.iter()
.any(|command| { matches!(command, edge::EdgeCommand::CancelQueuedLease { .. }) })
);
assert!(
harness
.commands()
.iter()
.any(|command| { matches!(command, edge::EdgeCommand::StopPump { .. }) })
);
assert!(
harness
.commands()
.iter()
.any(|command| { matches!(command, edge::EdgeCommand::UninstallWorkerRing { .. }) })
);
// Arena release is withheld until quiescence proof arrives.
assert!(!harness.commands().iter().any(|command| {
matches!(command, edge::EdgeCommand::ReleaseArenaLease { .. })
}));
assert!(
!harness
.commands()
.iter()
.any(|command| { matches!(command, edge::EdgeCommand::ReleaseArenaLease { .. }) })
);
harness.observe(edge::EdgeEvent::QuiescenceProven {
ring_id: edge::RingId(8001),
});

View file

@ -70,24 +70,35 @@ fn worker_lifecycle_runs_start_initialize_ready_and_fault_paths() {
// Start the worker process.
let mut harness = new_controller();
harness.observe(ctl::WorkerCtlEvent::StartWorker);
assert!(harness.commands().iter().any(|command| {
matches!(command, ctl::WorkerCtlCommand::SpawnProcessActor { .. })
}));
assert!(
harness
.commands()
.iter()
.any(|command| { matches!(command, ctl::WorkerCtlCommand::SpawnProcessActor { .. }) })
);
// Process start triggers InitializeWorker.
harness.observe(ctl::WorkerCtlEvent::ProcessStarted {
pid: ctl::ProcessId(1234),
});
assert!(harness.serialized_worker_commands().iter().any(|command| {
matches!(command, ctl::WorkerCommand::InitializeWorker { .. })
}));
assert!(
harness
.serialized_worker_commands()
.iter()
.any(|command| { matches!(command, ctl::WorkerCommand::InitializeWorker { .. }) })
);
// WorkerReady emits a running lifecycle event.
harness.observe(ctl::WorkerCtlEvent::WorkerReady {
generation: ctl::WorkerGeneration(1),
});
assert!(harness.events().iter().any(|event| {
matches!(event, ctl::WorkerCtlOut::WorkerRunning { generation: ctl::WorkerGeneration(1) })
matches!(
event,
ctl::WorkerCtlOut::WorkerRunning {
generation: ctl::WorkerGeneration(1)
}
)
}));
// Initialization timeout in a fresh controller is terminal failure.
@ -95,7 +106,13 @@ fn worker_lifecycle_runs_start_initialize_ready_and_fault_paths() {
timed_out.observe(ctl::WorkerCtlEvent::StartWorker);
timed_out.advance_time_ms(1_001);
assert!(timed_out.events().iter().any(|event| {
matches!(event, ctl::WorkerCtlOut::WorkerFailed { reason: ctl::WorkerFailure::InitializationTimeout, .. })
matches!(
event,
ctl::WorkerCtlOut::WorkerFailed {
reason: ctl::WorkerFailure::InitializationTimeout,
..
}
)
}));
}
@ -112,7 +129,10 @@ fn command_routing_requires_running_current_generation_and_is_payload_free() {
}
// All valid commands are serialized to the worker.
assert_eq!(harness.serialized_worker_commands().len(), current_generation_commands().len() + 1);
assert_eq!(
harness.serialized_worker_commands().len(),
current_generation_commands().len() + 1
);
// Serialized commands must not contain payload bytes.
for command in harness.serialized_worker_commands() {
@ -136,13 +156,21 @@ fn command_routing_requires_running_current_generation_and_is_payload_free() {
harness.observe(ctl::WorkerCtlEvent::WorkerReady {
generation: ctl::WorkerGeneration(2),
});
harness.observe(ctl::WorkerCtlEvent::ActorCommand(ctl::ActorCommand::ExecuteStep {
generation: ctl::WorkerGeneration(1),
step_id: ctl::StepId(9002),
input: ctl::DeviceHandle::new(ctl::WorkerGeneration(1), 42),
}));
harness.observe(ctl::WorkerCtlEvent::ActorCommand(
ctl::ActorCommand::ExecuteStep {
generation: ctl::WorkerGeneration(1),
step_id: ctl::StepId(9002),
input: ctl::DeviceHandle::new(ctl::WorkerGeneration(1), 42),
},
));
assert!(harness.events().iter().any(|event| {
matches!(event, ctl::WorkerCtlOut::CommandRejected { reason: ctl::CommandRejection::OldGenerationHandle, .. })
matches!(
event,
ctl::WorkerCtlOut::CommandRejected {
reason: ctl::CommandRejection::OldGenerationHandle,
..
}
)
}));
}
@ -154,31 +182,66 @@ fn parsed_worker_events_route_to_their_control_owners() {
let mut harness = running_controller();
// Deliver each worker event family through stdout parsing.
harness.observe(ctl::WorkerCtlEvent::StdoutEvent(ctl::WorkerEvent::RingInstalled {
ring_id: ctl::RingId(8001),
}));
harness.observe(ctl::WorkerCtlEvent::StdoutEvent(ctl::WorkerEvent::ObjectLoaded {
object_id: ctl::ObjectId(9000),
sequence: 0,
}));
harness.observe(ctl::WorkerCtlEvent::StdoutEvent(ctl::WorkerEvent::ObjectProduced {
object_id: ctl::ObjectId(9001),
sequence: 0,
}));
harness.observe(ctl::WorkerCtlEvent::StdoutEvent(ctl::WorkerEvent::StepCompleted {
step_id: ctl::StepId(77),
}));
harness.observe(ctl::WorkerCtlEvent::StdoutEvent(ctl::WorkerEvent::RingReadable {
ring_id: ctl::RingId(8001),
}));
harness.observe(ctl::WorkerCtlEvent::StdoutEvent(
ctl::WorkerEvent::RingInstalled {
ring_id: ctl::RingId(8001),
},
));
harness.observe(ctl::WorkerCtlEvent::StdoutEvent(
ctl::WorkerEvent::ObjectLoaded {
object_id: ctl::ObjectId(9000),
sequence: 0,
},
));
harness.observe(ctl::WorkerCtlEvent::StdoutEvent(
ctl::WorkerEvent::ObjectProduced {
object_id: ctl::ObjectId(9001),
sequence: 0,
},
));
harness.observe(ctl::WorkerCtlEvent::StdoutEvent(
ctl::WorkerEvent::StepCompleted {
step_id: ctl::StepId(77),
},
));
harness.observe(ctl::WorkerCtlEvent::StdoutEvent(
ctl::WorkerEvent::RingReadable {
ring_id: ctl::RingId(8001),
},
));
// Routing is proven by destination commands/events, not private dispatch
// tables.
assert!(harness.routed().iter().any(|route| matches!(route, ctl::RoutedEvent::ToEdgeEstablisher(_))));
assert!(harness.routed().iter().any(|route| matches!(route, ctl::RoutedEvent::ToRxOrRole(_))));
assert!(harness.routed().iter().any(|route| matches!(route, ctl::RoutedEvent::ToTxOrRole(_))));
assert!(harness.routed().iter().any(|route| matches!(route, ctl::RoutedEvent::ToStageController(_))));
assert!(harness.routed().iter().any(|route| matches!(route, ctl::RoutedEvent::ToDriverOrWorkerSide(_))));
assert!(
harness
.routed()
.iter()
.any(|route| matches!(route, ctl::RoutedEvent::ToEdgeEstablisher(_)))
);
assert!(
harness
.routed()
.iter()
.any(|route| matches!(route, ctl::RoutedEvent::ToRxOrRole(_)))
);
assert!(
harness
.routed()
.iter()
.any(|route| matches!(route, ctl::RoutedEvent::ToTxOrRole(_)))
);
assert!(
harness
.routed()
.iter()
.any(|route| matches!(route, ctl::RoutedEvent::ToStageController(_)))
);
assert!(
harness
.routed()
.iter()
.any(|route| matches!(route, ctl::RoutedEvent::ToDriverOrWorkerSide(_)))
);
}
// This proves worker crash invalidates old handles, roles, rings, in-flight
@ -188,15 +251,19 @@ fn parsed_worker_events_route_to_their_control_owners() {
fn crash_invalidates_generation_state_and_fans_out_faults() {
// Install one ring and start one step in generation 1.
let mut harness = running_controller();
harness.observe(ctl::WorkerCtlEvent::ActorCommand(ctl::ActorCommand::InstallRing {
generation: ctl::WorkerGeneration(1),
ring_id: ctl::RingId(8001),
}));
harness.observe(ctl::WorkerCtlEvent::ActorCommand(ctl::ActorCommand::ExecuteStep {
generation: ctl::WorkerGeneration(1),
step_id: ctl::StepId(9001),
input: ctl::DeviceHandle::new(ctl::WorkerGeneration(1), 42),
}));
harness.observe(ctl::WorkerCtlEvent::ActorCommand(
ctl::ActorCommand::InstallRing {
generation: ctl::WorkerGeneration(1),
ring_id: ctl::RingId(8001),
},
));
harness.observe(ctl::WorkerCtlEvent::ActorCommand(
ctl::ActorCommand::ExecuteStep {
generation: ctl::WorkerGeneration(1),
step_id: ctl::StepId(9001),
input: ctl::DeviceHandle::new(ctl::WorkerGeneration(1), 42),
},
));
// Crash the worker process.
harness.observe(ctl::WorkerCtlEvent::ProcessExited {
@ -205,10 +272,22 @@ fn crash_invalidates_generation_state_and_fans_out_faults() {
// Installed rings fault and affected pumps are stopped.
assert!(harness.events().iter().any(|event| {
matches!(event, ctl::WorkerCtlOut::RingFaulted { ring_id: ctl::RingId(8001), .. })
matches!(
event,
ctl::WorkerCtlOut::RingFaulted {
ring_id: ctl::RingId(8001),
..
}
)
}));
assert!(harness.commands().iter().any(|command| {
matches!(command, ctl::WorkerCtlCommand::StopDriverPump { ring_id: ctl::RingId(8001), .. })
matches!(
command,
ctl::WorkerCtlCommand::StopDriverPump {
ring_id: ctl::RingId(8001),
..
}
)
}));
// Restart increments worker generation.
@ -229,16 +308,21 @@ fn crash_invalidates_generation_state_and_fans_out_faults() {
fn graceful_shutdown_sends_worker_shutdown_and_reaches_terminal_stopped() {
// Start from a running worker with an installed ring.
let mut harness = running_controller();
harness.observe(ctl::WorkerCtlEvent::ActorCommand(ctl::ActorCommand::InstallRing {
generation: ctl::WorkerGeneration(1),
ring_id: ctl::RingId(8001),
}));
harness.observe(ctl::WorkerCtlEvent::ActorCommand(
ctl::ActorCommand::InstallRing {
generation: ctl::WorkerGeneration(1),
ring_id: ctl::RingId(8001),
},
));
// Request graceful shutdown.
harness.observe(ctl::WorkerCtlEvent::ShutdownRequested);
assert!(harness.serialized_worker_commands().iter().any(|command| {
matches!(command, ctl::WorkerCommand::ShutdownWorker { .. })
}));
assert!(
harness
.serialized_worker_commands()
.iter()
.any(|command| { matches!(command, ctl::WorkerCommand::ShutdownWorker { .. }) })
);
// WorkerStopped must precede terminal stopped.
harness.observe(ctl::WorkerCtlEvent::WorkerStopped {
@ -261,6 +345,12 @@ fn graceful_shutdown_sends_worker_shutdown_and_reaches_terminal_stopped() {
// Rings are marked quiesced on graceful stop.
assert!(harness.events().iter().any(|event| {
matches!(event, ctl::WorkerCtlOut::RingQuiesced { ring_id: ctl::RingId(8001), .. })
matches!(
event,
ctl::WorkerCtlOut::RingQuiesced {
ring_id: ctl::RingId(8001),
..
}
)
}));
}

View file

@ -109,7 +109,12 @@ fn header_is_written_and_committed_before_payload() {
// Payload bytes are not committed before the payload copy is valid.
assert_eq!(harness.committed_payload_bytes(egress::RingId(8002)), 0);
assert!(harness.wake_hints().iter().any(|wake| {
matches!(wake, egress::WakeHint::RingReadable { ring_id: egress::RingId(8002) })
matches!(
wake,
egress::WakeHint::RingReadable {
ring_id: egress::RingId(8002)
}
)
}));
}
@ -156,29 +161,37 @@ fn payload_copy_is_exact_extent_and_respects_backpressure() {
fn object_produced_precedes_step_completed_after_all_outputs() {
// Execute a step with two outputs.
let mut harness = installed_producer();
harness.observe(egress::WorkerEgressEvent::InstallRing(egress::InstallRing {
ring_id: egress::RingId(8003),
edge_id: egress::EdgeId(7003),
port_id: egress::PortId("out2".into()),
..egress_ring()
}));
harness.observe(egress::WorkerEgressEvent::InstallRing(
egress::InstallRing {
ring_id: egress::RingId(8003),
edge_id: egress::EdgeId(7003),
port_id: egress::PortId("out2".into()),
..egress_ring()
},
));
harness.observe(egress::WorkerEgressEvent::ExecuteStep {
step_id: egress::StepId(77),
outputs: vec![output_binding(0, 8), egress::OutputBinding {
ring_id: egress::RingId(8003),
object_id: egress::ObjectId(9100),
sequence: 0,
extent: 8,
flags: egress::ObjectFlags::default(),
device_source: egress::DeviceHandle::new(egress::WorkerGeneration(1), 55),
}],
outputs: vec![
output_binding(0, 8),
egress::OutputBinding {
ring_id: egress::RingId(8003),
object_id: egress::ObjectId(9100),
sequence: 0,
extent: 8,
flags: egress::ObjectFlags::default(),
device_source: egress::DeviceHandle::new(egress::WorkerGeneration(1), 55),
},
],
});
// Produce only the first output and prove StepCompleted is still absent.
harness.complete_output(egress::ObjectId(9000));
assert!(!harness.events().iter().any(|event| {
matches!(event, egress::WorkerEgressOut::StepCompleted { .. })
}));
assert!(
!harness
.events()
.iter()
.any(|event| { matches!(event, egress::WorkerEgressOut::StepCompleted { .. }) })
);
// Produce the second output and complete role state update.
harness.complete_output(egress::ObjectId(9100));
@ -190,17 +203,41 @@ fn object_produced_precedes_step_completed_after_all_outputs() {
let first_object_pos = harness
.events()
.iter()
.position(|event| matches!(event, egress::WorkerEgressOut::ObjectProduced { object_id: egress::ObjectId(9000), .. }))
.position(|event| {
matches!(
event,
egress::WorkerEgressOut::ObjectProduced {
object_id: egress::ObjectId(9000),
..
}
)
})
.expect("first object produced");
let second_object_pos = harness
.events()
.iter()
.position(|event| matches!(event, egress::WorkerEgressOut::ObjectProduced { object_id: egress::ObjectId(9100), .. }))
.position(|event| {
matches!(
event,
egress::WorkerEgressOut::ObjectProduced {
object_id: egress::ObjectId(9100),
..
}
)
})
.expect("second object produced");
let completed_pos = harness
.events()
.iter()
.position(|event| matches!(event, egress::WorkerEgressOut::StepCompleted { step_id: egress::StepId(77), .. }))
.position(|event| {
matches!(
event,
egress::WorkerEgressOut::StepCompleted {
step_id: egress::StepId(77),
..
}
)
})
.expect("step completed");
assert!(first_object_pos < completed_pos);
assert!(second_object_pos < completed_pos);
@ -220,7 +257,13 @@ fn egress_faults_are_visible_and_suppress_success_events() {
}],
});
assert!(invalid_ring.events().iter().any(|event| {
matches!(event, egress::WorkerEgressOut::StepFailed { reason: egress::StepFailureReason::InvalidOutputRing, .. })
matches!(
event,
egress::WorkerEgressOut::StepFailed {
reason: egress::StepFailureReason::InvalidOutputRing,
..
}
)
}));
// Extent violation fails the step.
@ -230,7 +273,13 @@ fn egress_faults_are_visible_and_suppress_success_events() {
outputs: vec![output_binding(0, 32)],
});
assert!(bad_extent.events().iter().any(|event| {
matches!(event, egress::WorkerEgressOut::StepFailed { reason: egress::StepFailureReason::OutputExtentViolation, .. })
matches!(
event,
egress::WorkerEgressOut::StepFailed {
reason: egress::StepFailureReason::OutputExtentViolation,
..
}
)
}));
// Copy failure faults the ring or fails the step, but must not emit
@ -247,7 +296,10 @@ fn egress_faults_are_visible_and_suppress_success_events() {
matches!(event, egress::WorkerEgressOut::StepFailed { .. })
|| matches!(event, egress::WorkerEgressOut::RingFault { .. })
}));
assert!(!copy_failed.events().iter().any(|event| {
matches!(event, egress::WorkerEgressOut::ObjectProduced { .. })
}));
assert!(
!copy_failed
.events()
.iter()
.any(|event| { matches!(event, egress::WorkerEgressOut::ObjectProduced { .. }) })
);
}

View file

@ -163,10 +163,7 @@ fn payload_loading_is_exact_extent_and_release_after_copy_completion() {
object_id: ingress::ObjectId(9000),
byte_count: 8,
});
assert_eq!(
harness.device_copy_log().last().unwrap().byte_count,
8
);
assert_eq!(harness.device_copy_log().last().unwrap().byte_count, 8);
assert!(harness.consume_cursor(ingress::RingId(8001)) > before_copy_complete);
// EOF mid-object faults the object.
@ -208,9 +205,12 @@ fn object_loaded_requires_complete_valid_object_and_current_handle() {
// No ObjectLoaded may appear before device copy completion and handle
// creation.
assert!(!harness.events().iter().any(|event| {
matches!(event, ingress::WorkerIngressOut::ObjectLoaded { .. })
}));
assert!(
!harness
.events()
.iter()
.any(|event| { matches!(event, ingress::WorkerIngressOut::ObjectLoaded { .. }) })
);
// Complete device work.
harness.observe(ingress::WorkerIngressEvent::DeviceCopyCompleted {

View file

@ -84,9 +84,12 @@ fn control_stream_is_line_framed_json_without_payload_bytes() {
// Stderr alone is diagnostic and does not define lifecycle state.
harness.receive_stderr_line("loading tinygrad backend");
assert!(!harness.events().iter().any(|event| {
matches!(event, adapter::AdapterEvent::WorkerFatal { .. })
}));
assert!(
!harness
.events()
.iter()
.any(|event| { matches!(event, adapter::AdapterEvent::WorkerFatal { .. }) })
);
}
// This proves worker initialization reads arena environment, waits for
@ -130,16 +133,24 @@ fn initialization_order_is_env_initialize_map_helper_backend_ready() {
// Successful initialization emits WorkerReady.
harness.receive_stdout_line(r#"{"type":"WorkerReady","generation":1}"#);
assert!(harness.events().iter().any(|event| {
matches!(event, adapter::AdapterEvent::WorkerReady { generation: adapter::WorkerGeneration(1) })
matches!(
event,
adapter::AdapterEvent::WorkerReady {
generation: adapter::WorkerGeneration(1)
}
)
}));
// Initialization failure emits WorkerFatal if possible and exits non-zero.
let mut failed = new_adapter();
failed.start_worker_process();
failed.inject_initialization_failure(adapter::InitializationFailure::BackendUnavailable);
assert!(failed.events().iter().any(|event| {
matches!(event, adapter::AdapterEvent::WorkerFatal { .. })
}));
assert!(
failed
.events()
.iter()
.any(|event| { matches!(event, adapter::AdapterEvent::WorkerFatal { .. }) })
);
assert_ne!(failed.exit_status(), Some(adapter::ExitStatus::Code(0)));
}
@ -151,14 +162,26 @@ fn parsing_and_abi_errors_fault_the_worker_process() {
let mut invalid_json = new_adapter();
invalid_json.receive_stdout_line("{not-json");
assert!(invalid_json.events().iter().any(|event| {
matches!(event, adapter::AdapterEvent::ProcessFault { reason: adapter::ProcessFaultReason::InvalidJson, .. })
matches!(
event,
adapter::AdapterEvent::ProcessFault {
reason: adapter::ProcessFaultReason::InvalidJson,
..
}
)
}));
// Unknown event shape faults the adapter.
let mut unknown = new_adapter();
unknown.receive_stdout_line(r#"{"type":"NotAWorkerEvent"}"#);
assert!(unknown.events().iter().any(|event| {
matches!(event, adapter::AdapterEvent::ProcessFault { reason: adapter::ProcessFaultReason::UnknownEventShape, .. })
matches!(
event,
adapter::AdapterEvent::ProcessFault {
reason: adapter::ProcessFaultReason::UnknownEventShape,
..
}
)
}));
// Unsupported helper ABI emits WorkerFatal.
@ -167,7 +190,13 @@ fn parsing_and_abi_errors_fault_the_worker_process() {
helper_abi_version: adapter::HelperAbiVersion(999),
});
assert!(abi.events().iter().any(|event| {
matches!(event, adapter::AdapterEvent::WorkerFatal { reason: adapter::WorkerFatalReason::UnsupportedHelperAbi, .. })
matches!(
event,
adapter::AdapterEvent::WorkerFatal {
reason: adapter::WorkerFatalReason::UnsupportedHelperAbi,
..
}
)
}));
}
@ -186,14 +215,20 @@ fn command_discipline_rejects_payload_bearing_control_messages() {
// Payload bytes in a control command are rejected at the adapter boundary.
harness.send_raw_json_command(r#"{"type":"ExecuteStep","step_id":1,"payload":[1,2,3]}"#);
assert!(harness.command_rejections().iter().any(|rejection| {
matches!(rejection.reason, adapter::CommandRejectionReason::PayloadBytesForbidden)
matches!(
rejection.reason,
adapter::CommandRejectionReason::PayloadBytesForbidden
)
}));
// Wake hints reload cursors; they do not carry byte ranges or credits.
harness.send_command(adapter::WorkerCommand::RingReadable {
ring_id: adapter::RingId(8001),
});
let wake_line = harness.stdin_lines().last().expect("wake command must be written");
let wake_line = harness
.stdin_lines()
.last()
.expect("wake command must be written");
let parsed = adapter::JsonLine::parse(wake_line).expect("wake line must parse");
assert_no_payload_bytes(&parsed);
assert!(!parsed.contains_key("range") && !parsed.contains_key("credits"));

View file

@ -34,10 +34,7 @@ fn readiness_config() -> membership::ReadinessConfig {
// This helper provides every required public fact for one node. Tests use it to
// build complete and deliberately incomplete pool views without observing any
// private readiness bookkeeping.
fn report_node_ready(
gate: &mut membership::ReadinessGateHarness,
node_id: membership::NodeId,
) {
fn report_node_ready(gate: &mut membership::ReadinessGateHarness, node_id: membership::NodeId) {
gate.observe(membership::Observation::NodeKnown { node_id });
gate.observe(membership::Observation::SwimLive { node_id });
gate.observe(membership::Observation::NodeAvailable { node_id });
@ -88,8 +85,14 @@ fn pool_ready_requires_complete_stable_candidate_pool() {
.expect("complete stable pool must emit PoolReady");
// Compare as sets so ordering is not part of the behavioral contract.
let observed = ready.iter().copied().collect::<std::collections::BTreeSet<_>>();
let expected = pool.iter().copied().collect::<std::collections::BTreeSet<_>>();
let observed = ready
.iter()
.copied()
.collect::<std::collections::BTreeSet<_>>();
let expected = pool
.iter()
.copied()
.collect::<std::collections::BTreeSet<_>>();
assert_eq!(observed, expected);
}
@ -142,9 +145,11 @@ fn planning_starts_only_after_pool_ready_and_stops_if_readiness_is_lost() {
gate.request_run_planning(membership::RunRequest::new(membership::RunId(7)));
// With no PoolReady event, there must be no planning command.
assert!(!gate.commands().iter().any(|command| {
matches!(command, membership::ReadinessCommand::StartPlanning { .. })
}));
assert!(
!gate.commands().iter().any(|command| {
matches!(command, membership::ReadinessCommand::StartPlanning { .. })
})
);
// Satisfy readiness, then immediately lose a required node before plan
// commit. The policy may wait or abort, but it must not commit placement.
@ -155,9 +160,11 @@ fn planning_starts_only_after_pool_ready_and_stops_if_readiness_is_lost() {
gate.observe(membership::Observation::SwimLost { node_id: pool[1] });
// No plan commitment command may be emitted from an unstable pool view.
assert!(!gate.commands().iter().any(|command| {
matches!(command, membership::ReadinessCommand::CommitRunPlan { .. })
}));
assert!(
!gate.commands().iter().any(|command| {
matches!(command, membership::ReadinessCommand::CommitRunPlan { .. })
})
);
}
// This proves membership loss after provisioning is a run fault, not active
@ -180,16 +187,20 @@ fn required_node_loss_after_provisioning_faults_without_replacement() {
gate.observe(membership::Observation::SwimLost { node_id: pool[1] });
// The run must fault with a membership reason.
assert!(gate.events().contains(&membership::ReadinessEvent::RunFaulted {
run_id: membership::RunId(7),
reason: membership::RunFaultReason::RequiredNodeLost {
node_id: pool[1],
},
}));
assert!(
gate.events()
.contains(&membership::ReadinessEvent::RunFaulted {
run_id: membership::RunId(7),
reason: membership::RunFaultReason::RequiredNodeLost { node_id: pool[1] },
})
);
// Re-placement would violate the committed-plan authority boundary.
assert!(!gate.commands().iter().any(|command| {
matches!(command, membership::ReadinessCommand::RecomputePlacement { .. })
matches!(
command,
membership::ReadinessCommand::RecomputePlacement { .. }
)
}));
}
@ -212,7 +223,9 @@ fn swim_observations_do_not_create_graph_assignments() {
| membership::ReadinessCommand::StartPlanning { .. }
| membership::ReadinessCommand::WaitForStability { .. }
| membership::ReadinessCommand::AbortPendingRun { .. } => {}
membership::ReadinessCommand::AssignStage { .. }
membership::ReadinessCommand::CommitRunPlan { .. }
| membership::ReadinessCommand::RecomputePlacement { .. }
| membership::ReadinessCommand::AssignStage { .. }
| membership::ReadinessCommand::AssignEdge { .. }
| membership::ReadinessCommand::AssignLayerRange { .. }
| membership::ReadinessCommand::AssignObjectSpec { .. } => {

View file

@ -0,0 +1,18 @@
mod arena_manager_guarantees;
mod device_bridge_guarantees;
mod edge_establisher_guarantees;
mod gpu_worker_ctl_guarantees;
mod gpu_worker_egress_producer_guarantees;
mod gpu_worker_ingress_parser_guarantees;
mod gpu_worker_process_adapter_guarantees;
mod membership_pool_readiness_guarantees;
mod node_boot_lifecycle_guarantees;
mod observability_surface_guarantees;
mod orchestrator_run_fsm_guarantees;
mod orchestrator_token_endpoint_guarantees;
mod resource_inventory_guarantees;
mod run_plan_guarantees;
mod shared_ring_helper_abi_guarantees;
mod stage_controller_guarantees;
mod tx_rx_edge_actor_guarantees;
mod weight_lifecycle_guarantees;

View file

@ -75,9 +75,13 @@ fn node_available_waits_for_all_required_readiness_facts() {
let final_fact = facts.pop().expect("fixture has a final fact");
for fact in facts {
harness.observe(fact);
assert!(!harness.events().contains(&boot::LifecycleEvent::NodeAvailable {
node_id: boot::NodeId(10),
}));
assert!(
!harness
.events()
.contains(&boot::LifecycleEvent::NodeAvailable {
node_id: boot::NodeId(10),
})
);
}
// Once the final required fact arrives, availability becomes observable.
@ -160,15 +164,22 @@ fn boot_resource_failure_faults_without_node_availability() {
// The emitted fault must carry the stable reason enum for operators and
// tests; logs are not part of the contract.
assert!(harness.events().contains(&boot::LifecycleEvent::NodeFaulted {
node_id: boot::NodeId(10),
kind: expected_kind,
}));
assert!(
harness
.events()
.contains(&boot::LifecycleEvent::NodeFaulted {
node_id: boot::NodeId(10),
kind: expected_kind,
})
);
// A faulted boot attempt must not also become available.
assert!(!harness.events().iter().any(|event| {
matches!(event, boot::LifecycleEvent::NodeAvailable { .. })
}));
assert!(
!harness
.events()
.iter()
.any(|event| { matches!(event, boot::LifecycleEvent::NodeAvailable { .. }) })
);
// The orchestrator-facing eligibility check must agree with the
// lifecycle transcript.
@ -193,7 +204,10 @@ fn boot_never_self_assigns_run_topology() {
boot::BootCommand::AdvertiseLifecycle { .. }
| boot::BootCommand::JoinMembership { .. }
| boot::BootCommand::OpenProvisioningInbox { .. } => {}
boot::BootCommand::AssignStage { .. }
boot::BootCommand::LoadWeights { .. }
| boot::BootCommand::ConfigureRole { .. }
| boot::BootCommand::EstablishEdge { .. }
| boot::BootCommand::AssignStage { .. }
| boot::BootCommand::AssignLayerRange { .. }
| boot::BootCommand::AssignEdge { .. }
| boot::BootCommand::AssignObjectSpec { .. } => {

View file

@ -54,7 +54,10 @@ fn fault_trace() -> Vec<obs::Event> {
obs::FaultReason::WorkerCrashed,
obs::Component::StageController,
)
.run_faulted(obs::FaultReason::WorkerCrashed, obs::Component::StageController)
.run_faulted(
obs::FaultReason::WorkerCrashed,
obs::Component::StageController,
)
.stop_run_sent(obs::StageIndex(0))
.stage_stopped(obs::StageIndex(0))
.run_torn_down()
@ -94,13 +97,19 @@ fn assert_required_identity(event: &obs::Event) {
sequence,
..
} => {
assert!([obs::ObjectId(9000), obs::ObjectId(9001), obs::ObjectId(9002)].contains(object_id));
assert!(
[
obs::ObjectId(9000),
obs::ObjectId(9001),
obs::ObjectId(9002)
]
.contains(object_id)
);
assert_eq!(*sequence, obs::Sequence(0));
}
obs::Event::StepScoped { step_id, .. } => assert_eq!(*step_id, obs::StepId(77)),
obs::Event::WorkerScoped {
worker_generation,
..
worker_generation, ..
} => assert_eq!(*worker_generation, obs::WorkerGeneration(1)),
}
}
@ -166,7 +175,10 @@ fn lifecycle_events_cover_successful_run_milestones() {
obs::EventKind::RunTornDown,
];
for kind in required {
assert!(observed.contains(&kind), "missing lifecycle event: {kind:?}");
assert!(
observed.contains(&kind),
"missing lifecycle event: {kind:?}"
);
}
}
@ -242,7 +254,10 @@ fn event_ordering_reflects_component_contracts_and_one_terminal_outcome() {
let terminal_count = events
.iter()
.filter(|event| {
matches!(event.kind(), obs::EventKind::RunCompleted | obs::EventKind::RunFaulted)
matches!(
event.kind(),
obs::EventKind::RunCompleted | obs::EventKind::RunFaulted
)
})
.count();
assert_eq!(terminal_count, 1);
@ -253,8 +268,10 @@ fn event_ordering_reflects_component_contracts_and_one_terminal_outcome() {
#[test]
fn event_contract_survives_transport_storage_and_batching_policy() {
// Build the same logical events under two batching policies.
let unbatched = obs::EventSubscriberHarness::collect(successful_run_trace(), obs::Batching::None);
let batched = obs::EventSubscriberHarness::collect(successful_run_trace(), obs::Batching::Fixed(8));
let unbatched =
obs::EventSubscriberHarness::collect(successful_run_trace(), obs::Batching::None);
let batched =
obs::EventSubscriberHarness::collect(successful_run_trace(), obs::Batching::Fixed(8));
// Flattened public event facts must match as an ordered stream.
let unbatched_kinds = unbatched

View file

@ -77,9 +77,12 @@ fn planning_and_provisioning_start_only_after_pool_ready() {
harness.observe(fsm::RunEvent::PlanAvailable(plan.clone()));
// Without PoolReady, provisioning must not begin.
assert!(!harness.commands().iter().any(|command| {
matches!(command, fsm::RunCommand::ProvisionStage { .. })
}));
assert!(
!harness
.commands()
.iter()
.any(|command| { matches!(command, fsm::RunCommand::ProvisionStage { .. }) })
);
// Once PoolReady is observed, the committed plan may be provisioned.
harness.observe(fsm::RunEvent::PoolReady {
@ -103,7 +106,10 @@ fn planning_and_provisioning_start_only_after_pool_ready() {
assert_eq!(provisioned, expected);
// Provisioning must not mention nodes outside the committed plan.
let plan_nodes = plan.stage_nodes().into_iter().collect::<std::collections::BTreeSet<_>>();
let plan_nodes = plan
.stage_nodes()
.into_iter()
.collect::<std::collections::BTreeSet<_>>();
for command in harness.commands() {
if let fsm::RunCommand::ProvisionStage { provision } = command {
assert!(plan_nodes.contains(&provision.node_id));
@ -143,9 +149,12 @@ fn readiness_barrier_controls_prompt_injection() {
});
harness.observe(fsm::RunEvent::TokenInEndpointReady);
harness.observe(fsm::RunEvent::TokenOutEndpointReady);
assert!(!harness.commands().iter().any(|command| {
matches!(command, fsm::RunCommand::InjectPrompt { .. })
}));
assert!(
!harness
.commands()
.iter()
.any(|command| { matches!(command, fsm::RunCommand::InjectPrompt { .. }) })
);
// Complete the remaining stage readiness facts.
for event in stage_ready_events(&plan).into_iter().skip(1) {
@ -198,9 +207,12 @@ fn execution_injects_next_sequence_only_after_consuming_previous_token() {
}
// No separate broadcast start command may exist alongside prompt injection.
assert!(!harness.commands().iter().any(|command| {
matches!(command, fsm::RunCommand::BroadcastStart { .. })
}));
assert!(
!harness
.commands()
.iter()
.any(|command| { matches!(command, fsm::RunCommand::BroadcastStart { .. }) })
);
// Sequence 0 must be injected first.
assert_eq!(harness.injected_sequences(), vec![0]);
@ -341,9 +353,12 @@ fn terminal_outcome_is_single_and_requires_teardown() {
.map(|stage| stage.stage_index)
.collect::<std::collections::BTreeSet<_>>();
assert_eq!(stopped_stages, expected_stages);
assert!(harness.commands().iter().any(|command| {
matches!(command, fsm::RunCommand::TearDownTokenEndpoints { .. })
}));
assert!(
harness
.commands()
.iter()
.any(|command| { matches!(command, fsm::RunCommand::TearDownTokenEndpoints { .. }) })
);
}
// This proves run_torn_down is emitted exactly once and only after teardown
@ -369,9 +384,12 @@ fn run_torn_down_is_emitted_once_after_teardown_terminal_state() {
run_id: fsm::RunId(7),
stage_index: 0,
});
assert!(!harness.events().iter().any(|event| {
matches!(event, fsm::LifecycleEvent::RunTornDown { .. })
}));
assert!(
!harness
.events()
.iter()
.any(|event| { matches!(event, fsm::LifecycleEvent::RunTornDown { .. }) })
);
// Finish teardown through remaining stopped events and local endpoint stop.
harness.observe(fsm::RunEvent::StageStopped {

View file

@ -18,8 +18,16 @@ fn token_plan() -> token::TokenEndpointPlan {
token::TokenEndpointPlan {
run_id: token::RunId(7),
orchestrator_node_id: token::NodeId(99),
token_in_edge: token::EdgePlan::token_in(token::EdgeId(7000), token::NodeId(99), token::NodeId(10)),
token_out_edge: token::EdgePlan::token_out(token::EdgeId(7003), token::NodeId(12), token::NodeId(99)),
token_in_edge: token::EdgePlan::token_in(
token::EdgeId(7000),
token::NodeId(99),
token::NodeId(10),
),
token_out_edge: token::EdgePlan::token_out(
token::EdgeId(7003),
token::NodeId(12),
token::NodeId(99),
),
token_spec: token::ObjectSpec::test_tokens(),
max_tokens: 4,
}
@ -95,9 +103,12 @@ fn prompt_injection_is_barrier_gated_sequence_zero_token_object() {
harness.request_prompt_injection(vec![101, 102, 103]);
// No prompt object may be written before the barrier.
assert!(!harness.commands().iter().any(|command| {
matches!(command, token::EndpointCommand::WriteTokenObject { .. })
}));
assert!(
!harness
.commands()
.iter()
.any(|command| { matches!(command, token::EndpointCommand::WriteTokenObject { .. }) })
);
// Passing the global barrier permits the first prompt write.
pass_global_barrier(&mut harness);
@ -117,9 +128,12 @@ fn prompt_injection_is_barrier_gated_sequence_zero_token_object() {
// Prompt injection is the start signal; there must not be a separate
// broadcast start command.
assert!(!harness.commands().iter().any(|command| {
matches!(command, token::EndpointCommand::BroadcastStart { .. })
}));
assert!(
!harness
.commands()
.iter()
.any(|command| { matches!(command, token::EndpointCommand::BroadcastStart { .. }) })
);
}
// This proves token-out is consumed in sequence order and sequence k + 1 is
@ -242,7 +256,10 @@ fn token_endpoint_failures_fault_the_run_or_teardown() {
malformed.observe(token::EndpointEvent::TeardownFailed {
edge_id: token::EdgeId(7003),
});
assert!(malformed.events().iter().any(|event| {
matches!(event, token::EndpointLifecycleEvent::TeardownFailed { .. })
}));
assert!(
malformed
.events()
.iter()
.any(|event| { matches!(event, token::EndpointLifecycleEvent::TeardownFailed { .. }) })
);
}

View file

@ -91,7 +91,10 @@ fn inventory_entries_are_known_before_planning_and_not_negotiated_by_nodes() {
.collect::<std::collections::BTreeSet<_>>();
assert_eq!(after_nodes, before_nodes);
assert!(!harness.commands().iter().any(|command| {
matches!(command, inventory::InventoryCommand::AcceptPlacementNegotiation { .. })
matches!(
command,
inventory::InventoryCommand::AcceptPlacementNegotiation { .. }
)
}));
}
@ -224,7 +227,8 @@ fn planning_preserves_candidate_pool_and_emits_no_hidden_nodes() {
#[test]
fn provisioned_stages_do_not_reinterpret_inventory() {
// Commit a plan and provision its stages.
let plan = inventory::plan_from_inventory(planning_request()).expect("valid inventory must plan");
let plan =
inventory::plan_from_inventory(planning_request()).expect("valid inventory must plan");
let mut harness = inventory::InventoryHarness::new(inventory_entries());
harness.commit_plan(plan.clone());
harness.provision_stages();
@ -238,7 +242,10 @@ fn provisioned_stages_do_not_reinterpret_inventory() {
// No provisioned stage may be asked to reinterpret its assignment.
assert!(!harness.commands().iter().any(|command| {
matches!(command, inventory::InventoryCommand::RewriteStagePlacement { .. })
matches!(
command,
inventory::InventoryCommand::RewriteStagePlacement { .. }
)
}));
// The committed plan remains the only graph-visible placement fact.

View file

@ -20,7 +20,7 @@ type EdgeKind = plan::EdgeKind;
type EdgePlan = plan::EdgePlan;
type InboundEdgeProvision = plan::InboundEdgeProvision;
type ModelFacts = plan::ModelFacts;
type NodeId = plan::NodeId;
use plan::NodeId;
type ObjectKind = plan::ObjectKind;
type OutboundEdgeProvision = plan::OutboundEdgeProvision;
type PlacementInput = plan::PlacementInput;
@ -283,7 +283,10 @@ fn edge_graph_is_exactly_the_linear_pipeline() {
.filter(|edge| edge.kind == EdgeKind::TokenOut)
.collect::<Vec<_>>();
assert_eq!(token_out.len(), 1);
assert_eq!(edge_stage_index(&token_out[0].producer), Some(stage_count - 1));
assert_eq!(
edge_stage_index(&token_out[0].producer),
Some(stage_count - 1)
);
assert!(matches!(
token_out[0].consumer,
EdgeEndpoint::Orchestrator { node_id } if node_id == node(99)

View file

@ -37,8 +37,7 @@ fn drain_committed(helper: &mut ring::RingHelperHarness) -> Vec<u8> {
// depending on scheduler internals.
fn assert_wake_is_payload_free(wake: &ring::WakeHint) {
match wake {
ring::WakeHint::RingReadable { ring_id }
| ring::WakeHint::RingWritable { ring_id } => {
ring::WakeHint::RingReadable { ring_id } | ring::WakeHint::RingWritable { ring_id } => {
assert_eq!(*ring_id, ring::RingId(7000));
}
}
@ -85,11 +84,16 @@ fn cursors_are_monotonic_and_wrap_by_modulo_capacity() {
assert_eq!(helper.cursor_snapshot().consume, 6);
// Wrap the physical index while logical cursors keep increasing.
let wrapped = helper.producer_reserve(5).expect("space must exist after consume");
let wrapped = helper
.producer_reserve(5)
.expect("space must exist after consume");
helper.producer_write(&wrapped, b"ghijk");
helper.producer_commit(wrapped);
assert_eq!(helper.cursor_snapshot().commit, 11);
assert_eq!(helper.cursor_snapshot().commit % helper.identity().capacity, 3);
assert_eq!(
helper.cursor_snapshot().commit % helper.identity().capacity,
3
);
assert_eq!(drain_committed(&mut helper), b"ghijk");
assert_eq!(helper.cursor_snapshot().consume, 11);
}
@ -101,7 +105,9 @@ fn cursors_are_monotonic_and_wrap_by_modulo_capacity() {
fn producer_respects_free_space_and_publishes_after_writing() {
// Reserve the full ring and publish it.
let mut helper = new_ring();
let reservation = helper.producer_reserve(8).expect("full ring reservation fits");
let reservation = helper
.producer_reserve(8)
.expect("full ring reservation fits");
helper.producer_write(&reservation, b"12345678");
helper.producer_commit(reservation);
@ -170,12 +176,22 @@ fn wake_hints_are_edge_hints_without_hiding_transitions() {
// The readable transition must be discoverable even if duplicate wakes are
// coalesced.
helper.coalesce_duplicate_wakes();
assert!(helper.scheduler_state().readable_rings.contains(&ring::RingId(7000)));
assert!(
helper
.scheduler_state()
.readable_rings
.contains(&ring::RingId(7000))
);
// Fill then release space to cause full-to-writable.
let _ = drain_committed(&mut helper);
helper.coalesce_duplicate_wakes();
assert!(helper.scheduler_state().writable_rings.contains(&ring::RingId(7000)));
assert!(
helper
.scheduler_state()
.writable_rings
.contains(&ring::RingId(7000))
);
// Every wake remains a payload-free hint.
for wake in helper.wake_hints() {

View file

@ -101,9 +101,12 @@ fn provisioning_validates_authority_and_assigned_shape_before_setup() {
// The controller must not emit any command that replaces the provisioned
// edge ids with a locally chosen edge.
assert!(!harness.commands().iter().any(|command| {
matches!(command, stage::StageCommand::RewireEdge { .. })
}));
assert!(
!harness
.commands()
.iter()
.any(|command| { matches!(command, stage::StageCommand::RewireEdge { .. }) })
);
// An unauthorized provision attempt must fault before setup can begin.
let mut unauthorized = new_controller();
@ -120,9 +123,12 @@ fn provisioning_validates_authority_and_assigned_shape_before_setup() {
}
)
}));
assert!(!unauthorized.commands().iter().any(|command| {
matches!(command, stage::StageCommand::ConfigureWorkerRole { .. })
}));
assert!(
!unauthorized
.commands()
.iter()
.any(|command| { matches!(command, stage::StageCommand::ConfigureWorkerRole { .. }) })
);
}
// This proves StageReady is emitted only after worker readiness, weight
@ -143,9 +149,12 @@ fn stage_ready_waits_for_worker_weights_and_both_edges() {
let final_event = events.pop().expect("fixture has final setup event");
for event in events {
harness.observe(event);
assert!(!harness.events().iter().any(|event| {
matches!(event, stage::StageLifecycleEvent::StageReady { .. })
}));
assert!(
!harness
.events()
.iter()
.any(|event| { matches!(event, stage::StageLifecycleEvent::StageReady { .. }) })
);
}
// The final prerequisite crosses the barrier.
@ -220,11 +229,7 @@ fn accepted_inbound_object_creates_one_same_sequence_execute_step() {
#[test]
fn duplicate_skipped_and_out_of_order_sequences_fault() {
// Each invalid trace starts from a freshly readied stage.
let invalid_traces = vec![
vec![0, 0],
vec![0, 2],
vec![0, 1, 0],
];
let invalid_traces = vec![vec![0, 0], vec![0, 2], vec![0, 1, 0]];
for trace in invalid_traces {
// Accept the first object and complete its step when needed so the next
@ -272,7 +277,10 @@ fn step_completed_releases_input_and_admits_next_object() {
// Before worker completion, no compute-complete lifecycle event is allowed.
assert!(!harness.events().iter().any(|event| {
matches!(event, stage::StageLifecycleEvent::StepAccepted { sequence: 1, .. })
matches!(
event,
stage::StageLifecycleEvent::StepAccepted { sequence: 1, .. }
)
}));
// Worker StepCompleted is the public completion signal.
@ -348,13 +356,21 @@ fn stage_fault_rejects_new_work_until_stopped() {
harness.observe(stage::StageEvent::StopRun {
run_id: stage::RunId(7),
});
assert!(harness.commands().iter().any(|command| {
matches!(command, stage::StageCommand::StopLocalEdges { .. })
}));
assert!(harness.commands().iter().any(|command| {
matches!(command, stage::StageCommand::ReleaseRunDeviceObjects { .. })
}));
assert!(harness.events().iter().any(|event| {
matches!(event, stage::StageLifecycleEvent::StageStopped { .. })
}));
assert!(
harness
.commands()
.iter()
.any(|command| { matches!(command, stage::StageCommand::StopLocalEdges { .. }) })
);
assert!(
harness.commands().iter().any(|command| {
matches!(command, stage::StageCommand::ReleaseRunDeviceObjects { .. })
})
);
assert!(
harness
.events()
.iter()
.any(|event| { matches!(event, stage::StageLifecycleEvent::StageStopped { .. }) })
);
}

View file

@ -0,0 +1,305 @@
#!/usr/bin/env python3
"""Tinygrad-backed device bridge probe for Rust MVP bridge tests.
Line-delimited JSON control only. Payload bytes live in the arena file whose
path Rust passes during initialize.
"""
from __future__ import annotations
import json
import os
import struct
import sys
from dataclasses import dataclass
from typing import Any
try:
from tinygrad import Tensor # type: ignore
except Exception as exc: # pragma: no cover - exercised from Rust process tests
print(
json.dumps(
{
"type": "worker_fatal",
"reason": "tinygrad_unavailable",
"message": str(exc),
}
),
flush=True,
)
raise SystemExit(2)
@dataclass
class DeviceObject:
dtype: str
shape: str
extent: int
values: list[int]
tensor: Any
@dataclass
class PendingCopy:
kind: str
copy_id: int
handle_id: int
host_offset: int
device_offset: int
length: int
arena_path: str | None = None
arena_bytes = 0
generation = 0
objects: dict[int, DeviceObject] = {}
pending_copies: dict[int, PendingCopy] = {}
fail_next: str | None = None
def emit(obj: dict[str, Any]) -> None:
print(json.dumps(obj, separators=(",", ":")), flush=True)
def backend_error(reason: str) -> None:
emit({"type": "backend_error", "reason": reason})
def fatal(reason: str, message: str) -> None:
emit({"type": "worker_fatal", "reason": reason, "message": message})
def require_arena() -> str:
if arena_path is None:
raise RuntimeError("arena is not initialized")
return arena_path
def consume_failure(reason: str) -> bool:
global fail_next
if fail_next == reason:
fail_next = None
backend_error(reason)
return True
return False
def check_u32_range(offset: int, length: int) -> None:
if offset < 0 or length < 0 or offset % 4 != 0 or length % 4 != 0:
raise ValueError("u32 ranges must be non-negative and 4-byte aligned")
def tensor_from_values(values: list[int]) -> Any:
try:
return Tensor(values, dtype="uint32").realize()
except Exception:
return Tensor(values, dtype="int32").realize()
def allocate_tensor(dtype: str, extent: int) -> tuple[list[int], Any]:
if dtype == "u32":
if extent % 4 != 0:
raise ValueError("u32 extent must be 4-byte aligned")
values = [0] * (extent // 4)
return values, tensor_from_values(values)
if dtype == "f16":
values = [0] * (extent // 2)
return values, Tensor(values, dtype="float16").realize()
raise ValueError(f"unsupported dtype {dtype!r}")
def update_tensor(obj: DeviceObject) -> None:
if obj.dtype == "u32":
obj.tensor = tensor_from_values(obj.values)
elif obj.dtype == "f16":
obj.tensor = Tensor(obj.values, dtype="float16").realize()
else:
raise ValueError(f"unsupported dtype {obj.dtype!r}")
def materialized_values(obj: DeviceObject) -> list[int]:
return [int(v) for v in obj.tensor.tolist()]
def read_arena(offset: int, length: int) -> bytes:
path = require_arena()
with open(path, "rb", buffering=0) as f:
f.seek(offset)
data = f.read(length)
if len(data) != length:
raise EOFError("short arena read")
return data
def write_arena(offset: int, data: bytes) -> None:
path = require_arena()
with open(path, "r+b", buffering=0) as f:
f.seek(offset)
f.write(data)
f.flush()
def perform_host_to_device(copy: PendingCopy) -> None:
obj = objects[copy.handle_id]
if obj.dtype != "u32":
raise ValueError("payload copy is implemented for u32 test objects only")
check_u32_range(copy.device_offset, copy.length)
payload = read_arena(copy.host_offset, copy.length)
words = list(struct.unpack("<" + "I" * (copy.length // 4), payload))
start = copy.device_offset // 4
end = start + len(words)
if end > len(obj.values):
raise ValueError("device range out of bounds")
obj.values[start:end] = words
update_tensor(obj)
def perform_device_to_host(copy: PendingCopy) -> None:
obj = objects[copy.handle_id]
if obj.dtype != "u32":
raise ValueError("payload copy is implemented for u32 test objects only")
check_u32_range(copy.device_offset, copy.length)
values = materialized_values(obj)
start = copy.device_offset // 4
end = start + (copy.length // 4)
if end > len(values):
raise ValueError("device range out of bounds")
payload = struct.pack("<" + "I" * (end - start), *values[start:end])
write_arena(copy.host_offset, payload)
def perform_copy(copy: PendingCopy) -> None:
if copy.kind == "host_to_device":
perform_host_to_device(copy)
elif copy.kind == "device_to_host":
perform_device_to_host(copy)
else:
raise ValueError(f"unknown copy kind {copy.kind!r}")
def handle(req: dict[str, Any]) -> bool:
global arena_path, arena_bytes, generation, fail_next
typ = req.get("type")
if typ == "initialize":
arena_path = str(req["arena_path"])
arena_bytes = int(req["arena_bytes"])
generation = int(req["generation"])
with open(arena_path, "r+b", buffering=0) as f:
f.truncate(arena_bytes)
emit({"type": "worker_ready", "generation": generation})
return True
if typ == "alloc":
if consume_failure("allocation_failed"):
return True
handle_id = int(req["handle_id"])
dtype = str(req["dtype"])
shape = str(req["shape"])
extent = int(req["extent"])
values, tensor = allocate_tensor(dtype, extent)
objects[handle_id] = DeviceObject(dtype=dtype, shape=shape, extent=extent, values=values, tensor=tensor)
emit({"type": "allocated", "handle_id": handle_id})
return True
if typ in ("host_to_device", "device_to_host"):
if consume_failure("copy_failed"):
return True
copy = PendingCopy(
kind=typ,
copy_id=int(req["copy_id"]),
handle_id=int(req["handle_id"]),
host_offset=int(req["host_offset"]),
device_offset=int(req["device_offset"]),
length=int(req["len"]),
)
if copy.handle_id not in objects:
backend_error("copy_failed")
return True
if bool(req.get("defer", False)):
pending_copies[copy.copy_id] = copy
emit({"type": "copy_started", "copy_id": copy.copy_id})
else:
perform_copy(copy)
emit({"type": "copy_completed", "copy_id": copy.copy_id})
return True
if typ == "complete_copy":
copy_id = int(req["copy_id"])
copy = pending_copies.pop(copy_id, None)
if copy is None:
backend_error("copy_failed")
return True
perform_copy(copy)
emit({"type": "copy_completed", "copy_id": copy_id})
return True
if typ == "wrap_for_tinygrad":
if consume_failure("invalid_view"):
return True
handle_id = int(req["handle_id"])
obj = objects.get(handle_id)
if obj is None:
backend_error("invalid_view")
return True
dtype = str(req["dtype"])
shape = str(req["shape"])
if obj.dtype != dtype or obj.shape != shape:
backend_error("invalid_view")
return True
# Force materialization at view time so success depends on live tensor state.
_ = obj.tensor.tolist()
emit({"type": "view", "handle_id": handle_id, "dtype": dtype, "shape": shape})
return True
if typ == "free":
handle_id = int(req["handle_id"])
if handle_id not in objects:
backend_error("invalid_view")
return True
del objects[handle_id]
emit({"type": "freed", "handle_id": handle_id})
return True
if typ == "restart":
generation = int(req["generation"])
objects.clear()
pending_copies.clear()
fail_next = None
emit({"type": "worker_ready", "generation": generation})
return True
if typ == "fail_next":
fail_next = str(req["failure"])
emit({"type": "ok"})
return True
if typ == "shutdown":
emit({"type": "worker_stopped"})
return False
fatal("protocol_error", f"unknown command type {typ!r}")
return False
def main() -> int:
for line in sys.stdin:
line = line.strip()
if not line:
continue
try:
req = json.loads(line)
if not isinstance(req, dict):
raise ValueError("request must be an object")
if not handle(req):
return 0
except SystemExit:
raise
except Exception as exc:
fatal("backend_exception", str(exc))
return 2
return 0
if __name__ == "__main__":
raise SystemExit(main())

View file

@ -82,8 +82,16 @@ fn edge_actor_messages_are_lifecycle_identity_and_handle_only() {
}
// Both actors remain tied to exactly one edge id.
assert!(tx.messages().iter().all(|message| message.edge_id() == edge_id()));
assert!(rx.messages().iter().all(|message| message.edge_id() == edge_id()));
assert!(
tx.messages()
.iter()
.all(|message| message.edge_id() == edge_id())
);
assert!(
rx.messages()
.iter()
.all(|message| message.edge_id() == edge_id())
);
}
// This proves Tx starts in provisioning, becomes ready only after EdgeReady,
@ -98,9 +106,11 @@ fn tx_lifecycle_gates_production_and_faults_on_stream_or_object_failure() {
object_id: edge_actor::ObjectId(9000),
sequence: 0,
});
assert!(!tx.messages().iter().any(|message| {
matches!(message, edge_actor::ActorMessage::ObjectIdentity { .. })
}));
assert!(
!tx.messages()
.iter()
.any(|message| { matches!(message, edge_actor::ActorMessage::ObjectIdentity { .. }) })
);
// EdgeReady admits production.
tx.observe(edge_actor::TxEvent::EdgeReady { edge_id: edge_id() });
@ -154,9 +164,11 @@ fn rx_lifecycle_gates_loaded_objects_and_faults_on_stream_or_object_failure() {
sequence: 0,
handle: edge_actor::OpaqueHandle::new(42),
});
assert!(!rx.messages().iter().any(|message| {
matches!(message, edge_actor::ActorMessage::OpaqueHandle { .. })
}));
assert!(
!rx.messages()
.iter()
.any(|message| { matches!(message, edge_actor::ActorMessage::OpaqueHandle { .. }) })
);
// EdgeReady admits ObjectLoaded exposure.
rx.observe(edge_actor::RxEvent::EdgeReady { edge_id: edge_id() });
@ -223,12 +235,16 @@ fn stop_and_mismatched_edge_events_do_not_create_run_work() {
sequence: 0,
handle: edge_actor::OpaqueHandle::new(42),
});
assert!(!tx.messages().iter().any(|message| {
matches!(message, edge_actor::ActorMessage::ObjectIdentity { .. })
}));
assert!(!rx.messages().iter().any(|message| {
matches!(message, edge_actor::ActorMessage::OpaqueHandle { .. })
}));
assert!(
!tx.messages()
.iter()
.any(|message| { matches!(message, edge_actor::ActorMessage::ObjectIdentity { .. }) })
);
assert!(
!rx.messages()
.iter()
.any(|message| { matches!(message, edge_actor::ActorMessage::OpaqueHandle { .. }) })
);
// A mismatched edge id must reject or fault, not create work on this actor.
let mut mismatched = new_tx();

View file

@ -169,9 +169,11 @@ fn weights_ready_requires_all_weight_facts_and_precedes_stage_ready() {
let final_event = events.pop().expect("fixture has final weight event");
for event in events {
harness.observe(event);
assert!(!harness.events().iter().any(|event| {
matches!(event, weights::WeightLifecycleEvent::WeightsReady { .. })
}));
assert!(
!harness.events().iter().any(|event| {
matches!(event, weights::WeightLifecycleEvent::WeightsReady { .. })
})
);
}
// The final weight prerequisite emits WeightsReady.

View file

@ -0,0 +1,324 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct EdgeId(pub u64);
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct PortId(pub String);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct ObjectId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct OpaqueHandle(pub u64);
impl OpaqueHandle {
pub fn new(id: u64) -> Self {
Self(id)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct TxConfig {
pub edge_id: EdgeId,
pub role_port: PortId,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RxConfig {
pub edge_id: EdgeId,
pub role_port: PortId,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ActorFaultReason {
MismatchedEdgeId,
StreamFault,
ObjectFault,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ActorMessage {
Lifecycle {
edge_id: EdgeId,
ready: bool,
},
ObjectIdentity {
edge_id: EdgeId,
object_id: ObjectId,
sequence: u64,
},
OpaqueHandle {
edge_id: EdgeId,
object_id: ObjectId,
sequence: u64,
handle: OpaqueHandle,
},
CoarseFault {
edge_id: EdgeId,
reason: ActorFaultReason,
},
PayloadBytes {
edge_id: EdgeId,
bytes: Vec<u8>,
},
HostPointer {
edge_id: EdgeId,
address: usize,
},
ByteRange {
edge_id: EdgeId,
start: u64,
len: u64,
},
CreditCount {
edge_id: EdgeId,
credits: u64,
},
}
impl ActorMessage {
pub fn edge_id(&self) -> EdgeId {
match self {
ActorMessage::Lifecycle { edge_id, .. }
| ActorMessage::ObjectIdentity { edge_id, .. }
| ActorMessage::OpaqueHandle { edge_id, .. }
| ActorMessage::CoarseFault { edge_id, .. }
| ActorMessage::PayloadBytes { edge_id, .. }
| ActorMessage::HostPointer { edge_id, .. }
| ActorMessage::ByteRange { edge_id, .. }
| ActorMessage::CreditCount { edge_id, .. } => *edge_id,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum TxEvent {
EdgeReady {
edge_id: EdgeId,
},
ObjectProduced {
edge_id: EdgeId,
object_id: ObjectId,
sequence: u64,
},
StreamFault {
edge_id: EdgeId,
},
ObjectFailed {
edge_id: EdgeId,
},
StopEdge {
edge_id: EdgeId,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum RxEvent {
EdgeReady {
edge_id: EdgeId,
},
ObjectLoaded {
edge_id: EdgeId,
object_id: ObjectId,
sequence: u64,
handle: OpaqueHandle,
},
StreamFault {
edge_id: EdgeId,
},
ObjectFailed {
edge_id: EdgeId,
},
StopEdge {
edge_id: EdgeId,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum ActorState {
Provisioning,
Ready,
Faulted,
Stopped,
}
#[cfg(test)]
pub struct TxActorHarness {
config: TxConfig,
state: ActorState,
messages: Vec<ActorMessage>,
}
#[cfg(test)]
impl TxActorHarness {
pub fn new(config: TxConfig) -> Self {
Self {
config,
state: ActorState::Provisioning,
messages: Vec::new(),
}
}
pub fn observe(&mut self, event: TxEvent) {
match event {
TxEvent::EdgeReady { edge_id } => {
if !self.check_edge(edge_id) {
return;
}
if self.state == ActorState::Provisioning {
self.state = ActorState::Ready;
self.messages.push(ActorMessage::Lifecycle {
edge_id,
ready: true,
});
}
}
TxEvent::ObjectProduced {
edge_id,
object_id,
sequence,
} => {
if self.state == ActorState::Ready && self.check_edge(edge_id) {
self.messages.push(ActorMessage::ObjectIdentity {
edge_id,
object_id,
sequence,
});
}
}
TxEvent::StreamFault { edge_id } => {
self.fault_if_edge(edge_id, ActorFaultReason::StreamFault)
}
TxEvent::ObjectFailed { edge_id } => {
self.fault_if_edge(edge_id, ActorFaultReason::ObjectFault)
}
TxEvent::StopEdge { edge_id } => {
if edge_id == self.config.edge_id {
self.state = ActorState::Stopped;
} else {
self.mismatched();
}
}
}
}
pub fn messages(&self) -> &[ActorMessage] {
&self.messages
}
fn check_edge(&mut self, edge_id: EdgeId) -> bool {
if edge_id == self.config.edge_id {
true
} else {
self.mismatched();
false
}
}
fn fault_if_edge(&mut self, edge_id: EdgeId, reason: ActorFaultReason) {
if self.check_edge(edge_id) && self.state != ActorState::Stopped {
self.state = ActorState::Faulted;
self.messages
.push(ActorMessage::CoarseFault { edge_id, reason });
}
}
fn mismatched(&mut self) {
self.state = ActorState::Faulted;
self.messages.push(ActorMessage::CoarseFault {
edge_id: self.config.edge_id,
reason: ActorFaultReason::MismatchedEdgeId,
});
}
}
#[cfg(test)]
pub struct RxActorHarness {
config: RxConfig,
state: ActorState,
messages: Vec<ActorMessage>,
}
#[cfg(test)]
impl RxActorHarness {
pub fn new(config: RxConfig) -> Self {
Self {
config,
state: ActorState::Provisioning,
messages: Vec::new(),
}
}
pub fn observe(&mut self, event: RxEvent) {
match event {
RxEvent::EdgeReady { edge_id } => {
if !self.check_edge(edge_id) {
return;
}
if self.state == ActorState::Provisioning {
self.state = ActorState::Ready;
self.messages.push(ActorMessage::Lifecycle {
edge_id,
ready: true,
});
}
}
RxEvent::ObjectLoaded {
edge_id,
object_id,
sequence,
handle,
} => {
if self.state == ActorState::Ready && self.check_edge(edge_id) {
self.messages.push(ActorMessage::OpaqueHandle {
edge_id,
object_id,
sequence,
handle,
});
}
}
RxEvent::StreamFault { edge_id } => {
self.fault_if_edge(edge_id, ActorFaultReason::StreamFault)
}
RxEvent::ObjectFailed { edge_id } => {
self.fault_if_edge(edge_id, ActorFaultReason::ObjectFault)
}
RxEvent::StopEdge { edge_id } => {
if edge_id == self.config.edge_id {
self.state = ActorState::Stopped;
} else {
self.mismatched();
}
}
}
}
pub fn messages(&self) -> &[ActorMessage] {
&self.messages
}
fn check_edge(&mut self, edge_id: EdgeId) -> bool {
if edge_id == self.config.edge_id {
true
} else {
self.mismatched();
false
}
}
fn fault_if_edge(&mut self, edge_id: EdgeId, reason: ActorFaultReason) {
if self.check_edge(edge_id) && self.state != ActorState::Stopped {
self.state = ActorState::Faulted;
self.messages
.push(ActorMessage::CoarseFault { edge_id, reason });
}
}
fn mismatched(&mut self) {
self.state = ActorState::Faulted;
self.messages.push(ActorMessage::CoarseFault {
edge_id: self.config.edge_id,
reason: ActorFaultReason::MismatchedEdgeId,
});
}
}

View file

@ -0,0 +1,203 @@
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct RunId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct NodeId(pub u64);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct LayerRange {
pub start: u32,
pub end_exclusive: u32,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WeightSource {
WholeGguf { uri: String },
ShardSet { uris: Vec<String> },
CachedArtifact { cache_key: String },
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct WeightAssignment {
pub run_id: RunId,
pub stage_index: u32,
pub plan_layer_range: LayerRange,
pub assigned_layer_range: LayerRange,
pub source: WeightSource,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ArtifactBytes {
Local,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WeightEvent {
Provisioned(WeightAssignment),
ArtifactAvailable { bytes: ArtifactBytes },
LayerRangeValidated,
WorkerRangeBound,
OtherStagePrerequisitesReady,
DownloadFailed,
ParseFailed,
DeviceAllocationFailed,
BindingFailed,
InvalidLayerRange,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum StageFaultReason {
WeightDownloadFailed,
WeightParseFailed,
DeviceAllocationFailed,
WeightBindingFailed,
InvalidLayerRange,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WeightLifecycleEvent {
WeightsReady {
run_id: RunId,
stage_index: u32,
},
StageReady {
run_id: RunId,
stage_index: u32,
},
StageFault {
run_id: RunId,
stage_index: u32,
reason: StageFaultReason,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum WeightCommand {
LoadOrBindRange {
source: WeightSource,
range: LayerRange,
},
AdvertiseLoadedLayerRange {
range: LayerRange,
},
}
#[cfg(test)]
pub struct WeightLifecycleHarness {
_node_id: NodeId,
assignment: Option<WeightAssignment>,
artifact: bool,
validated: bool,
bound: bool,
other_stage_prereqs: bool,
weights_ready: bool,
faulted: bool,
commands: Vec<WeightCommand>,
events: Vec<WeightLifecycleEvent>,
}
#[cfg(test)]
impl WeightLifecycleHarness {
pub fn new(node_id: NodeId) -> Self {
Self {
_node_id: node_id,
assignment: None,
artifact: false,
validated: false,
bound: false,
other_stage_prereqs: false,
weights_ready: false,
faulted: false,
commands: Vec::new(),
events: Vec::new(),
}
}
pub fn observe(&mut self, event: WeightEvent) {
match event {
WeightEvent::Provisioned(assignment) => {
self.commands.push(WeightCommand::LoadOrBindRange {
source: assignment.source.clone(),
range: assignment.assigned_layer_range,
});
if assignment.assigned_layer_range != assignment.plan_layer_range
|| assignment.assigned_layer_range.start
>= assignment.assigned_layer_range.end_exclusive
{
self.assignment = Some(assignment);
self.fault(StageFaultReason::InvalidLayerRange);
} else {
self.assignment = Some(assignment);
}
}
WeightEvent::ArtifactAvailable { .. } => self.artifact = true,
WeightEvent::LayerRangeValidated => self.validated = true,
WeightEvent::WorkerRangeBound => self.bound = true,
WeightEvent::OtherStagePrerequisitesReady => self.other_stage_prereqs = true,
WeightEvent::DownloadFailed => self.fault(StageFaultReason::WeightDownloadFailed),
WeightEvent::ParseFailed => self.fault(StageFaultReason::WeightParseFailed),
WeightEvent::DeviceAllocationFailed => {
self.fault(StageFaultReason::DeviceAllocationFailed)
}
WeightEvent::BindingFailed => self.fault(StageFaultReason::WeightBindingFailed),
WeightEvent::InvalidLayerRange => self.fault(StageFaultReason::InvalidLayerRange),
}
self.maybe_weights_ready();
self.maybe_stage_ready();
}
pub fn commands(&self) -> &[WeightCommand] {
&self.commands
}
pub fn events(&self) -> &[WeightLifecycleEvent] {
&self.events
}
fn maybe_weights_ready(&mut self) {
if self.faulted || self.weights_ready || self.assignment.is_none() {
return;
}
if self.artifact && self.validated && self.bound {
self.weights_ready = true;
let assignment = self.assignment.as_ref().unwrap();
self.events.push(WeightLifecycleEvent::WeightsReady {
run_id: assignment.run_id,
stage_index: assignment.stage_index,
});
}
}
fn maybe_stage_ready(&mut self) {
if self.faulted || !self.weights_ready || !self.other_stage_prereqs {
return;
}
let assignment = self.assignment.as_ref().unwrap();
if !self
.events
.iter()
.any(|event| matches!(event, WeightLifecycleEvent::StageReady { .. }))
{
self.events.push(WeightLifecycleEvent::StageReady {
run_id: assignment.run_id,
stage_index: assignment.stage_index,
});
}
}
fn fault(&mut self, reason: StageFaultReason) {
if self.faulted {
return;
}
self.faulted = true;
let (run_id, stage_index) = self
.assignment
.as_ref()
.map(|assignment| (assignment.run_id, assignment.stage_index))
.unwrap_or((RunId(0), 0));
self.events.push(WeightLifecycleEvent::StageFault {
run_id,
stage_index,
reason,
});
}
}