use vyre_foundation::ir::{DataType, Program};
use crate::graph::csr_frontier_step::{
csr_queue_step_program, CsrQueueEmit, CsrQueueInputs, CsrQueueLanes, CsrQueueRowPlan,
CsrQueueStepSpec,
};
mod strided;
pub use strided::{
csr_queue_delta_strided_dispatch_grid, csr_queue_delta_strided_enqueue,
csr_queue_delta_strided_enqueue_with, csr_queue_delta_strided_logical_lanes_per_launch,
csr_queue_delta_strided_source_slots_per_launch,
CSR_QUEUE_DELTA_STRIDED_CAPPED_LAUNCH_MIN_CAPACITY, CSR_QUEUE_DELTA_STRIDED_ENQUEUE_OP_ID,
CSR_QUEUE_DELTA_STRIDED_LANES_PER_SOURCE, CSR_QUEUE_DELTA_STRIDED_MAX_SOURCE_SLOTS_PER_LAUNCH,
};
pub const CSR_QUEUE_DELTA_ENQUEUE_OP_ID: &str = "vyre-primitives::graph::csr_queue_delta_enqueue";
pub const CSR_QUEUE_DELTA_ENQUEUE_WORKGROUP_SIZE: [u32; 3] = [256, 1, 1];
#[derive(Clone, Copy, Debug)]
pub struct CsrQueueDeltaEnqueueParams<'a> {
pub active_queue: &'a str,
pub active_len: &'a str,
pub edge_offsets: &'a str,
pub edge_targets: &'a str,
pub edge_kind_mask: &'a str,
pub accumulator: &'a str,
pub next_queue: &'a str,
pub next_len: &'a str,
pub node_count: u32,
pub edge_count: u32,
pub active_queue_capacity: u32,
pub next_queue_capacity: u32,
pub allow_mask: u32,
}
#[must_use]
#[allow(clippy::too_many_arguments)]
pub fn csr_queue_delta_enqueue(
active_queue: &str,
active_len: &str,
edge_offsets: &str,
edge_targets: &str,
edge_kind_mask: &str,
accumulator: &str,
next_queue: &str,
next_len: &str,
node_count: u32,
edge_count: u32,
active_queue_capacity: u32,
next_queue_capacity: u32,
allow_mask: u32,
) -> Program {
csr_queue_delta_enqueue_with(CsrQueueDeltaEnqueueParams {
active_queue,
active_len,
edge_offsets,
edge_targets,
edge_kind_mask,
accumulator,
next_queue,
next_len,
node_count,
edge_count,
active_queue_capacity,
next_queue_capacity,
allow_mask,
})
}
#[must_use]
pub fn csr_queue_delta_enqueue_with(params: CsrQueueDeltaEnqueueParams<'_>) -> Program {
let node_count = params.node_count;
let active_queue_capacity = params.active_queue_capacity;
let next_queue_capacity = params.next_queue_capacity;
if node_count == 0 || active_queue_capacity == 0 || next_queue_capacity == 0 {
return crate::invalid_output_program(CSR_QUEUE_DELTA_ENQUEUE_OP_ID,
params.next_len,
DataType::U32,
format!(
"Fix: csr_queue_delta_enqueue requires node_count > 0 and non-zero queue capacities, got node_count={node_count} active_queue_capacity={active_queue_capacity} next_queue_capacity={next_queue_capacity}."
),);
}
csr_queue_step_program(¶ms.spec(
CSR_QUEUE_DELTA_ENQUEUE_OP_ID,
"csr_queue_delta_enqueue",
"qd",
CsrQueueLanes::Scalar,
))
}
impl<'a> CsrQueueDeltaEnqueueParams<'a> {
fn spec(
&self,
op_id: &'static str,
builder_name: &'static str,
prefix: &'a str,
lanes: CsrQueueLanes,
) -> CsrQueueStepSpec<'a> {
CsrQueueStepSpec {
op_id,
builder_name,
prefix,
workgroup_size: CSR_QUEUE_DELTA_ENQUEUE_WORKGROUP_SIZE,
inputs: CsrQueueInputs {
active_queue: self.active_queue,
queue_len: self.active_len,
edge_offsets: self.edge_offsets,
edge_targets: self.edge_targets,
edge_kind_mask: self.edge_kind_mask,
},
lanes,
row_plan: CsrQueueRowPlan::ExpandAll,
emit: CsrQueueEmit::Delta {
accumulator: self.accumulator,
next_queue: self.next_queue,
next_len: self.next_len,
next_queue_capacity: self.next_queue_capacity,
},
node_count: self.node_count,
edge_count: self.edge_count,
queue_capacity: self.active_queue_capacity,
allow_mask: self.allow_mask,
}
}
}
#[must_use]
#[cfg(any(test, feature = "cpu-parity"))]
#[allow(clippy::too_many_arguments)]
pub fn csr_queue_delta_enqueue_cpu(
active_queue: &[u32],
active_len: u32,
edge_offsets: &[u32],
edge_targets: &[u32],
edge_kind_mask: &[u32],
accumulator: &[u32],
node_count: u32,
next_queue_capacity: usize,
allow_mask: u32,
) -> (Vec<u32>, Vec<u32>, u32) {
let mut accumulator = accumulator.to_vec();
let mut next_queue = Vec::new();
let next_len = try_csr_queue_delta_enqueue_cpu_into(
active_queue,
active_len,
edge_offsets,
edge_targets,
edge_kind_mask,
&mut accumulator,
node_count,
next_queue_capacity,
allow_mask,
&mut next_queue,
)
.unwrap_or_else(|err| {
panic!("csr_queue_delta_enqueue CPU oracle received malformed input. {err}")
});
(accumulator, next_queue, next_len)
}
#[cfg(any(test, feature = "cpu-parity"))]
#[allow(clippy::too_many_arguments)]
pub fn try_csr_queue_delta_enqueue_cpu_into(
active_queue: &[u32],
active_len: u32,
edge_offsets: &[u32],
edge_targets: &[u32],
edge_kind_mask: &[u32],
accumulator: &mut Vec<u32>,
node_count: u32,
next_queue_capacity: usize,
allow_mask: u32,
next_queue: &mut Vec<u32>,
) -> Result<u32, String> {
let layout = super::csr_frontier_queue::validate_csr_queue_graph(
node_count,
edge_offsets,
edge_targets,
edge_kind_mask,
)?;
if accumulator.len() != layout.words {
return Err(format!(
"Fix: csr_queue_delta_enqueue requires accumulator.len() == bitset_words(node_count), got len={} but expected {} for node_count={node_count}.",
accumulator.len(),
layout.words
));
}
crate::graph::scratch::reserve_graph_items(
next_queue,
next_queue_capacity,
"CSR queue delta CPU oracle",
"next active frontier queue",
)?;
let mut next_tmp = Vec::with_capacity(next_queue_capacity);
let mut accumulator_tmp = accumulator.clone();
let take = (active_len as usize).min(active_queue.len());
let mut next_seen = 0_u32;
for &src in &active_queue[..take] {
if src >= node_count {
continue;
}
let start = edge_offsets[src as usize] as usize;
let end = edge_offsets[src as usize + 1] as usize;
for edge in start..end {
if edge_kind_mask[edge] & allow_mask == 0 {
continue;
}
let dst = edge_targets[edge];
let word = dst as usize / 32;
let bit = 1_u32 << (dst % 32);
let old = accumulator_tmp[word];
if old & bit != 0 {
continue;
}
accumulator_tmp[word] = old | bit;
if next_tmp.len() < next_queue_capacity {
next_tmp.push(dst);
}
next_seen = next_seen.saturating_add(1);
}
}
*accumulator = accumulator_tmp;
next_queue.clear();
next_queue.extend_from_slice(&next_tmp);
Ok(next_seen)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn emitted_program_has_stable_delta_queue_shape() {
let program = csr_queue_delta_enqueue(
"active_queue",
"active_len",
"edge_offsets",
"edge_targets",
"edge_kind_mask",
"accumulator",
"next_queue",
"next_len",
64,
7,
8,
16,
1,
);
assert_eq!(
program.workgroup_size,
CSR_QUEUE_DELTA_ENQUEUE_WORKGROUP_SIZE
);
assert_eq!(program.buffers.len(), 8);
}
#[test]
fn delta_enqueue_rejects_offset_count_overflow_without_panic() {
let result = std::panic::catch_unwind(|| {
csr_queue_delta_enqueue(
"active_queue",
"active_len",
"edge_offsets",
"edge_targets",
"edge_kind_mask",
"accumulator",
"next_queue",
"next_len",
u32::MAX,
0,
1,
1,
1,
)
});
assert!(
result.is_ok(),
"CSR queue delta builder must reject offset-count overflow without panicking"
);
let program = result.unwrap();
assert!(program.stats().trap());
let entry = format!("{:?}", program.entry());
assert!(
entry.contains("node_count + 1 overflows u32"),
"Fix: trap must retain the CSR offset-count overflow diagnostic, got: {entry}"
);
}
#[test]
fn cpu_delta_enqueue_only_emits_first_time_discoveries() {
let edge_offsets = [0, 3, 4, 4, 4, 4];
let edge_targets = [1, 2, 3, 4];
let edge_kind_mask = [1, 1, 2, 1];
let accumulator = vec![0b00001];
let (accumulator, next_queue, next_len) = csr_queue_delta_enqueue_cpu(
&[0, 1],
2,
&edge_offsets,
&edge_targets,
&edge_kind_mask,
&accumulator,
5,
8,
1,
);
assert_eq!(accumulator, vec![0b10111]);
assert_eq!(next_queue, vec![1, 2, 4]);
assert_eq!(next_len, 3);
}
#[test]
fn cpu_delta_enqueue_reports_queue_pressure_without_clobbering_accumulator() {
let edge_offsets = [0, 3, 3, 3, 3];
let edge_targets = [1, 2, 3];
let edge_kind_mask = [1, 1, 1];
let mut accumulator = vec![0b0001];
let mut next_queue = Vec::new();
let next_len = try_csr_queue_delta_enqueue_cpu_into(
&[0],
1,
&edge_offsets,
&edge_targets,
&edge_kind_mask,
&mut accumulator,
4,
2,
1,
&mut next_queue,
)
.expect("Fix: canonical queue delta graph should enqueue bounded discoveries");
assert_eq!(accumulator, vec![0b1111]);
assert_eq!(next_queue, vec![1, 2]);
assert_eq!(next_len, 3);
}
#[test]
fn cpu_delta_enqueue_rejects_bad_accumulator_without_clobbering_outputs() {
let mut accumulator = vec![0xCAFE_BABE, 0xDEAD_BEEF];
let mut next_queue = vec![9, 8, 7];
let err = try_csr_queue_delta_enqueue_cpu_into(
&[0],
1,
&[0, 1],
&[0],
&[1],
&mut accumulator,
1,
4,
1,
&mut next_queue,
)
.expect_err("wrong accumulator width must fail before mutation");
assert!(err.contains("accumulator.len() == bitset_words(node_count)"));
assert_eq!(accumulator, vec![0xCAFE_BABE, 0xDEAD_BEEF]);
assert_eq!(next_queue, vec![9, 8, 7]);
}
}