diff --git a/CompElliptic.lean b/CompElliptic.lean index 3922c20..0c19d48 100644 --- a/CompElliptic.lean +++ b/CompElliptic.lean @@ -12,6 +12,11 @@ import CompElliptic.Encodings.Common import CompElliptic.Encodings.Pasta import CompElliptic.Fields.Pasta import CompElliptic.Fields.Residue +import CompElliptic.Fields.SafeGCD.Divstep +import CompElliptic.Fields.SafeGCD.Divsteps62 +import CompElliptic.Fields.SafeGCD.Limbs +import CompElliptic.Fields.SafeGCD.Pasta +import CompElliptic.Fields.SafeGCD.Reference import CompElliptic.Fields.Sqrt import CompElliptic.CurveForms.ShortWeierstrass import CompElliptic.CurveOrder diff --git a/CompElliptic/Fields/SafeGCD/Divstep.lean b/CompElliptic/Fields/SafeGCD/Divstep.lean new file mode 100644 index 0000000..f29585c --- /dev/null +++ b/CompElliptic/Fields/SafeGCD/Divstep.lean @@ -0,0 +1,341 @@ +/- +Copyright (c) 2026 CompElliptic Contributors. +Released under the Apache License, Version 2.0, or the MIT license, at your option, +as described in the files LICENSE-APACHE and LICENSE-MIT. +Authors: Danny Willems +-/ +import Mathlib.Data.Int.GCD +import Mathlib.Tactic.Ring + +/-! +# Bernstein-Yang divsteps over `ℤ` + +The mathematical core of *safegcd*: the divstep recurrence of Bernstein and Yang, +["Fast constant-time gcd computation and modular inversion"](https://eprint.iacr.org/2019/266), +in the `eta` form used by libsecp256k1's `modinv64`. + +A divstep acts on a state `(eta, f, g)` with `f` odd. Written as a matrix it is + +```text + ⌈f'⌉ 1 ⌈u v⌉ ⌈f⌉ + ⌊g'⌋ = ─ ⌊q r⌋ ⌊g⌋ + 2 +``` + +for one of three integer matrices of determinant `2`, selected by the parity of `g` and the sign +of `eta`. Composing `n` of them gives a single `Trans` of determinant `2 ^ n` that turns `n` +halvings into one exact division by `2 ^ n`, which is what makes a *batched* implementation +possible: read the branch decisions off the bottom limbs of `f` and `g`, accumulate the matrix, +then apply it once at full width. + +This module proves the properties such an implementation rests on, for the naive step-at-a-time +recurrence: + +* `run_det`: the accumulated determinant is exactly `2 ^ n`; +* `run_two_pow_mul_f` / `run_two_pow_mul_g`: the division by `2 ^ n` is exact, so no information + is lost; +* `run_emod_two_f`: `f` stays odd, the standing precondition of the recurrence; +* `run_inv_f` / `run_inv_g`: the *adjugate* recovers the original pair from the final one, since + the determinant is a unit after dividing through by `2 ^ n`. Consequently + `run_eq_one_or_neg_one`: once `g` reaches `0` and the inputs are coprime, the terminal `f` is + `±1`, because it divides both inputs. + +That last consequence is the usual "`gcd` is invariant along a divstep run", obtained here +without an invariance induction: invertibility of the transition over `ℤ` gives it directly. + +`CompElliptic.Fields.SafeGCD.Reference` turns these into the correctness of modular inversion. + +## Implementation notes + +Oddness of `f` is carried as `f % 2 = 1` rather than as `Odd f`, so that `omega` (which knows +`Int.emod` and `Int.ediv` by numerals) discharges the parity and exact-division side conditions +of all three branches directly. + +## Provenance + +Ported from `src/fields/modinv62.rs` of +[zcash/pasta_curves#119](https://github.com/zcash/pasta_curves/pull/119), pinned at commit +`9299dfbb19978428ba24229bc13dff25a424dc46` (branch `modinv62`). Every `Rust:` link on a +declaration below points into that commit, so the line numbers stay valid even after the branch +moves; to review against the tip of the branch instead, open the pull request and diff. + +The recurrence here is the one the port's batching kernel compresses; the kernel itself is +`divsteps_62_var`: + + +The module header stating the scaled invariants the driver maintains is: + + +## References + +* Daniel J. Bernstein, Bo-Yin Yang, *Fast constant-time gcd computation and modular inversion*, + IACR TCHES 2019(3). +* libsecp256k1, `src/modinv64_impl.h` and `doc/safegcd_implementation.md`. + +-/ + +namespace CompElliptic.Fields.SafeGCD + +/-- A 2×2 integer matrix, the transition of a run of divsteps: it is applied to the column +`(f, g)` and the result divided by the corresponding power of two. -/ +structure Trans where + /-- Top-left entry: the coefficient of `f` in the new `f`. -/ + u : Int + /-- Top-right entry: the coefficient of `g` in the new `f`. -/ + v : Int + /-- Bottom-left entry: the coefficient of `f` in the new `g`. -/ + q : Int + /-- Bottom-right entry: the coefficient of `g` in the new `g`. -/ + r : Int +deriving DecidableEq, Repr, Inhabited + +namespace Trans + +/-- The identity transition: the transition of zero divsteps. -/ +def id : Trans := ⟨1, 0, 0, 1⟩ + +/-- The determinant `u * r - v * q`. Every divstep matrix has determinant `2`, so a run of `n` +divsteps has determinant `2 ^ n`. -/ +def det (t : Trans) : Int := t.u * t.r - t.v * t.q + +/-- Matrix product: `a.comp b` applies `b` first, then `a`. -/ +def comp (a b : Trans) : Trans := + ⟨a.u * b.u + a.v * b.q, a.u * b.v + a.v * b.r, + a.q * b.u + a.r * b.q, a.q * b.v + a.r * b.r⟩ + +@[simp] theorem det_id : Trans.id.det = 1 := by simp [det, Trans.id] + +/-- Determinants are multiplicative, so a composed run's determinant is the product of the +determinants of its steps. -/ +theorem det_comp (a b : Trans) : (a.comp b).det = a.det * b.det := by + simp only [det, comp]; ring + +end Trans + +/-- The state of the divstep recurrence. -/ +structure DState where + /-- The Bernstein-Yang `eta` (the negation of their `delta`), driving the branch choice. -/ + eta : Int + /-- The first element of the pair, maintained odd. -/ + f : Int + /-- The second element of the pair, driven to zero. -/ + g : Int +deriving DecidableEq, Repr, Inhabited + +/-- The transition matrix of a single divstep, determined by `eta` and the parity of `g`: + +* `g` even: `(f, g) ↦ (f, g / 2)`; +* `g` odd and `eta < 0`: `(f, g) ↦ (g, (g - f) / 2)`, the swapping branch; +* `g` odd and `eta ≥ 0`: `(f, g) ↦ (f, (g + f) / 2)`. + +Each matrix has determinant `2`, and each division is exact whenever `f` is odd. -/ +def stepTrans (eta g : Int) : Trans := + if g % 2 = 0 then ⟨2, 0, 0, 1⟩ + else if eta < 0 then ⟨0, 2, -1, 1⟩ + else ⟨2, 0, 1, 1⟩ + +/-- The `eta` after a single divstep: it decreases by one, except on the swapping branch, where +it is reflected to `-eta - 1`. -/ +def stepEta (eta g : Int) : Int := + if g % 2 = 0 then eta - 1 + else if eta < 0 then -eta - 1 + else eta - 1 + +/-- The `f` after a single divstep. -/ +def stepF (eta f g : Int) : Int := ((stepTrans eta g).u * f + (stepTrans eta g).v * g) / 2 + +/-- The `g` after a single divstep. -/ +def stepG (eta f g : Int) : Int := ((stepTrans eta g).q * f + (stepTrans eta g).r * g) / 2 + +/-- One divstep on a whole state. -/ +def step (s : DState) : DState := + ⟨stepEta s.eta s.g, stepF s.eta s.f s.g, stepG s.eta s.f s.g⟩ + +/-- `n` divsteps from `s`, returning the final state together with the accumulated transition +`t`, which satisfies `2 ^ n * f' = t.u * f + t.v * g` and likewise for `g` (`run_two_pow_mul_f`, +`run_two_pow_mul_g`). -/ +def run : Nat → DState → DState × Trans + | 0, s => (s, Trans.id) + | n + 1, s => + let p := run n (step s) + (p.1, p.2.comp (stepTrans s.eta s.g)) + +@[simp] theorem run_zero (s : DState) : run 0 s = (s, Trans.id) := rfl + +theorem run_succ (n : Nat) (s : DState) : + run (n + 1) s = ((run n (step s)).1, (run n (step s)).2.comp (stepTrans s.eta s.g)) := rfl + +/-! ## Single-step lemmas + +Each is proved by splitting into the three branches of `stepTrans` and handing the resulting +concrete linear-arithmetic goal to `omega`, which knows `Int.ediv` and `Int.emod` by numerals. +-/ + +/-- Every divstep matrix has determinant `2`. -/ +theorem det_stepTrans (eta g : Int) : (stepTrans eta g).det = 2 := by + unfold stepTrans Trans.det + split_ifs <;> norm_num + +/-- The division defining the new `f` is exact. -/ +theorem two_mul_stepF (eta f g : Int) (hf : f % 2 = 1) : + 2 * stepF eta f g = (stepTrans eta g).u * f + (stepTrans eta g).v * g := by + unfold stepF stepTrans + split_ifs <;> dsimp only <;> omega + +/-- The division defining the new `g` is exact. -/ +theorem two_mul_stepG (eta f g : Int) (hf : f % 2 = 1) : + 2 * stepG eta f g = (stepTrans eta g).q * f + (stepTrans eta g).r * g := by + unfold stepG stepTrans + split_ifs <;> dsimp only <;> omega + +/-- A divstep preserves oddness of `f`: on the swapping branch the new `f` is the old `g`, which +that branch requires to be odd; otherwise `f` is unchanged. -/ +theorem stepF_emod_two (eta f g : Int) (hf : f % 2 = 1) : stepF eta f g % 2 = 1 := by + unfold stepF stepTrans + split_ifs <;> dsimp only <;> omega + +/-! ## The `n`-step specification -/ + +/-- The contract of a run of `n` divsteps, proved by induction on `n`: the accumulated +determinant is `2 ^ n`, both divisions by `2 ^ n` are exact, and `f` stays odd. -/ +theorem run_spec (n : Nat) (s : DState) (hf : s.f % 2 = 1) : + (run n s).2.det = 2 ^ n ∧ + 2 ^ n * (run n s).1.f = (run n s).2.u * s.f + (run n s).2.v * s.g ∧ + 2 ^ n * (run n s).1.g = (run n s).2.q * s.f + (run n s).2.r * s.g ∧ + (run n s).1.f % 2 = 1 := by + induction n generalizing s with + | zero => refine ⟨by simp, by simp [Trans.id], by simp [Trans.id], hf⟩ + | succ n ih => + have hstep : (step s).f % 2 = 1 := stepF_emod_two s.eta s.f s.g hf + obtain ⟨hdet, hfeq, hgeq, hodd⟩ := ih (step s) hstep + have hFf : 2 * (step s).f = (stepTrans s.eta s.g).u * s.f + (stepTrans s.eta s.g).v * s.g := + two_mul_stepF s.eta s.f s.g hf + have hFg : 2 * (step s).g = (stepTrans s.eta s.g).q * s.f + (stepTrans s.eta s.g).r * s.g := + two_mul_stepG s.eta s.f s.g hf + rw [run_succ] + refine ⟨?_, ?_, ?_, hodd⟩ + · simp only [Trans.det_comp, hdet, det_stepTrans]; ring + · simp only [Trans.comp] + calc (2 : Int) ^ (n + 1) * (run n (step s)).1.f + = 2 * (2 ^ n * (run n (step s)).1.f) := by ring + _ = 2 * ((run n (step s)).2.u * (step s).f + (run n (step s)).2.v * (step s).g) := by + rw [hfeq] + _ = (run n (step s)).2.u * (2 * (step s).f) + + (run n (step s)).2.v * (2 * (step s).g) := by ring + _ = _ := by rw [hFf, hFg]; ring + · simp only [Trans.comp] + calc (2 : Int) ^ (n + 1) * (run n (step s)).1.g + = 2 * (2 ^ n * (run n (step s)).1.g) := by ring + _ = 2 * ((run n (step s)).2.q * (step s).f + (run n (step s)).2.r * (step s).g) := by + rw [hgeq] + _ = (run n (step s)).2.q * (2 * (step s).f) + + (run n (step s)).2.r * (2 * (step s).g) := by ring + _ = _ := by rw [hFf, hFg]; ring + +/-- The accumulated determinant of `n` divsteps is `2 ^ n`. -/ +theorem run_det (n : Nat) (s : DState) (hf : s.f % 2 = 1) : (run n s).2.det = 2 ^ n := + (run_spec n s hf).1 + +/-- The division by `2 ^ n` producing the new `f` is exact. -/ +theorem run_two_pow_mul_f (n : Nat) (s : DState) (hf : s.f % 2 = 1) : + 2 ^ n * (run n s).1.f = (run n s).2.u * s.f + (run n s).2.v * s.g := + (run_spec n s hf).2.1 + +/-- The division by `2 ^ n` producing the new `g` is exact. -/ +theorem run_two_pow_mul_g (n : Nat) (s : DState) (hf : s.f % 2 = 1) : + 2 ^ n * (run n s).1.g = (run n s).2.q * s.f + (run n s).2.r * s.g := + (run_spec n s hf).2.2.1 + +/-- `f` stays odd across a run of divsteps. -/ +theorem run_emod_two_f (n : Nat) (s : DState) (hf : s.f % 2 = 1) : (run n s).1.f % 2 = 1 := + (run_spec n s hf).2.2.2 + +/-! ## Invertibility, and the terminal `f` + +The transition has determinant `2 ^ n` and scales the pair by `1 / 2 ^ n`, so the composite map +`(f, g) ↦ (f', g')` is an *integral* automorphism: its inverse is the adjugate, with the `2 ^ n` +cancelling. The original pair is therefore an integer combination of the final one, which is +what pins the terminal `f` to `± gcd (f, g)`. +-/ + +/-- The original `f` is recovered from the final pair by the adjugate row `(r, -v)`. -/ +theorem run_inv_f (n : Nat) (s : DState) (hf : s.f % 2 = 1) : + s.f = (run n s).2.r * (run n s).1.f - (run n s).2.v * (run n s).1.g := by + have hpow : (2 : Int) ^ n ≠ 0 := pow_ne_zero n (by norm_num) + refine mul_left_cancel₀ hpow ?_ + have hdet := run_det n s hf + have hfe := run_two_pow_mul_f n s hf + have hge := run_two_pow_mul_g n s hf + have expand : + (2 : Int) ^ n * ((run n s).2.r * (run n s).1.f - (run n s).2.v * (run n s).1.g) + = (run n s).2.r * (2 ^ n * (run n s).1.f) - (run n s).2.v * (2 ^ n * (run n s).1.g) := by + ring + rw [expand, hfe, hge] + have : (run n s).2.r * ((run n s).2.u * s.f + (run n s).2.v * s.g) + - (run n s).2.v * ((run n s).2.q * s.f + (run n s).2.r * s.g) + = (run n s).2.det * s.f := by simp only [Trans.det]; ring + rw [this, hdet] + +/-- The original `g` is recovered from the final pair by the adjugate row `(-q, u)`. -/ +theorem run_inv_g (n : Nat) (s : DState) (hf : s.f % 2 = 1) : + s.g = (run n s).2.u * (run n s).1.g - (run n s).2.q * (run n s).1.f := by + have hpow : (2 : Int) ^ n ≠ 0 := pow_ne_zero n (by norm_num) + refine mul_left_cancel₀ hpow ?_ + have hdet := run_det n s hf + have hfe := run_two_pow_mul_f n s hf + have hge := run_two_pow_mul_g n s hf + have expand : + (2 : Int) ^ n * ((run n s).2.u * (run n s).1.g - (run n s).2.q * (run n s).1.f) + = (run n s).2.u * (2 ^ n * (run n s).1.g) - (run n s).2.q * (2 ^ n * (run n s).1.f) := by + ring + rw [expand, hfe, hge] + have : (run n s).2.u * ((run n s).2.q * s.f + (run n s).2.r * s.g) + - (run n s).2.q * ((run n s).2.u * s.f + (run n s).2.v * s.g) + = (run n s).2.det * s.g := by simp only [Trans.det]; ring + rw [this, hdet] + +/-- **Invertibility.** Any common divisor of the final pair divides the original pair: the +adjugate rows express `f` and `g` as integer combinations of `f'` and `g'`. -/ +theorem run_dvd_of_dvd (n : Nat) (s : DState) (hf : s.f % 2 = 1) {k : Int} + (hkf : k ∣ (run n s).1.f) (hkg : k ∣ (run n s).1.g) : k ∣ s.f ∧ k ∣ s.g := by + obtain ⟨a, ha⟩ := hkf + obtain ⟨b, hb⟩ := hkg + refine ⟨⟨(run n s).2.r * a - (run n s).2.v * b, ?_⟩, + ⟨(run n s).2.u * b - (run n s).2.q * a, ?_⟩⟩ + · rw [run_inv_f n s hf, ha, hb]; ring + · rw [run_inv_g n s hf, ha, hb]; ring + +/-- Once `g` has reached `0`, the terminal `f` divides both original entries. -/ +theorem run_dvd_of_g_eq_zero (n : Nat) (s : DState) (hf : s.f % 2 = 1) + (hg : (run n s).1.g = 0) : + (run n s).1.f ∣ s.f ∧ (run n s).1.f ∣ s.g := + run_dvd_of_dvd n s hf dvd_rfl (hg ▸ dvd_zero _) + +/-- The gcd of the pair cannot grow along a run, and by invertibility it cannot shrink either. -/ +theorem run_gcd_dvd (n : Nat) (s : DState) (hf : s.f % 2 = 1) : + Int.gcd (run n s).1.f (run n s).1.g ∣ Int.gcd s.f s.g := by + obtain ⟨h1, h2⟩ := + run_dvd_of_dvd n s hf (Int.gcd_dvd_left _ _) (Int.gcd_dvd_right _ _) + exact_mod_cast Int.dvd_coe_gcd h1 h2 + +/-- Coprimality of the pair is preserved along a run: this is what lets the inversion driver +carry the hypothesis of `run_eq_one_or_neg_one` from one batch to the next. -/ +theorem run_gcd_eq_one (n : Nat) (s : DState) (hf : s.f % 2 = 1) (hcop : Int.gcd s.f s.g = 1) : + Int.gcd (run n s).1.f (run n s).1.g = 1 := + Nat.dvd_one.mp (hcop ▸ run_gcd_dvd n s hf) + +/-- **Termination value.** If the inputs are coprime and the run drives `g` to `0`, the terminal +`f` is a unit, that is `1` or `-1`. This is the fact the inversion driver reads the sign of `f` +for. -/ +theorem run_eq_one_or_neg_one (n : Nat) (s : DState) (hf : s.f % 2 = 1) + (hcop : Int.gcd s.f s.g = 1) (hg : (run n s).1.g = 0) : + (run n s).1.f = 1 ∨ (run n s).1.f = -1 := by + obtain ⟨hdf, hdg⟩ := run_dvd_of_g_eq_zero n s hf hg + have hdvd : (run n s).1.f ∣ ((Int.gcd s.f s.g : Nat) : Int) := Int.dvd_coe_gcd hdf hdg + rw [hcop] at hdvd + have habs : (run n s).1.f.natAbs ∣ 1 := by + simpa using Int.natAbs_dvd_natAbs.mpr hdvd + have := Nat.dvd_one.mp habs + omega + +end CompElliptic.Fields.SafeGCD diff --git a/CompElliptic/Fields/SafeGCD/Divsteps62.lean b/CompElliptic/Fields/SafeGCD/Divsteps62.lean new file mode 100644 index 0000000..cfe9a87 --- /dev/null +++ b/CompElliptic/Fields/SafeGCD/Divsteps62.lean @@ -0,0 +1,215 @@ +/- +Copyright (c) 2026 CompElliptic Contributors. +Released under the Apache License, Version 2.0, or the MIT license, at your option, +as described in the files LICENSE-APACHE and LICENSE-MIT. +Authors: Danny Willems +-/ +import CompElliptic.Fields.SafeGCD.Reference + +/-! +# The batched 62-divstep kernel, ported and checked against the recurrence + +`CompElliptic.Fields.SafeGCD.Divstep` runs divsteps one at a time over `ℤ`, which is what the +correctness proof is about. Implementations do not: they run a *batch* of 62 divsteps on the +bottom words of `f` and `g` alone, in 64-bit wrapping arithmetic, accumulating the transition +matrix, and only then apply that matrix once at full width. That batching kernel is +libsecp256k1's `secp256k1_modinv64_divsteps_62_var`, ported for the Pasta fields in +[zcash/pasta_curves#119](https://github.com/zcash/pasta_curves/pull/119) as +`divsteps_62_var`. + +`divsteps62Var` below is a transcription of that function into Lean, with `UInt64` standing for +`u64` (both wrap modulo `2 ^ 64`) and `Int` for `i64` where the Rust value is exact. + +It is not proved equal to `run 62`; it is *checked* against it. Three things make that check +sharp: + +* the kernel compresses a variable number of divsteps into each inner iteration: a + trailing-zero count batches the halvings, and a small modular-inverse identity cancels up to + 4 bits of `g` at once when `eta ≥ 0` (`f + (((f + 1) &&& 4) <<< 1)` inverts `f` modulo 16) or + up to 6 after the swap when `eta < 0` (`f * (f * f - 2)` inverts `-f` modulo 64). Those + identities are finite facts, so `magic_inverse_mod_16` and `magic_inverse_mod_64` below prove + them outright by `decide`; +* the kernel reads only the bottom words, so agreement with `run 62` on the *full* integers is + exactly the claim that 62 divsteps depend on no more than that; +* it returns the updated `eta` too, so a drift in the branch schedule cannot hide. + +`CompElliptic.Fields.SafeGCD.Pasta` additionally replays the check along the real trajectories +of the port's own regression vectors, where `f` and `g` are full 255-bit values rather than the +single-word samples used here. + +## Provenance + +Ported from `src/fields/modinv62.rs` of +[zcash/pasta_curves#119](https://github.com/zcash/pasta_curves/pull/119), pinned at commit +`9299dfbb19978428ba24229bc13dff25a424dc46` (branch `modinv62`). Every `Rust:` link on a +declaration below points into that commit, so the line numbers stay valid even after the branch +moves; to review against the tip of the branch instead, open the pull request and diff. + +## Scope + +The rest of the Rust module, the signed-radix-`2 ^ 62` limb arithmetic and the driver built on +it, is in `CompElliptic.Fields.SafeGCD.Limbs`. + +One thing here is deliberately *not* ported: `negative_eta_multiplier`, an +`#[inline(never)]` helper that exists only to stop the AArch64 backend speculatively +evaluating both cancellation formulas: + + It computes exactly the expression it replaces, so `divsteps62Step` inlines it; there +is no Lean counterpart to a codegen barrier, and nothing to check. +-/ + +namespace CompElliptic.Fields.SafeGCD + +/-! ## Word helpers -/ + +/-- `u64::MAX`. +Rust: -/ +def allOnes : UInt64 := 0xffffffffffffffff + +/-- Reinterpret a word as a signed 64-bit integer, the Rust `as i64`. +Rust: -/ +def toI64 (x : UInt64) : Int := + if x.toNat < 2 ^ 63 then (x.toNat : Int) else (x.toNat : Int) - 2 ^ 64 + +/-- Truncate an integer to its bottom 64 bits, the Rust `as u64`. +Rust: -/ +def toU64 (x : Int) : UInt64 := UInt64.ofNat (x % (2 ^ 64)).toNat + +/-- The number of trailing zero bits of a word; `64` for zero. The kernel only ever calls it on +a word made nonzero by a sentinel bit. -/ +def trailingZeros (x : UInt64) : Nat := + let rec go : Nat → UInt64 → Nat + | 0, _ => 64 + | fuel + 1, y => if y &&& 1 == 1 then 0 else 1 + go fuel (y >>> 1) + go 64 x + +/-! ## The magic modular inverses the kernel uses + +Both are finite claims about a word modulo a small power of two, so `decide` settles them for +every residue. They are what licenses cancelling several bits of `g` in one inner iteration +instead of one bit per divstep. +-/ + +/-- For odd `f`, `f + (((f + 1) &&& 4) <<< 1)` is the inverse of `f` modulo `16`: the identity +behind the `eta ≥ 0` branch, which cancels up to 4 bits of `g` at a time. +Rust: -/ +theorem magic_inverse_mod_16 : + ∀ f < 16, f % 2 = 1 → (f + ((f + 1) &&& 4) * 2) * f % 16 = 1 := by decide + +/-- For odd `f`, `f * (f * f - 2)` is the inverse of `-f` modulo `64`: the identity behind the +`eta < 0` branch, which cancels up to 6 bits of `g` at a time once `f` and `g` have swapped. +`62` stands for `-2` modulo `64`. +Rust: -/ +theorem magic_inverse_mod_64 : + ∀ f < 64, f % 2 = 1 → (f * (f * f + 62) % 64) * (64 - f) % 64 = 1 := by decide + +/-- The same fact in the form the cancellation actually needs: `g + f * w` with +`w = f * g * (f * f - 2)` is `g * (1 + f ^ 2 * (f ^ 2 - 2)) = g * (1 - (f ^ 2 - 1) ^ 2)`, and +`(f ^ 2 - 1) ^ 2` is divisible by `64` for odd `f` because `f ^ 2 - 1 = (f - 1) * (f + 1)` is +divisible by `8`. +Rust: -/ +theorem cancellation_mod_64 : + ∀ f < 64, f % 2 = 1 → (f * f * (f * f + 62) + 1) % 64 = 0 := by decide + +/-! ## The kernel -/ + +/-- The mutable state of `divsteps62Var`'s inner loop, mirroring the Rust locals. +Rust: -/ +structure Divsteps62State where + /-- The Bernstein-Yang `eta`, exact. -/ + eta : Int + /-- Top-left matrix accumulator. -/ + u : UInt64 + /-- Top-right matrix accumulator. -/ + v : UInt64 + /-- Bottom-left matrix accumulator. -/ + q : UInt64 + /-- Bottom-right matrix accumulator. -/ + r : UInt64 + /-- The working `f`, bottom word only. -/ + f : UInt64 + /-- The working `g`, bottom word only. -/ + g : UInt64 + /-- Divsteps still to perform in this batch. -/ + i : Nat +deriving Repr, Inhabited + +/-- One iteration of the Rust `loop`: batch away the trailing zeros of `g`, then (unless the +batch is complete) swap if `eta < 0` and cancel as many low bits of `g` as the current `eta` +allows. +Rust: -/ +def divsteps62Step (s : Divsteps62State) : Divsteps62State := + let zeros := trailingZeros (s.g ||| (allOnes <<< UInt64.ofNat s.i)) + let g := s.g >>> UInt64.ofNat zeros + let u := s.u <<< UInt64.ofNat zeros + let v := s.v <<< UInt64.ofNat zeros + let eta := s.eta - (zeros : Int) + let i := s.i - zeros + if i = 0 then + { s with eta := eta, u := u, v := v, g := g, i := 0 } + else + let etaNeg := eta < 0 + -- On the negative branch, negate `eta` and replace `f, g` by `g, -f` (and the matrix rows + -- likewise), so the cancellation below always works against an odd `f`. + let eta' := if etaNeg then -eta else eta + let f' := if etaNeg then g else s.f + let g' := if etaNeg then 0 - s.f else g + let u' := if etaNeg then s.q else u + let q' := if etaNeg then 0 - u else s.q + let v' := if etaNeg then s.r else v + let r' := if etaNeg then 0 - v else s.r + let limit : Nat := min (eta'.toNat + 1) i + let w := + if etaNeg then + let m := (allOnes >>> UInt64.ofNat (64 - limit)) &&& 63 + f' * g' * (f' * f' - 2) &&& m + else + let m := (allOnes >>> UInt64.ofNat (64 - limit)) &&& 15 + let w0 := f' + ((f' + 1 &&& 4) <<< 1) + (0 - w0) * g' &&& m + { eta := eta', u := u', v := v', q := q' + u' * w, r := r' + v' * w, + f := f', g := g' + f' * w, i := i } + +/-- Iterate `divsteps62Step` until the batch is complete. Each iteration after the first +consumes at least one divstep, so `70` is ample fuel for a 62-divstep batch. +Rust: -/ +def divsteps62Loop : Nat → Divsteps62State → Divsteps62State + | 0, s => s + | fuel + 1, s => if s.i = 0 then s else divsteps62Loop fuel (divsteps62Step s) + +/-- The port's `divsteps_62_var`: 62 divsteps read off the bottom words `f0` and `g0`, returning +the accumulated transition matrix and the updated `eta`. +Rust: -/ +def divsteps62Var (eta : Int) (f0 g0 : UInt64) : Trans × Int := + let s := divsteps62Loop 70 ⟨eta, 1, 0, 0, 1, f0, g0, 62⟩ + (⟨toI64 s.u, toI64 s.v, toI64 s.q, toI64 s.r⟩, s.eta) + +/-! ## Differential check against the recurrence -/ + +/-- The kernel agrees with `run 62` on the integers `f` and `g`. +Rust: -/ +def batchAgrees (eta f g : Int) : Bool := + let p := run 62 ⟨eta, f, g⟩ + divsteps62Var eta (toU64 f) (toU64 g) == (p.2, p.1.eta) + +/-- A pseudo-random word stream, for sampling the kernel's input space. -/ +def lcg (x : UInt64) : UInt64 := x * 6364136223846793005 + 1442695040888963407 + +/-- `n` sample triples `(eta, f, g)` with `f` odd, drawn from `lcg` starting at `seed`. -/ +def samples : Nat → UInt64 → List (Int × Int × Int) + | 0, _ => [] + | n + 1, x => + let a := lcg x + let b := lcg a + let c := lcg b + ((c.toNat % 61 : Nat) - 30, ((a ||| 1).toNat : Int), (b.toNat : Int)) :: samples n c + +-- The kernel reproduces the recurrence (matrix and `eta`) on 512 single-word samples. +#guard (samples 512 0x9e3779b97f4a7c15).all fun t => batchAgrees t.1 t.2.1 t.2.2 + +-- Boundary `eta` values, where the branch schedule and the `limit` clamp change behaviour. +#guard ([-62, -31, -2, -1, 0, 1, 2, 31, 62] : List Int).all fun e => + batchAgrees e 1 0 && batchAgrees e 1 1 && batchAgrees e (2 ^ 62 - 1) (2 ^ 62 - 2) + && batchAgrees e 0xffffffffffffffff 0 + +end CompElliptic.Fields.SafeGCD diff --git a/CompElliptic/Fields/SafeGCD/Limbs.lean b/CompElliptic/Fields/SafeGCD/Limbs.lean new file mode 100644 index 0000000..38c4bbf --- /dev/null +++ b/CompElliptic/Fields/SafeGCD/Limbs.lean @@ -0,0 +1,513 @@ +/- +Copyright (c) 2026 CompElliptic Contributors. +Released under the Apache License, Version 2.0, or the MIT license, at your option, +as described in the files LICENSE-APACHE and LICENSE-MIT. +Authors: Danny Willems +-/ +import CompElliptic.Fields.SafeGCD.Divsteps62 + +/-! +# The signed-radix-`2 ^ 62` limb layer, ported and checked + +The remaining kernels of the Rust module: the multi-limb arithmetic that `divsteps_62_var` +feeds, and the driver that ties it together. Where +`CompElliptic.Fields.SafeGCD.Reference` applies a batch's transition to `(f, g)` and `(d, e)` +as single integers, these apply it limb by limb, in signed radix `2 ^ 62` over five limbs, +exactly as the Rust does. + +## How the Rust types are modelled + +* `[i64; 5]` becomes `Signed62`, five `Int` limbs. Every value the Rust holds in an `i64` or an + `i128` is in range there, so plain `Int` arithmetic is faithful; the places where the Rust + genuinely relies on wrapping are the ones that go through `UInt64` below. +* `(x as u64 & MASK62) as i64` becomes `low62`, and the arithmetic shift `x >>= 62` becomes + `shr62`. Lean's `%` and `/` on `Int` are Euclidean, so for the positive divisor `2 ^ 62` they + are exactly the non-negative-remainder mask and the floor-division shift, and + `low62_add_shr62` records the decomposition they satisfy. +* the sign masks `x >> 63` (which yield `0` or `-1`) become the conditionals they stand for, and + `x & mask` with such a mask becomes the corresponding `if`. +* `pack62`, `unpack62` and the sign-folding in `shrink_len` operate on `u64` bit patterns, so + they go through `UInt64` and `toI64` / `toU64`. + +## Provenance + +Ported from `src/fields/modinv62.rs` of +[zcash/pasta_curves#119](https://github.com/zcash/pasta_curves/pull/119), pinned at commit +`9299dfbb19978428ba24229bc13dff25a424dc46` (branch `modinv62`). Every `Rust:` link on a +declaration below points into that commit, so the line numbers stay valid even after the branch +moves; to review against the tip of the branch instead, open the pull request and diff. + +## What is checked + +Nothing here is proved; it is checked against what is. The end-to-end check is +`limbInvertCounted` against the proved `invertCounted`, on the port's own regression vectors and +on pseudo-random inputs, both the value and the batch count. Beneath that, the Rust test +module's own differential obligations are reproduced: the sparse-modulus kernels against a +generic-modulus reference that multiplies by every limb (`refUpdateDe`), the first-batch and +terminal specializations against the general ones, `normalize62` against a reference reduction, +and the `pack62` / `unpack62` round trip. + +That split is deliberate. The limb layer is a *representation* of the integers the proof is +about, so the honest claim for it is agreement with the proved layer, not a second proof. +-/ + +namespace CompElliptic.Fields.SafeGCD + +/-! ## Words and limbs -/ + +/-- `MASK62`, the low 62 bits of a word. +Rust: -/ +def mask62 : UInt64 := ((1 : UInt64) <<< (62 : UInt64)) - 1 + +/-- `(x as u64 & MASK62) as i64`: the low 62 bits of a two's-complement value, as a +non-negative integer. -/ +def low62 (x : Int) : Int := x % 2 ^ 62 + +/-- `x >>= 62` on a signed value: an arithmetic shift, that is, floor division. -/ +def shr62 (x : Int) : Int := x / 2 ^ 62 + +/-- The two agree with the radix decomposition, which is what makes the limb loops exact. -/ +theorem low62_add_shr62 (x : Int) : 2 ^ 62 * shr62 x + low62 x = x := by + unfold low62 shr62; omega + +/-- A signed multi-word integer in radix `2 ^ 62`, least-significant limb first: the Rust +`Signed62([i64; 5])`. In canonical form every limb below the active length is in +`[0, 2 ^ 62)` and the top active limb carries the sign. +Rust: -/ +structure Signed62 where + /-- The five limbs, least-significant first. -/ + limbs : Array Int +deriving DecidableEq, Repr, Inhabited + +namespace Signed62 + +/-- Limb `i`, or `0` out of range. -/ +def get (s : Signed62) (i : Nat) : Int := s.limbs.getD i 0 + +/-- Replace limb `i`. -/ +def set (s : Signed62) (i : Nat) (v : Int) : Signed62 := ⟨s.limbs.set! i v⟩ + +/-- Build from a list of limbs. -/ +def ofList (l : List Int) : Signed62 := ⟨l.toArray⟩ + +/-- The all-zero value. -/ +def zero : Signed62 := ofList [0, 0, 0, 0, 0] + +/-- The integer denoted by the limbs. -/ +def val (s : Signed62) : Int := + s.get 0 + 2 ^ 62 * s.get 1 + 2 ^ 124 * s.get 2 + 2 ^ 186 * s.get 3 + 2 ^ 248 * s.get 4 + +end Signed62 + +/-- The per-field constants the kernels need: the Rust `InvParams` trait. +Rust: -/ +structure LimbParams where + /-- The modulus in signed radix `2 ^ 62`, least-significant limb first. -/ + modulus : Signed62 + /-- `m⁻¹ mod 2 ^ 62`. -/ + mu : UInt64 + /-- `R ^ 2 mod m` in signed radix `2 ^ 62`: the initial `e`. -/ + r2 : Signed62 +deriving Repr, Inhabited + +/-! ## `pack62` and `unpack62` -/ + +/-- Repack a canonical `4 x 64`-bit little-endian representation into signed radix `2 ^ 62`. +Rust: -/ +def pack62 (x : Array UInt64) : Signed62 := + let x0 := x.getD 0 0 + let x1 := x.getD 1 0 + let x2 := x.getD 2 0 + let x3 := x.getD 3 0 + Signed62.ofList + [ (x0 &&& mask62).toNat + , (((x0 >>> 62) ||| (x1 <<< 2)) &&& mask62).toNat + , (((x1 >>> 60) ||| (x2 <<< 4)) &&& mask62).toNat + , (((x2 >>> 58) ||| (x3 <<< 6)) &&& mask62).toNat + , (x3 >>> 56).toNat ] + +/-- Repack a canonical signed-62 value in `[0, m)` back into `4 x 64` bits. +Rust: -/ +def unpack62 (v : Signed62) : Array UInt64 := + let v0 := toU64 (v.get 0) + let v1 := toU64 (v.get 1) + let v2 := toU64 (v.get 2) + let v3 := toU64 (v.get 3) + let v4 := toU64 (v.get 4) + #[ v0 ||| (v1 <<< 62) + , (v1 >>> 2) ||| (v2 <<< 60) + , (v2 >>> 4) ||| (v3 <<< 58) + , (v3 >>> 6) ||| (v4 <<< 56) ] + +/-! ## Applying the transition to `f` and `g` -/ + +/-- The limb loop of `update_fg_62_var`, carrying the two `i128` accumulators. +Rust: -/ +def updateFgLoop (t : Trans) : Nat → Nat → Signed62 → Signed62 → Int → Int → + Signed62 × Signed62 × Int × Int + | 0, _, f, g, cf, cg => (f, g, cf, cg) + | n + 1, j, f, g, cf, cg => + let fi := f.get j + let gi := g.get j + let cf' := cf + t.u * fi + t.v * gi + let cg' := cg + t.q * fi + t.r * gi + updateFgLoop t n (j + 1) (f.set (j - 1) (low62 cf')) (g.set (j - 1) (low62 cg')) + (shr62 cf') (shr62 cg') + +/-- Apply the transition to the full-width `f` and `g` over `len` active limbs, dividing +exactly by `2 ^ 62`: the Rust `update_fg_62_var`. +Rust: -/ +def updateFg62Var (len : Nat) (f g : Signed62) (t : Trans) : Signed62 × Signed62 := + let cf := t.u * f.get 0 + t.v * g.get 0 + let cg := t.q * f.get 0 + t.r * g.get 0 + let (f, g, cf, cg) := updateFgLoop t (len - 1) 1 f g (shr62 cf) (shr62 cg) + (f.set (len - 1) cf, g.set (len - 1) cg) + +/-- `update_fg_62_var` specialized to the first batch, where `f` is still the modulus +`[m0, m1, 2, 0, 64]`: the three sparse `f` limbs become shifts of the matrix entries. +Rust: -/ +def updateFg62First (P : LimbParams) (f g : Signed62) (t : Trans) : Signed62 × Signed62 := + let m0 := P.modulus.get 0 + let m1 := P.modulus.get 1 + -- Limb 0. + let cf := t.u * m0 + t.v * g.get 0 + let cg := t.q * m0 + t.r * g.get 0 + let cf := shr62 cf + let cg := shr62 cg + -- Limb 1. + let cf := cf + t.u * m1 + t.v * g.get 1 + let cg := cg + t.q * m1 + t.r * g.get 1 + let f := f.set 0 (low62 cf) + let cf := shr62 cf + let g' := g.set 0 (low62 cg) + let cg := shr62 cg + -- Limb 2: `m2 = 2`. + let cf := cf + t.u * 2 + t.v * g.get 2 + let cg := cg + t.q * 2 + t.r * g.get 2 + let f := f.set 1 (low62 cf) + let cf := shr62 cf + let g' := g'.set 1 (low62 cg) + let cg := shr62 cg + -- Limb 3: `m3 = 0`. + let cf := cf + t.v * g.get 3 + let cg := cg + t.r * g.get 3 + let f := f.set 2 (low62 cf) + let cf := shr62 cf + let g' := g'.set 2 (low62 cg) + let cg := shr62 cg + -- Limb 4: `m4 = 64`. + let cf := cf + t.u * 64 + t.v * g.get 4 + let cg := cg + t.q * 64 + t.r * g.get 4 + let f := f.set 3 (low62 cf) + let cf := shr62 cf + let g' := g'.set 3 (low62 cg) + let cg := shr62 cg + (f.set 4 cf, g'.set 4 cg) + +/-! ## Applying the transition to the coefficients `d` and `e` -/ + +/-- The sign mask `x >> 63`, which is `-1` for a negative limb and `0` otherwise. -/ +def signMask (x : Int) : Int := if x < 0 then -1 else 0 + +/-- `x & signMask y`, the Rust `u & sd`. -/ +def andMask (x m : Int) : Int := if m = -1 then x else 0 + +/-- The `mu`-derived correction that makes a coefficient row divisible by `2 ^ 62`: +`md - ((MU * c + md) & MASK62)`, with the `u64` wrapping invisible past the 62-bit mask. +Rust: -/ +def muCorrect (mu : UInt64) (c md : Int) : Int := + md - ((toI64 mu * c + md) % 2 ^ 62) + +/-- Apply the transition to `d` and `e` modulo `m`, dividing exactly by `2 ^ 62`: the Rust +`update_de_62`, with the sparse modulus limbs 2 and 4 applied as shifts and limb 3 skipped. +Rust: -/ +def updateDe62 (P : LimbParams) (d e : Signed62) (t : Trans) : Signed62 × Signed62 := + let sd := signMask (d.get 4) + let se := signMask (e.get 4) + let md := andMask t.u sd + andMask t.v se + let me := andMask t.q sd + andMask t.r se + let cd := t.u * d.get 0 + t.v * e.get 0 + let ce := t.q * d.get 0 + t.r * e.get 0 + let md := muCorrect P.mu cd md + let me := muCorrect P.mu ce me + -- Limb 0 of the modulus contribution (general multiplication). + let cd := cd + P.modulus.get 0 * md + let ce := ce + P.modulus.get 0 * me + let cd := shr62 cd + let ce := shr62 ce + -- Limb 1 (general multiplication by `m1`). + let cd := cd + t.u * d.get 1 + t.v * e.get 1 + P.modulus.get 1 * md + let ce := ce + t.q * d.get 1 + t.r * e.get 1 + P.modulus.get 1 * me + let d' := d.set 0 (low62 cd) + let cd := shr62 cd + let e' := e.set 0 (low62 ce) + let ce := shr62 ce + -- Limb 2: `m2 = 2`, so the modulus contribution is a shift. + let cd := cd + t.u * d.get 2 + t.v * e.get 2 + md * 2 + let ce := ce + t.q * d.get 2 + t.r * e.get 2 + me * 2 + let d' := d'.set 1 (low62 cd) + let cd := shr62 cd + let e' := e'.set 1 (low62 ce) + let ce := shr62 ce + -- Limb 3: `m3 = 0`, no modulus contribution. + let cd := cd + t.u * d.get 3 + t.v * e.get 3 + let ce := ce + t.q * d.get 3 + t.r * e.get 3 + let d' := d'.set 2 (low62 cd) + let cd := shr62 cd + let e' := e'.set 2 (low62 ce) + let ce := shr62 ce + -- Limb 4: `m4 = 64`, so the modulus contribution is a shift. + let cd := cd + t.u * d.get 4 + t.v * e.get 4 + md * 64 + let ce := ce + t.q * d.get 4 + t.r * e.get 4 + me * 64 + let d' := d'.set 3 (low62 cd) + let cd := shr62 cd + let e' := e'.set 3 (low62 ce) + let ce := shr62 ce + (d'.set 4 cd, e'.set 4 ce) + +/-- `update_de_62` specialized to the first batch, where `d = 0` (so every product against `d` +vanishes and both sign masks are zero, since the initial `e = R ^ 2` is in `[0, m)`). +Rust: -/ +def updateDe62First (P : LimbParams) (e : Signed62) (t : Trans) : Signed62 × Signed62 := + let cd := t.v * e.get 0 + let ce := t.r * e.get 0 + let md := -((toI64 P.mu * cd) % 2 ^ 62) + let me := -((toI64 P.mu * ce) % 2 ^ 62) + let cd := cd + P.modulus.get 0 * md + let ce := ce + P.modulus.get 0 * me + let cd := shr62 cd + let ce := shr62 ce + -- Limb 1. + let cd := cd + t.v * e.get 1 + P.modulus.get 1 * md + let ce := ce + t.r * e.get 1 + P.modulus.get 1 * me + let dOut := Signed62.zero.set 0 (low62 cd) + let cd := shr62 cd + let eOut := Signed62.zero.set 0 (low62 ce) + let ce := shr62 ce + -- Limb 2: `m2 = 2`. + let cd := cd + t.v * e.get 2 + md * 2 + let ce := ce + t.r * e.get 2 + me * 2 + let dOut := dOut.set 1 (low62 cd) + let cd := shr62 cd + let eOut := eOut.set 1 (low62 ce) + let ce := shr62 ce + -- Limb 3: `m3 = 0`. + let cd := cd + t.v * e.get 3 + let ce := ce + t.r * e.get 3 + let dOut := dOut.set 2 (low62 cd) + let cd := shr62 cd + let eOut := eOut.set 2 (low62 ce) + let ce := shr62 ce + -- Limb 4: `m4 = 64`. + let cd := cd + t.v * e.get 4 + md * 64 + let ce := ce + t.r * e.get 4 + me * 64 + let dOut := dOut.set 3 (low62 cd) + let cd := shr62 cd + let eOut := eOut.set 3 (low62 ce) + let ce := shr62 ce + (dOut.set 4 cd, eOut.set 4 ce) + +/-- The terminal coefficient update: only the top row of the matrix, since after the final +batch only `d` feeds the result and `e` is dead. The Rust `update_d_only_62`. +Rust: -/ +def updateDOnly62 (P : LimbParams) (d e : Signed62) (u v : Int) : Signed62 := + let sd := signMask (d.get 4) + let se := signMask (e.get 4) + let md := andMask u sd + andMask v se + let cd := u * d.get 0 + v * e.get 0 + let md := muCorrect P.mu cd md + let cd := cd + P.modulus.get 0 * md + let cd := shr62 cd + let cd := cd + u * d.get 1 + v * e.get 1 + P.modulus.get 1 * md + let d' := d.set 0 (low62 cd) + let cd := shr62 cd + let cd := cd + u * d.get 2 + v * e.get 2 + md * 2 + let d' := d'.set 1 (low62 cd) + let cd := shr62 cd + let cd := cd + u * d.get 3 + v * e.get 3 + let d' := d'.set 2 (low62 cd) + let cd := shr62 cd + let cd := cd + u * d.get 4 + v * e.get 4 + md * 64 + let d' := d'.set 3 (low62 cd) + let cd := shr62 cd + d'.set 4 cd + +/-! ## Normalization, termination test, and length shrinking -/ + +/-- One pass of carry propagation over limbs `0` to `3`. +Rust: -/ +def propagate (r : Signed62) : Signed62 := + let r1 := r.get 1 + shr62 (r.get 0) + let r0 := low62 (r.get 0) + let r2 := r.get 2 + shr62 r1 + let r1 := low62 r1 + let r3 := r.get 3 + shr62 r2 + let r2 := low62 r2 + let r4 := r.get 4 + shr62 r3 + let r3 := low62 r3 + Signed62.ofList [r0, r1, r2, r3, r4] + +/-- Add the modulus to every limb when `cond` holds; the Rust `MODULUS[i] & cond_add`. -/ +def condAddModulus (P : LimbParams) (r : Signed62) (cond : Bool) : Signed62 := + if cond then + Signed62.ofList (List.range 5 |>.map fun i => r.get i + P.modulus.get i) + else r + +/-- Negate every limb; the Rust `(r ^ cond_negate) - cond_negate` with `cond_negate = -1`. -/ +def negateLimbs (r : Signed62) : Signed62 := + Signed62.ofList (List.range 5 |>.map fun i => -r.get i) + +/-- Normalize `r` from `(-2m, m)` to `[0, m)`, negating first if `sign` is negative: +the Rust `normalize_62`. +Rust: -/ +def normalize62 (P : LimbParams) (r : Signed62) (sign : Int) : Signed62 := + let r := condAddModulus P r (r.get 4 < 0) + let r := if sign < 0 then negateLimbs r else r + let r := propagate r + let r := condAddModulus P r (r.get 4 < 0) + propagate r + +/-- Whether the `len` active limbs of `g` are all zero. The Rust ORs the limbs together and +compares with zero, which is the same test. +Rust: -/ +def isZero (g : Signed62) (len : Nat) : Bool := + (List.range len).all fun i => g.get i == 0 + +/-- Shrink the active length by one when both top limbs are pure sign extensions of the limb +below, folding their signs down. The Rust `shrink_len`. +Rust: -/ +def shrinkLen (f g : Signed62) (len : Nat) : Signed62 × Signed62 × Nat := + let fn := f.get (len - 1) + let gn := g.get (len - 1) + if len ≥ 2 && (fn == 0 || fn == -1) && (gn == 0 || gn == -1) then + let foldDown (s : Signed62) (top : Int) : Signed62 := + s.set (len - 2) (toI64 (toU64 (s.get (len - 2)) ||| (toU64 top <<< 62))) + (foldDown f fn, foldDown g gn, len - 1) + else (f, g, len) + +/-! ## The driver -/ + +/-- The loop of `invert_counted` over batches `2 ..= 12`. +Rust: -/ +def limbLoop (P : LimbParams) : Nat → Nat → Int → Nat → Signed62 → Signed62 → Signed62 → + Signed62 → Option (Array UInt64 × Nat) + | 0, _, _, _, _, _, _, _ => none + | n + 1, batch, eta, len, f, g, d, e => + let (t, eta) := divsteps62Var eta (toU64 (f.get 0)) (toU64 (g.get 0)) + let (f, g) := updateFg62Var len f g t + if isZero g len then + let d := updateDOnly62 P d e t.u t.v + let d := normalize62 P d (f.get (len - 1)) + some (unpack62 d, batch) + else + let (d, e) := updateDe62 P d e t + let (f, g, len) := shrinkLen f g len + limbLoop P n (batch + 1) eta len f g d e + +/-- Invert a nonzero field element given by its canonical `4 x 64`-bit little-endian internal +(Montgomery) representation, reporting the batch count: the Rust `invert_counted`. The first +batch is specialized, as there `f` is the sparse modulus, `d` is zero and `0 ≤ e = R ^ 2 < m`. +Rust: -/ +def limbInvertCounted (P : LimbParams) (x : Array UInt64) : Option (Array UInt64 × Nat) := + if x.all (· == 0) then none + else + let f := P.modulus + let g := pack62 x + let e := P.r2 + let (t, eta) := divsteps62Var (-1) (toU64 (f.get 0)) (toU64 (g.get 0)) + let (f, g) := updateFg62First P f g t + let (d, e) := updateDe62First P e t + let (f, g, len) := shrinkLen f g 5 + limbLoop P 11 2 eta len f g d e + +/-- The Rust `invert`: `limbInvertCounted` without the batch count. +Rust: -/ +def limbInvert (P : LimbParams) (x : Array UInt64) : Option (Array UInt64) := + (limbInvertCounted P x).map Prod.fst + +/-! ## A generic-modulus reference for the coefficient update + +The Rust test module carries `ref_update_de`: the same coefficient update with no +specialization at all, multiplying by *every* modulus limb including the zero one. It exists so +the sparse kernels can be checked limb-exact against it rather than against themselves. Ported +here for the same purpose. -/ + +/-- The limb loop of `ref_update_de`, which reads `d` and `e` from captured copies. +Rust: -/ +def refUpdateDeLoop (t : Trans) (m : Signed62) (md me : Int) (dIn eIn : Signed62) : + Nat → Nat → Signed62 → Signed62 → Int → Int → Signed62 × Signed62 × Int × Int + | 0, _, d, e, cd, ce => (d, e, cd, ce) + | n + 1, i, d, e, cd, ce => + let cd' := cd + t.u * dIn.get i + t.v * eIn.get i + m.get i * md + let ce' := ce + t.q * dIn.get i + t.r * eIn.get i + m.get i * me + refUpdateDeLoop t m md me dIn eIn n (i + 1) + (d.set (i - 1) (low62 cd')) (e.set (i - 1) (low62 ce')) (shr62 cd') (shr62 ce') + +/-- The generic-modulus coefficient update: no sparse shifts, no skipped limb, no first-batch +or terminal specialization. The Rust `ref_update_de`. +Rust: -/ +def refUpdateDe (m : Signed62) (mu : UInt64) (d e : Signed62) (t : Trans) : Signed62 × Signed62 := + let sd := signMask (d.get 4) + let se := signMask (e.get 4) + let md := andMask t.u sd + andMask t.v se + let me := andMask t.q sd + andMask t.r se + let cd := t.u * d.get 0 + t.v * e.get 0 + let ce := t.q * d.get 0 + t.r * e.get 0 + let md := muCorrect mu cd md + let me := muCorrect mu ce me + let cd := shr62 (cd + m.get 0 * md) + let ce := shr62 (ce + m.get 0 * me) + let (d', e', cd, ce) := refUpdateDeLoop t m md me d e 4 1 d e cd ce + (d'.set 4 cd, e'.set 4 ce) + +/-! ## Value-level specification of `normalize62` + +Unlike the kernels above, normalization has a short independent specification (reduce to the +canonical residue, negating first if asked), so it is checked against that rather than against +a second copy of itself. -/ + +/-- `normalize62` produced the canonical residue of `± r` and left the limbs canonical. +Rust: -/ +def normalizeOk (P : LimbParams) (m : Nat) (r : Signed62) (sign : Int) : Bool := + let out := normalize62 P r sign + let expected := (if sign < 0 then -r.val else r.val) % (m : Int) + out.val == expected && (List.range 5).all fun i => 0 ≤ out.get i && out.get i < 2 ^ 62 + +/-! ## Walking a real inversion, checking every specialization + +Each specialized kernel is checked against the general one at every batch of a real inversion, +which is where the states that actually arise live: `update_fg_62_first` against +`update_fg_62_var`, `update_de_62_first` and `update_de_62` against `refUpdateDe`, and +`update_d_only_62` against the `d` row of `refUpdateDe` at the terminal batch. -/ + +/-- The batch loop of the specialization check; mirrors `limbLoop`. +Rust: -/ +def specializationsLoop (P : LimbParams) : Nat → Int → Nat → Signed62 → Signed62 → Signed62 → + Signed62 → Bool + | 0, _, _, _, _, _, _ => false + | n + 1, eta, len, f, g, d, e => + let (t, eta) := divsteps62Var eta (toU64 (f.get 0)) (toU64 (g.get 0)) + let (f, g) := updateFg62Var len f g t + let generic := refUpdateDe P.modulus P.mu d e t + if isZero g len then + updateDOnly62 P d e t.u t.v == generic.1 + else if updateDe62 P d e t != generic then false + else + let (f, g, len) := shrinkLen f g len + specializationsLoop P n eta len f g generic.1 generic.2 + +/-- Every specialized kernel is limb-exact against the general one, at every batch of the +inversion of `x`. +Rust: -/ +def specializationsAgree (P : LimbParams) (x : Array UInt64) : Bool := + let f := P.modulus + let g := pack62 x + let (t, eta) := divsteps62Var (-1) (toU64 (f.get 0)) (toU64 (g.get 0)) + let firstFgOk := updateFg62First P f g t == updateFg62Var 5 f g t + let firstDeOk := updateDe62First P P.r2 t == refUpdateDe P.modulus P.mu Signed62.zero P.r2 t + let (f, g) := updateFg62First P f g t + let (d, e) := updateDe62First P P.r2 t + let (f, g, len) := shrinkLen f g 5 + firstFgOk && firstDeOk && specializationsLoop P 11 eta len f g d e + +end CompElliptic.Fields.SafeGCD diff --git a/CompElliptic/Fields/SafeGCD/Pasta.lean b/CompElliptic/Fields/SafeGCD/Pasta.lean new file mode 100644 index 0000000..a1a00af --- /dev/null +++ b/CompElliptic/Fields/SafeGCD/Pasta.lean @@ -0,0 +1,348 @@ +/- +Copyright (c) 2026 CompElliptic Contributors. +Released under the Apache License, Version 2.0, or the MIT license, at your option, +as described in the files LICENSE-APACHE and LICENSE-MIT. +Authors: Danny Willems +-/ +import CompElliptic.Fields.SafeGCD.Limbs +import CompElliptic.Fields.Pasta + +/-! +# The Pasta divstep inversion parameters, and the Rust port checked against them + +Instantiates the proved driver of `CompElliptic.Fields.SafeGCD.Reference` at the two Pasta base +fields, and checks the Rust port of +[zcash/pasta_curves#119](https://github.com/zcash/pasta_curves/pull/119) against it. + +Everything the port hardcodes is *derived* here rather than restated, and then compared: + +* `pallasInv` / `vestaInv` take their modulus from `CompElliptic.Fields.Pasta`, their seed as + `2 ^ 512 mod m` (the Montgomery `R ^ 2`), and their `inv2w` as `((m + 1) / 2) ^ 62 mod m`, + which is `(2 ^ 62)⁻¹` because `(m + 1) / 2` is `2⁻¹` for odd `m`. No magic constant. +* `fp_modulus_limbs_eq` and friends check the port's radix-`2 ^ 62` `MODULUS`, its `MU` + (`m⁻¹ mod 2 ^ 62`, the multiplier that makes each coefficient update divisible by `2 ^ 62`) + and its `R2` against those derivations, by kernel computation. +* the `#guard`s at the end replay the port's own pinned regression vectors: for each, the + driver must return the Fermat inverse `R ^ 2 * x ^ (m - 2) mod m`, an oracle independent of + divsteps, *and* consume exactly the number of 62-divstep batches the port pins. A drifting + batch count means the control flow has diverged, which is the port's own stated rule. + +## Provenance + +Ported from `src/fields/modinv62.rs` of +[zcash/pasta_curves#119](https://github.com/zcash/pasta_curves/pull/119), pinned at commit +`9299dfbb19978428ba24229bc13dff25a424dc46` (branch `modinv62`). Every `Rust:` link on a +declaration below points into that commit, so the line numbers stay valid even after the branch +moves; to review against the tip of the branch instead, open the pull request and diff. + +The batch counts are the sharpest of these: they are a function of the whole divstep trajectory, +so reproducing all twelve of them from an independent implementation of the recurrence pins the +`eta` convention, the branch selection, the batch width, and the termination test at once. +-/ + +namespace CompElliptic.Fields.SafeGCD + +open CompElliptic.Fields.Pasta + +-- The closed numeric facts below evaluate 512-bit powers (`R ^ 2 = 2 ^ 512 mod m`) in the +-- kernel, which needs a deeper recursion budget and a raised `norm_num` exponent threshold. +set_option maxRecDepth 8000 +set_option exponentiation.threshold 600 + +/-! ## Helpers -/ + +/-- Recombine little-endian limbs of `bits` bits each into a natural number. -/ +def ofLimbs (bits : Nat) (l : List Nat) : Nat := + l.foldr (fun w acc => acc * 2 ^ bits + w) 0 + +/-- Binary modular exponentiation. `fuel` bounds the number of squarings; `powMod` supplies +`512`, ample for the 255-bit Pasta exponents. -/ +def powModAux (n : Nat) : Nat → Nat → Nat → Nat → Nat + | 0, _, _, acc => acc + | fuel + 1, b, e, acc => + if e = 0 then acc + else powModAux n fuel (b * b % n) (e / 2) (if e % 2 = 1 then acc * b % n else acc) + +/-- `b ^ e mod n`, by repeated squaring. -/ +def powMod (b e n : Nat) : Nat := powModAux n 512 (b % n) e 1 + +/-- The inverse of `2 ^ w` modulo an odd `m`, as `((m + 1) / 2) ^ w`: for odd `m` the element +`(m + 1) / 2` is `2⁻¹`, since `2 * ((m + 1) / 2) = m + 1 ≡ 1`. -/ +def invTwoPow (m w : Nat) : Nat := powMod ((m + 1) / 2) w m + +/-- The expected inverse, by Fermat's little theorem: `seed * x ^ (m - 2) mod m`. This is the +oracle the `#guard`s below compare the divstep driver against, and it shares no code with it. +Rust: -/ +def fermatInv (P : Params) (x : Nat) : Nat := P.seed * powMod x (P.m - 2) P.m % P.m + +/-! ## Parameters -/ + +/-- Divstep inversion parameters for the Pallas base field `𝔽ₚ` (`pasta_curves`' `Fp`), seeded +with `R ^ 2` so that the result is already in Montgomery form. -/ +def pallasInv : Params where + m := PALLAS_BASE_CARD + width := 62 + seed := 2 ^ 512 % PALLAS_BASE_CARD + inv2w := invTwoPow PALLAS_BASE_CARD 62 + +/-- Divstep inversion parameters for the Vesta base field `𝔽_q` (`pasta_curves`' `Fq`). -/ +def vestaInv : Params where + m := PALLAS_SCALAR_CARD + width := 62 + seed := 2 ^ 512 % PALLAS_SCALAR_CARD + inv2w := invTwoPow PALLAS_SCALAR_CARD 62 + +/-- `pallasInv` is well-formed, so `invert_spec` applies to it. -/ +theorem pallasInv_wf : pallasInv.WF where + odd := by decide + two_lt := by decide + inv2w_spec := by decide + +/-- `vestaInv` is well-formed, so `invert_spec` applies to it. -/ +theorem vestaInv_wf : vestaInv.WF where + odd := by decide + two_lt := by decide + inv2w_spec := by decide + +/-! ## The Rust port's hardcoded constants, re-derived + +Each `List` below is transcribed from `src/fields/modinv62.rs`; each theorem re-derives the +same number from the field constants of `CompElliptic.Fields.Pasta` and compares. +-/ + +/-- `FpParams::MODULUS`: `p` in signed radix `2 ^ 62`, least-significant limb first. +Rust: -/ +def fpModulusLimbs : List Nat := [0x192d30ed00000001, 0x091a63f02533e46e, 2, 0, 0x40] + +/-- `FqParams::MODULUS`. +Rust: -/ +def fqModulusLimbs : List Nat := [0x0c46eb2100000001, 0x091a63f02652a376, 2, 0, 0x40] + +/-- `FpParams::R2`: `2 ^ 512 mod p` in radix `2 ^ 62`. +Rust: -/ +def fpR2Limbs : List Nat := + [0x0c78ecb30000000f, 0x1f4c36f62c37839e, 0x397a99bc3c95d18d, 0x1b506bdee72dc51d, 9] + +/-- `FqParams::R2`. +Rust: -/ +def fqR2Limbs : List Nat := + [0x3c9678ff0000000f, 0x1eed0cf624685b8f, 0x3ae231004ccf5906, 0x1b506bdf33f6aa5f, 9] + +/-- `FpParams::MU`: `p⁻¹ mod 2 ^ 62`. +Rust: -/ +def fpMU : Nat := 0x26d2cf1300000001 + +/-- `FqParams::MU`. +Rust: -/ +def fqMU : Nat := 0x33b914df00000001 + +/-- The port's `Fp` modulus limbs recombine to the Pallas base field's order. +Rust: -/ +theorem fp_modulus_limbs_eq : ofLimbs 62 fpModulusLimbs = PALLAS_BASE_CARD := by decide + +/-- The port's `Fq` modulus limbs recombine to the Vesta base field's order. +Rust: -/ +theorem fq_modulus_limbs_eq : ofLimbs 62 fqModulusLimbs = PALLAS_SCALAR_CARD := by decide + +/-- The port's `Fp` `R2` limbs recombine to `2 ^ 512 mod p`, the seed `pallasInv` uses. +Rust: -/ +theorem fp_r2_limbs_eq : ofLimbs 62 fpR2Limbs = pallasInv.seed := by decide + +/-- The port's `Fq` `R2` limbs recombine to `2 ^ 512 mod q`, the seed `vestaInv` uses. +Rust: -/ +theorem fq_r2_limbs_eq : ofLimbs 62 fqR2Limbs = vestaInv.seed := by decide + +/-- The port's `Fp` `MU` inverts the bottom modulus limb modulo `2 ^ 62`; this is the +`const_assert!` the Rust module pins at compile time. +Rust: -/ +theorem fp_mu_spec : fpMU * fpModulusLimbs.headI % 2 ^ 62 = 1 := by decide + +/-- The port's `Fq` `MU` inverts the bottom modulus limb modulo `2 ^ 62`. +Rust: -/ +theorem fq_mu_spec : fqMU * fqModulusLimbs.headI % 2 ^ 62 = 1 := by decide + +/-- The sparse shape of both moduli that the port's kernels hardcode: limbs `2, 3, 4` are +`2, 0, 64`, so the `k * m` corrections need only two general multiplications. +Rust: -/ +theorem modulus_sparse_shape : + fpModulusLimbs.drop 2 = [2, 0, 64] ∧ fqModulusLimbs.drop 2 = [2, 0, 64] := by decide + +/-! ## The port's pinned regression vectors, replayed + +Each entry is `(x, batches)` with `x` the internal (Montgomery) representation as a natural +number and `batches` the 62-divstep batch count the Rust port pins. The check is that the proved +driver returns the Fermat inverse in exactly that many batches. +-/ + +/-- `PALLAS_VECTORS` from the Rust port, as `(value, pinned batch count)`. +Rust: -/ +def pallasVectors : List (Nat × Nat) := + [ (ofLimbs 64 [0x34786d38fffffffd, 0x992c350be41914ad, 0xffffffffffffffff, 0x3fffffffffffffff], 7) + , (ofLimbs 64 [0x5d1, 0, 0, 0], 8) + , (ofLimbs 64 [0xe096c0a18679d7ae, 0x2cc4e34bc6b6f06a, 0x20eddc12b5a7661d, 0x0c3c5e66537210ad], 8) + , (ofLimbs 64 [1, 0, 0, 0], 9) + , (ofLimbs 64 [0xa162a7d34ad63d62, 0xba71beb748b1fa25, 0x68dc1330fab3847b, 0x2dc205e082f2c197], 10) + , (ofLimbs 64 [0x7ceb4d8e9534361d, 0x8d7879adc97a8e79, 0x19ed0538ef7b15eb, 0x35fb58b0ee06c5b4], 10) ] + +/-- `VESTA_VECTORS` from the Rust port, as `(value, pinned batch count)`. +Rust: -/ +def vestaVectors : List (Nat × Nat) := + [ (ofLimbs 64 [0x5b2b3e9cfffffffd, 0x992c350be3420567, 0xffffffffffffffff, 0x3fffffffffffffff], 7) + , (ofLimbs 64 [0x3ea12df55a259593, 0xd9c58668e724391e, 0xfcf9d370dd78552a, 0x1f7fd517b4f7efeb], 8) + , (ofLimbs 64 [0x4973c2bd762bc27a, 0x9085d079ecab3a12, 0x9f66128174a6731a, 0x3e7853d1fb18fcca], 8) + , (ofLimbs 64 [1, 0, 0, 0], 9) + , (ofLimbs 64 [0xba5f061296c, 0, 0, 0], 10) + , (ofLimbs 64 [0xc00e79606e554fa8, 0x13f9947f444f41d3, 0xd5a780db63e83468, 0x337e275425a385f3], 10) ] + +/-- One vector passes when the driver returns the Fermat inverse in the pinned batch count. +Rust: -/ +def vectorOk (P : Params) (v : Nat × Nat) : Bool := + invertCounted P 12 v.1 == some ((fermatInv P v.1 : Int), v.2) + +-- The Rust port's twelve pinned vectors, replayed against the proved driver. +#guard pallasVectors.all (vectorOk pallasInv) +#guard vestaVectors.all (vectorOk vestaInv) + +-- The port's zero case returns nothing. +#guard invert pallasInv 12 0 == none +#guard invert vestaInv 12 0 == none + +/-! ## The batching kernel along real trajectories + +`CompElliptic.Fields.SafeGCD.Divsteps62` checks the ported `divsteps_62_var` against the +recurrence on single-word samples. Here the same check runs at every batch boundary of a real +inversion, where `f` and `g` are full 255-bit values (and negative for most of the run), so the +claim being tested is that reading only the bottom word is enough. -/ + +/-- Walk an inversion trajectory, checking at each batch that the ported kernel reproduces the +matrix and `eta` that `run 62` produces. Both Pasta parameter sets have `width = 62`, which is +what `batchAgrees` assumes. +Rust: -/ +def trajectoryAgrees (P : Params) : Nat → RState → Bool + | 0, _ => false + | n + 1, s => + if !batchAgrees s.eta s.f s.g then false + else + let s' := s.step P + if s'.g = 0 then true else trajectoryAgrees P n s' + +/-- The initial driver state for inverting `x`. -/ +def initialState (P : Params) (x : Nat) : RState := + ⟨-1, (P.m : Int), (x : Int), 0, (P.seed : Int)⟩ + +#guard pallasVectors.all fun v => trajectoryAgrees pallasInv 12 (initialState pallasInv v.1) +#guard vestaVectors.all fun v => trajectoryAgrees vestaInv 12 (initialState vestaInv v.1) + +/-! ## The limb layer, checked against the proved driver + +`CompElliptic.Fields.SafeGCD.Limbs` ports the radix-`2 ^ 62` kernels and the driver that uses +them. The parameter sets below are built from the Rust literals transcribed above, the very +ones `fp_modulus_limbs_eq`, `fp_r2_limbs_eq` and `fp_mu_spec` pin to the derived values, so +the limb layer runs on the port's own constants, and the checks compare its output with the +proved integer driver's. -/ + +/-- The Rust `FpParams`, as limb-layer parameters. +Rust: -/ +def pallasLimbParams : LimbParams where + modulus := Signed62.ofList (fpModulusLimbs.map Int.ofNat) + mu := UInt64.ofNat fpMU + r2 := Signed62.ofList (fpR2Limbs.map Int.ofNat) + +/-- The Rust `FqParams`, as limb-layer parameters. +Rust: -/ +def vestaLimbParams : LimbParams where + modulus := Signed62.ofList (fqModulusLimbs.map Int.ofNat) + mu := UInt64.ofNat fqMU + r2 := Signed62.ofList (fqR2Limbs.map Int.ofNat) + +/-- A natural number as four little-endian 64-bit limbs, the Rust internal representation. -/ +def toLimbs64 (x : Nat) : Array UInt64 := + #[UInt64.ofNat (x % 2 ^ 64), UInt64.ofNat (x / 2 ^ 64 % 2 ^ 64), + UInt64.ofNat (x / 2 ^ 128 % 2 ^ 64), UInt64.ofNat (x / 2 ^ 192 % 2 ^ 64)] + +/-- The natural number four little-endian 64-bit limbs denote. -/ +def ofLimbs64Array (a : Array UInt64) : Nat := + (a.getD 0 0).toNat + 2 ^ 64 * (a.getD 1 0).toNat + 2 ^ 128 * (a.getD 2 0).toNat + + 2 ^ 192 * (a.getD 3 0).toNat + +/-- The limb driver and the proved integer driver agree on `x`, value and batch count. +Rust: -/ +def limbAgrees (LP : LimbParams) (P : Params) (x : Nat) : Bool := + match limbInvertCounted LP (toLimbs64 x), invertCounted P 12 x with + | some (y, k), some (y', k') => (ofLimbs64Array y : Int) == y' && k == k' + | none, none => true + | _, _ => false + +-- The port's pinned vectors, through the limb layer. +#guard pallasVectors.all fun v => limbAgrees pallasLimbParams pallasInv v.1 +#guard vestaVectors.all fun v => limbAgrees vestaLimbParams vestaInv v.1 + +-- Zero has no inverse in either layer. +#guard limbInvert pallasLimbParams (toLimbs64 0) == none +#guard limbInvert vestaLimbParams (toLimbs64 0) == none + +/-- Pseudo-random field elements below `m`, for sampling the limb layer beyond the pinned +vectors. -/ +def randomElements (m : Nat) : Nat → UInt64 → List Nat + | 0, _ => [] + | n + 1, x => + let a := lcg x + let b := lcg a + let c := lcg b + let d := lcg c + let v := (a.toNat + 2 ^ 64 * b.toNat + 2 ^ 128 * c.toNat + 2 ^ 192 * d.toNat) % m + (if v == 0 then 1 else v) :: randomElements m n d + +#guard (randomElements pallasInv.m 64 0xda39a3ee5e6b4b0d).all + (limbAgrees pallasLimbParams pallasInv) +#guard (randomElements vestaInv.m 64 0x3f786850e387550f).all + (limbAgrees vestaLimbParams vestaInv) + +-- Small and boundary representations: 1, 2, m - 1, m - 2, and single set bits. +#guard ([1, 2, 3, pallasInv.m - 1, pallasInv.m - 2] ++ + (List.range 63).map (fun k => 2 ^ (4 * k + 1))).all + (limbAgrees pallasLimbParams pallasInv) +#guard ([1, 2, 3, vestaInv.m - 1, vestaInv.m - 2] ++ + (List.range 63).map (fun k => 2 ^ (4 * k + 1))).all + (limbAgrees vestaLimbParams vestaInv) + +/-! ### The Rust test module's own differential obligations + +Reproduced here against the same generic-modulus reference the Rust uses: the sparse-modulus +kernels must be limb-exact against a version that multiplies by every modulus limb, and the +first-batch and terminal specializations against the general ones. -/ + +/-- `pack62` and `unpack62` round-trip on a canonical representation. +Rust: -/ +def packRoundTrips (x : Nat) : Bool := ofLimbs64Array (unpack62 (pack62 (toLimbs64 x))) == x + +#guard ([0, 1, 2, pallasInv.m - 1] ++ (List.range 63).map (fun k => 2 ^ (4 * k + 1))).all + packRoundTrips +#guard (randomElements pallasInv.m 64 0x5ba93c9db0cff93f).all packRoundTrips + +-- Every sparse, first-batch and terminal specialization is limb-exact against the +-- generic-modulus reference, at every batch of a real inversion. +#guard pallasVectors.all fun v => + specializationsAgree pallasLimbParams (toLimbs64 v.1) +#guard vestaVectors.all fun v => + specializationsAgree vestaLimbParams (toLimbs64 v.1) +#guard (randomElements pallasInv.m 32 0x2cf24dba5fb0a30e).all fun x => + specializationsAgree pallasLimbParams (toLimbs64 x) +#guard (randomElements vestaInv.m 32 0x18ac3e7343f01690).all fun x => + specializationsAgree vestaLimbParams (toLimbs64 x) + +/-- Normalization inputs: the canonical residue of a random element, minus zero, one or two +copies of the modulus, which is the range `(-2m, m)` the kernel promises to handle. -/ +def normalizeCases (P : Params) (LP : LimbParams) : List Signed62 := + let shifted (x : Nat) (k : Nat) : Signed62 := + Signed62.ofList <| (List.range 5).map fun i => + (pack62 (toLimbs64 x)).get i - (k : Int) * LP.modulus.get i + (List.range 3).flatMap fun k => + (randomElements P.m 8 (0xfcde2b2edba56bf4 + UInt64.ofNat k)).map (shifted · k) + +#guard (normalizeCases pallasInv pallasLimbParams).all fun r => + normalizeOk pallasLimbParams pallasInv.m r 1 && normalizeOk pallasLimbParams pallasInv.m r (-1) +#guard (normalizeCases vestaInv vestaLimbParams).all fun r => + normalizeOk vestaLimbParams vestaInv.m r 1 && normalizeOk vestaLimbParams vestaInv.m r (-1) + +end CompElliptic.Fields.SafeGCD diff --git a/CompElliptic/Fields/SafeGCD/Reference.lean b/CompElliptic/Fields/SafeGCD/Reference.lean new file mode 100644 index 0000000..1a94e4d --- /dev/null +++ b/CompElliptic/Fields/SafeGCD/Reference.lean @@ -0,0 +1,321 @@ +/- +Copyright (c) 2026 CompElliptic Contributors. +Released under the Apache License, Version 2.0, or the MIT license, at your option, +as described in the files LICENSE-APACHE and LICENSE-MIT. +Authors: Danny Willems +-/ +import CompElliptic.Fields.SafeGCD.Divstep +import Mathlib.Data.Int.ModEq + +/-! +# Modular inversion by divsteps: the reference driver, proved correct + +The safegcd inversion of libsecp256k1's `secp256k1_modinv64_var`, and of the Pasta port in +[zcash/pasta_curves#119](https://github.com/zcash/pasta_curves/pull/119), specified and proved +at the level of integers: `f` and `g` run the divstep recurrence of +`CompElliptic.Fields.SafeGCD.Divstep` while a pair of *coefficients* `d` and `e` tracks, modulo +`m`, how the current `f` and `g` were built from the input. + +The scaled invariants carried across batches are + +```text + X * d ≡ seed * f (mod m) X * e ≡ seed * g (mod m) +``` + +where `X` is the value being inverted and `seed` is the initial `e`. They hold at the start for +trivial reasons (`d = 0`, `f = m`, `e = seed`, `g = X`), and `step_good` shows a batch preserves +them. At termination `g = 0` and, by `run_eq_one_or_neg_one` (since `gcd (m, X) = 1`), `f = ±1`, +so the sign-corrected `d` satisfies `X * d ≡ seed (mod m)`. That is `invert_spec`, the headline +result. + +The `seed` is a genuine parameter of the algorithm, and choosing it is what the Rust PR's +"no Montgomery conversion" trick amounts to: + +* `seed = 1` gives the classical inverse `X⁻¹ mod m`; +* `seed = R² mod m` gives `R² / X mod m`, which for a Montgomery representative `X = aR` is + exactly `a⁻¹R`, the Montgomery representative of the inverse. No conversion into or out of + Montgomery form is needed anywhere. + +Nothing here is specific to a radix, a limb count, or a word size: `Params.width` is the number +of divsteps per batch (62 in the port) and `batches` the cap on the number of batches (12). + +## Scope + +This is *partial* correctness: an answer, when produced, is right. Totality, that `g` reaches +`0` within `batches` batches, is the Bernstein-Yang iteration bound +(`⌊(49 · b + 57) / 17⌋` divsteps suffice for `b`-bit inputs, so `738 ≤ 12 · 62` for the Pasta +fields), which is not proved here; the driver returns `none` instead of looping, so the theorem +is unconditional on it. + +## Provenance + +Ported from `src/fields/modinv62.rs` of +[zcash/pasta_curves#119](https://github.com/zcash/pasta_curves/pull/119), pinned at commit +`9299dfbb19978428ba24229bc13dff25a424dc46` (branch `modinv62`). Every `Rust:` link on a +declaration below points into that commit, so the line numbers stay valid even after the branch +moves; to review against the tip of the branch instead, open the pull request and diff. +-/ + +namespace CompElliptic.Fields.SafeGCD + +/-- The parameters of a divstep inversion: the modulus, the number of divsteps per batch, the +seed for the coefficient `e`, and the modular inverse of `2 ^ width` used to undo each batch's +scaling. -/ +structure Params where + /-- The modulus, which must be odd and greater than `2`. -/ + m : Nat + /-- The number of divsteps per batch (`62` in the 64-bit ports). -/ + width : Nat + /-- The initial value of the coefficient `e`; `R ^ 2 mod m` for a Montgomery-native inverse, + `1` for the classical one. -/ + seed : Nat + /-- The inverse of `2 ^ width` modulo `m`, used to divide each batch's coefficient update. -/ + inv2w : Nat +deriving DecidableEq, Repr, Inhabited + +namespace Params + +/-- The modulus as an integer. -/ +def M (P : Params) : Int := (P.m : Int) + +/-- Well-formedness of a parameter set: an odd modulus above `2`, and `inv2w` really inverting +`2 ^ width`. Both are decidable, so a concrete instantiation discharges them by computation. -/ +structure WF (P : Params) : Prop where + /-- The modulus is odd, so every power of two is invertible modulo it. -/ + odd : P.m % 2 = 1 + /-- The modulus exceeds `2`; in particular `1` is its own canonical residue. -/ + two_lt : 2 < P.m + /-- `inv2w` is the inverse of `2 ^ width` modulo `m`. -/ + inv2w_spec : P.inv2w * 2 ^ P.width % P.m = 1 + +end Params + +/-- The state of the inversion driver: the divstep state `(eta, f, g)` together with the +coefficients `d` and `e`, kept reduced modulo `m`. -/ +structure RState where + /-- The divstep `eta`. -/ + eta : Int + /-- The divstep `f`; starts at the modulus. -/ + f : Int + /-- The divstep `g`; starts at the value being inverted and is driven to zero. -/ + g : Int + /-- The coefficient tracking `f`: `X * d ≡ seed * f (mod m)`. -/ + d : Int + /-- The coefficient tracking `g`: `X * e ≡ seed * g (mod m)`. -/ + e : Int +deriving DecidableEq, Repr, Inhabited + +/-- The divstep part of a driver state. -/ +def RState.toDState (s : RState) : DState := ⟨s.eta, s.f, s.g⟩ + +/-- One batch: `width` divsteps applied to `(f, g)`, and the same transition applied to the +coefficients `(d, e)` modulo `m`, undoing the batch's factor of `2 ^ width`. +Rust: -/ +def RState.step (P : Params) (s : RState) : RState := + let p := run P.width s.toDState + { eta := p.1.eta + f := p.1.f + g := p.1.g + d := (p.2.u * s.d + p.2.v * s.e) * (P.inv2w : Int) % P.M + e := (p.2.q * s.d + p.2.r * s.e) * (P.inv2w : Int) % P.M } + +/-- The driver loop: run batches until `g` reaches zero, then read off the sign-corrected `d`, +paired with the number of batches consumed. Returns `none` if the batch budget is exhausted, +which is how the Bernstein-Yang iteration bound is kept out of the trusted statement. +Rust: -/ +def loop (P : Params) : Nat → RState → Option (Int × Nat) + | 0, _ => none + | n + 1, s => + let s' := s.step P + if s'.g = 0 then + some ((if s'.f < 0 then (-s'.d) % P.M else s'.d % P.M), 1) + else (loop P n s').map fun z => (z.1, z.2 + 1) + +/-- Modular inversion by divsteps, reporting the number of batches consumed alongside the +answer. The count is what the Rust port pins in its regression-vector table. +Rust: -/ +def invertCounted (P : Params) (batches : Nat) (x : Nat) : Option (Int × Nat) := + if x = 0 then none + else loop P batches ⟨-1, (P.m : Int), (x : Int), 0, (P.seed : Int)⟩ + +/-- Modular inversion by divsteps: returns `y` with `x * y ≡ seed (mod m)`, given at most +`batches` batches of `P.width` divsteps. +Rust: -/ +def invert (P : Params) (batches : Nat) (x : Nat) : Option Int := + (invertCounted P batches x).map Prod.fst + +/-! ## The invariant -/ + +/-- The invariant carried across batches: `f` odd, the pair coprime, and the two scaled +congruences relating the coefficients to the divstep state. +Rust: -/ +structure Good (P : Params) (x : Int) (s : RState) : Prop where + /-- `f` is odd, the standing precondition of the divstep recurrence. -/ + odd : s.f % 2 = 1 + /-- The pair stays coprime, so the terminal `f` is a unit. -/ + cop : Int.gcd s.f s.g = 1 + /-- `X * d ≡ seed * f (mod m)`. -/ + dInv : x * s.d ≡ (P.seed : Int) * s.f [ZMOD P.M] + /-- `X * e ≡ seed * g (mod m)`. -/ + eInv : x * s.e ≡ (P.seed : Int) * s.g [ZMOD P.M] + +/-- Reducing modulo `n` does not change the residue class. -/ +theorem emod_modEq (a n : Int) : a % n ≡ a [ZMOD n] := Int.emod_emod_of_dvd _ dvd_rfl + +/-- `2 ^ width * inv2w ≡ 1 (mod m)`, the fact that makes each batch's division exact modulo the +modulus. -/ +theorem two_pow_mul_inv2w (P : Params) (hP : P.WF) : + (2 : Int) ^ P.width * (P.inv2w : Int) ≡ 1 [ZMOD P.M] := by + have h : ((P.inv2w * 2 ^ P.width % P.m : Nat) : Int) = ((1 : Nat) : Int) := by + exact_mod_cast congrArg (Nat.cast : Nat → Int) hP.inv2w_spec + have hcast : ((P.inv2w : Int) * 2 ^ P.width) % P.M = 1 := by + push_cast at h ⊢ + simpa [Params.M] using h + calc (2 : Int) ^ P.width * (P.inv2w : Int) + = (P.inv2w : Int) * 2 ^ P.width := by ring + _ ≡ ((P.inv2w : Int) * 2 ^ P.width) % P.M [ZMOD P.M] := (emod_modEq _ _).symm + _ = 1 := hcast + +/-- A batch preserves the invariant. The two congruences are pushed through the batch's matrix +and then rescaled by `inv2w`; `Divstep.run_two_pow_mul_f` is what makes the rescaling land back +on the new `f`. -/ +theorem step_good (P : Params) (hP : P.WF) (x : Int) (s : RState) (h : Good P x s) : + Good P x (s.step P) := by + set p := run P.width s.toDState with hp + have hfodd : s.toDState.f % 2 = 1 := h.odd + have hcop : Int.gcd s.toDState.f s.toDState.g = 1 := h.cop + have hfe : 2 ^ P.width * p.1.f = p.2.u * s.f + p.2.v * s.g := + run_two_pow_mul_f P.width s.toDState hfodd + have hge : 2 ^ P.width * p.1.g = p.2.q * s.f + p.2.r * s.g := + run_two_pow_mul_g P.width s.toDState hfodd + have hinv := two_pow_mul_inv2w P hP + -- One rescaling argument, instantiated at both the `d` row and the `e` row. + have row : ∀ (cu cv fnew : Int), 2 ^ P.width * fnew = cu * s.f + cv * s.g → + x * ((cu * s.d + cv * s.e) * (P.inv2w : Int) % P.M) ≡ (P.seed : Int) * fnew [ZMOD P.M] := by + intro cu cv fnew hnew + have step1 : x * ((cu * s.d + cv * s.e) * (P.inv2w : Int) % P.M) + ≡ x * ((cu * s.d + cv * s.e) * (P.inv2w : Int)) [ZMOD P.M] := + Int.ModEq.mul_left _ (emod_modEq _ _) + have step2 : x * ((cu * s.d + cv * s.e) * (P.inv2w : Int)) + ≡ ((P.seed : Int) * (2 ^ P.width * fnew)) * (P.inv2w : Int) [ZMOD P.M] := by + have hd : cu * (x * s.d) ≡ cu * ((P.seed : Int) * s.f) [ZMOD P.M] := + Int.ModEq.mul_left _ h.dInv + have he : cv * (x * s.e) ≡ cv * ((P.seed : Int) * s.g) [ZMOD P.M] := + Int.ModEq.mul_left _ h.eInv + have hsum : cu * (x * s.d) + cv * (x * s.e) + ≡ cu * ((P.seed : Int) * s.f) + cv * ((P.seed : Int) * s.g) [ZMOD P.M] := hd.add he + have hl : x * ((cu * s.d + cv * s.e) * (P.inv2w : Int)) + = (cu * (x * s.d) + cv * (x * s.e)) * (P.inv2w : Int) := by ring + have hr : cu * ((P.seed : Int) * s.f) + cv * ((P.seed : Int) * s.g) + = (P.seed : Int) * (cu * s.f + cv * s.g) := by ring + rw [hl, hnew, ← hr] + exact hsum.mul_right _ + have step3 : ((P.seed : Int) * (2 ^ P.width * fnew)) * (P.inv2w : Int) + ≡ (P.seed : Int) * fnew * 1 [ZMOD P.M] := by + have : ((P.seed : Int) * (2 ^ P.width * fnew)) * (P.inv2w : Int) + = ((P.seed : Int) * fnew) * (2 ^ P.width * (P.inv2w : Int)) := by ring + rw [this] + exact Int.ModEq.mul_left _ hinv + calc x * ((cu * s.d + cv * s.e) * (P.inv2w : Int) % P.M) + ≡ x * ((cu * s.d + cv * s.e) * (P.inv2w : Int)) [ZMOD P.M] := step1 + _ ≡ ((P.seed : Int) * (2 ^ P.width * fnew)) * (P.inv2w : Int) [ZMOD P.M] := step2 + _ ≡ (P.seed : Int) * fnew * 1 [ZMOD P.M] := step3 + _ = (P.seed : Int) * fnew := by ring + exact + { odd := run_emod_two_f P.width s.toDState hfodd + cop := run_gcd_eq_one P.width s.toDState hfodd hcop + dInv := row p.2.u p.2.v p.1.f hfe + eInv := row p.2.q p.2.r p.1.g hge } + +/-! ## Correctness of the driver -/ + +/-- **Partial correctness of the loop.** Whatever the loop returns is a canonical residue whose +product with `x` is the seed, modulo `m`. -/ +theorem loop_spec (P : Params) (hP : P.WF) (x : Int) : + ∀ (n : Nat) (s : RState) (y : Int) (k : Nat), Good P x s → loop P n s = some (y, k) → + 0 ≤ y ∧ y < P.M ∧ x * y ≡ (P.seed : Int) [ZMOD P.M] := by + have hMpos : (0 : Int) < P.M := by + have := hP.two_lt + simp only [Params.M] + omega + intro n + induction n with + | zero => intro s y k _ h; simp [loop] at h + | succ n ih => + intro s y k hgood heq + have hgood' : Good P x (s.step P) := step_good P hP x s hgood + by_cases hz : (s.step P).g = 0 + · -- Terminal batch: `f` is `±1`, so the sign-corrected `d` is the answer. + have hy : y = if (s.step P).f < 0 then (-(s.step P).d) % P.M else (s.step P).d % P.M := by + have := heq + simp only [loop, hz, if_pos] at this + exact (congrArg Prod.fst (Option.some.inj this)).symm + have hpm : (s.step P).f = 1 ∨ (s.step P).f = -1 := by + have hrun : (run P.width s.toDState).1.g = 0 := hz + have := run_eq_one_or_neg_one P.width s.toDState hgood.odd hgood.cop hrun + simpa [RState.step] using this + refine ⟨?_, ?_, ?_⟩ + · rw [hy]; split <;> exact Int.emod_nonneg _ (ne_of_gt hMpos) + · rw [hy]; split <;> exact Int.emod_lt_of_pos _ hMpos + · rcases hpm with h1 | h1 + · have : y = (s.step P).d % P.M := by rw [hy, if_neg (by omega)] + rw [this] + calc x * ((s.step P).d % P.M) + ≡ x * (s.step P).d [ZMOD P.M] := Int.ModEq.mul_left _ (emod_modEq _ _) + _ ≡ (P.seed : Int) * (s.step P).f [ZMOD P.M] := hgood'.dInv + _ = (P.seed : Int) := by rw [h1]; ring + · have : y = (-(s.step P).d) % P.M := by rw [hy, if_pos (by omega)] + rw [this] + calc x * ((-(s.step P).d) % P.M) + ≡ x * (-(s.step P).d) [ZMOD P.M] := Int.ModEq.mul_left _ (emod_modEq _ _) + _ = -(x * (s.step P).d) := by ring + _ ≡ -((P.seed : Int) * (s.step P).f) [ZMOD P.M] := hgood'.dInv.neg + _ = (P.seed : Int) := by rw [h1]; ring + · have hmap : (loop P n (s.step P)).map (fun z => (z.1, z.2 + 1)) = some (y, k) := by + simpa [loop, hz] using heq + cases hL : loop P n (s.step P) with + | none => rw [hL] at hmap; simp at hmap + | some z => + rw [hL] at hmap + simp only [Option.map_some] at hmap + have hz1 : z.1 = y := congrArg Prod.fst (Option.some.inj hmap) + exact ih (s.step P) z.1 z.2 hgood' (by rw [hL]) |>.imp (fun h => hz1 ▸ h) + (fun h => ⟨hz1 ▸ h.1, hz1 ▸ h.2⟩) + +/-- **Correctness of divstep modular inversion.** For a nonzero `x` below a modulus coprime to +it, any answer the driver produces is the canonical residue `y ∈ [0, m)` with +`x * y ≡ seed (mod m)`. + +With `seed = 1` this is the modular inverse; with `seed = R ^ 2 mod m` and `x` a Montgomery +representative `aR`, it is the Montgomery representative `a⁻¹R` of the inverse, which is what +the Pasta port computes, with no Montgomery reduction anywhere. + +This is the form that also reports the batch count; `invert_spec` drops it. -/ +theorem invertCounted_spec (P : Params) (hP : P.WF) (batches x : Nat) (y : Int) (k : Nat) + (hcop : Nat.gcd x P.m = 1) (h : invertCounted P batches x = some (y, k)) : + 0 ≤ y ∧ y < P.M ∧ (x : Int) * y ≡ (P.seed : Int) [ZMOD P.M] := by + rw [invertCounted] at h + split at h + · exact absurd h (by simp) + · refine loop_spec P hP (x : Int) batches _ y k ?_ h + refine { odd := ?_, cop := ?_, dInv := ?_, eInv := ?_ } + · have h2 : ((P.m % 2 : Nat) : Int) = 1 := by exact_mod_cast hP.odd + push_cast at h2 + simpa using h2 + · simpa [Int.gcd, Nat.gcd_comm] using hcop + · have hz : ((P.seed : Int) * (P.m : Int)) % P.M = 0 := by + simp [Params.M, Int.mul_emod_left] + simp [Int.ModEq, hz] + · simp [Int.ModEq, mul_comm] + +/-- **Correctness of divstep modular inversion**, in the form that drops the batch count. +This generality is what the port's `r2_seeding_equivalence` test asserts empirically. +Rust: -/ +theorem invert_spec (P : Params) (hP : P.WF) (batches x : Nat) (y : Int) + (hcop : Nat.gcd x P.m = 1) (h : invert P batches x = some y) : + 0 ≤ y ∧ y < P.M ∧ (x : Int) * y ≡ (P.seed : Int) [ZMOD P.M] := by + rw [invert, Option.map_eq_some_iff] at h + obtain ⟨z, hz, rfl⟩ := h + exact invertCounted_spec P hP batches x z.1 z.2 hcop hz + +end CompElliptic.Fields.SafeGCD diff --git a/CompElliptic/TrustBoundary.lean b/CompElliptic/TrustBoundary.lean index 9810f6b..5eb2e95 100644 --- a/CompElliptic/TrustBoundary.lean +++ b/CompElliptic/TrustBoundary.lean @@ -12,6 +12,7 @@ import CompElliptic.Curves.Pasta.Fast.Projective import CompElliptic.Curves.Pasta.Fast.Msm import CompElliptic.Curves.Pasta.Fast.ProjectiveMontEquiv import CompElliptic.Fields.Sqrt +import CompElliptic.Fields.SafeGCD.Pasta import CompElliptic.Meta.AxiomCheck /-! @@ -58,6 +59,34 @@ assert_axioms CompElliptic.Fields.TonelliShanks.sqrt?_isSome_of_isSquare assert_axioms CompElliptic.Fields.Pasta.PALLAS_BASE_is_prime assert_axioms CompElliptic.Fields.Pasta.PALLAS_SCALAR_is_prime +/-! ## Divstep (safegcd) modular inversion: standard axioms only + +The general theory (the divstep recurrence, its invertibility over `ℤ`, and the correctness of +the inversion driver) is quantified over every parameter set and input, so it must not reach +beyond the standard axioms. The concrete Pasta facts beneath it (the parameter sets' +well-formedness, the re-derivation of the Rust port's radix-`2 ^ 62` constants, and the two +magic modular inverses its cancellation steps use) are closed numeric claims the kernel evaluates +directly, so they add nothing either. -/ + +assert_axioms CompElliptic.Fields.SafeGCD.run_spec +assert_axioms CompElliptic.Fields.SafeGCD.run_dvd_of_dvd +assert_axioms CompElliptic.Fields.SafeGCD.run_eq_one_or_neg_one +assert_axioms CompElliptic.Fields.SafeGCD.step_good +assert_axioms CompElliptic.Fields.SafeGCD.invertCounted_spec +assert_axioms CompElliptic.Fields.SafeGCD.invert_spec +assert_axioms CompElliptic.Fields.SafeGCD.pallasInv_wf +assert_axioms CompElliptic.Fields.SafeGCD.vestaInv_wf +assert_axioms CompElliptic.Fields.SafeGCD.fp_modulus_limbs_eq +assert_axioms CompElliptic.Fields.SafeGCD.fq_modulus_limbs_eq +assert_axioms CompElliptic.Fields.SafeGCD.fp_r2_limbs_eq +assert_axioms CompElliptic.Fields.SafeGCD.fq_r2_limbs_eq +assert_axioms CompElliptic.Fields.SafeGCD.fp_mu_spec +assert_axioms CompElliptic.Fields.SafeGCD.fq_mu_spec +assert_axioms CompElliptic.Fields.SafeGCD.magic_inverse_mod_16 +assert_axioms CompElliptic.Fields.SafeGCD.magic_inverse_mod_64 +assert_axioms CompElliptic.Fields.SafeGCD.cancellation_mod_64 +assert_axioms CompElliptic.Fields.SafeGCD.low62_add_shr62 + /-! ## The isogeny layer's headline general theorems — standard axioms only -/ assert_axioms CompElliptic.Curves.Pasta.Pallas.iso_map_eq diff --git a/README.md b/README.md index e450750..4369a58 100644 --- a/README.md +++ b/README.md @@ -134,6 +134,15 @@ Early work in progress. Present so far: - a computable Tonelli–Shanks square root for prime fields, soundness and completeness proved, with `pallasBase` / `vestaBase` instances (`CompElliptic/Fields/Sqrt.lean`); - the compressed Pasta point encoding (`toBytes`) for Pallas and Vesta (`CompElliptic/Encodings/`); +- divstep (safegcd) modular inversion: the Bernstein-Yang recurrence with its transition + matrices, proved invertible over `ℤ`, and an inversion driver proved to return the canonical + `y` with `x · y ≡ seed (mod m)` for an arbitrary seed (so seeding with `R²` yields the + Montgomery representative of the inverse directly, with no Montgomery reduction). Instantiated + at both Pasta base fields, with the whole of + [zcash/pasta_curves#119](https://github.com/zcash/pasta_curves/pull/119) ported alongside it: + its radix-`2^62` constants re-derived, its batched 62-divstep kernel and its signed-62 limb + arithmetic checked against the proved recurrence and driver, and every declaration carrying a + line-anchored link into the Rust it came from (`CompElliptic/Fields/SafeGCD/`); - fast Vesta group arithmetic for computing with points rather than only reasoning about them — complete Renes–Costello–Batina projective addition, a windowed Pippenger multi-scalar multiplication and a scalar ladder, each proven to compute the affine group operation it replaces,