Skip to main content

rustpython_common/
wtf8_index.rs

1// spell-checker:ignore rpython rlib rutf
2//! Random access into a WTF-8 buffer.
3//!
4//! WTF-8 is variable width, so a buffer's n-th code point can only be found by
5//! decoding the n-1 before it: [`Wtf8`]'s iterators are sequential, and
6//! resolving an index through them is O(n). Code that indexes the same string
7//! repeatedly -- a regex scan restarting at successive positions, say -- then
8//! walks the whole buffer once per index, which is quadratic in its length.
9//!
10//! [`Wtf8Index`] is the side table that makes the lookup O(1): one 24-byte
11//! group per 64 code points, so 0.375 bytes per code point. It is a cache, and
12//! holds no state of its own beyond the buffer's shape -- building it twice for
13//! the same buffer yields the same table.
14//!
15//! The layout is PyPy's `UTF8_INDEX_STORAGE` (`rpython/rlib/rutf8.py`).
16
17use crate::wtf8::Wtf8;
18
19/// One group of 64 code points.
20#[derive(Clone, Copy)]
21struct Group {
22    /// The byte offset the group's first code point starts at.
23    base: usize,
24    /// `ofs[i]` is the byte offset of the group's `4 * i + 1`-th code point,
25    /// relative to `base`. One entry covers four code points, so the widest
26    /// offset an entry has to hold is that of the 61st code point of a group,
27    /// at most `61 * 4 = 244` bytes in -- inside a `u8`, which is what buys the
28    /// table its density.
29    ofs: [u8; 16],
30}
31
32/// A code-point-index to byte-offset table for one WTF-8 buffer.
33pub struct Wtf8Index {
34    groups: Box<[Group]>,
35}
36
37impl Wtf8Index {
38    /// Builds the table for `data`, whose code point count is `char_len`.
39    ///
40    /// O(`data.len()`), and touches every byte, so it pays for itself only when
41    /// the caller goes on to index the buffer more than a couple of times.
42    #[must_use]
43    pub fn new(data: &Wtf8, char_len: usize) -> Self {
44        let mut groups = vec![
45            Group {
46                base: 0,
47                ofs: [0; 16],
48            };
49            char_len / 64 + 1
50        ];
51        // Signed: the countdown overshoots the last group -- the loop stops on
52        // the first negative value rather than at a group boundary.
53        let mut remaining = char_len as isize;
54        let mut base = 0;
55        let mut current = 0;
56        loop {
57            groups[current].base = base;
58            let mut next = base;
59            let mut group_filled = true;
60            for i in 0..16 {
61                // Past the end, step as if one more single-byte code point
62                // followed, so the entry stays in range and is never read.
63                next = if remaining == 0 {
64                    next + 1
65                } else {
66                    next_pos(data, next)
67                };
68                groups[current].ofs[i] = (next - base) as u8;
69                remaining -= 4;
70                if remaining < 0 {
71                    debug_assert_eq!(current + 1, groups.len());
72                    group_filled = false;
73                    break;
74                }
75                next = next_pos(data, next_pos(data, next_pos(data, next)));
76            }
77            if !group_filled {
78                break;
79            }
80            current += 1;
81            base = next;
82        }
83        Self {
84            groups: groups.into_boxed_slice(),
85        }
86    }
87
88    /// The byte offset of `data`'s `index`-th code point.
89    ///
90    /// `data` must be the buffer the table was built for, and `index` must be
91    /// below its code point count.
92    #[inline]
93    #[must_use]
94    pub fn byte_offset(&self, data: &Wtf8, index: usize) -> usize {
95        let group = &self.groups[index >> 6];
96        // The entry sits on the 4k+1-th code point of the group, so a lookup is
97        // one table read plus at most two steps in either direction.
98        let pos = group.base + group.ofs[(index >> 2) & 0x0F] as usize;
99        match index & 0x3 {
100            0 => prev_pos(data, pos),
101            1 => pos,
102            2 => next_pos(data, pos),
103            _ => next_pos(data, next_pos(data, pos)),
104        }
105    }
106
107    /// The index of the code point starting at byte offset `bytepos`, the
108    /// inverse of [`Self::byte_offset`].
109    ///
110    /// `data` must be the buffer the table was built for, `char_len` its code
111    /// point count, and `bytepos` a code point boundary at or before its end.
112    ///
113    /// Logarithmic rather than constant: the table is keyed by code point
114    /// index, so going the other way is a search through it. The bracketing
115    /// below is what keeps that search short -- a code point occupies one to
116    /// four bytes, which pins the answer to a narrow band around `bytepos`
117    /// before the first comparison.
118    #[must_use]
119    pub fn char_index_at_byte(&self, data: &Wtf8, bytepos: usize, char_len: usize) -> usize {
120        let bytes_remaining = data.len() - bytepos;
121        // At least one byte per remaining code point, and at most four, so the
122        // group holding the answer lies between these.
123        let mut group_min =
124            usize::max(bytepos / 4, char_len.saturating_sub(bytes_remaining + 1)) >> 6;
125        let mut group_max = usize::min(bytepos, char_len.saturating_sub(bytes_remaining / 4)) >> 6;
126        while group_min < group_max {
127            let middle = group_min.midpoint(group_max) + 1;
128            if bytepos < self.groups[middle].base {
129                group_max = middle - 1;
130            } else {
131                group_min = middle;
132            }
133        }
134
135        let base = self.groups[group_min].base;
136        if base == bytepos {
137            return group_min << 6;
138        }
139        // Walk the group's entries to the last one at or before `bytepos`,
140        // then step the remaining code points, of which there are at most
141        // three -- an entry covers four.
142        let entries = if group_min == self.groups.len() - 1 {
143            ((char_len - 1) >> 2) & 0x0F
144        } else {
145            16
146        };
147        let mut index = group_min << 6;
148        let mut pos = base;
149        for (i, &entry) in self.groups[group_min].ofs.iter().enumerate().take(entries) {
150            let at = base + entry as usize;
151            if at >= bytepos {
152                break;
153            }
154            pos = at;
155            index = (group_min << 6) + (i << 2) + 1;
156        }
157        while pos < bytepos {
158            pos = next_pos(data, pos);
159            index += 1;
160        }
161        index
162    }
163
164    /// The table's heap footprint, in bytes.
165    #[must_use]
166    pub fn byte_size(&self) -> usize {
167        core::mem::size_of_val(&*self.groups)
168    }
169}
170
171/// The byte offset of the code point after the one at `pos`.
172///
173/// `data` must be well-formed WTF-8 and `pos` a code point boundary before its
174/// end -- reading only the lead byte is what makes this branch-light.
175#[inline]
176fn next_pos(data: &Wtf8, pos: usize) -> usize {
177    match data.as_bytes()[pos] {
178        0x00..=0x7F => pos + 1,
179        0x80..=0xDF => pos + 2,
180        0xE0..=0xEF => pos + 3,
181        _ => pos + 4,
182    }
183}
184
185/// The byte offset of the code point before the one at `pos`, which must not be
186/// zero.
187///
188/// A `pos` one past the end reads as the extra code point [`Wtf8Index::new`]
189/// steps over there.
190#[inline]
191fn prev_pos(data: &Wtf8, pos: usize) -> usize {
192    let data = data.as_bytes();
193    let mut pos = pos - 1;
194    if pos >= data.len() || data[pos] <= 0x7F {
195        return pos;
196    }
197    pos -= 1;
198    if data[pos] >= 0xC0 {
199        return pos;
200    }
201    pos -= 1;
202    if data[pos] >= 0xC0 {
203        return pos;
204    }
205    pos - 1
206}
207
208#[cfg(test)]
209mod tests {
210    use super::*;
211    use crate::wtf8::{CodePoint, Wtf8Buf};
212
213    /// Every index of `s`, both ways, against the offsets its own iterator
214    /// reports.
215    fn check(s: &Wtf8) {
216        let expected: Vec<usize> = s
217            .code_point_indices()
218            .map(|(byte_offset, _)| byte_offset)
219            .collect();
220        let char_len = expected.len();
221        let index = Wtf8Index::new(s, char_len);
222        for (i, &want) in expected.iter().enumerate() {
223            assert_eq!(
224                index.byte_offset(s, i),
225                want,
226                "index {i} of {s:?} ({char_len} code points)"
227            );
228            assert_eq!(
229                index.char_index_at_byte(s, want, char_len),
230                i,
231                "byte {want} of {s:?} ({char_len} code points)"
232            );
233        }
234        // One past the last code point is a boundary too, and the searches that
235        // use this ask for it as an end bound.
236        assert_eq!(
237            index.char_index_at_byte(s, s.len(), char_len),
238            char_len,
239            "end of {s:?}"
240        );
241    }
242
243    fn wtf8(s: &str) -> Wtf8Buf {
244        Wtf8Buf::from(s)
245    }
246
247    #[test]
248    fn empty() {
249        check(wtf8("").as_ref());
250    }
251
252    #[test]
253    fn widths() {
254        // One case per encoded width, and the boundaries between them.
255        check(wtf8("abc").as_ref());
256        check(wtf8("\u{80}\u{7ff}").as_ref());
257        check(wtf8("\u{800}\u{ffff}").as_ref());
258        check(wtf8("\u{10000}\u{10ffff}").as_ref());
259        check(wtf8("a\u{80}\u{800}\u{10000}").as_ref());
260    }
261
262    #[test]
263    fn group_boundaries() {
264        // A group covers 64 code points and an entry four, so the interesting
265        // lengths are the ones on and around both.
266        for len in [1, 3, 4, 5, 63, 64, 65, 127, 128, 129, 255, 256, 257] {
267            for unit in ["a", "\u{80}", "\u{800}", "\u{10000}"] {
268                check(wtf8(&unit.repeat(len)).as_ref());
269            }
270            // Mixed widths, so a group's entries do not share a stride.
271            check(wtf8(&"a\u{80}\u{800}\u{10000}".repeat(len)).as_ref());
272        }
273    }
274
275    #[test]
276    fn lone_surrogates() {
277        let mut s = wtf8("a");
278        for cp in [0xD800, 0xDBFF, 0xDC00, 0xDFFF] {
279            s.push(CodePoint::from_u32(cp).unwrap());
280            s.push_str("b");
281        }
282        check(s.as_ref());
283
284        // Surrogates only, spanning more than one group.
285        let mut s = wtf8("");
286        for i in 0..200 {
287            s.push(CodePoint::from_u32(0xD800 + (i % 0x400)).unwrap());
288        }
289        check(s.as_ref());
290    }
291
292    #[test]
293    fn byte_size_is_one_group_per_64_code_points() {
294        let s = wtf8(&"\u{10000}".repeat(200));
295        let index = Wtf8Index::new(s.as_ref(), 200);
296        assert_eq!(index.byte_size(), (200 / 64 + 1) * size_of::<Group>());
297        assert_eq!(size_of::<Group>(), 24);
298    }
299}