use cubecl_environment::stream::StreamId;
use cubecl_ir::MemoryDeviceProperties;
use cubecl_server::id::KernelId;
use cubecl_server::logging::ServerLogger;
use cubecl_server::memory_management::{
ErrorGraph, FailureId, MemoryConfiguration, MemoryManagement, MemoryManagementOptions,
};
use cubecl_server::server::{BufferBinding, Handle, ServerError};
use cubecl_server::storage::BytesStorage;
use cubecl_server::stream::{
ExecuteScope, FailureStore, Failures, ScopedOutcome, StreamCapture, StreamFactory,
StreamMemory, StreamPool, WriteScoped,
};
use std::sync::Arc;
fn service() -> cubecl_common::device::ServiceId {
cubecl_common::device::ServiceId::of::<()>(cubecl_common::device::DeviceId::new(0, 0))
}
const MAX_STREAMS: u8 = 4;
const OPS_PER_RUN: usize = 300;
const SEEDS: u64 = 40;
struct Buffer {
_handle: Handle,
binding: BufferBinding,
stale: Vec<bool>,
}
impl Buffer {
fn slice(&self, range: &core::ops::Range<u64>) -> BufferBinding {
let mut binding = self.binding.clone();
binding.offset_start = Some(range.start);
binding.offset_end = Some(binding.size - range.end);
binding
}
fn size(&self) -> u64 {
self.binding.size
}
}
struct TestStream {
memory: MemoryManagement<BytesStorage>,
}
impl core::fmt::Debug for TestStream {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("TestStream").finish()
}
}
impl StreamMemory for TestStream {
fn failure(&self, binding: &BufferBinding) -> Option<FailureId> {
self.memory.failure(&binding.memory, binding.range())
}
fn taint(&mut self, binding: &BufferBinding, failure: FailureId, failures: &mut ErrorGraph) {
self.memory
.taint(&binding.memory, binding.range(), failure, failures)
}
fn written(&mut self, binding: &BufferBinding, failures: &mut ErrorGraph) {
self.memory
.written(&binding.memory, binding.range(), failures)
}
}
struct Factory {
config: MemoryConfiguration,
properties: MemoryDeviceProperties,
logger: Arc<ServerLogger>,
}
impl StreamFactory for Factory {
type Stream = TestStream;
fn create(&mut self) -> Self::Stream {
TestStream {
memory: MemoryManagement::from_configuration(
BytesStorage::default(),
&self.properties,
self.config.clone(),
self.logger.clone(),
MemoryManagementOptions::new("property harness"),
),
}
}
}
struct Device {
pool: StreamPool<Factory>,
failures: Failures,
capture: StreamCapture,
}
impl FailureStore for Device {
type Factory = Factory;
fn split(&mut self) -> (&mut StreamPool<Factory>, &mut Failures) {
(&mut self.pool, &mut self.failures)
}
fn parts(&self) -> (&StreamPool<Factory>, &Failures) {
(&self.pool, &self.failures)
}
}
impl WriteScoped for Device {
type Streams = Self;
fn write_streams(&mut self) -> &mut Self::Streams {
self
}
fn capturing(&mut self, _stream: StreamId) -> Option<&mut StreamCapture> {
Some(&mut self.capture)
}
}
struct Rng(u64);
impl Rng {
fn next(&mut self) -> u64 {
self.0 = self.0.wrapping_add(0x9E3779B97F4A7C15);
let mut z = self.0;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
z ^ (z >> 31)
}
fn below(&mut self, bound: usize) -> usize {
(self.next() % bound.max(1) as u64) as usize
}
fn chance(&mut self, percent: u64) -> bool {
self.next() % 100 < percent
}
}
struct Harness {
device: Device,
buffers: Vec<Buffer>,
rng: Rng,
}
fn error(reason: &str) -> ServerError {
ServerError::Generic {
reason: reason.into(),
backtrace: Default::default(),
}
}
fn reason(error: &ServerError) -> String {
format!("{error}")
}
impl Harness {
fn new(config: MemoryConfiguration, seed: u64) -> Self {
let properties = MemoryDeviceProperties::new(128 * 1024, 32);
let logger = Arc::new(ServerLogger::default());
Self {
device: Device {
pool: StreamPool::new(
Factory {
config,
properties,
logger: logger.clone(),
},
MAX_STREAMS,
0,
),
failures: Failures::new(logger),
capture: StreamCapture::default(),
},
buffers: Vec::new(),
rng: Rng(seed),
}
}
fn stream_id(&mut self) -> StreamId {
StreamId {
value: self.rng.next() % (MAX_STREAMS as u64 * 2),
}
}
fn alloc(&mut self) {
let id = self.stream_id();
let size = 32 * (1 + self.rng.below(64)) as u64;
let handle = Handle::new(service(), id, size);
let device = &mut self.device;
let stream = device.pool.get_mut(&id);
let reserved = match stream.memory.reserve(size, device.failures.graph_mut()) {
Ok(reserved) => reserved,
Err(err) => panic!("the harness never outgrows its pools: {err}"),
};
stream
.memory
.bind(
reserved,
handle.memory.clone(),
0,
device.failures.graph_mut(),
)
.unwrap();
self.buffers.push(Buffer {
binding: handle.clone().binding(),
_handle: handle,
stale: vec![false; size as usize],
});
}
fn pick_range(&mut self, size: u64) -> core::ops::Range<u64> {
if self.rng.chance(50) {
return 0..size;
}
let start = self.rng.below(size as usize) as u64;
let end = start + 1 + self.rng.below((size - start) as usize) as u64;
start..end
}
fn launch(&mut self, fail: bool) {
if self.buffers.is_empty() {
return;
}
let _id = self.stream_id();
let count = 1 + self.rng.below(3.min(self.buffers.len()));
let mut indices = Vec::new();
for _ in 0..count {
let index = self.rng.below(self.buffers.len());
if !indices.contains(&index) {
indices.push(index);
}
}
let mut writes = Vec::new();
for index in &indices {
let range = self.pick_range(self.buffers[*index].size());
writes.push((*index, range));
}
let bindings: Vec<BufferBinding> = writes
.iter()
.map(|(index, range)| self.buffers[*index].slice(range))
.collect();
let _ = ExecuteScope::over(&mut self.device, StreamId::current(), bindings).execute(|_| {
match fail {
true => Err(error("launch")),
false => Ok(()),
}
});
for (index, range) in writes {
for byte in range.start as usize..range.end as usize {
self.buffers[index].stale[byte] = fail;
}
}
}
fn host_write(&mut self) {
if self.buffers.is_empty() {
return;
}
let index = self.rng.below(self.buffers.len());
let range = self.pick_range(self.buffers[index].size());
let binding = self.buffers[index].slice(&range);
let _id = binding.stream;
let _ = ExecuteScope::over(&mut self.device, StreamId::current(), vec![binding])
.execute(|_| Ok::<(), ServerError>(()));
for byte in range.start as usize..range.end as usize {
self.buffers[index].stale[byte] = false;
}
}
fn read(&mut self) {
if self.buffers.is_empty() {
return;
}
let index = self.rng.below(self.buffers.len());
let range = self.pick_range(self.buffers[index].size());
let buffer = &self.buffers[index];
let binding = buffer.slice(&range);
let result = self.device.ensure_written([&binding].into_iter());
let stale = buffer.stale[range.start as usize..range.end as usize]
.iter()
.any(|stale| *stale);
match stale {
false => assert!(
result.is_ok(),
"bytes whose last writer succeeded must read: {result:?}"
),
true => assert!(
result.is_err(),
"bytes whose last writer failed must not read"
),
}
}
fn free(&mut self) {
if self.buffers.is_empty() {
return;
}
let index = self.rng.below(self.buffers.len());
self.buffers.swap_remove(index);
}
fn cleanup(&mut self) {
let id = self.stream_id();
let explicit = self.rng.chance(50);
let device = &mut self.device;
let stream = device.pool.get_mut(&id);
stream.memory.cleanup(explicit, device.failures.graph_mut());
}
fn sweep(&mut self) {
let device = &mut self.device;
for value in 0..MAX_STREAMS as u64 {
let id = StreamId { value };
let stream = device.pool.get_mut(&id);
stream.memory.cleanup(true, device.failures.graph_mut());
}
}
}
fn run(config: MemoryConfiguration, seed: u64) {
let mut harness = Harness::new(config, seed);
for _ in 0..OPS_PER_RUN {
match harness.rng.below(100) {
0..=19 => harness.alloc(),
20..=39 => harness.launch(true),
40..=59 => harness.launch(false),
60..=69 => harness.host_write(),
70..=89 => harness.read(),
90..=95 => harness.free(),
_ => harness.cleanup(),
}
}
for index in 0..harness.buffers.len() {
let buffer = &harness.buffers[index];
let result = harness.device.ensure_written([&buffer.binding].into_iter());
match buffer.stale.iter().any(|stale| *stale) {
false => assert!(result.is_ok(), "trusted buffer failed at rest: {result:?}"),
true => assert!(result.is_err(), "stale buffer read clean at rest"),
}
}
harness.buffers.clear();
harness.sweep();
assert!(
harness.device.failures.graph().is_empty(),
"the graph held {} failure(s) after every buffer was dropped and every pool swept",
harness.device.failures.graph().len()
);
}
#[test]
fn a_read_returns_bytes_iff_their_last_writer_succeeded_subslices() {
#[cfg(not(exclusive_memory_only))]
for seed in 0..SEEDS {
run(MemoryConfiguration::SubSlices, seed);
}
}
#[test]
fn a_read_returns_bytes_iff_their_last_writer_succeeded_exclusive_pages() {
for seed in 0..SEEDS {
run(MemoryConfiguration::ExclusivePages, seed);
}
}
#[test]
fn a_scope_that_succeeds_releases_the_provisional_failure() {
let mut harness = Harness::new(MemoryConfiguration::ExclusivePages, 7);
harness.alloc();
let binding = harness.buffers[0].binding.clone();
let result = ExecuteScope::over(
&mut harness.device,
StreamId::current(),
vec![binding.clone()],
)
.execute(|_| Ok::<(), ServerError>(()));
assert!(matches!(result, ScopedOutcome::Executed(())));
assert!(
harness.device.failures.graph().is_empty(),
"success leaves no node"
);
harness
.device
.ensure_written([&binding].into_iter())
.expect("a buffer whose writer succeeded reads");
}
#[test]
fn a_scope_that_fails_names_the_real_error_and_logs_it() {
let mut harness = Harness::new(MemoryConfiguration::ExclusivePages, 7);
harness.alloc();
let binding = harness.buffers[0].binding.clone();
let _id = binding.stream;
let result = ExecuteScope::over(
&mut harness.device,
StreamId::current(),
vec![binding.clone()],
)
.execute(|_| Err::<(), ServerError>(error("the real failure")));
assert!(matches!(result, ScopedOutcome::Failed(_)));
let read = harness
.device
.ensure_written([&binding].into_iter())
.expect_err("a buffer whose writer failed must not read");
let read = reason(&read);
assert!(
read.contains("the real failure"),
"the read fails on the body's error, got: {read}"
);
assert!(
!read.contains("torn down"),
"the provisional error was replaced, got: {read}"
);
}
#[test]
fn a_mid_launch_panic_leaves_the_write_set_tainted() {
let mut harness = Harness::new(MemoryConfiguration::ExclusivePages, 7);
harness.alloc();
let binding = harness.buffers[0].binding.clone();
let panicked = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let _ = ExecuteScope::over(
&mut harness.device,
StreamId::current(),
vec![binding.clone()],
)
.execute(|_| -> Result<(), ServerError> {
panic!("mid-launch, before anything could report")
});
}));
assert!(panicked.is_err(), "the panic propagates");
let read = harness
.device
.ensure_written([&binding].into_iter())
.expect_err("the write set must be tainted after a mid-launch panic");
let read = reason(&read);
assert!(
read.contains("torn down"),
"the read fails on the provisional error, got: {read}"
);
let result = ExecuteScope::over(
&mut harness.device,
StreamId::current(),
vec![binding.clone()],
)
.execute(|_| Ok::<(), ServerError>(()));
assert!(matches!(result, ScopedOutcome::Executed(())));
assert!(harness.device.failures.graph().is_empty());
harness
.device
.ensure_written([&binding].into_iter())
.expect("a rewritten buffer reads again");
}
#[test]
fn a_skip_returns_the_pooled_write_set() {
let mut harness = Harness::new(MemoryConfiguration::ExclusivePages, 23);
harness.alloc();
harness.alloc();
let input = harness.buffers[0].binding.clone();
let output = harness.buffers[1].binding.clone();
let mut primed = FailureStore::write_set(&mut harness.device);
primed.push(input.clone());
primed.push(output.clone());
let pooled = primed.as_ptr();
ExecuteScope::over(&mut harness.device, StreamId::current(), primed)
.execute(|_| Ok::<(), ServerError>(()));
let mut failing = FailureStore::write_set(&mut harness.device);
assert_eq!(
failing.as_ptr(),
pooled,
"the clean exit returned the buffer"
);
failing.push(input.clone());
ExecuteScope::over(&mut harness.device, StreamId::current(), failing)
.execute(|_| Err::<(), ServerError>(error("the launch that left these bytes")));
let mut skipping = FailureStore::write_set(&mut harness.device);
assert_eq!(skipping.as_ptr(), pooled, "the failed exit returned it too");
skipping.push(output.clone());
let outcome = ExecuteScope::launching(
&mut harness.device,
KernelId::new::<()>(),
StreamId::current(),
[&input].into_iter(),
skipping,
)
.execute(|_| -> Result<(), ServerError> {
unreachable!("a skipped scope must not run its body")
});
assert!(matches!(outcome, ScopedOutcome::Skipped));
let returned = FailureStore::write_set(&mut harness.device);
assert_eq!(
returned.as_ptr(),
pooled,
"the skip path must hand the write set back, not drop it and allocate a new one"
);
}
#[test]
fn a_skipped_scope_claims_the_failure_its_input_carried_and_mints_none() {
let mut harness = Harness::new(MemoryConfiguration::ExclusivePages, 11);
harness.alloc();
harness.alloc();
let input = harness.buffers[0].binding.clone();
let output = harness.buffers[1].binding.clone();
let _ = ExecuteScope::over(
&mut harness.device,
StreamId::current(),
vec![input.clone()],
)
.execute(|_| Err::<(), ServerError>(error("the launch that left these bytes")));
let carried = harness.device.failures.graph().len();
assert_eq!(carried, 1, "one failure, on the input");
let outcome = ExecuteScope::launching(
&mut harness.device,
KernelId::new::<()>(),
StreamId::current(),
[&input].into_iter(),
vec![output.clone()],
)
.execute(|_| -> Result<(), ServerError> {
unreachable!("a skipped scope must not run its body")
});
assert!(matches!(outcome, ScopedOutcome::Skipped));
assert_eq!(
harness.device.failures.graph().len(),
carried,
"a skip reuses the failure it found, and a provisional would linger"
);
let read = harness
.device
.ensure_written([&output].into_iter())
.expect_err("a skipped launch's output must not read");
assert!(
reason(&read).contains("the launch that left these bytes"),
"the output names the root cause, got: {}",
reason(&read)
);
}
#[test]
fn a_partial_host_write_releases_only_the_bytes_it_covers() {
let mut harness = Harness::new(MemoryConfiguration::ExclusivePages, 7);
harness.alloc();
let buffer = &harness.buffers[0];
let size = buffer.size();
let whole = buffer.binding.clone();
let middle = buffer.slice(&(size / 4..size / 2));
let _id = whole.stream;
let _ = ExecuteScope::over(
&mut harness.device,
StreamId::current(),
vec![whole.clone()],
)
.execute(|_| Err::<(), ServerError>(error("the launch that left these bytes")));
let _ = ExecuteScope::over(
&mut harness.device,
StreamId::current(),
vec![middle.clone()],
)
.execute(|_| Ok::<(), ServerError>(()));
harness
.device
.ensure_written([&middle].into_iter())
.expect("the rewritten bytes have a writer");
let front = harness.buffers[0].slice(&(0..size / 4));
let read = harness
.device
.ensure_written([&front].into_iter())
.expect_err("the untouched bytes still carry the failure");
assert!(
reason(&read).contains("the launch that left these bytes"),
"the remainder still names the original failure"
);
let whole_read = harness.device.ensure_written([&whole].into_iter());
assert!(
whole_read.is_err(),
"a read spanning stale bytes fails however much of it was rewritten"
);
}
#[test]
fn a_failed_scope_dooms_the_recording_window() {
let mut harness = Harness::new(MemoryConfiguration::ExclusivePages, 7);
harness.alloc();
let binding = harness.buffers[0].binding.clone();
harness.device.capture.prepare(StreamId::current()).unwrap();
harness.device.capture.begin().unwrap();
let _ = ExecuteScope::over(
&mut harness.device,
StreamId::current(),
vec![binding.clone()],
)
.execute(|_| Err::<(), ServerError>(error("failed mid-window")));
harness.device.capture.end(StreamId::current()).unwrap();
let doomed = harness
.device
.capture
.take_failure()
.expect("a failure inside the window must doom the recording");
assert!(
doomed.to_string().contains("failed mid-window"),
"the window names the failure that doomed it, got: {doomed}"
);
}
#[test]
fn a_skipped_scope_dooms_the_recording_window() {
let mut harness = Harness::new(MemoryConfiguration::ExclusivePages, 11);
harness.alloc();
harness.alloc();
let input = harness.buffers[0].binding.clone();
let output = harness.buffers[1].binding.clone();
let _ = ExecuteScope::over(
&mut harness.device,
StreamId::current(),
vec![input.clone()],
)
.execute(|_| Err::<(), ServerError>(error("the writer that never wrote")));
harness.device.capture.prepare(StreamId::current()).unwrap();
harness.device.capture.begin().unwrap();
let outcome = ExecuteScope::launching(
&mut harness.device,
KernelId::new::<()>(),
StreamId::current(),
[&input].into_iter(),
vec![output],
)
.execute(|_| -> Result<(), ServerError> {
unreachable!("a skipped scope must not run its body")
});
assert!(matches!(outcome, ScopedOutcome::Skipped));
harness.device.capture.end(StreamId::current()).unwrap();
assert!(
harness.device.capture.take_failure().is_some(),
"a skip inside the window must doom the recording"
);
}
#[test]
fn a_clean_scope_hands_the_window_its_write_set() {
let mut harness = Harness::new(MemoryConfiguration::ExclusivePages, 13);
harness.alloc();
harness.alloc();
let recorded = harness.buffers[0].binding.clone();
let failed = harness.buffers[1].binding.clone();
harness.device.capture.prepare(StreamId::current()).unwrap();
harness.device.capture.begin().unwrap();
let _ = ExecuteScope::over(
&mut harness.device,
StreamId::current(),
vec![recorded.clone()],
)
.execute(|_| Ok::<(), ServerError>(()));
let _ = ExecuteScope::over(
&mut harness.device,
StreamId::current(),
vec![failed.clone()],
)
.execute(|_| Err::<(), ServerError>(error("enqueued nothing")));
harness.device.capture.end(StreamId::current()).unwrap();
let written = harness.device.capture.take_recorded();
assert_eq!(
written.len(),
1,
"the window holds exactly what clean scopes wrote"
);
assert_eq!(
written[0].claim_key(),
recorded.claim_key(),
"and it is the clean scope's write set, not the failed one's"
);
}
#[test]
fn a_scope_outside_a_window_neither_dooms_nor_records() {
let mut harness = Harness::new(MemoryConfiguration::ExclusivePages, 17);
harness.alloc();
let binding = harness.buffers[0].binding.clone();
let _ = ExecuteScope::over(
&mut harness.device,
StreamId::current(),
vec![binding.clone()],
)
.execute(|_| Ok::<(), ServerError>(()));
let _ = ExecuteScope::over(&mut harness.device, StreamId::current(), vec![binding])
.execute(|_| Err::<(), ServerError>(error("no window to doom")));
assert!(harness.device.capture.take_failure().is_none());
assert!(harness.device.capture.take_recorded().is_empty());
}