iuna

iuna

iuna - experimental mainnet-candidate protocol
git clone https://getiuna.org/git/iuna.git
Log | Files | Refs | README | LICENSE

arithmetic.rs (8468B)


      1 use kyn_vdf::Form;
      2 use num_bigint::BigInt;
      3 use num_integer::Integer;
      4 use num_traits::{One, Signed, Zero};
      5 
      6 #[cfg(test)]
      7 pub(super) fn nudupl(form: &Form, discriminant: &BigInt, threshold: &BigInt) -> Form {
      8     nudupl_owned(form.clone(), discriminant, threshold)
      9 }
     10 
     11 pub(super) fn nudupl_owned(form: Form, discriminant: &BigInt, threshold: &BigInt) -> Form {
     12     let two = BigInt::from(2);
     13     let four = BigInt::from(4);
     14     let Form { a, b, c } = form;
     15     let mut a1 = a;
     16     let mut c1 = c;
     17 
     18     let gcd = if b.is_negative() {
     19         let b_abs = -&b;
     20         let gcd = b_abs.extended_gcd(&a1);
     21         (-gcd.x, gcd.gcd)
     22     } else {
     23         let gcd = b.extended_gcd(&a1);
     24         (gcd.x, gcd.gcd)
     25     };
     26 
     27     let mut k = -(&gcd.0 * &c1);
     28     let s = gcd.1;
     29     if s != BigInt::one() {
     30         a1 /= &s;
     31         c1 *= &s;
     32     }
     33     k = mod_positive(k, &a1);
     34 
     35     if a1 < *threshold {
     36         let t = &a1 * &k;
     37         let result_a = &a1 * &a1;
     38         let result_b = &two * &t + &b;
     39         let result_c = ((&b + &t) * &k + &c1) / &a1;
     40         Form::new(result_a, result_b, result_c)
     41     } else {
     42         let (co2, co1, _r2, r1) = xgcd_partial(&a1, &k, threshold);
     43         let m2 = (&b * &r1 - &c1 * &co1) / &a1;
     44 
     45         let mut result_a = &r1 * &r1 - &co1 * &m2;
     46         if !co1.is_negative() {
     47             result_a = -result_a;
     48         }
     49 
     50         let result_b = mod_positive(
     51             (&two * (&a1 * &r1 - &result_a * &co2)) / &co1 - &b,
     52             &(&result_a * &two),
     53         );
     54         let mut result_c = (&result_b * &result_b - discriminant) / (&result_a * &four);
     55 
     56         if result_a.is_negative() {
     57             result_a = -result_a;
     58             result_c = -result_c;
     59         }
     60 
     61         Form::new(result_a, result_b, result_c)
     62     }
     63 }
     64 
     65 pub(super) fn nucomp(left: &Form, right: &Form, discriminant: &BigInt, threshold: &BigInt) -> Form {
     66     if left.a > right.a {
     67         return nucomp(right, left, discriminant, threshold);
     68     }
     69 
     70     let two = BigInt::from(2);
     71     let four = BigInt::from(4);
     72     let mut a1 = left.a.clone();
     73     let mut a2 = right.a.clone();
     74     let mut c2 = right.c.clone();
     75     let ss = (&left.b + &right.b) / &two;
     76     let m = (&left.b - &right.b) / &two;
     77 
     78     let t = &a2 % &a1;
     79     let (v1, sp) = if t.is_zero() {
     80         (BigInt::zero(), a1.clone())
     81     } else {
     82         let gcd = t.extended_gcd(&a1);
     83         (gcd.x, gcd.gcd)
     84     };
     85     let mut k = mod_positive(&m * &v1, &a1);
     86 
     87     if sp != BigInt::one() {
     88         let gcd = ss.extended_gcd(&sp);
     89         let v2 = gcd.x;
     90         let u2 = gcd.y;
     91         let s = gcd.gcd;
     92         k = &k * &u2 - &v2 * &c2;
     93         if s != BigInt::one() {
     94             a1 /= &s;
     95             a2 /= &s;
     96             c2 *= &s;
     97         }
     98         k = mod_positive(k, &a1);
     99     }
    100 
    101     if a1 < *threshold {
    102         let t = &a2 * &k;
    103         let result_a = &a2 * &a1;
    104         let result_b = &two * &t + &right.b;
    105         let result_c = ((&right.b + &t) * &k + &c2) / &a1;
    106         Form::new(result_a, result_b, result_c)
    107     } else {
    108         let (co2, co1, _r2, r1) = xgcd_partial(&a1, &k, threshold);
    109         let m1 = (&m * &co1 + &a2 * &r1) / &a1;
    110         let m2 = (&ss * &r1 - &c2 * &co1) / &a1;
    111 
    112         let mut result_a = &r1 * &m1 - &co1 * &m2;
    113         if !co1.is_negative() {
    114             result_a = -result_a;
    115         }
    116 
    117         let t = &a2 * &r1;
    118         let result_b = mod_positive(
    119             (&two * (&t - &result_a * &co2)) / &co1 - &right.b,
    120             &(&result_a * &two),
    121         );
    122         let mut result_c = (&result_b * &result_b - discriminant) / (&result_a * &four);
    123 
    124         if result_a.is_negative() {
    125             result_a = -result_a;
    126             result_c = -result_c;
    127         }
    128 
    129         Form::new(result_a, result_b, result_c)
    130     }
    131 }
    132 
    133 pub(super) fn xgcd_partial(
    134     r2: &BigInt,
    135     r1: &BigInt,
    136     threshold: &BigInt,
    137 ) -> (BigInt, BigInt, BigInt, BigInt) {
    138     let mut r2 = r2.clone();
    139     let mut r1 = r1.clone();
    140     let mut co2 = BigInt::zero();
    141     let mut co1 = BigInt::from(-1);
    142 
    143     while !r1.is_zero() && &r1 > threshold {
    144         let bits = r2.bits().max(r1.bits()).saturating_sub(63);
    145         let mut rr2 = shifted_low_word(&r2, bits);
    146         let mut rr1 = shifted_low_word(&r1, bits);
    147         let threshold_word = shifted_low_word(threshold, bits);
    148 
    149         let mut aa2 = 0_i128;
    150         let mut aa1 = 1_i128;
    151         let mut bb2 = 1_i128;
    152         let mut bb1 = 0_i128;
    153         let mut steps = 0_u32;
    154 
    155         while rr1 != 0 && rr1 > threshold_word {
    156             let q = rr2 / rr1;
    157             let next_r = rr2 - q * rr1;
    158             let next_a = aa2 - q * aa1;
    159             let next_b = bb2 - q * bb1;
    160 
    161             if steps & 1 == 1 {
    162                 if next_r < -next_b || rr1 - next_r < next_a - aa1 {
    163                     break;
    164                 }
    165             } else if next_r < -next_a || rr1 - next_r < next_b - bb1 {
    166                 break;
    167             }
    168 
    169             rr2 = rr1;
    170             rr1 = next_r;
    171             aa2 = aa1;
    172             aa1 = next_a;
    173             bb2 = bb1;
    174             bb1 = next_b;
    175             steps += 1;
    176         }
    177 
    178         if steps == 0 {
    179             let (q, next_r) = r2.div_rem(&r1);
    180             let next_co = &co2 - &q * &co1;
    181             r2 = r1;
    182             r1 = next_r;
    183             co2 = co1;
    184             co1 = next_co;
    185         } else {
    186             let old_r2 = r2;
    187             let old_r1 = r1;
    188             r2 = scaled(&old_r2, bb2) + scaled(&old_r1, aa2);
    189             r1 = scaled(&old_r1, aa1) + scaled(&old_r2, bb1);
    190 
    191             let old_co2 = co2;
    192             let old_co1 = co1;
    193             co2 = scaled(&old_co2, bb2) + scaled(&old_co1, aa2);
    194             co1 = scaled(&old_co1, aa1) + scaled(&old_co2, bb1);
    195 
    196             if r1.is_negative() {
    197                 r1 = -r1;
    198                 co1 = -co1;
    199             }
    200             if r2.is_negative() {
    201                 r2 = -r2;
    202                 co2 = -co2;
    203             }
    204         }
    205     }
    206 
    207     if r2.is_negative() {
    208         r2 = -r2;
    209         co2 = -co2;
    210         co1 = -co1;
    211     }
    212 
    213     (co2, co1, r2, r1)
    214 }
    215 
    216 fn scaled(value: &BigInt, scalar: i128) -> BigInt {
    217     value * BigInt::from(scalar)
    218 }
    219 
    220 fn mod_positive(mut value: BigInt, modulus: &BigInt) -> BigInt {
    221     value %= modulus;
    222     if value.is_negative() {
    223         value += modulus;
    224     }
    225     value
    226 }
    227 
    228 fn shifted_low_word(value: &BigInt, shift_bits: u64) -> i128 {
    229     debug_assert!(!value.is_negative());
    230     let mut digits = value.iter_u64_digits();
    231     let digit_index = usize::try_from(shift_bits / 64).unwrap_or(usize::MAX);
    232     let offset = (shift_bits % 64) as u32;
    233     let low = digits.nth(digit_index).unwrap_or(0);
    234     if offset == 0 {
    235         return i128::from(low);
    236     }
    237     let high = digits.next().unwrap_or(0);
    238     i128::from((low >> offset) | (high << (64 - offset)))
    239 }
    240 
    241 #[cfg(test)]
    242 mod tests {
    243     use kyn_vdf::{Form, create_discriminant, isqrt_fourth};
    244     use num_traits::Signed;
    245 
    246     use super::{nucomp, nudupl};
    247 
    248     #[test]
    249     fn optimized_nudupl_matches_kyn_across_sequential_squares() {
    250         let discriminant = create_discriminant(b"iuna-vdf-arithmetic-nudupl", 1024).unwrap();
    251         let threshold = isqrt_fourth(&discriminant.abs());
    252         let mut form = Form::generator(&discriminant).unwrap();
    253 
    254         for round in 1..=10_000 {
    255             let mut expected = form.nudupl(&discriminant, &threshold);
    256             expected.reduce(&discriminant);
    257             let mut actual = nudupl(&form, &discriminant, &threshold);
    258             actual.reduce(&discriminant);
    259 
    260             assert_eq!(actual, expected, "round={round}");
    261             form = actual;
    262         }
    263     }
    264 
    265     #[test]
    266     fn optimized_nucomp_matches_kyn_across_sequential_compositions() {
    267         let discriminant = create_discriminant(b"iuna-vdf-arithmetic-nucomp", 1024).unwrap();
    268         let threshold = isqrt_fourth(&discriminant.abs());
    269         let generator = Form::generator(&discriminant).unwrap();
    270         let mut left = generator.clone();
    271         let mut right = generator.nudupl(&discriminant, &threshold);
    272         right.reduce(&discriminant);
    273 
    274         for round in 1..=2_000 {
    275             let mut expected = left.nucomp(&right, &discriminant, &threshold);
    276             expected.reduce(&discriminant);
    277             let mut actual = nucomp(&left, &right, &discriminant, &threshold);
    278             actual.reduce(&discriminant);
    279 
    280             assert_eq!(actual, expected, "round={round}");
    281             left = actual;
    282             right = right.nudupl(&discriminant, &threshold);
    283             right.reduce(&discriminant);
    284         }
    285     }
    286 }