use std::time::{Duration, Instant};
use vyre_driver::{DispatchConfig, VyreBackend};
use vyre_driver_wgpu::WgpuBackend;
use vyre_foundation::ir::{BufferDecl, DataType, Expr, Node, Program};
fn long_running_program() -> Program {
const OUTPUT_WORDS: u32 = 16 * 1024 * 1024;
let mut body = Vec::with_capacity(515);
body.push(Node::let_bind("idx", Expr::gid_x()));
body.push(Node::let_bind("acc", Expr::var("idx")));
for round in 0..512u32 {
body.push(Node::assign(
"acc",
Expr::bitxor(
Expr::mul(Expr::var("acc"), Expr::u32(1_664_525)),
Expr::add(
Expr::var("idx"),
Expr::u32(1_013_904_223u32.wrapping_add(round)),
),
),
));
}
body.push(Node::if_then(
Expr::lt(Expr::var("idx"), Expr::buf_len("out")),
vec![Node::store("out", Expr::var("idx"), Expr::var("acc"))],
));
Program::wrapped(
vec![BufferDecl::output("out", 0, DataType::U32)
.with_count(OUTPUT_WORDS)
.with_output_byte_range(0..4)],
[256, 1, 1],
body,
)
}
#[test]
fn dispatch_cancels_within_deadline() {
let backend = WgpuBackend::acquire().expect("Fix: GPU required for pre-emption test");
let program = long_running_program();
let mut config = DispatchConfig::default();
config.timeout = Some(Duration::from_millis(100));
config.label = Some("dispatch-preemption".to_string());
let start = Instant::now();
let result = backend.dispatch(&program, &[], &config);
let elapsed = start.elapsed();
assert!(
result.is_err(),
"dispatch preemption: dispatch must return Err on timeout, got Ok"
);
assert!(
elapsed < Duration::from_secs(2),
"dispatch preemption: cancellation must complete within 2s of the deadline; took {:?}",
elapsed
);
let quick = vyre::Program::empty();
let _ = backend
.dispatch(&quick, &[], &DispatchConfig::default())
.expect("Fix: device must be usable after cancelled dispatch");
}