Skip to main content

kyn_vdf/
chia.rs

1//! Chia-compatible discriminant generation, BQFC form serialization, and Wesolowski
2//! VDF verification.
3//!
4//! This module implements:
5//! - [`create_discriminant`]: deterministic 1024-bit negative prime discriminant from a seed.
6//! - [`hash_prime`]: Chia's `HashPrime` algorithm (iterated SHA-256 with Miller-Rabin).
7//! - [`serialize_form`] / [`deserialize_form`]: Chia's Binary Quadratic Form Compression (BQFC).
8//! - [`verify_wesolowski`]: the Wesolowski VDF verification equation $\pi^B \cdot x^r = y$.
9//!
10//! All fallible operations return [`KynVdfError`] rather than panicking, ensuring safe
11//! execution in WASM, `no_std`, and adversarial input environments.
12
13use num_bigint::{BigInt, BigUint, Sign};
14use num_integer::Integer;
15use num_traits::{One, Signed, Zero};
16use sha2::{Digest, Sha256};
17
18use crate::error::KynVdfError;
19use crate::math::Form;
20
21/// Size of the Fiat-Shamir prime $B$ in bits.
22///
23/// $B$ is a 264-bit prime derived from the serialized forms $x$ and $y$ via [`hash_prime`].
24/// The choice of 264 bits provides ~128-bit security for the Wesolowski soundness argument.
25const B_BITS: usize = 264;
26
27/// Size in bytes of a single serialized binary quadratic form (Chia BQFC format).
28const BQFC_FORM_SIZE: usize = 100;
29
30// --- BQFC flag bits in the first byte of a serialized form ---
31
32/// Bit 0 of byte 0: set when the `b` coefficient of the original form is negative.
33const BQFC_B_SIGN: u8 = 1 << 0;
34/// Bit 1 of byte 0: set when the partial-XGCD quotient `t` is negative.
35const BQFC_T_SIGN: u8 = 1 << 1;
36/// Bit 2 of byte 0: set when the form is the principal identity (a=1, b=1).
37const BQFC_IS_1: u8 = 1 << 2;
38/// Bit 3 of byte 0: set when the form is the canonical generator (a=2, b=1).
39const BQFC_IS_GEN: u8 = 1 << 3;
40
41/// Computes the integer square root $\lfloor \sqrt{n} \rfloor$ for a non-negative `BigInt`.
42///
43/// # Errors
44/// Returns [`KynVdfError::ArithmeticError`] if `n` is negative, which indicates a
45/// programming error in the caller (reduced form coefficients `a` and `c` must be positive).
46fn isqrt(n: &BigInt) -> Result<BigInt, KynVdfError> {
47    if n.is_negative() {
48        return Err(KynVdfError::ArithmeticError(
49            "cannot compute integer square root of a negative number; \
50             form coefficient 'a' must be positive for a valid reduced form"
51                .to_string(),
52        ));
53    }
54    if n.is_zero() {
55        return Ok(BigInt::zero());
56    }
57    let uint_sqrt = n.to_biguint().unwrap().sqrt();
58    Ok(BigInt::from_biguint(Sign::Plus, uint_sqrt))
59}
60
61/// Miller-Rabin probabilistic primality test.
62///
63/// Uses a fixed set of deterministic small bases (2, 3, 5, …, 37) followed by
64/// additional rounds if `rounds > 12`. For 264-bit numbers, 12 rounds give a
65/// false-positive probability below $4^{-12} \approx 2^{-24}$.
66///
67/// # Parameters
68/// - `n`: The candidate integer to test (as `BigUint`).
69/// - `rounds`: Minimum number of Miller-Rabin witness rounds to perform.
70///
71/// # Returns
72/// `true` if `n` is probably prime; `false` if `n` is definitely composite.
73pub fn is_probable_prime(n: &BigUint, rounds: usize) -> bool {
74    if n < &BigUint::from(2u32) {
75        return false;
76    }
77    if n == &BigUint::from(2u32) || n == &BigUint::from(3u32) {
78        return true;
79    }
80    if n.is_even() {
81        return false;
82    }
83
84    // Fast elimination via small prime trial division
85    let small_primes = [3u32, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37, 41, 43, 47];
86    for &p in &small_primes {
87        let bp = BigUint::from(p);
88        if n == &bp {
89            return true;
90        }
91        if (n % &bp).is_zero() {
92            return false;
93        }
94    }
95
96    // Factor out powers of 2: write n - 1 = 2^r * d
97    let n_minus_1 = n - BigUint::one();
98    let mut d = n_minus_1.clone();
99    let mut r = 0usize;
100    while d.is_even() {
101        d >>= 1;
102        r += 1;
103    }
104
105    // Deterministic bases sufficient for numbers up to ~3.3 × 10^24; extend with more rounds
106    let bases = [2u32, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37];
107    let test_rounds = std::cmp::max(rounds, bases.len());
108
109    'outer: for i in 0..test_rounds {
110        let a = if i < bases.len() {
111            BigUint::from(bases[i])
112        } else {
113            BigUint::from((i as u32) * 2 + 39)
114        };
115        if &a >= n {
116            break;
117        }
118
119        let mut x = a.modpow(&d, n);
120        if x.is_one() || x == n_minus_1 {
121            continue;
122        }
123
124        for _ in 0..(r - 1) {
125            x = x.modpow(&BigUint::from(2u32), n);
126            if x == n_minus_1 {
127                continue 'outer;
128            }
129        }
130
131        return false; // Composite witness found
132    }
133
134    true
135}
136
137/// Generates a pseudoprime matching Chia's `HashPrime` algorithm.
138///
139/// Produces a prime of exactly `length_bits` bits by:
140/// 1. Incrementing `seed` byte-by-byte (big-endian counter) and hashing with SHA-256.
141/// 2. Concatenating hash outputs until `length_bits / 8` bytes are collected.
142/// 3. Setting specific bits from `bitmask` and forcing the lowest bit to 1 (odd).
143/// 4. Testing with Miller-Rabin (25 rounds); looping until a prime is found.
144///
145/// This is a deterministic function: the same seed always produces the same prime.
146///
147/// # Parameters
148/// - `seed`: Non-empty byte slice used as the initial hash input.
149/// - `length_bits`: Desired bit-length of the output prime. **Must be a non-zero multiple of 8.**
150/// - `bitmask`: Indices of bits to force to 1 in the candidate (used to set the MSB and
151///   ensure the prime has the required congruence properties).
152///
153/// # Errors
154/// Returns [`KynVdfError::InvalidDiscriminantSize`] if `length_bits` is zero or not a
155/// multiple of 8.
156pub fn hash_prime(
157    seed: &[u8],
158    length_bits: usize,
159    bitmask: &[usize],
160) -> Result<BigUint, KynVdfError> {
161    if length_bits == 0 || length_bits % 8 != 0 {
162        return Err(KynVdfError::InvalidDiscriminantSize(length_bits));
163    }
164
165    let mut sprout = seed.to_vec();
166
167    loop {
168        let mut blob = Vec::new();
169        while blob.len() * 8 < length_bits {
170            // Increment the counter (big-endian) before each hash
171            for i in (0..sprout.len()).rev() {
172                sprout[i] = sprout[i].wrapping_add(1);
173                if sprout[i] != 0 {
174                    break;
175                }
176            }
177            let hash = Sha256::digest(&sprout);
178            let needed = (length_bits / 8) - blob.len();
179            let take = std::cmp::min(hash.len(), needed);
180            blob.extend_from_slice(&hash[..take]);
181        }
182
183        let mut p = BigUint::from_bytes_be(&blob);
184        // Set required bits (MSB, congruence constraints)
185        for &b in bitmask {
186            p.set_bit(b as u64, true);
187        }
188        // Force odd — even numbers cannot be prime (except 2)
189        p.set_bit(0, true);
190
191        if is_probable_prime(&p, 25) {
192            return Ok(p);
193        }
194        // Not prime: loop with the incremented sprout to try the next candidate
195    }
196}
197
198/// Creates a negative fundamental prime discriminant $D = -p$ from a seed byte slice.
199///
200/// Matches Chia's `CreateDiscriminant` function exactly, producing a 1024-bit (or
201/// other size) negative prime $p \equiv 7 \pmod 8$, returned as $D = -p$.
202///
203/// The discriminant is deterministic: the same `seed` and `length_bits` always
204/// produce the same $D$.
205///
206/// # Parameters
207/// - `seed`: Non-empty byte slice (use ≥ 32 bytes for security).
208/// - `length_bits`: Bit-length of the discriminant. Must be a non-zero multiple of 8.
209///   Typical values: `512`, `1024`, `2048`.
210///
211/// # Errors
212/// - [`KynVdfError::InvalidSeed`] if `seed` is empty.
213/// - [`KynVdfError::InvalidDiscriminantSize`] if `length_bits` is zero or not a multiple of 8.
214///
215/// # Example
216/// ```rust
217/// use kyn_vdf::create_discriminant;
218///
219/// let d = create_discriminant(&[0x42u8; 32], 1024).expect("valid seed");
220/// assert!(d < num_bigint::BigInt::from(0)); // D is always negative
221/// ```
222pub fn create_discriminant(seed: &[u8], length_bits: usize) -> Result<BigInt, KynVdfError> {
223    if seed.is_empty() {
224        return Err(KynVdfError::InvalidSeed(
225            "seed must be a non-empty byte slice".to_string(),
226        ));
227    }
228    if length_bits == 0 || length_bits % 8 != 0 {
229        return Err(KynVdfError::InvalidDiscriminantSize(length_bits));
230    }
231
232    let p = hash_prime(seed, length_bits, &[0, 1, 2, length_bits - 1])?;
233    Ok(-BigInt::from_biguint(Sign::Plus, p))
234}
235
236/// Performs partial Extended Euclidean Algorithm for BQFC form compression.
237///
238/// Identical in semantics to [`crate::math::xgcd_partial`] but used internally
239/// within the BQFC compression pipeline on `BigInt` values.
240fn xgcd_partial_chia(a: &BigInt, b: &BigInt, l: &BigInt) -> (BigInt, BigInt, BigInt, BigInt) {
241    let mut r2 = a.clone();
242    let mut r1 = b.clone();
243    let mut co2 = BigInt::zero();
244    let mut co1 = BigInt::from(-1);
245
246    while r1 > BigInt::zero() && &r1 > l {
247        let q = &r2 / &r1;
248        let t1 = &r2 - &q * &r1;
249        let t2 = &co2 - &q * &co1;
250        r2 = r1;
251        r1 = t1;
252        co2 = co1;
253        co1 = t2;
254    }
255    (co2, co1, r2, r1)
256}
257
258/// Compressed representation of a binary quadratic form $(a, b)$ in Chia's BQFC format.
259///
260/// BQFC (Binary Quadratic Form Compression) reduces the storage of a 1024-bit form
261/// from 128+ bytes to exactly 100 bytes by encoding the partial XGCD decomposition of
262/// $(a, b)$ rather than storing the coefficients directly.
263#[derive(Debug, Clone)]
264pub struct CompressedForm {
265    /// Compressed form of the `a` coefficient (divided by `g` if `g > 1`).
266    pub a: BigInt,
267    /// Partial XGCD quotient $t$ such that $t \cdot b \equiv \pm\sqrt{D} \pmod{a}$.
268    pub t: BigInt,
269    /// Common divisor $g = \gcd(a, t)$; equal to 1 when no further factoring is needed.
270    pub g: BigInt,
271    /// High-order correction term $b_0 = b / a'$ (non-zero only when `g > 1`).
272    pub b0: BigInt,
273    /// Sign of the original `b` coefficient (`true` = negative).
274    pub b_sign: bool,
275}
276
277/// Compresses a reduced binary quadratic form $(a, b)$ into Chia's BQFC components.
278///
279/// This is the inverse of [`bqfc_decompr`]. The compressed form can be serialized
280/// to exactly 100 bytes by [`serialize_form`].
281///
282/// # Errors
283/// Returns [`KynVdfError::ArithmeticError`] if `a` is negative (which would indicate
284/// a non-reduced or invalid input form).
285pub fn bqfc_compr(a: &BigInt, b: &BigInt) -> Result<CompressedForm, KynVdfError> {
286    if a == b {
287        return Ok(CompressedForm {
288            a: a.clone(),
289            t: BigInt::zero(),
290            g: BigInt::zero(),
291            b0: BigInt::zero(),
292            b_sign: false,
293        });
294    }
295
296    let sign = b.is_negative();
297    let a_sqrt = isqrt(a)?; // a must be positive for a valid reduced form
298    let a_copy = a.clone();
299    let b_copy = if sign { -b } else { b.clone() };
300
301    let (_dummy, mut t, _r2, _r1) = xgcd_partial_chia(&a_copy, &b_copy, &a_sqrt);
302    t = -t;
303
304    let g = a.gcd(&t);
305    let (out_a, out_t, mut out_b0) = if g == BigInt::one() {
306        (a.clone(), t, BigInt::zero())
307    } else {
308        let out_a = a / &g;
309        let out_t = &t / &g;
310        let b0 = b / &out_a;
311        (out_a, out_t, b0)
312    };
313
314    if sign {
315        out_b0 = -out_b0;
316    }
317
318    Ok(CompressedForm {
319        a: out_a,
320        t: out_t,
321        g,
322        b0: out_b0,
323        b_sign: sign,
324    })
325}
326
327/// Decompresses Chia BQFC components back into the original $(a, b)$ coefficients.
328///
329/// This reconstructs $b$ from the partial XGCD decomposition stored in the
330/// [`CompressedForm`], using modular inversion and a square-root recovery step.
331///
332/// # Errors
333/// - [`KynVdfError::FormDeserializationError`] if:
334///   - The compressed `a` coefficient is zero (malformed input).
335///   - `t` and `a` are not coprime (modular inverse does not exist).
336///   - The discriminant residue $t^2 \cdot D \bmod a$ is not a perfect square
337///     (proof bytes are corrupted or tampered).
338pub fn bqfc_decompr(c: &CompressedForm, d: &BigInt) -> Result<(BigInt, BigInt), KynVdfError> {
339    // Special case: a == b (identity or generator detection)
340    if c.t.is_zero() {
341        return Ok((c.a.clone(), c.a.clone()));
342    }
343
344    if c.a.is_zero() {
345        return Err(KynVdfError::FormDeserializationError(
346            "compressed form has zero 'a' coefficient; the proof bytes are malformed".to_string(),
347        ));
348    }
349
350    let mut t = c.t.clone();
351    if t.is_negative() {
352        t += &c.a;
353    }
354
355    // Compute modular inverse of t modulo a via extended GCD
356    let ext = t.extended_gcd(&c.a);
357    if ext.gcd != BigInt::one() {
358        return Err(KynVdfError::FormDeserializationError(format!(
359            "partial quotient 't' (= {}) is not coprime with 'a' (= {}); \
360             the BQFC decompression inverse does not exist — proof bytes are corrupted",
361            c.t, c.a
362        )));
363    }
364    let mut t_inv = ext.x;
365    if t_inv.is_negative() {
366        t_inv += &c.a;
367    }
368
369    // Recover b from the quadratic residue: b ≡ sqrt(t² · D) · t⁻¹ (mod a)
370    let d_mod_a = d.mod_floor(&c.a);
371    let t_sq = (&c.t * &c.t).mod_floor(&c.a);
372    let tmp = (t_sq * d_mod_a).mod_floor(&c.a);
373
374    let root = isqrt(&tmp)?;
375    if &root * &root != tmp {
376        return Err(KynVdfError::FormDeserializationError(
377            "discriminant residue t²·D mod a is not a perfect square; \
378             the form cannot be reconstructed — proof bytes are corrupted or tampered"
379                .to_string(),
380        ));
381    }
382
383    let mut out_b = (&root * &t_inv).mod_floor(&c.a);
384    let out_a = if c.g > BigInt::one() {
385        &c.a * &c.g
386    } else {
387        c.a.clone()
388    };
389
390    if c.b0 > BigInt::zero() {
391        out_b += &c.a * &c.b0;
392    }
393
394    if c.b_sign {
395        out_b = -out_b;
396    }
397
398    Ok((out_a, out_b))
399}
400
401/// Writes `val` as little-endian bytes into `out_str[offset..offset+size]` with zero-padding.
402///
403/// Matches the byte layout of Chia's `bqfc.c` export routine.
404///
405/// # Errors
406/// Returns [`KynVdfError::FormDeserializationError`] if `val` requires more than `size` bytes.
407fn export_le(
408    val: &BigInt,
409    out_str: &mut [u8],
410    offset: &mut usize,
411    size: usize,
412) -> Result<(), KynVdfError> {
413    let bytes = val.to_biguint().unwrap_or_else(BigUint::zero).to_bytes_le();
414    if bytes.len() > size {
415        return Err(KynVdfError::FormDeserializationError(format!(
416            "integer value requires {} bytes but only {} bytes are available in the BQFC slot; \
417             the form coefficient is too large for the given discriminant size",
418            bytes.len(),
419            size
420        )));
421    }
422    out_str[*offset..*offset + bytes.len()].copy_from_slice(&bytes);
423    out_str[*offset + bytes.len()..*offset + size].fill(0);
424    *offset += size;
425    Ok(())
426}
427
428/// Reads a little-endian `BigInt` from a byte slice (positive, matching Chia's `bqfc.c`).
429fn import_le(data: &[u8]) -> BigInt {
430    BigInt::from_biguint(Sign::Plus, BigUint::from_bytes_le(data))
431}
432
433/// Serializes a reduced binary quadratic form into Chia's 100-byte BQFC wire format.
434///
435/// The output is always exactly [`BQFC_FORM_SIZE`] (100) bytes:
436/// - **Byte 0**: Flag byte (`BQFC_B_SIGN`, `BQFC_T_SIGN`, `BQFC_IS_1`, `BQFC_IS_GEN`).
437/// - **Byte 1**: `g_size` — the byte-length of the `g` coefficient minus 1.
438/// - **Bytes 2…end**: Little-endian packed fields `(a, t, g, b0)`.
439///
440/// Special cases (identity and generator) are encoded with a single flag byte.
441///
442/// # Parameters
443/// - `form`: The reduced binary quadratic form to serialize.
444/// - `d_bits`: The bit-size of the discriminant (e.g. `1024`).
445///
446/// # Errors
447/// Returns [`KynVdfError::FormDeserializationError`] if any coefficient overflows its
448/// allocated slot (which would indicate a form with a mismatched discriminant size).
449pub fn serialize_form(form: &Form, d_bits: usize) -> Result<Vec<u8>, KynVdfError> {
450    let mut res = vec![0u8; BQFC_FORM_SIZE];
451
452    // Fast path: identity (a=1, b=1) and generator (a=2, b=1) use a single flag byte
453    if form.b == BigInt::one() && form.a <= BigInt::from(2) {
454        res[0] = if form.a == BigInt::from(2) {
455            BQFC_IS_GEN
456        } else {
457            BQFC_IS_1
458        };
459        return Ok(res);
460    }
461
462    let d_bits_rounded = (d_bits + 31) & !31;
463    let compr = bqfc_compr(&form.a, &form.b)?;
464
465    // Encode sign flags
466    res[0] = if compr.b_sign { BQFC_B_SIGN } else { 0 };
467    if compr.t.is_negative() {
468        res[0] |= BQFC_T_SIGN;
469    }
470
471    // Compute field widths from discriminant size
472    let g_biguint = compr.g.to_biguint().unwrap_or_else(BigUint::zero);
473    let g_size = if compr.g.is_zero() {
474        0
475    } else {
476        (g_biguint.bits() as usize + 7) / 8 - 1
477    };
478    res[1] = g_size as u8;
479
480    let mut offset = 2;
481    let a_bytes_len = d_bits_rounded / 16 - g_size;
482    let t_bytes_len = d_bits_rounded / 32 - g_size;
483    let g_bytes_len = g_size + 1;
484
485    export_le(&compr.a, &mut res, &mut offset, a_bytes_len)?;
486    let t_abs = compr.t.abs();
487    export_le(&t_abs, &mut res, &mut offset, t_bytes_len)?;
488    export_le(&compr.g, &mut res, &mut offset, g_bytes_len)?;
489    let b0_abs = compr.b0.abs();
490    export_le(&b0_abs, &mut res, &mut offset, g_bytes_len)?;
491
492    Ok(res)
493}
494
495/// Deserializes a Chia 100-byte BQFC-compressed form into a reduced [`Form`].
496///
497/// Performs full validation after decompression:
498/// - Checks that the decompressed $(a, b)$ satisfy the discriminant identity $b^2 - 4ac = D$.
499/// - Checks that the resulting form is in reduced normal form.
500///
501/// # Parameters
502/// - `d`: The negative fundamental discriminant used to verify the form.
503/// - `bytes`: Exactly [`BQFC_FORM_SIZE`] (100) bytes in Chia BQFC wire format.
504///
505/// # Errors
506/// - [`KynVdfError::InvalidProofLength`] if `bytes.len() != 100`.
507/// - [`KynVdfError::FormDeserializationError`] for any structural corruption.
508/// - [`KynVdfError::InvalidDiscriminantIdentity`] if the form does not satisfy $b^2 - 4ac = D$.
509pub fn deserialize_form(d: &BigInt, bytes: &[u8]) -> Result<Form, KynVdfError> {
510    if bytes.len() != BQFC_FORM_SIZE {
511        return Err(KynVdfError::InvalidProofLength {
512            expected: BQFC_FORM_SIZE,
513            actual: bytes.len(),
514        });
515    }
516
517    // Fast path: identity and generator flags bypass full decompression
518    if bytes[0] & (BQFC_IS_1 | BQFC_IS_GEN) != 0 {
519        let a = if bytes[0] & BQFC_IS_GEN != 0 {
520            BigInt::from(2)
521        } else {
522            BigInt::from(1)
523        };
524        let b = BigInt::one();
525        return Form::from_abd(&a, &b, d).ok_or(KynVdfError::InvalidDiscriminantIdentity);
526    }
527
528    let d_bits = d.abs().to_biguint().unwrap().bits() as usize;
529    let d_bits_rounded = (d_bits + 31) & !31;
530
531    let g_size = bytes[1] as usize;
532    if g_size >= d_bits_rounded / 32 {
533        return Err(KynVdfError::FormDeserializationError(format!(
534            "g_size field ({}) exceeds the maximum allowed value ({}) for a {}-bit discriminant; \
535             the proof bytes are corrupted",
536            g_size,
537            d_bits_rounded / 32 - 1,
538            d_bits
539        )));
540    }
541
542    let mut offset = 2;
543    let a_len = d_bits_rounded / 16 - g_size;
544    let t_len = d_bits_rounded / 32 - g_size;
545    let g_len = g_size + 1;
546
547    if offset + a_len + t_len + 2 * g_len > bytes.len() {
548        return Err(KynVdfError::FormDeserializationError(
549            "encoded field sizes exceed the 100-byte form buffer; \
550             the proof bytes are truncated or corrupted"
551                .to_string(),
552        ));
553    }
554
555    // Parse each little-endian field
556    let a_part = import_le(&bytes[offset..offset + a_len]);
557    offset += a_len;
558
559    let mut t_part = import_le(&bytes[offset..offset + t_len]);
560    offset += t_len;
561
562    let g_part = import_le(&bytes[offset..offset + g_len]);
563    offset += g_len;
564
565    let b0_part = import_le(&bytes[offset..offset + g_len]);
566
567    // Decode sign flags
568    let b_sign = (bytes[0] & BQFC_B_SIGN) != 0;
569    if (bytes[0] & BQFC_T_SIGN) != 0 {
570        t_part = -t_part;
571    }
572
573    let compr = CompressedForm {
574        a: a_part,
575        t: t_part,
576        g: g_part,
577        b0: b0_part,
578        b_sign,
579    };
580
581    let (dec_a, dec_b) = bqfc_decompr(&compr, d)?;
582
583    // Validate the discriminant identity: b² - 4ac = D
584    let form = Form::from_abd(&dec_a, &dec_b, d).ok_or(KynVdfError::InvalidDiscriminantIdentity)?;
585
586    // Validate the form is in canonical reduced form
587    if !form.is_reduced() {
588        return Err(KynVdfError::FormDeserializationError(
589            "decompressed form is not in reduced normal form; \
590             the proof bytes may be corrupted or produced by an incompatible implementation"
591                .to_string(),
592        ));
593    }
594
595    Ok(form)
596}
597
598/// Derives the 264-bit Fiat-Shamir prime challenge $B$ from the serialized generator $x$
599/// and VDF output $y$.
600///
601/// $B = \text{HashPrime}(\text{serialize}(x) \| \text{serialize}(y), 264)$
602///
603/// This makes $B$ a deterministic function of the public inputs, binding the proof $\pi$
604/// to a specific $(x, y)$ pair and preventing the prover from choosing $B$ adaptively.
605///
606/// # Errors
607/// Propagates [`KynVdfError`] from [`serialize_form`] or [`hash_prime`].
608pub fn get_b(d: &BigInt, x: &Form, y: &Form) -> Result<BigUint, KynVdfError> {
609    let d_bits = d.abs().to_biguint().unwrap().bits() as usize;
610    let ser_x = serialize_form(x, d_bits)?;
611    let ser_y = serialize_form(y, d_bits)?;
612
613    let mut concat = ser_x;
614    concat.extend_from_slice(&ser_y);
615
616    hash_prime(&concat, B_BITS, &[B_BITS - 1])
617}
618
619/// Verifies a Wesolowski VDF proof.
620///
621/// Checks the Wesolowski verification equation:
622///
623/// $$\pi^B \cdot x^r = y$$
624///
625/// where:
626/// - $B = \text{HashPrime}(\text{ser}(x) \| \text{ser}(y), 264)$ is the Fiat-Shamir prime.
627/// - $r = 2^T \bmod B$ is the remainder term.
628/// - $\pi$ is the proof form.
629/// - $x$ is the generator (challenge) form.
630/// - $y = x^{2^T}$ is the claimed VDF output.
631///
632/// Verification runs in $\mathcal{O}(\log T)$ time because $B$ is a fixed-size (264-bit)
633/// prime regardless of $T$, so both `pow` calls are bounded by $\log_2(B) \approx 264$
634/// class group squarings.
635///
636/// # Parameters
637/// - `d`: Negative fundamental discriminant.
638/// - `x`: Generator form (challenge input).
639/// - `y`: Claimed VDF output form (result of $T$ sequential squarings of $x$).
640/// - `proof`: Wesolowski proof form $\pi$.
641/// - `iterations`: Number of sequential squarings $T$ that were evaluated.
642///
643/// # Returns
644/// - `Ok(true)` if $\pi^B \cdot x^r = y$ — proof is valid.
645/// - `Ok(false)` if the equation does not hold — proof is invalid.
646/// - `Err(KynVdfError)` if any input is malformed.
647pub fn verify_wesolowski(
648    d: &BigInt,
649    x: &Form,
650    y: &Form,
651    proof: &Form,
652    iterations: u64,
653) -> Result<bool, KynVdfError> {
654    // Derive the 264-bit Fiat-Shamir prime B = HashPrime(ser(x) || ser(y))
655    let b = get_b(d, x, y)?;
656
657    // r = 2^T mod B  (this is cheap: T is just a u64 used as exponent)
658    let r = BigUint::from(2u32).modpow(&BigUint::from(iterations), &b);
659
660    // f1 = π^B  (O(log B) ≈ O(264) squarings — independent of T)
661    let f1 = proof.pow(&b, d);
662    // f2 = x^r  (O(log r) ≤ O(264) squarings — independent of T)
663    let f2 = x.pow(&r, d);
664
665    // Check the verification equation: π^B · x^r == y
666    let result = f1.compose(&f2, d);
667    Ok(&result == y)
668}
669
670#[cfg(test)]
671mod tests {
672    use super::*;
673
674    #[test]
675    fn test_discriminant_chia_test_vector() {
676        let challenge = [42u8; 32];
677        let d = create_discriminant(&challenge, 1024).expect("valid seed and size");
678        assert!(d.is_negative());
679        let p_bytes = d.abs().to_biguint().unwrap().to_bytes_be();
680
681        let expected_prefix = [
682            237, 89, 165, 1, 5, 76, 207, 152, 207, 134, 182, 117, 254, 184, 124, 248,
683        ];
684        assert_eq!(&p_bytes[0..16], &expected_prefix[..]);
685    }
686
687    #[test]
688    fn test_create_discriminant_rejects_empty_seed() {
689        let res = create_discriminant(&[], 1024);
690        assert!(matches!(res, Err(KynVdfError::InvalidSeed(_))));
691    }
692
693    #[test]
694    fn test_create_discriminant_rejects_bad_size() {
695        let res = create_discriminant(&[1u8; 32], 0);
696        assert!(matches!(res, Err(KynVdfError::InvalidDiscriminantSize(0))));
697
698        let res2 = create_discriminant(&[1u8; 32], 100); // not multiple of 8? 100 is multiple of 8
699        // 100 IS a multiple of 8 (100 = 8 * 12 + 4... no wait: 100 / 8 = 12.5, not a whole number)
700        // Actually 100 % 8 = 4, so it's not a multiple of 8
701        assert!(matches!(res2, Err(KynVdfError::InvalidDiscriminantSize(100))));
702    }
703}