diff --git a/README.md b/README.md index 25a3fa678..1a8ace9e7 100644 --- a/README.md +++ b/README.md @@ -923,7 +923,7 @@ yours to keep: ✅ -✅ AVX2; SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2) +✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2) ✅ SHA extensions @@ -939,7 +939,7 @@ yours to keep: ✅ -✅ AVX2; SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2) +✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2) ✅ SHA extensions @@ -955,7 +955,7 @@ yours to keep: ✅ -✅ AVX2; SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2) +✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2) ✅ SHA extensions diff --git a/docs/algorithms/ml-dsa-44.toml b/docs/algorithms/ml-dsa-44.toml index 23c04ea14..77df162de 100644 --- a/docs/algorithms/ml-dsa-44.toml +++ b/docs/algorithms/ml-dsa-44.toml @@ -3,4 +3,4 @@ family = "Signatures" specs = ["MlDsa"] modules = ["src/mldsa44.rs"] asm = ["mldsa44", "mldsa"] -optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" } +optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" } diff --git a/docs/algorithms/ml-dsa-65.toml b/docs/algorithms/ml-dsa-65.toml index 04ce839cb..3d5f174db 100644 --- a/docs/algorithms/ml-dsa-65.toml +++ b/docs/algorithms/ml-dsa-65.toml @@ -3,4 +3,4 @@ family = "Signatures" specs = ["MlDsa"] modules = ["src/mldsa65.rs"] asm = ["mldsa65", "mldsa"] -optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" } +optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" } diff --git a/docs/algorithms/ml-dsa-87.toml b/docs/algorithms/ml-dsa-87.toml index 88e39d47a..75d9dc3d0 100644 --- a/docs/algorithms/ml-dsa-87.toml +++ b/docs/algorithms/ml-dsa-87.toml @@ -3,4 +3,4 @@ family = "Signatures" specs = ["MlDsa"] modules = ["src/mldsa87.rs"] asm = ["mldsa87", "mldsa"] -optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" } +optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" } diff --git a/lean/VerifiedGarbage/Artifacts/MlDsaRound/X86_64.lean b/lean/VerifiedGarbage/Artifacts/MlDsaRound/X86_64.lean index 09ad7c11e..69e4ffd2e 100644 --- a/lean/VerifiedGarbage/Artifacts/MlDsaRound/X86_64.lean +++ b/lean/VerifiedGarbage/Artifacts/MlDsaRound/X86_64.lean @@ -4,6 +4,8 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Round.Bits import VerifiedGarbage.Proof.MlDsa.X86_64.Round.NormLt 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 /-! # ML-DSA (FIPS 204) on x86-64: rounding and hints @@ -43,6 +45,26 @@ def artifacts : List Artifact := [ contract := Spec.MlDsa.lowBitsContract X86_64.abi verified := lowBits_verified spSafe := Code.all_of_allInstrs (by decide +kernel) }, + { Spec.MlDsa.highBitsApi with + name := Spec.MlDsa.highBitsApi.name ++ "_avx2" + target := X86_64.target + doc := Spec.MlDsa.highBitsApi.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.highBitsAvx2 + contract := Spec.MlDsa.highBitsContract X86_64.abi + verified := highBitsY_verified + spSafe := Code.all_of_allInstrs (by decide +kernel) + features := ["avx", "avx2"] }, + { Spec.MlDsa.lowBitsApi with + name := Spec.MlDsa.lowBitsApi.name ++ "_avx2" + target := X86_64.target + doc := Spec.MlDsa.lowBitsApi.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.lowBitsAvx2 + contract := Spec.MlDsa.lowBitsContract X86_64.abi + verified := lowBitsY_verified + spSafe := Code.all_of_allInstrs (by decide +kernel) + features := ["avx", "avx2"] }, { Spec.MlDsa.normLtApi with target := X86_64.target doc := Spec.MlDsa.normLtApi.doc @@ -57,6 +79,26 @@ def artifacts : List Artifact := [ contract := Spec.MlDsa.makeHintContract X86_64.abi verified := makeHint_verified spSafe := Code.all_of_allInstrs (by decide +kernel) }, + { Spec.MlDsa.normLtApi with + name := Spec.MlDsa.normLtApi.name ++ "_avx2" + target := X86_64.target + doc := Spec.MlDsa.normLtApi.doc (notes := ["The function compares eight coefficients at a time in AVX2 \ + registers; it needs AVX and AVX2."]) + code := Impl.MlDsa.X86_64.Round.normLtAvx2 + contract := Spec.MlDsa.normLtContract X86_64.abi + verified := normLtY_verified + spSafe := Code.all_of_allInstrs (by decide +kernel) + features := ["avx", "avx2"] }, + { Spec.MlDsa.makeHintApi with + name := Spec.MlDsa.makeHintApi.name ++ "_avx2" + target := X86_64.target + doc := Spec.MlDsa.makeHintApi.doc (notes := ["The function computes eight hints at a time in AVX2 \ + registers, multiplying by shifts and additions; it needs AVX and AVX2."]) + code := Impl.MlDsa.X86_64.Round.makeHintAvx2 + contract := Spec.MlDsa.makeHintContract X86_64.abi + verified := makeHintY_verified + spSafe := Code.all_of_allInstrs (by decide +kernel) + features := ["avx", "avx2"] }, { Spec.MlDsa.useHintApi with target := X86_64.target doc := Spec.MlDsa.useHintApi.doc diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean index 3c987cbb6..2d76aa1a0 100644 --- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean +++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean @@ -1,5 +1,6 @@ import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Backend import VerifiedGarbage.Impl.MlKem.X86_64.Avx +import VerifiedGarbage.Impl.MlDsa.X86_64.Round.Avx2 /-! # ML-DSA on x86-64: the polynomial arithmetic with AVX2 @@ -199,6 +200,7 @@ def subAvx2 : Prog isa := /-- The AVX2 code. -/ def Backend.avx2 : Backend := - ⟨nttAvx2, nttInvAvx2, mulAvx2, mulAddAvx2, addAvx2, subAvx2, Sample.Rej4.rejNTT4Avx2, "_avx2"⟩ + ⟨nttAvx2, nttInvAvx2, mulAvx2, mulAddAvx2, addAvx2, subAvx2, Round.highBitsAvx2, + Round.lowBitsAvx2, Round.normLtAvx2, Round.makeHintAvx2, 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 65d34d28d..b170d109b 100644 --- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean +++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean @@ -1,6 +1,7 @@ import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Ntt import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Mul import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.AddSub +import VerifiedGarbage.Impl.MlDsa.X86_64.Round.Round import VerifiedGarbage.Impl.MlDsa.X86_64.Sample.RejNtt4 /-! @@ -9,8 +10,9 @@ import VerifiedGarbage.Impl.MlDsa.X86_64.Sample.RejNtt4 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`, 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_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. `_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 @@ -29,17 +31,24 @@ structure Backend where mulAdd : Prog isa add : Prog isa sub : Prog isa + highBits : Prog isa + lowBits : Prog isa + normLt : Prog isa + makeHint : Prog isa rej4 : Prog isa /-- What the names of its functions, and of those calling them, end with. -/ sfx : String /-- The SSE2 code. -/ def Backend.sse2 : Backend := - ⟨Arith.ntt, Arith.nttInv, Arith.mul, Arith.mulAdd, Arith.add, Arith.sub, Sample.Rej4.rejNTT4, ""⟩ + ⟨Arith.ntt, Arith.nttInv, Arith.mul, Arith.mulAdd, Arith.add, Arith.sub, Round.highBits, + Round.lowBits, Round.normLt, Round.makeHint, 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 [], ""⟩ +def Backend.empty : Backend := + ⟨.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 new file mode 100644 index 000000000..79292bcbc --- /dev/null +++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Round/Avx2.lean @@ -0,0 +1,158 @@ +import VerifiedGarbage.Impl.MlDsa.X86_64.Round.Round +import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Vec +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 +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 +coefficients of `r` and `out` and `rcx` counting down the 32 vectors. + +`Decompose` is the reference implementation's, as in `Round.lean`: +`f = ⌊(⌊(a + 127)/2⁷⌋ · M + 2^(S-1)) / 2^S⌋` and `r₁ = f mod m`, but in 32 +bits, which hold every intermediate value, with the multiplication by `M` a +sum of shifts (`mulX`: `1025 = 2¹⁰ + 1` and `11275 = 2¹³ + 2¹¹ + 2¹⁰ + 2³ + +2 + 1`), and, as `f ≤ m`, `r₁` is `f` ANDed with the sign of `f - m` +(`psrad` by 31). `r₀ = a - r₁ · 2γ₂` likewise multiplies by shifts +(`2γ₂ = 2¹⁹ - 2⁹` or `2¹⁷ + 2¹⁵ + 2¹⁴ + 2¹³ + 2¹¹`), plus `q` if negative +(`vcadd`). There are no multiplication instructions, and no branch on data: +the functions branch once on the public `γ₂`, and every address depends +only on the pointers. Each clears the upper halves of the vector registers +before returning (`vzeroupper`). +-/ + +namespace VG.Impl.MlDsa.X86_64.Round + +open VG.X86_64 +open VG.Impl.MlKem.X86_64 (xb xmov rcxLoop toY yconst at_) +open VG.Impl.MlDsa.X86_64.Arith (vcadd vcsub) + +/-- The shifts whose sum, with 1, is `M`. -/ +def dSh (g : Nat) : List Nat := if g = 261888 then [10] else [1, 3, 10, 11, 13] + +/-- `xmm0 ← xmm0 · (1 + Σ 2^k)`, through `xmm1` and `xmm2`. -/ +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 ← 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] + +/-- `xmm1 ← xmm0 · 2γ₂`, through `xmm2`. -/ +def mul2X (g : Nat) : List Instr := + if g = 261888 then + [xmov .xmm1 .xmm0, .xop (.shift .pslld .xmm1 19), xmov .xmm2 .xmm0, .xop (.shift .pslld .xmm2 9), + xb .psubd .xmm1 .xmm2] + else + [xmov .xmm1 .xmm0, .xop (.shift .pslld .xmm1 11)] ++ [13, 14, 15, 17].flatMap fun k => + [xmov .xmm2 .xmm0, .xop (.shift .pslld .xmm2 (BitVec.ofNat 8 k)), xb .paddd .xmm1 .xmm2] + +/-- `xmm3 ← r₀` of `xmm0`, with `q` also in `xmm15`. -/ +def lbX (g : Nat) : List Instr := + xmov .xmm3 .xmm0 :: hbX g ++ mul2X g ++ xb .psubd .xmm3 .xmm1 :: vcadd .xmm3 .xmm1 + +/-- The constants of `hbX` and `lbX`. -/ +def yC (g : Nat) : List Instr := + yconst .xmm8 127 ++ yconst .xmm9 (BitVec.ofNat 32 (dAdd g)) ++ yconst .xmm10 (BitVec.ofNat 32 (dMod g)) ++ + yconst .xmm15 8380417 + +/-- Eight coefficients of `r` (at `rdi`) to `out` (at `r10`), through `x`, the result left in `ymm d`. -/ +def bitsBodyY (x : List Instr) (d : XReg) : List Instr := + [.vmovdquLoad .l256 .xmm0 (at_ .rdi 0)] ++ toY x ++ + [.vmovdquStore .l256 (at_ .r10 0) d, .alu .add .rdi (.imm 32), .alu .add .r10 (.imm 32)] + +/-- `γ₂` compared, the output pointer to `r10`, and the loop of `γ₂`, `x g` leaving its result in `ymm d`. -/ +def bitsY (x : Nat → List Instr) (d : XReg) : Prog isa := + .seq (.block (gammaCmp .rsi ++ [.mov .r10 (.reg .rdx)])) + (.seq (.ite .e (.seq (.block (yC g32)) (rcxLoop 32 (bitsBodyY (x g32) d))) + (.seq (.block (yC g88)) (rcxLoop 32 (bitsBodyY (x g88) d)))) (.block [.vop .vzeroupper])) + +def highBitsAvx2 : Prog isa := bitsY hbX .xmm0 + +def lowBitsAvx2 : Prog isa := bitsY lbX .xmm3 + +/-! ## `vg_mldsa_norm_lt_avx2` + +The bound is first clamped to `q` (a branch on the public bound), which +changes no result, as every reduced coefficient is less than `q`; then the +differences `a - b` and `(q - b) - a` fit in 32 bits, and one of them is +negative exactly when `a < b` or `q - a < b`. Each doubleword of `ymm10` +ANDs the ORs of the differences of its coefficients; its top bit stays set +while all of them are good. At the end the top bits are spread over their +doublewords (`vpsrad` by 31), and `vpmovmskb` gathers the top bits of the +32 bytes: the result is `(mask + 1) >> 32`, 1 exactly when every bit is set. -/ + +/-- `ymm10 ← ymm10 & ((a - b) | ((q - b) - a))` for `a` in `xmm0`, `b` in `xmm8`, `q - b` in `xmm9`. -/ +def nlX : List Instr := + [xmov .xmm1 .xmm0, xb .psubd .xmm1 .xmm8, xmov .xmm2 .xmm9, xb .psubd .xmm2 .xmm0, xb .por .xmm1 .xmm2, + xb .pand .xmm10 .xmm1] + +def nlBodyY : List Instr := [.vmovdquLoad .l256 .xmm0 (at_ .rdi 0)] ++ toY nlX ++ [.alu .add .rdi (.imm 32)] + +/-- The low doubleword of `rax` in each doubleword of `ymm r`. -/ +def ybcast (r : XReg) : List Instr := [.vop (.vmovq r .rax), .vop (.vpbroadcastd .l256 r r)] + +/-- `b` in `ymm8`, `q - b` in `ymm9` and all ones in `ymm10`. -/ +def nlConsts : List Instr := + [.mov32 .rax (.reg .rsi)] ++ ybcast .xmm8 ++ [.mov32 .rax (.imm qImm), .alu32 .sub .rax (.reg .rsi)] ++ + ybcast .xmm9 ++ yconst .xmm10 0xFFFFFFFF + +def nlEnd : List Instr := + toY [.xop (.shift .psrad .xmm10 31)] ++ [.vpmovmskb .l256 .rax .xmm10, .vop .vzeroupper] ++ + [.alu .add .rax (.imm 1), .shift .shr .rax 32] + +/-- The bound, clamped to `q`. -/ +def nlPro : Prog isa := + .seq (.block [.mov32 .rsi (.reg .rsi), .alu32 .cmp .rsi (.imm qImm)]) (.ite .b (.block []) (.block [.mov32 .rsi (.imm qImm)])) + +def normLtAvx2 : Prog isa := .seq nlPro (.seq (.block nlConsts) (.seq (rcxLoop 32 nlBodyY) (.block nlEnd))) + +/-! ## `vg_mldsa_make_hint_avx2` + +In each lane, `r₁` of `r` and of `(r + z) mod q` (`hbX` twice, with `vcsub` +between), XORed, is nonzero (`(x + 63) >> 6`, as `x < 64`) exactly when the +hint is 1 (`mhX`); the eight hints are stored. Their count is the sum of the +nibbles of the byte mask of the hints shifted to the top bit of their low +bytes (`vpmovmskb`, which has hint `i` at bit `4i`): three shifts and +additions put it in the low nibble (`cntH`), which is added to `r9`. -/ + +/-- `r₁` of `r` to `xmm3`, and `(r + z) mod q` to `xmm0`, with `r` in `xmm4` and `z` in `xmm5`. -/ +def mhMid : List Instr := [xmov .xmm3 .xmm0, xmov .xmm0 .xmm4, xb .paddd .xmm0 .xmm5] ++ vcsub .xmm0 .xmm1 + +/-- The hint from the two `r₁`, to `xmm0`, and shifted to bit 7, to `xmm1`. -/ +def mhTail : List Instr := + [xb .pxor .xmm0 .xmm3, xb .paddd .xmm0 .xmm11, .xop (.shift .psrld .xmm0 6), xmov .xmm1 .xmm0, + .xop (.shift .pslld .xmm1 7)] + +/-- `xmm0 ← the hints` of `r` in `xmm0` and `z` in `xmm5`, and `xmm1 ← xmm0 << 7`, with `63` in `xmm11`. -/ +def mhX (g : Nat) : List Instr := [xmov .xmm4 .xmm0] ++ hbX g ++ mhMid ++ hbX g ++ mhTail + +/-- `r9 ← r9 +` the sum of the nibbles of `eax`, through `rdx`. -/ +def cntH : List Instr := + [.mov32 .rdx (.reg .rax), .shift32 .shr .rdx 4, .alu32 .add .rax (.reg .rdx), + .mov32 .rdx (.reg .rax), .shift32 .shr .rdx 8, .alu32 .add .rax (.reg .rdx), + .mov32 .rdx (.reg .rax), .shift32 .shr .rdx 16, .alu32 .add .rax (.reg .rdx), + .alu32 .and .rax (.imm 15), .alu .add .r9 (.reg .rax)] + +/-- Eight coefficients of `r` (at `rsi`) and `z` (at `rdi`) to the hints at `r10`, and their count added to `r9`. -/ +def mhBodyY (g : Nat) : List Instr := + [.vmovdquLoad .l256 .xmm0 (at_ .rsi 0), .vmovdquLoad .l256 .xmm5 (at_ .rdi 0)] ++ toY (mhX g) ++ + [.vmovdquStore .l256 (at_ .r10 0) .xmm0, .vpmovmskb .l256 .rax .xmm1] ++ cntH ++ + [.alu .add .rdi (.imm 32), .alu .add .rsi (.imm 32), .alu .add .r10 (.imm 32)] + +def mhY (g : Nat) : Prog isa := .seq (.block (yC g ++ yconst .xmm11 63)) (rcxLoop 32 (mhBodyY g)) + +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])) + +end VG.Impl.MlDsa.X86_64.Round diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean index 4fa71ab43..90f290c61 100644 --- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean +++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean @@ -193,18 +193,18 @@ def ballAt (len tau : Nat) (c : Ptr) : Prog isa := callP "vg_mldsa_sample_in_ball" P.ball [.ptr (sc oCT), .imm len, .imm tau, .ptr c, .ptr (sc oPS)] def highBitsAt (r : Ptr) (gamma2 : Nat) (out : Ptr) : Prog isa := - callP "vg_mldsa_high_bits" P.highBits [.ptr r, .imm gamma2, .ptr out] + callP ("vg_mldsa_high_bits" ++ P.sfx) P.highBits [.ptr r, .imm gamma2, .ptr out] def lowBitsAt (r : Ptr) (gamma2 : Nat) (out : Ptr) : Prog isa := - callP "vg_mldsa_low_bits" P.lowBits [.ptr r, .imm gamma2, .ptr out] + callP ("vg_mldsa_low_bits" ++ P.sfx) P.lowBits [.ptr r, .imm gamma2, .ptr out] /-- `‖f‖∞ < bound`, and `r15 ← r15 ∧ result`. -/ def normAt (f : Ptr) (bound : Nat) : Prog isa := - .seq (callP "vg_mldsa_norm_lt" P.normLt [.ptr f, .imm bound]) (.block [.alu32 .and .r15 (.reg .rax)]) + .seq (callP ("vg_mldsa_norm_lt" ++ P.sfx) P.normLt [.ptr f, .imm bound]) (.block [.alu32 .and .r15 (.reg .rax)]) /-- `MakeHint` of `z` and `r` to `h`, and the number of 1s added to `ONES`. -/ def makeHintAt (z r : Ptr) (gamma2 : Nat) (h : Ptr) : Prog isa := - .seq (callP "vg_mldsa_make_hint" P.makeHint [.ptr z, .ptr r, .imm gamma2, .ptr h]) + .seq (callP ("vg_mldsa_make_hint" ++ P.sfx) P.makeHint [.ptr z, .ptr r, .imm gamma2, .ptr h]) (.block [.mov32 .rcx (.mem (at_ .rbx oONES)), .alu32 .add .rcx (.reg .rax), .store (at_ .rbx oONES) .rcx]) def simpleBitPackAt (f : Ptr) (b : Nat) (out : Ptr) (len : Nat) : Prog isa := diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean index 72ef6c175..b156b4e60 100644 --- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean +++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean @@ -189,7 +189,7 @@ def hintUnpackAt (y : Ptr) (len omega : Nat) (h : Ptr) (hlen : Nat) : Prog isa : [(.rdi, .ptr y), (.rsi, .imm len), (.rdx, .imm omega), (.rcx, .ptr h), (.r8, .imm hlen)] def normLtAt (f : Ptr) (bound : Nat) : Prog isa := - callAt "vg_mldsa_norm_lt" P.normLt [(.rdi, .ptr f), (.rsi, .imm bound)] + callAt ("vg_mldsa_norm_lt" ++ P.sfx) P.normLt [(.rdi, .ptr f), (.rsi, .imm bound)] end diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean index d8ad5b2bc..759bdcb13 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean @@ -4,6 +4,9 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.NttInv import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Mul import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.AddSub import VerifiedGarbage.Proof.MlKem.X86_64.ArithOk +import VerifiedGarbage.Proof.MlDsa.X86_64.Round.Bits +import VerifiedGarbage.Proof.MlDsa.X86_64.Round.MakeHint +import VerifiedGarbage.Proof.MlDsa.X86_64.Round.NormLt import VerifiedGarbage.Proof.MlDsa.X86_64.Sample.Rej4Verified /-! @@ -49,6 +52,10 @@ structure BackendOk (B : Backend) : Prop where mulAdd : FnOk (fun S => Spec.MlDsa.mulAddContract X86_64.abi S) B.mulAdd add : FnOk (fun S => Spec.MlDsa.addContract X86_64.abi S) B.add sub : FnOk (fun S => Spec.MlDsa.subContract X86_64.abi S) B.sub + highBits : FnOk (fun S => Spec.MlDsa.highBitsContract X86_64.abi S) B.highBits + 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 rej4 : Rej4Ok B.rej4 /-- An implementation of the polynomial arithmetic on x86-64. -/ @@ -80,6 +87,14 @@ def ArithImpl.sse2 : ArithImpl where (by decide +kernel) sub := FnOk.of Arith.sub_verified (by decide +kernel) (by decide +kernel) (by decide +kernel) (by decide +kernel) + highBits := FnOk.of Round.highBits_verified (by decide +kernel) (by decide +kernel) (by decide +kernel) + (by decide +kernel) + lowBits := FnOk.of Round.lowBits_verified (by decide +kernel) (by decide +kernel) (by decide +kernel) + (by decide +kernel) + normLt := FnOk.of Round.normLt_verified (by decide +kernel) (by decide +kernel) (by decide +kernel) + (by decide +kernel) + makeHint := FnOk.of Round.makeHint_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)⟩ } features := [] diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/BackendAvx2.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/BackendAvx2.lean index 0036f4cdb..f91cfc25f 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/BackendAvx2.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/BackendAvx2.lean @@ -2,6 +2,7 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Backend 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 /-! # ML-DSA on x86-64: the polynomial arithmetic with AVX2, as an `ArithImpl` @@ -31,6 +32,14 @@ def ArithImpl.avx2 : ArithImpl where (by decide +kernel) sub := FnOk.of Arith.subY_verified (by decide +kernel) (by decide +kernel) (by decide +kernel) (by decide +kernel) + highBits := FnOk.of Round.highBitsY_verified (by decide +kernel) (by decide +kernel) (by decide +kernel) + (by decide +kernel) + lowBits := FnOk.of Round.lowBitsY_verified (by decide +kernel) (by decide +kernel) (by decide +kernel) + (by decide +kernel) + normLt := FnOk.of Round.normLtY_verified (by decide +kernel) (by decide +kernel) (by decide +kernel) + (by decide +kernel) + makeHint := FnOk.of Round.makeHintY_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)⟩ } features := ["avx", "avx2"] diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBits.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBits.lean new file mode 100644 index 000000000..10064a572 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBits.lean @@ -0,0 +1,273 @@ +import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YBlock +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.YBase +import VerifiedGarbage.Proof.MlDsa.X86_64.Round.Bits + +/-! +# ML-DSA on x86-64: `vg_mldsa_high_bits_avx2` and `vg_mldsa_low_bits_avx2` + +Untrusted: everything here is checked by Lean. After the constants +(`yC_ok`), each iteration of the loop loads eight coefficients of `r`, does +in each lane what `hbX` or `lbX` does to a register (`YBlock.lean`), whose +proof holds of each lane (`ylanes`), and stores the eight results to `out` +(`YMap.step`); the loop leaves `out` with all 256 (`YMap.loop_ok`), for the +`γ₂` the function compared (`ybits_ok`). +-/ + +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 bc_ofDwords) +open VG.Proof.MlKem.X86_64 (Keep XOnly YOnly ylanes yld_ok yconst_ok WP.keep ifp ifn ptr_step GOnly wp_rcxLoopY + add_ofNat_zero lane_setReg lane_setFlags sx32 State.setMem_ymm) +open VG.Impl.MlKem.X86_64 (xb xmov toY yconst) +open VG.Spec.MlDsa (q n coeffAt) + +/-- Each doubleword of `v`. -/ +def bc (v : BitVec 32) : BitVec 128 := ofDwords v v v v + +/-- The constants of `hbX` and `lbX` in both lanes. -/ +def YC (g : Nat) (s : State) : Prop := + ∀ l < 2, s.lane .xmm8 l = bc 127 ∧ s.lane .xmm9 l = bc (BitVec.ofNat 32 (dAdd g)) ∧ + s.lane .xmm10 l = bc (BitVec.ofNat 32 (dMod g)) ∧ s.lane .xmm15 l = qV + +theorem YC.hbc {g : Nat} {s : State} (h : YC g s) {l : Nat} (hl : l < 2) : HbC g (s.proj l) := + ⟨by rw [State.proj_xmm, (h l hl).1]; exact bc_ofDwords _, by rw [State.proj_xmm, (h l hl).2.1]; exact bc_ofDwords _, + by rw [State.proj_xmm, (h l hl).2.2.1]; exact bc_ofDwords _⟩ + +theorem YC.q {g : Nat} {s : State} (h : YC g s) {l : Nat} (hl : l < 2) : (s.proj l).xmm .xmm15 = qV := by + rw [State.proj_xmm, (h l hl).2.2.2] + +/-- The constants are kept by code that writes only `xmm0` to `xmm3`. -/ +theorem YC.keep {g : Nat} {s s' : State} (h : YC g s) {rs : List XReg} (hr : ∀ r ∈ rs, r ∈ [.xmm0, .xmm1, .xmm2, .xmm3]) + (hl : ∀ r ∉ rs, ∀ l < 2, s'.lane r l = s.lane r l) : YC g s' := fun l hl' => by + have n : ∀ r ∈ [XReg.xmm8, .xmm9, .xmm10, .xmm15], r ∉ rs := fun r hr' hr'' => by + have := hr r hr''; simp only [List.mem_cons, List.not_mem_nil, or_false] at hr' this + rcases hr' with rfl | rfl | rfl | rfl <;> rcases this with h | h | h | h <;> cases h + rw [hl _ (n _ (by simp)) l hl', hl _ (n _ (by simp)) l hl', hl _ (n _ (by simp)) l hl', hl _ (n _ (by simp)) l hl'] + exact h l hl' + +theorem yC_ok (g : Nat) (s : State) : + WP isa (.block (yC g)) s fun s' => YC g s' ∧ Keep [.rax] s s' ∧ s'.mem = s.mem ∧ s'.mxcsr = s.mxcsr ∧ + ∀ r ∉ [XReg.xmm8, .xmm9, .xmm10, .xmm15], ∀ l < 2, s'.lane r l = s.lane r l := by + simp only [yC, List.append_assoc] + rw [WP.block_append_iff] + refine WP.mono (yconst_ok .xmm8 _ s) fun s1 ⟨l1, k1, m1, x1, o1⟩ => ?_ + rw [WP.block_append_iff] + refine WP.mono (yconst_ok .xmm9 _ s1) fun s2 ⟨l2, k2, m2, x2, o2⟩ => ?_ + rw [WP.block_append_iff] + refine WP.mono (yconst_ok .xmm10 _ s2) fun s3 ⟨l3, k3, m3, x3, o3⟩ => ?_ + refine WP.mono (yconst_ok .xmm15 _ s3) fun s4 ⟨l4, k4, m4, x4, o4⟩ => + ⟨fun l hl => ⟨?_, ?_, ?_, ?_⟩, (((k1.trans k2).trans k3).trans k4).mono (by simp), m4.trans (m3.trans (m2.trans m1)), + x4.trans (x3.trans (x2.trans x1)), fun r hr l hl => ?_⟩ + · rw [o4 _ (by decide) l hl, o3 _ (by decide) l hl, o2 _ (by decide) l hl, l1 l hl]; rfl + · rw [o4 _ (by decide) l hl, o3 _ (by decide) l hl, l2 l hl]; rfl + · rw [o4 _ (by decide) l hl, l3 l hl]; rfl + · rw [l4 l hl]; decide + · simp only [List.mem_cons, List.not_mem_nil, or_false, not_or] at hr + rw [o4 r hr.2.2.2 l hl, o3 r hr.2.2.1 l hl, o2 r hr.2.1 l hl, o1 r hr.1 l hl] + +namespace YMap + +/-- After `i` vectors of eight, each coefficient of `out` before `8i` is `v k`. -/ +structure Inv (s₀ : State) (g : Nat) (v : Nat → BitVec 32) (i : Nat) (s : State) : Prop where + rdi : s.gpr .rdi = s₀.gpr .rdi + BitVec.ofNat 64 (32 * i) + r10 : s.gpr .r10 = s₀.gpr .rdx + BitVec.ofNat 64 (32 * i) + rd : s.rd = s₀.rd + wr : s.wr = s₀.wr + yc : YC g s + frame : Frame [pR (s₀.gpr .rdx)] s₀.mem s.mem + coeff : ∀ k < 256, coeffAt s.mem (s₀.gpr .rdx) k = if k < 8 * i then v k else coeffAt s₀.mem (s₀.gpr .rdx) k + +section +variable {s₀ : State} (hrd : s₀.rd = [pR (s₀.gpr .rdi)]) (hwr : s₀.wr = [pR (s₀.gpr .rdx)]) + (hdis : (pR (s₀.gpr .rdi)).Disjoint (pR (s₀.gpr .rdx))) + {g : Nat} {x : List Instr} {d : XReg} {L : BitVec 32 → BitVec 32} + (hX : ∀ t : State, HbC g t → t.xmm .xmm15 = qV → + WP isa (.block x) t fun t' => (∀ e < 4, dword (t'.xmm d) e = L (dword (t.xmm .xmm0) e)) ∧ + XOnly [.xmm0, .xmm1, .xmm2, .xmm3] t t') + (hY : laneSseBlock (toY x) = some x) +include hrd hwr hdis hX hY + +/-- An iteration, which stores `L` of the eight coefficients of `r` to `out`. -/ +theorem step {i : Nat} (hi : i < 32) {s : State} + (hI : Inv s₀ g (fun k => L (coeffAt s₀.mem (s₀.gpr .rdi) k)) i s) : + WP isa (.block (bitsBodyY x d ++ ([.alu .sub .rcx (.imm 1)] : List Instr))) s fun s' => + Inv s₀ g (fun k => L (coeffAt s₀.mem (s₀.gpr .rdi) k)) (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 .rdx) ∈ s.wr := by rw [hI.wr, hwr]; simp + have hr : pR (s₀.gpr .rdi) ∈ s.rd ++ s.wr := by rw [hI.rd, hrd]; simp + have e1 : s.gpr .rdi + BitVec.ofNat 64 0 = coeffAddr (s₀.gpr .rdi) (8 * i) := by + rw [add_ofNat_zero, hI.rdi]; congr 2; omega + have e2 : s.gpr .r10 = coeffAddr (s₀.gpr .rdx) (8 * i) := by + rw [hI.r10]; congr 2; omega + rw [bitsBodyY, List.append_assoc, List.append_assoc, WP.block_append_iff] + refine WP.mono (yld_ok (by rw [e1]; exact f_in32 hr j0)) fun s1 ⟨L1, o1⟩ => ?_ + have yc1 : YC g s1 := hI.yc.keep (by simp) o1.lane + rw [WP.block_append_iff] + refine WP.mono (ylanes hY (P := fun l t => ∀ e < 4, dword (t.xmm d) e = L (dword ((s1.proj l).xmm .xmm0) e)) + fun l hl => hX _ (yc1.hbc hl) (yc1.q hl)) fun s3 ⟨B3, o3⟩ => ?_ + have o13 := o1.trans o3 + have yc3 : YC g s3 := yc1.keep (by simp) o3.lane + have g3 : s3.gpr .r10 = coeffAddr (s₀.gpr .rdx) (8 * i) := by rw [o13.gpr, e2] + have g3' : s3.gpr .rdi = coeffAddr (s₀.gpr .rdi) (8 * i) := by rw [o13.gpr, ← e1, add_ofNat_zero] + have w0 : InRegions s3.wr (s3.gpr .r10) 32 := by + rw [o13.wr, g3]; exact f_in32 hw j0 + 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.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 yc3.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 d l) e = L (coeffAt s₀.mem (s₀.gpr .rdi) (8 * i + 4 * l + e)) := + fun l hl e he => by + rw [← State.proj_xmm, B3 l hl e he, State.proj_xmm, L1 l hl, e1, dword_readW _ _ he, lane_load, + coeffAddr_add, ← coeffAt_eq, coeffAt_frame hI.frame (by simpa using hdis) (by rw [n_eq]; omega)] + 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. -/ +theorem loop_ok {s : State} (hs0 : s.gpr .rdi = s₀.gpr .rdi) (hs10 : s.gpr .r10 = s₀.gpr .rdx) + (hsrd : s.rd = s₀.rd) (hswr : s.wr = s₀.wr) (hsm : s.mem = s₀.mem) : + WP isa (.seq (.block (yC g)) (VG.Impl.MlKem.X86_64.rcxLoop 32 (bitsBodyY x d))) s fun s' => + Frame [pR (s₀.gpr .rdx)] s₀.mem s'.mem ∧ + ∀ k < 256, coeffAt s'.mem (s₀.gpr .rdx) k = L (coeffAt s₀.mem (s₀.gpr .rdi) 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, 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 hdis hX hY hi hI) fun u hI => ⟨hI.frame, fun k hk => by + rw [hI.coeff k hk, ifp (by omega)]⟩ + +end + +end YMap + +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 XOnly) +open VG.Impl.MlKem.X86_64 (toY) +open VG.Proof.MlDsa.X86_64.Arith (qV) + +section +variable {post : State → State → Prop} {s₀ : State} (hp : (bitsK post).pre s₀) +include hp + +/-- The prologue, the loop of the `γ₂` it compared, and the epilogue: `out` holds `L γ₂` of each +coefficient of `r`. -/ +theorem ybits_ok {x : Nat → List Instr} {d : XReg} {L : Nat → BitVec 32 → BitVec 32} + (hX : ∀ g, (g = g32 ∨ g = g88) → ∀ t : State, HbC g t → t.xmm .xmm15 = qV → + WP isa (.block (x g)) t fun t' => (∀ e < 4, dword (t'.xmm d) e = L g (dword (t.xmm .xmm0) e)) ∧ + XOnly [.xmm0, .xmm1, .xmm2, .xmm3] t t') + (hY : ∀ g, (g = g32 ∨ g = g88) → laneSseBlock (toY (x g)) = some (x g)) : + WP isa (bitsY x d) s₀ fun s' => + (∀ k < 256, coeffAt s'.mem (s₀.gpr .rdx) k = L (arg32 s₀ .rsi) (coeffAt s₀.mem (s₀.gpr .rdi) k)) ∧ + Frame [pR (s₀.gpr .rdx)] s₀.mem s'.mem := by + unfold bitsY + refine WP.seq (WP.mono (prologue_rsi_rdx s₀) fun s₁ ⟨⟨h10, hz, hm⟩, hk⟩ => ?_) + have h0 : s₁.gpr .rdi = s₀.gpr .rdi := hk.gpr (by decide) + have go : ∀ g, arg32 s₀ .rsi = g → (g = g32 ∨ g = g88) → + WP isa (.seq (.block (yC g)) (VG.Impl.MlKem.X86_64.rcxLoop 32 (bitsBodyY (x g) d))) s₁ fun s' => + (∀ k < 256, coeffAt s'.mem (s₀.gpr .rdx) k = L (arg32 s₀ .rsi) (coeffAt s₀.mem (s₀.gpr .rdi) k)) ∧ + Frame [pR (s₀.gpr .rdx)] s₀.mem s'.mem := fun g hge hg => + WP.mono (Arith.YMap.loop_ok hp.1 hp.2.1 hp.2.2.1 (hX g hg) (hY g hg) h0 h10 hk.2.1 hk.2.2 hm) + fun _ ⟨hf, hc⟩ => ⟨fun k hk => by rw [hc k hk, hge], hf⟩ + have hite : WP isa (.ite .e (.seq (.block (yC g32)) (VG.Impl.MlKem.X86_64.rcxLoop 32 (bitsBodyY (x g32) d))) + (.seq (.block (yC g88)) (VG.Impl.MlKem.X86_64.rcxLoop 32 (bitsBodyY (x g88) d)))) s₁ fun s' => + (∀ k < 256, coeffAt s'.mem (s₀.gpr .rdx) k = L (arg32 s₀ .rsi) (coeffAt s₀.mem (s₀.gpr .rdi) k)) ∧ + Frame [pR (s₀.gpr .rdx)] s₀.mem s'.mem := by + refine WP.ite (M := isa) _ (show isa.eval .e s₁ = _ from hz) (fun h => ?_) (fun h => ?_) + · rw [sub_beq_zero32, decide_eq_true_eq] at h + exact go _ ((gamma_cases hp.2.2.2.2.2.1).1 h) (.inl rfl) + · rw [sub_beq_zero32, decide_eq_false_iff_not] at h + exact go _ ((gamma_cases hp.2.2.2.2.2.1).2 h) (.inr rfl) + refine WP.seq (WP.mono hite fun u ⟨hc, hf⟩ => ?_) + refine WP.mono (Q := fun (u' : State) => u'.mem = u.mem) (by vrund; rfl) + fun u' hm' => ?_ + rw [hm'] + exact ⟨hc, hf⟩ + +end + +theorem hbX_okG (g : Nat) (hg : g = g32 ∨ g = g88) (t : State) (hc : HbC g t) (_ : t.xmm .xmm15 = qV) : + WP isa (.block (hbX g)) t fun t' => (∀ e < 4, dword (t'.xmm .xmm0) e = hbL g (dword (t.xmm .xmm0) e)) ∧ + XOnly [.xmm0, .xmm1, .xmm2, .xmm3] t t' := + WP.mono (hbX_ok hg t hc) fun _ ⟨h1, h2⟩ => ⟨h1, h2.mono (by simp)⟩ + +theorem lane_hbX : ∀ g, (g = g32 ∨ g = g88) → laneSseBlock (toY (hbX g)) = some (hbX g) := by + intro g hg; rcases hg with rfl | rfl <;> decide +kernel + +theorem lane_lbX : ∀ g, (g = g32 ∨ g = g88) → laneSseBlock (toY (lbX g)) = some (lbX g) := by + intro g hg; rcases hg with rfl | rfl <;> decide +kernel + +theorem highBitsY_correct (s₀ : State) (hp : highBitsK.pre s₀) : + ∃ t s', Exec isa highBitsAvx2 s₀ t s' ∧ abiPreserved s₀ s' ∧ highBitsK.post s₀ s' := by + obtain ⟨t, s', he, ⟨hv, hf⟩, hk⟩ := WP.keep [.rax, .rcx, .rsi, .rdi, .r10] + (ybits_ok hp (L := hbL) hbX_okG lane_hbX) (by decide +kernel) + have hr : Reduced s₀.mem (s₀.gpr .rdi) := hp.2.2.2.2.2.2 + refine ⟨t, s', he, abiPreserved_of_exec (by decide +kernel) he (gprPreserved_of hk (by decide) hf ?_), ?_⟩ + · simpa using hp.2.2.2.2.1 + · refine natPolyIs_of_toNat fun k hk => ?_ + rw [hv k hk, map_get _ _ hk, hbL_toNat hp.2.2.2.2.2.1 (hr k hk), highBits_eq hp.2.2.2.2.2.1, polyAt_val hr hk] + exact (Int.toNat_natCast _).symm + +theorem lowBitsY_correct (s₀ : State) (hp : lowBitsK.pre s₀) : + ∃ t s', Exec isa lowBitsAvx2 s₀ t s' ∧ abiPreserved s₀ s' ∧ lowBitsK.post s₀ s' := by + obtain ⟨t, s', he, ⟨hv, hf⟩, hk⟩ := WP.keep [.rax, .rcx, .rsi, .rdi, .r10] + (ybits_ok hp (L := lbL) (fun g hg t hc hq => lbX_ok hg t hc hq) lane_lbX) (by decide +kernel) + have hr : Reduced s₀.mem (s₀.gpr .rdi) := hp.2.2.2.2.2.2 + refine ⟨t, s', he, abiPreserved_of_exec (by decide +kernel) he (gprPreserved_of hk (by decide) hf ?_), ?_⟩ + · simpa using hp.2.2.2.2.1 + · refine polyIs_of_toNat fun k hk => ?_ + rw [hv k hk, map_get _ _ hk, lbL_toNat hp.2.2.2.2.2.1 (hr k hk), polyAt_get _ _ hk] + +theorem highBitsY_ct : ConstantTime isa highBitsK.pre highBitsK.pub highBitsAvx2 := + VG.Taint.constantTime (A := X86_64.taint) bitsτ bits_agree (by taint_decide) + +theorem lowBitsY_ct : ConstantTime isa lowBitsK.pre lowBitsK.pub lowBitsAvx2 := + VG.Taint.constantTime (A := X86_64.taint) bitsτ bits_agree (by taint_decide) + +theorem highBitsY_verified : Verified X86_64.target highBitsAvx2 (highBitsContract X86_64.abi) := + Verified.of_correct highBitsY_correct highBitsY_ct (by + round_implies [highBitsContract, bitsSig, highBitsK, bitsK, X86_64.abi, X86_64.argRegs] [bitsSat] + using bitsSat) + +theorem lowBitsY_verified : Verified X86_64.target lowBitsAvx2 (lowBitsContract X86_64.abi) := + Verified.of_correct lowBitsY_correct lowBitsY_ct (by + round_implies [lowBitsContract, bitsSig, lowBitsK, bitsK, X86_64.abi, X86_64.argRegs] [bitsSat] + using bitsSat) + +end VG.Proof.MlDsa.X86_64.Round diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBlock.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBlock.lean new file mode 100644 index 000000000..bc93dce0c --- /dev/null +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBlock.lean @@ -0,0 +1,69 @@ +import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YLane +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.VLanes + +/-! +# ML-DSA on x86-64: the SSE2 code of the AVX2 rounding, on a register + +Untrusted: everything here is checked by Lean. `hbX` and `lbX` compute +`hbL` and `lbL` in each doubleword of `xmm0` (`hbX_ok`, `lbX_ok`), with +the constants in `xmm8`, `xmm9`, `xmm10` (and `q` in `xmm15`). +-/ + +namespace VG.Proof.MlDsa.X86_64.Round + +open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round +open VG.Proof.MlDsa.X86_64.Arith (dword_psubd dword_pand dword_psrad caddL caddV dword_caddV qV + dword_qV) +open VG.Proof.MlKem.X86_64 (XOnly) +open VG.Impl.MlKem.X86_64 (xb xmov) +open VG.Spec.MlDsa (gamma2s) + +/-- Each doubleword of `x` is `v`. -/ +def Bc (x : BitVec 128) (v : BitVec 32) : Prop := ∀ e < 4, dword x e = v + +theorem bc_ofDwords (v : BitVec 32) : Bc (ofDwords v v v v) v := fun e he => by + rcases cases4 he with rfl | rfl | rfl | rfl <;> simp + +/-- The constants of `hbX`. -/ +structure HbC (g : Nat) (s : State) : Prop where + c8 : Bc (s.xmm .xmm8) (BitVec.ofNat 32 127) + c9 : Bc (s.xmm .xmm9) (BitVec.ofNat 32 (dAdd g)) + c10 : Bc (s.xmm .xmm10) (BitVec.ofNat 32 (dMod g)) + +theorem dSh_32 : dSh g32 = [10] := rfl +theorem dSh_88 : dSh g88 = [1, 3, 10, 11, 13] := rfl +theorem dShift_32 : dShift g32 = 22 := rfl +theorem dShift_88 : dShift g88 = 24 := rfl + +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, + 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_psrad, dword_psubd, dword_psrld, + dword_pslld, dword_paddd, hc.c8 e he, hc.c9 e he, hc.c10 e he, BitVec.toNat_ofNat] + rfl + +theorem mul2X_32 : mul2X g32 = [xmov .xmm1 .xmm0, .xop (.shift .pslld .xmm1 19), xmov .xmm2 .xmm0, + .xop (.shift .pslld .xmm2 9), xb .psubd .xmm1 .xmm2] := rfl + +theorem mul2X_88 : mul2X g88 = [xmov .xmm1 .xmm0, .xop (.shift .pslld .xmm1 11)] ++ [13, 14, 15, 17].flatMap fun k => + [xmov .xmm2 .xmm0, .xop (.shift .pslld .xmm2 (BitVec.ofNat 8 k)), xb .paddd .xmm1 .xmm2] := rfl + +theorem lbX_ok {g : Nat} (hg : g = g32 ∨ g = g88) (s : State) (hc : HbC g s) (hq : s.xmm .xmm15 = qV) : + 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, + 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] + refine ⟨fun e he => ?_, by xonly⟩ + simp (disch := first | decide | assumption) only [dword_pand, 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, hq, dword_qV he] + rfl + +end VG.Proof.MlDsa.X86_64.Round diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YHint.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YHint.lean new file mode 100644 index 000000000..9a5d3d831 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YHint.lean @@ -0,0 +1,508 @@ +import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YNorm +import VerifiedGarbage.Proof.MlDsa.X86_64.Round.MakeHint + +/-! +# ML-DSA on x86-64: `vg_mldsa_make_hint_avx2` + +Untrusted: everything here is checked by Lean. In each doubleword, `mhX` +computes the hint bit (`mhL`, `mhL_toNat`); the loop stores the eight +hints of an iteration and adds their count, the sum of the nibbles of the +byte mask of the hints shifted to bit 7 (`nib_count`), to `r9`. +-/ + +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 XOnly YOnly ylanes ifp ifn WP.keep) +open VG.Impl.MlKem.X86_64 (xb xmov toY) +open VG.Proof.MlDsa.X86_64.Arith (qV csubL csubL_toNat dword_csubV bc) + +/-! ## A doubleword -/ + +/-- The hint bit `mhX` computes from `z` and `r`. -/ +def mhL (g : Nat) (z r : BitVec 32) : BitVec 32 := + ((hbL g (csubL (r + z)) ^^^ hbL g r) + BitVec.ofNat 32 63) >>> 6 + +theorem mhL_toNat {g : Nat} (h : g ∈ gamma2s) {z r : BitVec 32} (hz : z.toNat < q) (hr : r.toNat < q) : + (mhL g z r).toNat = (makeHint g (Fin.ofNat q z.toNat) (Fin.ofNat q r.toNat)).toNat := by + have e1 : (r + z).toNat = r.toNat + z.toNat := by + rw [BitVec.toNat_add]; rw [q_eq] at hz hr; omega + have e2 : (csubL (r + z)).toNat = (r.toNat + z.toNat) % q := by + rw [csubL_toNat (by rw [e1]; omega), e1, VG.Proof.MlDsa.Arith.condSub] + split <;> rw [q_eq] at * <;> omega + have hs : (csubL (r + z)).toNat < q := by rw [e2]; exact Nat.mod_lt _ (by decide) + have hM : hbM g ≤ 44 ∧ 0 < hbM g := by rcases mem_gamma2s h with rfl | rfl <;> decide + have hu := Nat.mod_lt (hbF g ((r.toNat + z.toNat) % q)) hM.2 + have hv := Nat.mod_lt (hbF g r.toNat) hM.2 + have hx : hbF g ((r.toNat + z.toNat) % q) % hbM g ^^^ hbF g r.toNat % hbM g < 2 ^ 6 := + Nat.xor_lt_two_pow (by omega) (by omega) + have hvz : (Fin.ofNat q z.toNat).val = z.toNat := Nat.mod_eq_of_lt hz + have hvr : (Fin.ofNat q r.toNat).val = r.toNat := Nat.mod_eq_of_lt hr + rw [makeHint_eq h, hvz, hvr, mhL, BitVec.toNat_ushiftRight, BitVec.toNat_add, BitVec.toNat_xor, + hbL_toNat h hs, hbL_toNat h hr, e2, BitVec.toNat_ofNat, Nat.shiftRight_eq_div_pow] + by_cases e : hbF g r.toNat % hbM g = hbF g ((r.toNat + z.toNat) % q) % hbM g + · rw [decide_eq_false (fun h' => h' e), e, Nat.xor_self]; rfl + · rw [decide_eq_true e] + have : hbF g ((r.toNat + z.toNat) % q) % hbM g ^^^ hbF g r.toNat % hbM g ≠ 0 := + fun h' => e (xor_eq_zero h').symm + show _ = 1 + omega + +theorem mhL_le {g : Nat} (h : g ∈ gamma2s) {z r : BitVec 32} (hz : z.toNat < q) (hr : r.toNat < q) : + (mhL g z r).toNat ≤ 1 := by + rw [mhL_toNat h hz hr]; exact Bool.toNat_le _ + +/-! ## The code on a register -/ + +theorem mhMid_ok (s : State) (hq : s.xmm .xmm15 = qV) : + WP isa (.block mhMid) s fun s' => s'.xmm .xmm3 = s.xmm .xmm0 ∧ + (∀ e < 4, dword (s'.xmm .xmm0) e = csubL (dword (s.xmm .xmm4) e + dword (s.xmm .xmm5) e)) ∧ + XOnly [.xmm0, .xmm1, .xmm3] s s' := by + simp only [mhMid, VG.Impl.MlDsa.X86_64.Arith.vcsub, VG.Impl.MlDsa.X86_64.Arith.vcadd, xmov, xb, List.cons_append, + List.nil_append] + vrun [VG.X86_64.eval_movdqa] + refine ⟨trivial, fun e he => ?_, by xonly⟩ + rw [hq, ← dword_paddd _ _ he, ← dword_csubV _ he] + rfl + +theorem mhTail_ok (s : State) (h11 : Bc (s.xmm .xmm11) (BitVec.ofNat 32 63)) : + WP isa (.block mhTail) s fun s' => (∀ e < 4, + dword (s'.xmm .xmm0) e = ((dword (s.xmm .xmm0) e ^^^ dword (s.xmm .xmm3) e) + BitVec.ofNat 32 63) >>> 6 ∧ + dword (s'.xmm .xmm1) e = (((dword (s.xmm .xmm0) e ^^^ dword (s.xmm .xmm3) e) + BitVec.ofNat 32 63) >>> 6) <<< 7) ∧ + XOnly [.xmm0, .xmm1] s s' := by + simp only [mhTail, xmov, xb] + vrun [VG.X86_64.eval_movdqa] + refine ⟨fun e he => ?_, by xonly⟩ + simp (disch := first | decide | assumption) only [dword_pxor, dword_paddd, dword_psrld, dword_pslld, h11 e he] + exact ⟨rfl, rfl⟩ + +theorem mhX_ok {g : Nat} (hg : g = g32 ∨ g = g88) (s : State) (hc : HbC g s) (hq : s.xmm .xmm15 = qV) + (h11 : Bc (s.xmm .xmm11) (BitVec.ofNat 32 63)) : + WP isa (.block (mhX g)) s fun s' => (∀ e < 4, + dword (s'.xmm .xmm0) e = mhL g (dword (s.xmm .xmm5) e) (dword (s.xmm .xmm0) e) ∧ + dword (s'.xmm .xmm1) e = mhL g (dword (s.xmm .xmm5) e) (dword (s.xmm .xmm0) e) <<< 7) ∧ + XOnly [.xmm0, .xmm1, .xmm2, .xmm3, .xmm4] s s' := by + rw [mhX, List.append_assoc, List.append_assoc, List.append_assoc, WP.block_append_iff] + refine WP.mono (Q := fun (s1 : State) => s1.xmm .xmm4 = s.xmm .xmm0 ∧ XOnly [.xmm4] s s1) + (by simp only [xmov, xb]; vrun [VG.X86_64.eval_movdqa]; exact ⟨trivial, by xonly⟩) fun s1 ⟨a1, o1⟩ => ?_ + have hc1 : HbC g s1 := ⟨by rw [o1.xmm _ (by decide)]; exact hc.c8, by rw [o1.xmm _ (by decide)]; exact hc.c9, + by rw [o1.xmm _ (by decide)]; exact hc.c10⟩ + rw [WP.block_append_iff] + refine WP.mono (hbX_ok hg s1 hc1) fun s2 ⟨a2, o2⟩ => ?_ + rw [WP.block_append_iff] + refine WP.mono (mhMid_ok s2 (by rw [o2.xmm _ (by decide), o1.xmm _ (by decide), hq])) fun s3 ⟨b3, a3, o3⟩ => ?_ + have hc3 : HbC g s3 := ⟨by rw [o3.xmm _ (by decide), o2.xmm _ (by decide)]; exact hc1.c8, + by rw [o3.xmm _ (by decide), o2.xmm _ (by decide)]; exact hc1.c9, + by rw [o3.xmm _ (by decide), o2.xmm _ (by decide)]; exact hc1.c10⟩ + rw [WP.block_append_iff] + refine WP.mono (hbX_ok hg s3 hc3) fun s4 ⟨a4, o4⟩ => ?_ + refine WP.mono (mhTail_ok s4 (by + rw [o4.xmm _ (by decide), o3.xmm _ (by decide), o2.xmm _ (by decide), o1.xmm _ (by decide)]; exact h11)) + fun s5 ⟨a5, o5⟩ => ⟨fun e he => ?_, ?_⟩ + · have hz : dword (s4.xmm .xmm0) e = hbL g (csubL (dword (s.xmm .xmm0) e + dword (s.xmm .xmm5) e)) := by + rw [a4 e he, a3 e he, o2.xmm _ (by decide), a1, o2.xmm _ (by decide), o1.xmm _ (by decide)] + have hr1 : dword (s4.xmm .xmm3) e = hbL g (dword (s.xmm .xmm0) e) := by + rw [o4.xmm _ (by decide), b3, a2 e he, o1.xmm _ (by decide)] + rw [(a5 e he).1, (a5 e he).2, hz, hr1] + exact ⟨rfl, rfl⟩ + · exact ((((o1.trans o2).trans o3).trans o4).trans o5).mono (by simp) + +/-! ## The count -/ + +/-- What `cntH` leaves in `eax`: the sum of the nibbles of `x`, if it fits in one. -/ +def nib (x : BitVec 32) : BitVec 32 := + let x1 := x + x >>> 4 + let x2 := x1 + x1 >>> 8 + (x2 + x2 >>> 16) &&& 15 + +/-- The eight bits `b d` at bits `4d`. -/ +def sp (b : Nat → Bool) : Nat → Nat + | 0 => 0 + | k + 1 => sp b k + (b k).toNat * 16 ^ k + +/-- The number of the first `k` bits set. -/ +def cnt8 (b : Nat → Bool) : Nat → Nat + | 0 => 0 + | k + 1 => cnt8 b k + (b k).toNat + +theorem nib_bools : ∀ b0 b1 b2 b3 b4 b5 b6 b7 : Bool, + (nib (BitVec.ofNat 32 (b0.toNat + b1.toNat * 16 + b2.toNat * 16 ^ 2 + b3.toNat * 16 ^ 3 + b4.toNat * 16 ^ 4 + + b5.toNat * 16 ^ 5 + b6.toNat * 16 ^ 6 + b7.toNat * 16 ^ 7))).toNat = + b0.toNat + b1.toNat + b2.toNat + b3.toNat + b4.toNat + b5.toNat + b6.toNat + b7.toNat := by + decide +kernel + +theorem nib_sp (b : Nat → Bool) : (nib (BitVec.ofNat 32 (sp b 8))).toNat = cnt8 b 8 := by + have := nib_bools (b 0) (b 1) (b 2) (b 3) (b 4) (b 5) (b 6) (b 7) + simp only [sp, cnt8, Nat.pow_zero, Nat.mul_one, Nat.zero_add, Nat.pow_one] at this ⊢ + exact this + +/-- A byte mask with bit `4d` the bit `b d`, and the others clear. -/ +theorem bsum_sp {f : Nat → Bool} {b : Nat → Bool} : ∀ k, (∀ j < 4 * k, f j = (decide (j % 4 = 0) && b (j / 4))) → + VG.Proof.MlKem.X86_64.S4.bsum f (4 * k) = sp b k + | 0, _ => rfl + | k + 1, h => by + have ih := bsum_sp k fun j hj => h j (by omega) + rw [show 4 * (k + 1) = 4 * k + 1 + 1 + 1 + 1 by omega] + simp only [VG.Proof.MlKem.X86_64.S4.bsum, sp] + rw [ih, h _ (by omega), h _ (by omega), h _ (by omega), h _ (by omega)] + simp only [show (4 * k) % 4 = 0 by omega, show (4 * k + 1) % 4 = 1 by omega, show (4 * k + 2) % 4 = 2 by omega, + show (4 * k + 1 + 1 + 1) % 4 = 3 by omega, show 4 * k / 4 = k by omega, + decide_true, decide_false, Bool.true_and, Bool.false_and, Bool.toNat_false, Nat.zero_mul, Nat.add_zero, + show (1 : Nat) ≠ 0 by decide, show (2 : Nat) ≠ 0 by decide, show (3 : Nat) ≠ 0 by decide] + rw [show 2 ^ (4 * k) = 16 ^ k by rw [Nat.pow_mul]] + +theorem cntH_ok (s : State) : + WP isa (.block cntH) s fun s' => + (s'.gpr .r9 = s.gpr .r9 + BitVec.setWidth 64 (nib ((s.gpr .rax).setWidth 32)) ∧ s'.mem = s.mem ∧ + ∀ r l, s'.lane r l = s.lane r l) ∧ Keep [.rax, .rdx, .r9] s s' := by + refine WP.keep _ ?_ (by decide) + simp only [cntH] + xrun + exact ⟨rfl, fun _ _ => 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 Bc bc_ofDwords mhL mhL_le mhX_ok nib sp cnt8 nib_sp bsum_sp ymm_bit cntH_ok + sw3264) +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) + +/-- `Σ k < m, f k`. -/ +def csum (f : Nat → Nat) : Nat → Nat + | 0 => 0 + | m + 1 => csum f m + f m + +namespace YHintL + +/-- 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 + c11 : ∀ l < 2, s.lane .xmm11 l = bc (BitVec.ofNat 32 63) + 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 + mhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k) else coeffAt s₀.mem (s₀.gpr .rcx) k + r9 : (s.gpr .r9).toNat = + csum (fun k => (mhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)).toNat) (8 * i) + +theorem lane_mhX (g : Nat) (hg : g = g32 ∨ g = g88) : laneSseBlock (toY (mhX g)) = some (mhX g) := by + rcases hg with rfl | rfl <;> decide +kernel + +/-- The constants are kept by code that writes only `xmm0` to `xmm5`. -/ +theorem keep_consts {g : Nat} {s s' : State} (hyc : YC g s) (h11 : ∀ l < 2, s.lane .xmm11 l = bc (BitVec.ofNat 32 63)) + {rs : List XReg} (hr : ∀ r ∈ rs, r ∈ [.xmm0, .xmm1, .xmm2, .xmm3, .xmm4, .xmm5]) + (hl : ∀ r ∉ rs, ∀ l < 2, s'.lane r l = s.lane r l) : + YC g s' ∧ ∀ l < 2, s'.lane .xmm11 l = bc (BitVec.ofNat 32 63) := by + have n : ∀ r ∈ [XReg.xmm8, .xmm9, .xmm10, .xmm11, .xmm15], r ∉ rs := fun r hr' hr'' => by + have := hr r hr''; simp only [List.mem_cons, List.not_mem_nil, or_false] at hr' this + rcases hr' with rfl | rfl | rfl | rfl | rfl <;> rcases this with h | h | h | h | h | h <;> cases h + refine ⟨fun l hl' => ?_, fun l hl' => by rw [hl _ (n _ (by simp)) l hl']; exact h11 l hl'⟩ + rw [hl _ (n _ (by simp)) l hl', hl _ (n _ (by simp)) l hl', hl _ (n _ (by simp)) l hl', hl _ (n _ (by simp)) l hl'] + exact hyc l hl' + +/-- A hint shifted to bit 7: bit `8m + 7` is the hint's bit 0 if `m = 0`, and clear otherwise. -/ +theorem shl7_bit {v : BitVec 32} (hv : v.toNat ≤ 1) {m : Nat} (hm : m < 4) : + (v <<< 7).getLsbD (8 * m + 7) = (decide (m = 0) && v.getLsbD 0) := by + rw [BitVec.getLsbD_shiftLeft] + rcases (by omega : m = 0 ∨ 0 < m) with rfl | h + · simp + · have : v.getLsbD (8 * m) = false := by + rw [← BitVec.testBit_toNat, Nat.testBit_lt_two_pow (Nat.lt_of_le_of_lt hv + (Nat.one_lt_two_pow (by omega)))] + simp [show 8 * m + 7 - 7 = 8 * m by omega, this, show m ≠ 0 by omega] + +theorem bit0_toNat {v : BitVec 32} (hv : v.toNat ≤ 1) : (v.getLsbD 0).toNat = v.toNat := by + rw [← BitVec.testBit_toNat, Nat.testBit_zero] + rcases (by omega : v.toNat = 0 ∨ v.toNat = 1) with h | h <;> rw [h] <;> rfl + +theorem csum_le {f : Nat → Nat} : ∀ {m : Nat}, (∀ k < m, f k ≤ 1) → csum f m ≤ m + | 0, _ => Nat.le_refl _ + | m + 1, h => by + have := csum_le (m := m) fun k hk => h k (by omega) + have := h m (by omega) + simp only [csum]; omega + +theorem csum8 (f : Nat → Nat) (b : Nat → Bool) (m : Nat) (h : ∀ d < 8, (b d).toNat = f (m + d)) : + csum f (m + 8) = csum f m + cnt8 b 8 := by + rw [show csum f (m + 8) = csum f m + f m + f (m + 1) + f (m + 2) + f (m + 3) + f (m + 4) + f (m + 5) + + f (m + 6) + f (m + 7) from rfl, + show cnt8 b 8 = (b 0).toNat + (b 1).toNat + (b 2).toNat + (b 3).toNat + (b 4).toNat + (b 5).toNat + + (b 6).toNat + (b 7).toNat by simp [cnt8], + h 0 (by decide), h 1 (by decide), h 2 (by decide), h 3 (by decide), h 4 (by decide), h 5 (by decide), + h 6 (by decide), h 7 (by decide)] + simp only [Nat.add_zero] + omega + +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))) + (hz : Reduced s₀.mem (s₀.gpr .rdi)) (hr : Reduced s₀.mem (s₀.gpr .rsi)) {g : Nat} (hg : g = g32 ∨ g = g88) +include hrd hwr hdz hdr hz hr hg + +theorem step {i : Nat} (hi : i < 32) {s : State} (hI : Inv s₀ g i s) : + WP isa (.block (mhBodyY 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 hg' : g ∈ gamma2s := by rcases hg with rfl | rfl <;> decide + 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 [mhBodyY, List.append_assoc, List.append_assoc, 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 := keep_consts hI.yc hI.c11 (rs := [.xmm0, .xmm5]) (by simp) fun r hr' l hl => by + simp only [List.mem_cons, List.not_mem_nil, or_false, not_or] at hr' + rw [o2.lane r (by simp [hr'.2]) l hl, o1.lane r (by simp [hr'.1]) l hl] + rw [WP.block_append_iff] + refine WP.mono (ylanes (lane_mhX g hg) (P := fun l t => ∀ e < 4, + dword (t.xmm .xmm0) e = mhL g (dword ((s2.proj l).xmm .xmm5) e) (dword ((s2.proj l).xmm .xmm0) e) ∧ + dword (t.xmm .xmm1) e = mhL g (dword ((s2.proj l).xmm .xmm5) e) (dword ((s2.proj l).xmm .xmm0) e) <<< 7) + fun l hl => mhX_ok hg _ (k2.1.hbc hl) (k2.1.q hl) + (by rw [State.proj_xmm, k2.2 l hl]; exact bc_ofDwords _)) fun s3 ⟨B3, o3⟩ => ?_ + have o13 := (o1.trans o2).trans o3 + have k3 := keep_consts k2.1 k2.2 (rs := [.xmm0, .xmm1, .xmm2, .xmm3, .xmm4]) (by simp) o3.lane + 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 + -- The values of the lanes. + have hv : ∀ l < 2, ∀ e < 4, mhL g (dword ((s2.proj l).xmm .xmm5) e) (dword ((s2.proj l).xmm .xmm0) e) = + mhL 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)] + have hle : ∀ k < 256, (mhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)).toNat ≤ 1 := + fun k hk => mhL_le hg' (hz k (by rw [n_eq]; exact hk)) (hr k (by rw [n_eq]; exact hk)) + rw [WP.block_append_iff] + refine WP.mono (Q := fun (s4 : State) => s4.mem = s3.mem.writeW (s3.gpr .r10) (s3.ymm .xmm0) ∧ + s4.gpr = (s3.setReg .rax (byteMask (s3.ymm .xmm1) 32)).gpr ∧ s4.rd = s3.rd ∧ s4.wr = s3.wr ∧ + ∀ r l, s4.lane r l = s3.lane r l) + (by + vrund [State.store256_eq, State.setMem_gpr, State.setMem_wr, State.setMem_mem, State.setMem_rd, + State.setMem_ymm, w0] + exact ⟨rfl, fun r l => by simp only [lane_setReg, State.setMem_lane]⟩) fun s4 ⟨m4, g4, r4, w4, l4⟩ => ?_ + rw [WP.block_append_iff] + refine WP.mono (cntH_ok s4) fun s5 ⟨⟨h9, m5, l5⟩, k5⟩ => ?_ + have hax : s4.gpr .rax = byteMask (s3.ymm .xmm1) 32 := by rw [g4, RegUpd.gpr_setReg_self] + -- The count of the eight hints. + let b : Nat → Bool := fun d => + (mhL g (coeffAt s₀.mem (s₀.gpr .rdi) (8 * i + d)) (coeffAt s₀.mem (s₀.gpr .rsi) (8 * i + d))).getLsbD 0 + have hbits : ∀ j < 4 * 8, (s3.ymm .xmm1).getLsbD (8 * j + 7) = (decide (j % 4 = 0) && b (j / 4)) := fun j hj => by + have hl : j / 16 < 2 := by omega + have he : j % 16 / 4 < 4 := by omega + rw [ymm_bit _ _ (by omega), ← State.proj_xmm, (B3 _ hl _ he).2, hv _ hl _ he, + shl7_bit (hle _ (by omega)) (by omega)] + simp only [b, show 8 * i + 4 * (j / 16) + j % 16 / 4 = 8 * i + j / 4 by omega] + have hcnt : (BitVec.setWidth 64 (nib ((s4.gpr .rax).setWidth 32))).toNat = cnt8 b 8 := by + rw [hax, VG.Proof.MlKem.X86_64.S4.byteMask_eq _ (by decide), BitVec.toNat_setWidth, + show BitVec.setWidth 32 (BitVec.ofNat 64 (VG.Proof.MlKem.X86_64.S4.bsum + (fun i => (s3.ymm .xmm1).getLsbD (8 * i + 7)) 32)) = + BitVec.ofNat 32 (VG.Proof.MlKem.X86_64.S4.bsum (fun i => (s3.ymm .xmm1).getLsbD (8 * i + 7)) 32) by + apply BitVec.eq_of_toNat_eq + have := VG.Proof.MlKem.X86_64.S4.bsum_lt (fun i => (s3.ymm .xmm1).getLsbD (8 * i + 7)) 32 + simp only [BitVec.toNat_setWidth, BitVec.toNat_ofNat]; omega, + (bsum_sp 8 hbits : VG.Proof.MlKem.X86_64.S4.bsum _ 32 = _), nib_sp] + have : cnt8 b 8 ≤ 8 := by + simp only [cnt8] + have hb := fun d => Bool.toNat_le (b d) + have := hb 0; have := hb 1; have := hb 2; have := hb 3; have := hb 4; have := hb 5 + have := hb 6; have := hb 7 + omega + have := (nib_sp b).symm ▸ this + omega + simp only [List.cons_append, List.nil_append] + xrun + have G5 : ∀ r, r ∉ [Reg.rax, .rdx, .r9] → s5.gpr r = s.gpr r := fun r hr => by + rw [k5.gpr hr, g4, RegUpd.gpr_setReg_of_ne _ _ (fun h => hr (by simp [h])), o13.gpr] + refine ⟨⟨?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, fun k hk => ?_, ?_⟩, by rw [G5 .rcx (by simp)], by rw [G5 .rcx (by simp)]⟩ + · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_true, ite_false, reduceCtorEq] + rw [G5 .rdi (by simp), hI.rdi]; exact ptr_step _ i 32 + · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_true, ite_false, reduceCtorEq] + rw [G5 .rsi (by simp), hI.rsi]; exact ptr_step _ i 32 + · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_true, ite_false, reduceCtorEq] + rw [G5 .r10 (by simp), hI.r10]; exact ptr_step _ i 32 + · simp only [RegUpd.rd_setReg, RegUpd.rd_setFlags]; rw [k5.2.1, r4, o13.rd, hI.rd] + · simp only [RegUpd.wr_setReg, RegUpd.wr_setFlags]; rw [k5.2.2, w4, o13.wr, hI.wr] + · exact k3.1.keep (rs := []) (by simp) fun r _ l _ => by simp only [lane_setReg, lane_setFlags]; rw [l5, l4] + · intro l hl; simp only [lane_setReg, lane_setFlags]; rw [l5, l4]; exact k3.2 l hl + · simp only [RegUpd.mem_setReg, RegUpd.mem_setFlags] + rw [m5, m4, g3, o13.mem]; exact hI.frame.writeW (List.mem_singleton_self _) _ (pR_contains32 _ j0) + · simp only [RegUpd.mem_setReg, RegUpd.mem_setFlags] + rw [m5, m4, 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 .xmm0 l) e = + mhL 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).1, 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)] + · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_false, reduceCtorEq] + have h94 : s4.gpr .r9 = s.gpr .r9 := by + rw [g4, RegUpd.gpr_setReg_of_ne _ _ (by decide), o13.gpr] + have e8 := csum8 (fun k => (mhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)).toNat) b + (8 * i) fun d hd => bit0_toNat (hle _ (by omega)) + have hb := csum_le (f := fun k => (mhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)).toNat) + (m := 8 * i + 8) fun k hk => hle k (by omega) + rw [h9, BitVec.toNat_add, hcnt, h94, hI.r9, show 8 * (i + 1) = 8 * i + 8 by omega, e8] + rw [e8] at hb + exact Nat.mod_eq_of_lt (by omega) + +/-- The constants and the loop: the hints at `h`, and their count in `r9`. -/ +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) (hs9 : s.gpr .r9 = 0) (hsrd : s.rd = s₀.rd) (hswr : s.wr = s₀.wr) + (hsm : s.mem = s₀.mem) : + WP isa (mhY g) s fun s' => Frame [pR (s₀.gpr .rcx)] s₀.mem s'.mem ∧ + (∀ k < 256, coeffAt s'.mem (s₀.gpr .rcx) k = + mhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)) ∧ + (s'.gpr .r9).toNat = + csum (fun k => (mhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)).toNat) 256 := by + unfold mhY + refine WP.seq ?_ + rw [WP.block_append_iff] + refine WP.mono (yC_ok g s) fun w ⟨yc, k1, m1, _, _⟩ => + WP.mono (yconst_ok .xmm11 63 w) fun w2 ⟨l2, k2, m2, _, o2⟩ => ?_ + have yc2 : YC g w2 := fun l hl => by + rw [o2 .xmm8 (by decide) l hl, o2 .xmm9 (by decide) l hl, o2 .xmm10 (by decide) l hl, + o2 .xmm15 (by decide) l hl] + exact yc l hl + refine WP.mono (wp_rcxLoopY (N := 32) (by decide) (by decide) _ (fun u o hy _ => + ⟨by rw [o.keep.gpr (by decide), k2.gpr (by decide), k1.gpr (by decide), Nat.mul_zero, add_ofNat_zero, hs0], + by rw [o.keep.gpr (by decide), k2.gpr (by decide), k1.gpr (by decide), Nat.mul_zero, add_ofNat_zero, hs1], + by rw [o.keep.gpr (by decide), k2.gpr (by decide), k1.gpr (by decide), Nat.mul_zero, add_ofNat_zero, hs10], + by rw [o.keep.2.1, k2.2.1, k1.2.1, hsrd], by rw [o.keep.2.2, k2.2.2, k1.2.2, hswr], + fun l hl => by simp only [State.lane]; rw [o.xmm, hy]; exact yc2 l hl, + fun l hl => by simp only [State.lane]; rw [o.xmm, hy]; exact l2 l hl, + by rw [o.mem, m2, m1, hsm]; exact Frame.refl _ _, fun k _ => by rw [o.mem, m2, m1, hsm, ifn (by omega)], + by rw [o.keep.gpr (by decide), k2.gpr (by decide), k1.gpr (by decide), hs9]; rfl⟩) + fun i hi u hI => step hrd hwr hdz hdr hz hr hg hi hI) fun u hI => ⟨hI.frame, fun k hk => by + rw [hI.coeff k hk, ifp (by omega)], hI.r9⟩ + +end + +end YHintL + +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) +open VG.Proof.MlDsa.X86_64.Arith (csum) + +/-- The count of the hints is `onesFrom` them. -/ +theorem csum_onesFrom (v : Vector Bool n) (f : Nat → Nat) (hf : ∀ k < 256, f k = v[k]!.toNat) : + ∀ m ≤ 256, csum f m + onesFrom v m = onesFrom v 0 + | 0, _ => Nat.zero_add _ + | m + 1, hm => by + have ih := csum_onesFrom v f hf m (by omega) + rw [onesFrom_step v (by rw [n_eq]; omega), ← hf m (by omega)] at ih + simp only [csum]; omega + +section +variable {s₀ : State} (hp : makeHintK.pre s₀) +include hp + +theorem makeHintY_correct : ∃ t s', Exec isa makeHintAvx2 s₀ t s' ∧ abiPreserved s₀ s' ∧ makeHintK.post s₀ s' := by + have hg : arg32 s₀ .rdx ∈ gamma2s := hp.2.2.2.2.2.2.2.1 + have hz : Reduced s₀.mem (s₀.gpr .rdi) := hp.2.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.2 + let f : Nat → Nat := fun k => + (mhL (arg32 s₀ .rdx) (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)).toNat + have wp : WP isa makeHintAvx2 s₀ fun s' => Frame [pR (s₀.gpr .rcx)] s₀.mem s'.mem ∧ + (∀ k < 256, coeffAt s'.mem (s₀.gpr .rcx) k = + mhL (arg32 s₀ .rdx) (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)) ∧ + (s'.gpr .rax).toNat = csum f 256 := by + unfold makeHintAvx2 + refine WP.seq (WP.mono (mhPrologue_ok s₀) fun s₁ ⟨⟨h10, h9, hzf, hm⟩, hk⟩ => ?_) + have go : ∀ g, arg32 s₀ .rdx = g → (g = g32 ∨ g = g88) → WP isa (mhY g) s₁ fun s' => + Frame [pR (s₀.gpr .rcx)] s₀.mem s'.mem ∧ + (∀ k < 256, coeffAt s'.mem (s₀.gpr .rcx) k = + mhL (arg32 s₀ .rdx) (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)) ∧ + (s'.gpr .r9).toNat = csum f 256 := fun g hge hg' => by + subst hge + exact Arith.YHintL.loop_ok hp.1 hp.2.1 hp.2.2.1 hp.2.2.2.1 hz hr hg' (hk.gpr (by decide)) + (hk.gpr (by decide)) h10 h9 hk.2.1 hk.2.2 hm + have hite : WP isa (.ite .e (mhY g32) (mhY g88)) s₁ fun s' => Frame [pR (s₀.gpr .rcx)] s₀.mem s'.mem ∧ + (∀ k < 256, coeffAt s'.mem (s₀.gpr .rcx) k = + mhL (arg32 s₀ .rdx) (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)) ∧ + (s'.gpr .r9).toNat = csum f 256 := 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, h9'⟩ => ?_) + refine WP.mono (Q := fun (u' : State) => u'.mem = u.mem ∧ u'.gpr .rax = u.gpr .r9) (by vrund; exact ⟨rfl, rfl⟩) + fun u' ⟨hm', hax⟩ => ?_ + rw [hm', hax] + exact ⟨hf, hc, h9'⟩ + obtain ⟨t, s', he, ⟨hf, hv, hax⟩, hk⟩ := WP.keep [.rax, .rcx, .rdx, .rsi, .rdi, .r9, .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 hintIs_of_toNat fun k hk' => ?_ + apply BitVec.eq_of_toNat_eq + rw [hv k hk', mhL_toNat hg (hz k hk') (hr k hk'), zipWith_get _ _ _ hk', polyAt_get _ _ hk', polyAt_get _ _ hk'] + simp only [BitVec.natCast_eq_ofNat, BitVec.toNat_ofNat] + exact (Nat.mod_eq_of_lt (Nat.lt_of_le_of_lt (Bool.toNat_le _) (by decide))).symm + · have e := csum_onesFrom (Vector.zipWith (makeHint (arg32 s₀ .rdx)) (polyAt s₀.mem (s₀.gpr .rdi)) + (polyAt s₀.mem (s₀.gpr .rsi))) f (fun k hk' => by + rw [zipWith_get _ _ _ (by rw [n_eq]; exact hk'), polyAt_get _ _ (by rw [n_eq]; exact hk'), + polyAt_get _ _ (by rw [n_eq]; exact hk')] + exact mhL_toNat hg (hz k (by rw [n_eq]; exact hk')) (hr k (by rw [n_eq]; exact hk'))) 256 (Nat.le_refl _) + have e0 : onesFrom (Vector.zipWith (makeHint (arg32 s₀ .rdx)) (polyAt s₀.mem (s₀.gpr .rdi)) + (polyAt s₀.mem (s₀.gpr .rsi))) 256 = 0 := onesFrom_n _ + have : (s'.gpr .rax).toNat ≤ 256 := by + rw [hax] + exact Arith.YHintL.csum_le fun k hk' => mhL_le hg (hz k (by rw [n_eq]; exact hk')) (hr k (by rw [n_eq]; exact hk')) + rw [BitVec.toNat_setWidth, hintOnes_single] + omega + +end + +theorem makeHintY_ct : ConstantTime isa makeHintK.pre makeHintK.pub makeHintAvx2 := + 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 makeHintY_verified : Verified X86_64.target makeHintAvx2 (makeHintContract X86_64.abi) := + Verified.of_correct (fun _ hp => makeHintY_correct hp) makeHintY_ct (by + round_implies [makeHintContract, makeHintSig, makeHintK, 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/Round/YLane.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YLane.lean new file mode 100644 index 000000000..6c354f4b1 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YLane.lean @@ -0,0 +1,130 @@ +import VerifiedGarbage.Impl.MlDsa.X86_64.Round.Avx2 +import VerifiedGarbage.Proof.MlDsa.Round.Decompose +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.VArith + +/-! +# ML-DSA on x86-64: what the AVX2 rounding code computes in a doubleword + +Untrusted: everything here is checked by Lean. What `hbX` and `lbX` +(`Impl/MlDsa/X86_64/Round/Avx2.lean`) compute in each doubleword (`hbL`, +`lbL`), and that it is `r₁` and `r₀` of `Decompose` (`hbL_toNat`, +`lbL_toNat`): every intermediate value fits in 32 bits. +-/ + +namespace VG.Proof.MlDsa.X86_64.Round + +open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round +open VG.Spec.MlDsa (q gamma2s ofInt lowBits) +open VG.Proof.MlDsa.Round (hbF hbM q_eq mem_gamma2s hbF_le hbF_eq hbM_mul lowBits_val) +open VG.Proof.MlDsa.X86_64.Arith (caddL caddL_toNat sshiftRight31) + +/-- `t · (1 + Σ 2^k)` by shifts and additions, as `mulX` computes it. -/ +def mulL (sh : List Nat) (t : BitVec 32) : BitVec 32 := sh.foldl (fun acc k => acc + (t <<< k)) t + +theorem foldl_toNat (t : BitVec 32) : ∀ (sh : List Nat) (acc : BitVec 32), + acc.toNat + t.toNat * (sh.map (2 ^ ·)).sum < 2 ^ 32 → + (sh.foldl (fun acc k => acc + (t <<< k)) acc).toNat = acc.toNat + t.toNat * (sh.map (2 ^ ·)).sum + | [], acc, _ => by simp + | k :: sh, acc, h => by + simp only [List.map_cons, List.sum_cons, Nat.mul_add] at h ⊢ + have hk : (t <<< k).toNat = t.toNat * 2 ^ k := by + rw [BitVec.toNat_shiftLeft, Nat.shiftLeft_eq, Nat.mod_eq_of_lt (by omega)] + have ha : (acc + t <<< k).toNat = acc.toNat + t.toNat * 2 ^ k := by + rw [BitVec.toNat_add, hk, Nat.mod_eq_of_lt (by omega)] + rw [List.foldl_cons, foldl_toNat t sh _ (by rw [ha]; omega), ha] + omega + +theorem mulL_toNat (sh : List Nat) {t : BitVec 32} (h : t.toNat * (1 + (sh.map (2 ^ ·)).sum) < 2 ^ 32) : + (mulL sh t).toNat = t.toNat * (1 + (sh.map (2 ^ ·)).sum) := by + rw [mulL, foldl_toNat t sh t (by rw [Nat.mul_add] at h; omega), Nat.mul_add]; omega + +theorem dSh_sum {g : Nat} (h : g ∈ gamma2s) : 1 + ((dSh g).map (2 ^ ·)).sum = dMul g := by + rcases mem_gamma2s h with rfl | rfl <;> rfl + +/-- `f`, as `hbX` computes it in a doubleword. -/ +def hbFL (g : Nat) (x : BitVec 32) : BitVec 32 := + (mulL (dSh g) ((x + BitVec.ofNat 32 127) >>> 7) + BitVec.ofNat 32 (dAdd g)) >>> dShift g + +/-- What `hbX` computes in a doubleword. -/ +def hbL (g : Nat) (x : BitVec 32) : BitVec 32 := + hbFL g x &&& (hbFL g x - BitVec.ofNat 32 (dMod g)).sshiftRight (min (31 : BitVec 8).toNat 32) + +theorem dMod_eq' {g : Nat} (h : g ∈ gamma2s) : dMod g = hbM g := by + rcases mem_gamma2s h with rfl | rfl <;> rfl + +/-- `f` in 32 bits. -/ +theorem hbF32 {g : Nat} (h : g ∈ gamma2s) {x : BitVec 32} (hx : x.toNat < q) : + (hbFL g x).toNat = hbF g x.toNat := by + rw [q_eq] at hx + have e1 : (x + BitVec.ofNat 32 127).toNat = x.toNat + 127 := by + rw [BitVec.toNat_add, BitVec.toNat_ofNat]; omega + have e2 : ((x + BitVec.ofNat 32 127) >>> 7).toNat = (x.toNat + 127) / 128 := by + rw [BitVec.toNat_ushiftRight, e1, Nat.shiftRight_eq_div_pow] + have hM : dMul g ≤ 11275 := by unfold dMul; split <;> decide + have hA : dAdd g ≤ 2 ^ 23 := by unfold dAdd; split <;> decide + have hp : (x.toNat + 127) / 128 * dMul g ≤ 65473 * 11275 := Nat.mul_le_mul (by omega) hM + have e3 : (mulL (dSh g) ((x + BitVec.ofNat 32 127) >>> 7)).toNat = (x.toNat + 127) / 128 * dMul g := by + rw [mulL_toNat _ (by rw [e2, dSh_sum h]; omega), e2, dSh_sum h] + have e4 : (mulL (dSh g) ((x + BitVec.ofNat 32 127) >>> 7) + BitVec.ofNat 32 (dAdd g)).toNat = + (x.toNat + 127) / 128 * dMul g + dAdd g := by + rw [BitVec.toNat_add, e3, BitVec.toNat_ofNat, Nat.mod_eq_of_lt (show dAdd g < 2 ^ 32 by omega)]; omega + rw [hbFL, BitVec.toNat_ushiftRight, e4, Nat.shiftRight_eq_div_pow, hbF_eq h (by rw [q_eq]; exact hx)] + rfl + +theorem hbL_toNat {g : Nat} (h : g ∈ gamma2s) {x : BitVec 32} (hx : x.toNat < q) : + (hbL g x).toNat = hbF g x.toNat % hbM g := by + have hf := hbF32 h hx + have hle := hbF_le h hx + have hm : dMod g ≤ 44 ∧ 16 ≤ dMod g := by unfold dMod; split <;> decide + rw [← dMod_eq' h] at hle ⊢ + unfold hbL + generalize hbFL g x = f at hf ⊢ + rw [sshiftRight31] + have hsub : (f - BitVec.ofNat 32 (dMod g)).toNat = (f.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 + by_cases e : hbF g x.toNat = dMod g + · rw [ite_eq_left (by rw [hsub, hf, e]; omega), show f &&& (0 : BitVec 32) = 0 from BitVec.and_zero, e, Nat.mod_self]; rfl + · rw [ite_eq_right (by rw [hsub, hf]; omega), show (-1 : BitVec 32) = BitVec.allOnes 32 by decide, BitVec.and_allOnes, + hf, Nat.mod_eq_of_lt (by omega)] + +/-- `r₁ · 2γ₂` by shifts and additions, as `mul2X` computes it. -/ +def mul2L (g : Nat) (r : BitVec 32) : BitVec 32 := + if g = 261888 then r <<< 19 - r <<< 9 + else [13, 14, 15, 17].foldl (fun acc k => acc + (r <<< k)) (r <<< 11) + +theorem mul2L_toNat {g : Nat} (h : g ∈ gamma2s) {r : BitVec 32} (hr : r.toNat ≤ 44) : + (mul2L g r).toNat = r.toNat * (2 * g) := by + have s : ∀ k ≤ 19, (r <<< k).toNat = r.toNat * 2 ^ k := fun k hk => by + rw [BitVec.toNat_shiftLeft, Nat.shiftLeft_eq, Nat.mod_eq_of_lt] + calc r.toNat * 2 ^ k ≤ 44 * 2 ^ 19 := Nat.mul_le_mul hr (Nat.pow_le_pow_right (by decide) hk) + _ < 2 ^ 32 := by decide + rcases mem_gamma2s h with rfl | rfl + · unfold mul2L + rw [ite_eq_right (by decide)] + simp only [List.foldl_cons, List.foldl_nil] + rw [BitVec.toNat_add, BitVec.toNat_add, BitVec.toNat_add, BitVec.toNat_add, s 11 (by decide), s 13 (by decide), + s 14 (by decide), s 15 (by decide), s 17 (by decide)] + omega + · unfold mul2L + rw [ite_eq_left rfl, BitVec.toNat_sub, s 19 (by decide), s 9 (by decide)] + omega + +/-- What `lbX` computes in a doubleword. -/ +def lbL (g : Nat) (x : BitVec 32) : BitVec 32 := caddL (x - mul2L g (hbL g x)) + +theorem lbL_toNat {g : Nat} (h : g ∈ gamma2s) {x : BitVec 32} (hx : x.toNat < q) : + (lbL g x).toNat = (ofInt (lowBits g (Fin.ofNat q x.toNat))).val := by + have hv : (Fin.ofNat q x.toNat).val = x.toNat := Nat.mod_eq_of_lt hx + rw [lowBits_val h, hv] + have hlt : hbF g x.toNat % hbM g < hbM g := Nat.mod_lt _ (by rcases mem_gamma2s h with rfl | rfl <;> decide) + have hM : hbM g ≤ 44 := by rcases mem_gamma2s h with rfl | rfl <;> decide + have hle : hbF g x.toNat % hbM g * (2 * g) ≤ q - 1 := by + rw [← hbM_mul h]; exact Nat.mul_le_mul_right _ (Nat.le_of_lt hlt) + have e1 : (mul2L g (hbL g x)).toNat = hbF g x.toNat % hbM g * (2 * g) := by + rw [mul2L_toNat h (by rw [hbL_toNat h hx]; omega), hbL_toNat h hx] + unfold lbL + rw [caddL_toNat, BitVec.toNat_sub, e1] + rw [q_eq] at hx hle ⊢ + split <;> split <;> omega + +end VG.Proof.MlDsa.X86_64.Round diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YNorm.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YNorm.lean new file mode 100644 index 000000000..a94ca1e18 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YNorm.lean @@ -0,0 +1,403 @@ +import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YBits +import VerifiedGarbage.Proof.MlDsa.X86_64.Round.NormLt +import VerifiedGarbage.Proof.MlKem.X86_64.S4Vec + +/-! +# ML-DSA on x86-64: `vg_mldsa_norm_lt_avx2` + +Untrusted: everything here is checked by Lean. With the bound clamped to +`b ≤ q` (which changes no result, `good_clamp`), each doubleword of `ymm10` +keeps its top bit while every coefficient it has seen is good +(`Good b a`: `a < b` or `q - a < b`, `nlL_msb`); at the end the top bits +of the 32 bytes are all set exactly when every coefficient is good +(`allOnes_bsum`). +-/ + +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 XOnly YOnly ylanes ifp ifn) +open VG.Impl.MlKem.X86_64 (xb xmov toY) +open VG.Proof.MlDsa.X86_64.Arith (dword_psubd dword_pand dword_psrad) +open VG.Proof.MlKem.X86_64.S4 (bsum bsum_lt byteMask_eq) +open VG.Proof.MlDsa.X86_64.Arith (bc) + +/-! ## A doubleword -/ + +/-- The top bit of a difference of doublewords below `2³¹`: whether it borrows. -/ +theorem msb_sub32 {u v : BitVec 32} (hu : u.toNat < 2 ^ 31) (hv : v.toNat < 2 ^ 31) : + (u - v).msb = decide (u.toNat < v.toNat) := by + rw [BitVec.msb_eq_decide, BitVec.toNat_sub] + by_cases h : u.toNat < v.toNat + · rw [decide_eq_true h, decide_eq_true_iff]; omega + · rw [decide_eq_false h, decide_eq_false_iff_not]; omega + +/-- What `nlX` ORs into a doubleword: `(a - b) | ((q - b) - a)`. -/ +def nlL (a b c : BitVec 32) : BitVec 32 := (a - b) ||| (c - a) + +theorem nlL_msb {a b : BitVec 32} (ha : a.toNat < q) (hb : b.toNat ≤ q) : + (nlL a b (BitVec.ofNat 32 (q - b.toNat))).msb = decide (Good b.toNat a.toNat) := by + rw [q_eq] at ha hb + have hc : (BitVec.ofNat 32 (q - b.toNat)).toNat = q - b.toNat := by rw [BitVec.toNat_ofNat, q_eq]; omega + rw [nlL, BitVec.msb_or, msb_sub32 (by omega) (by omega), msb_sub32 (by rw [hc, q_eq]; omega) (by omega), hc] + simp only [Good, Bool.decide_or, q_eq] + congr 1 + apply Bool.eq_iff_iff.mpr + simp only [decide_eq_true_iff] + omega + +/-- Clamping the bound to `q` changes no coefficient's goodness. -/ +theorem good_clamp {B a : Nat} (ha : a < q) : Good (min B q) a ↔ Good B a := by + simp only [Good] + constructor + · rintro (h | h) <;> [left; right] <;> omega + · intro h + by_cases hB : B ≤ q + · rw [Nat.min_eq_left hB]; exact h + · rw [Nat.min_eq_right (by omega)]; left; exact ha + +/-! ## The mask -/ + +theorem bsum_allOnes (f : Nat → Bool) : ∀ n, bsum f n = 2 ^ n - 1 ↔ ∀ i < n, f i = true + | 0 => by simp [bsum] + | n + 1 => by + have hl := bsum_lt f n + have ih := bsum_allOnes f n + have h1 : 1 ≤ 2 ^ n := Nat.one_le_two_pow + simp only [bsum, Nat.pow_succ] + generalize 2 ^ n = N at hl ih h1 ⊢ + constructor + · intro h i hi + have hfn : f n = true := by + cases e : f n + · simp only [e, Bool.toNat_false, Nat.zero_mul, Nat.add_zero] at h; omega + · rfl + rw [hfn] at h + simp only [Bool.toNat_true, Nat.one_mul] at h + rcases (by omega : i < n ∨ i = n) with hi | rfl + · exact (ih.mp (by omega)) i hi + · exact hfn + · intro h + rw [ih.mpr fun i hi => h i (by omega), h n (by omega)] + simp only [Bool.toNat_true, Nat.one_mul] + omega + +/-- A bit of a 256-bit register, in its lanes and doublewords. -/ +theorem ymm_bit (s : State) (r : XReg) {i : Nat} (hi : i < 32) : + (s.ymm r).getLsbD (8 * i + 7) = (dword (s.lane r (i / 16)) (i % 16 / 4)).getLsbD (8 * (i % 4) + 7) := by + simp only [State.ymm, State.lane, dword, BitVec.getLsbD_append, BitVec.getLsbD_extractLsb', + show 8 * (i % 4) + 7 < 32 by omega, decide_true, Bool.true_and] + by_cases h : i < 16 + · rw [ite_eq_left (show 8 * i + 7 < 128 by omega), ite_eq_left (show i / 16 = 0 by omega), + show 32 * (i % 16 / 4) + (8 * (i % 4) + 7) = 8 * i + 7 by omega] + · rw [ite_eq_right (show ¬ 8 * i + 7 < 128 by omega), ite_eq_right (show ¬ i / 16 = 0 by omega), + show 32 * (i % 16 / 4) + (8 * (i % 4) + 7) = 8 * i + 7 - 128 by omega] + +/-- A doubleword spread from its top bit. -/ +theorem bit_of_spread {d : BitVec 32} (h : d = 0 ∨ d = BitVec.allOnes 32) {j : Nat} (hj : j < 32) : + d.getLsbD j = d.msb := by + rcases h with rfl | rfl + · simp + · rw [BitVec.msb_allOnes (by decide), BitVec.getLsbD_allOnes]; simp [hj] + +/-! ## The code on a register -/ + +theorem nlX_ok (s : State) : + WP isa (.block nlX) s fun s' => (∀ e < 4, dword (s'.xmm .xmm10) e = + dword (s.xmm .xmm10) e &&& nlL (dword (s.xmm .xmm0) e) (dword (s.xmm .xmm8) e) (dword (s.xmm .xmm9) e)) ∧ + XOnly [.xmm1, .xmm2, .xmm10] s s' := by + simp only [nlX, xmov, xb] + vrun [VG.X86_64.eval_movdqa] + refine ⟨fun e he => ?_, by xonly⟩ + simp only [dword_pand, dword_por, dword_psubd _ _ he] + rfl + +theorem lane_nlX : laneSseBlock (toY nlX) = some nlX := by decide +kernel + +/-- `rax`'s low doubleword in each doubleword of `ymm r`. -/ +theorem ybc_ok (r : XReg) (s : State) : + WP isa (.block (ybcast r)) s fun s' => + (∀ l < 2, s'.lane r l = bc ((s.gpr .rax).setWidth 32)) ∧ s'.gpr = s.gpr ∧ s'.mem = s.mem ∧ s'.rd = s.rd ∧ + s'.wr = s.wr ∧ s'.mxcsr = s.mxcsr ∧ ∀ r' ≠ r, ∀ l < 2, s'.lane r' l = s.lane r' l := by + apply WP.of_runBlock + simp only [ybcast, runBlock_cons, runStep_some, runBlock_nil, exec, VOp.exec, Option.some.injEq, + exists_eq_left'] + refine ⟨fun l hl => ?_, rfl, rfl, rfl, rfl, rfl, fun r' hr' l hl => ?_⟩ + · rw [State.lane_setV256, ifp rfl] + have : dword ((s.setV .l128 r ((0 : BitVec 64) ++ s.gpr .rax) 0).xmm r) 0 = (s.gpr .rax).setWidth 32 := by + rw [RegUpd.xmm_setV, ifp rfl] + apply BitVec.eq_of_getLsbD_eq; intro i hi + simp only [dword, BitVec.getLsbD_extractLsb', hi, decide_true, Bool.true_and] + rw [BitVec.getLsbD_append, ifp (by omega), BitVec.getLsbD_setWidth] + simp [show i < 64 by omega, hi] + rcases lane01 hl with rfl | rfl <;> simp only [ite_true, ite_false, this, Nat.one_ne_zero] <;> rfl + · rw [State.lane_setV256, ifn hr', State.lane_setV128, ifn hr'] + +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 (Good nlL nlL_msb nlX_ok lane_nlX bc_ofDwords) +open VG.Proof.MlKem.X86_64 (Keep XOnly YOnly ylanes yld_ok ifp ifn ptr_step add_ofNat_zero lane_setReg lane_setFlags) +open VG.Spec.MlDsa (q n coeffAt Reduced) + +namespace YNormL + +/-- After `i` vectors of eight: the top bit of doubleword `e` of lane `l` of +`ymm10` is whether coefficients `8j + 4l + e`, `j < i`, are good. -/ +structure Inv (s₀ : State) (b : BitVec 32) (i : Nat) (s : State) : Prop where + rdi : s.gpr .rdi = s₀.gpr .rdi + BitVec.ofNat 64 (32 * i) + rd : s.rd = s₀.rd + wr : s.wr = s₀.wr + mem : s.mem = s₀.mem + c8 : ∀ l < 2, s.lane .xmm8 l = bc b + c9 : ∀ l < 2, s.lane .xmm9 l = bc (BitVec.ofNat 32 (q - b.toNat)) + acc : ∀ l < 2, ∀ e < 4, (dword (s.lane .xmm10 l) e).msb = true ↔ + ∀ j < i, Good b.toNat (coeffAt s₀.mem (s₀.gpr .rdi) (8 * j + 4 * l + e)).toNat + +theorem step {s₀ : State} (hrd : pR (s₀.gpr .rdi) ∈ s₀.rd) (hr : Reduced s₀.mem (s₀.gpr .rdi)) {b : BitVec 32} + (hb : b.toNat ≤ q) {i : Nat} (hi : i < 32) {s : State} (hI : Inv s₀ b i s) : + WP isa (.block (nlBodyY ++ ([.alu .sub .rcx (.imm 1)] : List Instr))) s fun s' => + Inv s₀ b (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 hr' : pR (s₀.gpr .rdi) ∈ s.rd ++ s.wr := by rw [hI.rd]; exact List.mem_append_left _ hrd + have e1 : s.gpr .rdi + BitVec.ofNat 64 0 = coeffAddr (s₀.gpr .rdi) (8 * i) := by + rw [add_ofNat_zero, hI.rdi]; congr 2; omega + rw [nlBodyY, List.append_assoc, List.append_assoc, WP.block_append_iff] + refine WP.mono (yld_ok (by rw [e1]; exact f_in32 hr' j0)) fun s1 ⟨L1, o1⟩ => ?_ + rw [WP.block_append_iff] + refine WP.mono (ylanes lane_nlX (P := fun l t => ∀ e < 4, dword (t.xmm .xmm10) e = + dword ((s1.proj l).xmm .xmm10) e &&& nlL (dword ((s1.proj l).xmm .xmm0) e) (dword ((s1.proj l).xmm .xmm8) e) + (dword ((s1.proj l).xmm .xmm9) e)) fun l hl => nlX_ok _) fun s3 ⟨B3, o3⟩ => ?_ + have o13 := o1.trans o3 + vrund + refine ⟨⟨?_, ?_, ?_, ?_, fun l hl => ?_, fun l hl => ?_, fun l hl e he => ?_⟩, ?_⟩ + · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_true, ite_false, reduceCtorEq] + rw [o13.gpr, hI.rdi]; exact ptr_step _ i 32 + · simp only [RegUpd.rd_setReg, RegUpd.rd_setFlags]; rw [o13.rd, hI.rd] + · simp only [RegUpd.wr_setReg, RegUpd.wr_setFlags]; rw [o13.wr, hI.wr] + · simp only [RegUpd.mem_setReg, RegUpd.mem_setFlags]; rw [o13.mem, hI.mem] + · simp only [lane_setReg, lane_setFlags]; rw [o13.lane _ (by decide) l hl, hI.c8 l hl] + · simp only [lane_setReg, lane_setFlags]; rw [o13.lane _ (by decide) l hl, hI.c9 l hl] + · simp only [lane_setReg, lane_setFlags] + have hx : dword ((s1.proj l).xmm .xmm0) e = coeffAt s₀.mem (s₀.gpr .rdi) (8 * i + 4 * l + e) := by + rw [State.proj_xmm, L1 l hl, e1, dword_readW _ _ he, lane_load, coeffAddr_add, ← coeffAt_eq, hI.mem] + have h8 : dword ((s1.proj l).xmm .xmm8) e = b := by + rw [State.proj_xmm, o1.lane _ (by decide) l hl, hI.c8 l hl]; exact bc_ofDwords _ e he + have h9 : dword ((s1.proj l).xmm .xmm9) e = BitVec.ofNat 32 (q - b.toNat) := by + rw [State.proj_xmm, o1.lane _ (by decide) l hl, hI.c9 l hl]; exact bc_ofDwords _ e he + have h10 : (s1.proj l).xmm .xmm10 = s.lane .xmm10 l := by rw [State.proj_xmm, o1.lane _ (by decide) l hl] + rw [← State.proj_xmm, B3 l hl e he, BitVec.msb_and, Bool.and_eq_true, hx, h8, h9, h10, + nlL_msb (by rw [← VG.Proof.MlDsa.Round.n_eq] at *; exact hr _ (by rw [VG.Proof.MlDsa.Round.n_eq]; omega)) hb, + decide_eq_true_iff, hI.acc l hl e he] + refine ⟨fun ⟨h₁, h₂⟩ j hj => ?_, fun h => ⟨fun j hj => h j (by omega), h i (by omega)⟩⟩ + rcases (by omega : j < i ∨ j = i) with hj | rfl + exacts [h₁ j hj, h₂] + · exact ⟨by rw [o13.gpr], by rw [o13.gpr]⟩ + +end YNormL + +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 XOnly YOnly ylanes yconst_ok ifp ifn) +open VG.Proof.MlDsa.X86_64.Arith (bc dword_psrad sshiftRight31) +open VG.Impl.MlKem.X86_64 (toY) +open VG.Proof.MlKem.X86_64.S4 (bsum byteMask_eq) + +theorem sw3264 (y : BitVec 32) : BitVec.setWidth 32 (BitVec.setWidth 64 y) = y := by + apply BitVec.eq_of_toNat_eq; rw [BitVec.toNat_setWidth, BitVec.toNat_setWidth]; have := y.isLt; omega + +/-- The bound, clamped to `q`. -/ +theorem nlPro_ok (s₀ : State) : + WP isa nlPro s₀ fun s₂ => (s₂.mem = s₀.mem ∧ ((s₂.gpr .rsi).setWidth 32).toNat = min (arg32 s₀ .rsi) q) ∧ + Keep [.rsi] s₀ s₂ := by + unfold nlPro + refine WP.seq (WP.mono (Q := fun (s₁ : State) => (s₁.cf = some (decide (((s₀.gpr .rsi).setWidth 32).toNat < q)) ∧ + s₁.mem = s₀.mem ∧ s₁.gpr .rsi = BitVec.setWidth 64 ((s₀.gpr .rsi).setWidth 32)) ∧ Keep [.rsi] s₀ s₁) + (by refine WP.keep _ ?_ (by decide); xrun; rfl) + fun s₁ ⟨⟨hc, hm, hsi⟩, hk⟩ => ?_) + refine WP.ite (M := isa) _ (show isa.eval .b s₁ = _ from hc) (fun h => ?_) (fun h => ?_) + · rw [decide_eq_true_iff] at h + refine WP.mono (Q := fun s₂ => s₂ = s₁) (by vrund) fun s₂ e => ?_ + subst e + refine ⟨⟨hm, ?_⟩, hk⟩ + rw [hsi, sw3264, Nat.min_eq_left (by unfold arg32; omega)] + · rw [decide_eq_false_iff_not] at h + refine WP.mono (Q := fun (s₂ : State) => (s₂.gpr .rsi = BitVec.setWidth 64 qImm ∧ s₂.mem = s₁.mem) ∧ + Keep [.rsi] s₁ s₂) (by refine WP.keep _ ?_ (by decide); xrun) fun s₂ ⟨⟨h1, h2⟩, k2⟩ => ?_ + refine ⟨⟨h2.trans hm, ?_⟩, (hk.trans k2).mono (by simp)⟩ + rw [h1, Nat.min_eq_right (by unfold arg32; omega)]; rfl + +theorem qsub_eq {x : BitVec 32} (hx : x.toNat ≤ q) : qImm - x = BitVec.ofNat 32 (q - x.toNat) := by + apply BitVec.eq_of_toNat_eq + rw [BitVec.toNat_sub, BitVec.toNat_ofNat, show qImm.toNat = q from rfl] + rw [q_eq] at hx ⊢; omega + +/-- The constants of the loop: `b`, `q - b` and all ones. -/ +theorem nlConsts_ok (s : State) (hb : ((s.gpr .rsi).setWidth 32).toNat ≤ q) : + WP isa (.block nlConsts) s fun s' => + (∀ l < 2, s'.lane .xmm8 l = bc ((s.gpr .rsi).setWidth 32) ∧ + s'.lane .xmm9 l = bc (BitVec.ofNat 32 (q - ((s.gpr .rsi).setWidth 32).toNat)) ∧ + s'.lane .xmm10 l = bc (BitVec.allOnes 32)) ∧ Keep [.rax] s s' ∧ s'.mem = s.mem ∧ s'.mxcsr = s.mxcsr := by + simp only [nlConsts, List.append_assoc] + rw [WP.block_append_iff] + refine WP.mono (Q := fun (s1 : State) => (s1.gpr .rax = BitVec.setWidth 64 ((s.gpr .rsi).setWidth 32) ∧ + s1.mem = s.mem ∧ s1.mxcsr = s.mxcsr ∧ ∀ r l, s1.lane r l = s.lane r l) ∧ Keep [.rax] s s1) + (by refine WP.keep _ ?_ (by decide); xrun; exact ⟨rfl, fun _ _ => rfl⟩) fun s1 ⟨⟨a1, m1, x1, l1⟩, k1⟩ => ?_ + rw [WP.block_append_iff] + refine WP.mono (ybc_ok .xmm8 s1) fun s2 ⟨c2, g2, m2, r2, w2, x2, o2⟩ => ?_ + rw [WP.block_append_iff] + refine WP.mono (Q := fun (s3 : State) => (s3.gpr .rax = BitVec.setWidth 64 (qImm - (s2.gpr .rsi).setWidth 32) ∧ + s3.mem = s2.mem ∧ s3.mxcsr = s2.mxcsr ∧ ∀ r l, s3.lane r l = s2.lane r l) ∧ Keep [.rax] s2 s3) + (by refine WP.keep _ ?_ (by decide); xrun; exact ⟨rfl, fun _ _ => rfl⟩) fun s3 ⟨⟨a3, m3, x3, l3⟩, k3⟩ => ?_ + rw [WP.block_append_iff] + refine WP.mono (ybc_ok .xmm9 s3) fun s4 ⟨c4, g4, m4, r4, w4, x4, o4⟩ => ?_ + refine WP.mono (yconst_ok .xmm10 _ s4) fun s5 ⟨c5, k5, m5, x5, o5⟩ => ?_ + have e2 : s2.gpr .rsi = s.gpr .rsi := by rw [g2, k1.gpr (by decide)] + refine ⟨fun l hl => ⟨?_, ?_, ?_⟩, ?_, by rw [m5, m4, m3, m2, m1], by rw [x5, x4, x3, x2, x1]⟩ + · rw [o5 _ (by decide) l hl, o4 _ (by decide) l hl, l3, c2 l hl, a1, sw3264] + · rw [o5 _ (by decide) l hl, c4 l hl, a3, e2, sw3264, qsub_eq hb] + · rw [c5 l hl]; rfl + · refine ⟨fun r hr => ?_, ?_, ?_⟩ + · rw [k5.gpr hr, g4, k3.gpr hr, g2, k1.gpr hr] + · rw [k5.2.1, r4, k3.2.1, r2, k1.2.1] + · rw [k5.2.2, w4, k3.2.2, w2, k1.2.2] + +theorem lane_psrad : laneSseBlock (toY [.xop (.shift .psrad .xmm10 31)]) = some [.xop (.shift .psrad .xmm10 31)] := by + decide +kernel + +/-- Whether every doubleword of `ymm10` has its top bit set. -/ +def allMsb (s : State) : Bool := (List.range 2).all fun l => (List.range 4).all fun e => (dword (s.lane .xmm10 l) e).msb + +/-- The end: 1 if every doubleword of `ymm10` has its top bit set, and 0 otherwise. -/ +theorem nlEnd_ok (s : State) : + WP isa (.block nlEnd) s fun s' => + ((s'.gpr .rax).setWidth 32 = if allMsb s = true then 1 else 0) ∧ + s'.mem = s.mem ∧ s'.rd = s.rd ∧ s'.wr = s.wr ∧ ∀ r, r ≠ .rax → s'.gpr r = s.gpr r := by + rw [nlEnd, List.append_assoc, WP.block_append_iff] + refine WP.mono (ylanes lane_psrad (rs := [.xmm10]) (P := fun l t => ∀ e < 4, + dword (t.xmm .xmm10) e = (dword ((s.proj l).xmm .xmm10) e).sshiftRight (min (31 : BitVec 8).toNat 32)) + fun l hl => by vrun; exact ⟨fun e he => dword_psrad _ _ he, by xonly⟩) fun s1 ⟨B1, o1⟩ => ?_ + rw [WP.block_append_iff] + refine WP.mono (Q := fun (s2 : State) => s2.gpr = (s1.setReg .rax (byteMask (s1.ymm .xmm10) 32)).gpr ∧ + s2.mem = s1.mem ∧ s2.rd = s1.rd ∧ s2.wr = s1.wr) (by vrund; exact ⟨rfl, rfl, rfl, rfl⟩) fun s2 ⟨g2, m2, r2, w2⟩ => ?_ + xrun + have hf : ∀ i < 32, (s1.ymm .xmm10).getLsbD (8 * i + 7) = (dword (s.lane .xmm10 (i / 16)) (i % 16 / 4)).msb := + fun i hi => by + have hl : i / 16 < 2 := by omega + have he : i % 16 / 4 < 4 := by omega + have hd := B1 (i / 16) hl (i % 16 / 4) he + rw [State.proj_xmm, State.proj_xmm, sshiftRight31] at hd + rw [ymm_bit _ _ hi, hd] + split + · rename_i h + rw [BitVec.msb_eq_decide]; simp [h] + · rename_i h + rw [show (-1 : BitVec 32) = BitVec.allOnes 32 by decide, BitVec.getLsbD_allOnes, BitVec.msb_eq_decide] + simp only [decide_eq_true (show 8 * (i % 4) + 7 < 32 by omega)] + rw [decide_eq_true (by omega)] + have hall : allMsb s = true ↔ ∀ i < 32, (s1.ymm .xmm10).getLsbD (8 * i + 7) = true := by + simp only [allMsb, List.all_eq_true, List.mem_range] + constructor + · intro h i hi; rw [hf i hi]; exact h _ (by omega) _ (by omega) + · intro h l hl e he + have := h (16 * l + 4 * e) (by omega) + rwa [hf _ (by omega), show (16 * l + 4 * e) / 16 = l by omega, show (16 * l + 4 * e) % 16 / 4 = e by omega] + at this + have hb := VG.Proof.MlKem.X86_64.S4.bsum_lt (fun i => (s1.ymm .xmm10).getLsbD (8 * i + 7)) 32 + have hax : s2.gpr .rax = BitVec.ofNat 64 (bsum (fun i => (s1.ymm .xmm10).getLsbD (8 * i + 7)) 32) := by + rw [g2, RegUpd.gpr_setReg_self, byteMask_eq _ (by decide)] + refine ⟨?_, by rw [m2, o1.mem], by rw [r2, o1.rd], by rw [w2, o1.wr], fun r hr => ?_⟩ + · rw [hax] + apply BitVec.eq_of_toNat_eq + rw [BitVec.toNat_setWidth, BitVec.toNat_ushiftRight, BitVec.toNat_add, BitVec.toNat_ofNat, + Nat.shiftRight_eq_div_pow] + have e1 : (1 : BitVec 64).toNat = 1 := rfl + rw [e1] + by_cases h : allMsb s = true + · rw [ite_eq_left h, (bsum_allOnes _ 32).mpr (hall.mp h)]; rfl + · rw [ite_eq_right h] + have : bsum (fun i => (s1.ymm .xmm10).getLsbD (8 * i + 7)) 32 ≠ 2 ^ 32 - 1 := fun e => + h (hall.mpr ((bsum_allOnes _ 32).mp e)) + show _ = 0 + omega + · rw [ite_eq_right hr, ite_eq_right hr, g2, RegUpd.gpr_setReg_of_ne _ _ hr, o1.gpr] + +section +variable {s₀ : State} (hp : normLtK.pre s₀) +include hp + +theorem normLtY_wp : WP isa normLtAvx2 s₀ fun s' => + ((s'.gpr .rax).setWidth 32 = if normRq [polyAt s₀.mem (s₀.gpr .rdi)] < arg32 s₀ .rsi then 1 else 0) ∧ + Frame [] s₀.mem s'.mem := by + have hr : Reduced s₀.mem (s₀.gpr .rdi) := hp.2.2.2 + have hrd : pR (s₀.gpr .rdi) ∈ s₀.rd := by rw [hp.1]; simp + unfold normLtAvx2 + refine WP.seq (WP.mono (nlPro_ok s₀) fun s₂ ⟨⟨hm2, hb2⟩, k2⟩ => ?_) + have hbq : ((s₂.gpr .rsi).setWidth 32).toNat ≤ q := by rw [hb2]; exact Nat.min_le_right _ _ + refine WP.seq (WP.mono (nlConsts_ok s₂ hbq) fun s₃ ⟨hc, k3, m3, _⟩ => ?_) + refine WP.seq (WP.mono (VG.Proof.MlKem.X86_64.wp_rcxLoopY (N := 32) (by decide) (by decide) + (Arith.YNormL.Inv s₀ ((s₂.gpr .rsi).setWidth 32)) (fun u o hy _ => ⟨?_, ?_, ?_, ?_, fun l hl => ?_, fun l hl => ?_, + fun l hl e he => ?_⟩) fun i hi u hI => Arith.YNormL.step hrd hr hbq hi hI) fun s₄ hI => ?_) + · rw [o.keep.gpr (by decide), k3.gpr (by decide), k2.gpr (by decide), Nat.mul_zero, + VG.Proof.MlKem.X86_64.add_ofNat_zero] + · rw [o.keep.2.1, k3.2.1, k2.2.1] + · rw [o.keep.2.2, k3.2.2, k2.2.2] + · rw [o.mem, m3, hm2] + · simp only [State.lane]; rw [o.xmm, hy]; exact (hc l hl).1 + · simp only [State.lane]; rw [o.xmm, hy]; exact (hc l hl).2.1 + · have : u.lane .xmm10 l = bc (BitVec.allOnes 32) := by simp only [State.lane]; rw [o.xmm, hy]; exact (hc l hl).2.2 + rw [this, show dword (bc (BitVec.allOnes 32)) e = BitVec.allOnes 32 from bc_ofDwords _ e he] + exact iff_of_true (by decide) fun j hj => absurd hj (Nat.not_lt_zero j) + refine WP.mono (nlEnd_ok s₄) fun s' ⟨hv, hm', _, _, _⟩ => ⟨?_, ?_⟩ + · have hnorm : normRq [polyAt s₀.mem (s₀.gpr .rdi)] < arg32 s₀ .rsi ↔ allMsb s₄ = true := by + rw [normRq_lt] + simp only [allMsb, List.all_eq_true, List.mem_range] + constructor + · intro h l hl e he + rw [(hI.acc l hl e he)] + intro j hj + have hk : 8 * j + 4 * l + e < 256 := by omega + have := h _ hk + rw [normZq_lt, polyAt_val hr hk] at this + rw [hb2] + exact (good_clamp (by rw [← VG.Proof.MlDsa.Round.n_eq] at *; exact hr _ hk)).mpr this + · intro h k hk + have hk' : k < 256 := by rw [VG.Proof.MlDsa.Round.n_eq] at hk; exact hk + have := (hI.acc (k % 8 / 4) (by omega) (k % 4) (by omega)).mp (h _ (by omega) _ (by omega)) (k / 8) (by omega) + rw [show 8 * (k / 8) + 4 * (k % 8 / 4) + k % 4 = k by omega, hb2] at this + rw [normZq_lt, polyAt_val hr hk] + exact (good_clamp (by rw [← VG.Proof.MlDsa.Round.n_eq] at *; exact hr _ hk)).mp this + rw [hv] + by_cases h : allMsb s₄ = true + · rw [ite_eq_left h, ite_eq_left (hnorm.mpr h)] + · rw [ite_eq_right h, ite_eq_right (fun h' => h (hnorm.mp h'))] + · rw [hm', hI.mem]; exact Frame.refl _ _ + +theorem normLtY_correct : ∃ t s', Exec isa normLtAvx2 s₀ t s' ∧ abiPreserved s₀ s' ∧ normLtK.post s₀ s' := by + obtain ⟨t, s', he, ⟨hv, hf⟩, hk⟩ := WP.keep [.rax, .rcx, .rsi, .rdi] (normLtY_wp hp) (by decide +kernel) + exact ⟨t, s', he, abiPreserved_of_exec (by decide +kernel) he (gprPreserved_of hk (by decide) hf (by simp)), hv⟩ + +end + +theorem normLtY_ct : ConstantTime isa normLtK.pre normLtK.pub normLtAvx2 := + VG.Taint.constantTime (A := X86_64.taint) (regsLo [.rdi, .rsp] [.rsi]) + (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 + exacts [hp.1, hp.2.1]) fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact hp.2.2) + (by taint_decide) + +theorem normLtY_verified : Verified X86_64.target normLtAvx2 (normLtContract X86_64.abi) := + Verified.of_correct (fun _ hp => normLtY_correct hp) normLtY_ct (by + round_implies [normLtContract, normLtSig, normLtK, X86_64.abi, X86_64.argRegs] [normSat] using normSat) + +end VG.Proof.MlDsa.X86_64.Round diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Inst.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Inst.lean index 0d9b4dcd2..938cf8fa4 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Inst.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Inst.lean @@ -63,6 +63,10 @@ def primsWith (B : Impl.MlDsa.X86_64.Arith.Backend) : Prims := mulAdd := B.mulAdd add := B.add sub := B.sub + highBits := B.highBits + lowBits := B.lowBits + normLt := B.normLt + makeHint := B.makeHint sfx := B.sfx } theorem nosp_of {c : Prog isa} (h : c.allInstrs (fun i => !Taint.clobbers i .rsp) = true) : NoSp c := by @@ -141,14 +145,10 @@ def prims_okWith (v : ArithImpl) : PrimsOk (primsWith v.code) signStack where ball := (⟨16, by decide, Proof.MlDsa.X86_64.Sample.sampleInBall_verified, nosp_of (by decide +kernel), by decide +kernel⟩ : Callee _ signStack prims.ball) - highBits := (⟨0, by decide, Proof.MlDsa.X86_64.Round.highBits_verified, nosp_of (by decide +kernel), by decide +kernel⟩ : - Callee _ signStack prims.highBits) - lowBits := (⟨0, by decide, Proof.MlDsa.X86_64.Round.lowBits_verified, nosp_of (by decide +kernel), by decide +kernel⟩ : - Callee _ signStack prims.lowBits) - normLt := (⟨0, by decide, Proof.MlDsa.X86_64.Round.normLt_verified, nosp_of (by decide +kernel), by decide +kernel⟩ : - Callee _ signStack prims.normLt) - makeHint := (⟨0, by decide, Proof.MlDsa.X86_64.Round.makeHint_verified, nosp_of (by decide +kernel), by decide +kernel⟩ : - Callee _ signStack prims.makeHint) + highBits := calleeOf v.ok.highBits + lowBits := calleeOf v.ok.lowBits + normLt := calleeOf v.ok.normLt + makeHint := calleeOf v.ok.makeHint simpleBitPack := (⟨0, by decide, Proof.MlDsa.X86_64.Pack.simpleBitPack_verified, nosp_of (by decide +kernel), by decide +kernel⟩ : Callee _ signStack prims.simpleBitPack) diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/PrimsC.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/PrimsC.lean index 8560f2aa7..450aeefea 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/PrimsC.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/PrimsC.lean @@ -160,9 +160,9 @@ theorem normArgs_ok {bs : List (Reg × Nat)} {f : Ptr} {B : Nat} (hB : B < 2 ^ 3 simp only [normChk, Bool.and_eq_true, decide_eq_true_eq] at hc simp only [List.all_cons, List.all_nil, Arg.ok, hc.1.2, hc.2, decide_true, Bool.and_true, decide_eq_true hB] -theorem normCall_ok {P : Prims} (hP : PrimsOk P D) {s : State} (L : Lay D rbs wbs s) {f : Ptr} {B : Nat} +theorem normCall_ok {nm : String} {P : Prims} (hP : PrimsOk P D) {s : State} (L : Lay D rbs wbs s) {f : Ptr} {B : Nat} (hB : B < 2 ^ 32) (hc : normChk (rbs ++ wbs) f = true) (hr : Reduced s.mem (pa s f)) : - WP isa (callP "vg_mldsa_norm_lt" P.normLt [.ptr f, .imm B]) s fun s' => PPostB D s s' [] ∧ + WP isa (callP nm P.normLt [.ptr f, .imm B]) s fun s' => PPostB D s s' [] ∧ (∀ r ∈ calleeSaved, s'.gpr r = s.gpr r) ∧ (s'.gpr .rax).setWidth 32 = if normRq [polyAt s.mem (pa s f)] < B then 1 else 0 := by have i1 : inB (rbs ++ wbs) f 1024 = true := by @@ -179,10 +179,10 @@ theorem normCall_ok {P : Prims} (hP : PrimsOk P D) {s : State} (L : Lay D rbs wb simp only [e1, e2, Arg.val, er, A.poly' i1 hD, sw32_ofNat hB] at hq exact hq -theorem normCall_tr {P : Prims} (hP : PrimsOk P D) {f : Ptr} {B : Nat} (hB : B < 2 ^ 32) +theorem normCall_tr {nm : String} {P : Prims} (hP : PrimsOk P D) {f : Ptr} {B : Nat} (hB : B < 2 ^ 32) (hc : normChk (rbs ++ wbs) f = true) : RelCT isa (fun x y => LRel D rbs wbs x y ∧ Reduced x.mem (pa x f) ∧ Reduced y.mem (pa y f)) - (callP "vg_mldsa_norm_lt" P.normLt [.ptr f, .imm B]) fun _ _ => True := by + (callP nm P.normLt [.ptr f, .imm B]) fun _ _ => True := by have i1 : inB (rbs ++ wbs) f 1024 = true := by simp only [normChk, Bool.and_eq_true] at hc; exact hc.1.1 refine callP_tr hP.normLt.ver.1 hP.normLt.ver.2.1 (normArgs_ok hB hc) @@ -238,10 +238,10 @@ theorem hintPre {S : Nat} (hS : S + 8 ≤ D) {s s1 : State} (A : At D rbs wbs s simp only [List.mem_cons, List.mem_nil_iff, or_false, forall_eq_or_imp, forall_eq] exact ⟨A.stk hS i1, A.stk hS i2, A.stk hS i3⟩ -theorem hintCall_ok {P : Prims} (hP : PrimsOk P D) {s : State} (L : Lay D rbs wbs s) {z r h : Ptr} {γ : Nat} +theorem hintCall_ok {nm : String} {P : Prims} (hP : PrimsOk P D) {s : State} (L : Lay D rbs wbs s) {z r h : Ptr} {γ : Nat} (hγ : γ ∈ gamma2s) (hc : hintChk (rbs ++ wbs) wbs z r h = true) (rz : Reduced s.mem (pa s z)) (rr : Reduced s.mem (pa s r)) : - WP isa (callP "vg_mldsa_make_hint" P.makeHint [.ptr z, .ptr r, .imm γ, .ptr h]) s fun s' => + WP isa (callP nm P.makeHint [.ptr z, .ptr r, .imm γ, .ptr h]) s fun s' => PPostB D s s' [(h, 1024)] ∧ (∀ r ∈ calleeSaved, s'.gpr r = s.gpr r) ∧ HintIs s'.mem (pa s h) 1 [Vector.zipWith (makeHint γ) (polyAt s.mem (pa s z)) (polyAt s.mem (pa s r))] ∧ ((s'.gpr .rax).setWidth 32).toNat = @@ -259,11 +259,11 @@ theorem hintCall_ok {P : Prims} (hP : PrimsOk P D) {s : State} (L : Lay D rbs wb simp only [e1, e2, e3, e4, Arg.val, hm₂, er, A.poly' i1 hD, A.poly' i2 hD, sw32_ofNat (gamma2_lt hγ)] at hq exact hq -theorem hintCall_tr {P : Prims} (hP : PrimsOk P D) {z r h : Ptr} {γ : Nat} (hγ : γ ∈ gamma2s) +theorem hintCall_tr {nm : String} {P : Prims} (hP : PrimsOk P D) {z r h : Ptr} {γ : Nat} (hγ : γ ∈ gamma2s) (hc : hintChk (rbs ++ wbs) wbs z r h = true) : RelCT isa (fun x y => LRel D rbs wbs x y ∧ (Reduced x.mem (pa x z) ∧ Reduced x.mem (pa x r)) ∧ (Reduced y.mem (pa y z) ∧ Reduced y.mem (pa y r))) - (callP "vg_mldsa_make_hint" P.makeHint [.ptr z, .ptr r, .imm γ, .ptr h]) fun _ _ => True := by + (callP nm P.makeHint [.ptr z, .ptr r, .imm γ, .ptr h]) fun _ _ => True := by obtain ⟨w1, i1, i2, i3, _⟩ := hintChk_spec hc refine callP_tr hP.makeHint.ver.1 hP.makeHint.ver.2.1 (hintArgs_ok hγ hc) fun x y x1 y1 ⟨R, ⟨rzx, rrx⟩, ⟨rzy, rry⟩⟩ ⟨⟨hAx, hmx⟩, kx⟩ ⟨⟨hAy, hmy⟩, ky⟩ => diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean index 3b2dc921e..8d11355f3 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean @@ -76,7 +76,8 @@ with no implementation of it. -/ theorem sign_same {m mc : Prog isa → Bool} (hm : Comp m mc) {B : Backend} (h1 : mc B.ntt = true) (h2 : mc B.invNtt = true) (h3 : mc B.mul = true) (h4 : mc B.mulAdd = true) (h5 : mc B.add = true) - (h6 : mc B.sub = true) (p : Params) : + (h6 : mc B.sub = true) (h8 : mc B.highBits = true) (h9 : mc B.lowBits = true) + (h10 : mc B.normLt = true) (h11 : mc B.makeHint = true) (p : Params) : Same m (Impl.MlDsa.X86_64.Sign.sign (primsWith B) p) (Impl.MlDsa.X86_64.Sign.sign (primsWith .empty) p) := by unfold Impl.MlDsa.X86_64.Sign.sign same_tac hm @@ -93,13 +94,15 @@ include h3 theorem sign_ctl : ctlOk (Impl.MlDsa.X86_64.Sign.sign (primsWith v.code) p) = true := ctlOk_of_ctlC (Same.ok (sign_same Comp.ctlC v.ok.ntt.ctl v.ok.invNtt.ctl v.ok.mul.ctl v.ok.mulAdd.ctl - v.ok.add.ctl v.ok.sub.ctl p) (sign0_ctlC h3)) + v.ok.add.ctl v.ok.sub.ctl v.ok.highBits.ctl v.ok.lowBits.ctl v.ok.normLt.ctl v.ok.makeHint.ctl p) (sign0_ctlC h3)) theorem sign_spSafe : (Impl.MlDsa.X86_64.Sign.sign (primsWith v.code) p).all (fun i => !isa.writesSp i) = true := Code.all_of_allInstrs (Same.ok (sign_same (Comp.all _) (Code.allInstrs_of_all v.ok.ntt.sp) (Code.allInstrs_of_all v.ok.invNtt.sp) (Code.allInstrs_of_all v.ok.mul.sp) (Code.allInstrs_of_all v.ok.mulAdd.sp) (Code.allInstrs_of_all v.ok.add.sp) - (Code.allInstrs_of_all v.ok.sub.sp) p) (sign0_sp h3)) + (Code.allInstrs_of_all v.ok.sub.sp) + (Code.allInstrs_of_all v.ok.highBits.sp) (Code.allInstrs_of_all v.ok.lowBits.sp) + (Code.allInstrs_of_all v.ok.normLt.sp) (Code.allInstrs_of_all v.ok.makeHint.sp) p) (sign0_sp h3)) theorem sign_verified : Verified X86_64.target (Impl.MlDsa.X86_64.Sign.sign (primsWith v.code) p) (signContractT p X86_64.abi signStack) := diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean index ef0be7ae9..744298288 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean @@ -54,6 +54,7 @@ def primsWith (B : Arith.Backend) : Prims := mul := B.mul mulAdd := B.mulAdd sub := B.sub + normLt := B.normLt rej4 := B.rej4 sfx := B.sfx } @@ -104,10 +105,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.hintUnpack _) - normLt := (CalleeOk.of_verified Proof.MlDsa.X86_64.Round.normLt_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.normLt _) + normLt := calleeOf v.ok.normLt rej4 := ⟨v.ok.rej4.ver.1, v.ok.rej4.ver.2.1, v.ok.rej4.nosp, v.ok.rej4.depth, v.ok.rej4.ctl, v.ok.rej4.sp⟩ /-- `vg_mldsa*_verify` for the parameter set `p`, calling the x86-64 diff --git a/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Avx2.lean b/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Avx2.lean index 5b74ee5fe..5f57b2be7 100644 --- a/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Avx2.lean +++ b/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Avx2.lean @@ -6,8 +6,10 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.BackendAvx2 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` and `vg_mldsa_sub_avx2`, on eight coefficients at a time -in AVX2 registers, which need AVX and AVX2; key generation, signing and +`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 +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 51dccc5dd..fd2c251ab 100644 --- a/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Sse2.lean +++ b/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Sse2.lean @@ -5,8 +5,9 @@ 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` and `vg_mldsa_sub`, in the baseline ISA (SSE2), which key -generation, signing and verification call. +`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 +baseline ISA (SSE2), which key generation, signing and verification call. -/ namespace VG.Variants.MlDsaArith.X86_64.Sse2 diff --git a/src/asm/x86_64/mldsa.rs b/src/asm/x86_64/mldsa.rs index 2997f10a4..41484c2ca 100644 --- a/src/asm/x86_64/mldsa.rs +++ b/src/asm/x86_64/mldsa.rs @@ -4588,6 +4588,230 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_low_bits(r: *const [u32; 256], gam ) } +/// The CPU features `vg_mldsa_high_bits_avx2` requires (`Artifact.features`). +pub(crate) const VG_MLDSA_HIGH_BITS_AVX2_FEATURES: &[&str] = &["avx", "avx2"]; + +/// `HighBits` (FIPS 204 Algorithm 37) of each coefficient of `*r`, with `gamma2` = `γ₂`: writes the `r1`s to `*out`. +/// +/// Contract: `VG.Spec.MlDsa.highBitsContract`. 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 +/// +/// * `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 `r` (distinct Rust objects never do). +/// * Neither `r` nor `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_high_bits_avx2(r: *const [u32; 256], gamma2: u32, out: *mut [u32; 256]) { + core::arch::naked_asm!( + "mov esi, esi", + "cmp esi, 261888", + "mov r10, rdx", + "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 [rdi]", + "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", + "vpsubd ymm1, ymm0, ymm10", + "vpsrad ymm1, ymm1, 31", + "vpand ymm0, ymm0, ymm1", + "vmovdqu YMMWORD PTR [r10], ymm0", + "add rdi, 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 [rdi]", + "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", + "vpsubd ymm1, ymm0, ymm10", + "vpsrad ymm1, ymm1, 31", + "vpand ymm0, ymm0, ymm1", + "vmovdqu YMMWORD PTR [r10], ymm0", + "add rdi, 32", + "add r10, 32", + "sub rcx, 1", + "jne 23b", + "21:", + "vzeroupper", + "ret", + ) +} + +/// The CPU features `vg_mldsa_low_bits_avx2` requires (`Artifact.features`). +pub(crate) const VG_MLDSA_LOW_BITS_AVX2_FEATURES: &[&str] = &["avx", "avx2"]; + +/// `LowBits` (FIPS 204 Algorithm 38) of each coefficient of `*r`, with `gamma2` = `γ₂`: writes the `r0`s, modulo `q` = 8380417, to `*out`. +/// +/// Contract: `VG.Spec.MlDsa.lowBitsContract`. 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 +/// +/// * `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 `r` (distinct Rust objects never do). +/// * Neither `r` nor `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_low_bits_avx2(r: *const [u32; 256], gamma2: u32, out: *mut [u32; 256]) { + core::arch::naked_asm!( + "mov esi, esi", + "cmp esi, 261888", + "mov r10, rdx", + "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 [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", + "vpsubd ymm1, ymm0, ymm10", + "vpsrad ymm1, ymm1, 31", + "vpand ymm0, ymm0, ymm1", + "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 ymm3, ymm3, ymm1", + "vpsrad ymm1, ymm3, 31", + "vpand ymm1, ymm1, ymm15", + "vpaddd ymm3, ymm3, ymm1", + "vmovdqu YMMWORD PTR [r10], ymm3", + "add rdi, 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 [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", + "vpsubd ymm1, ymm0, ymm10", + "vpsrad ymm1, ymm1, 31", + "vpand ymm0, ymm0, ymm1", + "vpslld ymm1, ymm0, 19", + "vpslld ymm2, ymm0, 9", + "vpsubd ymm1, ymm1, ymm2", + "vpsubd ymm3, ymm3, ymm1", + "vpsrad ymm1, ymm3, 31", + "vpand ymm1, ymm1, ymm15", + "vpaddd ymm3, ymm3, ymm1", + "vmovdqu YMMWORD PTR [r10], ymm3", + "add rdi, 32", + "add r10, 32", + "sub rcx, 1", + "jne 23b", + "21:", + "vzeroupper", + "ret", + ) +} + /// Returns 1 if the infinity norm of the polynomial `*f` (FIPS 204 §2.3: the largest `|fᵢ mod± q|`) is less than `bound`, and 0 otherwise. /// /// Contract: `VG.Spec.MlDsa.normLtContract`. Constant time: only the pointer and `bound` may affect timing, not the data. @@ -4723,6 +4947,249 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_make_hint(z: *const [u32; 256], r: ) } +/// The CPU features `vg_mldsa_norm_lt_avx2` requires (`Artifact.features`). +pub(crate) const VG_MLDSA_NORM_LT_AVX2_FEATURES: &[&str] = &["avx", "avx2"]; + +/// Returns 1 if the infinity norm of the polynomial `*f` (FIPS 204 §2.3: the largest `|fᵢ mod± q|`) is less than `bound`, and 0 otherwise. +/// +/// Contract: `VG.Spec.MlDsa.normLtContract`. Constant time: only the pointer and `bound` may affect timing, not the data. +/// +/// The function compares eight coefficients at a time in AVX2 registers; it needs AVX and AVX2. +/// +/// # Safety +/// +/// * `f` must be valid for reads of 1024 bytes. +/// * Each of the 256 `u32`s of `f` must be less than `q` = 8380417. +/// * `f` must not 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_norm_lt_avx2(f: *const [u32; 256], bound: u32) -> u32 { + core::arch::naked_asm!( + "mov esi, esi", + "cmp esi, 8380417", + "jb 20f", + "mov esi, 8380417", + "jmp 21f", + "20:", + "21:", + "mov eax, esi", + "vmovq xmm8, rax", + "vpbroadcastd ymm8, xmm8", + "mov eax, 8380417", + "sub eax, esi", + "vmovq xmm9, rax", + "vpbroadcastd ymm9, xmm9", + "mov eax, -1", + "vmovq xmm10, rax", + "vpbroadcastd ymm10, xmm10", + "mov ecx, 32", + "22:", + "vmovdqu ymm0, YMMWORD PTR [rdi]", + "vpsubd ymm1, ymm0, ymm8", + "vpsubd ymm2, ymm9, ymm0", + "vpor ymm1, ymm1, ymm2", + "vpand ymm10, ymm10, ymm1", + "add rdi, 32", + "sub rcx, 1", + "jne 22b", + "vpsrad ymm10, ymm10, 31", + "vpmovmskb eax, ymm10", + "vzeroupper", + "add rax, 1", + "shr rax, 32", + "ret", + ) +} + +/// The CPU features `vg_mldsa_make_hint_avx2` requires (`Artifact.features`). +pub(crate) const VG_MLDSA_MAKE_HINT_AVX2_FEATURES: &[&str] = &["avx", "avx2"]; + +/// `MakeHint` (FIPS 204 Algorithm 39) of each pair of coefficients of `*z` and `*r`, with `gamma2` = `γ₂`: writes 1 for true and 0 for false to `*h`, and returns the number of 1s. +/// +/// Contract: `VG.Spec.MlDsa.makeHintContract`. Constant time: only the pointers and `gamma2` may affect timing, not the data. +/// +/// The function computes eight hints at a time in AVX2 registers, multiplying by shifts and additions; it needs AVX and AVX2. +/// +/// # Safety +/// +/// * `z` must be valid for reads of 1024 bytes. +/// * `r` must be valid for reads of 1024 bytes. +/// * `h` 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 `z` must be less than `q` = 8380417. +/// * Each of the 256 `u32`s of `r` must be less than `q` = 8380417. +/// * `h` must not overlap `z` or `r` (distinct Rust objects never do). +/// * None of `z`, `r` and `h` 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_make_hint_avx2(z: *const [u32; 256], r: *const [u32; 256], gamma2: u32, h: *mut [u32; 256]) -> u32 { + core::arch::naked_asm!( + "mov edx, edx", + "cmp edx, 261888", + "mov r10, rcx", + "mov r9d, 0", + "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 eax, 63", + "vmovq xmm11, rax", + "vpbroadcastd ymm11, xmm11", + "mov ecx, 32", + "22:", + "vmovdqu ymm0, YMMWORD PTR [rsi]", + "vmovdqu ymm5, YMMWORD PTR [rdi]", + "vmovdqa ymm4, 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", + "vpsubd ymm1, ymm0, ymm10", + "vpsrad ymm1, ymm1, 31", + "vpand ymm0, ymm0, ymm1", + "vmovdqa ymm3, ymm0", + "vpaddd ymm0, ymm4, ymm5", + "vpsubd ymm0, ymm0, ymm15", + "vpsrad ymm1, ymm0, 31", + "vpand ymm1, ymm1, ymm15", + "vpaddd ymm0, ymm0, ymm1", + "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", + "vpsubd ymm1, ymm0, ymm10", + "vpsrad ymm1, ymm1, 31", + "vpand ymm0, ymm0, ymm1", + "vpxor ymm0, ymm0, ymm3", + "vpaddd ymm0, ymm0, ymm11", + "vpsrld ymm0, ymm0, 6", + "vpslld ymm1, ymm0, 7", + "vmovdqu YMMWORD PTR [r10], ymm0", + "vpmovmskb eax, ymm1", + "mov edx, eax", + "shr edx, 4", + "add eax, edx", + "mov edx, eax", + "shr edx, 8", + "add eax, edx", + "mov edx, eax", + "shr edx, 16", + "add eax, edx", + "and eax, 15", + "add r9, rax", + "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 eax, 63", + "vmovq xmm11, rax", + "vpbroadcastd ymm11, xmm11", + "mov ecx, 32", + "23:", + "vmovdqu ymm0, YMMWORD PTR [rsi]", + "vmovdqu ymm5, YMMWORD PTR [rdi]", + "vmovdqa ymm4, 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", + "vpsubd ymm1, ymm0, ymm10", + "vpsrad ymm1, ymm1, 31", + "vpand ymm0, ymm0, ymm1", + "vmovdqa ymm3, ymm0", + "vpaddd ymm0, ymm4, ymm5", + "vpsubd ymm0, ymm0, ymm15", + "vpsrad ymm1, ymm0, 31", + "vpand ymm1, ymm1, ymm15", + "vpaddd ymm0, ymm0, ymm1", + "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", + "vpsubd ymm1, ymm0, ymm10", + "vpsrad ymm1, ymm1, 31", + "vpand ymm0, ymm0, ymm1", + "vpxor ymm0, ymm0, ymm3", + "vpaddd ymm0, ymm0, ymm11", + "vpsrld ymm0, ymm0, 6", + "vpslld ymm1, ymm0, 7", + "vmovdqu YMMWORD PTR [r10], ymm0", + "vpmovmskb eax, ymm1", + "mov edx, eax", + "shr edx, 4", + "add eax, edx", + "mov edx, eax", + "shr edx, 8", + "add eax, edx", + "mov edx, eax", + "shr edx, 16", + "add eax, edx", + "and eax, 15", + "add r9, rax", + "add rdi, 32", + "add rsi, 32", + "add r10, 32", + "sub rcx, 1", + "jne 23b", + "21:", + "mov rax, r9", + "vzeroupper", + "ret", + ) +} + /// `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. diff --git a/src/asm/x86_64/mldsa44.rs b/src/asm/x86_64/mldsa44.rs index b51c4cdf8..b8c0ab667 100644 --- a/src/asm/x86_64/mldsa44.rs +++ b/src/asm/x86_64/mldsa44.rs @@ -1685,7 +1685,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov esi, 95232", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 43", @@ -1698,7 +1698,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov esi, 95232", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 43", @@ -1711,7 +1711,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov esi, 95232", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 43", @@ -1724,7 +1724,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov esi, 95232", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 43", @@ -1840,7 +1840,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov rdi, rbx", "add rdi, 14336", "mov esi, 130994", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -1862,7 +1862,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov rdi, rbx", "add rdi, 15360", "mov esi, 130994", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -1884,7 +1884,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov rdi, rbx", "add rdi, 16384", "mov esi, 130994", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -1906,7 +1906,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov rdi, rbx", "add rdi, 17408", "mov esi, 130994", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -1930,11 +1930,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov esi, 95232", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 95154", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -1958,11 +1958,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov esi, 95232", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 95154", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -1986,11 +1986,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov esi, 95232", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 95154", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -2014,11 +2014,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov esi, 95232", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 95154", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 8192", @@ -2035,7 +2035,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov rdi, rbx", "add rdi, 8192", "mov esi, 95232", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -2066,7 +2066,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov edx, 95232", "mov rcx, rbx", "add rcx, 10240", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -2085,7 +2085,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov rdi, rbx", "add rdi, 8192", "mov esi, 95232", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -2116,7 +2116,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov edx, 95232", "mov rcx, rbx", "add rcx, 11264", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -2135,7 +2135,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov rdi, rbx", "add rdi, 8192", "mov esi, 95232", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -2166,7 +2166,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov edx, 95232", "mov rcx, rbx", "add rcx, 12288", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -2185,7 +2185,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov rdi, rbx", "add rdi, 8192", "mov esi, 95232", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -2216,7 +2216,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], "mov edx, 95232", "mov rcx, rbx", "add rcx, 13312", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -2315,14 +2315,14 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560], vg_mldsa_multiply_ntt_avx2 = sym super::mldsa::vg_mldsa_multiply_ntt_avx2, vg_mldsa_multiply_add_ntt_avx2 = sym super::mldsa::vg_mldsa_multiply_add_ntt_avx2, vg_mldsa_inv_ntt_avx2 = sym super::mldsa::vg_mldsa_inv_ntt_avx2, - vg_mldsa_high_bits = sym super::mldsa::vg_mldsa_high_bits, + vg_mldsa_high_bits_avx2 = sym super::mldsa::vg_mldsa_high_bits_avx2, vg_mldsa_simple_bit_pack = sym super::mldsa::vg_mldsa_simple_bit_pack, vg_mldsa_sample_in_ball = sym super::mldsa::vg_mldsa_sample_in_ball, vg_mldsa_add_avx2 = sym super::mldsa::vg_mldsa_add_avx2, - vg_mldsa_norm_lt = sym super::mldsa::vg_mldsa_norm_lt, + vg_mldsa_norm_lt_avx2 = sym super::mldsa::vg_mldsa_norm_lt_avx2, vg_mldsa_sub_avx2 = sym super::mldsa::vg_mldsa_sub_avx2, - vg_mldsa_low_bits = sym super::mldsa::vg_mldsa_low_bits, - vg_mldsa_make_hint = sym super::mldsa::vg_mldsa_make_hint, + vg_mldsa_low_bits_avx2 = sym super::mldsa::vg_mldsa_low_bits_avx2, + vg_mldsa_make_hint_avx2 = sym super::mldsa::vg_mldsa_make_hint_avx2, vg_mldsa_bit_pack = sym super::mldsa::vg_mldsa_bit_pack, vg_mldsa_hint_bit_pack = sym super::mldsa::vg_mldsa_hint_bit_pack, ) @@ -2385,7 +2385,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312 "mov rdi, rbx", "add rdi, 16384", "mov esi, 130994", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, r13", "add rdi, 608", @@ -2398,7 +2398,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312 "mov rdi, rbx", "add rdi, 17408", "mov esi, 130994", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, r13", "add rdi, 1184", @@ -2411,7 +2411,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312 "mov rdi, rbx", "add rdi, 18432", "mov esi, 130994", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, r13", "add rdi, 1760", @@ -2424,7 +2424,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312 "mov rdi, rbx", "add rdi, 19456", "mov esi, 130994", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "test r15d, r15d", "jne 22f", @@ -3055,7 +3055,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312 "ret", vg_mldsa_hint_bit_unpack = sym super::mldsa::vg_mldsa_hint_bit_unpack, vg_mldsa_bit_unpack = sym super::mldsa::vg_mldsa_bit_unpack, - vg_mldsa_norm_lt = sym super::mldsa::vg_mldsa_norm_lt, + vg_mldsa_norm_lt_avx2 = sym super::mldsa::vg_mldsa_norm_lt_avx2, vg_mldsa_rej_ntt_poly4_avx2 = sym super::mldsa::vg_mldsa_rej_ntt_poly4_avx2, vg_mldsa_sample_in_ball = sym super::mldsa::vg_mldsa_sample_in_ball, vg_mldsa_ntt_avx2 = sym super::mldsa::vg_mldsa_ntt_avx2, diff --git a/src/asm/x86_64/mldsa65.rs b/src/asm/x86_64/mldsa65.rs index 7d80620cf..6feef85c9 100644 --- a/src/asm/x86_64/mldsa65.rs +++ b/src/asm/x86_64/mldsa65.rs @@ -2473,7 +2473,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov esi, 261888", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 15", @@ -2486,7 +2486,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov esi, 261888", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 15", @@ -2499,7 +2499,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov esi, 261888", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 15", @@ -2512,7 +2512,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov esi, 261888", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 15", @@ -2525,7 +2525,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov esi, 261888", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 15", @@ -2538,7 +2538,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov esi, 261888", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 15", @@ -2654,7 +2654,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov rdi, rbx", "add rdi, 16384", "mov esi, 524092", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -2676,7 +2676,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov rdi, rbx", "add rdi, 17408", "mov esi, 524092", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -2698,7 +2698,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov rdi, rbx", "add rdi, 18432", "mov esi, 524092", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -2720,7 +2720,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov rdi, rbx", "add rdi, 19456", "mov esi, 524092", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -2742,7 +2742,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov rdi, rbx", "add rdi, 20480", "mov esi, 524092", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -2766,11 +2766,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov esi, 261888", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 261692", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -2794,11 +2794,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov esi, 261888", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 261692", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -2822,11 +2822,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov esi, 261888", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 261692", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -2850,11 +2850,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov esi, 261888", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 261692", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -2878,11 +2878,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov esi, 261888", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 261692", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -2906,11 +2906,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov esi, 261888", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 261692", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 8192", @@ -2927,7 +2927,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov rdi, rbx", "add rdi, 8192", "mov esi, 261888", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -2958,7 +2958,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov edx, 261888", "mov rcx, rbx", "add rcx, 10240", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -2977,7 +2977,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov rdi, rbx", "add rdi, 8192", "mov esi, 261888", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -3008,7 +3008,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov edx, 261888", "mov rcx, rbx", "add rcx, 11264", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -3027,7 +3027,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov rdi, rbx", "add rdi, 8192", "mov esi, 261888", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -3058,7 +3058,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov edx, 261888", "mov rcx, rbx", "add rcx, 12288", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -3077,7 +3077,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov rdi, rbx", "add rdi, 8192", "mov esi, 261888", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -3108,7 +3108,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov edx, 261888", "mov rcx, rbx", "add rcx, 13312", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -3127,7 +3127,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov rdi, rbx", "add rdi, 8192", "mov esi, 261888", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -3158,7 +3158,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov edx, 261888", "mov rcx, rbx", "add rcx, 14336", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -3177,7 +3177,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov rdi, rbx", "add rdi, 8192", "mov esi, 261888", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -3208,7 +3208,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], "mov edx, 261888", "mov rcx, rbx", "add rcx, 15360", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -3315,14 +3315,14 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032], vg_mldsa_multiply_ntt_avx2 = sym super::mldsa::vg_mldsa_multiply_ntt_avx2, vg_mldsa_multiply_add_ntt_avx2 = sym super::mldsa::vg_mldsa_multiply_add_ntt_avx2, vg_mldsa_inv_ntt_avx2 = sym super::mldsa::vg_mldsa_inv_ntt_avx2, - vg_mldsa_high_bits = sym super::mldsa::vg_mldsa_high_bits, + vg_mldsa_high_bits_avx2 = sym super::mldsa::vg_mldsa_high_bits_avx2, vg_mldsa_simple_bit_pack = sym super::mldsa::vg_mldsa_simple_bit_pack, vg_mldsa_sample_in_ball = sym super::mldsa::vg_mldsa_sample_in_ball, vg_mldsa_add_avx2 = sym super::mldsa::vg_mldsa_add_avx2, - vg_mldsa_norm_lt = sym super::mldsa::vg_mldsa_norm_lt, + vg_mldsa_norm_lt_avx2 = sym super::mldsa::vg_mldsa_norm_lt_avx2, vg_mldsa_sub_avx2 = sym super::mldsa::vg_mldsa_sub_avx2, - vg_mldsa_low_bits = sym super::mldsa::vg_mldsa_low_bits, - vg_mldsa_make_hint = sym super::mldsa::vg_mldsa_make_hint, + vg_mldsa_low_bits_avx2 = sym super::mldsa::vg_mldsa_low_bits_avx2, + vg_mldsa_make_hint_avx2 = sym super::mldsa::vg_mldsa_make_hint_avx2, vg_mldsa_bit_pack = sym super::mldsa::vg_mldsa_bit_pack, vg_mldsa_hint_bit_pack = sym super::mldsa::vg_mldsa_hint_bit_pack, ) @@ -3385,7 +3385,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952 "mov rdi, rbx", "add rdi, 16384", "mov esi, 524092", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, r13", "add rdi, 688", @@ -3398,7 +3398,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952 "mov rdi, rbx", "add rdi, 17408", "mov esi, 524092", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, r13", "add rdi, 1328", @@ -3411,7 +3411,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952 "mov rdi, rbx", "add rdi, 18432", "mov esi, 524092", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, r13", "add rdi, 1968", @@ -3424,7 +3424,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952 "mov rdi, rbx", "add rdi, 19456", "mov esi, 524092", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, r13", "add rdi, 2608", @@ -3437,7 +3437,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952 "mov rdi, rbx", "add rdi, 20480", "mov esi, 524092", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "test r15d, r15d", "jne 22f", @@ -4471,7 +4471,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952 "ret", vg_mldsa_hint_bit_unpack = sym super::mldsa::vg_mldsa_hint_bit_unpack, vg_mldsa_bit_unpack = sym super::mldsa::vg_mldsa_bit_unpack, - vg_mldsa_norm_lt = sym super::mldsa::vg_mldsa_norm_lt, + vg_mldsa_norm_lt_avx2 = sym super::mldsa::vg_mldsa_norm_lt_avx2, vg_mldsa_rej_ntt_poly4_avx2 = sym super::mldsa::vg_mldsa_rej_ntt_poly4_avx2, vg_mldsa_rej_ntt_poly = sym super::mldsa::vg_mldsa_rej_ntt_poly, vg_mldsa_sample_in_ball = sym super::mldsa::vg_mldsa_sample_in_ball, diff --git a/src/asm/x86_64/mldsa87.rs b/src/asm/x86_64/mldsa87.rs index f8dfeb95d..8e96a3ffe 100644 --- a/src/asm/x86_64/mldsa87.rs +++ b/src/asm/x86_64/mldsa87.rs @@ -3698,7 +3698,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 15", @@ -3711,7 +3711,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 15", @@ -3724,7 +3724,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 15", @@ -3737,7 +3737,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 15", @@ -3750,7 +3750,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 15", @@ -3763,7 +3763,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 15", @@ -3776,7 +3776,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 15", @@ -3789,7 +3789,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 6144", - "call {vg_mldsa_high_bits}", + "call {vg_mldsa_high_bits_avx2}", "mov rdi, rbx", "add rdi, 6144", "mov esi, 15", @@ -3905,7 +3905,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 18432", "mov esi, 524168", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -3927,7 +3927,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 19456", "mov esi, 524168", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -3949,7 +3949,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 20480", "mov esi, 524168", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -3971,7 +3971,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 21504", "mov esi, 524168", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -3993,7 +3993,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 22528", "mov esi, 524168", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -4015,7 +4015,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 23552", "mov esi, 524168", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -4037,7 +4037,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 24576", "mov esi, 524168", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -4061,11 +4061,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 261768", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -4089,11 +4089,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 261768", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -4117,11 +4117,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 261768", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -4145,11 +4145,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 261768", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -4173,11 +4173,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 261768", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -4201,11 +4201,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 261768", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -4229,11 +4229,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 261768", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 6144", @@ -4257,11 +4257,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov esi, 261888", "mov rdx, rbx", "add rdx, 7168", - "call {vg_mldsa_low_bits}", + "call {vg_mldsa_low_bits_avx2}", "mov rdi, rbx", "add rdi, 7168", "mov esi, 261768", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 8192", @@ -4278,7 +4278,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 8192", "mov esi, 261888", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -4309,7 +4309,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov edx, 261888", "mov rcx, rbx", "add rcx, 10240", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -4328,7 +4328,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 8192", "mov esi, 261888", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -4359,7 +4359,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov edx, 261888", "mov rcx, rbx", "add rcx, 11264", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -4378,7 +4378,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 8192", "mov esi, 261888", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -4409,7 +4409,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov edx, 261888", "mov rcx, rbx", "add rcx, 12288", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -4428,7 +4428,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 8192", "mov esi, 261888", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -4459,7 +4459,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov edx, 261888", "mov rcx, rbx", "add rcx, 13312", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -4478,7 +4478,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 8192", "mov esi, 261888", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -4509,7 +4509,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov edx, 261888", "mov rcx, rbx", "add rcx, 14336", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -4528,7 +4528,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 8192", "mov esi, 261888", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -4559,7 +4559,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov edx, 261888", "mov rcx, rbx", "add rcx, 15360", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -4578,7 +4578,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 8192", "mov esi, 261888", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -4609,7 +4609,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov edx, 261888", "mov rcx, rbx", "add rcx, 16384", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -4628,7 +4628,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov rdi, rbx", "add rdi, 8192", "mov esi, 261888", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, rbx", "add rdi, 9216", @@ -4659,7 +4659,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], "mov edx, 261888", "mov rcx, rbx", "add rcx, 17408", - "call {vg_mldsa_make_hint}", + "call {vg_mldsa_make_hint_avx2}", "mov ecx, DWORD PTR [rbx+904]", "add ecx, eax", "mov QWORD PTR [rbx+904], rcx", @@ -4782,14 +4782,14 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896], vg_mldsa_multiply_ntt_avx2 = sym super::mldsa::vg_mldsa_multiply_ntt_avx2, vg_mldsa_multiply_add_ntt_avx2 = sym super::mldsa::vg_mldsa_multiply_add_ntt_avx2, vg_mldsa_inv_ntt_avx2 = sym super::mldsa::vg_mldsa_inv_ntt_avx2, - vg_mldsa_high_bits = sym super::mldsa::vg_mldsa_high_bits, + vg_mldsa_high_bits_avx2 = sym super::mldsa::vg_mldsa_high_bits_avx2, vg_mldsa_simple_bit_pack = sym super::mldsa::vg_mldsa_simple_bit_pack, vg_mldsa_sample_in_ball = sym super::mldsa::vg_mldsa_sample_in_ball, vg_mldsa_add_avx2 = sym super::mldsa::vg_mldsa_add_avx2, - vg_mldsa_norm_lt = sym super::mldsa::vg_mldsa_norm_lt, + vg_mldsa_norm_lt_avx2 = sym super::mldsa::vg_mldsa_norm_lt_avx2, vg_mldsa_sub_avx2 = sym super::mldsa::vg_mldsa_sub_avx2, - vg_mldsa_low_bits = sym super::mldsa::vg_mldsa_low_bits, - vg_mldsa_make_hint = sym super::mldsa::vg_mldsa_make_hint, + vg_mldsa_low_bits_avx2 = sym super::mldsa::vg_mldsa_low_bits_avx2, + vg_mldsa_make_hint_avx2 = sym super::mldsa::vg_mldsa_make_hint_avx2, vg_mldsa_bit_pack = sym super::mldsa::vg_mldsa_bit_pack, vg_mldsa_hint_bit_pack = sym super::mldsa::vg_mldsa_hint_bit_pack, ) @@ -4852,7 +4852,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov rdi, rbx", "add rdi, 16384", "mov esi, 524168", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, r13", "add rdi, 704", @@ -4865,7 +4865,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov rdi, rbx", "add rdi, 17408", "mov esi, 524168", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, r13", "add rdi, 1344", @@ -4878,7 +4878,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov rdi, rbx", "add rdi, 18432", "mov esi, 524168", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, r13", "add rdi, 1984", @@ -4891,7 +4891,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov rdi, rbx", "add rdi, 19456", "mov esi, 524168", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, r13", "add rdi, 2624", @@ -4904,7 +4904,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov rdi, rbx", "add rdi, 20480", "mov esi, 524168", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, r13", "add rdi, 3264", @@ -4917,7 +4917,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov rdi, rbx", "add rdi, 21504", "mov esi, 524168", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "mov rdi, r13", "add rdi, 3904", @@ -4930,7 +4930,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "mov rdi, rbx", "add rdi, 22528", "mov esi, 524168", - "call {vg_mldsa_norm_lt}", + "call {vg_mldsa_norm_lt_avx2}", "and r15d, eax", "test r15d, r15d", "jne 22f", @@ -6456,7 +6456,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592 "ret", vg_mldsa_hint_bit_unpack = sym super::mldsa::vg_mldsa_hint_bit_unpack, vg_mldsa_bit_unpack = sym super::mldsa::vg_mldsa_bit_unpack, - vg_mldsa_norm_lt = sym super::mldsa::vg_mldsa_norm_lt, + vg_mldsa_norm_lt_avx2 = sym super::mldsa::vg_mldsa_norm_lt_avx2, vg_mldsa_rej_ntt_poly4_avx2 = sym super::mldsa::vg_mldsa_rej_ntt_poly4_avx2, vg_mldsa_sample_in_ball = sym super::mldsa::vg_mldsa_sample_in_ball, vg_mldsa_ntt_avx2 = sym super::mldsa::vg_mldsa_ntt_avx2,