#![forbid(unsafe_code)]
#![feature(portable_simd)]
use std::simd::cmp::{SimdPartialEq, SimdPartialOrd};
use std::simd::{Mask, Simd, u8x64, u64x8};
const CELLS_PER_CHUNK: usize = 4;
const BYTES_PER_CHUNK: usize = 64;
#[must_use]
pub fn first_mismatch_u128(a: &[u128], b: &[u128]) -> Option<usize> {
let len = a.len().min(b.len());
let mut offset = 0;
while offset + CELLS_PER_CHUNK <= len {
let lhs = load_cells(&a[offset..offset + CELLS_PER_CHUNK]);
let rhs = load_cells(&b[offset..offset + CELLS_PER_CHUNK]);
let differing: Mask<i64, 8> = lhs.simd_ne(rhs);
if differing.any() {
let lane = differing.first_set().unwrap_or(0);
return Some(offset + lane / 2);
}
offset += CELLS_PER_CHUNK;
}
(offset..len).find(|&i| a[i] != b[i])
}
#[must_use]
pub fn first_mismatch_u128_scalar(a: &[u128], b: &[u128]) -> Option<usize> {
let len = a.len().min(b.len());
(0..len).find(|&i| a[i] != b[i])
}
#[must_use]
pub fn rows_equal_u128(a: &[u128], b: &[u128]) -> bool {
a.len() == b.len() && first_mismatch_u128(a, b).is_none()
}
#[must_use]
pub fn rows_equal_u128_scalar(a: &[u128], b: &[u128]) -> bool {
a.len() == b.len() && first_mismatch_u128_scalar(a, b).is_none()
}
#[must_use]
pub fn all_ascii(bytes: &[u8]) -> bool {
let mut offset = 0;
while offset + BYTES_PER_CHUNK <= bytes.len() {
let chunk = u8x64::from_slice(&bytes[offset..offset + BYTES_PER_CHUNK]);
if (chunk & u8x64::splat(0x80)).simd_ne(u8x64::splat(0)).any() {
return false;
}
offset += BYTES_PER_CHUNK;
}
bytes[offset..].iter().all(u8::is_ascii)
}
#[must_use]
pub fn all_ascii_scalar(bytes: &[u8]) -> bool {
bytes.iter().all(u8::is_ascii)
}
#[must_use]
pub fn ascii_width(bytes: &[u8]) -> Option<usize> {
let mut offset = 0;
while offset + BYTES_PER_CHUNK <= bytes.len() {
let chunk = u8x64::from_slice(&bytes[offset..offset + BYTES_PER_CHUNK]);
let shifted = chunk - u8x64::splat(0x20);
if shifted.simd_gt(u8x64::splat(0x5e)).any() {
return None;
}
offset += BYTES_PER_CHUNK;
}
if bytes[offset..].iter().all(|b| (0x20..=0x7e).contains(b)) {
Some(bytes.len())
} else {
None
}
}
#[must_use]
pub fn ascii_width_scalar(bytes: &[u8]) -> Option<usize> {
if bytes.iter().all(|b| (0x20..=0x7e).contains(b)) {
Some(bytes.len())
} else {
None
}
}
#[inline]
fn load_cells(cells: &[u128]) -> u64x8 {
debug_assert_eq!(cells.len(), CELLS_PER_CHUNK);
let mut lanes = [0_u64; 8];
for (i, cell) in cells.iter().enumerate() {
lanes[i * 2] = (*cell & u128::from(u64::MAX)) as u64;
lanes[i * 2 + 1] = (*cell >> 64) as u64;
}
Simd::from_array(lanes)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn first_mismatch_reports_none_for_identical_rows() {
let row: Vec<u128> = (0..200).collect();
assert_eq!(first_mismatch_u128(&row, &row), None);
assert_eq!(first_mismatch_u128_scalar(&row, &row), None);
}
#[test]
fn first_mismatch_finds_the_earliest_difference_at_every_position() {
for len in [1_usize, 3, 4, 5, 8, 63, 64, 65, 200] {
let base: Vec<u128> = (0..len as u128).collect();
for idx in 0..len {
let mut changed = base.clone();
changed[idx] = u128::MAX;
assert_eq!(
first_mismatch_u128(&base, &changed),
Some(idx),
"len {len}, idx {idx}"
);
assert_eq!(
first_mismatch_u128_scalar(&base, &changed),
Some(idx),
"scalar len {len}, idx {idx}"
);
}
}
}
#[test]
fn first_mismatch_detects_a_difference_in_either_half_of_a_cell() {
for delta in [1_u128, 1_u128 << 64, (1_u128 << 64) | 1] {
let base = vec![0_u128; 8];
let mut changed = base.clone();
changed[5] = delta;
assert_eq!(
first_mismatch_u128(&base, &changed),
Some(5),
"delta {delta}"
);
}
}
#[test]
fn first_mismatch_stops_at_the_shorter_slice() {
let long: Vec<u128> = (0..16).collect();
let short = &long[..6];
assert_eq!(first_mismatch_u128(&long, short), None);
assert_eq!(first_mismatch_u128(short, &long), None);
}
#[test]
fn first_mismatch_handles_empty_input() {
assert_eq!(first_mismatch_u128(&[], &[]), None);
assert_eq!(first_mismatch_u128(&[], &[1, 2]), None);
}
#[test]
fn rows_equal_requires_matching_length() {
assert!(rows_equal_u128(&[1, 2, 3], &[1, 2, 3]));
assert!(!rows_equal_u128(&[1, 2, 3], &[1, 2]));
assert!(!rows_equal_u128(&[1, 2, 3], &[1, 2, 4]));
assert!(rows_equal_u128(&[], &[]));
}
#[test]
fn all_ascii_agrees_with_the_scalar_twin_across_lengths() {
for len in [0_usize, 1, 63, 64, 65, 127, 128, 1000] {
let ascii = vec![b'a'; len];
assert!(all_ascii(&ascii), "len {len}");
assert_eq!(all_ascii(&ascii), all_ascii_scalar(&ascii));
if len > 0 {
for idx in [0, len / 2, len - 1] {
let mut probe = ascii.clone();
probe[idx] = 0xC3;
assert!(!all_ascii(&probe), "len {len}, idx {idx}");
assert_eq!(all_ascii(&probe), all_ascii_scalar(&probe));
}
}
}
}
#[test]
fn all_ascii_accepts_control_bytes() {
assert!(all_ascii(b"\x00\x09\x1b\x7f"));
}
#[test]
fn ascii_width_counts_printable_runs_and_rejects_the_rest() {
assert_eq!(ascii_width(b""), Some(0));
assert_eq!(ascii_width(b" "), Some(1));
assert_eq!(ascii_width(b"~"), Some(1));
assert_eq!(ascii_width(b"hello world"), Some(11));
let long = vec![b'x'; 300];
assert_eq!(ascii_width(&long), Some(300));
assert_eq!(ascii_width(b"\x1f"), None);
assert_eq!(ascii_width(b"\x7f"), None);
assert_eq!(ascii_width(b"\xc3\xa9"), None);
}
#[test]
fn ascii_width_rejects_a_bad_byte_anywhere_including_past_a_full_chunk() {
for len in [64_usize, 65, 129, 300] {
for idx in [0, len / 2, len - 1] {
for bad in [0x00_u8, 0x1f, 0x7f, 0x80, 0xff] {
let mut probe = vec![b'a'; len];
probe[idx] = bad;
assert_eq!(ascii_width(&probe), None, "len {len}, idx {idx}, bad {bad}");
assert_eq!(ascii_width(&probe), ascii_width_scalar(&probe));
}
}
}
}
#[test]
fn kernels_agree_with_their_twins_over_a_deterministic_sweep() {
let mut state = 0x2545_F491_4F6C_DD1D_u64;
let mut next = move || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
state
};
for len in 0..200_usize {
let a: Vec<u128> = (0..len)
.map(|_| (u128::from(next()) << 64) | u128::from(next()))
.collect();
let mut b = a.clone();
if len > 0 {
let idx = (next() as usize) % len;
if next() % 2 == 0 {
b[idx] ^= 1 << (next() % 128);
}
}
assert_eq!(
first_mismatch_u128(&a, &b),
first_mismatch_u128_scalar(&a, &b),
"len {len}"
);
assert_eq!(rows_equal_u128(&a, &b), rows_equal_u128_scalar(&a, &b));
let bytes: Vec<u8> = (0..len).map(|_| (next() % 256) as u8).collect();
assert_eq!(all_ascii(&bytes), all_ascii_scalar(&bytes), "len {len}");
assert_eq!(ascii_width(&bytes), ascii_width_scalar(&bytes), "len {len}");
}
}
#[test]
fn ascii_kernels_agree_with_their_twins_on_runs_that_reach_the_chunk_loop() {
for len in 0..=(2 * BYTES_PER_CHUNK + 5) {
let mut run: Vec<u8> = (0..len).map(|i| 0x20 + (i % 0x5f) as u8).collect();
assert!(all_ascii(&run), "len {len}");
assert_eq!(all_ascii(&run), all_ascii_scalar(&run), "len {len}");
assert_eq!(ascii_width(&run), Some(len), "len {len}");
assert_eq!(ascii_width(&run), ascii_width_scalar(&run), "len {len}");
for pos in 0..len {
for bad in [0x00_u8, 0x1f, 0x7f, 0x80, 0xff] {
let saved = run[pos];
run[pos] = bad;
assert_eq!(
all_ascii(&run),
all_ascii_scalar(&run),
"all_ascii len {len} pos {pos} byte {bad:#04x}"
);
assert_eq!(
ascii_width(&run),
ascii_width_scalar(&run),
"ascii_width len {len} pos {pos} byte {bad:#04x}"
);
run[pos] = saved;
}
}
}
}
}