use num_bigint::{BigInt, BigUint, Sign};
use num_integer::Integer;
use num_traits::{One, Signed, Zero};
use sha2::{Digest, Sha256};
use crate::error::KynVdfError;
use crate::math::Form;
const B_BITS: usize = 264;
const BQFC_FORM_SIZE: usize = 100;
const BQFC_B_SIGN: u8 = 1 << 0;
const BQFC_T_SIGN: u8 = 1 << 1;
const BQFC_IS_1: u8 = 1 << 2;
const BQFC_IS_GEN: u8 = 1 << 3;
fn isqrt(n: &BigInt) -> Result<BigInt, KynVdfError> {
if n.is_negative() {
return Err(KynVdfError::ArithmeticError(
"cannot compute integer square root of a negative number; \
form coefficient 'a' must be positive for a valid reduced form"
.to_string(),
));
}
if n.is_zero() {
return Ok(BigInt::zero());
}
let uint_sqrt = n.to_biguint().unwrap().sqrt();
Ok(BigInt::from_biguint(Sign::Plus, uint_sqrt))
}
pub fn is_probable_prime(n: &BigUint, rounds: usize) -> bool {
if n < &BigUint::from(2u32) {
return false;
}
if n == &BigUint::from(2u32) || n == &BigUint::from(3u32) {
return true;
}
if n.is_even() {
return false;
}
let small_primes = [3u32, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37, 41, 43, 47];
for &p in &small_primes {
let bp = BigUint::from(p);
if n == &bp {
return true;
}
if (n % &bp).is_zero() {
return false;
}
}
let n_minus_1 = n - BigUint::one();
let mut d = n_minus_1.clone();
let mut r = 0usize;
while d.is_even() {
d >>= 1;
r += 1;
}
let bases = [2u32, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37];
let test_rounds = std::cmp::max(rounds, bases.len());
'outer: for i in 0..test_rounds {
let a = if i < bases.len() {
BigUint::from(bases[i])
} else {
BigUint::from((i as u32) * 2 + 39)
};
if &a >= n {
break;
}
let mut x = a.modpow(&d, n);
if x.is_one() || x == n_minus_1 {
continue;
}
for _ in 0..(r - 1) {
x = x.modpow(&BigUint::from(2u32), n);
if x == n_minus_1 {
continue 'outer;
}
}
return false; }
true
}
pub fn hash_prime(
seed: &[u8],
length_bits: usize,
bitmask: &[usize],
) -> Result<BigUint, KynVdfError> {
if length_bits == 0 || length_bits % 8 != 0 {
return Err(KynVdfError::InvalidDiscriminantSize(length_bits));
}
let mut sprout = seed.to_vec();
loop {
let mut blob = Vec::new();
while blob.len() * 8 < length_bits {
for i in (0..sprout.len()).rev() {
sprout[i] = sprout[i].wrapping_add(1);
if sprout[i] != 0 {
break;
}
}
let hash = Sha256::digest(&sprout);
let needed = (length_bits / 8) - blob.len();
let take = std::cmp::min(hash.len(), needed);
blob.extend_from_slice(&hash[..take]);
}
let mut p = BigUint::from_bytes_be(&blob);
for &b in bitmask {
p.set_bit(b as u64, true);
}
p.set_bit(0, true);
if is_probable_prime(&p, 25) {
return Ok(p);
}
}
}
pub fn create_discriminant(seed: &[u8], length_bits: usize) -> Result<BigInt, KynVdfError> {
if seed.is_empty() {
return Err(KynVdfError::InvalidSeed(
"seed must be a non-empty byte slice".to_string(),
));
}
if length_bits == 0 || length_bits % 8 != 0 {
return Err(KynVdfError::InvalidDiscriminantSize(length_bits));
}
let p = hash_prime(seed, length_bits, &[0, 1, 2, length_bits - 1])?;
Ok(-BigInt::from_biguint(Sign::Plus, p))
}
fn xgcd_partial_chia(a: &BigInt, b: &BigInt, l: &BigInt) -> (BigInt, BigInt, BigInt, BigInt) {
let mut r2 = a.clone();
let mut r1 = b.clone();
let mut co2 = BigInt::zero();
let mut co1 = BigInt::from(-1);
while r1 > BigInt::zero() && &r1 > l {
let q = &r2 / &r1;
let t1 = &r2 - &q * &r1;
let t2 = &co2 - &q * &co1;
r2 = r1;
r1 = t1;
co2 = co1;
co1 = t2;
}
(co2, co1, r2, r1)
}
#[derive(Debug, Clone)]
pub struct CompressedForm {
pub a: BigInt,
pub t: BigInt,
pub g: BigInt,
pub b0: BigInt,
pub b_sign: bool,
}
pub fn bqfc_compr(a: &BigInt, b: &BigInt) -> Result<CompressedForm, KynVdfError> {
if a == b {
return Ok(CompressedForm {
a: a.clone(),
t: BigInt::zero(),
g: BigInt::zero(),
b0: BigInt::zero(),
b_sign: false,
});
}
let sign = b.is_negative();
let a_sqrt = isqrt(a)?; let a_copy = a.clone();
let b_copy = if sign { -b } else { b.clone() };
let (_dummy, mut t, _r2, _r1) = xgcd_partial_chia(&a_copy, &b_copy, &a_sqrt);
t = -t;
let g = a.gcd(&t);
let (out_a, out_t, mut out_b0) = if g == BigInt::one() {
(a.clone(), t, BigInt::zero())
} else {
let out_a = a / &g;
let out_t = &t / &g;
let b0 = b / &out_a;
(out_a, out_t, b0)
};
if sign {
out_b0 = -out_b0;
}
Ok(CompressedForm {
a: out_a,
t: out_t,
g,
b0: out_b0,
b_sign: sign,
})
}
pub fn bqfc_decompr(c: &CompressedForm, d: &BigInt) -> Result<(BigInt, BigInt), KynVdfError> {
if c.t.is_zero() {
return Ok((c.a.clone(), c.a.clone()));
}
if c.a.is_zero() {
return Err(KynVdfError::FormDeserializationError(
"compressed form has zero 'a' coefficient; the proof bytes are malformed".to_string(),
));
}
let mut t = c.t.clone();
if t.is_negative() {
t += &c.a;
}
let ext = t.extended_gcd(&c.a);
if ext.gcd != BigInt::one() {
return Err(KynVdfError::FormDeserializationError(format!(
"partial quotient 't' (= {}) is not coprime with 'a' (= {}); \
the BQFC decompression inverse does not exist — proof bytes are corrupted",
c.t, c.a
)));
}
let mut t_inv = ext.x;
if t_inv.is_negative() {
t_inv += &c.a;
}
let d_mod_a = d.mod_floor(&c.a);
let t_sq = (&c.t * &c.t).mod_floor(&c.a);
let tmp = (t_sq * d_mod_a).mod_floor(&c.a);
let root = isqrt(&tmp)?;
if &root * &root != tmp {
return Err(KynVdfError::FormDeserializationError(
"discriminant residue t²·D mod a is not a perfect square; \
the form cannot be reconstructed — proof bytes are corrupted or tampered"
.to_string(),
));
}
let mut out_b = (&root * &t_inv).mod_floor(&c.a);
let out_a = if c.g > BigInt::one() {
&c.a * &c.g
} else {
c.a.clone()
};
if c.b0 > BigInt::zero() {
out_b += &c.a * &c.b0;
}
if c.b_sign {
out_b = -out_b;
}
Ok((out_a, out_b))
}
fn export_le(
val: &BigInt,
out_str: &mut [u8],
offset: &mut usize,
size: usize,
) -> Result<(), KynVdfError> {
let bytes = val.to_biguint().unwrap_or_else(BigUint::zero).to_bytes_le();
if bytes.len() > size {
return Err(KynVdfError::FormDeserializationError(format!(
"integer value requires {} bytes but only {} bytes are available in the BQFC slot; \
the form coefficient is too large for the given discriminant size",
bytes.len(),
size
)));
}
out_str[*offset..*offset + bytes.len()].copy_from_slice(&bytes);
out_str[*offset + bytes.len()..*offset + size].fill(0);
*offset += size;
Ok(())
}
fn import_le(data: &[u8]) -> BigInt {
BigInt::from_biguint(Sign::Plus, BigUint::from_bytes_le(data))
}
pub fn serialize_form(form: &Form, d_bits: usize) -> Result<Vec<u8>, KynVdfError> {
let mut res = vec![0u8; BQFC_FORM_SIZE];
if form.b == BigInt::one() && form.a <= BigInt::from(2) {
res[0] = if form.a == BigInt::from(2) {
BQFC_IS_GEN
} else {
BQFC_IS_1
};
return Ok(res);
}
let d_bits_rounded = (d_bits + 31) & !31;
let compr = bqfc_compr(&form.a, &form.b)?;
res[0] = if compr.b_sign { BQFC_B_SIGN } else { 0 };
if compr.t.is_negative() {
res[0] |= BQFC_T_SIGN;
}
let g_biguint = compr.g.to_biguint().unwrap_or_else(BigUint::zero);
let g_size = if compr.g.is_zero() {
0
} else {
(g_biguint.bits() as usize + 7) / 8 - 1
};
res[1] = g_size as u8;
let mut offset = 2;
let a_bytes_len = d_bits_rounded / 16 - g_size;
let t_bytes_len = d_bits_rounded / 32 - g_size;
let g_bytes_len = g_size + 1;
export_le(&compr.a, &mut res, &mut offset, a_bytes_len)?;
let t_abs = compr.t.abs();
export_le(&t_abs, &mut res, &mut offset, t_bytes_len)?;
export_le(&compr.g, &mut res, &mut offset, g_bytes_len)?;
let b0_abs = compr.b0.abs();
export_le(&b0_abs, &mut res, &mut offset, g_bytes_len)?;
Ok(res)
}
pub fn deserialize_form(d: &BigInt, bytes: &[u8]) -> Result<Form, KynVdfError> {
if bytes.len() != BQFC_FORM_SIZE {
return Err(KynVdfError::InvalidProofLength {
expected: BQFC_FORM_SIZE,
actual: bytes.len(),
});
}
if bytes[0] & (BQFC_IS_1 | BQFC_IS_GEN) != 0 {
let a = if bytes[0] & BQFC_IS_GEN != 0 {
BigInt::from(2)
} else {
BigInt::from(1)
};
let b = BigInt::one();
return Form::from_abd(&a, &b, d).ok_or(KynVdfError::InvalidDiscriminantIdentity);
}
let d_bits = d.abs().to_biguint().unwrap().bits() as usize;
let d_bits_rounded = (d_bits + 31) & !31;
let g_size = bytes[1] as usize;
if g_size >= d_bits_rounded / 32 {
return Err(KynVdfError::FormDeserializationError(format!(
"g_size field ({}) exceeds the maximum allowed value ({}) for a {}-bit discriminant; \
the proof bytes are corrupted",
g_size,
d_bits_rounded / 32 - 1,
d_bits
)));
}
let mut offset = 2;
let a_len = d_bits_rounded / 16 - g_size;
let t_len = d_bits_rounded / 32 - g_size;
let g_len = g_size + 1;
if offset + a_len + t_len + 2 * g_len > bytes.len() {
return Err(KynVdfError::FormDeserializationError(
"encoded field sizes exceed the 100-byte form buffer; \
the proof bytes are truncated or corrupted"
.to_string(),
));
}
let a_part = import_le(&bytes[offset..offset + a_len]);
offset += a_len;
let mut t_part = import_le(&bytes[offset..offset + t_len]);
offset += t_len;
let g_part = import_le(&bytes[offset..offset + g_len]);
offset += g_len;
let b0_part = import_le(&bytes[offset..offset + g_len]);
let b_sign = (bytes[0] & BQFC_B_SIGN) != 0;
if (bytes[0] & BQFC_T_SIGN) != 0 {
t_part = -t_part;
}
let compr = CompressedForm {
a: a_part,
t: t_part,
g: g_part,
b0: b0_part,
b_sign,
};
let (dec_a, dec_b) = bqfc_decompr(&compr, d)?;
let form = Form::from_abd(&dec_a, &dec_b, d).ok_or(KynVdfError::InvalidDiscriminantIdentity)?;
if !form.is_reduced() {
return Err(KynVdfError::FormDeserializationError(
"decompressed form is not in reduced normal form; \
the proof bytes may be corrupted or produced by an incompatible implementation"
.to_string(),
));
}
Ok(form)
}
pub fn get_b(d: &BigInt, x: &Form, y: &Form) -> Result<BigUint, KynVdfError> {
let d_bits = d.abs().to_biguint().unwrap().bits() as usize;
let ser_x = serialize_form(x, d_bits)?;
let ser_y = serialize_form(y, d_bits)?;
let mut concat = ser_x;
concat.extend_from_slice(&ser_y);
hash_prime(&concat, B_BITS, &[B_BITS - 1])
}
pub fn verify_wesolowski(
d: &BigInt,
x: &Form,
y: &Form,
proof: &Form,
iterations: u64,
) -> Result<bool, KynVdfError> {
let b = get_b(d, x, y)?;
let r = BigUint::from(2u32).modpow(&BigUint::from(iterations), &b);
let f1 = proof.pow(&b, d);
let f2 = x.pow(&r, d);
let result = f1.compose(&f2, d);
Ok(&result == y)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_discriminant_chia_test_vector() {
let challenge = [42u8; 32];
let d = create_discriminant(&challenge, 1024).expect("valid seed and size");
assert!(d.is_negative());
let p_bytes = d.abs().to_biguint().unwrap().to_bytes_be();
let expected_prefix = [
237, 89, 165, 1, 5, 76, 207, 152, 207, 134, 182, 117, 254, 184, 124, 248,
];
assert_eq!(&p_bytes[0..16], &expected_prefix[..]);
}
#[test]
fn test_create_discriminant_rejects_empty_seed() {
let res = create_discriminant(&[], 1024);
assert!(matches!(res, Err(KynVdfError::InvalidSeed(_))));
}
#[test]
fn test_create_discriminant_rejects_bad_size() {
let res = create_discriminant(&[1u8; 32], 0);
assert!(matches!(res, Err(KynVdfError::InvalidDiscriminantSize(0))));
let res2 = create_discriminant(&[1u8; 32], 100); assert!(matches!(res2, Err(KynVdfError::InvalidDiscriminantSize(100))));
}
}