mod support;
use lampshade::{Error, KeyValue, KeyValueSoaSorter, KeyValueSorter};
use wgpu::util::DeviceExt;
const SORT_SIZES: [usize; 18] = [
0, 1, 2, 31, 32, 33, 127, 128, 129, 511, 512, 513, 2_047, 2_048, 2_049, 4_097, 65_537, 17,
];
fn cpu_stable_sort(input: &[KeyValue]) -> Vec<KeyValue> {
let mut expected = input.to_vec();
expected.sort_by_key(|item| item.key);
expected
}
fn duplicate_keys(len: usize, seed: u64) -> Vec<KeyValue> {
support::random_u32(len, seed)
.into_iter()
.enumerate()
.map(|(index, key)| KeyValue::new(key & 0x0f, index as u32))
.collect()
}
#[tokio::test]
async fn key_value_sort_matches_stable_cpu_sort_across_boundaries() {
let Some(context) = support::gpu_context().await else {
return;
};
let mut sorter = KeyValueSorter::from_context(&context);
for (case, size) in SORT_SIZES.into_iter().enumerate() {
let input = duplicate_keys(size, case as u64);
let actual = sorter
.sort(&input)
.await
.expect("GPU key-value sort failed");
assert_eq!(
actual,
cpu_stable_sort(&input),
"key-value sort mismatch for size {size}"
);
}
}
#[tokio::test]
async fn key_value_sort_preserves_value_order_for_equal_keys() {
let Some(context) = support::gpu_context().await else {
return;
};
let mut sorter = KeyValueSorter::from_context(&context);
let input: Vec<_> = (0..4_097).map(|value| KeyValue::new(42, value)).collect();
let actual = sorter
.sort(&input)
.await
.expect("GPU key-value sort failed");
assert_eq!(actual, input);
let edge_keys = [
KeyValue::new(u32::MAX, 0),
KeyValue::new(0, 1),
KeyValue::new(u32::MAX, 2),
KeyValue::new(0, 3),
KeyValue::new(1, 4),
];
assert_eq!(
sorter
.sort(&edge_keys)
.await
.expect("GPU key-value sort failed"),
cpu_stable_sort(&edge_keys)
);
}
#[tokio::test]
async fn key_value_sort_handles_full_width_keys_across_many_tiles() {
let Some(context) = support::gpu_context().await else {
return;
};
let mut sorter = KeyValueSorter::from_context(&context);
let input: Vec<_> = support::random_u32(262_147, 0x00F0_1132)
.into_iter()
.enumerate()
.map(|(index, key)| KeyValue::new(key, index as u32))
.collect();
let actual = sorter
.sort(&input)
.await
.expect("full-width GPU key-value sort failed");
assert_eq!(actual, cpu_stable_sort(&input));
}
#[tokio::test]
async fn key_value_sort_is_stable_for_large_duplicate_heavy_input() {
let Some(context) = support::gpu_context().await else {
return;
};
let mut sorter = KeyValueSorter::from_context(&context);
let input = duplicate_keys(1_000_003, 0xD001_1CA7);
let actual = sorter
.sort(&input)
.await
.expect("large duplicate-heavy GPU key-value sort failed");
assert_eq!(actual, cpu_stable_sort(&input));
}
#[tokio::test]
async fn bounded_key_value_sort_is_stable_across_pass_parities() {
let Some(context) = support::gpu_context().await else {
return;
};
let mut portable_sorter = KeyValueSorter::new(&context.device, &context.queue);
let mut adapter_sorter = KeyValueSorter::from_context(&context);
for key_bits in [0, 1, 8, 9, 16, 17, 24, 25, 32] {
let mask = if key_bits == 32 {
u32::MAX
} else if key_bits == 0 {
0
} else {
(1_u32 << key_bits) - 1
};
let input: Vec<_> = support::random_u32(4_097, u64::from(key_bits) + 900)
.into_iter()
.enumerate()
.map(|(index, key)| KeyValue::new(key & mask, index as u32))
.collect();
let expected = cpu_stable_sort(&input);
let portable_actual = portable_sorter
.sort_with_key_bits(&input, key_bits)
.await
.expect("portable bounded key-value sort failed");
assert_eq!(
portable_actual, expected,
"portable mismatch for {key_bits} key bits"
);
let adapter_actual = adapter_sorter
.sort_with_key_bits(&input, key_bits)
.await
.expect("adapter bounded key-value sort failed");
assert_eq!(
adapter_actual, expected,
"adapter mismatch for {key_bits} key bits"
);
}
}
#[tokio::test]
async fn bounded_key_value_sort_validates_host_keys() {
let Some(context) = support::gpu_context().await else {
return;
};
let mut sorter = KeyValueSorter::new(&context.device, &context.queue);
let error = sorter
.sort_with_key_bits(&[KeyValue::new(256, 0)], 8)
.await
.expect_err("a nine-bit key must not satisfy an eight-bit bound");
assert!(matches!(error, Error::KeyExceedsBitRange { .. }));
}
#[tokio::test]
async fn key_value_sort_gpu_to_gpu_writes_the_caller_output_buffer() {
let Some(context) = support::gpu_context().await else {
return;
};
let mut sorter = KeyValueSorter::from_context(&context);
let input = duplicate_keys(4_097, 100);
let input_buffer = create_sort_input(&context.device, &input);
let output = create_sort_output(&context.device, input.len());
sorter
.sort_gpu_to_gpu(&input_buffer, &output, input.len() as u32)
.expect("GPU key-value sort failed");
let actual = support::read_pod::<KeyValue>(&context, &output, input.len()).await;
assert_eq!(actual, cpu_stable_sort(&input));
}
#[tokio::test]
async fn counted_key_value_sort_uses_the_gpu_resident_prefix_and_stays_stable() {
let Some(context) = support::gpu_context().await else {
return;
};
let input = [
KeyValue::new(7, 70),
KeyValue::new(2, 20),
KeyValue::new(2, 21),
KeyValue::new(1, 10),
KeyValue::new(0, 0),
];
let selected = 4_u32;
let input_buffer = create_sort_input(&context.device, &input);
let output = create_sort_output(&context.device, input.len());
let count = context
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("Key-Value Sort GPU Count"),
contents: bytemuck::bytes_of(&selected),
usage: wgpu::BufferUsages::STORAGE,
});
let mut sorter = KeyValueSorter::from_context(&context);
sorter
.sort_counted_gpu_to_gpu_with_key_bits(
&input_buffer,
&output,
&count,
input.len() as u32,
3,
)
.expect("counted key-value sort failed");
let actual = support::read_pod::<KeyValue>(&context, &output, selected as usize).await;
assert_eq!(
actual,
[
KeyValue::new(1, 10),
KeyValue::new(2, 20),
KeyValue::new(2, 21),
KeyValue::new(7, 70),
]
);
}
#[tokio::test]
async fn native_soa_counted_sort_is_stable_for_a_gpu_selected_prefix() {
let Some(context) = support::gpu_context().await else {
return;
};
let Some(mut sorter) =
KeyValueSoaSorter::new_for_adapter(&context.device, &context.adapter_info)
else {
return;
};
let capacity = 65_537_u32;
let active = 65_521_u32;
let keys: Vec<_> = support::random_u32(capacity as usize, 0x50A5_0A11)
.into_iter()
.map(|key| key & 0xff)
.collect();
let values: Vec<_> = (0..capacity).collect();
let key_buffer = context
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("SoA Sort Keys"),
contents: bytemuck::cast_slice(&keys),
usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
});
let value_buffer = context
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("SoA Sort Values"),
contents: bytemuck::cast_slice(&values),
usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
});
let count = context
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("SoA Sort Count"),
contents: bytemuck::bytes_of(&active),
usage: wgpu::BufferUsages::STORAGE,
});
let mut encoder = context
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
sorter
.record_sort_counted(&mut encoder, &key_buffer, &value_buffer, &count, capacity)
.expect("native SoA counted sort records");
context.queue.submit(Some(encoder.finish()));
let actual_keys = support::read_pod::<u32>(&context, &key_buffer, active as usize).await;
let actual_values = support::read_pod::<u32>(&context, &value_buffer, active as usize).await;
let mut expected: Vec<_> = keys[..active as usize]
.iter()
.copied()
.zip(values[..active as usize].iter().copied())
.collect();
expected.sort_by_key(|&(key, _)| key);
assert_eq!(
actual_keys
.into_iter()
.zip(actual_values)
.collect::<Vec<_>>(),
expected
);
}
#[tokio::test]
async fn native_soa_counted_sort_clamps_oversized_counts_and_accepts_zero() {
let Some(context) = support::gpu_context().await else {
return;
};
let Some(mut sorter) =
KeyValueSoaSorter::new_native_for_adapter(&context.device, &context.adapter_info)
else {
return;
};
let keys = [9_u32, 1, 1, 7, 0, 7, 3];
let values: Vec<u32> = (0..keys.len() as u32).collect();
let key_buffer = context
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("SoA Clamp Keys"),
contents: bytemuck::cast_slice(&keys),
usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_SRC
| wgpu::BufferUsages::COPY_DST,
});
let value_buffer = context
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("SoA Clamp Values"),
contents: bytemuck::cast_slice(&values),
usage: wgpu::BufferUsages::STORAGE
| wgpu::BufferUsages::COPY_SRC
| wgpu::BufferUsages::COPY_DST,
});
let count = context
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("SoA Clamp Count"),
contents: bytemuck::bytes_of(&u32::MAX),
usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
});
let mut encoder = context
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
sorter
.record_sort_counted(
&mut encoder,
&key_buffer,
&value_buffer,
&count,
keys.len() as u32,
)
.expect("oversized count is clamped");
context.queue.submit(Some(encoder.finish()));
let actual_keys = support::read_pod::<u32>(&context, &key_buffer, keys.len()).await;
let actual_values = support::read_pod::<u32>(&context, &value_buffer, keys.len()).await;
let mut expected: Vec<_> = keys.into_iter().zip(values.iter().copied()).collect();
expected.sort_by_key(|&(key, _)| key);
assert_eq!(
actual_keys
.into_iter()
.zip(actual_values)
.collect::<Vec<_>>(),
expected
);
context
.queue
.write_buffer(&count, 0, bytemuck::bytes_of(&0_u32));
let mut encoder = context
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
sorter
.record_sort_counted(
&mut encoder,
&key_buffer,
&value_buffer,
&count,
keys.len() as u32,
)
.expect("zero GPU count records");
context.queue.submit(Some(encoder.finish()));
assert_eq!(
support::read_pod::<u32>(&context, &key_buffer, keys.len()).await,
expected.iter().map(|&(key, _)| key).collect::<Vec<_>>()
);
}
#[tokio::test]
async fn native_soa_reserved_sort_accepts_zero_capacity_without_workspace() {
let Some(context) = support::gpu_context().await else {
return;
};
let Some(mut sorter) =
KeyValueSoaSorter::new_native_for_adapter(&context.device, &context.adapter_info)
else {
return;
};
let keys = context.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("Zero-Capacity SoA Keys"),
size: 4,
usage: wgpu::BufferUsages::STORAGE,
mapped_at_creation: false,
});
let values = context.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("Zero-Capacity SoA Values"),
size: 4,
usage: wgpu::BufferUsages::STORAGE,
mapped_at_creation: false,
});
let count = context
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("Zero-Capacity SoA Count"),
contents: bytemuck::bytes_of(&u32::MAX),
usage: wgpu::BufferUsages::STORAGE,
});
sorter
.prepare_counted_from_word(&keys, &values, &count, 0, 0)
.expect("zero-capacity plan requires no workspace");
let mut encoder = context
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
sorter
.record_reserved_sort_counted_from_word(&mut encoder, &keys, &values, &count, 0, 0)
.expect("zero-capacity reserved recording is a no-op");
let submission = context.queue.submit(Some(encoder.finish()));
context
.device
.poll(wgpu::PollType::Wait {
submission_index: Some(submission),
timeout: None,
})
.expect("zero-capacity submission completes");
}
#[tokio::test]
async fn soa_fixed_sort_hides_the_count_buffer_and_preserves_stability() {
let Some(context) = support::gpu_context().await else {
return;
};
let mut sorter = KeyValueSoaSorter::from_context(&context);
let keys: Vec<_> = support::random_u32(65_537, 0xF1CE_D50A)
.into_iter()
.map(|key| key & 0xff)
.collect();
let values: Vec<_> = (0..keys.len() as u32).collect();
let (key_buffer, value_buffer) = create_soa_buffers(&context.device, &keys, &values);
sorter
.prepare_sort(&key_buffer, &value_buffer, keys.len() as u32)
.expect("fixed SoA sort prepares without a caller count buffer");
let mut encoder = context
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
sorter
.record_reserved_sort(&mut encoder, &key_buffer, &value_buffer, keys.len() as u32)
.expect("prepared fixed SoA sort records");
context.queue.submit(Some(encoder.finish()));
assert_soa_prefix(
&context,
&key_buffer,
&value_buffer,
&keys,
&values,
keys.len(),
)
.await;
}
#[tokio::test]
async fn portable_soa_backend_matches_fixed_and_gpu_counted_contracts() {
let Some(context) = support::gpu_context_without_optional_features().await else {
return;
};
let mut sorter = KeyValueSoaSorter::new(&context.device, &context.queue);
assert!(!sorter.is_accelerated());
let keys: Vec<_> = support::random_u32(4_097, 0xB81D_6E50)
.into_iter()
.map(|key| key & 0x3f)
.collect();
let values: Vec<_> = (0..keys.len() as u32).collect();
let (key_buffer, value_buffer) = create_soa_buffers(&context.device, &keys, &values);
let mut encoder = context
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
sorter
.record_sort(&mut encoder, &key_buffer, &value_buffer, keys.len() as u32)
.expect("fixed sort uses the portable bridge");
context.queue.submit(Some(encoder.finish()));
assert_soa_prefix(
&context,
&key_buffer,
&value_buffer,
&keys,
&values,
keys.len(),
)
.await;
let counted_keys: Vec<_> = support::random_u32(4_097, 0xC0A7_ED50)
.into_iter()
.map(|key| key & 0xff)
.collect();
let counted_values: Vec<_> = (0..counted_keys.len() as u32).collect();
let active = 2_049_u32;
let (counted_key_buffer, counted_value_buffer) =
create_soa_buffers(&context.device, &counted_keys, &counted_values);
let count_words = [u32::MAX, active, 17, 23];
let count = context
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("Portable SoA Count Words"),
contents: bytemuck::cast_slice(&count_words),
usage: wgpu::BufferUsages::STORAGE,
});
let mut encoder = context
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
sorter
.record_sort_counted_from_word(
&mut encoder,
&counted_key_buffer,
&counted_value_buffer,
&count,
1,
counted_keys.len() as u32,
)
.expect("GPU-counted sort uses the portable bridge");
context.queue.submit(Some(encoder.finish()));
assert_soa_prefix(
&context,
&counted_key_buffer,
&counted_value_buffer,
&counted_keys,
&counted_values,
active as usize,
)
.await;
let mut encoder = context
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
sorter
.record_sort(&mut encoder, &key_buffer, &value_buffer, keys.len() as u32)
.expect("fixed plan is rebuilt after a counted sort changes bindings");
context.queue.submit(Some(encoder.finish()));
assert_soa_prefix(
&context,
&key_buffer,
&value_buffer,
&keys,
&values,
keys.len(),
)
.await;
}
#[tokio::test]
async fn portable_soa_backend_rejects_all_caller_buffer_aliases() {
let Some(context) = support::gpu_context_without_optional_features().await else {
return;
};
let mut sorter = KeyValueSoaSorter::new_portable(&context.device, &context.queue);
let data = context
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("Aliased Portable SoA Data"),
contents: bytemuck::cast_slice(&[1_u32]),
usage: wgpu::BufferUsages::STORAGE,
});
let other = context
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("Distinct Portable SoA Data"),
contents: bytemuck::cast_slice(&[1_u32]),
usage: wgpu::BufferUsages::STORAGE,
});
for (keys, values, count, first, second) in [
(&data, &data, &other, "sort keys", "sort values"),
(&data, &other, &data, "sort keys", "sort item count"),
(&other, &data, &data, "sort values", "sort item count"),
] {
let mut encoder = context
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
let error = sorter
.record_sort_counted(&mut encoder, keys, values, count, 1)
.expect_err("aliased SoA buffers must be rejected before recording");
assert!(matches!(
error,
Error::BufferAlias {
first: actual_first,
second: actual_second,
} if actual_first == first && actual_second == second
));
}
}
#[tokio::test]
async fn portable_soa_consumes_a_count_written_earlier_in_the_same_encoder() {
let Some(context) = support::gpu_context_without_optional_features().await else {
return;
};
let keys = [9_u32, 2, 2, 1, 0, 8, 7];
let values: Vec<_> = (0..keys.len() as u32).collect();
let active = 4_u32;
let (key_buffer, value_buffer) = create_soa_buffers(&context.device, &keys, &values);
let produced_count = context
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("Portable SoA Produced Count"),
contents: bytemuck::bytes_of(&active),
usage: wgpu::BufferUsages::COPY_SRC,
});
let count = context
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("Portable SoA Consumed Count"),
contents: bytemuck::bytes_of(&0_u32),
usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
});
let mut sorter = KeyValueSoaSorter::new_portable(&context.device, &context.queue);
sorter
.prepare_counted_from_word(&key_buffer, &value_buffer, &count, 0, keys.len() as u32)
.expect("portable SoA counted plan prepares");
let mut encoder = context
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
encoder.copy_buffer_to_buffer(&produced_count, 0, &count, 0, size_of::<u32>() as u64);
sorter
.record_reserved_sort_counted_from_word(
&mut encoder,
&key_buffer,
&value_buffer,
&count,
0,
keys.len() as u32,
)
.expect("portable SoA consumes the earlier count write");
context.queue.submit(Some(encoder.finish()));
assert_soa_prefix(
&context,
&key_buffer,
&value_buffer,
&keys,
&values,
active as usize,
)
.await;
}
#[tokio::test]
async fn bounded_key_value_gpu_sort_writes_output_after_one_and_three_byte_passes() {
let Some(context) = support::gpu_context().await else {
return;
};
let mut sorter = KeyValueSorter::from_context(&context);
for (key_bits, input) in [
(
8,
vec![
KeyValue::new(255, 0),
KeyValue::new(0, 1),
KeyValue::new(17, 2),
KeyValue::new(3, 3),
KeyValue::new(17, 4),
],
),
(
17,
vec![
KeyValue::new(0x1ffff, 0),
KeyValue::new(0, 1),
KeyValue::new(0x10001, 2),
KeyValue::new(0xff, 3),
KeyValue::new(0x10001, 4),
],
),
] {
let input_buffer = create_sort_input(&context.device, &input);
let output = create_sort_output(&context.device, input.len());
sorter
.sort_gpu_to_gpu_with_key_bits(&input_buffer, &output, input.len() as u32, key_bits)
.expect("bounded key-value GPU sort failed");
let actual = support::read_pod::<KeyValue>(&context, &output, input.len()).await;
assert_eq!(
actual,
cpu_stable_sort(&input),
"caller output mismatch for {key_bits} key bits"
);
}
}
#[tokio::test]
async fn bounded_key_value_sort_rebuilds_cached_bindings_when_pass_count_changes() {
let Some(context) = support::gpu_context().await else {
return;
};
let mut sorter = KeyValueSorter::from_context(&context);
let input = [
KeyValue::new(255, 0),
KeyValue::new(0, 1),
KeyValue::new(17, 2),
KeyValue::new(3, 3),
KeyValue::new(17, 4),
];
let input_buffer = create_sort_input(&context.device, &input);
let output = create_sort_output(&context.device, input.len());
let expected = cpu_stable_sort(&input);
for key_bits in [17, 8, 16, 24, 32, 8] {
sorter
.sort_gpu_to_gpu_with_key_bits(&input_buffer, &output, input.len() as u32, key_bits)
.expect("bounded key-value GPU sort failed");
assert_eq!(
support::read_pod::<KeyValue>(&context, &output, input.len()).await,
expected,
"cache rebuild mismatch for {key_bits} key bits"
);
}
}
#[tokio::test]
async fn zero_bit_gpu_sort_stably_overwrites_the_caller_output() {
let Some(context) = support::gpu_context().await else {
return;
};
let mut sorter = KeyValueSorter::from_context(&context);
let input = [
KeyValue::new(0, 41),
KeyValue::new(0, 7),
KeyValue::new(0, 99),
KeyValue::new(0, 3),
];
let input_buffer = create_sort_input(&context.device, &input);
let sentinels = [KeyValue::new(u32::MAX, u32::MAX); 4];
let output = context
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("Initialized Key-Value Sort Output"),
contents: bytemuck::cast_slice(&sentinels),
usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
});
sorter
.sort_gpu_to_gpu_with_key_bits(&input_buffer, &output, input.len() as u32, 0)
.expect("zero-bit key-value GPU sort failed");
assert_eq!(
support::read_pod::<KeyValue>(&context, &output, input.len()).await,
input
);
}
#[tokio::test]
async fn key_value_sort_rejects_aliased_input_and_output() {
let Some(context) = support::gpu_context().await else {
return;
};
let mut sorter = KeyValueSorter::from_context(&context);
let input = create_sort_input(&context.device, &[KeyValue::new(1, 0)]);
let error = sorter
.sort_gpu_to_gpu_with_key_bits(&input, &input, 1, 8)
.expect_err("in-place key-value scatter must be rejected");
assert!(matches!(error, Error::BufferAlias { .. }));
}
#[tokio::test]
async fn record_key_value_sort_composes_multiple_invocations_in_one_encoder() {
let Some(context) = support::gpu_context().await else {
return;
};
let mut sorter = KeyValueSorter::from_context(&context);
let first = [KeyValue::new(7, 11)];
let second = duplicate_keys(4_097, 200);
let first_input = create_sort_input(&context.device, &first);
let second_input = create_sort_input(&context.device, &second);
let first_output = create_sort_output(&context.device, first.len());
let second_output = create_sort_output(&context.device, second.len());
let mut encoder = context
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
sorter
.record_sort(
&mut encoder,
&first_input,
&first_output,
first.len() as u32,
)
.expect("first key-value sort recording failed");
sorter
.record_sort(
&mut encoder,
&second_input,
&second_output,
second.len() as u32,
)
.expect("second key-value sort recording failed");
context.queue.submit(Some(encoder.finish()));
assert_eq!(
support::read_pod::<KeyValue>(&context, &first_output, first.len()).await,
cpu_stable_sort(&first)
);
assert_eq!(
support::read_pod::<KeyValue>(&context, &second_output, second.len()).await,
cpu_stable_sort(&second)
);
}
#[tokio::test]
async fn record_key_value_sort_rejects_short_pair_buffers() {
let Some(context) = support::gpu_context().await else {
return;
};
let mut sorter = KeyValueSorter::from_context(&context);
let input = context.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("Short Key-Value Sort Input"),
size: 24,
usage: wgpu::BufferUsages::STORAGE,
mapped_at_creation: false,
});
let output = context.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("Key-Value Sort Output"),
size: 32,
usage: wgpu::BufferUsages::STORAGE,
mapped_at_creation: false,
});
let mut encoder = context
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
let error = sorter
.record_sort(&mut encoder, &input, &output, 4)
.expect_err("short key-value input must be rejected");
assert!(matches!(error, Error::BufferTooSmall { .. }));
}
fn create_sort_input(device: &wgpu::Device, input: &[KeyValue]) -> wgpu::Buffer {
device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("Key-Value Sort Input"),
contents: bytemuck::cast_slice(input),
usage: wgpu::BufferUsages::STORAGE,
})
}
fn create_sort_output(device: &wgpu::Device, len: usize) -> wgpu::Buffer {
device.create_buffer(&wgpu::BufferDescriptor {
label: Some("Key-Value Sort Output"),
size: (len * size_of::<KeyValue>()) as u64,
usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
mapped_at_creation: false,
})
}
fn create_soa_buffers(
device: &wgpu::Device,
keys: &[u32],
values: &[u32],
) -> (wgpu::Buffer, wgpu::Buffer) {
assert_eq!(keys.len(), values.len());
let usage = wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC;
(
device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("SoA Keys"),
contents: bytemuck::cast_slice(keys),
usage,
}),
device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("SoA Values"),
contents: bytemuck::cast_slice(values),
usage,
}),
)
}
async fn assert_soa_prefix(
context: &lampshade::Context,
key_buffer: &wgpu::Buffer,
value_buffer: &wgpu::Buffer,
keys: &[u32],
values: &[u32],
active: usize,
) {
let actual_keys = support::read_pod::<u32>(context, key_buffer, active).await;
let actual_values = support::read_pod::<u32>(context, value_buffer, active).await;
let mut expected: Vec<_> = keys[..active]
.iter()
.copied()
.zip(values[..active].iter().copied())
.collect();
expected.sort_by_key(|&(key, _)| key);
assert_eq!(
actual_keys
.into_iter()
.zip(actual_values)
.collect::<Vec<_>>(),
expected
);
}