use group::ff::PrimeField;
use sigma_proofs::errors::Error;
use subtle::Choice;
pub fn bit_decomp_vartime<S: PrimeField>(mut s: S) -> Option<(u128, u32)> {
let mut val = 0u128;
let mut bitnum = 0u32;
let mut bitval = 1u128; while bitnum < 127 && !s.is_zero_vartime() {
if s.is_odd().into() {
val += bitval;
s -= S::ONE;
}
bitnum += 1;
bitval <<= 1;
s *= S::TWO_INV;
}
if s.is_zero_vartime() {
Some((val, bitnum))
} else {
None
}
}
pub fn bit_decomp<S: PrimeField>(mut s: S, nbits: u32) -> Vec<Choice> {
let mut bits = Vec::with_capacity(nbits as usize);
let mut bitnum = 0u32;
while bitnum < nbits && bitnum < 127 {
let lowbit = s.is_odd();
s -= S::conditional_select(&S::ZERO, &S::ONE, lowbit);
s *= S::TWO_INV;
bits.push(lowbit);
bitnum += 1;
}
bits
}
pub fn bitrep_scalars_vartime<S: PrimeField>(upper: S) -> Result<Vec<S>, Error> {
let (upper_val, mut nbits) = bit_decomp_vartime(upper).ok_or(Error::VerificationFailure)?;
if nbits < 2 {
return Err(Error::VerificationFailure);
}
if upper_val == 1u128 << (nbits - 1) {
nbits -= 1;
}
Ok((0..nbits)
.map(|i| {
if i < nbits - 1 {
S::from_u128(1u128 << i)
} else {
S::from_u128(upper_val - (1u128 << (nbits - 1)))
}
})
.collect())
}
pub fn compute_bitrep<S: PrimeField>(mut x: S, bitrep_scalars: &[S]) -> Vec<Choice> {
let nbits: u32 = bitrep_scalars.len().try_into().unwrap();
let x_raw_bits = bit_decomp(x, nbits);
let high_bit = x_raw_bits[(nbits as usize) - 1];
x -= S::conditional_select(&S::ZERO, &bitrep_scalars[(nbits as usize) - 1], high_bit);
let mut x_bits = bit_decomp(x, nbits - 1);
x_bits.push(high_bit);
x_bits
}
#[cfg(test)]
mod tests {
use super::*;
use curve25519_dalek::scalar::Scalar;
use std::ops::Neg;
use subtle::ConditionallySelectable;
fn bit_decomp_tester(s: Scalar, nbits: u32, expect_bitstr: &str) {
assert_eq!(
bit_decomp(s, nbits)
.into_iter()
.map(|c| char::from(u8::conditional_select(&b'0', &b'1', c)))
.collect::<String>(),
expect_bitstr
);
}
#[test]
fn bit_decomp_test() {
assert_eq!(bit_decomp_vartime(Scalar::from(0u32)), Some((0, 0)));
assert_eq!(bit_decomp_vartime(Scalar::from(1u32)), Some((1, 1)));
assert_eq!(bit_decomp_vartime(Scalar::from(2u32)), Some((2, 2)));
assert_eq!(bit_decomp_vartime(Scalar::from(3u32)), Some((3, 2)));
assert_eq!(bit_decomp_vartime(Scalar::from(4u32)), Some((4, 3)));
assert_eq!(bit_decomp_vartime(Scalar::from(5u32)), Some((5, 3)));
assert_eq!(bit_decomp_vartime(Scalar::from(6u32)), Some((6, 3)));
assert_eq!(bit_decomp_vartime(Scalar::from(7u32)), Some((7, 3)));
assert_eq!(bit_decomp_vartime(Scalar::from(8u32)), Some((8, 4)));
assert_eq!(bit_decomp_vartime(Scalar::from(1u32).neg()), None);
assert_eq!(
bit_decomp_vartime(Scalar::from((1u128 << 127) - 2)),
Some(((i128::MAX - 1) as u128, 127))
);
assert_eq!(
bit_decomp_vartime(Scalar::from((1u128 << 127) - 1)),
Some((i128::MAX as u128, 127))
);
assert_eq!(bit_decomp_vartime(Scalar::from(1u128 << 127)), None);
bit_decomp_tester(Scalar::from(0u32), 0, "");
bit_decomp_tester(Scalar::from(0u32), 5, "00000");
bit_decomp_tester(Scalar::from(1u32), 0, "");
bit_decomp_tester(Scalar::from(1u32), 1, "1");
bit_decomp_tester(Scalar::from(2u32), 1, "0");
bit_decomp_tester(Scalar::from(2u32), 2, "01");
bit_decomp_tester(Scalar::from(3u32), 1, "1");
bit_decomp_tester(Scalar::from(3u32), 2, "11");
bit_decomp_tester(Scalar::from(5u32), 8, "10100000");
bit_decomp_tester(
Scalar::from(1u32).neg(),
32,
"00110111110010111010111100111010",
);
bit_decomp_tester(Scalar::from((1u128 << 127) - 2), 127,
"0111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111"
);
bit_decomp_tester(Scalar::from((1u128 << 127) - 1), 127,
"1111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111111"
);
bit_decomp_tester(Scalar::from(1u128 << 127), 127,
"0000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000"
);
bit_decomp_tester(Scalar::from(1u128 << 127), 128,
"0000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000"
);
}
fn bitrep_tester(upper: Scalar, x: Scalar, expected: bool) -> Result<(), Error> {
let rep_scalars = bitrep_scalars_vartime(upper)?;
let bitrep = compute_bitrep(x, &rep_scalars);
let nbits = bitrep.len();
assert!(nbits == rep_scalars.len());
let mut x_out = Scalar::ZERO;
for i in 0..nbits {
x_out += Scalar::conditional_select(&Scalar::ZERO, &rep_scalars[i], bitrep[i]);
}
if (x == x_out) != expected {
return Err(Error::VerificationFailure);
}
Ok(())
}
#[test]
fn bitrep_test() {
bitrep_tester(Scalar::from(0u32), Scalar::from(0u32), false).unwrap_err();
bitrep_tester(Scalar::from(1u32), Scalar::from(0u32), true).unwrap_err();
bitrep_tester(Scalar::from(2u32), Scalar::from(1u32), true).unwrap();
bitrep_tester(Scalar::from(3u32), Scalar::from(1u32), true).unwrap();
bitrep_tester(Scalar::from(100u32), Scalar::from(99u32), true).unwrap();
bitrep_tester(Scalar::from(127u32), Scalar::from(126u32), true).unwrap();
bitrep_tester(Scalar::from(128u32), Scalar::from(127u32), true).unwrap();
bitrep_tester(Scalar::from(128u32), Scalar::from(128u32), false).unwrap();
bitrep_tester(Scalar::from(129u32), Scalar::from(128u32), true).unwrap();
bitrep_tester(Scalar::from(129u32), Scalar::from(0u32), true).unwrap();
bitrep_tester(Scalar::from(129u32), Scalar::from(129u32), false).unwrap();
}
}