use acvm_blackbox_solver::blake3;
use ark_ec::{short_weierstrass::Affine, AffineRepr, CurveConfig};
use ark_ff::Field;
use ark_ff::{BigInteger, PrimeField};
use grumpkin::GrumpkinParameters;
pub(crate) fn hash_to_curve(seed: &[u8], attempt_count: u8) -> Affine<GrumpkinParameters> {
let seed_size = seed.len();
let mut target_seed = seed.to_vec();
target_seed.extend_from_slice(&[0u8; 2]);
target_seed[seed_size] = attempt_count;
target_seed[seed_size + 1] = 0;
let hash_hi = blake3(&target_seed).expect("hash should succeed");
target_seed[seed_size + 1] = 1;
let hash_lo = blake3(&target_seed).expect("hash should succeed");
let mut hash = hash_hi.to_vec();
hash.extend_from_slice(&hash_lo);
let x = <<GrumpkinParameters as CurveConfig>::BaseField as Field>::BasePrimeField::from_be_bytes_mod_order(&hash);
let x = <GrumpkinParameters as CurveConfig>::BaseField::from_base_prime_field(x);
if let Some(point) = Affine::<GrumpkinParameters>::get_point_from_x_unchecked(x, false) {
let parity_bit = hash_hi[0] > 127;
let y_bit_set = point.y().unwrap().into_bigint().get_bit(0);
if (parity_bit && !y_bit_set) || (!parity_bit && y_bit_set) {
-point
} else {
point
}
} else {
hash_to_curve(seed, attempt_count + 1)
}
}
#[cfg(test)]
mod test {
use ark_ec::AffineRepr;
use ark_ff::{BigInteger, PrimeField};
use super::hash_to_curve;
#[test]
fn smoke_test() {
let test_cases: [(&[u8], u8, (&str, &str)); 4] = [
(
&[],
0,
(
"24c4cb9c1206ab5470592f237f1698abe684dadf0ab4d7a132c32b2134e2c12e",
"0668b8d61a317fb34ccad55c930b3554f1828a0e5530479ecab4defe6bbc0b2e",
),
),
(
&[],
1,
(
"24c4cb9c1206ab5470592f237f1698abe684dadf0ab4d7a132c32b2134e2c12e",
"0668b8d61a317fb34ccad55c930b3554f1828a0e5530479ecab4defe6bbc0b2e",
),
),
(
&[1],
0,
(
"107f1b633c6113f3222f39f6256f0546b41a4880918c86864b06471afb410454",
"050cd3823d0c01590b6a50adcc85d2ee4098668fd28805578aa05a423ea938c6",
),
),
(
&[0x68, 0x65, 0x6c, 0x6c, 0x6f, 0x20, 0x77, 0x6f, 0x72, 0x6c, 0x64],
0,
(
"037c5c229ae495f6e8d1b4bf7723fafb2b198b51e27602feb8a4d1053d685093",
"10cf9596c5b2515692d930efa2cf3817607e4796856a79f6af40c949b066969f",
),
),
];
for (seed, attempt_count, expected_point) in test_cases {
let point = hash_to_curve(seed, attempt_count);
assert!(point.is_on_curve());
assert_eq!(
hex::encode(point.x().unwrap().into_bigint().to_bytes_be()),
expected_point.0,
"Failed on x component with seed {seed:?}, attempt_count {attempt_count}"
);
assert_eq!(
hex::encode(point.y().unwrap().into_bigint().to_bytes_be()),
expected_point.1,
"Failed on y component with seed {seed:?}, attempt_count {attempt_count}"
);
}
}
}