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}