use memchr::memmem::FinderRev;
use simdeez::{prelude::*, simd_runtime_generate};
use std::sync::OnceLock;
const UTF_8_CONTINUATION_PATTERN: i8 = 0b1000_0000_u8 as i8;
const NON_UTF_8_CONTINUATION_PATTERN: i8 = 0b1100_0000_u8 as i8;
static REV_LINE_FINDER: OnceLock<FinderRev> = OnceLock::new();
simd_runtime_generate!(
pub fn is_ascii_simd(text: &str) -> bool {
let bytes = text.as_bytes();
let len = bytes.len();
let bytes_i8 = unsafe { std::slice::from_raw_parts(bytes.as_ptr().cast::<i8>(), len) };
let mut remainder = bytes_i8;
while remainder.len() >= S::Vi8::WIDTH {
let chunk = &remainder[..S::Vi8::WIDTH];
let v = S::Vi8::load_from_slice(chunk);
let mask = v.cmp_lt(S::Vi8::set1(0));
if mask.get_mask() != 0 {
return false;
}
remainder = &remainder[S::Vi8::WIDTH..];
}
remainder.iter().all(|&b| b >= 0)
}
);
simd_runtime_generate!(
fn find_last_byte_simd(haystack: &[u8], needle: u8, is_eol: bool) -> Option<usize> {
if haystack.is_empty() {
return None;
}
if is_eol {
let line_finder =
REV_LINE_FINDER.get_or_init(|| FinderRev::new(&[needle]).into_owned());
return line_finder.rfind(haystack);
}
let bound_needle = &[needle];
let finder = FinderRev::new(bound_needle);
finder.rfind(haystack)
}
);
simd_runtime_generate!(
fn count_utf8_chars_simd(bytes: &[u8]) -> usize {
let len = bytes.len();
if len == 0 {
return 0;
}
let bytes_i8 = unsafe { std::slice::from_raw_parts(bytes.as_ptr().cast::<i8>(), len) };
let mut remainder = bytes_i8;
let mut char_count = 0;
let continuation_pattern = S::Vi8::set1(UTF_8_CONTINUATION_PATTERN);
let mask_pattern = S::Vi8::set1(NON_UTF_8_CONTINUATION_PATTERN);
while remainder.len() >= S::Vi8::WIDTH {
let chunk = &remainder[..S::Vi8::WIDTH];
let v = S::Vi8::load_from_slice(chunk);
let masked = v & mask_pattern;
let is_continuation = masked.cmp_eq(continuation_pattern);
let mask = is_continuation.get_mask();
char_count += S::Vi8::WIDTH - mask.count_ones() as usize;
remainder = &remainder[S::Vi8::WIDTH..];
}
for &byte in remainder {
if (byte as u8) & NON_UTF_8_CONTINUATION_PATTERN as u8
!= UTF_8_CONTINUATION_PATTERN as u8
{
char_count += 1;
}
}
char_count
}
);
#[inline]
#[must_use]
pub fn get_char_column_simd(text: &str, offset: usize) -> usize {
if offset == 0 {
return 0;
}
let bytes = text.as_bytes();
if offset > bytes.len() {
return 0;
}
let search_slice = &bytes[..offset];
if is_ascii_simd(text) {
match find_last_byte_simd(search_slice, b'\n', true) {
Some(newline_pos) => offset - newline_pos - 1,
None => offset, }
} else {
match find_last_byte_simd(search_slice, b'\n', true) {
Some(newline_pos) => {
let line_start = newline_pos + 1;
let line_bytes = &search_slice[line_start..];
count_utf8_chars_simd(line_bytes)
}
None => {
count_utf8_chars_simd(search_slice)
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_empty_string() {
assert!(is_ascii_simd(""));
}
#[test]
fn test_pure_ascii() {
assert!(is_ascii_simd("Hello, World!"));
assert!(is_ascii_simd("123456789"));
assert!(is_ascii_simd("ABCDEFGHIJKLMNOPQRSTUVWXYZ"));
assert!(is_ascii_simd("abcdefghijklmnopqrstuvwxyz"));
assert!(is_ascii_simd("!@#$%^&*()_+-=[]{}|;':\",./<>?"));
}
#[test]
fn test_ascii_with_newlines_and_tabs() {
assert!(is_ascii_simd("Hello\nWorld\t!"));
assert!(is_ascii_simd("\t\n\r"));
}
#[test]
fn test_ascii_control_characters() {
assert!(is_ascii_simd("\x00\x01\x02\x03\x04\x05\x06\x07"));
assert!(is_ascii_simd("\x08\x09\x0A\x0B\x0C\x0D\x0E\x0F"));
assert!(is_ascii_simd("\x10\x11\x12\x13\x14\x15\x16\x17"));
assert!(is_ascii_simd("\x18\x19\x1A\x1B\x1C\x1D\x1E\x1F"));
assert!(is_ascii_simd("\x7F")); }
#[test]
fn test_non_ascii_characters() {
assert!(!is_ascii_simd("café")); assert!(!is_ascii_simd("naïve")); assert!(!is_ascii_simd("résumé")); assert!(!is_ascii_simd("🚀")); assert!(!is_ascii_simd("こんにちは")); assert!(!is_ascii_simd("Привет")); assert!(!is_ascii_simd("مرحبا")); assert!(!is_ascii_simd(
"café مرحبا こんにちは 🚀 Привет résumé naïve"
));
}
#[test]
fn test_mixed_ascii_non_ascii() {
assert!(!is_ascii_simd("Hello café"));
assert!(!is_ascii_simd("ASCII and 🚀"));
assert!(!is_ascii_simd("test\u{200B}")); }
#[test]
fn test_long_ascii_strings() {
let long_ascii = "a".repeat(1000);
assert!(is_ascii_simd(&long_ascii));
let long_ascii_mixed = "ABC123!@#".repeat(100);
assert!(is_ascii_simd(&long_ascii_mixed));
}
#[test]
fn test_long_non_ascii_strings() {
let long_non_ascii = "café".repeat(100);
assert!(!is_ascii_simd(&long_non_ascii));
}
#[test]
fn test_ascii_boundary_values() {
assert!(is_ascii_simd("\x00")); assert!(is_ascii_simd("\x7F"));
assert!(!is_ascii_simd("ü")); assert!(!is_ascii_simd("€")); }
#[test]
fn test_various_lengths() {
for i in 1..=100 {
let ascii_string = "a".repeat(i);
assert!(is_ascii_simd(&ascii_string), "Failed for length {i}");
}
}
#[test]
fn test_non_ascii_at_different_positions() {
assert!(!is_ascii_simd("éabc"));
assert!(!is_ascii_simd("abéc"));
assert!(!is_ascii_simd("abcé"));
assert!(!is_ascii_simd("éabcé"));
}
#[test]
fn test_consistency_with_str_is_ascii() {
let test_strings = vec![
"",
"Hello",
"café",
"🚀",
"ASCII123!@#",
"test\u{200B}",
"\x00\x7F",
];
for test_str in &test_strings {
assert_eq!(
is_ascii_simd(test_str),
test_str.is_ascii(),
"Mismatch for string: {test_str:?}"
);
}
let long_string = "a".repeat(1000);
assert_eq!(
is_ascii_simd(&long_string),
long_string.is_ascii(),
"Mismatch for long string"
);
let non_ascii_chars = ["ü", "€", "漢", "🎉"];
for ch in &non_ascii_chars {
assert_eq!(
is_ascii_simd(ch),
ch.is_ascii(),
"Mismatch for non-ASCII character: {ch:?}"
);
}
}
#[test]
fn test_simd_vector_width_boundaries() {
for width in [16, 32, 64] {
let exact = "a".repeat(width);
assert!(is_ascii_simd(&exact));
let one_less = "a".repeat(width - 1);
assert!(is_ascii_simd(&one_less));
let one_more = "a".repeat(width + 1);
assert!(is_ascii_simd(&one_more));
let mut boundary_test = "a".repeat(width - 1);
boundary_test.push('é');
assert!(!is_ascii_simd(&boundary_test));
}
}
#[test]
fn test_all_ascii_characters() {
let mut all_ascii = String::new();
for i in 0u8..=127 {
all_ascii.push(i as char);
}
assert!(is_ascii_simd(&all_ascii));
}
#[test]
fn debug_simple_case() {
assert!(is_ascii_simd("a"));
assert!(is_ascii_simd("aa"));
assert!(is_ascii_simd("aaa"));
assert!(!is_ascii_simd("é"));
println!("Simple cases work");
}
#[test]
fn test_find_last_byte_empty() {
assert_eq!(find_last_byte_simd(&[], b'a', false), None);
}
#[test]
fn test_find_last_byte_single() {
assert_eq!(find_last_byte_simd(b"a", b'a', false), Some(0));
assert_eq!(find_last_byte_simd(b"a", b'b', false), None);
}
#[test]
fn test_find_last_byte_multiple() {
let haystack = b"hello world hello";
assert_eq!(find_last_byte_simd(haystack, b'l', false), Some(15)); assert_eq!(find_last_byte_simd(haystack, b'h', false), Some(12)); assert_eq!(find_last_byte_simd(haystack, b'o', false), Some(16)); assert_eq!(find_last_byte_simd(haystack, b'x', false), None); }
#[test]
fn test_find_last_byte_newlines() {
let text = b"line1\nline2\nline3";
assert_eq!(find_last_byte_simd(text, b'\n', true), Some(11));
let single_line = b"no newlines here";
assert_eq!(find_last_byte_simd(single_line, b'\n', true), None);
}
#[test]
fn test_find_last_byte_long() {
let long_text = "a".repeat(100) + "b" + &"a".repeat(100);
let bytes = long_text.as_bytes();
assert_eq!(find_last_byte_simd(bytes, b'b', false), Some(100));
}
#[test]
fn test_count_utf8_chars_empty() {
assert_eq!(count_utf8_chars_simd(&[]), 0);
}
#[test]
fn test_count_utf8_chars_ascii() {
assert_eq!(count_utf8_chars_simd(b"hello"), 5);
assert_eq!(count_utf8_chars_simd(b"Hello, World!"), 13);
assert_eq!(count_utf8_chars_simd(b"123"), 3);
}
#[test]
fn test_count_utf8_chars_utf8() {
assert_eq!(count_utf8_chars_simd("café".as_bytes()), 4);
assert_eq!(count_utf8_chars_simd("🚀".as_bytes()), 1);
assert_eq!(count_utf8_chars_simd("Hello🚀".as_bytes()), 6);
}
#[test]
fn test_count_utf8_chars_consistency() {
let test_strings = vec!["Hello", "café", "🚀", "Hello, 世界!", "résumé", "测试", ""];
for test_str in test_strings {
let simd_count = count_utf8_chars_simd(test_str.as_bytes());
let std_count = test_str.chars().count();
assert_eq!(simd_count, std_count, "Mismatch for string: {test_str:?}");
}
}
#[test]
fn test_get_char_column_simple() {
assert_eq!(get_char_column_simd("hello", 5), 5);
assert_eq!(get_char_column_simd("hello", 3), 3);
assert_eq!(get_char_column_simd("hello", 0), 0);
}
#[test]
fn test_get_char_column_with_newlines() {
let text = "line1\nline2\nline3";
assert_eq!(get_char_column_simd(text, 0), 0); assert_eq!(get_char_column_simd(text, 6), 0); assert_eq!(get_char_column_simd(text, 12), 0);
assert_eq!(get_char_column_simd(text, 3), 3); assert_eq!(get_char_column_simd(text, 9), 3); assert_eq!(get_char_column_simd(text, 15), 3); }
#[test]
fn test_get_char_column_utf8() {
let text = "café\nnaïve";
assert_eq!(get_char_column_simd(text, 2), 2);
assert_eq!(get_char_column_simd(text, 6), 0);
assert_eq!(get_char_column_simd(text, 8), 2);
}
#[test]
fn test_get_char_column_consistency_with_original() {
fn original_get_char_column(text: &str, offset: usize) -> usize {
let src = text.as_bytes();
let mut col = 0;
for &b in src[..offset].iter().rev() {
if b == b'\n' {
break;
}
if b & 0b1100_0000 != 0b1000_0000 {
col += 1;
}
}
col
}
let test_cases = vec![
("hello", vec![0, 1, 3, 5]),
("line1\nline2", vec![0, 3, 5, 6, 9]),
("café\nworld", vec![0, 2, 5, 6, 8]),
("🚀test\nnew", vec![0, 1, 3, 6, 7]),
("", vec![0]),
("a", vec![0, 1]),
];
for (text, offsets) in test_cases {
for offset in offsets {
if offset <= text.len() {
let original = original_get_char_column(text, offset);
let simd = get_char_column_simd(text, offset);
assert_eq!(
original, simd,
"Mismatch for text: {text:?}, offset: {offset}"
);
}
}
}
}
#[test]
fn test_get_char_column_edge_cases() {
assert_eq!(get_char_column_simd("", 0), 0);
assert_eq!(get_char_column_simd("test", 0), 0);
assert_eq!(get_char_column_simd("test", 100), 0);
assert_eq!(get_char_column_simd("\n\n\n", 1), 0);
assert_eq!(get_char_column_simd("\n\n\n", 2), 0);
let long_line = "a".repeat(1000);
assert_eq!(get_char_column_simd(&long_line, 500), 500);
let long_with_newline = "a".repeat(500) + "\n" + &"b".repeat(300);
assert_eq!(get_char_column_simd(&long_with_newline, 800), 299);
}
}