use super::*;
#[test]
fn chained_unary_ops_with_no_intermediate_sync_stay_ordered() {
let Some(b) = try_init() else { return };
let input = [-3.5f32, 0.0, 2.25, -100.0, 7.0, -0.001, 42.0, -8.0];
let n = input.len();
let mut bufs: Vec<u64> = vec![upload_f32(&b, &input)];
bufs.extend((0..8).map(|_| upload_f32(&b, &vec![0.0f32; n])));
for i in 0..8 {
b.unary(UnaryOp::Neg, bufs[i], bufs[i + 1], n)
.unwrap_or_else(|e| panic!("unary Neg step {i} failed: {e:?}"));
}
let result = download_f32(&b, bufs[8], n);
for (r, e) in result.iter().zip(input.iter()) {
assert!(
(r - e).abs() < 1e-5,
"chained 8x-negate result {result:?} does not match original input {input:?} \
(got {r}, expected {e}) — dispatch ordering broke without a per-op poll"
);
}
for h in bufs {
b.free(h).expect("free");
}
}
#[test]
fn chained_gemm_reusing_same_handles_observes_latest_uniform_each_call() {
let Some(b) = try_init() else { return };
let a = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0];
let bm = [7.0f32, 8.0, 9.0, 10.0, 11.0, 12.0];
let base = [58.0f32, 64.0, 139.0, 154.0];
let a_h = upload_f32(&b, &a);
let b_h = upload_f32(&b, &bm);
let c_h = upload_f32(&b, &[0.0f32; 4]);
let nt = BackendTranspose::NoTrans;
const ITERS: usize = 30;
let mut last_alpha = 0.0f64;
for i in 0..ITERS {
let alpha = (i as f64) * 1000.0 + 7.0;
last_alpha = alpha;
b.gemm(nt, nt, 2, 2, 3, alpha, a_h, 3, b_h, 2, 0.0, c_h, 2)
.unwrap_or_else(|e| panic!("gemm iteration {i} (alpha={alpha}) failed: {e:?}"));
}
let result = download_f32(&b, c_h, 4);
for (r, base_v) in result.iter().zip(base.iter()) {
let expected = base_v * (last_alpha as f32);
let tol = (expected.abs() * 1e-5).max(0.5);
assert!(
(r - expected).abs() < tol,
"reused-handle gemm chain result {result:?} does not reflect the last \
alpha={last_alpha} (expected ~{expected}, got {r}) — a stale cached uniform \
buffer or bind group would produce exactly this symptom"
);
}
b.free(a_h).expect("free");
b.free(b_h).expect("free");
b.free(c_h).expect("free");
}
#[test]
fn bind_group_cache_entry_is_evicted_on_free() {
let Some(b) = try_init() else { return };
let nt = BackendTranspose::NoTrans;
let a1 = upload_f32(&b, &[1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0]);
let b1 = upload_f32(&b, &[1.0f32, 0.0, 0.0, 1.0, 1.0, 1.0]);
let c1 = upload_f32(&b, &[0.0f32; 4]);
b.gemm(nt, nt, 2, 2, 3, 1.0, a1, 3, b1, 2, 0.0, c1, 2)
.expect("gemm entry 1");
let a2 = upload_f32(&b, &[7.0f32, 8.0, 9.0, 10.0, 11.0, 12.0]);
let b2 = upload_f32(&b, &[1.0f32, 0.0, 0.0, 1.0, 1.0, 1.0]);
let c2 = upload_f32(&b, &[0.0f32; 4]);
b.gemm(nt, nt, 2, 2, 3, 1.0, a2, 3, b2, 2, 0.0, c2, 2)
.expect("gemm entry 2");
let entries_after_two_gemms = b
.bind_group_cache
.lock()
.expect("bind_group_cache lock")
.len();
assert_eq!(
entries_after_two_gemms, 2,
"two gemm calls with disjoint operand handles should populate two distinct cache entries"
);
b.free(c1).expect("free c1");
let entries_after_free = b
.bind_group_cache
.lock()
.expect("bind_group_cache lock")
.len();
assert_eq!(
entries_after_free, 1,
"freeing a handle referenced by one cached bind group must evict only that entry"
);
b.gemm(nt, nt, 2, 2, 3, 1.0, a2, 3, b2, 2, 0.0, c2, 2)
.expect("gemm reusing entry 2's cached bind group after an unrelated eviction");
let result = download_f32(&b, c2, 4);
let expected = [16.0f32, 17.0, 22.0, 23.0];
for (r, e) in result.iter().zip(expected.iter()) {
assert!((r - e).abs() < 1e-3, "got {r}, expected {e}");
}
b.free(a1).expect("free a1");
b.free(b1).expect("free b1");
b.free(a2).expect("free a2");
b.free(b2).expect("free b2");
b.free(c2).expect("free c2");
}
#[test]
fn synchronize_surfaces_recorded_uncaptured_error() {
let Some(b) = try_init() else { return };
let dev = b.device().expect("device() after successful init");
assert!(
dev.poll_error().is_none(),
"precondition: no error recorded yet"
);
let _bogus = dev.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("oxicuda-webgpu-test-oversize"),
size: u64::MAX,
usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_SRC
| wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let err = b
.synchronize()
.expect_err("synchronize() should surface the recorded uncaptured error, not Ok(())");
assert!(
matches!(err, BackendError::DeviceError(_)),
"got {err:?}, expected a DeviceError wrapping the uncaptured wgpu error"
);
}
#[test]
fn synchronize_is_ok_with_nothing_pending() {
let Some(b) = try_init() else { return };
let a_h = upload_f32(&b, &[1.0f32, 2.0, 3.0, 4.0]);
let out_h = upload_f32(&b, &[0.0f32; 4]);
b.unary(UnaryOp::Relu, a_h, out_h, 4)
.expect("unary before synchronize");
b.synchronize()
.expect("synchronize() must be Ok when nothing failed");
b.free(a_h).expect("free");
b.free(out_h).expect("free");
}
#[test]
fn bind_group_cache_order_does_not_grow_unbounded_across_alloc_dispatch_free() {
let Some(b) = try_init() else { return };
let nt = BackendTranspose::NoTrans;
const ITERS: usize = 200;
for _ in 0..ITERS {
let a_h = upload_f32(&b, &[1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0]);
let b_h = upload_f32(&b, &[1.0f32, 0.0, 0.0, 1.0, 1.0, 1.0]);
let c_h = upload_f32(&b, &[0.0f32; 4]);
b.gemm(nt, nt, 2, 2, 3, 1.0, a_h, 3, b_h, 2, 0.0, c_h, 2)
.expect("gemm");
b.free(a_h).expect("free a");
b.free(b_h).expect("free b");
b.free(c_h).expect("free c");
}
let cache = b.bind_group_cache.lock().expect("bind_group_cache lock");
assert_eq!(
cache.len(),
0,
"every entry's handles were freed every iteration; none should remain live"
);
assert_eq!(
cache.order_len(),
0,
"order should track exactly the live entries (0), not accumulate one \
dead key per iteration ({ITERS} would indicate the pre-fix unbounded-growth bug)"
);
}