1use crypto_primitives::SemiringConfig;
7use zinc_poly::utils::next_mle_eval;
8
9#[allow(clippy::arithmetic_side_effects)]
19pub fn eval_shift_predicate<C: SemiringConfig>(
20 cfg: &C,
21 x: &[C::Element],
22 y: &[C::Element],
23 c: usize,
24) -> C::Element {
25 let m = x.len();
26 assert_eq!(y.len(), m);
27
28 if c == 0 {
30 return eval_eq_poly(cfg, x, y);
31 }
32
33 if c == 1 {
35 return next_mle_eval(cfg, x, y);
36 }
37
38 assert!(c < (1usize << m), "shift c must satisfy c < 2^m");
39 let k = (2 * c).next_power_of_two().trailing_zeros() as usize;
41 if k >= m {
42 return eval_shift_small(cfg, x, y, c, m);
43 }
44
45 let (x_lo, x_hi) = x.split_at(k);
47 let (y_lo, y_hi) = y.split_at(k);
48
49 let l0 = eval_l0(cfg, x_lo, y_lo, c, k);
50 let l1 = eval_l1(cfg, x_lo, y_lo, c, k);
51 let eq = eval_eq_poly(cfg, x_hi, y_hi);
52 let next = next_mle_eval(cfg, x_hi, y_hi);
53
54 cfg.add(&cfg.mul(&l0, &eq), &cfg.mul(&l1, &next))
55}
56
57pub(crate) fn eval_eq_poly<C: SemiringConfig>(
61 cfg: &C,
62 u: &[C::Element],
63 v: &[C::Element],
64) -> C::Element {
65 let one = cfg.one();
66 u.iter()
67 .zip(v.iter())
68 .map(|(u_i, v_i)| {
69 cfg.add(
70 &cfg.mul(u_i, v_i),
71 &cfg.mul(&cfg.sub(&one, u_i), &cfg.sub(&one, v_i)),
72 )
73 })
74 .fold(one.clone(), |acc, term| cfg.mul(&acc, &term))
75}
76
77pub(crate) fn eval_delta<C: SemiringConfig>(
84 cfg: &C,
85 u: &[C::Element],
86 a: usize,
87 k: usize,
88) -> C::Element {
89 let one = cfg.one();
90 let mut result = one.clone();
91 for (i, u) in u.iter().take(k).enumerate() {
92 let bit = (a >> i) & 1;
93 if bit == 1 {
94 cfg.mul_assign(&mut result, u);
95 } else {
96 let term = cfg.sub(&one, u);
97 cfg.mul_assign(&mut result, &term);
98 }
99 }
100 result
101}
102
103#[allow(clippy::arithmetic_side_effects)]
110pub(crate) fn eval_l0<C: SemiringConfig>(
111 cfg: &C,
112 x_lo: &[C::Element],
113 y_lo: &[C::Element],
114 c: usize,
115 k: usize,
116) -> C::Element {
117 let upper = (1 << k) - c;
118 (0..upper).fold(cfg.zero(), |acc, a| {
119 cfg.add(
120 &acc,
121 &cfg.mul(
122 &eval_delta(cfg, x_lo, a, k),
123 &eval_delta(cfg, y_lo, a + c, k),
124 ),
125 )
126 })
127}
128
129#[allow(clippy::arithmetic_side_effects)]
135pub(crate) fn eval_l1<C: SemiringConfig>(
136 cfg: &C,
137 x_lo: &[C::Element],
138 y_lo: &[C::Element],
139 c: usize,
140 k: usize,
141) -> C::Element {
142 let two_k = 1 << k;
143 ((two_k - c)..two_k).fold(cfg.zero(), |acc, a| {
144 cfg.add(
145 &acc,
146 &cfg.mul(
147 &eval_delta(cfg, x_lo, a, k),
148 &eval_delta(cfg, y_lo, a + c - two_k, k),
149 ),
150 )
151 })
152}
153
154#[allow(clippy::arithmetic_side_effects)]
158fn eval_shift_small<C: SemiringConfig>(
159 cfg: &C,
160 x: &[C::Element],
161 y: &[C::Element],
162 c: usize,
163 m: usize,
164) -> C::Element {
165 let upper = (1 << m) - c;
166 (0..upper).fold(cfg.zero(), |acc, a| {
167 cfg.add(
168 &acc,
169 &cfg.mul(&eval_delta(cfg, x, a, m), &eval_delta(cfg, y, a + c, m)),
170 )
171 })
172}
173
174#[cfg(test)]
175mod tests {
176 use super::*;
177 use crate::test_utils::test_config;
178 use crypto_primitives::{
179 ProjectElementWithConfig,
180 crypto_bigint_monty::{MontyField, MontyFieldElement},
181 };
182 use rand::prelude::*;
183 use zinc_poly::utils::{build_eq_x_r, build_next_c_r_mle};
184
185 type F = MontyField<4>;
186 type E = MontyFieldElement<4>;
187
188 fn to_bin(cfg: &F, val: usize, bit: usize) -> E {
190 if (val >> bit) & 1 == 1 {
191 cfg.one()
192 } else {
193 cfg.zero()
194 }
195 }
196
197 fn rand_field(cfg: &F, rng: &mut impl Rng) -> E {
198 cfg.project(&rng.random::<u32>())
199 }
200
201 #[test]
204 fn test_shift_predicate_boolean() {
205 let cfg = test_config();
206 let m = 4;
207 let n = 1usize << m;
208
209 for c in [1, 2, 5] {
210 for a in 0..n {
211 for b in 0..n {
212 let x: Vec<E> = (0..m).map(|i| to_bin(&cfg, a, i)).collect();
213 let y: Vec<E> = (0..m).map(|i| to_bin(&cfg, b, i)).collect();
214 let val = eval_shift_predicate(&cfg, &x, &y, c);
215
216 if b == a + c && a + c < n {
217 assert_eq!(val, cfg.one(), "S_c({a},{b}) should be 1 for c={c}");
218 } else {
219 assert_eq!(val, cfg.zero(), "S_c({a},{b}) should be 0 for c={c}");
220 }
221 }
222 }
223 }
224 }
225
226 #[test]
228 fn test_next_boolean() {
229 let cfg = test_config();
230 let m = 4;
231 let n = 1usize << m;
232 for a in 0..n {
233 for b in 0..n {
234 let u: Vec<E> = (0..m).map(|i| to_bin(&cfg, a, i)).collect();
235 let v: Vec<E> = (0..m).map(|i| to_bin(&cfg, b, i)).collect();
236 let val = next_mle_eval(&cfg, &u, &v);
237
238 if b == a + 1 && a + 1 < n {
239 assert_eq!(val, cfg.one(), "Next({a},{b}) should be 1");
240 } else {
241 assert_eq!(val, cfg.zero(), "Next({a},{b}) should be 0");
242 }
243 }
244 }
245 }
246
247 #[test]
251 fn test_shift_predicate_vs_prover_mle() {
252 let cfg = test_config();
253 let mut rng = rand::rng();
254 let m = 4;
255 let n = 1usize << m;
256 let c = 3;
257
258 let r: Vec<E> = (0..m).map(|_| rand_field(&cfg, &mut rng)).collect();
259 let next_c = build_next_c_r_mle(&cfg, &r, c).unwrap();
260
261 for b in 0..n {
262 let b_bin: Vec<E> = (0..m).map(|i| to_bin(&cfg, b, i)).collect();
263 let val = eval_shift_predicate(&cfg, &r, &b_bin, c);
264 assert_eq!(
265 val, next_c.evaluations[b],
266 "S_{c}(r, bin({b})) mismatch with prover MLE"
267 );
268 }
269 }
270
271 #[test]
275 fn test_shift_predicate_random_points() {
276 let cfg = test_config();
277 let mut rng = rand::rng();
278 let m = 4;
279 let c = 3;
280
281 for _ in 0..8 {
282 let r: Vec<E> = (0..m).map(|_| rand_field(&cfg, &mut rng)).collect();
283 let y: Vec<E> = (0..m).map(|_| rand_field(&cfg, &mut rng)).collect();
284
285 let next_c = build_next_c_r_mle(&cfg, &r, c).unwrap();
286 let eq_y = build_eq_x_r(&cfg, &y).unwrap();
287 let rhs = next_c
288 .evaluations
289 .iter()
290 .zip(eq_y.evaluations.iter())
291 .fold(cfg.zero(), |acc, (ni, ei)| cfg.add(&acc, &cfg.mul(ni, ei)));
292 let lhs = eval_shift_predicate(&cfg, &r, &y, c);
293
294 assert_eq!(lhs, rhs, "random-point MLE mismatch");
295 }
296 }
297
298 #[test]
302 fn test_fast_paths_and_multi_c() {
303 let cfg = test_config();
304 let mut rng = rand::rng();
305 let m = 4;
306 let n = 1usize << m;
307
308 for c in [0, 1, 2, 5, 7] {
309 let r: Vec<E> = (0..m).map(|_| rand_field(&cfg, &mut rng)).collect();
310 let next_c = build_next_c_r_mle(&cfg, &r, c).unwrap();
311
312 for b in 0..n {
314 let b_bin: Vec<E> = (0..m).map(|i| to_bin(&cfg, b, i)).collect();
315 let val = eval_shift_predicate(&cfg, &r, &b_bin, c);
316 assert_eq!(
317 val, next_c.evaluations[b],
318 "S_{c}(r, bin({b})) mismatch with prover MLE"
319 );
320 }
321
322 for _ in 0..4 {
324 let y: Vec<E> = (0..m).map(|_| rand_field(&cfg, &mut rng)).collect();
325 let eq_y = build_eq_x_r(&cfg, &y).unwrap();
326 let rhs = next_c
327 .evaluations
328 .iter()
329 .zip(eq_y.evaluations.iter())
330 .fold(cfg.zero(), |acc, (ni, ei)| cfg.add(&acc, &cfg.mul(ni, ei)));
331 let lhs = eval_shift_predicate(&cfg, &r, &y, c);
332 assert_eq!(lhs, rhs, "random-point MLE mismatch for c={c}");
333 }
334 }
335 }
336
337 #[test]
339 fn test_shift_predicate_boundary() {
340 let cfg = test_config();
341 let m = 3;
342 let n = 1usize << m; for c in [n / 2, n - 1] {
345 for a in 0..n {
347 for b in 0..n {
348 let x: Vec<E> = (0..m).map(|i| to_bin(&cfg, a, i)).collect();
349 let y: Vec<E> = (0..m).map(|i| to_bin(&cfg, b, i)).collect();
350 let val = eval_shift_predicate(&cfg, &x, &y, c);
351
352 if b == a + c && a + c < n {
353 assert_eq!(val, cfg.one(), "S_{c}(bin({a}), bin({b})) should be 1");
354 } else {
355 assert_eq!(val, cfg.zero(), "S_{c}(bin({a}), bin({b})) should be 0");
356 }
357 }
358 }
359
360 let mut rng = rand::rng();
362 let r: Vec<E> = (0..m).map(|_| rand_field(&cfg, &mut rng)).collect();
363 let next_c = build_next_c_r_mle(&cfg, &r, c).unwrap();
364 let zero = cfg.zero();
365
366 for b in 0..c {
368 assert_eq!(
369 next_c.evaluations[b], zero,
370 "next_c[{b}] should be zero for c={c}"
371 );
372 }
373 let nonzero_count = next_c.evaluations[c..]
375 .iter()
376 .filter(|e| **e != zero)
377 .count();
378 assert_eq!(
379 nonzero_count,
380 n - c,
381 "expected {} nonzero entries for c={c}",
382 n - c
383 );
384 }
385 }
386
387 #[test]
390 fn test_prover_mle_inner_product() {
391 let cfg = test_config();
392 let mut rng = rand::rng();
393 let m = 4;
394 let n = 1usize << m;
395
396 for c in [1, 2, 3, 7] {
397 let v: Vec<E> = (0..n).map(|_| rand_field(&cfg, &mut rng)).collect();
398 let r: Vec<E> = (0..m).map(|_| rand_field(&cfg, &mut rng)).collect();
399
400 let eq_r = build_eq_x_r(&cfg, &r).unwrap();
402 let expected = (c..n).fold(cfg.zero(), |acc, b| {
403 cfg.add(&acc, &cfg.mul(&eq_r.evaluations[b - c], &v[b]))
404 });
405
406 let next_c = build_next_c_r_mle(&cfg, &r, c).unwrap();
408 let got = next_c
409 .evaluations
410 .iter()
411 .zip(v.iter())
412 .fold(cfg.zero(), |acc, (ni, vi)| cfg.add(&acc, &cfg.mul(vi, ni)));
413
414 assert_eq!(got, expected, "prover MLE inner product mismatch for c={c}");
415 }
416 }
417}