Skip to main content

rostl_sort/
compaction.rs

1//! Implements oblivious compaction algorithms.
2//!
3
4use assume::assume;
5use rostl_primitives::{
6  traits::{Cmov, CswapIndex},
7  utils::get_smaller_or_equal_power_of_two,
8};
9// use rostl_primitives::indexable::{Indexable, Length};
10
11/// Computes the prefix sum of valid elements in arr.
12/// # Behavior
13/// * returns `ret` - the prefix sum array of length `arr.len() + 1`
14/// # Oblivious
15/// * Fully data-independent memory access pattern.
16/// * Leaks: `arr.len()` - the full length of the original array
17pub fn compute_prefix_sum<T, F>(arr: &[T], is_dummy: F) -> Vec<usize>
18where
19  F: Fn(&T) -> bool,
20{
21  let size = arr.len();
22  let mut sarr = vec![0; size + 1];
23  for i in 0..size {
24    let mut adder = 1usize;
25    adder.cmov(&0, is_dummy(&arr[i]));
26    sarr[i + 1] = sarr[i] + adder;
27  }
28  sarr
29}
30
31/// Stably compacts an array `arr` of length n in place using nlogn oblivious compaction.
32/// Uses `https://arxiv.org/pdf/1103.5102`
33/// # Behavior
34/// * returns `ret` - the number of non-dummy elements (new real length of the array).
35/// * first `ret` elements of `arr` are the non-dummy elements in the same order as they were in the original array.
36/// * the rest of the elements in `arr` are dummy elements.
37/// # Oblivious
38/// * Fully data-independent memory access pattern.
39/// * Leaks: `arr.len()` - the full length of the original array
40#[deprecated(note = "use compact instead, it's faster")]
41pub fn compact_goodrich<T, F>(arr: &mut [T], is_dummy: F) -> usize
42where
43  F: Fn(&T) -> bool,
44  T: Cmov + Copy,
45{
46  if arr.is_empty() {
47    return 0;
48  }
49  let l2len = arr.len().next_power_of_two().trailing_zeros() as usize;
50  let mut csum = vec![0; arr.len()];
51  let mut dummy_count = 0;
52
53  csum[0] = 0;
54  let pred = is_dummy(&arr[0]);
55  dummy_count.cmov(&1, pred);
56  for i in 1..arr.len() {
57    csum[i] = 0;
58    let pred = is_dummy(&arr[i]);
59    dummy_count.cmov(&(dummy_count + 1), pred);
60    csum[i].cmov(&(dummy_count), !pred);
61  }
62  let ret = arr.len() - dummy_count;
63
64  for i in 0..l2len {
65    let offset = 1 << i;
66    for j in 0..(arr.len() - offset) {
67      let a = j;
68      let b = j + offset;
69      let pred = (csum[b] & offset) != 0;
70      arr.cswap(a, b, pred);
71      let newacsum = csum[b].wrapping_sub(offset);
72      csum[a].cmov(&newacsum, pred);
73      csum[b].cmov(&0, pred);
74    }
75  }
76
77  ret
78}
79
80/// Compacts arr, as marked by the prefixsum payload, to an offset z.
81/// # Requires
82/// * `arr.len()` is a power of two.
83/// * `payload.len() == arr.len() + 1`
84/// * payload is a prefix sum of valid elements in arr.
85/// * `0 <= z < arr.len()`
86fn compact_payload_offset<T>(arr: &mut [T], payload: &[usize], z: usize)
87where
88  T: Cmov + Copy,
89{
90  assume!(unsafe: arr.len()+1 == payload.len());
91  let n = arr.len();
92  let half_n = n / 2;
93  let m = payload[half_n] - payload[0];
94  if n == 2 {
95    let should_swap = ((!m) & (payload[2] - payload[1])) != z;
96    arr.cswap(0, 1, should_swap);
97    return;
98  }
99  let zleft = z % half_n;
100  let zright = (z + m) % half_n;
101  compact_payload_offset(&mut arr[..half_n], &payload[..half_n + 1], zleft);
102  compact_payload_offset(&mut arr[half_n..], &payload[half_n..], zright);
103
104  let s_a = zleft + m >= half_n;
105  let s_b = z >= half_n;
106  let s = s_a ^ s_b;
107
108  for i in 0..half_n {
109    let left = i;
110    let right = i + half_n;
111    let cond = s ^ (i >= zright);
112    assume!(unsafe: left < arr.len());
113    assume!(unsafe: right < arr.len());
114    arr.cswap(left, right, cond);
115  }
116}
117
118/// Stably compacts an array `arr` of length n using oblivious compaction.
119/// The payload array `payload` is the prefix sum of valid elements.
120/// Uses `https://eprint.iacr.org/2022/1333.pdf`
121/// # Requires
122/// * `payload.len() == arr.len() + 1`
123/// * payload is a prefix sum of valid elements in arr.
124/// * first elements of `arr` are the non-dummy elements in the same order as they were in the original array.
125/// * the rest of the elements in `arr` are the dummy elements in no particular order.
126/// # Oblivious
127/// * Fully data-independent memory access pattern.
128/// * Leaks: `arr.len()` - the full length of the original array
129pub fn compact_payload<T>(arr: &mut [T], payload: &[usize])
130where
131  T: Cmov + Copy,
132{
133  assume!(unsafe: arr.len() + 1 == payload.len());
134  let n = arr.len();
135  if n <= 1 {
136    return;
137  }
138
139  let n1 = get_smaller_or_equal_power_of_two(n);
140  let n2 = n - n1;
141
142  if n2 == 0 {
143    compact_payload_offset(arr, payload, 0);
144    return;
145  }
146
147  let m = payload[n2] - payload[0];
148  compact_payload(arr[..n2].as_mut(), &payload[..n2 + 1]);
149  compact_payload_offset(arr[n2..].as_mut(), &payload[n2..], (n1 - n2 + m) % n1);
150
151  for i in 0..n2 {
152    let left = i;
153    let right = i + n1;
154    assume!(unsafe: left < arr.len());
155    assume!(unsafe: right < arr.len());
156    arr.cswap(left, right, i >= m);
157  }
158}
159
160/// Stably compacts an array `arr` of length n using oblivious compaction.
161/// The payload array `payload` is the prefix sum of valid elements.
162/// Uses `https://eprint.iacr.org/2022/1333.pdf`
163/// # Requires
164/// # Behavior
165/// * returns `ret` - the number of non-dummy elements (new real length of the array).
166/// * first `ret` elements of `arr` are the non-dummy elements in the same order as they were in the original array.
167/// * the rest of the elements in `arr` are the dummy elements in no particular order.
168/// # Oblivious
169/// * Fully data-independent memory access pattern.
170/// * Leaks: `arr.len()` - the full length of the original array
171/// # Returns the number of non-dummy elements in the array after compaction.
172pub fn compact<T, F>(arr: &mut [T], is_dummy: F) -> usize
173where
174  F: Fn(&T) -> bool,
175  T: Cmov + Copy,
176{
177  let payload = compute_prefix_sum(arr, is_dummy);
178  compact_payload(arr, &payload);
179  payload[payload.len() - 1]
180}
181
182fn distribute_payload_offset<T>(arr: &mut [T], payload: &[usize], z: usize)
183where
184  T: Cmov + Copy,
185{
186  assume!(unsafe: arr.len()+1 == payload.len());
187  let n = arr.len();
188  let half_n = n / 2;
189  let m = payload[half_n] - payload[0];
190  if n == 2 {
191    let should_swap = ((!m) & (payload[2] - payload[1])) != z;
192    arr.cswap(0, 1, should_swap);
193    return;
194  }
195  let zleft = z % half_n;
196  let zright = (z + m) % half_n;
197  let s_a = zleft + m >= half_n;
198  let s_b = z >= half_n;
199  let s = s_a ^ s_b;
200
201  for i in 0..half_n {
202    let left = i;
203    let right = i + half_n;
204    let cond = s ^ (i >= zright);
205    assume!(unsafe: left < arr.len());
206    assume!(unsafe: right < arr.len());
207    arr.cswap(left, right, cond);
208  }
209  distribute_payload_offset(&mut arr[..half_n], &payload[..half_n + 1], zleft);
210  distribute_payload_offset(&mut arr[half_n..], &payload[half_n..], zright);
211}
212
213/// Distributes the elements of arr according to the prefix sum payload (reverse of compaction for the same payload).
214/// The payload array `payload` is the prefix sum of valid elements.
215/// Uses `https://eprint.iacr.org/2022/1333.pdf`
216pub fn distribute_payload<T>(arr: &mut [T], payload: &[usize])
217where
218  T: Cmov + Copy,
219{
220  assume!(unsafe: arr.len() + 1 == payload.len());
221  let n = arr.len();
222  if n <= 1 {
223    return;
224  }
225
226  let n1 = get_smaller_or_equal_power_of_two(n);
227  let n2 = n - n1;
228
229  if n2 == 0 {
230    distribute_payload_offset(arr, payload, 0);
231    return;
232  }
233
234  let m = payload[n2] - payload[0];
235
236  for i in 0..n2 {
237    let left = i;
238    let right = i + n1;
239    assume!(unsafe: left < arr.len());
240    assume!(unsafe: right < arr.len());
241    arr.cswap(left, right, i >= m);
242  }
243
244  distribute_payload(arr[..n2].as_mut(), &payload[..n2 + 1]);
245  distribute_payload_offset(arr[n2..].as_mut(), &payload[n2..], (n1 - n2 + m) % n1);
246}
247
248#[cfg(test)]
249#[allow(deprecated)]
250mod tests {
251  use rand::Rng;
252
253  use super::*;
254
255  #[test]
256  fn test_compact() {
257    let mut arr = [1, 2, 3, 4, 5];
258    let new_len = compact(&mut arr, |x| *x % 2 == 0);
259    assert_eq!(new_len, 3);
260    assert_eq!(&arr[..new_len], &[1, 3, 5]);
261
262    let mut arr = [1, 2, 3, 4, 5];
263    compact_goodrich(&mut arr, |x| *x % 2 == 0);
264    assert_eq!(&arr[..3], &[1, 3, 5]);
265  }
266
267  #[test]
268  fn test_small() {
269    let mut arr: Vec<i32> = vec![1];
270    let new_len = compact(&mut arr, |x| *x % 2 == 0);
271    assert_eq!(new_len, 1);
272    assert_eq!(&arr[..new_len], &[1]);
273    let mut arr: Vec<i32> = vec![1];
274    compact_goodrich(&mut arr, |x| *x % 2 == 0);
275    assert_eq!(&arr[..1], &[1]);
276
277    let mut arr: Vec<i32> = vec![2];
278    let new_len = compact(&mut arr, |x| *x % 2 == 0);
279    assert_eq!(new_len, 0);
280    assert_eq!(&arr[..new_len], &[]);
281
282    let mut arr: Vec<i32> = vec![1, 2];
283    let new_len = compact(&mut arr, |x| *x % 2 == 0);
284    assert_eq!(new_len, 1);
285    assert_eq!(&arr[..new_len], &[1]);
286    let mut arr: Vec<i32> = vec![1, 2];
287    compact_goodrich(&mut arr, |x| *x % 2 == 0);
288    assert_eq!(&arr[..1], &[1]);
289
290    let mut arr: Vec<i32> = vec![];
291    let new_len = compact(&mut arr, |x| *x % 2 == 0);
292    assert_eq!(new_len, 0);
293    assert_eq!(&arr[..new_len], &[]);
294    let mut arr: Vec<i32> = vec![];
295    compact_goodrich(&mut arr, |x| *x % 2 == 0);
296    assert_eq!(&arr[..0], &[]);
297  }
298
299  #[test]
300  fn test_many_sizes() {
301    // Picks a random array size and fills with random values and checks if it's correct via a non oblivious comparison
302    let mut rng = rand::rng();
303    for _i in 0..100 {
304      let size = rng.random_range(0..2050);
305      let arr: Vec<i32> = (0..size).map(|_| rng.random_range(0..100)).collect();
306      let mut arr1 = arr.clone();
307      let new_len = compact(&mut arr1, |x| *x % 2 == 0);
308      for itm in arr1.iter().take(new_len) {
309        assert!(itm % 2 != 0);
310      }
311      for itm in arr1.iter().skip(new_len) {
312        assert!(itm % 2 == 0);
313      }
314      let mut arr2 = arr.clone();
315      compact_goodrich(&mut arr2, |x| *x % 2 == 0);
316      for itm in arr2.iter().take(new_len) {
317        assert!(itm % 2 != 0);
318      }
319      for itm in arr2.iter().skip(new_len) {
320        assert!(itm % 2 == 0);
321      }
322    }
323  }
324
325  #[test]
326  fn test_distribute() {
327    let mut arr = [1, 3, 5, 0, 2, 4];
328    let payload = [0, 1, 2, 3, 3, 4, 5];
329    distribute_payload(&mut arr, &payload);
330    assert_eq!(&arr, &[1, 3, 5, 4, 0, 2]);
331
332    let mut arr = [1, 2, 3, 4, 5];
333    let payload = [0, 1, 1, 2, 2, 3];
334    compact_payload(&mut arr, &payload);
335    assert_eq!(&arr[..3], &[1, 3, 5]);
336    distribute_payload(&mut arr, &payload);
337    assert_eq!(&arr, &[1, 2, 3, 4, 5]);
338  }
339
340  #[test]
341  fn test_distribute_after_compact_rands() {
342    let mut rng = rand::rng();
343    for _i in 0..100 {
344      let size = rng.random_range(0..2050);
345      let arr: Vec<i32> = (0..size).map(|_| rng.random_range(0..100)).collect();
346      let mut arr1 = arr.clone();
347      let mut payload = vec![0; size + 1];
348      for i in 0..size {
349        let mut adder = 1usize;
350        adder.cmov(&0, arr[i] % 2 == 0);
351        payload[i + 1] = payload[i] + adder;
352      }
353      compact_payload(&mut arr1, &payload);
354      distribute_payload(&mut arr1, &payload);
355      assert_eq!(&arr1, &arr);
356    }
357  }
358}