1use geo_types::Coord;
2use wide::u32x8;
3
4use crate::decoder::Morton;
5use crate::encoder::model::CurveParams;
6use crate::{Decoder, MltError, MltResult};
7
8const LANES: usize = 8;
9
10#[must_use]
18#[inline]
19pub fn interleave_bits(coord: Coord<u32>) -> u32 {
20 let mut sx = coord.x & 0xFFFF;
24 sx = (sx | (sx << 8)) & 0x00FF_00FF;
25 sx = (sx | (sx << 4)) & 0x0F0F_0F0F;
26 sx = (sx | (sx << 2)) & 0x3333_3333;
27 sx = (sx | (sx << 1)) & 0x5555_5555;
28
29 let mut sy = coord.y & 0xFFFF;
30 sy = (sy | (sy << 8)) & 0x00FF_00FF;
31 sy = (sy | (sy << 4)) & 0x0F0F_0F0F;
32 sy = (sy | (sy << 2)) & 0x3333_3333;
33 sy = (sy | (sy << 1)) & 0x5555_5555;
34
35 sx | (sy << 1)
36}
37
38#[must_use]
49#[inline]
50pub fn morton_sort_key(c: Coord<i32>, params: CurveParams) -> u32 {
51 debug_assert!((1..=16).contains(¶ms.bits));
52 #[expect(
53 clippy::cast_possible_truncation,
54 clippy::cast_sign_loss,
55 reason = "shift brings value into [0, extent]; masked to 16 bits immediately after"
56 )]
57 let sx = ((i64::from(c.x) + i64::from(params.shift)) as u32) & 0xFFFF;
58 #[expect(
59 clippy::cast_possible_truncation,
60 clippy::cast_sign_loss,
61 reason = "shift brings value into [0, extent]; masked to 16 bits immediately after"
62 )]
63 let sy = ((i64::from(c.y) + i64::from(params.shift)) as u32) & 0xFFFF;
64 interleave_bits((sx, sy).into())
65}
66
67impl Morton {
69 pub fn from_vertices(vertices: &[i32]) -> MltResult<Self> {
74 let min_v = vertices.iter().copied().min().unwrap_or(0);
75 let max_v = vertices.iter().copied().max().unwrap_or(0);
76 let shift: u32 = if min_v < 0 { min_v.unsigned_abs() } else { 0 };
77 let tile_extent = i64::from(max_v) + i64::from(shift);
78 let bits = if let Ok(extent) = u32::try_from(tile_extent) {
79 let required_bits = u32::BITS - extent.leading_zeros();
83 if required_bits > 16 {
84 return Err(MltError::VertexMortonNotCompatibleWithExtent {
85 extent,
86 required_bits,
87 });
88 }
89 required_bits
90 } else {
91 0u32
92 };
93 Self::new(bits, shift)
94 }
95
96 #[inline]
101 pub fn encode_morton(self, x: i32, y: i32) -> MltResult<u32> {
102 let sx = u32::try_from(i64::from(x) + i64::from(self.shift))?;
103 let sy = u32::try_from(i64::from(y) + i64::from(self.shift))?;
104 let mut code = 0u32;
105 for i in 0..self.bits {
106 code |= ((sx >> i) & 1) << (2 * i);
108 code |= ((sy >> i) & 1) << (2 * i + 1);
109 }
110 Ok(code)
111 }
112}
113
114impl Morton {
115 #[inline]
117 fn decode_one(self, morton_code: u32) -> Coord<i32> {
118 let mut x = 0u32;
119 let mut y = 0u32;
120 for i in 0..self.bits {
121 let bit_mask = 1u32 << (2 * i);
122 x |= (morton_code & bit_mask) >> i;
123 y |= ((morton_code >> 1) & bit_mask) >> i;
124 }
125 Coord::<i32> {
126 x: x.wrapping_sub(self.shift).cast_signed(),
127 y: y.wrapping_sub(self.shift).cast_signed(),
128 }
129 }
130
131 pub fn decode_codes(self, data: &[u32], dec: &mut Decoder) -> MltResult<Vec<i32>> {
137 let alloc_size = data.len() * 2;
138 let mut out = dec.alloc(alloc_size)?;
139 let shift_vec = u32x8::splat(self.shift);
140
141 let mut chunks = data.chunks_exact(LANES);
142
143 for chunk in chunks.by_ref() {
144 let buf = [
145 chunk[0], chunk[1], chunk[2], chunk[3], chunk[4], chunk[5], chunk[6], chunk[7],
146 ];
147 self.decode_chunk(buf, shift_vec, &mut out);
148 }
149
150 for &code in chunks.remainder() {
152 let coord = self.decode_one(code);
153 out.push(coord.x);
154 out.push(coord.y);
155 }
156
157 dec.adjust_alloc(&out, alloc_size)?;
158 Ok(out)
159 }
160
161 pub fn decode_delta(self, data: &[u32], dec: &mut Decoder) -> MltResult<Vec<i32>> {
168 let alloc_size = data.len() * 2;
169 let mut out = dec.alloc(alloc_size)?;
170 let shift_vec = u32x8::splat(self.shift);
171
172 let mut prev = 0i32;
173 let mut chunks = data.chunks_exact(LANES);
174
175 for chunk in chunks.by_ref() {
176 let mut buf = [0u32; LANES];
178 for (b, &d) in buf.iter_mut().zip(chunk.iter()) {
179 prev = prev.wrapping_add(d.cast_signed());
180 *b = prev.cast_unsigned();
181 }
182 self.decode_chunk(buf, shift_vec, &mut out);
183 }
184
185 for &d in chunks.remainder() {
187 prev = prev.wrapping_add(d.cast_signed());
188 let coord = self.decode_one(prev.cast_unsigned());
189 out.push(coord.x);
190 out.push(coord.y);
191 }
192
193 dec.adjust_alloc(&out, alloc_size)?;
194 Ok(out)
195 }
196
197 #[inline]
202 fn decode_chunk(self, buf: [u32; LANES], shift_vec: u32x8, out: &mut Vec<i32>) {
203 let codes = u32x8::from(buf);
204 let codes_y = codes >> 1;
206
207 let mut x_vec = u32x8::ZERO;
208 let mut y_vec = u32x8::ZERO;
209
210 for i in 0..self.bits {
211 let bit_mask = u32x8::splat(1u32 << (2 * i));
213 x_vec |= (codes & bit_mask) >> i;
215 y_vec |= (codes_y & bit_mask) >> i;
216 }
217
218 let xs: [u32; LANES] = (x_vec - shift_vec).into();
219 let ys: [u32; LANES] = (y_vec - shift_vec).into();
220
221 for lane in 0..LANES {
222 out.push(xs[lane].cast_signed());
223 out.push(ys[lane].cast_signed());
224 }
225 }
226}
227
228#[cfg(test)]
229mod tests {
230
231 use super::*;
232 use crate::test_helpers::dec;
233
234 const fn c(x: i32, y: i32) -> Coord<i32> {
235 Coord::<i32> { x, y }
236 }
237
238 const fn p(shift: u32, bits: u32) -> CurveParams {
239 CurveParams { shift, bits }
240 }
241
242 fn spread_bits(mut tx: u32) -> u32 {
247 tx = (tx | (tx << 8)) & 0x00FF_00FF;
248 tx = (tx | (tx << 4)) & 0x0F0F_0F0F;
249 tx = (tx | (tx << 2)) & 0x3333_3333;
250 tx = (tx | (tx << 1)) & 0x5555_5555;
251 tx
252 }
253
254 fn compact_bits(mut tx: u32) -> u32 {
257 tx &= 0x5555_5555;
258 tx = (tx | (tx >> 1)) & 0x3333_3333;
259 tx = (tx | (tx >> 2)) & 0x0F0F_0F0F;
260 tx = (tx | (tx >> 4)) & 0x00FF_00FF;
261 tx = (tx | (tx >> 8)) & 0x0000_FFFF;
262 tx
263 }
264
265 #[test]
266 fn spread_then_compact_is_identity() {
267 for x in 0u32..=0xFFFF {
268 assert_eq!(compact_bits(spread_bits(x)), x, "round-trip failed for {x}");
269 }
270 }
271
272 #[test]
273 fn spread_bits_places_bit0_at_position0() {
274 assert_eq!(spread_bits(1), 1);
275 }
276
277 #[test]
278 fn spread_bits_places_bit1_at_position2() {
279 assert_eq!(spread_bits(2), 4);
280 }
281
282 #[test]
283 fn spread_bits_places_bit2_at_position4() {
284 assert_eq!(spread_bits(4), 16);
285 }
286
287 #[test]
288 fn origin_maps_to_zero() {
289 assert_eq!(morton_sort_key(c(0, 0), p(0, 16)), 0);
290 }
291
292 #[test]
293 fn x_axis_produces_even_bits() {
294 assert_eq!(morton_sort_key(c(1, 0), p(0, 16)), 1);
296 assert_eq!(morton_sort_key(c(2, 0), p(0, 16)), 4);
298 }
299
300 #[test]
301 fn y_axis_produces_odd_bits() {
302 assert_eq!(morton_sort_key(c(0, 1), p(0, 16)), 2);
304 assert_eq!(morton_sort_key(c(0, 2), p(0, 16)), 8);
306 }
307
308 #[test]
309 fn negative_coords_shift_correctly() {
310 assert_eq!(morton_sort_key(c(-1, -1), p(1, 16)), 0);
312 assert_eq!(morton_sort_key(c(-1, 0), p(1, 16)), 2);
314 }
315
316 #[test]
317 fn spatial_locality_z_order() {
318 let k00 = morton_sort_key(c(0, 0), p(0, 16));
320 let k10 = morton_sort_key(c(1, 0), p(0, 16));
321 let k01 = morton_sort_key(c(0, 1), p(0, 16));
322 let k11 = morton_sort_key(c(1, 1), p(0, 16));
323 assert!(k00 < k10);
324 assert!(k10 < k01);
325 assert!(k01 < k11);
326 }
327
328 #[test]
329 fn interleave_round_trips_via_deinterleave() {
330 for x in 0u32..16 {
332 for y in 0u32..16 {
333 let code = interleave_bits((x, y).into());
334 let mut rx = 0u32;
335 let mut ry = 0u32;
336 for bit in 0..16 {
337 rx |= ((code >> (2 * bit)) & 1) << bit;
338 ry |= ((code >> (2 * bit + 1)) & 1) << bit;
339 }
340 assert_eq!(rx, x, "x mismatch for ({x}, {y})");
341 assert_eq!(ry, y, "y mismatch for ({x}, {y})");
342 }
343 }
344 }
345
346 const NUM_BITS: u32 = 15;
349 const COORD_SHIFT: u32 = 1 << (NUM_BITS - 1); const MORTON: Morton = Morton {
351 bits: NUM_BITS,
352 shift: COORD_SHIFT,
353 };
354
355 #[must_use]
360 #[inline]
361 pub fn encode_morton_15(coord: Coord<u32>) -> u32 {
362 let mut code = 0u32;
363 for bit in 0..15 {
364 code |= ((coord.x >> bit) & 1) << (2 * bit);
365 code |= ((coord.y >> bit) & 1) << (2 * bit + 1);
366 }
367 code
368 }
369
370 #[test]
371 fn test_decode_morton_codes_empty() {
372 assert!(MORTON.decode_codes(&[], &mut dec()).unwrap().is_empty());
373 }
374
375 #[test]
376 fn test_decode_morton_codes_origin() {
377 let code = encode_morton_15((COORD_SHIFT, COORD_SHIFT).into());
379 let decoded = MORTON.decode_codes(&[code], &mut dec()).unwrap();
380 assert_eq!(decoded, [0, 0]);
381 }
382
383 #[test]
384 fn test_decode_morton_codes_known_values() {
385 let x: u32 = 1;
387 let y: u32 = 2;
388 let code = encode_morton_15((x, y).into());
389 let expected_x = x.cast_signed() - COORD_SHIFT.cast_signed();
390 let expected_y = y.cast_signed() - COORD_SHIFT.cast_signed();
391 let decoded = MORTON.decode_codes(&[code], &mut dec()).unwrap();
392 assert_eq!(decoded, [expected_x, expected_y]);
393 }
394
395 #[test]
396 fn test_decode_morton_codes_scalar_tail() {
397 let pairs: [Coord<u32>; _] = [(0, 1).into(), (2, 3).into(), (4, 5).into()];
399 let codes: Vec<u32> = pairs.iter().map(|&c| encode_morton_15(c)).collect();
400 let result = MORTON.decode_codes(&codes, &mut dec()).unwrap();
401 let expected = expected_coords(&pairs);
402 assert_eq!(result, expected);
403 }
404
405 #[test]
406 fn test_decode_morton_codes_full_simd_chunk() {
407 let pairs: [Coord<u32>; _] = [
409 (0, 0).into(),
410 (1, 0).into(),
411 (0, 1).into(),
412 (1, 1).into(),
413 (2, 3).into(),
414 (7, 5).into(),
415 (10, 9).into(),
416 (15, 15).into(),
417 ];
418 let codes: Vec<u32> = pairs.iter().map(|&c| encode_morton_15(c)).collect();
419 let result = MORTON.decode_codes(&codes, &mut dec()).unwrap();
420 let expected = expected_coords(&pairs);
421 assert_eq!(result, expected);
422 }
423
424 #[test]
425 fn test_decode_morton_codes_simd_plus_tail() {
426 let pairs: Vec<Coord<u32>> = (0..11u32)
428 .map(|i| (i * 3 % 100, i * 7 % 100).into())
429 .collect();
430 let codes: Vec<u32> = pairs.iter().map(|&c| encode_morton_15(c)).collect();
431 let result = MORTON.decode_codes(&codes, &mut dec()).unwrap();
432 let expected = expected_coords(&pairs);
433 assert_eq!(result, expected);
434 }
435
436 #[test]
439 fn test_decode_morton_delta_empty() {
440 assert!(MORTON.decode_delta(&[], &mut dec()).unwrap().is_empty());
441 }
442
443 #[test]
444 fn test_decode_morton_delta_identity_with_zero_deltas() {
445 let deltas = vec![0u32; 3];
447 let result = MORTON.decode_delta(&deltas, &mut dec()).unwrap();
448 let shift = -COORD_SHIFT.cast_signed();
449 assert_eq!(result, vec![shift, shift, shift, shift, shift, shift]);
450 }
451
452 #[test]
453 fn test_decode_morton_delta_matches_codes_after_prefix_sum() {
454 let pairs: Vec<Coord<u32>> = (0..11u32)
457 .map(|i| (i * 5 % 200, i * 9 % 200).into())
458 .collect();
459 let codes: Vec<u32> = pairs.iter().map(|&c| encode_morton_15(c)).collect();
460 let deltas = signed_deltas(&codes);
461
462 let from_codes = MORTON.decode_codes(&codes, &mut dec()).unwrap();
463 let from_deltas = MORTON.decode_delta(&deltas, &mut dec()).unwrap();
464 assert_eq!(from_codes, from_deltas);
465 }
466
467 #[test]
468 fn test_decode_morton_delta_scalar_tail() {
469 let codes: Vec<u32> = vec![
471 encode_morton_15((10, 20).into()),
472 encode_morton_15((30, 40).into()),
473 encode_morton_15((50, 60).into()),
474 ];
475 let deltas = signed_deltas(&codes);
476 let from_codes = MORTON.decode_codes(&codes, &mut dec()).unwrap();
477 let from_deltas = MORTON.decode_delta(&deltas, &mut dec()).unwrap();
478 assert_eq!(from_codes, from_deltas);
479 }
480
481 #[test]
482 fn test_decode_morton_delta_wrapping() {
483 let code_a = encode_morton_15((500, 300).into());
486 let code_b = encode_morton_15((10, 10).into()); let delta_b = code_b
488 .cast_signed()
489 .wrapping_sub(code_a.cast_signed())
490 .cast_unsigned();
491 assert_eq!(
492 MORTON.decode_delta(&[code_a, delta_b], &mut dec()).unwrap(),
493 MORTON.decode_codes(&[code_a, code_b], &mut dec()).unwrap()
494 );
495 }
496
497 fn expected_coords(pairs: &[Coord<u32>]) -> Vec<i32> {
499 pairs
500 .iter()
501 .flat_map(|&Coord { x, y }| {
502 [
503 x.cast_signed() - COORD_SHIFT.cast_signed(),
504 y.cast_signed() - COORD_SHIFT.cast_signed(),
505 ]
506 })
507 .collect()
508 }
509
510 fn signed_deltas(codes: &[u32]) -> Vec<u32> {
512 let mut prev = 0i32;
513 codes
514 .iter()
515 .map(|&c| {
516 let delta = c.cast_signed().wrapping_sub(prev).cast_unsigned();
517 prev = c.cast_signed();
518 delta
519 })
520 .collect()
521 }
522}