Skip to main content

ftui_simd/
lib.rs

1#![forbid(unsafe_code)]
2#![feature(portable_simd)]
3
4//! Portable-SIMD kernels for FrankenTUI's two lane-parallel hot loops.
5//!
6//! # Role in FrankenTUI
7//! Two places in the render path do the same narrow thing over a long run of
8//! bytes: the diff compares one row of 16-byte cells against the previous
9//! frame's row, and the width fast path asks whether a string is entirely
10//! ASCII. Both are a wide equality test followed by "where was the first
11//! difference", which is what SIMD lanes are for.
12//!
13//! # How it fits in the system
14//! Every kernel here is safe `std::simd` and carries a `_scalar` twin with
15//! identical semantics. The twins are not dead weight: they are what the
16//! parity tests compare against, what the benches measure against, and what
17//! callers use when the `simd` feature is off. Nothing in the workspace
18//! depends on this crate unless that feature is enabled, so the scalar paths
19//! stay the default until benches justify otherwise.
20//!
21//! # Why `u128` and not `Cell`
22//! `ftui_render::Cell` is `#[repr(C, align(16))]` over four `u32` fields, so a
23//! cell is bit-for-bit a `u128`. Keeping the kernels on `u128` leaves this
24//! crate free of a dependency on the render crate, and leaves the conversion
25//! (which must stay safe, so it composes the fields rather than transmuting)
26//! on the caller's side.
27//!
28//! # Lane widths
29//! The kernels use 512-bit vectors (`u64x8`, `u8x64`). That is wider than most
30//! targets execute natively; `std::simd` splits them into whatever the target
31//! has, which keeps one code path across x86-64 and aarch64 and gives the
32//! compiler a full unrolled chunk to work with.
33//!
34//! # Which of these are worth calling
35//! Measured 2026-09-18, full numbers in
36//! `docs/perf/simd_kernels_2026-09-18.md`:
37//!
38//! - [`all_ascii`] and [`ascii_width`] beat their twins by 5x at 64 bytes and
39//!   by 20-44x from a kilobyte up. Call these.
40//! - [`first_mismatch_u128`] and [`rows_equal_u128`] are **3-4x slower** than
41//!   their twins and are deliberately not wired into the diff. `#![forbid(unsafe_code)]`
42//!   leaves no way to view `&[u128]` as lanes, so the kernel must build each
43//!   vector with shifts and masks, while the scalar `a[i] != b[i]` over `u128`
44//!   is already one 128-bit compare after LLVM is done with it. They are kept
45//!   because the measurement is worth keeping, and because the parity tests
46//!   over them are what prove the lane indexing is right.
47//!
48//! Prefer the `_scalar` twin for cell comparison. That is not a placeholder.
49
50use std::simd::cmp::{SimdPartialEq, SimdPartialOrd};
51use std::simd::{Mask, Simd, u8x64, u64x8};
52
53/// Cells compared per vector chunk: eight `u64` lanes is four `u128` cells.
54const CELLS_PER_CHUNK: usize = 4;
55
56/// Bytes compared per vector chunk.
57const BYTES_PER_CHUNK: usize = 64;
58
59/// Index of the first element where `a` and `b` differ, or `None` when the
60/// common prefix runs to the end of the shorter slice.
61///
62/// Only the first `a.len().min(b.len())` elements are examined; a length
63/// difference past that point is not a mismatch this reports, because the
64/// diff's callers size both rows to the same width and a shorter slice means
65/// a clipped row rather than a changed cell.
66///
67/// # Examples
68/// ```
69/// # use ftui_simd::first_mismatch_u128;
70/// assert_eq!(first_mismatch_u128(&[1, 2, 3], &[1, 9, 3]), Some(1));
71/// assert_eq!(first_mismatch_u128(&[1, 2, 3], &[1, 2, 3]), None);
72/// ```
73#[must_use]
74pub fn first_mismatch_u128(a: &[u128], b: &[u128]) -> Option<usize> {
75    let len = a.len().min(b.len());
76    let mut offset = 0;
77
78    while offset + CELLS_PER_CHUNK <= len {
79        let lhs = load_cells(&a[offset..offset + CELLS_PER_CHUNK]);
80        let rhs = load_cells(&b[offset..offset + CELLS_PER_CHUNK]);
81        let differing: Mask<i64, 8> = lhs.simd_ne(rhs);
82        if differing.any() {
83            // Each cell is two consecutive u64 lanes, so the first differing
84            // lane identifies the cell by halving its index. Both halves of a
85            // cell can differ; the earlier lane is the one `first_set` gives.
86            let lane = differing.first_set().unwrap_or(0);
87            return Some(offset + lane / 2);
88        }
89        offset += CELLS_PER_CHUNK;
90    }
91
92    // Tail shorter than one chunk.
93    (offset..len).find(|&i| a[i] != b[i])
94}
95
96/// Scalar twin of [`first_mismatch_u128`], and the reference its parity tests
97/// compare against.
98#[must_use]
99pub fn first_mismatch_u128_scalar(a: &[u128], b: &[u128]) -> Option<usize> {
100    let len = a.len().min(b.len());
101    (0..len).find(|&i| a[i] != b[i])
102}
103
104/// Whether two equal-length runs of cells are bitwise identical.
105///
106/// Returns `false` for slices of differing length: callers use this to decide
107/// whether a row can be skipped entirely, and a row whose width changed cannot.
108///
109/// # Examples
110/// ```
111/// # use ftui_simd::rows_equal_u128;
112/// assert!(rows_equal_u128(&[7, 7], &[7, 7]));
113/// assert!(!rows_equal_u128(&[7, 7], &[7, 7, 7]));
114/// ```
115#[must_use]
116pub fn rows_equal_u128(a: &[u128], b: &[u128]) -> bool {
117    a.len() == b.len() && first_mismatch_u128(a, b).is_none()
118}
119
120/// Scalar twin of [`rows_equal_u128`].
121#[must_use]
122pub fn rows_equal_u128_scalar(a: &[u128], b: &[u128]) -> bool {
123    a.len() == b.len() && first_mismatch_u128_scalar(a, b).is_none()
124}
125
126/// Whether every byte is ASCII, i.e. has its high bit clear.
127///
128/// This is the question the width fast path actually asks: if no byte is
129/// continuation or lead, the run needs no Unicode width lookup and its display
130/// width is its length.
131///
132/// # Examples
133/// ```
134/// # use ftui_simd::all_ascii;
135/// assert!(all_ascii(b"plain text"));
136/// assert!(!all_ascii("caf\u{e9}".as_bytes()));
137/// assert!(all_ascii(b""));
138/// ```
139#[must_use]
140pub fn all_ascii(bytes: &[u8]) -> bool {
141    let mut offset = 0;
142
143    while offset + BYTES_PER_CHUNK <= bytes.len() {
144        let chunk = u8x64::from_slice(&bytes[offset..offset + BYTES_PER_CHUNK]);
145        if (chunk & u8x64::splat(0x80)).simd_ne(u8x64::splat(0)).any() {
146            return false;
147        }
148        offset += BYTES_PER_CHUNK;
149    }
150
151    bytes[offset..].iter().all(u8::is_ascii)
152}
153
154/// Scalar twin of [`all_ascii`].
155#[must_use]
156pub fn all_ascii_scalar(bytes: &[u8]) -> bool {
157    bytes.iter().all(u8::is_ascii)
158}
159
160/// Display width of a run that is entirely printable ASCII, or `None`.
161///
162/// `None` means "this run needs the real width tables", and covers both
163/// non-ASCII bytes and ASCII control characters: C0 (`0x00..=0x1f`) and DEL
164/// (`0x7f`) have no single agreed display width, so they are handed back to
165/// the caller rather than counted as one column each. Space through `~` are
166/// one column apiece, so the width of an accepted run is its length.
167///
168/// # Examples
169/// ```
170/// # use ftui_simd::ascii_width;
171/// assert_eq!(ascii_width(b"hello"), Some(5));
172/// assert_eq!(ascii_width(b"tab\there"), None);
173/// assert_eq!(ascii_width("\u{e9}".as_bytes()), None);
174/// ```
175#[must_use]
176pub fn ascii_width(bytes: &[u8]) -> Option<usize> {
177    let mut offset = 0;
178
179    while offset + BYTES_PER_CHUNK <= bytes.len() {
180        let chunk = u8x64::from_slice(&bytes[offset..offset + BYTES_PER_CHUNK]);
181        // Printable ASCII is 0x20..=0x7e, so one unsigned-wrapping subtract
182        // puts every acceptable byte in 0..=0x5e and everything else above it:
183        // 0x7f wraps to 0x5f, and any high-bit byte lands at 0x60 or more.
184        let shifted = chunk - u8x64::splat(0x20);
185        if shifted.simd_gt(u8x64::splat(0x5e)).any() {
186            return None;
187        }
188        offset += BYTES_PER_CHUNK;
189    }
190
191    if bytes[offset..].iter().all(|b| (0x20..=0x7e).contains(b)) {
192        Some(bytes.len())
193    } else {
194        None
195    }
196}
197
198/// Scalar twin of [`ascii_width`].
199#[must_use]
200pub fn ascii_width_scalar(bytes: &[u8]) -> Option<usize> {
201    if bytes.iter().all(|b| (0x20..=0x7e).contains(b)) {
202        Some(bytes.len())
203    } else {
204        None
205    }
206}
207
208/// Reinterpret four `u128` cells as the eight `u64` lanes of one vector.
209///
210/// Splitting each cell into two `u64` halves keeps the whole thing in safe
211/// code: there is no 128-bit lane type to compare with, and `u64` is the
212/// widest lane every supported target handles well.
213#[inline]
214fn load_cells(cells: &[u128]) -> u64x8 {
215    debug_assert_eq!(cells.len(), CELLS_PER_CHUNK);
216    let mut lanes = [0_u64; 8];
217    for (i, cell) in cells.iter().enumerate() {
218        lanes[i * 2] = (*cell & u128::from(u64::MAX)) as u64;
219        lanes[i * 2 + 1] = (*cell >> 64) as u64;
220    }
221    Simd::from_array(lanes)
222}
223
224#[cfg(test)]
225mod tests {
226    use super::*;
227
228    #[test]
229    fn first_mismatch_reports_none_for_identical_rows() {
230        let row: Vec<u128> = (0..200).collect();
231        assert_eq!(first_mismatch_u128(&row, &row), None);
232        assert_eq!(first_mismatch_u128_scalar(&row, &row), None);
233    }
234
235    #[test]
236    fn first_mismatch_finds_the_earliest_difference_at_every_position() {
237        // Every index matters: inside the first chunk, straddling a chunk
238        // boundary, and in the scalar tail.
239        for len in [1_usize, 3, 4, 5, 8, 63, 64, 65, 200] {
240            let base: Vec<u128> = (0..len as u128).collect();
241            for idx in 0..len {
242                let mut changed = base.clone();
243                changed[idx] = u128::MAX;
244                assert_eq!(
245                    first_mismatch_u128(&base, &changed),
246                    Some(idx),
247                    "len {len}, idx {idx}"
248                );
249                assert_eq!(
250                    first_mismatch_u128_scalar(&base, &changed),
251                    Some(idx),
252                    "scalar len {len}, idx {idx}"
253                );
254            }
255        }
256    }
257
258    #[test]
259    fn first_mismatch_detects_a_difference_in_either_half_of_a_cell() {
260        // The low half alone, the high half alone, and both: all three must
261        // report the same cell, since a cell is two lanes.
262        for delta in [1_u128, 1_u128 << 64, (1_u128 << 64) | 1] {
263            let base = vec![0_u128; 8];
264            let mut changed = base.clone();
265            changed[5] = delta;
266            assert_eq!(
267                first_mismatch_u128(&base, &changed),
268                Some(5),
269                "delta {delta}"
270            );
271        }
272    }
273
274    #[test]
275    fn first_mismatch_stops_at_the_shorter_slice() {
276        let long: Vec<u128> = (0..16).collect();
277        let short = &long[..6];
278        assert_eq!(first_mismatch_u128(&long, short), None);
279        assert_eq!(first_mismatch_u128(short, &long), None);
280    }
281
282    #[test]
283    fn first_mismatch_handles_empty_input() {
284        assert_eq!(first_mismatch_u128(&[], &[]), None);
285        assert_eq!(first_mismatch_u128(&[], &[1, 2]), None);
286    }
287
288    #[test]
289    fn rows_equal_requires_matching_length() {
290        assert!(rows_equal_u128(&[1, 2, 3], &[1, 2, 3]));
291        assert!(!rows_equal_u128(&[1, 2, 3], &[1, 2]));
292        assert!(!rows_equal_u128(&[1, 2, 3], &[1, 2, 4]));
293        assert!(rows_equal_u128(&[], &[]));
294    }
295
296    #[test]
297    fn all_ascii_agrees_with_the_scalar_twin_across_lengths() {
298        for len in [0_usize, 1, 63, 64, 65, 127, 128, 1000] {
299            let ascii = vec![b'a'; len];
300            assert!(all_ascii(&ascii), "len {len}");
301            assert_eq!(all_ascii(&ascii), all_ascii_scalar(&ascii));
302
303            if len > 0 {
304                // A single non-ASCII byte must be found wherever it sits.
305                for idx in [0, len / 2, len - 1] {
306                    let mut probe = ascii.clone();
307                    probe[idx] = 0xC3;
308                    assert!(!all_ascii(&probe), "len {len}, idx {idx}");
309                    assert_eq!(all_ascii(&probe), all_ascii_scalar(&probe));
310                }
311            }
312        }
313    }
314
315    #[test]
316    fn all_ascii_accepts_control_bytes() {
317        // Control characters are ASCII; only ascii_width is fussy about them.
318        assert!(all_ascii(b"\x00\x09\x1b\x7f"));
319    }
320
321    #[test]
322    fn ascii_width_counts_printable_runs_and_rejects_the_rest() {
323        assert_eq!(ascii_width(b""), Some(0));
324        assert_eq!(ascii_width(b" "), Some(1));
325        assert_eq!(ascii_width(b"~"), Some(1));
326        assert_eq!(ascii_width(b"hello world"), Some(11));
327
328        let long = vec![b'x'; 300];
329        assert_eq!(ascii_width(&long), Some(300));
330
331        // The boundaries either side of printable ASCII, and a high byte.
332        assert_eq!(ascii_width(b"\x1f"), None);
333        assert_eq!(ascii_width(b"\x7f"), None);
334        assert_eq!(ascii_width(b"\xc3\xa9"), None);
335    }
336
337    #[test]
338    fn ascii_width_rejects_a_bad_byte_anywhere_including_past_a_full_chunk() {
339        for len in [64_usize, 65, 129, 300] {
340            for idx in [0, len / 2, len - 1] {
341                for bad in [0x00_u8, 0x1f, 0x7f, 0x80, 0xff] {
342                    let mut probe = vec![b'a'; len];
343                    probe[idx] = bad;
344                    assert_eq!(ascii_width(&probe), None, "len {len}, idx {idx}, bad {bad}");
345                    assert_eq!(ascii_width(&probe), ascii_width_scalar(&probe));
346                }
347            }
348        }
349    }
350
351    #[test]
352    fn kernels_agree_with_their_twins_over_a_deterministic_sweep() {
353        // A cheap xorshift keeps this reproducible without a dev-dependency.
354        let mut state = 0x2545_F491_4F6C_DD1D_u64;
355        let mut next = move || {
356            state ^= state << 13;
357            state ^= state >> 7;
358            state ^= state << 17;
359            state
360        };
361
362        for len in 0..200_usize {
363            // Fill both halves of each cell. Widening a single u64 left every
364            // base value with a zero high half, so a kernel that only ever
365            // compared the low 64 bits of equal cells would have passed —
366            // the perturbation below can flip a high bit, but two *equal*
367            // cells were never wide.
368            let a: Vec<u128> = (0..len)
369                .map(|_| (u128::from(next()) << 64) | u128::from(next()))
370                .collect();
371            let mut b = a.clone();
372            if len > 0 {
373                let idx = (next() as usize) % len;
374                // Half the time perturb a cell, half the time leave them equal.
375                if next() % 2 == 0 {
376                    b[idx] ^= 1 << (next() % 128);
377                }
378            }
379            assert_eq!(
380                first_mismatch_u128(&a, &b),
381                first_mismatch_u128_scalar(&a, &b),
382                "len {len}"
383            );
384            assert_eq!(rows_equal_u128(&a, &b), rows_equal_u128_scalar(&a, &b));
385
386            let bytes: Vec<u8> = (0..len).map(|_| (next() % 256) as u8).collect();
387            assert_eq!(all_ascii(&bytes), all_ascii_scalar(&bytes), "len {len}");
388            assert_eq!(ascii_width(&bytes), ascii_width_scalar(&bytes), "len {len}");
389        }
390    }
391
392    /// The sweep above draws bytes uniformly from `0..=255`, so a run of
393    /// `BYTES_PER_CHUNK` bytes is entirely ASCII with probability 2^-64. Every
394    /// chunk it ever evaluates therefore takes the rejecting branch on its
395    /// first iteration: the accepting path through the chunk loop - advancing
396    /// `offset`, and the tail that is measured from wherever it stopped - is
397    /// never reached for any input long enough to have a chunk at all.
398    ///
399    /// These runs are printable ASCII by construction, then poisoned one byte
400    /// at a time so the rejecting branch is exercised at every position rather
401    /// than only near the front.
402    #[test]
403    fn ascii_kernels_agree_with_their_twins_on_runs_that_reach_the_chunk_loop() {
404        // Past two full chunks, so a bad byte can land in the first chunk, a
405        // later chunk, or the tail.
406        for len in 0..=(2 * BYTES_PER_CHUNK + 5) {
407            // 0x20..=0x7e, cycled: printable, and not a single repeated byte.
408            let mut run: Vec<u8> = (0..len).map(|i| 0x20 + (i % 0x5f) as u8).collect();
409
410            assert!(all_ascii(&run), "len {len}");
411            assert_eq!(all_ascii(&run), all_ascii_scalar(&run), "len {len}");
412            assert_eq!(ascii_width(&run), Some(len), "len {len}");
413            assert_eq!(ascii_width(&run), ascii_width_scalar(&run), "len {len}");
414
415            for pos in 0..len {
416                // NUL and 0x1f are ASCII but not printable, so the two kernels
417                // disagree with each other by design - each is compared only
418                // against its own twin. 0x7f is the byte the wrapping subtract
419                // in `ascii_width` folds to the top of the accepted range.
420                for bad in [0x00_u8, 0x1f, 0x7f, 0x80, 0xff] {
421                    let saved = run[pos];
422                    run[pos] = bad;
423                    assert_eq!(
424                        all_ascii(&run),
425                        all_ascii_scalar(&run),
426                        "all_ascii len {len} pos {pos} byte {bad:#04x}"
427                    );
428                    assert_eq!(
429                        ascii_width(&run),
430                        ascii_width_scalar(&run),
431                        "ascii_width len {len} pos {pos} byte {bad:#04x}"
432                    );
433                    run[pos] = saved;
434                }
435            }
436        }
437    }
438}