use vyre_foundation::ir::{BufferAccess, BufferDecl, DataType, Expr, Node, Program};
use crate::region::wrap_anonymous;
use crate::scan::builders::{
append_match, append_match_subgroup, load_packed_byte, load_packed_byte_expr,
};
use vyre_primitives::matching::CompiledDfa;
#[cfg(any(test, feature = "cpu-parity"))]
use super::ClassicAcAutomaton;
mod prefilter;
#[cfg(all(feature = "matching-regex", feature = "matching-dfa"))]
mod regex_exact;
pub use prefilter::{
build_ac_bounded_ranges_prefilter_program, build_ac_bounded_ranges_prefilter_program_with_subgroup_coalesce,
build_ac_bounded_ranges_suffix3_prefilter_program,
build_ac_bounded_ranges_suffix3_prefilter_program_with_subgroup_coalesce,
classic_ac_bounded_ranges_prefilter_program, classic_ac_bounded_ranges_prefilter_program_with_subgroup_coalesce,
classic_ac_bounded_ranges_suffix3_prefilter_program,
classic_ac_bounded_ranges_suffix3_prefilter_program_with_subgroup_coalesce,
classic_ac_bounded_ranges_suffix3_presence_and_positions_by_region_program,
classic_ac_bounded_ranges_suffix3_presence_and_positions_by_region_program_filtered,
classic_ac_bounded_ranges_suffix3_presence_by_region_program,
classic_ac_bounded_ranges_suffix3_presence_program, presence_bitmap_words,
presence_by_region_words, try_build_ac_bounded_ranges_prefilter_program,
try_build_ac_bounded_ranges_prefilter_program_with_subgroup_coalesce,
try_build_ac_bounded_ranges_suffix3_prefilter_program,
try_build_ac_bounded_ranges_suffix3_prefilter_program_with_subgroup_coalesce,
try_build_ac_bounded_ranges_suffix3_presence_and_positions_by_region_program,
try_build_ac_bounded_ranges_suffix3_presence_and_positions_by_region_program_filtered,
try_build_ac_bounded_ranges_suffix3_presence_by_region_program,
try_build_ac_bounded_ranges_suffix3_presence_program,
};
#[cfg(all(feature = "matching-regex", feature = "matching-dfa"))]
pub(in crate::scan) use regex_exact::regex_exact_ranges_program;
pub(in crate::scan) fn ac_advance_state_node(transitions: &str, byte: Expr) -> Node {
Node::assign(
"state",
Expr::load(
transitions,
Expr::add(Expr::mul(Expr::var("state"), Expr::u32(256)), byte),
),
)
}
pub(in crate::scan) fn ac_transition_step_nodes(
haystack: &str,
transitions: &str,
idx: Expr,
) -> Vec<Node> {
let (load_byte, byte) = load_packed_byte(haystack, idx);
vec![load_byte, ac_advance_state_node(transitions, byte)]
}
pub(in crate::scan) fn ac_output_span_nodes(output_offsets: &str) -> Vec<Node> {
vec![
Node::let_bind("out_begin", Expr::load(output_offsets, Expr::var("state"))),
Node::let_bind(
"out_end",
Expr::load(output_offsets, Expr::add(Expr::var("state"), Expr::u32(1))),
),
]
}
pub(in crate::scan) fn bounded_walk_prologue_nodes(
haystack: &str,
transitions: &str,
output_offsets: &str,
max_pattern_len: u32,
) -> Vec<Node> {
let max_pattern_len = max_pattern_len.max(1);
let i = Expr::var("i");
let end = Expr::add(i.clone(), Expr::u32(1));
let scan_start = Expr::select(
Expr::lt(i, Expr::u32(max_pattern_len - 1)),
Expr::u32(0),
Expr::sub(end.clone(), Expr::u32(max_pattern_len)),
);
let mut nodes = vec![
Node::let_bind("state", Expr::u32(0)),
Node::let_bind("scan_start", scan_start),
Node::let_bind("scan_end", end),
Node::loop_for(
"step",
Expr::var("scan_start"),
Expr::var("scan_end"),
ac_transition_step_nodes(haystack, transitions, Expr::var("step")),
),
];
nodes.extend(ac_output_span_nodes(output_offsets));
nodes
}
pub(in crate::scan) fn candidate_end_gate_nodes(
haystack: &str,
haystack_len: &str,
candidate_end_mask: &str,
accepted: Vec<Node>,
) -> Vec<Node> {
let i = Expr::var("i");
vec![
Node::let_bind("i", Expr::InvocationId { axis: 0 }),
Node::if_then(
Expr::lt(i.clone(), Expr::load(haystack_len, Expr::u32(0))),
vec![
Node::let_bind("candidate_byte", load_packed_byte_expr(haystack, i)),
Node::let_bind(
"candidate_word",
Expr::load(
candidate_end_mask,
Expr::shr(Expr::var("candidate_byte"), Expr::u32(5)),
),
),
Node::let_bind(
"candidate_bit",
Expr::shl(
Expr::u32(1),
Expr::bitand(Expr::var("candidate_byte"), Expr::u32(31)),
),
),
Node::if_then(
Expr::ne(
Expr::bitand(Expr::var("candidate_word"), Expr::var("candidate_bit")),
Expr::u32(0),
),
accepted,
),
],
),
]
}
pub(in crate::scan) fn classic_ac_dfa_buffer_decls(
haystack: &str,
transitions: &str,
output_offsets: &str,
state_count: u32,
) -> Vec<BufferDecl> {
vec![
BufferDecl::storage(haystack, 0, BufferAccess::ReadOnly, DataType::U32),
BufferDecl::storage(transitions, 1, BufferAccess::ReadOnly, DataType::U32)
.with_count(state_count.saturating_mul(256)),
BufferDecl::storage(output_offsets, 2, BufferAccess::ReadOnly, DataType::U32)
.with_count(state_count.saturating_add(1)),
]
}
#[derive(Clone, Copy)]
pub(in crate::scan) struct AcInputBindings<'a> {
pub haystack: &'a str,
pub transitions: &'a str,
pub output_offsets: &'a str,
pub output_records: &'a str,
pub pattern_lengths: &'a str,
pub haystack_len: &'a str,
pub state_count: u32,
pub output_records_len: u32,
pub pattern_count: u32,
}
impl AcInputBindings<'_> {
pub(in crate::scan) fn decls(&self) -> Vec<BufferDecl> {
let mut decls = classic_ac_dfa_buffer_decls(
self.haystack,
self.transitions,
self.output_offsets,
self.state_count,
);
decls.reserve(3);
decls.extend([
BufferDecl::storage(self.output_records, 3, BufferAccess::ReadOnly, DataType::U32)
.with_count(self.output_records_len),
BufferDecl::storage(self.pattern_lengths, 4, BufferAccess::ReadOnly, DataType::U32)
.with_count(self.pattern_count),
BufferDecl::storage(self.haystack_len, 5, BufferAccess::ReadOnly, DataType::U32)
.with_count(1),
]);
decls
}
}
#[must_use]
#[allow(clippy::too_many_arguments)]
pub fn classic_ac_bounded_ranges_program(
haystack: &str,
transitions: &str,
output_offsets: &str,
output_records: &str,
pattern_lengths: &str,
haystack_len: &str,
match_count: &str,
matches: &str,
state_count: u32,
output_records_len: u32,
pattern_count: u32,
max_matches: u32,
max_pattern_len: u32,
) -> Program {
classic_ac_bounded_ranges_program_with_subgroup_coalesce(
haystack,
transitions,
output_offsets,
output_records,
pattern_lengths,
haystack_len,
match_count,
matches,
state_count,
output_records_len,
pattern_count,
max_matches,
max_pattern_len,
true,
)
}
#[must_use]
#[allow(clippy::too_many_arguments)]
pub fn classic_ac_bounded_ranges_program_with_subgroup_coalesce(
haystack: &str,
transitions: &str,
output_offsets: &str,
output_records: &str,
pattern_lengths: &str,
haystack_len: &str,
match_count: &str,
matches: &str,
state_count: u32,
output_records_len: u32,
pattern_count: u32,
max_matches: u32,
max_pattern_len: u32,
use_subgroup_coalesce: bool,
) -> Program {
let max_pattern_len = max_pattern_len.max(1);
let i = Expr::var("i");
let walk_body = vec![
Node::let_bind("i", Expr::InvocationId { axis: 0 }),
Node::if_then(
Expr::lt(i.clone(), Expr::load(haystack_len, Expr::u32(0))),
bounded_ranges_scan_nodes(
haystack,
transitions,
output_offsets,
output_records,
pattern_lengths,
match_count,
matches,
max_pattern_len,
use_subgroup_coalesce,
),
),
];
let mut buffers = AcInputBindings {
haystack,
transitions,
output_offsets,
output_records,
pattern_lengths,
haystack_len,
state_count,
output_records_len,
pattern_count,
}
.decls();
buffers.extend([
BufferDecl::read_write(match_count, 6, DataType::U32).with_count(1),
BufferDecl::output(matches, 7, DataType::U32).with_count(max_matches.saturating_mul(3)),
]);
Program::wrapped(
buffers,
[128, 1, 1],
vec![wrap_anonymous(
"vyre-libs::matching::classic_ac_bounded_ranges",
walk_body,
)],
)
}
#[allow(clippy::too_many_arguments)]
fn bounded_ranges_scan_nodes(
haystack: &str,
transitions: &str,
output_offsets: &str,
output_records: &str,
pattern_lengths: &str,
match_count: &str,
matches: &str,
max_pattern_len: u32,
use_subgroup_coalesce: bool,
) -> Vec<Node> {
let mut nodes =
bounded_walk_prologue_nodes(haystack, transitions, output_offsets, max_pattern_len);
nodes.push(Node::loop_for(
"out_idx",
Expr::var("out_begin"),
Expr::var("out_end"),
{
let mut body = vec![
Node::let_bind(
"pattern_id",
Expr::load(output_records, Expr::var("out_idx")),
),
Node::let_bind(
"pat_len",
Expr::load(pattern_lengths, Expr::var("pattern_id")),
),
Node::let_bind(
"match_start",
Expr::select(
Expr::lt(Expr::var("scan_end"), Expr::var("pat_len")),
Expr::u32(0),
Expr::sub(Expr::var("scan_end"), Expr::var("pat_len")),
),
),
];
if use_subgroup_coalesce {
body.extend(append_match_subgroup(
matches,
match_count,
Expr::var("pattern_id"),
Expr::var("match_start"),
Expr::var("scan_end"),
Expr::bool(true),
));
} else {
body.push(append_match(
matches,
match_count,
Expr::var("pattern_id"),
Expr::var("match_start"),
Expr::var("scan_end"),
));
}
body
},
));
nodes
}
fn bounded_ranges_presence_nodes(
haystack: &str,
transitions: &str,
output_offsets: &str,
output_records: &str,
presence: &str,
max_pattern_len: u32,
) -> Vec<Node> {
let mut nodes =
bounded_walk_prologue_nodes(haystack, transitions, output_offsets, max_pattern_len);
nodes.push(Node::loop_for(
"out_idx",
Expr::var("out_begin"),
Expr::var("out_end"),
vec![
Node::let_bind(
"pattern_id",
Expr::load(output_records, Expr::var("out_idx")),
),
Node::let_bind(
"_vyre_presence_prev",
Expr::atomic_or(
presence,
Expr::shr(Expr::var("pattern_id"), Expr::u32(5)),
Expr::shl(
Expr::u32(1),
Expr::bitand(Expr::var("pattern_id"), Expr::u32(31)),
),
),
),
],
));
nodes
}
#[allow(clippy::too_many_arguments)]
fn bounded_ranges_presence_by_region_nodes(
haystack: &str,
transitions: &str,
output_offsets: &str,
output_records: &str,
presence: &str,
region_starts: &str,
region_base: &str,
max_pattern_len: u32,
presence_words: u32,
log2_max_regions: u32,
) -> Vec<Node> {
let mut region_and_emit =
region_search_prologue_nodes(region_starts, region_base, presence_words, log2_max_regions);
region_and_emit.push(Node::loop_for(
"out_idx",
Expr::var("out_begin"),
Expr::var("out_end"),
vec![
Node::let_bind(
"pattern_id",
Expr::load(output_records, Expr::var("out_idx")),
),
Node::let_bind(
"_vyre_presence_prev",
Expr::atomic_or(
presence,
Expr::add(
Expr::var("rs_base"),
Expr::shr(Expr::var("pattern_id"), Expr::u32(5)),
),
Expr::shl(
Expr::u32(1),
Expr::bitand(Expr::var("pattern_id"), Expr::u32(31)),
),
),
),
],
));
let mut nodes =
bounded_walk_prologue_nodes(haystack, transitions, output_offsets, max_pattern_len);
nodes.push(Node::if_then(
Expr::lt(Expr::var("out_begin"), Expr::var("out_end")),
region_and_emit,
));
nodes
}
pub(in crate::scan) fn region_search_prologue_nodes(
region_starts: &str,
region_base: &str,
presence_words: u32,
log2_max_regions: u32,
) -> Vec<Node> {
vec![
Node::let_bind(
"rs_pos",
Expr::add(Expr::var("i"), Expr::load(region_base, Expr::u32(0))),
),
Node::let_bind("rs_lo", Expr::u32(0)),
Node::let_bind(
"rs_hi",
Expr::sub(Expr::buf_len(region_starts), Expr::u32(1)),
),
Node::loop_for(
"rs_step",
Expr::u32(0),
Expr::u32(log2_max_regions.max(1)),
vec![
Node::let_bind(
"rs_mid",
Expr::div(
Expr::add(
Expr::add(Expr::var("rs_lo"), Expr::var("rs_hi")),
Expr::u32(1),
),
Expr::u32(2),
),
),
Node::let_bind(
"rs_cond",
Expr::le(
Expr::load(region_starts, Expr::var("rs_mid")),
Expr::var("rs_pos"),
),
),
Node::assign(
"rs_lo",
Expr::select(
Expr::var("rs_cond"),
Expr::var("rs_mid"),
Expr::var("rs_lo"),
),
),
Node::assign(
"rs_hi",
Expr::select(
Expr::var("rs_cond"),
Expr::var("rs_hi"),
Expr::sub(Expr::var("rs_mid"), Expr::u32(1)),
),
),
],
),
Node::let_bind(
"rs_base",
Expr::mul(Expr::var("rs_lo"), Expr::u32(presence_words.max(1))),
),
]
}
#[allow(clippy::too_many_arguments)]
fn bounded_ranges_presence_and_positions_by_region_nodes(
haystack: &str,
transitions: &str,
output_offsets: &str,
output_records: &str,
pattern_lengths: &str,
presence: &str,
region_starts: &str,
region_base: &str,
match_count: &str,
matches: &str,
max_pattern_len: u32,
presence_words: u32,
log2_max_regions: u32,
first_positioned_pattern_id: u32,
) -> Vec<Node> {
let mut region_and_emit =
region_search_prologue_nodes(region_starts, region_base, presence_words, log2_max_regions);
region_and_emit.push(Node::loop_for(
"out_idx",
Expr::var("out_begin"),
Expr::var("out_end"),
vec![
Node::let_bind(
"pattern_id",
Expr::load(output_records, Expr::var("out_idx")),
),
Node::let_bind(
"_vyre_presence_prev",
Expr::atomic_or(
presence,
Expr::add(
Expr::var("rs_base"),
Expr::shr(Expr::var("pattern_id"), Expr::u32(5)),
),
Expr::shl(
Expr::u32(1),
Expr::bitand(Expr::var("pattern_id"), Expr::u32(31)),
),
),
),
Node::if_then(
Expr::ge(
Expr::var("pattern_id"),
Expr::u32(first_positioned_pattern_id),
),
vec![
Node::let_bind(
"pat_len",
Expr::load(pattern_lengths, Expr::var("pattern_id")),
),
Node::let_bind(
"match_start",
Expr::select(
Expr::lt(Expr::var("scan_end"), Expr::var("pat_len")),
Expr::u32(0),
Expr::sub(Expr::var("scan_end"), Expr::var("pat_len")),
),
),
append_match(
matches,
match_count,
Expr::var("pattern_id"),
Expr::var("match_start"),
Expr::var("scan_end"),
),
],
),
],
));
let mut nodes =
bounded_walk_prologue_nodes(haystack, transitions, output_offsets, max_pattern_len);
nodes.push(Node::if_then(
Expr::lt(Expr::var("out_begin"), Expr::var("out_end")),
region_and_emit,
));
nodes
}
#[must_use]
pub fn build_ac_bounded_ranges_program(
dfa: &CompiledDfa,
pattern_count: u32,
max_matches: u32,
) -> Program {
build_ac_bounded_ranges_program_with_subgroup_coalesce(dfa, pattern_count, max_matches, true)
}
#[must_use]
pub fn build_ac_bounded_ranges_program_with_subgroup_coalesce(
dfa: &CompiledDfa,
pattern_count: u32,
max_matches: u32,
use_subgroup_coalesce: bool,
) -> Program {
match try_build_ac_bounded_ranges_program_with_subgroup_coalesce(
dfa,
pattern_count,
max_matches,
use_subgroup_coalesce,
) {
Ok(program) => program,
Err(error) => {
panic!(
"AC bounded-ranges program build failed: {error}. \
returning an empty rejecting automaton would silently drop every match; \
use try_build_ac_bounded_ranges_program_with_subgroup_coalesce and shard oversized DFAs."
)
}
}
}
pub fn try_build_ac_bounded_ranges_program(
dfa: &CompiledDfa,
pattern_count: u32,
max_matches: u32,
) -> Result<Program, String> {
try_build_ac_bounded_ranges_program_with_subgroup_coalesce(dfa, pattern_count, max_matches, true)
}
pub fn try_build_ac_bounded_ranges_program_with_subgroup_coalesce(
dfa: &CompiledDfa,
pattern_count: u32,
max_matches: u32,
use_subgroup_coalesce: bool,
) -> Result<Program, String> {
let output_records_len = u32::try_from(dfa.output_records.len()).map_err(|source| {
format!(
"AC bounded-ranges DFA output record count {} exceeds u32 GPU buffer metadata: {source}. Fix: shard the pattern set or lower the DFA budget before dispatch.",
dfa.output_records.len()
)
})?;
Ok(classic_ac_bounded_ranges_program_with_subgroup_coalesce(
"haystack",
"transitions",
"output_offsets",
"output_records",
"pattern_lengths",
"haystack_len",
"match_count",
"matches",
dfa.state_count,
output_records_len,
pattern_count,
max_matches,
dfa.max_pattern_len,
use_subgroup_coalesce,
))
}
#[must_use]
#[cfg(any(test, feature = "cpu-parity"))]
pub fn classic_ac_bounded_ranges_scan(
ac: &ClassicAcAutomaton,
pattern_lengths: &[u32],
haystack: &[u8],
) -> Vec<(u32, u32, u32)> {
let dfa = &ac.dfa;
let mut state = 0u32;
let mut out = Vec::new();
for (pos, &b) in haystack.iter().enumerate() {
state = dfa.transitions[(state as usize) * 256 + (b as usize)];
let begin = dfa.output_offsets[state as usize] as usize;
let end_off = dfa.output_offsets[state as usize + 1] as usize;
for &pid in &dfa.output_records[begin..end_off] {
let pat_len = pattern_lengths[pid as usize];
let end_pos = (pos as u32).saturating_add(1);
let start = end_pos.saturating_sub(pat_len);
out.push((pid, start, end_pos));
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
use crate::scan::classic_ac::classic_ac_compile;
#[test]
fn infallible_builder_uses_real_dfa_not_empty_fallback() {
let ac = classic_ac_compile(&[b"abc", b"de", b"abcd"]);
let via_infallible = build_ac_bounded_ranges_program_with_subgroup_coalesce(&ac.dfa, 3, 128, false);
let via_try = try_build_ac_bounded_ranges_program_with_subgroup_coalesce(&ac.dfa, 3, 128, false)
.expect("valid DFA must build");
let records = via_infallible.buffers()[3].count;
assert_eq!(records as usize, ac.dfa.output_records.len());
assert!(
records > 0,
"infallible builder must not emit the empty-records fallback program"
);
assert_eq!(
via_infallible.buffers()[1].count,
ac.dfa.state_count.saturating_mul(256)
);
assert_eq!(via_infallible.buffers().len(), via_try.buffers().len());
assert_eq!(
via_infallible.buffers()[3].count,
via_try.buffers()[3].count
);
}
#[test]
fn try_build_ac_bounded_ranges_program_ext_succeeds_for_valid_dfa() {
let ac = classic_ac_compile(&[b"abc", b"de"]);
let result = try_build_ac_bounded_ranges_program_with_subgroup_coalesce(&ac.dfa, 2, 128, false);
assert!(
result.is_ok(),
"try_build must succeed for a valid small DFA: {:?}",
result.err()
);
let program = result.unwrap();
assert_eq!(
program.workgroup_size(),
[128, 1, 1],
"workgroup size must be [128, 1, 1]"
);
}
#[test]
#[should_panic]
fn classic_ac_bounded_ranges_scan_panics_on_oob_pid() {
use vyre_primitives::matching::CompiledDfa;
let transitions: Vec<u32> = {
let mut t = vec![0u32; 2 * 256]; t[0 * 256 + b'A' as usize] = 1; t
};
let accept = vec![0u32, 6u32]; let output_offsets = vec![0u32, 0u32, 1u32]; let output_records = vec![5u32];
let dfa = CompiledDfa {
transitions,
accept,
state_count: 2,
max_pattern_len: 1,
output_offsets,
output_records,
};
let ac = crate::scan::classic_ac::ClassicAcAutomaton { dfa };
let _result = classic_ac_bounded_ranges_scan(&ac, &[1u32, 2u32, 3u32], b"A");
}
}