1use crate::{DelimiterResult, classify_byte};
7
8pub unsafe fn find_delimiters(haystack: &[u8]) -> DelimiterResult {
18 find_delimiters_safe(haystack)
19}
20
21#[inline]
23pub fn find_delimiters_safe(haystack: &[u8]) -> DelimiterResult {
24 for (i, &b) in haystack.iter().enumerate() {
25 if is_delimiter(b) {
26 return DelimiterResult::Found { pos: i, byte: b };
27 }
28 }
29 DelimiterResult::NotFound
30}
31
32pub unsafe fn classify_bytes(input: &[u8]) -> Vec<u8> {
42 classify_bytes_safe(input)
43}
44
45#[inline]
47pub fn classify_bytes_safe(input: &[u8]) -> Vec<u8> {
48 input.iter().map(|&b| classify_byte(b)).collect()
49}
50
51pub unsafe fn skip_whitespace(input: &[u8]) -> usize {
59 skip_whitespace_safe(input)
60}
61
62#[inline]
64pub fn skip_whitespace_safe(input: &[u8]) -> usize {
65 input
66 .iter()
67 .position(|&b| !b.is_ascii_whitespace())
68 .unwrap_or(input.len())
69}
70
71pub unsafe fn compute_byte_mask(block: &[u8], byte: u8) -> u64 {
80 compute_byte_mask_safe(block, byte)
81}
82
83#[inline]
88pub fn compute_byte_mask_safe(block: &[u8], byte: u8) -> u64 {
89 let mut mask = 0u64;
90 for (i, &b) in block.iter().take(64).enumerate() {
91 if b == byte {
92 mask |= 1u64 << i;
93 }
94 }
95 mask
96}
97
98pub unsafe fn compute_all_masks(block: &[u8]) -> crate::AllMasks {
105 compute_all_masks_safe(block)
106}
107
108#[inline]
113pub fn compute_all_masks_safe(block: &[u8]) -> crate::AllMasks {
114 let mut masks = crate::AllMasks::default();
115 for (i, &b) in block.iter().take(64).enumerate() {
116 let bit = 1u64 << i;
117 match b {
118 b'<' => masks.lt |= bit,
119 b'>' => masks.gt |= bit,
120 b'"' => masks.quot |= bit,
121 b'\'' => masks.apos |= bit,
122 _ => {}
123 }
124 }
125 masks
126}
127
128#[inline(always)]
130fn is_delimiter(b: u8) -> bool {
131 matches!(b, b'<' | b'>' | b'&' | b'"' | b'\'' | b'=' | b'/')
132}
133
134#[cfg(test)]
135mod tests {
136 use super::*;
137 use crate::class;
138
139 #[test]
140 fn find_delimiters_lt() {
141 let input = b"hello <world>";
142 let result = unsafe { find_delimiters(input) };
143 assert_eq!(result, DelimiterResult::Found { pos: 6, byte: b'<' });
144 }
145
146 #[test]
147 fn find_delimiters_amp() {
148 let input = b"a & b";
149 let result = unsafe { find_delimiters(input) };
150 assert_eq!(result, DelimiterResult::Found { pos: 2, byte: b'&' });
151 }
152
153 #[test]
154 fn find_delimiters_none() {
155 let input = b"hello world";
156 let result = unsafe { find_delimiters(input) };
157 assert_eq!(result, DelimiterResult::NotFound);
158 }
159
160 #[test]
161 fn find_delimiters_empty() {
162 let result = unsafe { find_delimiters(b"") };
163 assert_eq!(result, DelimiterResult::NotFound);
164 }
165
166 #[test]
167 fn find_delimiters_first_byte() {
168 let result = unsafe { find_delimiters(b"<html>") };
169 assert_eq!(result, DelimiterResult::Found { pos: 0, byte: b'<' });
170 }
171
172 #[test]
173 fn find_delimiters_all_types() {
174 for &delim in b"<>&\"'=/" {
175 let input = [b'x', b'x', delim, b'x'];
176 let result = unsafe { find_delimiters(&input) };
177 assert_eq!(
178 result,
179 DelimiterResult::Found {
180 pos: 2,
181 byte: delim
182 },
183 "failed for delimiter 0x{delim:02X}"
184 );
185 }
186 }
187
188 #[test]
189 fn classify_bytes_mixed() {
190 let input = b"a1 <";
191 let result = unsafe { classify_bytes(input) };
192 assert_eq!(result[0], class::ALPHA); assert_eq!(result[1], class::DIGIT); assert_eq!(result[2], class::WHITESPACE); assert_eq!(result[3], class::DELIMITER); }
197
198 #[test]
199 fn classify_bytes_empty() {
200 let result = unsafe { classify_bytes(b"") };
201 assert!(result.is_empty());
202 }
203
204 #[test]
205 fn skip_whitespace_leading() {
206 let result = unsafe { skip_whitespace(b" hello") };
207 assert_eq!(result, 3);
208 }
209
210 #[test]
211 fn skip_whitespace_mixed() {
212 let result = unsafe { skip_whitespace(b" \t\n\rX") };
213 assert_eq!(result, 4);
214 }
215
216 #[test]
217 fn skip_whitespace_all() {
218 let result = unsafe { skip_whitespace(b" ") };
219 assert_eq!(result, 3);
220 }
221
222 #[test]
223 fn skip_whitespace_none() {
224 let result = unsafe { skip_whitespace(b"hello") };
225 assert_eq!(result, 0);
226 }
227
228 #[test]
229 fn skip_whitespace_empty() {
230 let result = unsafe { skip_whitespace(b"") };
231 assert_eq!(result, 0);
232 }
233
234 #[test]
235 fn compute_byte_mask_basic() {
236 let input = b"hello <world>";
237 let mask = unsafe { compute_byte_mask(input, b'<') };
238 assert_eq!(mask, 1 << 6);
239 }
240
241 #[test]
242 fn compute_byte_mask_multiple() {
243 let input = b"a<b<c";
244 let mask = unsafe { compute_byte_mask(input, b'<') };
245 assert_eq!(mask, (1 << 1) | (1 << 3));
246 }
247
248 #[test]
249 fn compute_byte_mask_none() {
250 let input = b"hello world";
251 let mask = unsafe { compute_byte_mask(input, b'<') };
252 assert_eq!(mask, 0);
253 }
254
255 #[test]
256 fn compute_byte_mask_empty() {
257 let mask = unsafe { compute_byte_mask(b"", b'<') };
258 assert_eq!(mask, 0);
259 }
260
261 #[test]
262 fn compute_byte_mask_over_64_bytes_does_not_overflow() {
263 let input = vec![b'<'; 80];
266 let mask = compute_byte_mask_safe(&input, b'<');
267 assert_eq!(mask, u64::MAX, "first 64 '<' set, rest ignored");
268 }
269
270 #[test]
271 fn compute_all_masks_over_64_bytes_does_not_overflow() {
272 let input = vec![b'<'; 80];
273 let masks = compute_all_masks_safe(&input);
274 assert_eq!(masks.lt, u64::MAX);
275 }
276
277 #[test]
278 fn compute_all_masks_basic() {
279 let input = b"<div class=\"foo\">";
280 let masks = unsafe { compute_all_masks(input) };
281 assert_eq!(masks.lt, 1 << 0); assert_eq!(masks.gt, 1 << 16); assert_eq!(masks.quot, (1 << 11) | (1 << 15)); }
285
286 #[test]
287 fn compute_all_masks_empty() {
288 let masks = unsafe { compute_all_masks(b"") };
289 assert_eq!(masks.lt, 0);
290 assert_eq!(masks.gt, 0);
291 }
292
293 #[test]
294 fn compute_all_masks_matches_individual() {
295 let input = b"Hello <World> & \"test\" = 'value' / 123\n\r\t end!!";
296 let masks = unsafe { compute_all_masks(input) };
297 assert_eq!(masks.lt, unsafe { compute_byte_mask(input, b'<') });
298 assert_eq!(masks.gt, unsafe { compute_byte_mask(input, b'>') });
299 assert_eq!(masks.quot, unsafe { compute_byte_mask(input, b'"') });
300 assert_eq!(masks.apos, unsafe { compute_byte_mask(input, b'\'') });
301 }
302}