1use std::iter;
4
5use crate::{Field, PackedField, Underlier, field::FieldOps};
6
7pub trait FieldFn<F: Field> {
14 fn call<E: FieldOps<Scalar = F> + From<F>>(&self, inputs: &[E]) -> E;
19
20 fn call_native(&self, inputs: &[F]) -> F {
27 self.call::<F>(inputs)
28 }
29}
30
31pub fn powers<F: FieldOps>(val: F) -> impl Iterator<Item = F> {
33 iter::successors(Some(F::one()), move |power| Some(power.clone() * val.clone()))
34}
35
36pub fn expand_subset_sums_array<P: PackedField, const N: usize, const N_EXP2: usize>(
71 elems: [P; N],
72) -> [P; N_EXP2] {
73 const {
74 assert!(N_EXP2 == 1 << N, "N_EXP2 must equal 2^N");
75 }
76
77 let mut expanded = [P::zero(); N_EXP2];
78 for (i, elem_i) in elems.into_iter().enumerate() {
79 let span = &mut expanded[..1 << (i + 1)];
80 let (lo_half, hi_half) = span.split_at_mut(1 << i);
81 for (lo_half_i, hi_half_i) in iter::zip(lo_half, hi_half) {
82 *hi_half_i = *lo_half_i + elem_i;
83 }
84 }
85 expanded
86}
87
88pub fn expand_subset_sums<P: PackedField>(elems: &[P]) -> Vec<P> {
100 assert!(elems.len() < usize::BITS as usize); let mut expanded = vec![P::zero(); 1 << elems.len()];
103 for (i, &elem_i) in elems.iter().enumerate() {
104 let (lo_half, hi_half) = expanded[..1 << (i + 1)].split_at_mut(1 << i);
105 for (lo_half_i, hi_half_i) in iter::zip(lo_half, hi_half) {
106 *hi_half_i = *lo_half_i + elem_i;
107 }
108 }
109 expanded
110}
111
112pub fn expand_subset_products<P: PackedField>(elems: &[P]) -> Vec<P> {
128 assert!(elems.len() < usize::BITS as usize); let mut expanded = vec![P::one(); 1 << elems.len()];
131 for (i, &elem_i) in elems.iter().enumerate() {
132 let (lo_half, hi_half) = expanded[..1 << (i + 1)].split_at_mut(1 << i);
133 for (lo_half_i, hi_half_i) in iter::zip(lo_half, hi_half) {
134 *hi_half_i = *lo_half_i * elem_i;
135 }
136 }
137 expanded
138}
139
140pub fn expand_subset_xors<U: Underlier, const N: usize, const N_EXP2: usize>(
150 elems: [U; N],
151) -> [U; N_EXP2] {
152 const {
153 assert!(N_EXP2 == 1 << N, "N_EXP2 must equal 2^N");
154 }
155
156 let mut expanded = [U::ZERO; N_EXP2];
157 for (i, elem_i) in elems.into_iter().enumerate() {
158 let span = &mut expanded[..1 << (i + 1)];
159 let (lo_half, hi_half) = span.split_at_mut(1 << i);
160 for (lo_half_i, hi_half_i) in iter::zip(lo_half, hi_half) {
161 *hi_half_i = *lo_half_i ^ elem_i;
162 }
163 }
164 expanded
165}
166
167#[cfg(test)]
168mod tests {
169 use std::array;
170
171 use proptest::prelude::*;
172 use rand::{SeedableRng, rngs::StdRng};
173
174 use super::*;
175 use crate::{Ghash128b, Random};
176
177 #[test]
178 fn test_powers_against_pow() {
179 let generator = Ghash128b::MULTIPLICATIVE_GENERATOR;
181 let power_values: Vec<_> = powers(generator).take(10).collect();
182
183 for (i, power) in power_values.iter().enumerate() {
184 assert_eq!(*power, generator.pow(i as u64));
185 }
186 }
187
188 type F = Ghash128b;
189
190 fn check_subset_sums<const N: usize, const N_EXP2: usize>(seed: u64, index: usize) {
193 let mut rng = StdRng::seed_from_u64(seed);
194 let elems: [F; N] = array::from_fn(|_| F::random(&mut rng));
195
196 let result = expand_subset_sums_array::<_, N, N_EXP2>(elems);
197 assert_eq!(result.len(), N_EXP2);
198
199 let index = index % N_EXP2;
201 let mut expected = F::ZERO;
202 for (bit_pos, &elem) in elems.iter().enumerate() {
203 if (index >> bit_pos) & 1 == 1 {
204 expected += elem;
205 }
206 }
207
208 assert_eq!(
209 result[index], expected,
210 "index {index} should hold the subset sum for its binary representation"
211 );
212 }
213
214 proptest! {
215 #[test]
216 fn test_expand_subset_sums_array_correctness(
217 n in 0usize..=8, index in 0usize..256, ) {
220 match n {
223 0 => check_subset_sums::<0, 1>(n as u64, index),
224 1 => check_subset_sums::<1, 2>(n as u64, index),
225 2 => check_subset_sums::<2, 4>(n as u64, index),
226 3 => check_subset_sums::<3, 8>(n as u64, index),
227 4 => check_subset_sums::<4, 16>(n as u64, index),
228 5 => check_subset_sums::<5, 32>(n as u64, index),
229 6 => check_subset_sums::<6, 64>(n as u64, index),
230 7 => check_subset_sums::<7, 128>(n as u64, index),
231 8 => check_subset_sums::<8, 256>(n as u64, index),
232 _ => unreachable!("n is constrained to 0..=8"),
233 }
234 }
235 }
236 proptest! {
237 #[test]
238 fn expand_subset_products_selects_the_product_over_set_bits(seed: u64, n in 0usize..=8) {
239 let mut rng = StdRng::seed_from_u64(seed);
240 let elems = (0..n)
241 .map(|_| F::random(&mut rng))
242 .collect::<Vec<_>>();
243
244 let expanded = expand_subset_products(&elems);
245 prop_assert_eq!(expanded.len(), 1 << n);
246
247 for (mask, &entry) in expanded.iter().enumerate() {
249 let expected = (0..n)
250 .filter(|i| (mask >> i) & 1 == 1)
251 .fold(F::ONE, |acc, i| acc * elems[i]);
252 prop_assert_eq!(entry, expected);
253 }
254 }
255 }
256}