diff --git a/lean/VerifiedGarbage/Artifacts/MlDsaRound/X86_64.lean b/lean/VerifiedGarbage/Artifacts/MlDsaRound/X86_64.lean index 69e4ffd2e..8e76071c7 100644 --- a/lean/VerifiedGarbage/Artifacts/MlDsaRound/X86_64.lean +++ b/lean/VerifiedGarbage/Artifacts/MlDsaRound/X86_64.lean @@ -6,6 +6,7 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Round.UseHint import VerifiedGarbage.Proof.MlDsa.X86_64.Round.MakeHint import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YBits import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YHint +import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YUse /-! # ML-DSA (FIPS 204) on x86-64: rounding and hints @@ -105,6 +106,16 @@ def artifacts : List Artifact := [ code := Impl.MlDsa.X86_64.Round.useHint contract := Spec.MlDsa.useHintContract X86_64.abi verified := useHint_verified - spSafe := Code.all_of_allInstrs (by decide +kernel) }] + spSafe := Code.all_of_allInstrs (by decide +kernel) }, + { Spec.MlDsa.useHintApi with + name := Spec.MlDsa.useHintApi.name ++ "_avx2" + target := X86_64.target + doc := Spec.MlDsa.useHintApi.doc (notes := ["The function computes on eight coefficients at a time in AVX2 \ + registers, multiplying by shifts and additions; it needs AVX and AVX2."]) + code := Impl.MlDsa.X86_64.Round.useHintAvx2 + contract := Spec.MlDsa.useHintContract X86_64.abi + verified := useHintY_verified + spSafe := Code.all_of_allInstrs (by decide +kernel) + features := ["avx", "avx2"] }] end VG.Artifacts.MlDsaRound.X86_64 diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean index 2d76aa1a0..46db69523 100644 --- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean +++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean @@ -201,6 +201,7 @@ def subAvx2 : Prog isa := /-- The AVX2 code. -/ def Backend.avx2 : Backend := ⟨nttAvx2, nttInvAvx2, mulAvx2, mulAddAvx2, addAvx2, subAvx2, Round.highBitsAvx2, - Round.lowBitsAvx2, Round.normLtAvx2, Round.makeHintAvx2, Sample.Rej4.rejNTT4Avx2, "_avx2"⟩ + Round.lowBitsAvx2, Round.normLtAvx2, Round.makeHintAvx2, Round.useHintAvx2, Sample.Rej4.rejNTT4Avx2, + "_avx2"⟩ end VG.Impl.MlDsa.X86_64.Arith diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean index b170d109b..e823e7fbe 100644 --- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean +++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean @@ -11,8 +11,8 @@ Key generation, signing and verification call the polynomial arithmetic of one implementation, a `Backend`: the code of `vg_mldsa_ntt`, `vg_mldsa_inv_ntt`, `vg_mldsa_multiply_ntt`, `vg_mldsa_multiply_add_ntt`, `vg_mldsa_add` and `vg_mldsa_sub`, of `vg_mldsa_high_bits`, `vg_mldsa_low_bits`, -`vg_mldsa_norm_lt` and `vg_mldsa_make_hint`, and of `vg_mldsa_rej_ntt_poly4` (which samples four -entries of the matrix `Â` at once), whose names end with `sfx` (e.g. +`vg_mldsa_norm_lt`, `vg_mldsa_make_hint` and `vg_mldsa_use_hint`, and of `vg_mldsa_rej_ntt_poly4` +(which samples four entries of the matrix `Â` at once), whose names end with `sfx` (e.g. `_avx2`; nothing for the SSE2 code, `sse2`). Each is a variant of the interface `MlDsaArith` on x86-64 (`Variants/MlDsaArith/X86_64/`), and the functions that call them are emitted once for each @@ -35,6 +35,7 @@ structure Backend where lowBits : Prog isa normLt : Prog isa makeHint : Prog isa + useHint : Prog isa rej4 : Prog isa /-- What the names of its functions, and of those calling them, end with. -/ sfx : String @@ -42,13 +43,13 @@ structure Backend where /-- The SSE2 code. -/ def Backend.sse2 : Backend := ⟨Arith.ntt, Arith.nttInv, Arith.mul, Arith.mulAdd, Arith.add, Arith.sub, Round.highBits, - Round.lowBits, Round.normLt, Round.makeHint, Sample.Rej4.rejNTT4, ""⟩ + Round.lowBits, Round.normLt, Round.makeHint, Round.useHint, Sample.Rej4.rejNTT4, ""⟩ /-- Every function empty, which the proofs that the functions calling a backend never write `rsp` (and load MXCSR only to restore it) evaluate in its place. -/ def Backend.empty : Backend := ⟨.block [], .block [], .block [], .block [], .block [], .block [], .block [], .block [], .block [], - .block [], .block [], ""⟩ + .block [], .block [], .block [], ""⟩ end VG.Impl.MlDsa.X86_64.Arith diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Round/Avx2.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Round/Avx2.lean index 79292bcbc..b03427169 100644 --- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Round/Avx2.lean +++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Round/Avx2.lean @@ -6,12 +6,13 @@ import VerifiedGarbage.Impl.MlKem.X86_64.Avx # ML-DSA on x86-64: rounding with AVX2 `vg_mldsa_high_bits_avx2`, `vg_mldsa_low_bits_avx2`, -`vg_mldsa_norm_lt_avx2` and `vg_mldsa_make_hint_avx2` are -`vg_mldsa_high_bits`, `vg_mldsa_low_bits`, `vg_mldsa_norm_lt` and -`vg_mldsa_make_hint` (`Round.lean`) on eight +`vg_mldsa_norm_lt_avx2`, `vg_mldsa_make_hint_avx2` and +`vg_mldsa_use_hint_avx2` are `vg_mldsa_high_bits`, `vg_mldsa_low_bits`, +`vg_mldsa_norm_lt`, `vg_mldsa_make_hint` and `vg_mldsa_use_hint` +(`Round.lean`) on eight coefficients at a time, in the doublewords of `ymm` registers: in each 128-bit lane, the VEX.256 form (`toY`) of SSE2 code on the four doublewords -of an `xmm` register (`hbX`, `lbX`, `nlX`, `mhX`), with `rdi` and `r10` at the eight +of an `xmm` register (`hbX`, `lbX`, `nlX`, `mhX`, `uhX`), with `rdi` and `r10` at the eight coefficients of `r` and `out` and `rcx` counting down the 32 vectors. `Decompose` is the reference implementation's, as in `Round.lean`: @@ -41,11 +42,14 @@ def mulX (sh : List Nat) : List Instr := xmov .xmm1 .xmm0 :: sh.flatMap fun k => [xmov .xmm2 .xmm1, .xop (.shift .pslld .xmm2 (BitVec.ofNat 8 k)), xb .paddd .xmm0 .xmm2] +/-- `xmm0 ← f` of `xmm0`, with `127` and `2^(S-1)` in `xmm8` and `xmm9`. -/ +def hfX (g : Nat) : List Instr := + [xb .paddd .xmm0 .xmm8, .xop (.shift .psrld .xmm0 7)] ++ mulX (dSh g) ++ + [xb .paddd .xmm0 .xmm9, .xop (.shift .psrld .xmm0 (BitVec.ofNat 8 (dShift g)))] + /-- `xmm0 ← r₁` of `xmm0`, with `127`, `2^(S-1)` and `m` in `xmm8`, `xmm9` and `xmm10`. -/ def hbX (g : Nat) : List Instr := - [xb .paddd .xmm0 .xmm8, .xop (.shift .psrld .xmm0 7)] ++ mulX (dSh g) ++ - [xb .paddd .xmm0 .xmm9, .xop (.shift .psrld .xmm0 (BitVec.ofNat 8 (dShift g))), xmov .xmm1 .xmm0, - xb .psubd .xmm1 .xmm10, .xop (.shift .psrad .xmm1 31), xb .pand .xmm0 .xmm1] + hfX g ++ [xmov .xmm1 .xmm0, xb .psubd .xmm1 .xmm10, .xop (.shift .psrad .xmm1 31), xb .pand .xmm0 .xmm1] /-- `xmm1 ← xmm0 · 2γ₂`, through `xmm2`. -/ def mul2X (g : Nat) : List Instr := @@ -155,4 +159,36 @@ def makeHintAvx2 : Prog isa := .seq (.block (gammaCmp .rdx ++ [.mov .r10 (.reg .rcx), .mov32 .r9 (.imm 0)])) (.seq (.ite .e (mhY g32) (mhY g88)) (.block [.mov .rax (.reg .r9), .vop .vzeroupper])) +/-! ## `vg_mldsa_use_hint_avx2` + +In each lane, `f` of `r` (`hfX`, which `hbX` reduces modulo `m`), the sign +`P` of `f · 2γ₂ - r` (all ones exactly when `r₀ > 0`), and the sign `H` of +`h | -h` (all ones exactly when the hint is not 0); then `δ = ¬(P + P) ∧ H` +is `1` if `P` is set, `-1` if not, and `0` if the hint is 0, and the result +is `(f + δ + m) mod m`, by two conditional subtractions of `m` (`msub`), as +in `vg_mldsa_use_hint` (`Round.lean`). -/ + +/-- `x ← x - m`, plus `m` if negative, with `m` in `xmm10`, through `t`. -/ +def msub (x t : XReg) : List Instr := + [xb .psubd x .xmm10, xmov t x, .xop (.shift .psrad t 31), xb .pand t .xmm10, xb .paddd x t] + +/-- `xmm4 ← UseHint` of the hint in `xmm5` and `r` in `xmm0`. -/ +def uhX (g : Nat) : List Instr := + xmov .xmm3 .xmm0 :: hfX g ++ xmov .xmm4 .xmm0 :: mul2X g ++ + [xb .psubd .xmm1 .xmm3, .xop (.shift .psrad .xmm1 31), xb .paddd .xmm1 .xmm1, xb .pxor .xmm2 .xmm2, + xb .psubd .xmm2 .xmm5, xb .por .xmm2 .xmm5, .xop (.shift .psrad .xmm2 31), xb .pandn .xmm1 .xmm2, + xb .paddd .xmm4 .xmm1, xb .paddd .xmm4 .xmm10] ++ msub .xmm4 .xmm1 ++ msub .xmm4 .xmm1 + +/-- Eight hints (at `rdi`) and coefficients of `r` (at `rsi`) to `out` (at `r10`). -/ +def uhBodyY (g : Nat) : List Instr := + [.vmovdquLoad .l256 .xmm0 (at_ .rsi 0), .vmovdquLoad .l256 .xmm5 (at_ .rdi 0)] ++ toY (uhX g) ++ + [.vmovdquStore .l256 (at_ .r10 0) .xmm4, .alu .add .rdi (.imm 32), .alu .add .rsi (.imm 32), + .alu .add .r10 (.imm 32)] + +def uhY (g : Nat) : Prog isa := .seq (.block (yC g)) (rcxLoop 32 (uhBodyY g)) + +def useHintAvx2 : Prog isa := + .seq (.block (gammaCmp .rdx ++ [.mov .r10 (.reg .rcx)])) + (.seq (.ite .e (uhY g32) (uhY g88)) (.block [.vop .vzeroupper])) + end VG.Impl.MlDsa.X86_64.Round diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean index b156b4e60..4e38a0278 100644 --- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean +++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean @@ -171,7 +171,7 @@ def ballAt (ct : Ptr) (len tau : Nat) (c : Ptr) : Prog isa := [(.rdi, .ptr ct), (.rsi, .imm len), (.rdx, .imm tau), (.rcx, .ptr c), (.r8, .ptr (sc oSS))] def useHintAt (h r : Ptr) (g2 : Nat) (out : Ptr) : Prog isa := - callAt "vg_mldsa_use_hint" P.useHint [(.rdi, .ptr h), (.rsi, .ptr r), (.rdx, .imm g2), (.rcx, .ptr out)] + callAt ("vg_mldsa_use_hint" ++ P.sfx) P.useHint [(.rdi, .ptr h), (.rsi, .ptr r), (.rdx, .imm g2), (.rcx, .ptr out)] def sbpAt (f : Ptr) (b : Nat) (out : Ptr) (len : Nat) : Prog isa := callAt "vg_mldsa_simple_bit_pack" P.simpleBitPack diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean index ec62cfd5b..7cc453a26 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean @@ -60,6 +60,7 @@ structure BackendOk (B : Backend) : Prop where lowBits : FnOk (fun S => Spec.MlDsa.lowBitsContract X86_64.abi S) B.lowBits normLt : FnOk (fun S => Spec.MlDsa.normLtContract X86_64.abi S) B.normLt makeHint : FnOk (fun S => Spec.MlDsa.makeHintContract X86_64.abi S) B.makeHint + useHint : FnOk (fun S => Spec.MlDsa.useHintContract X86_64.abi S) B.useHint rej4 : Rej4Ok B.rej4 /-- An implementation of the polynomial arithmetic on x86-64. -/ @@ -99,6 +100,8 @@ def ArithImpl.sse2 : ArithImpl where (by decide +kernel) makeHint := FnOk.of Round.makeHint_verified (by decide +kernel) (by decide +kernel) (by decide +kernel) (by decide +kernel) + useHint := FnOk.of Round.useHint_verified (by decide +kernel) (by decide +kernel) (by decide +kernel) + (by decide +kernel) rej4 := ⟨Rej4.rejNTT4_verified, Proof.MlKem.X86_64.nosp_of (by decide +kernel), by decide +kernel, by decide +kernel, Code.all_of_allInstrs (by decide +kernel), fun _ _ _ => Rej4.rejNTT4_ret⟩ } features := [] diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/BackendAvx2.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/BackendAvx2.lean index ec836ecc4..22b9585af 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/BackendAvx2.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/BackendAvx2.lean @@ -3,6 +3,7 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.YNtt import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.YMul import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.YAddSub import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YHint +import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YUse /-! # ML-DSA on x86-64: the polynomial arithmetic with AVX2, as an `ArithImpl` @@ -40,6 +41,8 @@ def ArithImpl.avx2 : ArithImpl where (by decide +kernel) makeHint := FnOk.of Round.makeHintY_verified (by decide +kernel) (by decide +kernel) (by decide +kernel) (by decide +kernel) + useHint := FnOk.of Round.useHintY_verified (by decide +kernel) (by decide +kernel) (by decide +kernel) + (by decide +kernel) rej4 := ⟨Rej4.rejNTT4Avx2_verified, Proof.MlKem.X86_64.nosp_of (by decide +kernel), by decide +kernel, by decide +kernel, Code.all_of_allInstrs (by decide +kernel), fun _ _ _ => Rej4.rejNTT4Avx2_ret⟩ } features := ["avx", "avx2"] diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBlock.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBlock.lean index bc93dce0c..84e8e3739 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBlock.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBlock.lean @@ -39,7 +39,7 @@ theorem hbX_ok {g : Nat} (hg : g = g32 ∨ g = g88) (s : State) (hc : HbC g s) : WP isa (.block (hbX g)) s fun s' => (∀ e < 4, dword (s'.xmm .xmm0) e = hbL g (dword (s.xmm .xmm0) e)) ∧ XOnly [.xmm0, .xmm1, .xmm2] s s' := by rcases hg with rfl | rfl <;> - · simp only [hbX, mulX, dSh_32, dSh_88, dShift_32, dShift_88, xmov, xb, List.cons_append, List.nil_append, + · simp only [hbX, hfX, mulX, dSh_32, dSh_88, dShift_32, dShift_88, xmov, xb, List.cons_append, List.nil_append, List.flatMap_cons, List.flatMap_nil, List.append_nil] vrun [VG.X86_64.eval_movdqa] refine ⟨fun e he => ?_, by xonly⟩ @@ -57,7 +57,7 @@ theorem lbX_ok {g : Nat} (hg : g = g32 ∨ g = g88) (s : State) (hc : HbC g s) ( WP isa (.block (lbX g)) s fun s' => (∀ e < 4, dword (s'.xmm .xmm3) e = lbL g (dword (s.xmm .xmm0) e)) ∧ XOnly [.xmm0, .xmm1, .xmm2, .xmm3] s s' := by rcases hg with rfl | rfl <;> - · simp only [lbX, hbX, mulX, mul2X_32, mul2X_88, dSh_32, dSh_88, dShift_32, dShift_88, + · simp only [lbX, hbX, hfX, mulX, mul2X_32, mul2X_88, dSh_32, dSh_88, dShift_32, dShift_88, VG.Impl.MlDsa.X86_64.Arith.vcadd, xmov, xb, List.cons_append, List.nil_append, List.flatMap_cons, List.flatMap_nil, List.append_nil] vrun [VG.X86_64.eval_movdqa] diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YUse.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YUse.lean new file mode 100644 index 000000000..0b2fece2a --- /dev/null +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YUse.lean @@ -0,0 +1,340 @@ +import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YBits +import VerifiedGarbage.Proof.MlDsa.X86_64.Round.UseHint + +/-! +# ML-DSA on x86-64: `vg_mldsa_use_hint_avx2` + +Untrusted: everything here is checked by Lean. In each doubleword, `uhX` +computes `uhL` of the hint and the coefficient of `r` (`uhX_ok`), which is +the value `vg_mldsa_use_hint` stores (`uhL_toNat`, as `uhS_toNat`); the loop +stores eight of them in each iteration (`YUseL.step`). +-/ + +namespace VG.Proof.MlDsa.X86_64.Round + +open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round VG.Proof.MlDsa.Round +open VG.Spec.MlDsa +open VG.Proof.MlKem.X86_64 (XOnly ifp ifn) +open VG.Impl.MlKem.X86_64 (xb xmov toY) +open VG.Proof.MlDsa.X86_64.Arith (dword_psubd dword_pand dword_psrad sshiftRight31) + +/-! ## A doubleword -/ + +/-- `x - m`, plus `m` if negative, as `msub` computes it. -/ +def msubL (g : Nat) (x : BitVec 32) : BitVec 32 := + (x - BitVec.ofNat 32 (dMod g)) + + ((x - BitVec.ofNat 32 (dMod g)).sshiftRight (min (31 : BitVec 8).toNat 32) &&& BitVec.ofNat 32 (dMod g)) + +/-- What `uhX` computes from the hint `h` and the coefficient `a`. -/ +def uhL (g : Nat) (h a : BitVec 32) : BitVec 32 := + msubL g (msubL g (hbFL g a + (~~~((mul2L g (hbFL g a) - a).sshiftRight (min (31 : BitVec 8).toNat 32) + + (mul2L g (hbFL g a) - a).sshiftRight (min (31 : BitVec 8).toNat 32)) &&& + ((0 - h) ||| h).sshiftRight (min (31 : BitVec 8).toNat 32)) + BitVec.ofNat 32 (dMod g))) + +theorem msubL_toNat (g : Nat) (hm : dMod g ≤ 44) {y : BitVec 32} (hy : y.toNat < 2 ^ 31) : + (msubL g y).toNat = if y.toNat < dMod g then y.toNat else y.toNat - dMod g := by + have hsub : (y - BitVec.ofNat 32 (dMod g)).toNat = (y.toNat + 2 ^ 32 - dMod g) % 2 ^ 32 := by + rw [BitVec.toNat_sub, BitVec.toNat_ofNat, Nat.mod_eq_of_lt (show dMod g < 2 ^ 32 by omega)]; omega + unfold msubL + rw [sshiftRight31] + split + · rename_i h + rw [hsub] at h + rw [show (0 : BitVec 32) &&& BitVec.ofNat 32 (dMod g) = 0 from BitVec.zero_and, + show ∀ x : BitVec 32, x + 0 = x from fun x => BitVec.add_zero x, hsub, + ifn (by omega)] + omega + · rename_i h + rw [hsub] at h + rw [show (-1 : BitVec 32) = BitVec.allOnes 32 by decide, BitVec.allOnes_and, BitVec.sub_add_cancel, + ifp (by omega)] + +/-- The sign of `h | -h` is set exactly when `h` is not 0. -/ +theorem nz_mask (h : BitVec 32) : + ((0 - h) ||| h).sshiftRight (min (31 : BitVec 8).toNat 32) = if h ≠ 0 then -1 else 0 := by + rw [sshiftRight31] + by_cases e : h = 0 + · subst e; rfl + · have h0 : h.toNat ≠ 0 := fun h' => e (BitVec.eq_of_toNat_eq h') + have en : (0 - h).toNat = 2 ^ 32 - h.toNat := by + rw [BitVec.toNat_sub, show (0 : BitVec 32).toNat = 0 from rfl]; have := h.isLt; omega + have hm : (0 - h ||| h).msb = true := by + rw [BitVec.msb_or, BitVec.msb_eq_decide, BitVec.msb_eq_decide, en] + have := h.isLt + by_cases hb : 2 ^ 31 ≤ h.toNat + · simp [hb] + · simp only [Bool.or_eq_true, decide_eq_true_eq]; omega + have h2 := BitVec.msb_eq_decide (0 - h ||| h) + rw [hm] at h2 + have h3 : 2 ^ 31 ≤ (0 - h ||| h).toNat := of_decide_eq_true h2.symm + rw [ite_eq_right (fun h' => absurd h' (by omega)), ifp e] + +theorem uhL_toNat {g : Nat} (h : g ∈ gamma2s) (hv : BitVec 32) {a : BitVec 32} (ha : a.toNat < q) : + (uhL g hv a).toNat = (if hv ≠ 0 then (if hbF g a.toNat * (2 * g) < a.toNat then + hbF g a.toNat + Proof.MlDsa.Round.hbM g + 1 else hbF g a.toNat + Proof.MlDsa.Round.hbM g - 1) + else hbF g a.toNat + Proof.MlDsa.Round.hbM g) % Proof.MlDsa.Round.hbM g := by + have hf := hbF32 h ha + have hle := hbF_le h ha + have hm : dMod g ≤ 44 ∧ 16 ≤ dMod g := by unfold dMod; split <;> decide + have hM := dMod_eq' h + have hfg : hbF g a.toNat * (2 * g) ≤ q - 1 := by rw [← hbM_mul h]; exact Nat.mul_le_mul_right _ hle + rw [← hM] at hle ⊢ + have hmul : (mul2L g (hbFL g a)).toNat = hbF g a.toNat * (2 * g) := by + rw [mul2L_toNat h (by omega), hf] + have hP : (mul2L g (hbFL g a) - a).sshiftRight (min (31 : BitVec 8).toNat 32) = + if hbF g a.toNat * (2 * g) < a.toNat then -1 else 0 := by + rw [sshiftRight31] + have e : (mul2L g (hbFL g a) - a).toNat = (hbF g a.toNat * (2 * g) + 2 ^ 32 - a.toNat) % 2 ^ 32 := by + rw [BitVec.toNat_sub, hmul]; omega + rw [q_eq] at ha hfg + by_cases c : hbF g a.toNat * (2 * g) < a.toNat + · rw [ifn (by rw [e]; omega), ifp c] + · rw [ifp (by rw [e]; omega), ifn c] + unfold uhL + rw [hP, nz_mask] + generalize hx : hbFL g a = F at hf + have hF := hf + by_cases hb : hv ≠ 0 <;> by_cases hp : hbF g a.toNat * (2 * g) < a.toNat + · rw [ifp hb, ifp hp, ifp hb, ifp hp, show ~~~((-1 : BitVec 32) + -1) &&& -1 = 1 by decide] + have e1 : (F + 1 + BitVec.ofNat 32 (dMod g)).toNat = hbF g a.toNat + dMod g + 1 := by + rw [BitVec.toNat_add, BitVec.toNat_add, hF, BitVec.toNat_ofNat, show (1 : BitVec 32).toNat = 1 from rfl]; omega + rw [msubL_toNat g hm.1 (by rw [msubL_toNat g hm.1 (by omega), e1]; split <;> omega), msubL_toNat g hm.1 (by omega), + e1] + rcases (by unfold dMod; split <;> simp : dMod g = 16 ∨ dMod g = 44) with e | e <;> rw [e] at hle ⊢ <;> + (repeat' split) <;> omega + · rw [ifp hb, ifn hp, ifp hb, ifn hp, show ~~~((0 : BitVec 32) + 0) &&& -1 = -1 by decide] + have e1 : (F + -1 + BitVec.ofNat 32 (dMod g)).toNat = hbF g a.toNat + dMod g - 1 := by + rw [BitVec.toNat_add, BitVec.toNat_add, hF, BitVec.toNat_ofNat, show (-1 : BitVec 32).toNat = 2 ^ 32 - 1 from rfl] + omega + rw [msubL_toNat g hm.1 (by rw [msubL_toNat g hm.1 (by omega), e1]; split <;> omega), msubL_toNat g hm.1 (by omega), + e1] + rcases (by unfold dMod; split <;> simp : dMod g = 16 ∨ dMod g = 44) with e | e <;> rw [e] at hle ⊢ <;> + (repeat' split) <;> omega + all_goals + rw [ifn hb, ifn hb, show ∀ x : BitVec 32, x &&& 0 = 0 from fun _ => BitVec.and_zero, + show ∀ x : BitVec 32, x + 0 = x from fun x => BitVec.add_zero x] + have e1 : (F + BitVec.ofNat 32 (dMod g)).toNat = hbF g a.toNat + dMod g := by + rw [BitVec.toNat_add, hF, BitVec.toNat_ofNat]; omega + rw [msubL_toNat g hm.1 (by rw [msubL_toNat g hm.1 (by omega), e1]; split <;> omega), msubL_toNat g hm.1 (by omega), + e1] + rcases (by unfold dMod; split <;> simp : dMod g = 16 ∨ dMod g = 44) with e | e <;> rw [e] at hle ⊢ <;> + (repeat' split) <;> omega + +/-! ## The code on a register -/ + +theorem dword_pandn (a b : BitVec 128) {i : Nat} (hi : i < 4) : + dword (XBinOp.eval .pandn a b) i = ~~~dword a i &&& dword b i := by + apply BitVec.eq_of_getLsbD_eq; intro j hj + simp [XBinOp.eval, dword, hj, show 32 * i + j < 128 by omega] + +theorem uhX_ok {g : Nat} (hg : g = g32 ∨ g = g88) (s : State) (hc : HbC g s) : + WP isa (.block (uhX g)) s fun s' => + (∀ e < 4, dword (s'.xmm .xmm4) e = uhL g (dword (s.xmm .xmm5) e) (dword (s.xmm .xmm0) e)) ∧ + XOnly [.xmm0, .xmm1, .xmm2, .xmm3, .xmm4] s s' := by + rcases hg with rfl | rfl <;> + · simp only [uhX, hfX, msub, mulX, mul2X_32, mul2X_88, dSh_32, dSh_88, dShift_32, dShift_88, xmov, xb, + List.cons_append, List.nil_append, List.flatMap_cons, List.flatMap_nil, List.append_nil] + vrun [VG.X86_64.eval_movdqa] + refine ⟨fun e he => ?_, by xonly⟩ + simp (disch := first | decide | assumption) only [dword_pand, dword_pandn, dword_por, dword_pxor, dword_psrad, + dword_psubd, dword_psrld, dword_pslld, dword_paddd, hc.c8 e he, hc.c9 e he, hc.c10 e he, BitVec.toNat_ofNat, + BitVec.xor_self] + rfl + +end VG.Proof.MlDsa.X86_64.Round + +namespace VG.Proof.MlDsa.X86_64.Arith + +open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round +open VG.Proof.MlDsa.Arith +open VG.Proof.MlDsa.X86_64.Round (HbC uhL uhX_ok) +open VG.Proof.MlKem.X86_64 (Keep XOnly YOnly ylanes yld_ok yconst_ok wp_rcxLoopY ifp ifn ptr_step add_ofNat_zero + lane_setReg lane_setFlags sx32 State.setMem_ymm) +open VG.Impl.MlKem.X86_64 (toY) +open VG.Spec.MlDsa (q n coeffAt Reduced gamma2s) + +/-- `YC` holds while the lanes of its registers are kept. -/ +theorem YC.keep' {g : Nat} {s s' : State} (h : YC g s) + (hl : ∀ r ∈ [XReg.xmm8, .xmm9, .xmm10, .xmm15], ∀ l < 2, s'.lane r l = s.lane r l) : YC g s' := fun l hl' => by + rw [hl _ (by simp) l hl', hl _ (by simp) l hl', hl _ (by simp) l hl', hl _ (by simp) l hl'] + exact h l hl' + +theorem lane_uhX (g : Nat) (hg : g = g32 ∨ g = g88) : laneSseBlock (toY (uhX g)) = some (uhX g) := by + rcases hg with rfl | rfl <;> decide +kernel + +namespace YUseL + +/-- After `i` vectors of eight. -/ +structure Inv (s₀ : State) (g : Nat) (i : Nat) (s : State) : Prop where + rdi : s.gpr .rdi = s₀.gpr .rdi + BitVec.ofNat 64 (32 * i) + rsi : s.gpr .rsi = s₀.gpr .rsi + BitVec.ofNat 64 (32 * i) + r10 : s.gpr .r10 = s₀.gpr .rcx + BitVec.ofNat 64 (32 * i) + rd : s.rd = s₀.rd + wr : s.wr = s₀.wr + yc : YC g s + frame : Frame [pR (s₀.gpr .rcx)] s₀.mem s.mem + coeff : ∀ k < 256, coeffAt s.mem (s₀.gpr .rcx) k = if k < 8 * i then + uhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k) else coeffAt s₀.mem (s₀.gpr .rcx) k + +section +variable {s₀ : State} (hrd : s₀.rd = [pR (s₀.gpr .rdi), pR (s₀.gpr .rsi)]) (hwr : s₀.wr = [pR (s₀.gpr .rcx)]) + (hdz : (pR (s₀.gpr .rdi)).Disjoint (pR (s₀.gpr .rcx))) (hdr : (pR (s₀.gpr .rsi)).Disjoint (pR (s₀.gpr .rcx))) + {g : Nat} (hg : g = g32 ∨ g = g88) +include hrd hwr hdz hdr hg + +theorem step {i : Nat} (hi : i < 32) {s : State} (hI : Inv s₀ g i s) : + WP isa (.block (uhBodyY g ++ ([.alu .sub .rcx (.imm 1)] : List Instr))) s fun s' => + Inv s₀ g (i + 1) s' ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ s'.zf = some (s.gpr .rcx - 1 == 0) := by + have j0 : 8 * i + 8 ≤ 256 := by omega + have hw : pR (s₀.gpr .rcx) ∈ s.wr := by rw [hI.wr, hwr]; simp + have hrz : pR (s₀.gpr .rdi) ∈ s.rd ++ s.wr := by rw [hI.rd, hrd]; simp + have hrr : pR (s₀.gpr .rsi) ∈ s.rd ++ s.wr := by rw [hI.rd, hrd]; simp + have e1 : s.gpr .rsi + BitVec.ofNat 64 0 = coeffAddr (s₀.gpr .rsi) (8 * i) := by + rw [add_ofNat_zero, hI.rsi]; congr 2; omega + have e2 : s.gpr .rdi + BitVec.ofNat 64 0 = coeffAddr (s₀.gpr .rdi) (8 * i) := by + rw [add_ofNat_zero, hI.rdi]; congr 2; omega + have e3 : s.gpr .r10 = coeffAddr (s₀.gpr .rcx) (8 * i) := by rw [hI.r10]; congr 2; omega + rw [uhBodyY, List.append_assoc, List.append_assoc, + show ∀ a b : Instr, [a, b] = [a] ++ [b] from fun _ _ => rfl, List.append_assoc, WP.block_append_iff] + refine WP.mono (yld_ok (by rw [e1]; exact f_in32 hrr j0)) fun s1 ⟨L1, o1⟩ => ?_ + rw [WP.block_append_iff] + refine WP.mono (yld_ok (by rw [o1.rd, o1.wr, o1.gpr, e2]; exact f_in32 hrz j0)) fun s2 ⟨L2, o2⟩ => ?_ + have k2 : YC g s2 := hI.yc.keep' fun r hr' l hl => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr' + rw [o2.lane r (by rcases hr' with rfl | rfl | rfl | rfl <;> decide) l hl, + o1.lane r (by rcases hr' with rfl | rfl | rfl | rfl <;> decide) l hl] + rw [WP.block_append_iff] + refine WP.mono (ylanes (lane_uhX g hg) (P := fun l t => ∀ e < 4, + dword (t.xmm .xmm4) e = uhL g (dword ((s2.proj l).xmm .xmm5) e) (dword ((s2.proj l).xmm .xmm0) e)) + fun l hl => uhX_ok hg _ (k2.hbc hl)) fun s3 ⟨B3, o3⟩ => ?_ + have o13 := (o1.trans o2).trans o3 + have k3 : YC g s3 := k2.keep' fun r hr' l hl => o3.lane r (by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr' + rcases hr' with rfl | rfl | rfl | rfl <;> decide) l hl + have g3 : s3.gpr .r10 = coeffAddr (s₀.gpr .rcx) (8 * i) := by rw [o13.gpr, e3] + have w0 : InRegions s3.wr (s3.gpr .r10) 32 := by rw [o13.wr, g3]; exact f_in32 hw j0 + have hv : ∀ l < 2, ∀ e < 4, uhL g (dword ((s2.proj l).xmm .xmm5) e) (dword ((s2.proj l).xmm .xmm0) e) = + uhL g (coeffAt s₀.mem (s₀.gpr .rdi) (8 * i + 4 * l + e)) (coeffAt s₀.mem (s₀.gpr .rsi) (8 * i + 4 * l + e)) := + fun l hl e he => by + rw [State.proj_xmm, State.proj_xmm, L2 l hl, o2.lane _ (by decide) l hl, L1 l hl, o1.mem, o1.gpr, e1, e2, + dword_readW _ _ he, dword_readW _ _ he, lane_load, lane_load, coeffAddr_add, coeffAddr_add, ← coeffAt_eq, + ← coeffAt_eq, coeffAt_frame hI.frame (by simpa using hdz) (by rw [n_eq]; omega), + coeffAt_frame hI.frame (by simpa using hdr) (by rw [n_eq]; omega)] + vrund [State.store256_eq, State.setMem_gpr, State.setMem_wr, State.setMem_mem, State.setMem_rd, + State.setMem_ymm, w0, sx32] + refine ⟨⟨?_, ?_, ?_, ?_, ?_, ?_, ?_, fun k hk => ?_⟩, ?_⟩ + · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_true, ite_false, reduceCtorEq, State.setMem_gpr] + rw [o13.gpr, hI.rdi]; exact ptr_step _ i 32 + · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_true, ite_false, reduceCtorEq, State.setMem_gpr] + rw [o13.gpr, hI.rsi]; exact ptr_step _ i 32 + · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_true, ite_false, reduceCtorEq, State.setMem_gpr] + rw [o13.gpr, hI.r10]; exact ptr_step _ i 32 + · simp only [RegUpd.rd_setReg, RegUpd.rd_setFlags, State.setMem_rd]; rw [o13.rd, hI.rd] + · simp only [RegUpd.wr_setReg, RegUpd.wr_setFlags, State.setMem_wr]; rw [o13.wr, hI.wr] + · exact k3.keep (rs := []) (by simp) fun r _ l _ => by simp only [lane_setReg, lane_setFlags, State.setMem_lane] + · simp only [RegUpd.mem_setReg, RegUpd.mem_setFlags, State.setMem_mem] + rw [g3, o13.mem]; exact hI.frame.writeW (List.mem_singleton_self _) _ (pR_contains32 _ j0) + · simp only [RegUpd.mem_setReg, RegUpd.mem_setFlags, State.setMem_mem] + rw [g3, o13.mem, coeffAt_write256 _ _ j0 _ hk] + split + · rename_i h + rw [ifp (show k < 8 * (i + 1) by omega), State.ymm, extract_ymm _ _ (by omega)] + have hc : ∀ l < 2, ∀ e < 4, dword (s3.lane .xmm4 l) e = + uhL g (coeffAt s₀.mem (s₀.gpr .rdi) (8 * i + 4 * l + e)) (coeffAt s₀.mem (s₀.gpr .rsi) (8 * i + 4 * l + e)) := + fun l hl e he => by rw [← State.proj_xmm, B3 l hl e he, hv l hl e he] + split + · rename_i h4 + have := hc 0 (by decide) (k - 8 * i) h4 + rw [show 8 * i + 4 * 0 + (k - 8 * i) = k by omega] at this + exact this + · rename_i h4 + have := hc 1 (by decide) (k - 8 * i - 4) (by omega) + rw [show 8 * i + 4 * 1 + (k - 8 * i - 4) = k by omega] at this + exact this + · rename_i h + rw [hI.coeff k hk] + by_cases h' : k < 8 * i + · rw [ifp h', ifp (by omega)] + · rw [ifn h', ifn (by omega)] + · exact ⟨by rw [o13.gpr], by rw [o13.gpr]⟩ + +/-- The constants and the loop: `UseHint` of each coefficient at `out`. -/ +theorem loop_ok {s : State} (hs0 : s.gpr .rdi = s₀.gpr .rdi) (hs1 : s.gpr .rsi = s₀.gpr .rsi) + (hs10 : s.gpr .r10 = s₀.gpr .rcx) (hsrd : s.rd = s₀.rd) (hswr : s.wr = s₀.wr) (hsm : s.mem = s₀.mem) : + WP isa (uhY g) s fun s' => Frame [pR (s₀.gpr .rcx)] s₀.mem s'.mem ∧ + ∀ k < 256, coeffAt s'.mem (s₀.gpr .rcx) k = + uhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k) := by + refine WP.seq (WP.mono (yC_ok g s) fun w ⟨yc, k1, m1, _, _⟩ => ?_) + refine WP.mono (wp_rcxLoopY (N := 32) (by decide) (by decide) _ (fun u o hy _ => + ⟨by rw [o.keep.gpr (by decide), k1.gpr (by decide), Nat.mul_zero, add_ofNat_zero, hs0], + by rw [o.keep.gpr (by decide), k1.gpr (by decide), Nat.mul_zero, add_ofNat_zero, hs1], + by rw [o.keep.gpr (by decide), k1.gpr (by decide), Nat.mul_zero, add_ofNat_zero, hs10], + by rw [o.keep.2.1, k1.2.1, hsrd], by rw [o.keep.2.2, k1.2.2, hswr], + fun l hl => by simp only [State.lane]; rw [o.xmm, hy]; exact yc l hl, + by rw [o.mem, m1, hsm]; exact Frame.refl _ _, fun k _ => by rw [o.mem, m1, hsm, ifn (by omega)]⟩) + fun i hi u hI => step hrd hwr hdz hdr hg hi hI) fun u hI => ⟨hI.frame, fun k hk => by + rw [hI.coeff k hk, ifp (by omega)]⟩ + +end + +end YUseL + +end VG.Proof.MlDsa.X86_64.Arith + +namespace VG.Proof.MlDsa.X86_64.Round + +open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round VG.Proof.MlDsa.Round +open VG.Spec.MlDsa +open VG.Proof.MlKem.X86_64 (Keep Keep.gpr WP.keep gprPreserved_of) + +theorem useHintY_correct (s₀ : State) (hp : useHintK.pre s₀) : + ∃ t s', Exec isa useHintAvx2 s₀ t s' ∧ abiPreserved s₀ s' ∧ useHintK.post s₀ s' := by + have hg : arg32 s₀ .rdx ∈ gamma2s := hp.2.2.2.2.2.2.2.1 + have hr : Reduced s₀.mem (s₀.gpr .rsi) := hp.2.2.2.2.2.2.2.2 + have wp : WP isa useHintAvx2 s₀ fun s' => Frame [pR (s₀.gpr .rcx)] s₀.mem s'.mem ∧ + ∀ k < 256, coeffAt s'.mem (s₀.gpr .rcx) k = + uhL (arg32 s₀ .rdx) (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k) := by + unfold useHintAvx2 + refine WP.seq (WP.mono (prologue_rdx_rcx s₀) fun s₁ ⟨⟨h10, hzf, hm⟩, hk⟩ => ?_) + have go : ∀ g, arg32 s₀ .rdx = g → (g = g32 ∨ g = g88) → WP isa (uhY g) s₁ fun s' => + Frame [pR (s₀.gpr .rcx)] s₀.mem s'.mem ∧ ∀ k < 256, coeffAt s'.mem (s₀.gpr .rcx) k = + uhL (arg32 s₀ .rdx) (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k) := + fun g hge hg' => by + subst hge + exact Arith.YUseL.loop_ok hp.1 hp.2.1 hp.2.2.1 hp.2.2.2.1 hg' (hk.gpr (by decide)) (hk.gpr (by decide)) h10 + hk.2.1 hk.2.2 hm + have hite : WP isa (.ite .e (uhY g32) (uhY g88)) s₁ fun s' => Frame [pR (s₀.gpr .rcx)] s₀.mem s'.mem ∧ + ∀ k < 256, coeffAt s'.mem (s₀.gpr .rcx) k = + uhL (arg32 s₀ .rdx) (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k) := by + refine WP.ite (M := isa) _ (show isa.eval .e s₁ = _ from hzf) (fun h => ?_) (fun h => ?_) + · rw [sub_beq_zero32, decide_eq_true_eq] at h + exact go _ ((gamma_cases hg).1 h) (.inl rfl) + · rw [sub_beq_zero32, decide_eq_false_iff_not] at h + exact go _ ((gamma_cases hg).2 h) (.inr rfl) + refine WP.seq (WP.mono hite fun u ⟨hf, hc⟩ => ?_) + refine WP.mono (Q := fun (u' : State) => u'.mem = u.mem) (by vrund; rfl) fun u' hm' => ?_ + rw [hm'] + exact ⟨hf, hc⟩ + obtain ⟨t, s', he, ⟨hf, hv⟩, hk⟩ := WP.keep [.rax, .rcx, .rdx, .rsi, .rdi, .r10] wp (by decide +kernel) + refine ⟨t, s', he, abiPreserved_of_exec (by decide +kernel) he (gprPreserved_of hk (by decide) hf ?_), ?_⟩ + · simpa using hp.2.2.2.2.2.2.1 + · refine natPolyIs_of_toNat fun k hk => ?_ + rw [hv k hk, zipWith_get _ _ _ hk, hintAt_get _ _ hk, useHint_eq hg, polyAt_val hr hk, Int.toNat_natCast] + rw [uhL_toNat hg _ (hr k hk)] + simp only [decide_eq_true_eq] + +theorem useHintY_ct : ConstantTime isa useHintK.pre useHintK.pub useHintAvx2 := + VG.Taint.constantTime (A := X86_64.taint) (regsLo [.rdi, .rsi, .rcx, .rsp] [.rdx]) + (fun _ _ _ _ hp => agree_regsLo (fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl + exacts [hp.1, hp.2.1, hp.2.2.1, hp.2.2.2.1]) fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact hp.2.2.2.2) + (by taint_decide) + +theorem useHintY_verified : Verified X86_64.target useHintAvx2 (useHintContract X86_64.abi) := + Verified.of_correct useHintY_correct useHintY_ct (by + round_implies [useHintContract, useHintSig, useHintK, hintK, X86_64.abi, X86_64.argRegs] [hintSat] + using hintSat) + +end VG.Proof.MlDsa.X86_64.Round diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean index 744298288..a86c99ea2 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean @@ -55,6 +55,7 @@ def primsWith (B : Arith.Backend) : Prims := mulAdd := B.mulAdd sub := B.sub normLt := B.normLt + useHint := B.useHint rej4 := B.rej4 sfx := B.sfx } @@ -85,10 +86,7 @@ theorem prims_okWith (v : ArithImpl) : PrimsOk (primsWith v.code) where (Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide) (Code.all_of_allInstrs (by lit_decide)) : CalleeOk prims.ball _) - useHint := (CalleeOk.of_verified Proof.MlDsa.X86_64.Round.useHint_verified (by decide) - (Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide) - (Code.all_of_allInstrs (by lit_decide)) : - CalleeOk prims.useHint _) + useHint := calleeOf v.ok.useHint simpleBitPack := (CalleeOk.of_verified Proof.MlDsa.X86_64.Pack.simpleBitPack_verified (by decide) (Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide) (Code.all_of_allInstrs (by lit_decide)) : diff --git a/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Avx2.lean b/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Avx2.lean index 5f57b2be7..a81361808 100644 --- a/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Avx2.lean +++ b/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Avx2.lean @@ -7,8 +7,8 @@ A variant of `MlDsaArith` on x86-64 (see `TCB/Emit.lean`): `vg_mldsa_ntt_avx2`, `vg_mldsa_inv_ntt_avx2`, `vg_mldsa_multiply_ntt_avx2`, `vg_mldsa_multiply_add_ntt_avx2`, `vg_mldsa_add_avx2`, `vg_mldsa_sub_avx2`, `vg_mldsa_high_bits_avx2`, -`vg_mldsa_low_bits_avx2`, `vg_mldsa_norm_lt_avx2` and -`vg_mldsa_make_hint_avx2`, on eight coefficients at a time in AVX2 +`vg_mldsa_low_bits_avx2`, `vg_mldsa_norm_lt_avx2`, +`vg_mldsa_make_hint_avx2` and `vg_mldsa_use_hint_avx2`, on eight coefficients at a time in AVX2 registers, which need AVX and AVX2; key generation, signing and verification calling them need them too. -/ diff --git a/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Sse2.lean b/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Sse2.lean index fd2c251ab..beed231a8 100644 --- a/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Sse2.lean +++ b/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Sse2.lean @@ -6,7 +6,8 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Backend A variant of `MlDsaArith` on x86-64 (see `TCB/Emit.lean`): `vg_mldsa_ntt`, `vg_mldsa_inv_ntt`, `vg_mldsa_multiply_ntt`, `vg_mldsa_multiply_add_ntt`, `vg_mldsa_add`, `vg_mldsa_sub`, `vg_mldsa_high_bits`, -`vg_mldsa_low_bits`, `vg_mldsa_norm_lt` and `vg_mldsa_make_hint`, in the +`vg_mldsa_low_bits`, `vg_mldsa_norm_lt`, `vg_mldsa_make_hint` and +`vg_mldsa_use_hint`, in the baseline ISA (SSE2), which key generation, signing and verification call. -/ diff --git a/src/asm/x86_64/mldsa.rs b/src/asm/x86_64/mldsa.rs index 41484c2ca..360b60b48 100644 --- a/src/asm/x86_64/mldsa.rs +++ b/src/asm/x86_64/mldsa.rs @@ -5287,6 +5287,158 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_use_hint(h: *const [u32; 256], r: ) } +/// The CPU features `vg_mldsa_use_hint_avx2` requires (`Artifact.features`). +pub(crate) const VG_MLDSA_USE_HINT_AVX2_FEATURES: &[&str] = &["avx", "avx2"]; + +/// `UseHint` (FIPS 204 Algorithm 40) of each pair of coefficients of `*h` (a hint bit: true if it is not 0) and `*r`, with `gamma2` = `γ₂`: writes the results to `*out`. +/// +/// Contract: `VG.Spec.MlDsa.useHintContract`. Constant time: only the pointers and `gamma2` may affect timing, not the data. +/// +/// The function computes on eight coefficients at a time in AVX2 registers, multiplying by shifts and additions; it needs AVX and AVX2. +/// +/// # Safety +/// +/// * `h` must be valid for reads of 1024 bytes. +/// * `r` must be valid for reads of 1024 bytes. +/// * `out` must be valid for reads and writes of 1024 bytes. +/// * `gamma2` must be (q - 1)/88 = 95232 or (q - 1)/32 = 261888. +/// * Each of the 256 `u32`s of `r` must be less than `q` = 8380417. +/// * `out` must not overlap `h` or `r` (distinct Rust objects never do). +/// * None of `h`, `r` and `out` may overlap the return address on the stack, or wrap around the end of the address space (no Rust object does). +/// * The CPU must support the `avx` and `avx2` target features. +#[unsafe(naked)] +pub(crate) unsafe extern "sysv64" fn vg_mldsa_use_hint_avx2(h: *const [u32; 256], r: *const [u32; 256], gamma2: u32, out: *mut [u32; 256]) { + core::arch::naked_asm!( + "mov edx, edx", + "cmp edx, 261888", + "mov r10, rcx", + "je 20f", + "mov eax, 127", + "vmovq xmm8, rax", + "vpbroadcastd ymm8, xmm8", + "mov eax, 8388608", + "vmovq xmm9, rax", + "vpbroadcastd ymm9, xmm9", + "mov eax, 44", + "vmovq xmm10, rax", + "vpbroadcastd ymm10, xmm10", + "mov eax, 8380417", + "vmovq xmm15, rax", + "vpbroadcastd ymm15, xmm15", + "mov ecx, 32", + "22:", + "vmovdqu ymm0, YMMWORD PTR [rsi]", + "vmovdqu ymm5, YMMWORD PTR [rdi]", + "vmovdqa ymm3, ymm0", + "vpaddd ymm0, ymm0, ymm8", + "vpsrld ymm0, ymm0, 7", + "vmovdqa ymm1, ymm0", + "vpslld ymm2, ymm1, 1", + "vpaddd ymm0, ymm0, ymm2", + "vpslld ymm2, ymm1, 3", + "vpaddd ymm0, ymm0, ymm2", + "vpslld ymm2, ymm1, 10", + "vpaddd ymm0, ymm0, ymm2", + "vpslld ymm2, ymm1, 11", + "vpaddd ymm0, ymm0, ymm2", + "vpslld ymm2, ymm1, 13", + "vpaddd ymm0, ymm0, ymm2", + "vpaddd ymm0, ymm0, ymm9", + "vpsrld ymm0, ymm0, 24", + "vmovdqa ymm4, ymm0", + "vpslld ymm1, ymm0, 11", + "vpslld ymm2, ymm0, 13", + "vpaddd ymm1, ymm1, ymm2", + "vpslld ymm2, ymm0, 14", + "vpaddd ymm1, ymm1, ymm2", + "vpslld ymm2, ymm0, 15", + "vpaddd ymm1, ymm1, ymm2", + "vpslld ymm2, ymm0, 17", + "vpaddd ymm1, ymm1, ymm2", + "vpsubd ymm1, ymm1, ymm3", + "vpsrad ymm1, ymm1, 31", + "vpaddd ymm1, ymm1, ymm1", + "vpxor ymm2, ymm2, ymm2", + "vpsubd ymm2, ymm2, ymm5", + "vpor ymm2, ymm2, ymm5", + "vpsrad ymm2, ymm2, 31", + "vpandn ymm1, ymm1, ymm2", + "vpaddd ymm4, ymm4, ymm1", + "vpaddd ymm4, ymm4, ymm10", + "vpsubd ymm4, ymm4, ymm10", + "vpsrad ymm1, ymm4, 31", + "vpand ymm1, ymm1, ymm10", + "vpaddd ymm4, ymm4, ymm1", + "vpsubd ymm4, ymm4, ymm10", + "vpsrad ymm1, ymm4, 31", + "vpand ymm1, ymm1, ymm10", + "vpaddd ymm4, ymm4, ymm1", + "vmovdqu YMMWORD PTR [r10], ymm4", + "add rdi, 32", + "add rsi, 32", + "add r10, 32", + "sub rcx, 1", + "jne 22b", + "jmp 21f", + "20:", + "mov eax, 127", + "vmovq xmm8, rax", + "vpbroadcastd ymm8, xmm8", + "mov eax, 2097152", + "vmovq xmm9, rax", + "vpbroadcastd ymm9, xmm9", + "mov eax, 16", + "vmovq xmm10, rax", + "vpbroadcastd ymm10, xmm10", + "mov eax, 8380417", + "vmovq xmm15, rax", + "vpbroadcastd ymm15, xmm15", + "mov ecx, 32", + "23:", + "vmovdqu ymm0, YMMWORD PTR [rsi]", + "vmovdqu ymm5, YMMWORD PTR [rdi]", + "vmovdqa ymm3, ymm0", + "vpaddd ymm0, ymm0, ymm8", + "vpsrld ymm0, ymm0, 7", + "vmovdqa ymm1, ymm0", + "vpslld ymm2, ymm1, 10", + "vpaddd ymm0, ymm0, ymm2", + "vpaddd ymm0, ymm0, ymm9", + "vpsrld ymm0, ymm0, 22", + "vmovdqa ymm4, ymm0", + "vpslld ymm1, ymm0, 19", + "vpslld ymm2, ymm0, 9", + "vpsubd ymm1, ymm1, ymm2", + "vpsubd ymm1, ymm1, ymm3", + "vpsrad ymm1, ymm1, 31", + "vpaddd ymm1, ymm1, ymm1", + "vpxor ymm2, ymm2, ymm2", + "vpsubd ymm2, ymm2, ymm5", + "vpor ymm2, ymm2, ymm5", + "vpsrad ymm2, ymm2, 31", + "vpandn ymm1, ymm1, ymm2", + "vpaddd ymm4, ymm4, ymm1", + "vpaddd ymm4, ymm4, ymm10", + "vpsubd ymm4, ymm4, ymm10", + "vpsrad ymm1, ymm4, 31", + "vpand ymm1, ymm1, ymm10", + "vpaddd ymm4, ymm4, ymm1", + "vpsubd ymm4, ymm4, ymm10", + "vpsrad ymm1, ymm4, 31", + "vpand ymm1, ymm1, ymm10", + "vpaddd ymm4, ymm4, ymm1", + "vmovdqu YMMWORD PTR [r10], ymm4", + "add rdi, 32", + "add rsi, 32", + "add r10, 32", + "sub rcx, 1", + "jne 23b", + "21:", + "vzeroupper", + "ret", + ) +} + /// `RejNTTPoly` (FIPS 204 Algorithm 30): writes the element of `T_q` sampled from the SHAKE128 output of the 34 bytes `*seed` to `*a` (256 coefficients less than `q` = 8380417), and returns 1. Returns 0 if the loop reaches its bound, which is at least 894 bytes of SHAKE128 output (FIPS 204 Appendix C; this happens with probability about 2^-256 or less): `*a` is then unspecified, and the caller must destroy it and treat the operation as failed. /// /// Contract: `VG.Spec.MlDsa.rejNTTContract`. Not constant time in the seed: timing may depend on the pointers and on `*seed` (public in ML-DSA: the seed `ρ` of the matrix and two indices), but not on anything else. diff --git a/src/asm/x86_64/mldsa44.rs b/src/asm/x86_64/mldsa44.rs index c62980516..60c5509fb 100644 --- a/src/asm/x86_64/mldsa44.rs +++ b/src/asm/x86_64/mldsa44.rs @@ -2695,7 +2695,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312 "mov edx, 95232", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 43", @@ -2765,7 +2765,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312 "mov edx, 95232", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 43", @@ -2835,7 +2835,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312 "mov edx, 95232", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 43", @@ -2905,7 +2905,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312 "mov edx, 95232", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 43", @@ -3016,7 +3016,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312 vg_mldsa_unpack_t1 = sym super::mldsa::vg_mldsa_unpack_t1, vg_mldsa_sub_avx2 = sym super::mldsa::vg_mldsa_sub_avx2, vg_mldsa_inv_ntt_avx2 = sym super::mldsa::vg_mldsa_inv_ntt_avx2, - vg_mldsa_use_hint = sym super::mldsa::vg_mldsa_use_hint, + vg_mldsa_use_hint_avx2 = sym super::mldsa::vg_mldsa_use_hint_avx2, vg_mldsa_simple_bit_pack = sym super::mldsa::vg_mldsa_simple_bit_pack, vg_keccak_absorb = sym super::sha3::vg_keccak_absorb, vg_keccak_pad = sym super::sha3::vg_keccak_pad, diff --git a/src/asm/x86_64/mldsa65.rs b/src/asm/x86_64/mldsa65.rs index e83f0160d..b4481c257 100644 --- a/src/asm/x86_64/mldsa65.rs +++ b/src/asm/x86_64/mldsa65.rs @@ -3865,7 +3865,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952 "mov edx, 261888", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 15", @@ -3942,7 +3942,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952 "mov edx, 261888", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 15", @@ -4019,7 +4019,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952 "mov edx, 261888", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 15", @@ -4096,7 +4096,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952 "mov edx, 261888", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 15", @@ -4173,7 +4173,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952 "mov edx, 261888", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 15", @@ -4250,7 +4250,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952 "mov edx, 261888", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 15", @@ -4362,7 +4362,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952 vg_mldsa_unpack_t1 = sym super::mldsa::vg_mldsa_unpack_t1, vg_mldsa_sub_avx2 = sym super::mldsa::vg_mldsa_sub_avx2, vg_mldsa_inv_ntt_avx2 = sym super::mldsa::vg_mldsa_inv_ntt_avx2, - vg_mldsa_use_hint = sym super::mldsa::vg_mldsa_use_hint, + vg_mldsa_use_hint_avx2 = sym super::mldsa::vg_mldsa_use_hint_avx2, vg_mldsa_simple_bit_pack = sym super::mldsa::vg_mldsa_simple_bit_pack, vg_keccak_absorb = sym super::sha3::vg_keccak_absorb, vg_keccak_pad = sym super::sha3::vg_keccak_pad, diff --git a/src/asm/x86_64/mldsa87.rs b/src/asm/x86_64/mldsa87.rs index ed91837d4..c6da2a94a 100644 --- a/src/asm/x86_64/mldsa87.rs +++ b/src/asm/x86_64/mldsa87.rs @@ -5429,7 +5429,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov edx, 261888", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 15", @@ -5520,7 +5520,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov edx, 261888", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 15", @@ -5611,7 +5611,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov edx, 261888", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 15", @@ -5702,7 +5702,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov edx, 261888", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 15", @@ -5793,7 +5793,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov edx, 261888", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 15", @@ -5884,7 +5884,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov edx, 261888", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 15", @@ -5975,7 +5975,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov edx, 261888", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 15", @@ -6066,7 +6066,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov edx, 261888", "mov rcx, rbx", "add rcx, 27648", - "call {vg_mldsa_use_hint}", + "call {vg_mldsa_use_hint_avx2}", "mov rdi, rbx", "add rdi, 27648", "mov esi, 15", @@ -6177,7 +6177,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 vg_mldsa_unpack_t1 = sym super::mldsa::vg_mldsa_unpack_t1, vg_mldsa_sub_avx2 = sym super::mldsa::vg_mldsa_sub_avx2, vg_mldsa_inv_ntt_avx2 = sym super::mldsa::vg_mldsa_inv_ntt_avx2, - vg_mldsa_use_hint = sym super::mldsa::vg_mldsa_use_hint, + vg_mldsa_use_hint_avx2 = sym super::mldsa::vg_mldsa_use_hint_avx2, vg_mldsa_simple_bit_pack = sym super::mldsa::vg_mldsa_simple_bit_pack, vg_keccak_absorb = sym super::sha3::vg_keccak_absorb, vg_keccak_pad = sym super::sha3::vg_keccak_pad,