From c8a898673abe4aefd7f6253bf95e739b94b96c5c Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 13:33:05 +0000 Subject: [PATCH 1/2] =?UTF-8?q?ML-DSA=20on=20x86-64:=20NTT=20and=20NTT?= =?UTF-8?q?=E2=81=BB=C2=B9=20in=20SSE2?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit vg_mldsa_ntt and vg_mldsa_inv_ntt now compute on four coefficients at a time in SSE2 registers, as ML-KEM's x86-64 NTT does: a Montgomery multiplication with pmuludq (the even doublewords, then the odd ones moved down by pshufd), conditional additions of q with psrad masks, and for the layers with len 2 and 1 the coefficients of two or four blocks gathered with punpck{l,h}qdq (and pshufd) and interleaved back. The zetas are a table in Montgomery form that the prologue stores in scratch; the multiplications run inside ML-KEM's withMxcsr, so Intel's MCDT mitigation holds. Proofs: the lanes' arithmetic (VArith), the butterflies on registers (VLanes), the loads and stores of four coefficients and the zetas (VMem), the layers (VLay, VLay21), and the functions (Ntt, NttInv). Signing and verification now check that their code loads MXCSR only to restore it (ctlOk, through a compositional ctlC for verification) rather than never, since their primitives now do. ML-DSA-65 on this machine: sign 1.53 ms -> 0.81 ms, verify 281 us -> 181 us, keygen 276 us -> 248 us. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01Ddof3szoTi7HB8iCsM2MCr --- README.md | 6 +- docs/algorithms/ml-dsa-44.toml | 1 + docs/algorithms/ml-dsa-65.toml | 1 + docs/algorithms/ml-dsa-87.toml | 1 + .../Artifacts/MlDsaArith/X86_64.lean | 10 +- .../Impl/MlDsa/X86_64/Arith/Common.lean | 12 +- .../Impl/MlDsa/X86_64/Arith/Ntt.lean | 177 +- .../Impl/MlDsa/X86_64/Arith/Vec.lean | 88 + .../Proof/Framework/X86_64/Mxcsr.lean | 25 + .../Proof/MlDsa/X86_64/Arith/Ntt.lean | 323 +- .../Proof/MlDsa/X86_64/Arith/NttBfly.lean | 201 -- .../Proof/MlDsa/X86_64/Arith/NttInv.lean | 327 +- .../Proof/MlDsa/X86_64/Arith/NttLoop.lean | 165 - .../Proof/MlDsa/X86_64/Arith/Table.lean | 37 +- .../Proof/MlDsa/X86_64/Arith/VArith.lean | 285 ++ .../Proof/MlDsa/X86_64/Arith/VLanes.lean | 138 + .../Proof/MlDsa/X86_64/Arith/VLay.lean | 322 ++ .../Proof/MlDsa/X86_64/Arith/VLay21.lean | 458 +++ .../Proof/MlDsa/X86_64/Arith/VMem.lean | 104 + .../Proof/MlDsa/X86_64/Sign/Correct.lean | 5 +- .../Proof/MlDsa/X86_64/Sign/Verified.lean | 2 +- .../Proof/MlDsa/X86_64/Verify/Correct.lean | 2 +- .../Proof/MlDsa/X86_64/Verify/Entry.lean | 9 +- .../Proof/MlDsa/X86_64/Verify/Instrs.lean | 85 +- src/asm/x86_64/mldsa.rs | 2920 ++++++++--------- 25 files changed, 3323 insertions(+), 2381 deletions(-) create mode 100644 lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Vec.lean delete mode 100644 lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/NttBfly.lean delete mode 100644 lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/NttLoop.lean create mode 100644 lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VArith.lean create mode 100644 lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VLanes.lean create mode 100644 lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VLay.lean create mode 100644 lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VLay21.lean create mode 100644 lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VMem.lean diff --git a/README.md b/README.md index c4ed49c34..00030ddd4 100644 --- a/README.md +++ b/README.md @@ -774,7 +774,7 @@ yours to keep: ✅ -✅ +✅ SSE2 NTT ✅ SHA extensions @@ -790,7 +790,7 @@ yours to keep: ✅ -✅ +✅ SSE2 NTT ✅ SHA extensions @@ -806,7 +806,7 @@ yours to keep: ✅ -✅ +✅ SSE2 NTT ✅ SHA extensions diff --git a/docs/algorithms/ml-dsa-44.toml b/docs/algorithms/ml-dsa-44.toml index b87af90f8..7f316830d 100644 --- a/docs/algorithms/ml-dsa-44.toml +++ b/docs/algorithms/ml-dsa-44.toml @@ -3,3 +3,4 @@ family = "Signatures" specs = ["MlDsa"] modules = ["src/mldsa44.rs"] asm = ["mldsa44", "mldsa"] +optimized = { x86_64 = "SSE2 NTT" } diff --git a/docs/algorithms/ml-dsa-65.toml b/docs/algorithms/ml-dsa-65.toml index 1220edd86..26de285e2 100644 --- a/docs/algorithms/ml-dsa-65.toml +++ b/docs/algorithms/ml-dsa-65.toml @@ -3,3 +3,4 @@ family = "Signatures" specs = ["MlDsa"] modules = ["src/mldsa65.rs"] asm = ["mldsa65", "mldsa"] +optimized = { x86_64 = "SSE2 NTT" } diff --git a/docs/algorithms/ml-dsa-87.toml b/docs/algorithms/ml-dsa-87.toml index d463e224c..1f76a641b 100644 --- a/docs/algorithms/ml-dsa-87.toml +++ b/docs/algorithms/ml-dsa-87.toml @@ -3,3 +3,4 @@ family = "Signatures" specs = ["MlDsa"] modules = ["src/mldsa87.rs"] asm = ["mldsa87", "mldsa"] +optimized = { x86_64 = "SSE2 NTT" } diff --git a/lean/VerifiedGarbage/Artifacts/MlDsaArith/X86_64.lean b/lean/VerifiedGarbage/Artifacts/MlDsaArith/X86_64.lean index 8cd4766b7..f8fdd427a 100644 --- a/lean/VerifiedGarbage/Artifacts/MlDsaArith/X86_64.lean +++ b/lean/VerifiedGarbage/Artifacts/MlDsaArith/X86_64.lean @@ -22,7 +22,10 @@ def artifacts : List Artifact := [ { Spec.MlDsa.nttApi with target := X86_64.target doc := Spec.MlDsa.nttApi.doc - (notes := ["The function stores a table of the 256 zetas in `scratch`."]) + (notes := ["The function computes on four coefficients at a time in SSE2 registers, with a table of \ + the 256 zetas that it stores in `scratch`. It sets MXCSR to `0x1FBF` around its multiplications \ + (Intel's mitigation of MXCSR-configuration-dependent timing) and loads the caller's MXCSR back \ + before returning."]) code := Impl.MlDsa.X86_64.Arith.ntt contract := Spec.MlDsa.nttContract X86_64.abi verified := Proof.MlDsa.X86_64.Arith.ntt_verified @@ -31,7 +34,10 @@ def artifacts : List Artifact := [ { Spec.MlDsa.nttInvApi with target := X86_64.target doc := Spec.MlDsa.nttInvApi.doc - (notes := ["The function stores a table of the 256 negated zetas in `scratch`."]) + (notes := ["The function computes on four coefficients at a time in SSE2 registers, with a table of \ + the 256 zetas that it stores in `scratch`. It sets MXCSR to `0x1FBF` around its multiplications \ + (Intel's mitigation of MXCSR-configuration-dependent timing) and loads the caller's MXCSR back \ + before returning."]) code := Impl.MlDsa.X86_64.Arith.nttInv contract := Spec.MlDsa.nttInvContract X86_64.abi verified := Proof.MlDsa.X86_64.Arith.nttInv_verified diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Common.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Common.lean index e808dbeac..ca4ccf001 100644 --- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Common.lean +++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Common.lean @@ -16,10 +16,7 @@ Pieces of code that the ML-DSA arithmetic functions share, for `⌊rax / q⌋ - 1`, so `rax` less that quotient times `q` (a second `mul`) is less than `2q`, and `csubQ` reduces it. `mul` is the only multiplication of the model, and its timing does not depend on its operands (it is on - Intel's DOIT list). It uses `rax`, `rdx` and `r11`; -* `storeTab t n`: the table `t 0, …, t (n - 1)` of constants stored as - `u32`s at `r9` (in the working space: the code has no other memory), with - immediates. It uses `rax`. + Intel's DOIT list). It uses `rax`, `rdx` and `r11`. -/ namespace VG.Impl.MlDsa.X86_64.Arith @@ -47,11 +44,4 @@ def reduce : List Instr := [.mov .r10 (.reg .rax), .movImm64 .r11 barrettImm, .mul .r11, .mov .rax (.reg .rdx), .mov .r11 (.imm qImm), .mul .r11, .alu .sub .r10 (.reg .rax)] ++ csubQ .r10 .r11 -/-- `t i` to `[r9 + 4i]`. -/ -def tabStep (t : Nat → Nat) (i : Nat) : List Instr := - [.mov32 .rax (.imm (BitVec.ofNat 32 (t i))), .store32 (at_ .r9 (4 * i)) .rax] - -/-- The table `t 0, …, t (n - 1)` at `r9`. -/ -def storeTab (t : Nat → Nat) (n : Nat) : List Instr := (List.range n).flatMap (tabStep t) - end VG.Impl.MlDsa.X86_64.Arith diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Ntt.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Ntt.lean index 6c0808ddf..97f85c3ad 100644 --- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Ntt.lean +++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Ntt.lean @@ -1,97 +1,104 @@ -import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Common +import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Vec import VerifiedGarbage.Spec.MlDsa /-! # ML-DSA on x86-64: `vg_mldsa_ntt` and `vg_mldsa_inv_ntt` -`ntt(f = rdi, scratch = rsi)` and `nttInv(f = rdi, scratch = rsi)`: the -prologue stores a table of 256 zetas to `scratch` as `u32`s (`storeTab`, -with `r9` = `scratch`): `ζ^BitRev8(m) mod q` for `NTT`, and its negation -`-ζ^BitRev8(m) mod q` for `NTT⁻¹` (the `z` of Algorithm 42). Then each of -the eight layers runs its blocks, with `rsi` pointing at coefficient `j` of -`f`, `r8` at the zeta of the block, `rdi` counting the blocks down and `rcx` -the butterflies of a block; the zeta of the block is in `r9`. - -* `NTT` (Algorithm 41): the layers with `len` = 128, 64, …, 1, whose zetas - are consecutive, from `m = 1` up. A butterfly computes - `t = ζ · f[j + len] mod q` (with `reduce`, a Barrett reduction with - `mul`), and stores `f[j] - t` (`f[j] + q - t`, reduced with `csubQ`) to - `f[j + len]` and `f[j] + t` (reduced) to `f[j]`. -* `NTT⁻¹` (Algorithm 42): the layers with `len` = 1, 2, …, 128, whose zetas - are consecutive from `m = 255` down. A butterfly stores `f[j] + f[j + len]` - (reduced) to `f[j]` and `z · (f[j] - f[j + len]) mod q` to `f[j + len]`. - Then every coefficient is multiplied by `8347681 = 256⁻¹ mod q` and - reduced. - -A block ends with `rsi` advanced past its upper half, so a layer ends with -`rsi` at `f + 1024`, and moves it back. Every address and branch depends -only on the pointers. +`ntt(f = rdi, scratch = rsi)` and `nttInv(f = rdi, scratch = rsi)` compute +on four coefficients of `f` at a time, in place, as doublewords of SSE +registers (see `Vec.lean`). `scratch` holds the 256 `u32`s +`ζ^BitRev8(m) · 2³² mod q` (`zmTab`, stored with immediates: the code has +no other memory), and MXCSR's at bytes 768 to 775 before and after +(`withMxcsr`, which saves the caller's MXCSR in `r11` in between). + +Within `withMxcsr`, the prologue stores the table and the constants. Then +the layers, each a pass over `f` with `rdx` pointing at the coefficients it +loads and `r8` at the zetas of the table: + +* `NTT` (Algorithm 41): the layers with `len` = 128, 64, 32, 16, 8 and 4 + (`vlay`) run their blocks (counted in `rax`), each its zeta in the + doublewords of `xmm13` (`vzeta`), and `len / 4` times the butterflies of + four coefficients `w[j]` and of the four `w[j + len]` (`vbfly`, counted + in `rcx`). The layer with `len = 2` (`vlay2`) loads the 8 coefficients of + two blocks, gathers their lower and upper halves into `xmm0` and `xmm1` + (`punpcklqdq`, `punpckhqdq`), with the two zetas in the halves of + `xmm13`; the layer with `len = 1` (`vlay1`) those of four blocks, their + pairs gathered (with `pshufd` first), with the four zetas in the + doublewords of `xmm13`. +* `NTT⁻¹` (Algorithm 42): the same layers in the opposite order, with the + zetas from `m = 255` down and the inverse butterflies (`vibfly`); then + every coefficient is multiplied by `8347681 = 256⁻¹ mod q` (`vscale`, + with `vmont` by `8347681 · 2³² mod q = 16382`). + +Every address and branch depends only on the pointers. -/ namespace VG.Impl.MlDsa.X86_64.Arith open VG.X86_64 - -/-- `ζ^BitRev8(m) mod q`. -/ -def zetaTab (m : Nat) : Nat := 1753 ^ Spec.MlDsa.bitRev8 m % 8380417 - -/-- `-ζ^BitRev8(m) mod q`. -/ -def negZetaTab (m : Nat) : Nat := (8380417 - zetaTab m) % 8380417 - -/-- A butterfly of `NTT` on `[rsi]` and `[rsi + 4len]`, with the zeta in `r9`. -/ -def bfly (len : Nat) : List Instr := - [.mov32 .rax (.mem (at_ .rsi (4 * len))), .mul .r9] ++ reduce ++ - [.mov32 .rax (.mem (at_ .rsi 0)), .mov32 .rdx (.reg .rax), .alu32 .add .rdx (.imm qImm), - .alu32 .sub .rdx (.reg .r10)] ++ csubQ .rdx .r11 ++ - [.store32 (at_ .rsi (4 * len)) .rdx, .alu32 .add .rax (.reg .r10)] ++ csubQ .rax .r11 ++ - [.store32 (at_ .rsi 0) .rax, .alu .add .rsi (.imm 4), .alu .sub .rcx (.imm 1)] - -/-- A butterfly of `NTT⁻¹` on `[rsi]` and `[rsi + 4len]`, with the zeta in `r9`. -/ -def bflyInv (len : Nat) : List Instr := - [.mov32 .rax (.mem (at_ .rsi 0)), .mov32 .r10 (.mem (at_ .rsi (4 * len))), .mov32 .rdx (.reg .rax), - .alu32 .add .rdx (.reg .r10)] ++ csubQ .rdx .r11 ++ - [.store32 (at_ .rsi 0) .rdx, .alu32 .add .rax (.imm qImm), .alu32 .sub .rax (.reg .r10)] ++ - csubQ .rax .r11 ++ [.mul .r9] ++ reduce ++ - [.store32 (at_ .rsi (4 * len)) .r10, .alu .add .rsi (.imm 4), .alu .sub .rcx (.imm 1)] - -/-- A block of `len` butterflies `b`, with the zeta at `r8`, which then moves -by `dz` bytes (4 or -4). -/ -def nttBlk (b : List Instr) (len : Nat) (dz : BitVec 32) : Prog isa := - .seq (.block [.mov32 .r9 (.mem (at_ .r8 0)), .alu .add .r8 (.imm dz), .mov32 .rcx (.imm (BitVec.ofNat 32 len))]) - (.seq (.loop (.block b) .ne) - (.block [.alu .add .rsi (.imm (BitVec.ofNat 32 (4 * len))), .alu .sub .rdi (.imm 1)])) - -/-- A layer: its `128 / len` blocks, then `rsi` back to `f`. -/ -def nttLay (b : List Instr) (len : Nat) (dz : BitVec 32) : Prog isa := - .seq (.block [.mov32 .rdi (.imm (BitVec.ofNat 32 (128 / len)))]) - (.seq (.loop (nttBlk b len dz) .ne) (.block [.alu .sub .rsi (.imm 1024)])) - -/-- The layers of `NTT` with `len` in `lens`. -/ -def nttLays : List Nat → Prog isa - | [] => .block [] - | len :: lens => .seq (nttLay (bfly len) len 4) (nttLays lens) - -/-- The layers of `NTT⁻¹` with `len` in `lens`. -/ -def nttInvLays : List Nat → Prog isa - | [] => .block [] - | len :: lens => .seq (nttLay (bflyInv len) len (-4)) (nttInvLays lens) - -/-- The table `t` to `scratch`, and `rsi` = `f`. -/ -def nttPro (t : Nat → Nat) : List Instr := - [.mov .r9 (.reg .rsi)] ++ storeTab t 256 ++ [.mov .rsi (.reg .rdi), .mov .r8 (.reg .r9)] - -def ntt : Prog isa := - .seq (.block (nttPro zetaTab ++ [.alu .add .r8 (.imm 4)])) (nttLays [128, 64, 32, 16, 8, 4, 2, 1]) - -/-- A coefficient times `8347681`, reduced. -/ -def scaleBody : List Instr := - [.mov32 .rax (.mem (at_ .rsi 0)), .mul .r9] ++ reduce ++ - [.store32 (at_ .rsi 0) .r10, .alu .add .rsi (.imm 4), .alu .sub .rcx (.imm 1)] - -def nttInv : Prog isa := - .seq (.block (nttPro negZetaTab ++ [.alu .add .r8 (.imm (4 * 255))])) - (.seq (nttInvLays [1, 2, 4, 8, 16, 32, 64, 128]) - (.seq (.block [.mov32 .r9 (.imm 8347681)]) - (.seq (.block [.mov32 .rcx (.imm 256)]) (.loop (.block scaleBody) .ne)))) +open VG.Impl.MlKem.X86_64 (xb xmov withMxcsr rcxLoop) + +/-- `ζ^BitRev8(m) · 2³² mod q`. -/ +def zmTab (m : Nat) : Nat := 1753 ^ Spec.MlDsa.bitRev8 m * 2 ^ 32 % 8380417 + +/-- `d ← r + off`. -/ +def leaR (d r : Reg) (off : Nat) : List Instr := [.mov d (.reg r), .alu .add d (.imm (BitVec.ofNat 32 off))] + +/-- The zetas at `[r8]` in the doublewords of `xmm13`, arranged by `pshufd` +with `o`, and its odd doublewords in the even ones of `xmm12`. -/ +def vzeta (o : BitVec 8) : List Instr := + [.movdquLoad .xmm13 (at_ .r8 0), .xop (.pshufd .xmm13 .xmm13 o), .xop (.pshufd .xmm12 .xmm13 0xF5)] + +/-- A layer with `len ≥ 4` and butterflies `bf`: its `128 / len` blocks, the +first with the zeta `k`, the zeta pointer moving by `dz` bytes. -/ +def vlay (bf : List Instr) (len k : Nat) (dz : BitVec 32) : Prog isa := + .seq (.block ([.mov .rdx (.reg .rdi)] ++ leaR .r8 .rsi (4 * k) ++ + [.mov32 .rax (.imm (BitVec.ofNat 32 (128 / len)))])) <| + .loop (.seq (.block (vzeta 0 ++ [.alu .add .r8 (.imm dz)])) + (.seq (rcxLoop (len / 4) ([.movdquLoad .xmm0 (at_ .rdx 0), .movdquLoad .xmm1 (at_ .rdx (4 * len))] ++ + bf ++ [.movdquStore (at_ .rdx 0) .xmm0, .movdquStore (at_ .rdx (4 * len)) .xmm3, + .alu .add .rdx (.imm 16)])) + (.block [.alu .add .rdx (.imm (BitVec.ofNat 32 (4 * len))), .alu .sub .rax (.imm 1)]))) .ne + +/-- The layer with `len = 2`, two blocks at a time: the zetas at `[r8]` +arranged by `pshufd` with `o`, the zeta pointer moving by `dz` bytes. -/ +def vlay2 (bf : List Instr) (k : Nat) (o : BitVec 8) (dz : BitVec 32) : Prog isa := + .seq (.block ([.mov .rdx (.reg .rdi)] ++ leaR .r8 .rsi (4 * k))) <| + rcxLoop 32 ([.movdquLoad .xmm0 (at_ .rdx 0), .movdquLoad .xmm1 (at_ .rdx 16)] ++ vzeta o ++ + [.alu .add .r8 (.imm dz), xmov .xmm2 .xmm0, xb .punpcklqdq .xmm0 .xmm1, xb .punpckhqdq .xmm2 .xmm1, + xmov .xmm1 .xmm2] ++ bf ++ + [xmov .xmm1 .xmm0, xb .punpcklqdq .xmm0 .xmm3, xb .punpckhqdq .xmm1 .xmm3, + .movdquStore (at_ .rdx 0) .xmm0, .movdquStore (at_ .rdx 16) .xmm1, .alu .add .rdx (.imm 32)]) + +/-- The layer with `len = 1`, four blocks at a time: the zetas at `[r8]` +arranged by `pshufd` with `o`, the zeta pointer moving by `dz` bytes. -/ +def vlay1 (bf : List Instr) (k : Nat) (o : BitVec 8) (dz : BitVec 32) : Prog isa := + .seq (.block ([.mov .rdx (.reg .rdi)] ++ leaR .r8 .rsi (4 * k))) <| + rcxLoop 32 ([.movdquLoad .xmm0 (at_ .rdx 0), .movdquLoad .xmm2 (at_ .rdx 16)] ++ vzeta o ++ + [.alu .add .r8 (.imm dz), .xop (.pshufd .xmm0 .xmm0 0xD8), .xop (.pshufd .xmm2 .xmm2 0xD8), + xmov .xmm1 .xmm0, xb .punpcklqdq .xmm0 .xmm2, xb .punpckhqdq .xmm1 .xmm2] ++ bf ++ + [xmov .xmm1 .xmm0, xb .punpckldq .xmm0 .xmm3, xb .punpckhdq .xmm1 .xmm3, + .movdquStore (at_ .rdx 0) .xmm0, .movdquStore (at_ .rdx 16) .xmm1, .alu .add .rdx (.imm 32)]) + +/-- Every coefficient times `8347681 = 256⁻¹ mod q`, reduced. -/ +def vscale : Prog isa := + .seq (.block [.mov .rdx (.reg .rdi), .mov32 .rax (.imm 16382), .xop (.movq .xmm13 .rax), + .xop (.pshufd .xmm13 .xmm13 0), xmov .xmm12 .xmm13]) + (rcxLoop 64 ([.movdquLoad .xmm3 (at_ .rdx 0)] ++ vmont .xmm3 .xmm13 .xmm12 .xmm2 .xmm4 ++ + vcsub .xmm3 .xmm2 ++ [.movdquStore (at_ .rdx 0) .xmm3, .alu .add .rdx (.imm 16)])) + +/-- The table and the constants. -/ +def vpro : List Instr := dwordTab zmTab 256 .rsi ++ vconsts + +def ntt : Prog isa := withMxcsr .rsi 768 <| + .seq (.block vpro) (.seq (vlay vbfly 128 1 4) (.seq (vlay vbfly 64 2 4) (.seq (vlay vbfly 32 4 4) + (.seq (vlay vbfly 16 8 4) (.seq (vlay vbfly 8 16 4) (.seq (vlay vbfly 4 32 4) + (.seq (vlay2 vbfly 64 0x50 8) (vlay1 vbfly 128 0xE4 16)))))))) + +def nttInv : Prog isa := withMxcsr .rsi 768 <| + .seq (.block vpro) (.seq (vlay1 vibfly 252 0x1B (-16)) (.seq (vlay2 vibfly 126 0x05 (-8)) + (.seq (vlay vibfly 4 63 (-4)) (.seq (vlay vibfly 8 31 (-4)) (.seq (vlay vibfly 16 15 (-4)) + (.seq (vlay vibfly 32 7 (-4)) (.seq (vlay vibfly 64 3 (-4)) (.seq (vlay vibfly 128 1 (-4)) + vscale)))))))) end VG.Impl.MlDsa.X86_64.Arith diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Vec.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Vec.lean new file mode 100644 index 000000000..1feaab831 --- /dev/null +++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Vec.lean @@ -0,0 +1,88 @@ +import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Common +import VerifiedGarbage.Impl.MlKem.X86_64.Vec + +/-! +# ML-DSA on x86-64: arithmetic modulo `q` in the doublewords of SSE registers + +The NTT and its inverse compute on four coefficients at a time, as the +doublewords of SSE2 registers, with `q` in the doublewords of `xmm15` and +`-q⁻¹ mod 2³² = 4236238847` in those of `xmm14` (`vconsts`). + +* `vmont d z zo t u`: `d ← d · z · 2⁻³² mod q`, in `[0, 2q)`, for any + doublewords `d` and `z < q` (a Montgomery reduction). `pmuludq` multiplies + the even doublewords of its operands into quadwords, so the even + doublewords of `d` are multiplied by those of `z`, and the odd ones, + moved to the even places of `u` by `pshufd`, by the even doublewords of + `zo`, which hold the odd doublewords of `z` (`pshufd` with `0xF5`). For + each product `P < 2³² · q`, `m = (P mod 2³²) · (-q⁻¹) mod 2³²` makes + `P + m · q` a multiple of `2³²` (with `pmuludq` by `xmm14`, then by + `xmm15`, which use the low doubleword of the product), less than `2⁶⁴`, + and its high doubleword is `(P + m · q) / 2³² < 2q`, congruent to + `P · 2⁻³²` modulo `q`. The quotients of the even doublewords are moved + down to their places by `psrlq`; those of the odd ones are in place, and + the low doublewords of their quadwords are 0, so `por` merges them. +* `vcadd d t`: `d ← d + q` for the doublewords of `d` that are negative + (with `psrad` by 31, a mask), from `(-q, q)` to `[0, q)`; `vcsub d t`: + `d ← d - q`, then `vcadd`, from `[0, 2q)` to `[0, q)`. + +A coefficient `x` is multiplied by `ζ` as `vmont` with `ζ · 2³² mod q`, +which the tables hold. + +`pmuludq` has data-dependent timing on processors with MCDT unless MXCSR is +`0x1FBF` (see `TCB/X86_64/Isa.lean`): the functions run inside ML-KEM's +`withMxcsr`. +-/ + +namespace VG.Impl.MlDsa.X86_64.Arith + +open VG.X86_64 +open VG.Impl.MlKem.X86_64 (xb xmov) + +/-- `q` in the doublewords of `xmm15` and `-q⁻¹ mod 2³² = 4236238847` in +those of `xmm14`, through `rax`. -/ +def vconsts : List Instr := + [.mov32 .rax (.imm 8380417), .xop (.movq .xmm15 .rax), .xop (.pshufd .xmm15 .xmm15 0), + .mov32 .rax (.imm 4236238847), .xop (.movq .xmm14 .rax), .xop (.pshufd .xmm14 .xmm14 0)] + +/-- The Montgomery reductions of the quadword products in `d`, with a +temporary `t`: each quadword becomes `P + m · q`. -/ +def vredc (d t : XReg) : List Instr := + [xmov t d, xb .pmuludq t .xmm14, xb .pmuludq t .xmm15, xb .paddq d t] + +/-- `d ← d · z · 2⁻³² mod q`, in `[0, 2q)`, with the odd doublewords of `z` +in the even doublewords of `zo`, and temporaries `t` and `u`. -/ +def vmont (d z zo t u : XReg) : List Instr := + [.xop (.pshufd u d 0xF5), xb .pmuludq d z, xb .pmuludq u zo] ++ vredc d t ++ + [.xop (.shift .psrlq d 32)] ++ vredc u t ++ [xb .por d u] + +/-- `d ← d + q` for the negative doublewords of `d`, with a temporary `t`. -/ +def vcadd (d t : XReg) : List Instr := + [xmov t d, .xop (.shift .psrad t 31), xb .pand t .xmm15, xb .paddd d t] + +/-- `d ← d mod q` for doublewords in `[0, 2q)`, with a temporary `t`. -/ +def vcsub (d t : XReg) : List Instr := xb .psubd d .xmm15 :: vcadd d t + +/-- The butterflies of Algorithm 41 on the doublewords of `xmm0` (`w[j]`) +and `xmm1` (`w[j + len]`) with the zetas `ζ · 2³² mod q` in `xmm13` (and +its odd doublewords in the even ones of `xmm12`): `xmm0 ← xmm0 + ζ · xmm1` +and `xmm3 ← xmm0 - ζ · xmm1`. -/ +def vbfly : List Instr := + vmont .xmm1 .xmm13 .xmm12 .xmm2 .xmm4 ++ vcsub .xmm1 .xmm2 ++ + (xmov .xmm3 .xmm0 :: xb .paddd .xmm0 .xmm1 :: vcsub .xmm0 .xmm2) ++ + (xb .psubd .xmm3 .xmm1 :: vcadd .xmm3 .xmm2) + +/-- The butterflies of Algorithm 42 on the doublewords of `xmm0` (`w[j]`) +and `xmm1` (`w[j + len]`) with the zetas `ζ · 2³² mod q` in `xmm13` (and +`xmm12`): `xmm0 ← xmm0 + xmm1` and `xmm3 ← ζ · (xmm1 - xmm0)` (Algorithm +42 multiplies `w[j] - w[j + len]` by `-ζ`), from `xmm1 - xmm0 + q`. -/ +def vibfly : List Instr := + (xmov .xmm3 .xmm1 :: xb .psubd .xmm3 .xmm0 :: xb .paddd .xmm3 .xmm15 :: xb .paddd .xmm0 .xmm1 :: + vcsub .xmm0 .xmm2) ++ vmont .xmm3 .xmm13 .xmm12 .xmm2 .xmm4 ++ vcsub .xmm3 .xmm2 + +/-- The `u64`s `t (2i) + 2³² · t (2i + 1)` for `i < n / 2` at `[r + 8i]`, +through `r9`. -/ +def dwordTab (t : Nat → Nat) (n : Nat) (r : Reg) : List Instr := + (List.range (n / 2)).flatMap fun i => + [.movImm64 .r9 (BitVec.ofNat 64 (t (2 * i) + 2 ^ 32 * t (2 * i + 1))), .store (at_ r (8 * i)) .r9] + +end VG.Impl.MlDsa.X86_64.Arith diff --git a/lean/VerifiedGarbage/Proof/Framework/X86_64/Mxcsr.lean b/lean/VerifiedGarbage/Proof/Framework/X86_64/Mxcsr.lean index d6590c716..c276269aa 100644 --- a/lean/VerifiedGarbage/Proof/Framework/X86_64/Mxcsr.lean +++ b/lean/VerifiedGarbage/Proof/Framework/X86_64/Mxcsr.lean @@ -54,6 +54,31 @@ def ctlOk : Prog isa → Bool | .call _ b => ctlOk b | .frame i b j => !loadsMxcsr i && ctlOk b && !loadsMxcsr j +/-- A check that implies `ctlOk` and composes like `Code.allInstrs`: `c` +itself never loads MXCSR, and the functions it calls satisfy `ctlOk`. -/ +def ctlC : Prog isa → Bool + | .block is => is.all fun i => !loadsMxcsr i + | .seq a b => ctlC a && ctlC b + | .ite _ t e => ctlC t && ctlC e + | .loop b _ => ctlC b + | .call _ b => ctlOk b + | .frame i b j => !loadsMxcsr i && ctlC b && !loadsMxcsr j + +theorem ctlOk_of_ctlC {c : Prog isa} (h : ctlC c = true) : ctlOk c = true := by + induction c with + | block _ => exact h + | seq _ _ iha ihb => + simp only [ctlC, Bool.and_eq_true] at h + simp only [ctlOk, iha h.1, ihb h.2, Bool.and_self, Bool.or_true] + | ite _ _ _ iht ihe => + simp only [ctlC, Bool.and_eq_true] at h + simp only [ctlOk, iht h.1, ihe h.2, Bool.and_self] + | loop _ _ ih => exact ih h + | call _ _ _ => exact h + | frame _ _ _ ih => + simp only [ctlC, Bool.and_eq_true] at h + simp only [ctlOk, h.1.1, h.2, ih h.1.2, Bool.and_self] + /-- MXCSR's control bits, which the calling convention preserves. -/ abbrev ctl (v : BitVec 32) : BitVec 10 := v.extractLsb' 6 10 diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Ntt.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Ntt.lean index 0b1f7e992..fbd108838 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Ntt.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Ntt.lean @@ -1,12 +1,17 @@ -import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.NttLoop +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.VLay21 +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Basic +import VerifiedGarbage.Proof.MlKem.X86_64.VMxcsr import VerifiedGarbage.Proof.Framework.X86_64.Abi +import VerifiedGarbage.Proof.Framework.Range /-! # ML-DSA on x86-64: `vg_mldsa_ntt` -Untrusted: everything here is checked by Lean. The butterfly's code does -what `bfly` does (`bfly_spec`), so each layer is `nttLayer` (`lay_ok`), and -the eight layers are `NTT` (`ntt_eq_layers`). `Ntt.LI`, `Ntt.pro_ok` and +Untrusted: everything here is checked by Lean. ML-KEM's `withMxcsr` runs +its code from any MXCSR and keeps what it does (`withMxcsr_ok`); the +prologue leaves the table of zetas in `scratch` and the constants +(`vpro_ok`), each layer is `nttLayer` (`vlay_ok`, `vlay2_ok`, `vlay1_ok`), +and the eight layers are `NTT` (`ntt_eq_layers`). `LI`, `vpro_ok` and `inPlaceSat` serve `NTT⁻¹` too. -/ @@ -14,99 +19,233 @@ namespace VG.Proof.MlDsa.X86_64.Arith open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Arith open VG.Proof.MlDsa.Arith -open VG.Proof.MlKem.X86_64 (Keep WP.keep writesOnly gprPreserved_of) +open VG.Proof.MlKem.X86_64 (Keep WP.keep writesOnly gprPreserved_of withMxcsr_ok mxR mx_sub xmm_setXmm + GOnly add_ofNat_zero sel) open VG.Spec.MlDsa (q n Poly Zq coeffAt polyAt Reduced PolyIs zetas ntt) -namespace Ntt - -/-- Between layers: the polynomial `F` at `fP`, the table `tab` at `zP`, and -entry `k` of the table at `r8`. -/ -structure LI (tab : Nat → Nat) (s₀ : State) (fP zP : Addr) (F : Poly) (k : Nat) (s : State) : Prop where - rsi : s.gpr .rsi = fP - r8 : s.gpr .r8 = coeffAddr zP k - poly : PolyIs s.mem fP F - rd : s.rd = s₀.rd - wr : s.wr = s₀.wr - tab : Tab tab s.mem zP 256 - frame : Frame [pR fP, pR zP] s₀.mem s.mem - keep : Keep [.rax, .rcx, .rdx, .rsi, .rdi, .r8, .r9, .r10, .r11] s₀ s - -theorem LI.step {tab : Nat → Nat} {s₀ : State} {fP zP : Addr} {F F' : Poly} {k k' : Nat} {s s' : State} - (hI : LI tab s₀ fP zP F k s) (hP : PolyIs s'.mem fP F') (hf : Frame [pR fP] s.mem s'.mem) - (hsi : s'.gpr .rsi = fP) (h8 : s'.gpr .r8 = coeffAddr zP k') - (hk : Keep [.rdi, .r9, .r8, .rcx, .rax, .rdx, .rsi, .rcx, .r10, .r11, .rsi, .rdi, .rsi] s s') - (hd : (pR zP).Disjoint (pR fP)) : LI tab s₀ fP zP F' k' s' := - ⟨hsi, h8, hP, hk.2.1.trans hI.rd, hk.2.2.trans hI.wr, hI.tab.frame hf (by simpa using hd) (by decide), - hI.frame.trans (hf.mono (by simp)), (hI.keep.trans hk).mono (by decide)⟩ - -/-- The chain of zeta indices of the layers `ls` of `NTT`, from `k`. -/ -def Chain : Nat → List Nat → Prop - | _, [] => True - | k, len :: ls => k = 128 / len ∧ Chain (256 / len) ls - -theorem lens_fwd : ∀ len ∈ nttLens, 2 * (128 / len) ≤ 256 ∧ 128 / len + 128 / len = 256 / len := by decide - -theorem zetaTab_of : TabOf zetaTab zetas := fun k _ => zetaNat_eq k - -theorem lays_ok {s₀ : State} {fP zP : Addr} (hw : pR fP ∈ s₀.wr) (hz : pR zP ∈ s₀.rd ++ s₀.wr) - (hd : (pR zP).Disjoint (pR fP)) : - ∀ (ls : List Nat) (F : Poly) (k : Nat) (s : State), (∀ len ∈ ls, len ∈ nttLens) → Chain k ls → - LI zetaTab s₀ fP zP F k s → - WP isa (nttLays ls) s fun s' => ∃ k', LI zetaTab s₀ fP zP (ls.foldl nttLayer F) k' s' - | [], F, k, s, _, _, hI => WP.block_nil ⟨k, hI⟩ - | len :: ls, F, k, s, hls, ⟨hk, hc⟩, hI => by - have hlen := hls len (List.mem_cons_self ..) - obtain ⟨h1, h2⟩ := lens_fwd len hlen - refine WP.seq (WP.mono (lay_ok bfly_spec zetaTab_of hlen 4 (fun c => 128 / len + c) (fun c hc => by omega) - (fun c _ => by rw [show BitVec.signExtend 64 (4 : BitVec 32) = 4 by decide, coeffAddr_succ]; rfl) - F s hI.rsi (by rw [hI.r8, hk]; rfl) hI.poly (by rw [hI.wr]; exact hw) (by rw [hI.rd, hI.wr]; exact hz) hd - hI.tab) fun s' ⟨⟨hP, hf, hsi, h8⟩, hk'⟩ => ?_) - exact lays_ok hw hz hd ls _ (256 / len) s' (fun l hl => hls l (List.mem_cons_of_mem _ hl)) hc - (hI.step hP hf hsi (by rw [h8, h2]) hk' hd) - -/-- The table `tab` to `scratch`, `rsi` = `f` and `r8` at entry `k`. -/ -theorem pro_ok {s₀ : State} {t : Poly → Poly} (hp : (inPlaceK t).pre s₀) (tab : Nat → Nat) (d : BitVec 32) - (k : Nat) (hd : BitVec.signExtend 64 d = BitVec.ofNat 64 (4 * k)) : - WP isa (.block (nttPro tab ++ ([.alu .add .r8 (.imm d)] : List Instr))) s₀ - (LI tab s₀ (s₀.gpr .rdi) (s₀.gpr .rsi) (polyAt s₀.mem (s₀.gpr .rdi)) k) := by - simp only [nttPro, List.append_assoc] - rw [WP.block_append_iff] - refine WP.mono (WP.keep [.r9] (Q := fun s => s.mem = s₀.mem ∧ s.gpr .r9 = s₀.gpr .rsi) (by xrund) (by decide)) - fun s1 ⟨⟨hm1, h9⟩, k1⟩ => ?_ - rw [WP.block_append_iff] - refine WP.mono (storeTab_ok tab (by decide) s1 (by rw [k1.2.2, hp.2.1, h9]; simp)) - fun s2 ⟨ht, hf, k2⟩ => ?_ - have k12 := k1.trans k2 - refine WP.mono (WP.keep [.rsi, .r8] (Q := fun s => s.mem = s2.mem ∧ s.gpr .rsi = s2.gpr .rdi ∧ - s.gpr .r8 = s2.gpr .r9 + BitVec.signExtend 64 d) (by xrund [List.cons_append, List.nil_append]) (by rfl)) - fun s3 ⟨⟨hm3, hsi, h8⟩, k3⟩ => ?_ - have hd' : (pR (s₀.gpr .rdi)).Disjoint (pR (s₀.gpr .rsi)) := hp.2.2.1 - rw [h9] at hf - refine ⟨by rw [hsi, k12.gpr (by decide)], by rw [h8, k2.gpr (by decide), h9, hd], ?_, by rw [k3.2.1, k12.2.1], - by rw [k3.2.2, k12.2.2], by rw [hm3, ← h9]; exact ht, - by rw [hm3, ← hm1]; exact hf.mono (by simp), ((k12.trans k3)).mono (by decide)⟩ - rw [hm3] - exact ⟨reduced_frame (by rw [← hm1]; exact hf) (by simpa using hd') hp.2.2.2.2.2, - polyAt_frame (by rw [← hm1]; exact hf) (by simpa using hd')⟩ - -end Ntt - -theorem chain_fwd : Ntt.Chain 1 nttLens := by - simp only [nttLens, Ntt.Chain]; decide +/-! ## The prologue -/ + +/-- A table of 256 `u32`s `t k`, stored at `sP` (in `r`), two at a time +through `r9`. -/ +theorem dwordTab_ok (t : Nat → Nat) (ht : ∀ k, t k < 2 ^ 32) {r : Reg} (hr : r ≠ .r9) {sP : Addr} {s : State} + (hsi : s.gpr r = sP) (hw : pR sP ∈ s.wr) : + WP isa (.block (dwordTab t 256 r)) s fun s' => Tab t s'.mem sP 256 ∧ + Frame [pR sP] s.mem s'.mem ∧ Keep [.r9] s s' ∧ s'.mxcsr = s.mxcsr ∧ s'.xmm = s.xmm := by + refine WP.mono (wp_range_flatMap (M := isa) (N := 128) (fun i w => + Tab t w.mem sP (2 * i) ∧ Frame [pR sP] s.mem w.mem ∧ + Keep [.r9] s w ∧ w.mxcsr = s.mxcsr ∧ w.xmm = s.xmm) + (fun i w hi ⟨hT, hf, hk, hm, hx⟩ => ?_) 128 (Nat.le_refl _) s + ⟨fun _ h => absurd h (by omega), Frame.refl _ _, Keep.refl _ _, rfl, rfl⟩) + fun w ⟨hT, hf, hk, hm, hx⟩ => ⟨hT, hf, hk, hm, hx⟩ + have hsi' : w.gpr r = sP := by + rw [hk.gpr (by simp only [List.mem_singleton]; exact hr), hsi] + have w0 : InRegions w.wr (sP + BitVec.ofNat 64 (8 * i)) 8 := + ⟨_, by rw [hk.2.2]; exact hw, Offset.contains_base sP (by omega) (by omega)⟩ + have hV : ∀ e < 2, (BitVec.ofNat 64 (t (2 * i) + 2 ^ 32 * t (2 * i + 1))).extractLsb' (32 * e) 32 = + BitVec.ofNat 32 (t (2 * i + e)) := fun e he => by + apply BitVec.eq_of_toNat_eq + have h0 := ht (2 * i) + have h1 := ht (2 * i + 1) + rw [BitVec.extractLsb'_toNat, BitVec.toNat_ofNat, BitVec.toNat_ofNat, Nat.shiftRight_eq_div_pow] + rcases (by omega : e = 0 ∨ e = 1) with rfl | rfl + · simp only [Nat.mul_zero, Nat.pow_zero, Nat.div_one, Nat.add_zero]; omega + · simp only [Nat.mul_one]; omega + vrund [hsi', w0, hr] + generalize BitVec.ofNat 64 (t (2 * i) + 2 ^ 32 * t (2 * i + 1)) = V at hV ⊢ + refine ⟨fun k hk' => ?_, hf.writeW (List.mem_singleton_self _) _ + (Offset.contains_base sP (by omega) (by omega)), + ⟨fun r hr => ?_, hk.2.1, hk.2.2⟩, hm, hx⟩ + · by_cases h : 2 * i ≤ k + · rw [coeffAt_eq, coeffAddr, show sP + BitVec.ofNat 64 (4 * k) = + sP + BitVec.ofNat 64 (8 * i) + BitVec.ofNat 64 (4 * (k - 2 * i)) by + rw [BitVec.add_assoc, ← BitVec.ofNat_add]; exact congrArg _ (congrArg _ (by omega)), + show 32 = 8 * 4 from rfl, readW_writeW_inside _ _ _ (by omega) (by decide), + show 8 * (4 * (k - 2 * i)) = 32 * (k - 2 * i) by omega, hV _ (by omega), + show 2 * i + (k - 2 * i) = k by omega] + · rw [coeffAt_eq, Mem.readW_writeW_sep (Offset.sep sP (by omega) (by omega) (by omega)) (by decide)] + exact hT k (by omega) + · simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + simp only [RegUpd.gpr_setReg, hr, ite_false] + exact hk.1 r (by simp [hr]) + +theorem vconsts_ok (s : State) : + WP isa (.block vconsts) s fun s' => VConsts s' ∧ Keep [.rax] s s' ∧ s'.mem = s.mem ∧ + s'.mxcsr = s.mxcsr := by + simp only [vconsts] + vrund + refine ⟨⟨?_, ?_⟩, ⟨fun r hr => ?_, rfl, rfl⟩⟩ + · simp only [RegUpd.xmm_setReg, xmm_setXmm, ite_true]; decide + · simp only [RegUpd.xmm_setReg, xmm_setXmm, ite_true, ite_false, reduceCtorEq]; decide + · simp only [List.mem_singleton] at hr + simp only [RegUpd.gpr_setReg, RegUpd.gpr_setXmm, hr, ite_false] + +/-- The table of zetas and the constants. -/ +theorem vpro_ok {sP : Addr} {s : State} (hsi : s.gpr .rsi = sP) (hw : pR sP ∈ s.wr) : + WP isa (.block vpro) s fun s' => Tab zmTab s'.mem sP 256 ∧ VConsts s' ∧ + Frame [pR sP] s.mem s'.mem ∧ Keep [.r9, .rax] s s' ∧ s'.mxcsr = s.mxcsr := by + rw [vpro, WP.block_append_iff] + refine WP.mono (dwordTab_ok zmTab (fun k => Nat.lt_trans (zmTab_lt k) (by decide)) (by decide) hsi hw) + fun s1 ⟨hT, hf, k1, x1, _⟩ => WP.mono (vconsts_ok s1) fun s2 ⟨hc, k2, m2, x2⟩ => + ⟨by rw [m2]; exact hT, hc, by rw [m2]; exact hf, (k1.trans k2).mono (by simp), by rw [x2, x1]⟩ + +/-! ## The layers -/ + +/-- Between the layers: the polynomial `F` at `fP`, the table at `sP`, and +the constants. -/ +structure LI (fP sP : Addr) (s₀ : State) (F : Poly) (s : State) : Prop where + P : PolyIs s.mem fP F + T : Tab zmTab s.mem sP 256 + c : VConsts s + keep : Keep [.rax, .rcx, .rdx, .r8] s₀ s + frame : Frame [pR fP] s₀.mem s.mem + +/-- The last layer. -/ +theorem LI.last {fP sP : Addr} {s₀ : State} (hdi : s₀.gpr .rdi = fP) (hsi : s₀.gpr .rsi = sP) + (hwf : pR fP ∈ s₀.wr) (hw : pR sP ∈ s₀.wr) (hd : (pR sP).Disjoint (pR fP)) {l : Prog isa} + {F F' : Poly} + (hl : ∀ s, VConsts s → s.gpr .rdi = fP → s.gpr .rsi = sP → PolyIs s.mem fP F → Tab zmTab s.mem sP 256 → + pR fP ∈ s.wr → pR sP ∈ s.wr → WP isa l s fun s' => PolyIs s'.mem fP F' ∧ BInv fP s s') + {s : State} (hI : LI fP sP s₀ F s) : WP isa l s (LI fP sP s₀ F') := + WP.mono (hl s hI.c (by rw [hI.keep.gpr (by decide), hdi]) (by rw [hI.keep.gpr (by decide), hsi]) hI.P + hI.T (by rw [hI.keep.2.2]; exact hwf) (by rw [hI.keep.2.2]; exact hw)) + fun s' ⟨hS, hb⟩ => ⟨hS, hI.T.frame hb.frame (fun r hr => by rw [List.mem_singleton.mp hr]; exact hd) + (by decide), hb.consts, (hI.keep.trans hb.keep).mono (by decide), hI.frame.trans hb.frame⟩ + +/-- A layer, then `c`. -/ +theorem LI.seq {fP sP : Addr} {s₀ : State} (hdi : s₀.gpr .rdi = fP) (hsi : s₀.gpr .rsi = sP) + (hwf : pR fP ∈ s₀.wr) (hw : pR sP ∈ s₀.wr) (hd : (pR sP).Disjoint (pR fP)) {l c : Prog isa} + {F F' : Poly} {Q : State → Prop} + (hl : ∀ s, VConsts s → s.gpr .rdi = fP → s.gpr .rsi = sP → PolyIs s.mem fP F → Tab zmTab s.mem sP 256 → + pR fP ∈ s.wr → pR sP ∈ s.wr → WP isa l s fun s' => PolyIs s'.mem fP F' ∧ BInv fP s s') + (hc : ∀ s, LI fP sP s₀ F' s → WP isa c s Q) {s : State} (hI : LI fP sP s₀ F s) : + WP isa (.seq l c) s Q := + WP.seq (WP.mono (hl s hI.c (by rw [hI.keep.gpr (by decide), hdi]) (by rw [hI.keep.gpr (by decide), hsi]) hI.P + hI.T (by rw [hI.keep.2.2]; exact hwf) (by rw [hI.keep.2.2]; exact hw)) + fun s' ⟨hS, hb⟩ => hc s' ⟨hS, hI.T.frame hb.frame (fun r hr => by rw [List.mem_singleton.mp hr]; exact hd) + (by decide), hb.consts, (hI.keep.trans hb.keep).mono (by decide), hI.frame.trans hb.frame⟩) + +theorem step_fwd (p : Addr) (a d : Nat) {dz : BitVec 32} (h : BitVec.signExtend 64 dz = BitVec.ofNat 64 (4 * d)) : + coeffAddr p a + BitVec.signExtend 64 dz = coeffAddr p (a + d) := by + rw [h, coeffAddr_add] + +theorem step_bwd (p : Addr) (a d : Nat) {dz : BitVec 32} + (h : BitVec.ofNat 64 (4 * d) + BitVec.signExtend 64 dz = 0) : + coeffAddr p (a + d) + BitVec.signExtend 64 dz = coeffAddr p a := by + rw [← coeffAddr_add, BitVec.add_assoc, h]; exact BitVec.add_zero _ + +/-- The block of `NTT`. -/ +abbrev fwdBlk : Poly → Nat → Nat → Nat → Nat → Poly := fun f len k st t => blockN bfly f len (zetas k) st t + +theorem nttLayer_eq (F : Poly) (len : Nat) : + nttLayer F len = layF fwdBlk F len (fun c => 128 / len + c) (128 / len) := rfl + +/-- A layer of `NTT` with `len ≥ 4`, whose first zeta is `zetas k`. -/ +theorem fwdLay_ok {fP sP : Addr} (hd : (pR sP).Disjoint (pR fP)) (len k : Nat) (hlen : len ∈ [4, 8, 16, 32, 64, 128]) (hk : 128 / len = k) + {F : Poly} (s : State) (hc : VConsts s) (hdi : s.gpr .rdi = fP) (hsi : s.gpr .rsi = sP) + (hS : PolyIs s.mem fP F) (hT : Tab zmTab s.mem sP 256) (hwf : pR fP ∈ s.wr) (hw : pR sP ∈ s.wr) : + WP isa (vlay vbfly len k 4) s fun s' => PolyIs s'.mem fP (nttLayer F len) ∧ BInv fP s s' := by + rw [nttLayer_eq] + exact vlay_ok vbfly_spec nttBlk_ok hlen 4 (fun c => 128 / len + c) (by rw [hk]; rfl) + (fun c hc => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hlen + rcases hlen with rfl | rfl | rfl | rfl | rfl | rfl <;> omega) + (fun c _ => step_fwd _ _ 1 (by decide)) hc hdi hsi hS hT hwf hw hd + +theorem fwdLay2_ok {fP sP : Addr} (hd : (pR sP).Disjoint (pR fP)) {F : Poly} (s : State) (hc : VConsts s) (hdi : s.gpr .rdi = fP) + (hsi : s.gpr .rsi = sP) (hS : PolyIs s.mem fP F) (hT : Tab zmTab s.mem sP 256) (hwf : pR fP ∈ s.wr) + (hw : pR sP ∈ s.wr) : + WP isa (vlay2 vbfly 64 0x50 8) s fun s' => PolyIs s'.mem fP (nttLayer F 2) ∧ BInv fP s s' := by + have hs : ∀ e < 4, sel 0x50 e = e / 2 := by decide + rw [nttLayer_eq, show 128 / 2 = 64 from rfl] + exact vlay2_ok vbfly_spec nttBlk_ok 64 0x50 8 (fun c => 64 + c) (fun i => 64 + 2 * i) rfl + (fun i _ => by omega) (fun i _ e he => by rw [hs e he]; omega) + (fun i _ => (step_fwd _ _ 2 (by decide)).trans (congrArg _ (by omega))) hc hdi hsi hS hT hwf hw hd + +theorem fwdLay1_ok {fP sP : Addr} (hd : (pR sP).Disjoint (pR fP)) {F : Poly} (s : State) (hc : VConsts s) (hdi : s.gpr .rdi = fP) + (hsi : s.gpr .rsi = sP) (hS : PolyIs s.mem fP F) (hT : Tab zmTab s.mem sP 256) (hwf : pR fP ∈ s.wr) + (hw : pR sP ∈ s.wr) : + WP isa (vlay1 vbfly 128 0xE4 16) s fun s' => PolyIs s'.mem fP (nttLayer F 1) ∧ BInv fP s s' := by + have hs : ∀ e < 4, sel 0xE4 e = e := by decide + rw [nttLayer_eq, show 128 / 1 = 128 from rfl] + exact vlay1_ok vbfly_spec nttBlk_ok 128 0xE4 16 (fun c => 128 + c) (fun i => 128 + 4 * i) rfl + (fun i _ => by omega) (fun i _ e he => by rw [hs e he]; omega) + (fun i _ => (step_fwd _ _ 4 (by decide)).trans (congrArg _ (by omega))) hc hdi hsi hS hT hwf hw hd + +/-! ## `vg_mldsa_ntt` -/ + +theorem mx_sub' (sP : Addr) : Region.Sub (mxR sP) (pR sP) := mx_sub sP + +/-- The regions of `scratch` within it, and `f`. -/ +theorem frame_fs {fP sP : Addr} {m m' : Mem} {rs : List Region} (h : Frame rs m m') + (hs : ∀ r ∈ rs, Region.Sub r (pR fP) ∨ Region.Sub r (pR sP)) : Frame [pR fP, pR sP] m m' := + h.sub fun r hr => (hs r hr).elim (fun h => ⟨_, List.mem_cons_self .., h⟩) + fun h => ⟨_, List.mem_cons_of_mem _ (List.mem_singleton_self _), h⟩ + +/-- The code in `withMxcsr`, from its state `s1`: the prologue, then the +layers `l`, which leave `F`, then `NTT⁻¹`'s scaling or nothing. -/ +theorem nttBody_ok {t : Poly → Poly} {s s1 : State} (hs : (inPlaceK t).pre s) {l : Prog isa} {G : Poly} + (k1 : Keep [.rax, .r11] s s1) (f1 : Frame [mxR (s.gpr .rsi)] s.mem s1.mem) + (hl : ∀ s2, LI (s.gpr .rdi) (s.gpr .rsi) s2 (polyAt s.mem (s.gpr .rdi)) s2 → s2.gpr .rdi = s.gpr .rdi → + s2.gpr .rsi = s.gpr .rsi → pR (s.gpr .rdi) ∈ s2.wr → pR (s.gpr .rsi) ∈ s2.wr → + WP isa l s2 fun s3 => LI (s.gpr .rdi) (s.gpr .rsi) s2 G s3) : + WP isa (.seq (.block vpro) l) s1 fun s' => + PolyIs s'.mem (s.gpr .rdi) G ∧ Frame [pR (s.gpr .rdi), pR (s.gpr .rsi)] s.mem s'.mem := by + have hw : pR (s.gpr .rsi) ∈ s.wr := by rw [hs.2.1]; simp + have hwf : pR (s.gpr .rdi) ∈ s.wr := by rw [hs.2.1]; simp + have hd : (pR (s.gpr .rdi)).Disjoint (pR (s.gpr .rsi)) := hs.2.2.1 + have hdi1 : s1.gpr .rdi = s.gpr .rdi := k1.gpr (by decide) + have hsi1 : s1.gpr .rsi = s.gpr .rsi := k1.gpr (by decide) + have hF1 : PolyIs s1.mem (s.gpr .rdi) (polyAt s.mem (s.gpr .rdi)) := + polyIs_frame f1 (fun r hr => by rw [List.mem_singleton.mp hr]; exact hd.sub_right (mx_sub' _)) + ⟨hs.2.2.2.2.2, rfl⟩ + refine WP.seq (WP.mono (vpro_ok hsi1 (by rw [k1.2.2]; exact hw)) fun s2 ⟨hT, hc, hf2, k2, _⟩ => ?_) + have hF2 : PolyIs s2.mem (s.gpr .rdi) (polyAt s.mem (s.gpr .rdi)) := + polyIs_frame hf2 (fun r hr => by rw [List.mem_singleton.mp hr]; exact hd) hF1 + refine WP.mono (hl s2 ⟨hF2, hT, hc, Keep.refl _ _, Frame.refl _ _⟩ + (by rw [k2.gpr (by decide), hdi1]) (by rw [k2.gpr (by decide), hsi1]) + (by rw [k2.2.2, k1.2.2]; exact hwf) (by rw [k2.2.2, k1.2.2]; exact hw)) fun s3 hI => ⟨hI.P, ?_⟩ + refine (frame_fs f1 ?_).trans ((frame_fs hf2 ?_).trans (frame_fs hI.frame ?_)) <;> + intro r hr <;> simp only [List.mem_singleton] at hr <;> subst hr + exacts [.inr (mx_sub' _), .inr fun _ h => h, .inl fun _ h => h] + +/-- `withMxcsr` around the body, and the ABI. -/ +theorem mx_correct {t : Poly → Poly} {l : Prog isa} (s : State) (hs : (inPlaceK t).pre s) + (hk : writesOnly [.rax, .rcx, .rdx, .r8, .r9] (.seq (.block vpro) l) = true) + (hctl : ctlOk (VG.Impl.MlKem.X86_64.withMxcsr .rsi 768 (.seq (.block vpro) l)) = true) + (hk' : writesOnly [.rax, .rcx, .rdx, .r8, .r9, .r11] + (VG.Impl.MlKem.X86_64.withMxcsr .rsi 768 (.seq (.block vpro) l)) = true) + (hl : ∀ s1, Keep [.rax, .r11] s s1 → Frame [mxR (s.gpr .rsi)] s.mem s1.mem → + WP isa (.seq (.block vpro) l) s1 fun s' => PolyIs s'.mem (s.gpr .rdi) (t (polyAt s.mem (s.gpr .rdi))) ∧ + Frame [pR (s.gpr .rdi), pR (s.gpr .rsi)] s.mem s'.mem) : + ∃ tr s', Exec isa (VG.Impl.MlKem.X86_64.withMxcsr .rsi 768 (.seq (.block vpro) l)) s tr s' ∧ + abiPreserved s s' ∧ (inPlaceK t).post s s' := by + have hw : pR (s.gpr .rsi) ∈ s.wr := by rw [hs.2.1]; simp + have hd : (pR (s.gpr .rdi)).Disjoint (pR (s.gpr .rsi)) := hs.2.2.1 + have hW := withMxcsr_ok (c := .seq (.block vpro) l) (by decide) [.rax, .rcx, .rdx, .r8, .r9] (by decide) rfl hw + hk (hl) + obtain ⟨tr, s', he, ⟨s2, ⟨hP, hf⟩, hf', -⟩, hk⟩ := WP.keep [.rax, .rcx, .rdx, .r8, .r9, .r11] hW hk' + refine ⟨tr, s', he, abiPreserved_of_ctl hctl he (gprPreserved_of hk (by decide) + (hf.trans (hf'.sub fun r hr => ⟨_, List.mem_cons_of_mem _ (List.mem_singleton_self _), ?_⟩)) + (by simpa using ⟨hs.2.2.2.1, hs.2.2.2.2.1⟩)), ?_⟩ + · rw [List.mem_singleton.mp hr]; exact mx_sub' _ + · exact polyIs_frame hf' (fun r hr => by + rw [List.mem_singleton.mp hr]; exact hd.sub_right (mx_sub' _)) hP theorem ntt_correct (s : State) (hs : (inPlaceK ntt).pre s) : ∃ t s', Exec isa Impl.MlDsa.X86_64.Arith.ntt s t s' ∧ abiPreserved s s' ∧ (inPlaceK ntt).post s s' := by - have hw : pR (s.gpr .rdi) ∈ s.wr := by rw [hs.2.1]; simp - have hz : pR (s.gpr .rsi) ∈ s.rd ++ s.wr := by rw [hs.1, hs.2.1]; simp - obtain ⟨t, s', he, ⟨k, hI⟩, hk⟩ := WP.keep (c := Impl.MlDsa.X86_64.Arith.ntt) - [.rax, .rcx, .rdx, .rsi, .rdi, .r8, .r9, .r10, .r11] - (WP.seq (WP.mono (Ntt.pro_ok hs zetaTab 4 1 (by decide)) fun s1 hI => - Ntt.lays_ok hw hz hs.2.2.1.symm nttLens _ 1 s1 (fun _ h => h) chain_fwd hI)) (by decide +kernel) - refine ⟨t, s', he, abiPreserved_of_exec (by decide +kernel) he (gprPreserved_of hk (by decide) hI.frame - (by simpa using ⟨hs.2.2.2.1, hs.2.2.2.2.1⟩)), ?_⟩ - show PolyIs _ _ _ - rw [ntt_eq_layers] - exact hI.poly + have hd : (pR (s.gpr .rsi)).Disjoint (pR (s.gpr .rdi)) := hs.2.2.1.symm + refine mx_correct s hs (by decide +kernel) (by decide +kernel) (by decide +kernel) fun s1 k1 f1 => + WP.mono (nttBody_ok hs k1 f1 (G := nttLens.foldl nttLayer (polyAt s.mem (s.gpr .rdi))) + fun s2 hI hdi hsi hwf hw => ?_) fun s' ⟨hP, hf⟩ => ⟨by rw [ntt_eq_layers]; exact hP, hf⟩ + simp only [nttLens, List.foldl_cons, List.foldl_nil] + refine LI.seq hdi hsi hwf hw hd (fwdLay_ok hd 128 1 (by decide) (by decide)) ?_ hI + refine fun _ hI => LI.seq hdi hsi hwf hw hd (fwdLay_ok hd 64 2 (by decide) (by decide)) ?_ hI + refine fun _ hI => LI.seq hdi hsi hwf hw hd (fwdLay_ok hd 32 4 (by decide) (by decide)) ?_ hI + refine fun _ hI => LI.seq hdi hsi hwf hw hd (fwdLay_ok hd 16 8 (by decide) (by decide)) ?_ hI + refine fun _ hI => LI.seq hdi hsi hwf hw hd (fwdLay_ok hd 8 16 (by decide) (by decide)) ?_ hI + refine fun _ hI => LI.seq hdi hsi hwf hw hd (fwdLay_ok hd 4 32 (by decide) (by decide)) ?_ hI + refine fun _ hI => LI.seq hdi hsi hwf hw hd (fwdLay2_ok hd) ?_ hI + exact fun _ hI => LI.last hdi hsi hwf hw hd (fwdLay1_ok hd) hI /-- The pointers and `rsp` are public. -/ theorem inPlace_agree {t : Poly → Poly} (s₁ s₂ : State) (_ : (inPlaceK t).pre s₁) (_ : (inPlaceK t).pre s₂) diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/NttBfly.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/NttBfly.lean deleted file mode 100644 index d836ba0a2..000000000 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/NttBfly.lean +++ /dev/null @@ -1,201 +0,0 @@ -import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Ntt -import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Basic -import VerifiedGarbage.Proof.MlDsa.Arith.Ntt - -/-! -# ML-DSA on x86-64: the butterflies of `NTT` and `NTT⁻¹` - -Untrusted: everything here is checked by Lean. What one butterfly's code -stores (`bfly_ok`, `bflyInv_ok`), for any `len`, from the words it reads -and the zeta in `r9`; and that it does what the butterfly of the -specification does (`bfly_spec`, `bflyInv_spec`). --/ - -namespace VG.Proof.MlDsa.X86_64.Arith - -open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Arith -open VG.Proof.MlDsa.Arith -open VG.Proof.MlKem.X86_64 (Keep WP.keep) -open VG.Spec.MlDsa (q n Poly Zq coeffAt polyAt Reduced PolyIs) - -/-! ## `NTT` -/ - -/-- The part of `bfly` after the product is reduced. -/ -def bflyTail (len : Nat) : List Instr := - [.mov32 .rax (.mem (at_ .rsi 0)), .mov32 .rdx (.reg .rax), .alu32 .add .rdx (.imm qImm), - .alu32 .sub .rdx (.reg .r10)] ++ csubQ .rdx .r11 ++ - [.store32 (at_ .rsi (4 * len)) .rdx, .alu32 .add .rax (.reg .r10)] ++ csubQ .rax .r11 ++ - [.store32 (at_ .rsi 0) .rax, .alu .add .rsi (.imm 4), .alu .sub .rcx (.imm 1)] - -theorem bfly_eq (len : Nat) : - Impl.MlDsa.X86_64.Arith.bfly len = - ([.mov32 .rax (.mem (at_ .rsi (4 * len))), .mul .r9] : List Instr) ++ (reduce ++ bflyTail len) := by - simp only [Impl.MlDsa.X86_64.Arith.bfly, bflyTail, List.append_assoc] - -theorem bflyHead_ok (len : Nat) (s : State) - (h : InRegions (s.rd ++ s.wr) (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 4) : - WP isa (.block ([.mov32 .rax (.mem (at_ .rsi (4 * len))), .mul .r9] : List Instr)) s fun s' => - (s'.gpr .rax = prodW (s.mem.readW (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 32) (s.gpr .r9) ∧ - s'.mem = s.mem) ∧ Keep [.rax, .rdx] s s' := by - refine WP.keep _ ?_ (by rfl) - xrund [h] - -theorem bflyTail_ok (len : Nat) (s : State) (h0 : InRegions (s.rd ++ s.wr) (s.gpr .rsi) 4) - (w0 : InRegions s.wr (s.gpr .rsi) 4) (w1 : InRegions s.wr (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 4) : - WP isa (.block (bflyTail len)) s fun s' => - (s'.mem = (s.mem.writeW (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) - (csubD (s.mem.readW (s.gpr .rsi) 32 + qImm - BitVec.setWidth 32 (s.gpr .r10)))).writeW (s.gpr .rsi) - (csubD (s.mem.readW (s.gpr .rsi) 32 + BitVec.setWidth 32 (s.gpr .r10))) ∧ - s'.gpr .rsi = s.gpr .rsi + 4 ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ - s'.zf = some (s.gpr .rcx - 1 == 0)) ∧ Keep [.rax, .rdx, .rsi, .rcx, .r11] s s' := by - refine WP.keep _ ?_ (by rfl) - unfold bflyTail csubQ - xrund [h0, w0, w1, List.cons_append, List.nil_append, csubD] - -theorem bfly_ok (len : Nat) (s : State) (h0 : InRegions (s.rd ++ s.wr) (s.gpr .rsi) 4) - (h1 : InRegions (s.rd ++ s.wr) (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 4) - (w0 : InRegions s.wr (s.gpr .rsi) 4) (w1 : InRegions s.wr (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 4) : - WP isa (.block (Impl.MlDsa.X86_64.Arith.bfly len)) s fun s' => - (s'.mem = (s.mem.writeW (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) - (csubD (s.mem.readW (s.gpr .rsi) 32 + qImm - BitVec.setWidth 32 - (redD (prodW (s.mem.readW (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 32) (s.gpr .r9)))))).writeW - (s.gpr .rsi) (csubD (s.mem.readW (s.gpr .rsi) 32 + BitVec.setWidth 32 - (redD (prodW (s.mem.readW (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 32) (s.gpr .r9))))) ∧ - s'.gpr .rsi = s.gpr .rsi + 4 ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ - s'.zf = some (s.gpr .rcx - 1 == 0)) ∧ Keep [.rax, .rdx, .rsi, .rcx, .r10, .r11] s s' := by - rw [bfly_eq, WP.block_append_iff] - refine WP.mono (bflyHead_ok len s h1) fun s1 ⟨⟨ha, hm1⟩, k1⟩ => ?_ - rw [WP.block_append_iff] - refine WP.mono (reduce_ok s1) fun s2 ⟨⟨hr, hm2⟩, k2⟩ => ?_ - have k12 := k1.trans k2 - have hsi : s2.gpr .rsi = s.gpr .rsi := k12.gpr (by decide) - have hrr : s2.rd ++ s2.wr = s.rd ++ s.wr := by rw [k12.2.1, k12.2.2] - refine WP.mono (bflyTail_ok len s2 (by rw [hrr, hsi]; exact h0) (by rw [k12.2.2, hsi]; exact w0) - (by rw [k12.2.2, hsi]; exact w1)) fun s3 ⟨⟨hm3, h3si, h3cx, h3z⟩, k3⟩ => ⟨?_, (k12.trans k3).mono (by decide)⟩ - have hcx : s2.gpr .rcx = s.gpr .rcx := k12.gpr (by decide) - rw [hm3, h3si, h3cx, h3z, hcx, hsi, hr, hm2, hm1, ha] - exact ⟨rfl, rfl, rfl, rfl⟩ - -/-! ## `NTT⁻¹` -/ - -/-- The part of `bflyInv` up to the product. -/ -def bflyInvHead (len : Nat) : List Instr := - [.mov32 .rax (.mem (at_ .rsi 0)), .mov32 .r10 (.mem (at_ .rsi (4 * len))), .mov32 .rdx (.reg .rax), - .alu32 .add .rdx (.reg .r10)] ++ csubQ .rdx .r11 ++ - [.store32 (at_ .rsi 0) .rdx, .alu32 .add .rax (.imm qImm), .alu32 .sub .rax (.reg .r10)] ++ - csubQ .rax .r11 ++ [.mul .r9] - -theorem bflyInv_eq (len : Nat) : - Impl.MlDsa.X86_64.Arith.bflyInv len = bflyInvHead len ++ (reduce ++ ([.store32 (at_ .rsi (4 * len)) .r10, - .alu .add .rsi (.imm 4), .alu .sub .rcx (.imm 1)] : List Instr)) := by - simp only [Impl.MlDsa.X86_64.Arith.bflyInv, bflyInvHead, List.append_assoc] - -theorem bflyInvHead_ok (len : Nat) (s : State) (h0 : InRegions (s.rd ++ s.wr) (s.gpr .rsi) 4) - (h1 : InRegions (s.rd ++ s.wr) (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 4) - (w0 : InRegions s.wr (s.gpr .rsi) 4) : - WP isa (.block (bflyInvHead len)) s fun s' => - (s'.mem = s.mem.writeW (s.gpr .rsi) (csubD (s.mem.readW (s.gpr .rsi) 32 + - s.mem.readW (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 32)) ∧ - s'.gpr .rax = prodW (csubD (s.mem.readW (s.gpr .rsi) 32 + qImm - - s.mem.readW (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 32)) (s.gpr .r9)) ∧ - Keep [.rax, .rdx, .r10, .r11] s s' := by - refine WP.keep _ ?_ (by rfl) - unfold bflyInvHead csubQ - xrund [h0, h1, w0, List.cons_append, List.nil_append, csubD] - -theorem bflyInv_ok (len : Nat) (s : State) (h0 : InRegions (s.rd ++ s.wr) (s.gpr .rsi) 4) - (h1 : InRegions (s.rd ++ s.wr) (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 4) - (w0 : InRegions s.wr (s.gpr .rsi) 4) (w1 : InRegions s.wr (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 4) : - WP isa (.block (Impl.MlDsa.X86_64.Arith.bflyInv len)) s fun s' => - (s'.mem = (s.mem.writeW (s.gpr .rsi) (csubD (s.mem.readW (s.gpr .rsi) 32 + - s.mem.readW (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 32))).writeW - (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) (BitVec.setWidth 32 (redD (prodW - (csubD (s.mem.readW (s.gpr .rsi) 32 + qImm - s.mem.readW (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 32)) - (s.gpr .r9)))) ∧ - s'.gpr .rsi = s.gpr .rsi + 4 ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ - s'.zf = some (s.gpr .rcx - 1 == 0)) ∧ Keep [.rax, .rdx, .rsi, .rcx, .r10, .r11] s s' := by - rw [bflyInv_eq, WP.block_append_iff] - refine WP.mono (bflyInvHead_ok len s h0 h1 w0) fun s1 ⟨⟨hm1, ha⟩, k1⟩ => ?_ - rw [WP.block_append_iff] - refine WP.mono (reduce_ok s1) fun s2 ⟨⟨hr, hm2⟩, k2⟩ => ?_ - have k12 := k1.trans k2 - have hsi : s2.gpr .rsi = s.gpr .rsi := k12.gpr (by decide) - have hcx : s2.gpr .rcx = s.gpr .rcx := k12.gpr (by decide) - refine WP.mono (WP.keep [.rsi, .rcx] (Q := fun s' => s'.mem = s2.mem.writeW (s2.gpr .rsi + BitVec.ofNat 64 (4 * len)) - (BitVec.setWidth 32 (s2.gpr .r10)) ∧ s'.gpr .rsi = s2.gpr .rsi + 4 ∧ s'.gpr .rcx = s2.gpr .rcx - 1 ∧ - s'.zf = some (s2.gpr .rcx - 1 == 0)) (by - xrund [show InRegions s2.wr (s2.gpr .rsi + BitVec.ofNat 64 (4 * len)) 4 by rw [k12.2.2, hsi]; exact w1]) - (by rfl)) fun s3 ⟨⟨hm3, h3si, h3cx, h3z⟩, k3⟩ => ⟨?_, (k12.trans k3).mono (by decide)⟩ - rw [hm3, h3si, h3cx, h3z, hcx, hsi, hr, hm2, hm1, ha] - exact ⟨rfl, rfl, rfl, rfl⟩ - -/-! ## What they do to a stored polynomial -/ - -/-- The code `code len` of a butterfly does what `op` does. -/ -def BflyOk (code : Nat → List Instr) (op : Poly → Nat → Nat → Zq → Poly) : Prop := - ∀ (fP : Addr) (len j : Nat), 0 < len → j + len < 256 → ∀ (z : Zq) (F : Poly) (s : State), - s.gpr .rsi = coeffAddr fP j → s.gpr .r9 = BitVec.ofNat 64 z.val → PolyIs s.mem fP F → pR fP ∈ s.wr → - WP isa (.block (code len)) s fun s' => (PolyIs s'.mem fP (op F j len z) ∧ Frame [pR fP] s.mem s'.mem ∧ - s'.gpr .rsi = s.gpr .rsi + 4 ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ s'.zf = some (s.gpr .rcx - 1 == 0)) ∧ - Keep [.rax, .rdx, .rsi, .rcx, .r10, .r11] s s' - -/-- The words a butterfly reads. -/ -theorem bfly_regions {fP : Addr} {len j : Nat} (hj : j + len < 256) {s : State} - (hsi : s.gpr .rsi = coeffAddr fP j) (hw : pR fP ∈ s.wr) : - InRegions (s.rd ++ s.wr) (s.gpr .rsi) 4 ∧ InRegions (s.rd ++ s.wr) (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 4 ∧ - InRegions s.wr (s.gpr .rsi) 4 ∧ InRegions s.wr (s.gpr .rsi + BitVec.ofNat 64 (4 * len)) 4 := by - rw [hsi, coeffAddr_add] - exact ⟨⟨_, List.mem_append_right _ hw, coeff_contains _ (show j < 256 by omega)⟩, - ⟨_, List.mem_append_right _ hw, coeff_contains _ hj⟩, ⟨_, hw, coeff_contains _ (show j < 256 by omega)⟩, - ⟨_, hw, coeff_contains _ hj⟩⟩ - -/-- The two writes of a butterfly are within the polynomial. -/ -theorem bfly_frame {fP : Addr} {len j : Nat} (hj : j + len < 256) (m : Mem) (a b : BitVec 32) : - Frame [pR fP] m ((m.writeW (coeffAddr fP (j + len)) a).writeW (coeffAddr fP j) b) ∧ - Frame [pR fP] m ((m.writeW (coeffAddr fP j) a).writeW (coeffAddr fP (j + len)) b) := - ⟨(Frame.refl _ _ |>.writeW (List.mem_singleton_self _) _ (coeff_contains _ hj)).writeW - (List.mem_singleton_self _) _ (coeff_contains _ (show j < 256 by omega)), - (Frame.refl _ _ |>.writeW (List.mem_singleton_self _) _ (coeff_contains _ (show j < 256 by omega))).writeW - (List.mem_singleton_self _) _ (coeff_contains _ hj)⟩ - -theorem bfly_spec : BflyOk Impl.MlDsa.X86_64.Arith.bfly Arith.bfly := by - intro fP len j hlen hj z F s hsi h9 hF hw - obtain ⟨r0, r1, w0, w1⟩ := bfly_regions hj hsi hw - refine WP.mono (bfly_ok len s r0 r1 w0 w1) fun s' ⟨⟨hm, hsi', hcx, hz⟩, hk⟩ => ⟨⟨?_, ?_, hsi', hcx, hz⟩, hk⟩ - · rw [hm, h9, hsi, coeffAddr_add, ← coeffAt_eq, ← coeffAt_eq] - have hj' : j < 256 := by omega - have ha := polyIs_toNat hF hj' - have hT := redD_prodW (z := z) (polyIs_toNat hF hj) - show PolyIs _ _ ((F.set! (j + len) (F[j]! - z * F[j + len]!)).set! j - ((F.set! (j + len) (F[j]! - z * F[j + len]!))[j]! + z * F[j + len]!)) - rw [getElem!_set!_ne _ hj' (by omega)] - exact polyIs_writeW (polyIs_writeW hF hj _ (csubD_sub_val ha hT)) hj' _ (csubD_add_val ha hT) - · rw [hm, hsi, coeffAddr_add] - exact (bfly_frame hj _ _ _).1 - -theorem bflyInv_spec : BflyOk Impl.MlDsa.X86_64.Arith.bflyInv Arith.bflyInv := by - intro fP len j hlen hj z F s hsi h9 hF hw - obtain ⟨r0, r1, w0, w1⟩ := bfly_regions hj hsi hw - refine WP.mono (bflyInv_ok len s r0 r1 w0 w1) fun s' ⟨⟨hm, hsi', hcx, hz⟩, hk⟩ => ⟨⟨?_, ?_, hsi', hcx, hz⟩, hk⟩ - · rw [hm, h9, hsi, coeffAddr_add, ← coeffAt_eq, ← coeffAt_eq] - have hj' : j < 256 := by omega - have ha := polyIs_toNat hF hj' - have hu := polyIs_toNat hF hj - have hne : j ≠ j + len := by omega - show PolyIs _ _ (((F.set! j (F[j]! + F[j + len]!)).set! (j + len) - (F[j]! - (F.set! j (F[j]! + F[j + len]!))[j + len]!)).set! (j + len) - (z * ((F.set! j (F[j]! + F[j + len]!)).set! (j + len) - (F[j]! - (F.set! j (F[j]! + F[j + len]!))[j + len]!))[j + len]!)) - rw [getElem!_set!_ne _ hj hne, getElem!_set!_self _ hj] - have e : ∀ (G : Poly) (x y : Zq), (G.set! (j + len) x).set! (j + len) y = G.set! (j + len) y := fun G x y => - ext_getElem! fun i hi => by - by_cases h : i = j + len - · subst h; rw [getElem!_set!_self _ hi, getElem!_set!_self _ hi] - · rw [getElem!_set!_ne _ hi (Ne.symm h), getElem!_set!_ne _ hi (Ne.symm h), getElem!_set!_ne _ hi (Ne.symm h)] - rw [e] - exact polyIs_writeW (polyIs_writeW hF hj' _ (csubD_add_val ha hu)) hj _ - (redD_prodW (csubD_sub_val ha hu)) - · rw [hm, hsi, coeffAddr_add] - exact (bfly_frame hj _ _ _).2 - -end VG.Proof.MlDsa.X86_64.Arith diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/NttInv.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/NttInv.lean index bed819100..f2597066f 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/NttInv.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/NttInv.lean @@ -3,163 +3,200 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Ntt /-! # ML-DSA on x86-64: `vg_mldsa_inv_ntt` -Untrusted: everything here is checked by Lean. The butterfly's code does -what `bflyInv` does (`bflyInv_spec`), with the negated zetas of the table -(`negZetaTab_of`), so each layer is `nttInvLayer`, the eight layers are -those of `NTT⁻¹` (`nttInv_eq_layers`), and the last loop multiplies each -coefficient by `8347681 = 256⁻¹ mod q`. +Untrusted: everything here is checked by Lean. As `vg_mldsa_ntt` +(`Ntt.lean`): each layer is `nttInvLayer` (with `vibfly`, which multiplies +by `ζ` the difference the other way round: `nttInvBlk_ok`), the eight +layers are those of `NTT⁻¹` (`nttInv_eq_layers`), and the last pass +multiplies each coefficient by `8347681 = 256⁻¹ mod q` (`vscale_ok`). -/ namespace VG.Proof.MlDsa.X86_64.Arith open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Arith open VG.Proof.MlDsa.Arith -open VG.Proof.MlKem.X86_64 (Keep WP.keep writesOnly gprPreserved_of wp_counted ifp ifn) +open VG.Proof.MlKem.X86_64 (Keep XOnly xmm_setXmm ifp ifn sel GOnly wp_rcxLoop add_ofNat_zero) +open VG.Impl.MlKem.X86_64 (xb xmov) open VG.Spec.MlDsa (q n Poly Zq coeffAt polyAt Reduced PolyIs zetas nttInv) -namespace NttInv - -open Ntt (LI) - -/-- The chain of zeta indices of the layers `ls` of `NTT⁻¹`, from `k`. -/ -def Chain : Nat → List Nat → Prop - | _, [] => True - | k, len :: ls => k = 256 / len - 1 ∧ Chain (256 / len - 1 - 128 / len) ls - -theorem lens_inv : ∀ len ∈ nttLens, 128 / len + 1 ≤ 256 / len ∧ 256 / len ≤ 256 := by decide - -theorem sx_m4 : BitVec.signExtend 64 (-4 : BitVec 32) = -4 := by decide - -theorem step_down (zP : Addr) (m : Nat) : - coeffAddr zP (m + 1) + BitVec.signExtend 64 (-4 : BitVec 32) = coeffAddr zP m := by - rw [sx_m4, ← coeffAddr_succ, BitVec.add_assoc, show (4 : BitVec 64) + -4 = 0#64 by decide, BitVec.add_zero] - -theorem negZetaTab_of : TabOf negZetaTab fun k => -zetas k := fun k _ => negZetaNat_eq k - -theorem lays_ok {s₀ : State} {fP zP : Addr} (hw : pR fP ∈ s₀.wr) (hz : pR zP ∈ s₀.rd ++ s₀.wr) - (hd : (pR zP).Disjoint (pR fP)) : - ∀ (ls : List Nat) (F : Poly) (k : Nat) (s : State), (∀ len ∈ ls, len ∈ nttLens) → Chain k ls → - LI negZetaTab s₀ fP zP F k s → - WP isa (nttInvLays ls) s fun s' => ∃ k', LI negZetaTab s₀ fP zP (ls.foldl nttInvLayer F) k' s' - | [], F, k, s, _, _, hI => WP.block_nil ⟨k, hI⟩ - | len :: ls, F, k, s, hls, ⟨hk, hc⟩, hI => by - have hlen := hls len (List.mem_cons_self ..) - obtain ⟨h1, h2⟩ := lens_inv len hlen - refine WP.seq (WP.mono (lay_ok bflyInv_spec negZetaTab_of hlen (-4) (fun c => 256 / len - 1 - c) - (fun c hc => by omega) - (fun c hc => by - rw [show 256 / len - 1 - c = (256 / len - 1 - (c + 1)) + 1 by omega] - exact step_down zP _) - F s hI.rsi (by rw [hI.r8, hk]; rfl) hI.poly (by rw [hI.wr]; exact hw) (by rw [hI.rd, hI.wr]; exact hz) hd - hI.tab) fun s' ⟨⟨hP, hf, hsi, h8⟩, hk'⟩ => ?_) - exact lays_ok hw hz hd ls _ _ s' (fun l hl => hls l (List.mem_cons_of_mem _ hl)) hc - (hI.step hP hf hsi h8 hk' hd) - -theorem chain_inv : Chain 255 nttInvLens := by - simp only [nttInvLens, Chain]; decide - -/-! ## The multiplication by 8347681 -/ - -theorem scaleHead_ok (s : State) (h : InRegions (s.rd ++ s.wr) (s.gpr .rsi) 4) : - WP isa (.block ([.mov32 .rax (.mem (at_ .rsi 0)), .mul .r9] : List Instr)) s fun s' => - (s'.gpr .rax = prodW (s.mem.readW (s.gpr .rsi) 32) (s.gpr .r9) ∧ s'.mem = s.mem) ∧ - Keep [.rax, .rdx] s s' := by - refine WP.keep _ ?_ (by decide) - xrund [h] - -theorem scaleBody_ok (s : State) (h : InRegions (s.rd ++ s.wr) (s.gpr .rsi) 4) (w : InRegions s.wr (s.gpr .rsi) 4) : - WP isa (.block scaleBody) s fun s' => - (s'.mem = s.mem.writeW (s.gpr .rsi) (BitVec.setWidth 32 (redD (prodW (s.mem.readW (s.gpr .rsi) 32) - (s.gpr .r9)))) ∧ s'.gpr .rsi = s.gpr .rsi + 4 ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ - s'.zf = some (s.gpr .rcx - 1 == 0)) ∧ Keep [.rax, .rdx, .rsi, .rcx, .r10, .r11] s s' := by - unfold scaleBody - simp only [List.append_assoc] +/-! ## The layers -/ + +/-- The block of `NTT⁻¹`. -/ +abbrev invBlk : Poly → Nat → Nat → Nat → Nat → Poly := fun f len k st t => blockN bflyInv f len (-zetas k) st t + +theorem nttInvLayer_eq (F : Poly) (len : Nat) : + nttInvLayer F len = layF invBlk F len (fun c => 256 / len - 1 - c) (128 / len) := rfl + +/-- A layer of `NTT⁻¹` with `len ≥ 4`, whose first zeta is `zetas k`. -/ +theorem invLay_ok {fP sP : Addr} (hd : (pR sP).Disjoint (pR fP)) (len k : Nat) + (hlen : len ∈ [4, 8, 16, 32, 64, 128]) (hk : 256 / len - 1 = k) + {F : Poly} (s : State) (hc : VConsts s) (hdi : s.gpr .rdi = fP) (hsi : s.gpr .rsi = sP) + (hS : PolyIs s.mem fP F) (hT : Tab zmTab s.mem sP 256) (hwf : pR fP ∈ s.wr) (hw : pR sP ∈ s.wr) : + WP isa (vlay vibfly len k (-4)) s fun s' => PolyIs s'.mem fP (nttInvLayer F len) ∧ BInv fP s s' := by + have hl : 128 / len ≥ 1 ∧ 256 / len = 2 * (128 / len) ∧ 256 / len ≤ 64 := by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hlen + rcases hlen with rfl | rfl | rfl | rfl | rfl | rfl <;> decide + rw [nttInvLayer_eq] + exact vlay_ok vibfly_spec nttInvBlk_ok hlen (-4) (fun c => 256 / len - 1 - c) (by rw [hk]; rfl) + (fun c _ => by omega) + (fun c hc' => (congrArg (· + _) (congrArg (coeffAddr sP) (show 256 / len - 1 - c = + 256 / len - 1 - (c + 1) + 1 by omega))).trans (step_bwd _ _ 1 (by decide))) + hc hdi hsi hS hT hwf hw hd + +theorem invLay2_ok {fP sP : Addr} (hd : (pR sP).Disjoint (pR fP)) {F : Poly} (s : State) (hc : VConsts s) + (hdi : s.gpr .rdi = fP) (hsi : s.gpr .rsi = sP) (hS : PolyIs s.mem fP F) (hT : Tab zmTab s.mem sP 256) + (hwf : pR fP ∈ s.wr) (hw : pR sP ∈ s.wr) : + WP isa (vlay2 vibfly 126 0x05 (-8)) s fun s' => PolyIs s'.mem fP (nttInvLayer F 2) ∧ BInv fP s s' := by + have hs : ∀ e < 4, sel 0x05 e = 1 - e / 2 := by decide + rw [nttInvLayer_eq, show 256 / 2 - 1 = 127 from rfl, show 128 / 2 = 64 from rfl] + exact vlay2_ok vibfly_spec nttInvBlk_ok 126 0x05 (-8) (fun c => 127 - c) (fun i => 126 - 2 * i) rfl + (fun i _ => by omega) (fun i hi e he => by rw [hs e he]; omega) + (fun i hi => (congrArg (· + _) (congrArg (coeffAddr sP) (show 126 - 2 * i = 126 - 2 * (i + 1) + 2 by + omega))).trans (step_bwd _ _ 2 (by decide))) hc hdi hsi hS hT hwf hw hd + +theorem invLay1_ok {fP sP : Addr} (hd : (pR sP).Disjoint (pR fP)) {F : Poly} (s : State) (hc : VConsts s) + (hdi : s.gpr .rdi = fP) (hsi : s.gpr .rsi = sP) (hS : PolyIs s.mem fP F) (hT : Tab zmTab s.mem sP 256) + (hwf : pR fP ∈ s.wr) (hw : pR sP ∈ s.wr) : + WP isa (vlay1 vibfly 252 0x1B (-16)) s fun s' => PolyIs s'.mem fP (nttInvLayer F 1) ∧ BInv fP s s' := by + have hs : ∀ e < 4, sel 0x1B e = 3 - e := by decide + rw [nttInvLayer_eq, show 256 / 1 - 1 = 255 from rfl, show 128 / 1 = 128 from rfl] + exact vlay1_ok vibfly_spec nttInvBlk_ok 252 0x1B (-16) (fun c => 255 - c) (fun i => 252 - 4 * i) rfl + (fun i _ => by omega) (fun i hi e he => by rw [hs e he]; omega) + (fun i hi => (congrArg (· + _) (congrArg (coeffAddr sP) (show 252 - 4 * i = 252 - 4 * (i + 1) + 4 by + omega))).trans (step_bwd _ _ 4 (by decide))) hc hdi hsi hS hT hwf hw hd + +/-! ## The multiplication by `256⁻¹` -/ + +/-- The coefficients of `G` before `4i` multiplied by `8347681`. -/ +def Scaled (m : Mem) (fP : Addr) (G : Poly) (i : Nat) : Prop := + ∀ k < 256, (coeffAt m fP k).toNat = (if k < 4 * i then G[k]! * 8347681 else G[k]!).val + +/-- `8347681 · 2³² mod q` in the doublewords of `xmm13` and `xmm12`. -/ +theorem scale_zlanes (x : BitVec 128) (hx : x = shufDwords ((0 : BitVec 64) ++ BitVec.setWidth 64 16382#32) 0) : + ZLanes x (fun _ => 8347681) ∧ ZOdd x x := by + subst hx + refine ⟨fun i hi => ?_, fun j hj => ?_⟩ + · rcases cases4 hi with rfl | rfl | rfl | rfl <;> decide + · rcases (by omega : j = 0 ∨ j = 1) with rfl | rfl <;> decide + +/-- The body of the loop of `vscale`. -/ +abbrev sbody : List Instr := + [.movdquLoad .xmm3 (at_ .rdx 0)] ++ vmont .xmm3 .xmm13 .xmm12 .xmm2 .xmm4 ++ + vcsub .xmm3 .xmm2 ++ [.movdquStore (at_ .rdx 0) .xmm3, .alu .add .rdx (.imm 16)] ++ + [.alu .sub .rcx (.imm 1)] + +/-- The product of the doublewords of `xmm3` by the constant, reduced. -/ +theorem vmul3_ok {s : State} (hc : VConsts s) : + WP isa (.block (vmont .xmm3 .xmm13 .xmm12 .xmm2 .xmm4 ++ vcsub .xmm3 .xmm2)) s fun s' => + s'.xmm .xmm3 = csubV (montV (s.xmm .xmm3) (s.xmm .xmm13) (s.xmm .xmm12)) ∧ + XOnly [.xmm3, .xmm2, .xmm4] s s' := by + simp only [vmont, vredc, vcsub, vcadd, xmov, xb, List.cons_append, List.nil_append] + vrun [eval_movdqa] + rw [hc.q, hc.qinv] + exact ⟨rfl, by xonly⟩ + +theorem vscale_step {fP : Addr} {G : Poly} {i : Nat} (hi : i < 64) {s : State} (hc : VConsts s) + (hz : ZLanes (s.xmm .xmm13) (fun _ => 8347681)) (ho : ZOdd (s.xmm .xmm13) (s.xmm .xmm12)) + (hdx : s.gpr .rdx = coeffAddr fP (4 * i)) (hS : Scaled s.mem fP G i) (hw : pR fP ∈ s.wr) : + WP isa (.block sbody) s fun s' => + Scaled s'.mem fP G (i + 1) ∧ s'.gpr .rdx = coeffAddr fP (4 * (i + 1)) ∧ + Frame [pR fP] s.mem s'.mem ∧ VConsts s' ∧ s'.xmm .xmm13 = s.xmm .xmm13 ∧ + s'.xmm .xmm12 = s.xmm .xmm12 ∧ Keep [.rdx, .rcx] s s' ∧ + s'.gpr .rcx = s.gpr .rcx - 1 ∧ s'.zf = some (s.gpr .rcx - 1 == 0) ∧ s'.mxcsr = s.mxcsr := by + have j0 : 4 * i + 4 ≤ 256 := by omega + have r0 : InRegions (s.rd ++ s.wr) (coeffAddr fP (4 * i)) 16 := f_in (List.mem_append_right _ hw) j0 + have w0 := f_in hw j0 + have hx : DLanes (s.mem.readW (coeffAddr fP (4 * i)) 128) (fun e => G[4 * i + e]!) := fun e he => by + rw [dword_readW _ _ he, coeffAddr_add, ← coeffAt_eq, hS _ (by omega), ifn (by omega)] + rw [show sbody = [.movdquLoad .xmm3 (at_ .rdx 0)] ++ ((vmont .xmm3 .xmm13 .xmm12 .xmm2 .xmm4 ++ + vcsub .xmm3 .xmm2) ++ [.movdquStore (at_ .rdx 0) .xmm3, .alu .add .rdx (.imm 16), .alu .sub .rcx (.imm 1)]) + from rfl, WP.block_append_iff] + vrund [hdx, r0] rw [WP.block_append_iff] - refine WP.mono (scaleHead_ok s h) fun s1 ⟨⟨ha, hm1⟩, k1⟩ => ?_ - rw [WP.block_append_iff] - refine WP.mono (reduce_ok s1) fun s2 ⟨⟨hr, hm2⟩, k2⟩ => ?_ - have k12 := k1.trans k2 - have hsi : s2.gpr .rsi = s.gpr .rsi := k12.gpr (by decide) - have hcx : s2.gpr .rcx = s.gpr .rcx := k12.gpr (by decide) - refine WP.mono (WP.keep [.rsi, .rcx] (Q := fun s' => s'.mem = s2.mem.writeW (s2.gpr .rsi) - (BitVec.setWidth 32 (s2.gpr .r10)) ∧ s'.gpr .rsi = s2.gpr .rsi + 4 ∧ s'.gpr .rcx = s2.gpr .rcx - 1 ∧ - s'.zf = some (s2.gpr .rcx - 1 == 0)) (by - xrund [show InRegions s2.wr (s2.gpr .rsi) 4 by rw [k12.2.2, hsi]; exact w]) - (by decide)) fun s3 ⟨⟨hm3, h3si, h3cx, h3z⟩, k3⟩ => ⟨?_, (k12.trans k3).mono (by decide)⟩ - rw [hm3, h3si, h3cx, h3z, hcx, hsi, hr, hm2, hm1, ha] - exact ⟨rfl, rfl, rfl, rfl⟩ - -/-- After `i` coefficients of `G` multiplied by 8347681, from the state `sL`. -/ -structure SInv (sL : State) (fP : Addr) (G : Poly) (i : Nat) (s : State) : Prop where - rsi : s.gpr .rsi = coeffAddr fP i - r9 : s.gpr .r9 = BitVec.ofNat 64 (8347681 : Zq).val - rd : s.rd = sL.rd - wr : s.wr = sL.wr - frame : Frame [pR fP] sL.mem s.mem - coeff : ∀ k < 256, (coeffAt s.mem fP k).toNat = if k < i then (G[k]! * 8347681).val else (G[k]!).val - keep : Keep [.r9, .rcx, .rax, .rdx, .rsi, .rcx, .r10, .r11] sL s - -theorem scale_step {sL : State} {fP : Addr} {G : Poly} (hw : pR fP ∈ sL.wr) {i : Nat} (hi : i < 256) - {s : State} (hI : SInv sL fP G i s) : - WP isa (.block scaleBody) s fun s' => SInv sL fP G (i + 1) s' ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ - s'.zf = some (s.gpr .rcx - 1 == 0) := by - have hw' : pR fP ∈ s.wr := by rw [hI.wr]; exact hw - refine WP.mono (scaleBody_ok s (by rw [hI.rsi]; exact ⟨_, List.mem_append_right _ hw', coeff_contains _ hi⟩) - (by rw [hI.rsi]; exact ⟨_, hw', coeff_contains _ hi⟩)) fun s' ⟨⟨hm, hsi, hcx, hz⟩, hk⟩ => ⟨?_, hcx, hz⟩ - have hv : (BitVec.setWidth 32 (redD (prodW (coeffAt s.mem fP i) (s.gpr .r9)))).toNat = - (G[i]! * 8347681).val := by - rw [hI.r9, redD_prodW (x := G[i]!) (by rw [hI.coeff i hi, ifn (Nat.lt_irrefl i)]), val_mul, val_mul, - Nat.mul_comm] - rw [hI.rsi, ← coeffAt_eq] at hm - refine ⟨by rw [hsi, hI.rsi, coeffAddr_succ], by rw [hk.gpr (by decide), hI.r9], hk.2.1.trans hI.rd, - hk.2.2.trans hI.wr, by rw [hm]; exact hI.frame.writeW (List.mem_singleton_self _) _ (coeff_contains _ hi), - fun k hk' => ?_, (hI.keep.trans hk).mono (by decide)⟩ - rw [hm, coeffAt_writeW _ _ hk' hi] - by_cases e : i = k - · subst e; rw [ifp rfl, ifp (Nat.lt_succ_self _)]; exact hv - · rw [ifn e, hI.coeff k hk'] - by_cases h : k < i - · rw [ifp h, ifp (by omega)] - · rw [ifn h, ifn (by omega)] - -/-- The multiplication of every coefficient by 8347681. -/ -theorem scale_ok {fP : Addr} {G : Poly} (sL : State) (hsi : sL.gpr .rsi = fP) (hG : PolyIs sL.mem fP G) - (hw : pR fP ∈ sL.wr) : - WP isa (.seq (.block [.mov32 .r9 (.imm 8347681)]) - (.seq (.block [.mov32 .rcx (.imm 256)]) (.loop (.block scaleBody) .ne))) sL fun s' => - PolyIs s'.mem fP (G.map (· * 8347681)) ∧ Frame [pR fP] sL.mem s'.mem ∧ - Keep [.r9, .rcx, .rax, .rdx, .rsi, .rcx, .r10, .r11] sL s' := by - refine WP.seq (WP.mono (WP.keep [.r9] - (Q := fun s => s.mem = sL.mem ∧ s.gpr .r9 = BitVec.ofNat 64 (8347681 : Zq).val) - (by xrund) (by decide)) fun s3 ⟨⟨hm3, h93⟩, k3⟩ => ?_) - refine WP.mono (wp_counted (s₀ := s3) (N := 256) (v := 256) rfl (by decide) (SInv sL fP G) - (fun s4 hm4 hk4 => ⟨by rw [hk4.gpr (by decide), k3.gpr (by decide), hsi]; simp, - by rw [hk4.gpr (by decide), h93], hk4.2.1.trans k3.2.1, hk4.2.2.trans k3.2.2, - by rw [hm4, hm3]; exact Frame.refl _ _, - fun k hk => by rw [ifn (Nat.not_lt_zero _), hm4, hm3]; exact polyIs_toNat hG hk, - (k3.trans hk4).mono (by decide)⟩) - fun i hi s hI => scale_step hw hi hI) fun s' hI => ⟨?_, hI.frame, hI.keep⟩ - refine polyIs_of_toNat fun k hk => ?_ - rw [hI.coeff k hk, ifp hk, map_mul_get _ _ hk] - -end NttInv + refine WP.mono (vmul3_ok (hc.setXmm (by decide) (by decide) _)) fun s2 ⟨h3, o2⟩ => ?_ + have c2 := xonly_vconsts o2 (hc.setXmm (by decide) (by decide) _) (by decide) (by decide) + have g2 : s2.gpr = s.gpr := o2.gpr + have m2 : s2.mem = s.mem := o2.mem + have e2 : s2.rd = s.rd ∧ s2.wr = s.wr := ⟨o2.rd, o2.wr⟩ + have x2 : s2.mxcsr = s.mxcsr := o2.mxcsr + have z2 : s2.xmm .xmm13 = s.xmm .xmm13 := by rw [o2.xmm _ (by decide), xmm_setXmm]; rfl + have z2' : s2.xmm .xmm12 = s.xmm .xmm12 := by rw [o2.xmm _ (by decide), xmm_setXmm]; rfl + rw [xmm_setXmm, ifp rfl, xmm_setXmm, ifn (by decide), xmm_setXmm, ifn (by decide)] at h3 + generalize hV : csubV (montV (s.mem.readW (coeffAddr fP (4 * i)) 128) (s.xmm .xmm13) (s.xmm .xmm12)) = V at h3 + vrund [g2, m2, e2.1, e2.2, hdx, w0, x2, h3] + refine ⟨fun k hk => ?_, by rw [show (16 : BitVec 64) = BitVec.ofNat 64 (4 * 4) from rfl, coeffAddr_add, + Nat.mul_succ], (Frame.refl _ _).writeW (List.mem_singleton_self _) _ (pR_contains fP j0), ⟨?_, ?_⟩, ?_, ?_, + ⟨fun r hr => ?_, rfl, rfl⟩⟩ + · rw [coeffAt_write128 _ _ j0 _ (by omega)] + split + · rename_i h + rw [ifp (by omega), ← hV, dword_csubV _ (by omega), mulZ (hx _ (by omega)) (hz _ (by omega)) + (dword_montV ho (prod_lt hx hz) (by omega))] + dsimp only + rw [show 4 * i + (k - 4 * i) = k by omega, Fin.mul_comm] + · rename_i h + rw [hS k hk] + by_cases h' : k < 4 * i + · rw [ifp h', ifp (by omega)] + · rw [ifn h', ifn (by omega)] + · simp only [RegUpd.xmm_setReg, RegUpd.xmm_setFlags]; exact c2.q + · simp only [RegUpd.xmm_setReg, RegUpd.xmm_setFlags]; exact c2.qinv + · exact z2 + · exact z2' + · simp only [List.mem_cons, List.not_mem_nil, or_false, not_or] at hr + simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, hr.1, hr.2, ite_false] + +/-- Every coefficient times `8347681`. -/ +theorem vscale_ok {fP sP : Addr} (_hd : (pR sP).Disjoint (pR fP)) {G : Poly} (s : State) (hc : VConsts s) + (hdi : s.gpr .rdi = fP) (_hsi : s.gpr .rsi = sP) (hS : PolyIs s.mem fP G) (_hT : Tab zmTab s.mem sP 256) + (hwf : pR fP ∈ s.wr) (_hw : pR sP ∈ s.wr) : + WP isa vscale s fun s' => PolyIs s'.mem fP (G.map (· * 8347681)) ∧ BInv fP s s' := by + refine WP.seq (WP.mono (Q := fun (w : State) => w.gpr .rdx = fP ∧ + w.xmm .xmm13 = shufDwords ((0 : BitVec 64) ++ BitVec.setWidth 64 16382#32) 0 ∧ + w.xmm .xmm12 = w.xmm .xmm13 ∧ VConsts w ∧ Keep [.rdx, .rax] s w ∧ w.mem = s.mem ∧ w.mxcsr = s.mxcsr) + (by + vrund [hdi, eval_movdqa] + refine ⟨⟨?_, ?_⟩, ⟨fun r hr => ?_, rfl, rfl⟩⟩ + · simp only [RegUpd.xmm_setReg, xmm_setXmm, reduceCtorEq, ite_false]; exact hc.q + · simp only [RegUpd.xmm_setReg, xmm_setXmm, reduceCtorEq, ite_false]; exact hc.qinv + · simp only [List.mem_cons, List.not_mem_nil, or_false, not_or] at hr + simp only [RegUpd.gpr_setReg, RegUpd.gpr_setXmm, hr.1, hr.2, ite_false]) fun w ⟨hdx, h13, h12, hcw, kw, mw, xw⟩ => ?_) + obtain ⟨hz, ho⟩ := scale_zlanes _ h13 + replace ho : ZOdd (w.xmm .xmm13) (w.xmm .xmm12) := by rw [h12]; exact ho + refine WP.mono (wp_rcxLoop (N := 64) (by decide) (by decide) + (fun i u => Scaled u.mem fP G i ∧ u.gpr .rdx = coeffAddr fP (4 * i) ∧ VConsts u ∧ + u.xmm .xmm13 = w.xmm .xmm13 ∧ u.xmm .xmm12 = w.xmm .xmm12 ∧ Keep [.rcx, .rdx] w u ∧ + Frame [pR fP] w.mem u.mem ∧ u.mxcsr = w.mxcsr) + (fun u o _ => ⟨fun k hk => by rw [o.mem, mw, ifn (by omega)]; exact polyIs_toNat hS (by rw [n_eq]; exact hk), + by rw [o.keep.gpr (by decide), hdx, Nat.mul_zero, coeffAddr, Nat.mul_zero, add_ofNat_zero], + gonly_vconsts o hcw, by rw [o.xmm], by rw [o.xmm], o.keep.mono (by simp), by rw [o.mem]; exact Frame.refl _ _, + o.mxcsr⟩) + (fun i hi u ⟨hS', hdx', hc', hz', hzo', hk', hf', hx'⟩ => WP.mono (vscale_step hi hc' (by rw [hz']; exact hz) + (by rw [hz', hzo']; exact ho) hdx' hS' (by rw [hk'.2.2, kw.2.2]; exact hwf)) + fun u' ⟨hS'', hdx'', hf'', hc'', hz'', hzo'', hk'', hcx, hzf, hx''⟩ => + ⟨⟨hS'', hdx'', hc'', by rw [hz'', hz'], by rw [hzo'', hzo'], (hk'.trans hk'').mono (by simp), + hf'.trans hf'', by rw [hx'', hx']⟩, hcx, hzf⟩)) fun u ⟨hS', _, hc', _, _, hk', hf', hx'⟩ => ?_ + refine ⟨polyIs_of_toNat fun k hk => ?_, ⟨(kw.trans hk').mono (by simp), by rw [← mw]; exact hf', hc', + by rw [hx', xw]⟩⟩ + rw [n_eq] at hk + rw [hS' k hk, ifp (by omega), map_mul_get _ _ (by rw [n_eq]; exact hk)] theorem nttInv_correct (s : State) (hs : (inPlaceK nttInv).pre s) : ∃ t s', Exec isa Impl.MlDsa.X86_64.Arith.nttInv s t s' ∧ abiPreserved s s' ∧ (inPlaceK nttInv).post s s' := by - have hw : pR (s.gpr .rdi) ∈ s.wr := by rw [hs.2.1]; simp - have hz : pR (s.gpr .rsi) ∈ s.rd ++ s.wr := by rw [hs.1, hs.2.1]; simp - obtain ⟨t, s', he, ⟨hP, hf⟩, hk⟩ := WP.keep (c := Impl.MlDsa.X86_64.Arith.nttInv) - [.rax, .rcx, .rdx, .rsi, .rdi, .r8, .r9, .r10, .r11] - (Q := fun s' => PolyIs s'.mem (s.gpr .rdi) (nttInv (polyAt s.mem (s.gpr .rdi))) ∧ - Frame [pR (s.gpr .rdi), pR (s.gpr .rsi)] s.mem s'.mem) - (WP.seq (WP.mono (Ntt.pro_ok hs negZetaTab (4 * 255) 255 (by decide)) fun s1 hI => - WP.seq (WP.mono (NttInv.lays_ok hw hz hs.2.2.1.symm nttInvLens _ 255 s1 (fun _ h => by - simp only [nttInvLens, nttLens, List.mem_cons, List.not_mem_nil, or_false] at h ⊢; omega) - NttInv.chain_inv hI) fun s2 ⟨k, hI2⟩ => - WP.mono (NttInv.scale_ok s2 hI2.rsi hI2.poly (by rw [hI2.wr]; exact hw)) fun s' ⟨hP, hf, _⟩ => - ⟨by rw [nttInv_eq_layers]; exact hP, hI2.frame.trans (hf.mono (by simp))⟩))) (by decide +kernel) - exact ⟨t, s', he, abiPreserved_of_exec (by decide +kernel) he (gprPreserved_of hk (by decide) hf - (by simpa using ⟨hs.2.2.2.1, hs.2.2.2.2.1⟩)), hP⟩ + have hd : (pR (s.gpr .rsi)).Disjoint (pR (s.gpr .rdi)) := hs.2.2.1.symm + refine mx_correct s hs (by decide +kernel) (by decide +kernel) (by decide +kernel) fun s1 k1 f1 => + WP.mono (nttBody_ok hs k1 f1 + (G := (nttInvLens.foldl nttInvLayer (polyAt s.mem (s.gpr .rdi))).map (· * 8347681)) + fun s2 hI hdi hsi hwf hw => ?_) fun s' ⟨hP, hf⟩ => ⟨by rw [nttInv_eq_layers]; exact hP, hf⟩ + simp only [nttInvLens, List.foldl_cons, List.foldl_nil] + refine LI.seq hdi hsi hwf hw hd (invLay1_ok hd) ?_ hI + refine fun _ hI => LI.seq hdi hsi hwf hw hd (invLay2_ok hd) ?_ hI + refine fun _ hI => LI.seq hdi hsi hwf hw hd (invLay_ok hd 4 63 (by decide) (by decide)) ?_ hI + refine fun _ hI => LI.seq hdi hsi hwf hw hd (invLay_ok hd 8 31 (by decide) (by decide)) ?_ hI + refine fun _ hI => LI.seq hdi hsi hwf hw hd (invLay_ok hd 16 15 (by decide) (by decide)) ?_ hI + refine fun _ hI => LI.seq hdi hsi hwf hw hd (invLay_ok hd 32 7 (by decide) (by decide)) ?_ hI + refine fun _ hI => LI.seq hdi hsi hwf hw hd (invLay_ok hd 64 3 (by decide) (by decide)) ?_ hI + refine fun _ hI => LI.seq hdi hsi hwf hw hd (invLay_ok hd 128 1 (by decide) (by decide)) ?_ hI + exact fun _ hI => LI.last hdi hsi hwf hw hd (vscale_ok hd) hI theorem nttInv_ct : ConstantTime isa (inPlaceK nttInv).pre (inPlaceK nttInv).pub Impl.MlDsa.X86_64.Arith.nttInv := VG.Taint.constantTime (A := taint) (X86_64.Taint.ofRegs [.rdi, .rsi, .rsp]) inPlace_agree (by taint_decide) diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/NttLoop.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/NttLoop.lean deleted file mode 100644 index eac150643..000000000 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/NttLoop.lean +++ /dev/null @@ -1,165 +0,0 @@ -import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.NttBfly -import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Table - -/-! -# ML-DSA on x86-64: the blocks and layers of `NTT` and `NTT⁻¹` - -Untrusted: everything here is checked by Lean. The loops of `nttBlk` and -`nttLay`, for any butterfly code that does what a butterfly `op` of the -specification does (`BflyOk`): a block runs `len` butterflies (`blockN`), -and a layer its `128 / len` blocks (`layerN`), with the zetas `Z (zi c)`, -whose values `tab` the table at `zP` holds. --/ - -namespace VG.Proof.MlDsa.X86_64.Arith - -open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Arith -open VG.Proof.MlDsa.Arith -open VG.Proof.MlKem.X86_64 (Keep WP.keep wp_countdown toNat_setWidth64) -open VG.Spec.MlDsa (q n Poly Zq coeffAt polyAt Reduced PolyIs) - -/-- The table `tab` holds the values of the zetas `Z`. -/ -def TabOf (tab : Nat → Nat) (Z : Nat → Zq) : Prop := ∀ k < 256, tab k = (Z k).val - -/-! ## A block -/ - -section -variable {code : Nat → List Instr} {op : Poly → Nat → Nat → Zq → Poly} (hb : BflyOk code op) -include hb - -/-- The `len` butterflies of a block. -/ -theorem bflys_ok {fP : Addr} {len start : Nat} (hlen : 0 < len) (hs : start + 2 * len ≤ 256) (z : Zq) - (G : Poly) (s : State) (hsi : s.gpr .rsi = coeffAddr fP start) (h9 : s.gpr .r9 = BitVec.ofNat 64 z.val) - (hG : PolyIs s.mem fP G) (hw : pR fP ∈ s.wr) (hc : s.gpr .rcx = BitVec.ofNat 64 len) : - WP isa (.loop (.block (code len)) .ne) s fun s' => - PolyIs s'.mem fP (blockN op G len z start len) ∧ Frame [pR fP] s.mem s'.mem ∧ - s'.gpr .rsi = coeffAddr fP (start + len) ∧ Keep [.rax, .rdx, .rsi, .rcx, .r10, .r11] s s' := by - refine wp_countdown (cnt := .rcx) (N := len) (by omega) hlen (fun t s' => - PolyIs s'.mem fP (blockN op G len z start t) ∧ Frame [pR fP] s.mem s'.mem ∧ - s'.gpr .rsi = coeffAddr fP (start + t) ∧ Keep [.rax, .rdx, .rsi, .rcx, .r10, .r11] s s') - (fun t ht s' ⟨hP, hf, hs', hk⟩ _ => ?_) (fun _ h => h) ⟨hG, Frame.refl _ _, hsi, Keep.refl _ _⟩ hc - refine WP.mono (hb fP len (start + t) hlen (by omega) z _ s' hs' (by rw [hk.gpr (by decide), h9]) hP - (by rw [hk.2.2]; exact hw)) fun s'' ⟨⟨hP', hf', hsi', hcx, hz⟩, hk'⟩ => - ⟨⟨?_, hf.trans hf', by rw [hsi', hs', coeffAddr_succ, Nat.add_assoc], (hk.trans hk').mono (by decide)⟩, - hcx, hz⟩ - rw [blockN_succ]; exact hP' - -omit hb in -theorem blkPre_ok (len : Nat) (dz : BitVec 32) (s : State) (h : InRegions (s.rd ++ s.wr) (s.gpr .r8) 4) : - WP isa (.block [.mov32 .r9 (.mem (at_ .r8 0)), .alu .add .r8 (.imm dz), - .mov32 .rcx (.imm (BitVec.ofNat 32 len))]) s fun s' => - (s'.gpr .r9 = BitVec.setWidth 64 (s.mem.readW (s.gpr .r8) 32) ∧ s'.gpr .r8 = s.gpr .r8 + BitVec.signExtend 64 dz ∧ - s'.gpr .rcx = BitVec.setWidth 64 (BitVec.ofNat 32 len) ∧ s'.mem = s.mem) ∧ Keep [.r9, .r8, .rcx] s s' := by - refine WP.keep _ ?_ (by rfl) - xrund [h] - -omit hb in -theorem blkPost_ok (len : Nat) (hl : 4 * len < 2 ^ 31) (s : State) : - WP isa (.block [.alu .add .rsi (.imm (BitVec.ofNat 32 (4 * len))), .alu .sub .rdi (.imm 1)]) s fun s' => - (s'.gpr .rsi = s.gpr .rsi + BitVec.ofNat 64 (4 * len) ∧ s'.gpr .rdi = s.gpr .rdi - 1 ∧ - s'.zf = some (s.gpr .rdi - 1 == 0) ∧ s'.mem = s.mem) ∧ Keep [.rsi, .rdi] s s' := by - refine WP.keep _ ?_ (by rfl) - xrund [sx_ofNat hl] - -/-- A block, with the zeta `Z k` at `r8`. -/ -theorem blk_ok {tab : Nat → Nat} {Z : Nat → Zq} (hZ : TabOf tab Z) {fP zP : Addr} {len start k : Nat} - (hlen : 0 < len) (hl : len ≤ 128) (hs : start + 2 * len ≤ 256) - (hk : k < 256) (dz : BitVec 32) (G : Poly) (s : State) (hsi : s.gpr .rsi = coeffAddr fP start) - (h8 : s.gpr .r8 = coeffAddr zP k) (hG : PolyIs s.mem fP G) (hw : pR fP ∈ s.wr) - (hz : pR zP ∈ s.rd ++ s.wr) (ht : Tab tab s.mem zP 256) : - WP isa (nttBlk (code len) len dz) s fun s' => - (PolyIs s'.mem fP (blockN op G len (Z k) start len) ∧ Frame [pR fP] s.mem s'.mem ∧ - s'.gpr .rsi = coeffAddr fP (start + 2 * len) ∧ s'.gpr .r8 = s.gpr .r8 + BitVec.signExtend 64 dz ∧ - s'.gpr .rdi = s.gpr .rdi - 1 ∧ s'.zf = some (s.gpr .rdi - 1 == 0)) ∧ - Keep [.r9, .r8, .rcx, .rax, .rdx, .rsi, .rcx, .r10, .r11, .rsi, .rdi] s s' := by - refine WP.seq (WP.mono (blkPre_ok len dz s (by rw [h8]; exact ⟨_, hz, coeff_contains _ (show k < 256 by omega)⟩)) - fun s1 ⟨⟨h9, h8', hc, hm⟩, k1⟩ => ?_) - have hz9 : s1.gpr .r9 = BitVec.ofNat 64 (Z k).val := by - rw [h9, h8, ← coeffAt_eq, ht k hk, ← hZ k hk] - apply BitVec.eq_of_toNat_eq - rw [toNat_setWidth64, BitVec.toNat_ofNat, BitVec.toNat_ofNat] - have := val_lt (Z k) - rw [hZ k hk, Nat.mod_eq_of_lt (by omega), Nat.mod_eq_of_lt (by omega)] - refine WP.seq (WP.mono (bflys_ok hb (fP := fP) hlen hs (Z k) G s1 (by rw [k1.gpr (by decide), hsi]) hz9 - (by rw [hm]; exact hG) (by rw [k1.2.2]; exact hw) (by - rw [hc]; apply BitVec.eq_of_toNat_eq - rw [toNat_setWidth64, BitVec.toNat_ofNat, BitVec.toNat_ofNat]; omega)) fun s2 ⟨hP, hf, hsi2, k2⟩ => ?_) - refine WP.mono (blkPost_ok len (Nat.lt_of_le_of_lt (Nat.mul_le_mul_left 4 hl) (by decide)) s2) - fun s3 ⟨⟨hsi3, hdi, hz3, hm3⟩, k3⟩ => - ⟨⟨by rw [hm3]; exact hP, by rw [hm3, ← hm]; exact hf, ?_, ?_, ?_, ?_⟩, ((k1.trans k2).trans k3).mono (by decide)⟩ - · rw [hsi3, hsi2, coeffAddr_add, show start + len + len = start + 2 * len by omega] - · rw [k3.gpr (by decide), k2.gpr (by decide), h8'] - · rw [hdi, k2.gpr (by decide), k1.gpr (by decide)] - · rw [hz3, k2.gpr (by decide), k1.gpr (by decide)] - -/-! ## A layer -/ - -omit hb in -theorem layPre_ok (c : Nat) (hl : c < 2 ^ 31) (s : State) : - WP isa (.block [.mov32 .rdi (.imm (BitVec.ofNat 32 c))]) s fun s' => - (s'.gpr .rdi = BitVec.ofNat 64 c ∧ s'.mem = s.mem) ∧ Keep [.rdi] s s' := by - refine WP.keep _ ?_ (by rfl) - xrund - apply BitVec.eq_of_toNat_eq - rw [toNat_setWidth64, BitVec.toNat_ofNat, BitVec.toNat_ofNat]; omega - -omit hb in -theorem layPost_ok (s : State) : - WP isa (.block [.alu .sub .rsi (.imm 1024)]) s fun s' => - (s'.gpr .rsi = s.gpr .rsi - 1024 ∧ s'.mem = s.mem) ∧ Keep [.rsi] s s' := by - refine WP.keep _ ?_ (by rfl) - xrund [show BitVec.signExtend 64 (1024 : BitVec 32) = 1024 by decide] - -omit hb in -/-- The facts about the lengths of the layers. -/ -theorem lens_facts : ∀ len ∈ nttLens, 0 < len ∧ len ≤ 128 ∧ 2 * len * (128 / len) = 256 ∧ 0 < 128 / len := by - decide - -/-- A layer, from `rsi` = `f` and the zeta of its first block at `r8`. -/ -theorem lay_ok {tab : Nat → Nat} {Z : Nat → Zq} (hZ : TabOf tab Z) {fP zP : Addr} {len : Nat} - (hlen : len ∈ nttLens) (dz : BitVec 32) (zi : Nat → Nat) - (hzi : ∀ c < 128 / len, zi c < 256) - (hstep : ∀ c < 128 / len, coeffAddr zP (zi c) + BitVec.signExtend 64 dz = coeffAddr zP (zi (c + 1))) - (F : Poly) (s : State) (hsi : s.gpr .rsi = fP) (h8 : s.gpr .r8 = coeffAddr zP (zi 0)) - (hF : PolyIs s.mem fP F) (hw : pR fP ∈ s.wr) (hz : pR zP ∈ s.rd ++ s.wr) - (hd : (pR zP).Disjoint (pR fP)) (ht : Tab tab s.mem zP 256) : - WP isa (nttLay (code len) len dz) s fun s' => - (PolyIs s'.mem fP (layerN op F len (fun c => Z (zi c)) (128 / len)) ∧ Frame [pR fP] s.mem s'.mem ∧ - s'.gpr .rsi = fP ∧ s'.gpr .r8 = coeffAddr zP (zi (128 / len))) ∧ - Keep [.rdi, .r9, .r8, .rcx, .rax, .rdx, .rsi, .rcx, .r10, .r11, .rsi, .rdi, .rsi] s s' := by - obtain ⟨hl0, hl1, hl2, hl3⟩ := lens_facts len hlen - refine WP.seq (WP.mono (layPre_ok (128 / len) (by have := Nat.div_le_self 128 len; omega) s) - fun s1 ⟨⟨hdi, hm⟩, k1⟩ => ?_) - refine WP.seq (WP.mono (Q := fun (s' : State) => - PolyIs s'.mem fP (layerN op F len (fun c => Z (zi c)) (128 / len)) ∧ - Frame [pR fP] s.mem s'.mem ∧ s'.gpr .rsi = coeffAddr fP 256 ∧ s'.gpr .r8 = coeffAddr zP (zi (128 / len)) ∧ - Keep [.rdi, .r9, .r8, .rcx, .rax, .rdx, .rsi, .rcx, .r10, .r11, .rsi, .rdi] s s') ?_ - fun s2 ⟨hP, hf, hsi2, h82, k2⟩ => ?_) - · refine wp_countdown (cnt := .rdi) (N := 128 / len) (by have := Nat.div_le_self 128 len; omega) hl3 - (fun c (s' : State) => - PolyIs s'.mem fP (layerN op F len (fun c => Z (zi c)) c) ∧ Frame [pR fP] s.mem s'.mem ∧ - s'.gpr .rsi = coeffAddr fP (2 * len * c) ∧ s'.gpr .r8 = coeffAddr zP (zi c) ∧ - Keep [.rdi, .r9, .r8, .rcx, .rax, .rdx, .rsi, .rcx, .r10, .r11, .rsi, .rdi] s s') - (fun c hc s' ⟨hP, hf, hs', h8', hk⟩ _ => ?_) - (fun s' ⟨hP, hf, hs', h8', hk⟩ => ⟨hP, hf, by rw [hs', hl2], h8', hk⟩) - ⟨by rw [hm]; exact hF, by rw [hm]; exact Frame.refl _ _, by rw [k1.gpr (by decide), hsi]; simp, - by rw [k1.gpr (by decide), h8], k1.mono (by decide)⟩ hdi - have hcm : 2 * len * c + 2 * len ≤ 256 := by - have : 2 * len * (c + 1) ≤ 2 * len * (128 / len) := Nat.mul_le_mul_left _ (by omega) - rw [Nat.mul_succ] at this; omega - refine WP.mono (blk_ok hb hZ hl0 hl1 hcm (hzi c hc) dz _ s' hs' h8' hP (by rw [hk.2.2]; exact hw) - (by rw [hk.2.1, hk.2.2]; exact hz) (ht.frame hf (by simpa using hd) (by decide))) - fun s'' ⟨⟨hP', hf', hsi', h8'', hdi', hz'⟩, hk'⟩ => ⟨⟨by rw [layerN_succ]; exact hP', - hf.trans hf', ?_, ?_, (hk.trans hk').mono (by decide)⟩, hdi', hz'⟩ - · rw [hsi', Nat.mul_succ] - · rw [h8'', h8', hstep c hc] - · refine WP.mono (layPost_ok s2) fun s3 ⟨⟨hsi3, hm3⟩, k3⟩ => - ⟨⟨by rw [hm3]; exact hP, by rw [hm3]; exact hf, ?_, by rw [k3.gpr (by decide), h82]⟩, - (k2.trans k3).mono (by decide)⟩ - rw [hsi3, hsi2, coeffAddr] - show fP + BitVec.ofNat 64 1024 - 1024 = fP - exact BitVec.add_sub_cancel _ _ - -end - -end VG.Proof.MlDsa.X86_64.Arith diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Table.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Table.lean index 4501f89d9..fa3dc3c46 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Table.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Table.lean @@ -1,55 +1,22 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Basic -import VerifiedGarbage.Proof.Framework.Range /-! # ML-DSA on x86-64: tables of constants in the working space -Untrusted: everything here is checked by Lean. `storeTab t n` leaves the -`u32`s `t 0, …, t (n - 1)` at `r9` (`Tab`), and writes nothing else -(`storeTab_ok`). +Untrusted: everything here is checked by Lean. A table of `u32`s in the +working space (`Tab`), which writes elsewhere keep (`Tab.frame`). -/ namespace VG.Proof.MlDsa.X86_64.Arith open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Arith open VG.Proof.MlDsa.Arith -open VG.Proof.MlKem.X86_64 (Keep WP.keep ifp ifn) open VG.Spec.MlDsa (coeffAt) /-- The first `n` entries of the table `t` are the `u32`s at `p`. -/ def Tab (t : Nat → Nat) (m : Mem) (p : Addr) (n : Nat) : Prop := ∀ k < n, coeffAt m p k = BitVec.ofNat 32 (t k) -theorem tabStep_ok (t : Nat → Nat) (i : Nat) (s : State) - (hw : InRegions s.wr (s.gpr .r9 + BitVec.ofNat 64 (4 * i)) 4) : - WP isa (.block (tabStep t i)) s fun s' => - s'.mem = s.mem.writeW (s.gpr .r9 + BitVec.ofNat 64 (4 * i)) (BitVec.ofNat 32 (t i)) ∧ - Keep [.rax] s s' := by - refine WP.keep _ ?_ (by rfl) - unfold tabStep - xrund [hw] - -/-- The table, stored in the 1024 bytes at `r9`. -/ -theorem storeTab_ok (t : Nat → Nat) {n : Nat} (hn : n ≤ 256) (s : State) - (hw : pR (s.gpr .r9) ∈ s.wr) : - WP isa (.block (storeTab t n)) s fun s' => - Tab t s'.mem (s.gpr .r9) n ∧ Frame [pR (s.gpr .r9)] s.mem s'.mem ∧ Keep [.rax] s s' := by - refine WP.mono (wp_range_flatMap (M := isa) (fun k s' => Keep [.rax] s s' ∧ - Frame [pR (s.gpr .r9)] s.mem s'.mem ∧ Tab t s'.mem (s.gpr .r9) k) - (fun k s' hk ⟨hk', hf, ht⟩ => ?_) n (Nat.le_refl _) s - ⟨Keep.refl _ _, Frame.refl _ _, fun _ h => absurd h (Nat.not_lt_zero _)⟩) - fun s' ⟨hk, hf, ht⟩ => ⟨ht, hf, hk⟩ - have h9 : s'.gpr .r9 = s.gpr .r9 := hk'.gpr (by decide) - refine WP.mono (tabStep_ok t k s' (by - rw [hk'.2.2, h9]; exact ⟨_, hw, coeff_contains _ (show k < 256 by omega)⟩)) - fun s'' ⟨hm', hk''⟩ => ⟨(hk'.trans hk'').mono (by decide), ?_, fun j hj => ?_⟩ - · rw [hm', h9] - exact hf.writeW (List.mem_singleton_self _) _ (coeff_contains _ (show k < 256 by omega)) - · rw [hm', h9, ← coeffAddr, coeffAt_writeW _ _ (show j < 256 by omega) (show k < 256 by omega)] - by_cases e : k = j - · subst e; rw [ifp rfl] - · rw [ifn e]; exact ht j (by omega) - /-- Writes elsewhere keep the table. -/ theorem Tab.frame {t : Nat → Nat} {m m' : Mem} {p : Addr} {n : Nat} (h : Tab t m p n) {rs : List Region} (hf : Frame rs m m') (hd : ∀ r ∈ rs, (pR p).Disjoint r) (hn : n ≤ 256) : Tab t m' p n := diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VArith.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VArith.lean new file mode 100644 index 000000000..f8add0c32 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VArith.lean @@ -0,0 +1,285 @@ +import VerifiedGarbage.Proof.MlDsa.Arith.Mont +import VerifiedGarbage.Proof.Framework.X86_64.Avx +import VerifiedGarbage.Proof.MlKem.X86_64.VArith + +/-! +# ML-DSA on x86-64: arithmetic modulo `q` in doublewords + +Untrusted: everything here is checked by Lean. What `vmont`, `vcadd` and +`vcsub` (`Impl/MlDsa/X86_64/Arith/Vec.lean`) compute in each doubleword: + +* `montV d z zo`, the register `vmont` leaves: each doubleword is `mont` of + the product of those of `d` and `z` (`dword_montV`), if the even + doublewords of `zo` are the odd ones of `z` and each product is less than + `q · 2³²` (each quadword product `P` becomes `P + m · q`, which is + `mont P · 2³²`: `redc_toNat`); +* `caddL`, `csubL`: a doubleword plus `q` if it is negative, and less `q` + first, which `condSub` describes (`csubL_toNat`, `subD_toNat`, + `addD_toNat`). +-/ + +namespace VG.Proof.MlDsa.X86_64.Arith + +open VG VG.X86_64 +open VG.Proof.MlDsa.Arith +open VG.Spec.MlDsa (q) + +/-! ## Quadwords -/ + +theorem qword_app0 (a b : BitVec 64) : qword (a ++ b) 0 = b := by + apply BitVec.eq_of_getLsbD_eq; intro i hi + simp only [qword, BitVec.getLsbD_extractLsb', BitVec.getLsbD_append] + simp [hi] + +theorem qword_app1 (a b : BitVec 64) : qword (a ++ b) 1 = a := by + apply BitVec.eq_of_getLsbD_eq; intro i hi + simp only [qword, BitVec.getLsbD_extractLsb', BitVec.getLsbD_append] + simp [hi] + +theorem dword_lo (x : BitVec 128) (i : Nat) : dword x (2 * i) = (qword x i).extractLsb' 0 32 := by + apply BitVec.eq_of_getLsbD_eq; intro j hj + simp only [dword, qword, BitVec.getLsbD_extractLsb', hj, decide_true, Bool.true_and, + decide_eq_true (show 0 + j < 64 by omega)] + exact congrArg _ (by omega) + +theorem dword_hi (x : BitVec 128) (i : Nat) : dword x (2 * i + 1) = (qword x i).extractLsb' 32 32 := by + apply BitVec.eq_of_getLsbD_eq; intro j hj + simp only [dword, qword, BitVec.getLsbD_extractLsb', hj, decide_true, Bool.true_and, + decide_eq_true (show 32 + j < 64 by omega)] + exact congrArg _ (by omega) + +/-- The low doubleword of a quadword, zero-extended. -/ +def lo32 (x : BitVec 64) : BitVec 64 := (x.extractLsb' 0 32).setWidth 64 + +theorem lo32_toNat (x : BitVec 64) : (lo32 x).toNat = x.toNat % 2 ^ 32 := by + rw [lo32, BitVec.toNat_setWidth, BitVec.extractLsb'_toNat, Nat.shiftRight_zero, + Nat.mod_eq_of_lt (Nat.lt_of_lt_of_le (Nat.mod_lt _ (by decide)) (by decide))] + +theorem qword_paddq (x y : BitVec 128) {i : Nat} (hi : i < 2) : + qword (XBinOp.eval .paddq x y) i = qword x i + qword y i := by + rcases (by omega : i = 0 ∨ i = 1) with rfl | rfl + · simp only [XBinOp.eval, qword_app0] + · simp only [XBinOp.eval, qword_app1] + +theorem qword_pmuludq (x y : BitVec 128) {i : Nat} (hi : i < 2) : + qword (XBinOp.eval .pmuludq x y) i = lo32 (qword x i) * lo32 (qword y i) := by + rcases (by omega : i = 0 ∨ i = 1) with rfl | rfl + · simp only [XBinOp.eval, qword_app0, lo32, ← dword_lo] + · simp only [XBinOp.eval, qword_app1, lo32, ← dword_lo] + +theorem qword_psrlq32 (x : BitVec 128) {i : Nat} (hi : i < 2) : + qword (XShiftOp.eval .psrlq x 32) i = qword x i >>> 32 := by + rcases (by omega : i = 0 ∨ i = 1) with rfl | rfl + · simp only [XShiftOp.eval, show ¬ 63 < (32 : BitVec 8).toNat by decide, ite_false, qword_app0]; rfl + · simp only [XShiftOp.eval, show ¬ 63 < (32 : BitVec 8).toNat by decide, ite_false, qword_app1]; rfl + +theorem toNat_lo32_mul (x y : BitVec 64) : (lo32 x * lo32 y).toNat = x.toNat % 2 ^ 32 * (y.toNat % 2 ^ 32) := by + rw [BitVec.toNat_mul, lo32_toNat, lo32_toNat] + have h1 := Nat.mod_lt x.toNat (show 0 < 2 ^ 32 by decide) + have h2 := Nat.mod_lt y.toNat (show 0 < 2 ^ 32 by decide) + exact Nat.mod_eq_of_lt (Nat.lt_of_lt_of_le (Nat.mul_lt_mul'' h1 h2) (by decide)) + +/-! ## Montgomery reduction of a quadword -/ + +/-- `q` in each doubleword. -/ +def qV : BitVec 128 := 0x007FE001007FE001007FE001007FE001#128 + +/-- `-q⁻¹ mod 2³²` in each doubleword. -/ +def qinvV : BitVec 128 := 0xFC7FDFFFFC7FDFFFFC7FDFFFFC7FDFFF#128 + +theorem lo32_qword_qV {i : Nat} (hi : i < 2) : (lo32 (qword qV i)).toNat = q := by + rcases (by omega : i = 0 ∨ i = 1) with rfl | rfl <;> decide + +theorem lo32_qword_qinvV {i : Nat} (hi : i < 2) : (lo32 (qword qinvV i)).toNat = montQInv := by + rcases (by omega : i = 0 ∨ i = 1) with rfl | rfl <;> decide + +/-- `vredc` on a quadword `x`: `x + m · q`. -/ +def redc (x qi qq : BitVec 64) : BitVec 64 := x + lo32 (lo32 x * lo32 qi) * lo32 qq + +theorem redc_toNat {x qi qq : BitVec 64} (hqi : (lo32 qi).toNat = montQInv) (hqq : (lo32 qq).toNat = q) + (hx : x.toNat < q * 2 ^ 32) : (redc x qi qq).toNat = mont x.toNat * 2 ^ 32 := by + have hm : (lo32 (lo32 x * lo32 qi)).toNat = montM x.toNat := by + rw [lo32_toNat, BitVec.toNat_mul, lo32_toNat, hqi, montM, Nat.mod_mod_of_dvd _ (by decide)] + have hmq : (lo32 (lo32 x * lo32 qi) * lo32 qq).toNat = montM x.toNat * q := by + rw [BitVec.toNat_mul, hm, hqq] + exact Nat.mod_eq_of_lt (Nat.lt_of_lt_of_le (Nat.mul_lt_mul_of_pos_right (montM_lt _) (by decide)) + (by decide)) + rw [redc, BitVec.toNat_add, hmq, ← mont_mul] + have := mont_lt hx + rw [q_eq] at this + exact Nat.mod_eq_of_lt (by omega) + +theorem redc_hi {x qi qq : BitVec 64} (hqi : (lo32 qi).toNat = montQInv) (hqq : (lo32 qq).toNat = q) + (hx : x.toNat < q * 2 ^ 32) : ((redc x qi qq).extractLsb' 32 32).toNat = mont x.toNat := by + rw [BitVec.extractLsb'_toNat, redc_toNat hqi hqq hx, Nat.shiftRight_eq_div_pow, + Nat.mul_div_cancel _ (by decide)] + have := mont_lt hx + rw [q_eq] at this + exact Nat.mod_eq_of_lt (by omega) + +theorem redc_lo {x qi qq : BitVec 64} (hqi : (lo32 qi).toNat = montQInv) (hqq : (lo32 qq).toNat = q) + (hx : x.toNat < q * 2 ^ 32) : (redc x qi qq).extractLsb' 0 32 = 0 := by + apply BitVec.eq_of_toNat_eq + rw [BitVec.extractLsb'_toNat, redc_toNat hqi hqq hx, Nat.shiftRight_zero, Nat.mul_mod_left] + rfl + +theorem or_zero_toNat (x : BitVec 32) : (x ||| 0).toNat = x.toNat := by + simp + +theorem zero_or_toNat (x : BitVec 32) : ((0 : BitVec 32) ||| x).toNat = x.toNat := by + simp + +theorem lo_shr32 (x : BitVec 64) : ((x >>> 32).extractLsb' 0 32).toNat = (x.extractLsb' 32 32).toNat := by + rw [BitVec.extractLsb'_toNat, BitVec.extractLsb'_toNat, BitVec.toNat_ushiftRight, Nat.shiftRight_zero] + +theorem hi_shr32 (x : BitVec 64) : (x >>> 32).extractLsb' 32 32 = 0 := by + apply BitVec.eq_of_toNat_eq + rw [BitVec.extractLsb'_toNat, BitVec.toNat_ushiftRight, Nat.shiftRight_eq_div_pow, + Nat.shiftRight_eq_div_pow, Nat.div_div_eq_div_mul] + have := x.isLt + rw [Nat.div_eq_of_lt (by omega)]; rfl + +/-! ## `vmont` -/ + +/-- The register `vmont d z zo` leaves in `d`. -/ +def montV (d z zo : BitVec 128) : BitVec 128 := + let u := shufDwords d 0xF5 + let a := XBinOp.eval .pmuludq d z + let b := XBinOp.eval .pmuludq u zo + XBinOp.eval .por + (XShiftOp.eval .psrlq (XBinOp.eval .paddq a (XBinOp.eval .pmuludq (XBinOp.eval .pmuludq a qinvV) qV)) 32) + (XBinOp.eval .paddq b (XBinOp.eval .pmuludq (XBinOp.eval .pmuludq b qinvV) qV)) + +theorem qword_vredc (a : BitVec 128) {j : Nat} (hj : j < 2) : + qword (XBinOp.eval .paddq a (XBinOp.eval .pmuludq (XBinOp.eval .pmuludq a qinvV) qV)) j = + redc (qword a j) (qword qinvV j) (qword qV j) := by + rw [qword_paddq _ _ hj, qword_pmuludq _ _ hj, qword_pmuludq _ _ hj]; rfl + +theorem toNat_qword_pmuludq (x y : BitVec 128) {j : Nat} (hj : j < 2) : + (qword (XBinOp.eval .pmuludq x y) j).toNat = (dword x (2 * j)).toNat * (dword y (2 * j)).toNat := by + rw [qword_pmuludq _ _ hj, toNat_lo32_mul, dword_lo, dword_lo, BitVec.extractLsb'_toNat, + BitVec.extractLsb'_toNat, Nat.shiftRight_zero, Nat.shiftRight_zero] + +/-- Each doubleword of `montV d z zo` is `mont` of the product of those of +`d` and `z`. -/ +theorem dword_montV {d z zo : BitVec 128} (hzo : ∀ j < 2, dword zo (2 * j) = dword z (2 * j + 1)) + (hb : ∀ i < 4, (dword d i).toNat * (dword z i).toNat < q * 2 ^ 32) {i : Nat} (hi : i < 4) : + (dword (montV d z zo) i).toNat = mont ((dword d i).toNat * (dword z i).toNat) := by + have hqi : ∀ j < 2, (lo32 (qword qinvV j)).toNat = montQInv := fun j hj => lo32_qword_qinvV hj + have hqq : ∀ j < 2, (lo32 (qword qV j)).toNat = q := fun j hj => lo32_qword_qV hj + rw [montV, dword_por] + obtain ⟨j, hj, rfl | rfl⟩ : ∃ j < 2, i = 2 * j ∨ i = 2 * j + 1 := ⟨i / 2, by omega, by omega⟩ + all_goals + have hx : (qword (XBinOp.eval .pmuludq d z) j).toNat = (dword d (2 * j)).toNat * (dword z (2 * j)).toNat := + toNat_qword_pmuludq d z hj + have hy : (qword (XBinOp.eval .pmuludq (shufDwords d 0xF5) zo) j).toNat = + (dword d (2 * j + 1)).toNat * (dword z (2 * j + 1)).toNat := by + rw [toNat_qword_pmuludq _ _ hj, hzo j hj, dword_shufDwords _ _ (by omega)] + rcases (by omega : j = 0 ∨ j = 1) with rfl | rfl <;> rfl + · -- an even doubleword: the quotient of the even product, moved down + rw [dword_lo, dword_lo, qword_psrlq32 _ hj, qword_vredc _ hj, qword_vredc _ hj, + redc_lo (hqi j hj) (hqq j hj) (by rw [hy]; exact hb _ (by omega)), or_zero_toNat, lo_shr32, + redc_hi (hqi j hj) (hqq j hj) (by rw [hx]; exact hb _ (by omega)), hx] + · -- an odd doubleword: the quotient of the odd product, in place + rw [dword_hi, dword_hi, qword_psrlq32 _ hj, hi_shr32, zero_or_toNat, qword_vredc _ hj, + redc_hi (hqi j hj) (hqq j hj) (by rw [hy]; exact hb _ (by omega)), hy] + +/-! ## Conditional additions and subtractions of `q` -/ + +/-- `q` as a doubleword. -/ +def qB : BitVec 32 := 8380417#32 + +theorem dword_qV {i : Nat} (hi : i < 4) : dword qV i = qB := by + rcases cases4 hi with rfl | rfl | rfl | rfl <;> decide + +/-- `vcadd` on a doubleword: `d + q` if `d` is negative (as a signed doubleword). -/ +def caddL (d : BitVec 32) : BitVec 32 := d + (d.sshiftRight (min (31 : BitVec 8).toNat 32) &&& qB) + +/-- `vcsub` on a doubleword. -/ +def csubL (d : BitVec 32) : BitVec 32 := caddL (d - qB) + +theorem sshiftRight31 (d : BitVec 32) : + d.sshiftRight (min (31 : BitVec 8).toNat 32) = if d.toNat < 2 ^ 31 then 0 else -1 := by + rw [show min (31 : BitVec 8).toNat 32 = 31 from rfl] + apply BitVec.eq_of_toInt_eq + rw [BitVec.toInt_sshiftRight, MlKem.X86_64.W.toInt32] + have := d.isLt + split + · rw [show (0 : BitVec 32).toInt = 0 by decide, Int.shiftRight_eq_div_pow]; omega + · rw [show (-1 : BitVec 32).toInt = -1 by decide, Int.shiftRight_eq_div_pow]; omega + +theorem caddL_toNat (d : BitVec 32) : + (caddL d).toNat = if d.toNat < 2 ^ 31 then d.toNat else (d.toNat + q) % 2 ^ 32 := by + rw [caddL, sshiftRight31] + split + · rw [show (0 : BitVec 32) &&& qB = 0 by decide]; exact congrArg BitVec.toNat (BitVec.add_zero d) + · rw [show (-1 : BitVec 32) &&& qB = qB by decide, BitVec.toNat_add]; rfl + +/-- `vcsub` reduces a doubleword less than `2q`. -/ +theorem csubL_toNat {d : BitVec 32} (h : d.toNat < 2 * q) : (csubL d).toNat = condSub d.toNat := by + have e : (d - qB).toNat = (d.toNat + 2 ^ 32 - q) % 2 ^ 32 := by + rw [BitVec.toNat_sub]; rw [q_eq] at *; simp only [qB, BitVec.toNat_ofNat]; omega + rw [csubL, caddL_toNat, e, condSub]; rw [q_eq] at * + split <;> split <;> omega + +/-- The sum of two reduced doublewords, reduced by `vcsub`. -/ +theorem addD_toNat {a b : BitVec 32} (ha : a.toNat < q) (hb : b.toNat < q) : + (csubL (a + b)).toNat = condSub (a.toNat + b.toNat) := by + have e : (a + b).toNat = a.toNat + b.toNat := by + rw [BitVec.toNat_add]; rw [q_eq] at *; omega + rw [csubL_toNat (by rw [e]; omega), e] + +/-- The difference of two reduced doublewords, reduced by `vcadd`. -/ +theorem subD_toNat {a b : BitVec 32} (ha : a.toNat < q) (hb : b.toNat < q) : + (caddL (a - b)).toNat = condSub (a.toNat + q - b.toNat) := by + have e : (a - b).toNat = (a.toNat + 2 ^ 32 - b.toNat) % 2 ^ 32 := by + rw [BitVec.toNat_sub]; have := b.isLt; omega + rw [caddL_toNat, e, condSub]; rw [q_eq] at * + split <;> split <;> omega + +/-- `b - a + q` of two reduced doublewords, which `vibfly` multiplies by the zeta. -/ +theorem subq_toNat {a b : BitVec 32} (ha : a.toNat < q) (hb : b.toNat < q) : + (b - a + qB).toNat = b.toNat + q - a.toNat := by + rw [BitVec.toNat_add, BitVec.toNat_sub]; rw [q_eq] at *; simp only [qB, BitVec.toNat_ofNat]; omega + +/-! ## Butterflies, lane by lane -/ + +/-- A product by a zeta in Montgomery form, reduced: `ζ · y`. -/ +theorem mulZ {b z m : BitVec 32} {y ζ : Spec.MlDsa.Zq} (hb : b.toNat = y.val) (hz : z.toNat = ζ.val * 2 ^ 32 % q) + (hm : m.toNat = mont (b.toNat * z.toNat)) : (csubL m).toNat = (ζ * y).val := by + have hx : b.toNat * z.toNat < q * 2 ^ 32 := by + rw [hb, hz] + exact Nat.lt_of_lt_of_le (Nat.mul_lt_mul_of_lt_of_le (val_lt y) (Nat.le_of_lt (Nat.mod_lt _ (by decide))) + (by decide)) (by decide) + rw [csubL_toNat (by rw [hm]; exact mont_lt hx), hm, condSub_mont hx, hb, hz, mont_mulR, val_mul, + Nat.mul_comm] + +/-- The lanes of `vbfly`: `x + ζ · y` and `x - ζ · y`, from `t = ζ · y`. -/ +theorem bflyD {a t : BitVec 32} {x y ζ : Spec.MlDsa.Zq} (ha : a.toNat = x.val) (ht : t.toNat = (ζ * y).val) : + (csubL (a + t)).toNat = (x + ζ * y).val ∧ (caddL (a - t)).toNat = (x - ζ * y).val := by + have hx := val_lt x + have hzy := val_lt (ζ * y) + rw [addD_toNat (by rw [ha]; exact hx) (by rw [ht]; exact hzy), + subD_toNat (by rw [ha]; exact hx) (by rw [ht]; exact hzy), ha, ht, val_add, val_sub] + exact ⟨rfl, rfl⟩ + +/-- The lanes of `vibfly`: `x + y` and `ζ · (y - x)`. -/ +theorem ibflyD {a b z m : BitVec 32} {x y ζ : Spec.MlDsa.Zq} (ha : a.toNat = x.val) (hb : b.toNat = y.val) + (hz : z.toNat = ζ.val * 2 ^ 32 % q) (hm : m.toNat = mont ((b - a + qB).toNat * z.toNat)) : + (csubL (a + b)).toNat = (x + y).val ∧ (csubL m).toNat = (ζ * (y - x)).val := by + have hx := val_lt x + have hy := val_lt y + have e := subq_toNat (a := a) (b := b) (by rw [ha]; exact hx) (by rw [hb]; exact hy) + refine ⟨by rw [addD_toNat (by rw [ha]; exact hx) (by rw [hb]; exact hy), ha, hb, val_add], ?_⟩ + have hb' : (b - a + qB).toNat < q * 2 ^ 32 := by rw [e, ha, hb]; rw [q_eq] at *; omega + have hx' : (b - a + qB).toNat * z.toNat < q * 2 ^ 32 := by + rw [hz]; rw [e, ha, hb] at hb' ⊢ + have := Nat.mod_lt (ζ.val * 2 ^ 32) (show 0 < q by decide) + rw [q_eq] at * + exact Nat.lt_of_lt_of_le (Nat.mul_lt_mul_of_lt_of_le (show y.val + 8380417 - x.val < 2 * 8380417 by omega) + (Nat.le_of_lt this) (by decide)) (by decide) + rw [csubL_toNat (by rw [hm]; exact mont_lt hx'), hm, condSub_mont hx', hz, mont_mulR, e, ha, hb, + val_mul, val_sub', Nat.mul_mod_mod, Nat.mul_comm ζ.val, + Nat.add_sub_assoc (Nat.le_of_lt x.isLt) y.val] + +end VG.Proof.MlDsa.X86_64.Arith diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VLanes.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VLanes.lean new file mode 100644 index 000000000..a6e0c8ede --- /dev/null +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VLanes.lean @@ -0,0 +1,138 @@ +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.VArith +import VerifiedGarbage.Proof.MlKem.X86_64.VLanes +import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Ntt + +/-! +# ML-DSA on x86-64: coefficients in the doublewords of SSE registers + +Untrusted: everything here is checked by Lean. A register holds four +coefficients (`DLanes`), and the butterflies `vbfly` and `vibfly` compute +four butterflies of the specification at once (`vbfly_ok`, `vibfly_ok`), +from `q` and `-q⁻¹` in `xmm15` and `xmm14` (`VConsts`), with the zetas in +Montgomery form in `xmm13` (`ZLanes`) and its odd doublewords in the even +ones of `xmm12` (`ZOdd`). +-/ + +namespace VG.Proof.MlDsa.X86_64.Arith + +open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Arith +open VG.Proof.MlDsa.Arith +open VG.Proof.MlKem.X86_64 (XOnly XKeep xmm_setXmm mxcsr_setXmm ifp ifn) +open VG.Impl.MlKem.X86_64 (xb xmov) +open VG.Spec.MlDsa (q Zq) + +/-- The doublewords of `x` are the values of the coefficients `f 0, …, f 3`. -/ +def DLanes (x : BitVec 128) (f : Nat → Zq) : Prop := ∀ i < 4, (dword x i).toNat = (f i).val + +/-- The doublewords of `x` are the zetas `ζ i · 2³² mod q` (Montgomery form). -/ +def ZLanes (x : BitVec 128) (ζ : Nat → Zq) : Prop := ∀ i < 4, (dword x i).toNat = (ζ i).val * 2 ^ 32 % q + +/-- The even doublewords of `zo` are the odd ones of `z`. -/ +def ZOdd (z zo : BitVec 128) : Prop := ∀ j < 2, dword zo (2 * j) = dword z (2 * j + 1) + +/-- The constants of the vector code are in place. -/ +structure VConsts (s : State) : Prop where + q : s.xmm .xmm15 = qV + qinv : s.xmm .xmm14 = qinvV + +theorem VConsts.setXmm {s : State} (hc : VConsts s) {d : XReg} (h14 : XReg.xmm14 ≠ d) + (h15 : XReg.xmm15 ≠ d) (v : BitVec 128) : VConsts (s.setXmm d v) := + ⟨by rw [xmm_setXmm, ifn h15]; exact hc.q, by rw [xmm_setXmm, ifn h14]; exact hc.qinv⟩ + +theorem xonly_vconsts {rs : List XReg} {s s' : State} (h : XOnly rs s s') (hc : VConsts s) + (h14 : XReg.xmm14 ∉ rs) (h15 : XReg.xmm15 ∉ rs) : VConsts s' := + ⟨by rw [h.xmm _ h15, hc.q], by rw [h.xmm _ h14, hc.qinv]⟩ + +/-! ## Registers -/ + +/-- `vcadd` on a register. -/ +def caddV (d : BitVec 128) : BitVec 128 := + XBinOp.eval .paddd d (XBinOp.eval .pand (XShiftOp.eval .psrad d 31) qV) + +/-- `vcsub` on a register. -/ +def csubV (d : BitVec 128) : BitVec 128 := caddV (XBinOp.eval .psubd d qV) + +theorem dword_psubd (a b : BitVec 128) {i : Nat} (hi : i < 4) : + dword (XBinOp.eval .psubd a b) i = dword a i - dword b i := by + rcases cases4 hi with rfl | rfl | rfl | rfl <;> simp [XBinOp.eval] + +theorem dword_pand (a b : BitVec 128) (i : Nat) : + dword (XBinOp.eval .pand a b) i = dword a i &&& dword b i := by + apply BitVec.eq_of_getLsbD_eq; intro j hj + simp [XBinOp.eval, dword, hj] + +theorem dword_psrad (a : BitVec 128) (n : BitVec 8) {i : Nat} (hi : i < 4) : + dword (XShiftOp.eval .psrad a n) i = (dword a i).sshiftRight (min n.toNat 32) := by + rcases cases4 hi with rfl | rfl | rfl | rfl <;> simp [XShiftOp.eval] + +theorem dword_caddV (d : BitVec 128) {i : Nat} (hi : i < 4) : dword (caddV d) i = caddL (dword d i) := by + rw [caddV, dword_paddd _ _ hi, dword_pand, dword_psrad _ _ hi, dword_qV hi]; rfl + +theorem dword_csubV (d : BitVec 128) {i : Nat} (hi : i < 4) : dword (csubV d) i = csubL (dword d i) := by + rw [csubV, dword_caddV _ hi, dword_psubd _ _ hi, dword_qV hi]; rfl + +/-- A reduced coefficient times a zeta in Montgomery form is less than `q · 2³²`. -/ +theorem prod_lt {y z : BitVec 128} {f ζ : Nat → Zq} (hy : DLanes y f) (hz : ZLanes z ζ) : + ∀ i < 4, (dword y i).toNat * (dword z i).toNat < q * 2 ^ 32 := fun i hi => by + rw [hy i hi, hz i hi] + exact Nat.lt_of_lt_of_le (Nat.mul_lt_mul_of_lt_of_le (val_lt (f i)) (Nat.le_of_lt (Nat.mod_lt _ (by decide))) + (by decide)) (by decide) + +theorem subq_mul_lt {a b c : Nat} (ha : a < 8380417) (hb : b < 8380417) (hc : c < q) : + (b + q - a) * c < q * 2 ^ 32 := by + rw [q_eq] at * + exact Nat.lt_of_lt_of_le (Nat.mul_lt_mul_of_lt_of_le (show b + 8380417 - a < 2 * 8380417 by omega) + (Nat.le_of_lt hc) (by decide)) (by decide) + +theorem eval_movdqa (a b : BitVec 128) : XBinOp.eval .movdqa a b = b := rfl + +/-! ## Butterflies -/ + +theorem vbfly_ok {s : State} (hc : VConsts s) {x y ζ : Nat → Zq} (hx : DLanes (s.xmm .xmm0) x) + (hy : DLanes (s.xmm .xmm1) y) (hz : ZLanes (s.xmm .xmm13) ζ) (ho : ZOdd (s.xmm .xmm13) (s.xmm .xmm12)) : + WP isa (.block vbfly) s fun s' => DLanes (s'.xmm .xmm0) (fun i => x i + ζ i * y i) ∧ + DLanes (s'.xmm .xmm3) (fun i => x i - ζ i * y i) ∧ + XOnly [.xmm1, .xmm2, .xmm4, .xmm0, .xmm3] s s' := by + simp only [vbfly, vmont, vredc, vcsub, vcadd, xmov, xb, List.cons_append, List.nil_append] + vrun [eval_movdqa] + rw [hc.q, hc.qinv] + refine ⟨?_, ?_, by xonly⟩ + · change DLanes (csubV (XBinOp.eval .paddd (s.xmm .xmm0) + (csubV (montV (s.xmm .xmm1) (s.xmm .xmm13) (s.xmm .xmm12))))) _ + intro i hi + rw [dword_csubV _ hi, dword_paddd _ _ hi, dword_csubV _ hi] + exact (bflyD (hx i hi) (mulZ (hy i hi) (hz i hi) (dword_montV ho (prod_lt hy hz) hi))).1 + · change DLanes (caddV (XBinOp.eval .psubd (s.xmm .xmm0) + (csubV (montV (s.xmm .xmm1) (s.xmm .xmm13) (s.xmm .xmm12))))) _ + intro i hi + rw [dword_caddV _ hi, dword_psubd _ _ hi, dword_csubV _ hi] + exact (bflyD (hx i hi) (mulZ (hy i hi) (hz i hi) (dword_montV ho (prod_lt hy hz) hi))).2 + +theorem vibfly_ok {s : State} (hc : VConsts s) {x y ζ : Nat → Zq} (hx : DLanes (s.xmm .xmm0) x) + (hy : DLanes (s.xmm .xmm1) y) (hz : ZLanes (s.xmm .xmm13) ζ) (ho : ZOdd (s.xmm .xmm13) (s.xmm .xmm12)) : + WP isa (.block vibfly) s fun s' => DLanes (s'.xmm .xmm0) (fun i => x i + y i) ∧ + DLanes (s'.xmm .xmm3) (fun i => ζ i * (y i - x i)) ∧ + XOnly [.xmm1, .xmm2, .xmm4, .xmm0, .xmm3] s s' := by + simp only [vibfly, vmont, vredc, vcsub, vcadd, xmov, xb, List.cons_append, List.nil_append] + vrun [eval_movdqa] + rw [hc.q, hc.qinv] + have e : ∀ i < 4, dword (XBinOp.eval .paddd (XBinOp.eval .psubd (s.xmm .xmm1) (s.xmm .xmm0)) qV) i = + dword (s.xmm .xmm1) i - dword (s.xmm .xmm0) i + qB := fun i hi => by + rw [dword_paddd _ _ hi, dword_psubd _ _ hi, dword_qV hi] + have hb : ∀ i < 4, (dword (XBinOp.eval .paddd (XBinOp.eval .psubd (s.xmm .xmm1) (s.xmm .xmm0)) qV) i).toNat * + (dword (s.xmm .xmm13) i).toNat < q * 2 ^ 32 := fun i hi => by + rw [e i hi, subq_toNat (by rw [hx i hi]; exact val_lt _) (by rw [hy i hi]; exact val_lt _), hx i hi, hy i hi, + hz i hi] + exact subq_mul_lt (val_lt (x i)) (val_lt (y i)) (Nat.mod_lt _ (by decide)) + refine ⟨?_, ?_, by xonly⟩ + · change DLanes (csubV (XBinOp.eval .paddd (s.xmm .xmm0) (s.xmm .xmm1))) _ + intro i hi + rw [dword_csubV _ hi, dword_paddd _ _ hi, addD_toNat (by rw [hx i hi]; exact val_lt _) + (by rw [hy i hi]; exact val_lt _), hx i hi, hy i hi, val_add] + · change DLanes (csubV (montV (XBinOp.eval .paddd (XBinOp.eval .psubd (s.xmm .xmm1) (s.xmm .xmm0)) qV) + (s.xmm .xmm13) (s.xmm .xmm12))) _ + intro i hi + rw [dword_csubV _ hi] + exact (ibflyD (hx i hi) (hy i hi) (hz i hi) (by rw [dword_montV ho hb hi, e i hi])).2 + +end VG.Proof.MlDsa.X86_64.Arith diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VLay.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VLay.lean new file mode 100644 index 000000000..698814fa2 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VLay.lean @@ -0,0 +1,322 @@ +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.VMem +import VerifiedGarbage.Proof.MlKem.X86_64.VLay + +/-! +# ML-DSA on x86-64: the layers of the NTT and its inverse with `len ≥ 4` + +Untrusted: everything here is checked by Lean. For any butterfly code `bf` +that does what `op` does to the doublewords of two registers (`VBflyOk`), +and any block of the specification whose butterflies do `op` (`BlkOk`): +four butterflies of a block (`vstep`), the `len / 4` of them of a block +(`vblock_ok`), and the `128 / len` blocks of a layer (`vlay_ok`), on the +polynomial at `fP`, with the zetas from the table at `sP`. +-/ + +namespace VG.Proof.MlDsa.X86_64.Arith + +open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Arith +open VG.Proof.MlDsa.Arith +open VG.Proof.MlKem.X86_64 (XOnly xmm_setXmm ifp ifn sel sel_lt sel_zero add_ofNat_zero Keep wp_countdown + GOnly wp_rcxLoop sx1 sx16) +open VG.Spec.MlDsa (q n Poly Zq coeffAt polyAt Reduced PolyIs zetas) + +/-- Runs a block of general-purpose and SSE instructions. -/ +syntax "vrund" (" [" Lean.Parser.Tactic.simpLemma,* "]")? : tactic +macro_rules + | `(tactic| vrund) => `(tactic| vrund []) + | `(tactic| vrund [$ls,*]) => `(tactic| vrunm [ea_atD, $ls,*]) + +/-- `GOnly` of a chain of `setReg` and `setFlags`. -/ +macro "gonlyd" : tactic => `(tactic| exact ⟨⟨fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false, not_or] at hr + simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, hr, ite_false], rfl, rfl⟩, rfl, rfl, rfl⟩) + +theorem gonly_vconsts {rs : List Reg} {s s' : State} (h : GOnly rs s s') (hc : VConsts s) : VConsts s' := + ⟨by rw [h.xmm]; exact hc.q, by rw [h.xmm]; exact hc.qinv⟩ + +/-- The code `bf` of four butterflies does what `op` does to each pair of +doublewords of `xmm0` and `xmm1`, with the zetas in `xmm13` (and `xmm12`), +leaving the results in `xmm0` and `xmm3`. -/ +def VBflyOk (bf : List Instr) (op : Zq → Zq → Zq → Zq × Zq) : Prop := + ∀ s : State, VConsts s → ∀ x y ζ : Nat → Zq, DLanes (s.xmm .xmm0) x → DLanes (s.xmm .xmm1) y → + ZLanes (s.xmm .xmm13) ζ → ZOdd (s.xmm .xmm13) (s.xmm .xmm12) → + WP isa (.block bf) s fun s' => DLanes (s'.xmm .xmm0) (fun i => (op (x i) (y i) (ζ i)).1) ∧ + DLanes (s'.xmm .xmm3) (fun i => (op (x i) (y i) (ζ i)).2) ∧ + XOnly [.xmm1, .xmm2, .xmm4, .xmm0, .xmm3] s s' + +theorem vbfly_spec : VBflyOk vbfly (fun x y z => (x + z * y, x - z * y)) := + fun _ hc _ _ _ hx hy hz ho => vbfly_ok hc hx hy hz ho + +theorem vibfly_spec : VBflyOk vibfly (fun x y z => (x + y, z * (y - x))) := + fun _ hc _ _ _ hx hy hz ho => vibfly_ok hc hx hy hz ho + +/-- The first `t` butterflies of a block of the specification, with the +zeta of index `k` of the table, do `op` to `(j, j + len)`. -/ +structure BlkOk (blk : Poly → Nat → Nat → Nat → Nat → Poly) (op : Zq → Zq → Zq → Zq × Zq) : Prop where + zero : ∀ f len k st, blk f len k st 0 = f + add : ∀ f len k st t t', blk f len k st (t + t') = blk (blk f len k st t) len k (st + t) t' + get : ∀ f len k st t, 0 < len → t ≤ len → st + len + t ≤ n → ∀ i < n, + (blk f len k st t)[i]! = if st ≤ i ∧ i < st + t then (op f[i]! f[i + len]! (zetas k)).1 + else if st + len ≤ i ∧ i < st + len + t then (op f[i - len]! f[i]! (zetas k)).2 else f[i]! + +theorem blockN_add (op : Poly → Nat → Nat → Zq → Poly) (f : Poly) (len : Nat) (z : Zq) (st t t' : Nat) : + blockN op f len z st (t + t') = blockN op (blockN op f len z st t) len z (st + t) t' := by + simp only [blockN]; rw [← List.foldl_append, List.range'_append_1] + +/-- `blockN_bfly_get`, for the butterflies of the block up to `start + t` +only. -/ +theorem blockN_bfly_get' (w : Poly) {len : Nat} {z : Zq} {start t : Nat} (hlen : 0 < len) (ht : t ≤ len) + (hs : start + len + t ≤ n) {i : Nat} (hi : i < n) : + (blockN bfly w len z start t)[i]! = + if start ≤ i ∧ i < start + t then w[i]! + z * w[i + len]! + else if start + len ≤ i ∧ i < start + len + t then w[i - len]! - z * w[i]! + else w[i]! := by + induction t generalizing i with + | zero => + rw [blockN_zero, ite_eq_right (by omega), ite_eq_right (by omega)] + | succ t ih => + rw [blockN_succ, bfly_get _ hlen (by omega) _ hi, ih (i := start + t) (by omega) (by omega) (by omega), + ih (i := start + t + len) (by omega) (by omega) (by omega), ih (by omega) (by omega) hi] + rcases (by omega : i < start ∨ (start ≤ i ∧ i < start + t) ∨ i = start + t ∨ + (start + t < i ∧ i < start + len) ∨ (start + len ≤ i ∧ i < start + t + len) ∨ + i = start + t + len ∨ start + t + len < i) with h | h | rfl | h | h | rfl | h <;> + simp (disch := omega) only [ite_eq_left, ite_eq_right, Nat.add_sub_cancel, ↓reduceIte] + +/-- `blockN_bflyInv_get`, for the butterflies of the block up to +`start + t` only. -/ +theorem blockN_bflyInv_get' (w : Poly) {len : Nat} {z : Zq} {start t : Nat} (hlen : 0 < len) (ht : t ≤ len) + (hs : start + len + t ≤ n) {i : Nat} (hi : i < n) : + (blockN bflyInv w len z start t)[i]! = + if start ≤ i ∧ i < start + t then w[i]! + w[i + len]! + else if start + len ≤ i ∧ i < start + len + t then z * (w[i - len]! - w[i]!) + else w[i]! := by + induction t generalizing i with + | zero => + rw [blockN_zero, ite_eq_right (by omega), ite_eq_right (by omega)] + | succ t ih => + rw [blockN_succ, bflyInv_get _ hlen (by omega) _ hi, ih (i := start + t) (by omega) (by omega) (by omega), + ih (i := start + t + len) (by omega) (by omega) (by omega), ih (by omega) (by omega) hi] + rcases (by omega : i < start ∨ (start ≤ i ∧ i < start + t) ∨ i = start + t ∨ + (start + t < i ∧ i < start + len) ∨ (start + len ≤ i ∧ i < start + t + len) ∨ + i = start + t + len ∨ start + t + len < i) with h | h | rfl | h | h | rfl | h <;> + simp (disch := omega) only [ite_eq_left, ite_eq_right, Nat.add_sub_cancel, ↓reduceIte] + +theorem nttBlk_ok : BlkOk (fun f len k st t => blockN bfly f len (zetas k) st t) + (fun x y z => (x + z * y, x - z * y)) := + ⟨fun _ _ _ _ => rfl, fun _ _ _ _ _ _ => blockN_add _ _ _ _ _ _ _, + fun f _ _ _ _ hl ht hs _ hi => blockN_bfly_get' f hl ht hs hi⟩ + +/-- Algorithm 42 multiplies by `-ζ`; `vibfly` by `ζ`, the other way round. -/ +theorem neg_mul_sub (z x y : Zq) : -z * (x - y) = z * (y - x) := by + grind + +theorem nttInvBlk_ok : BlkOk (fun f len k st t => blockN bflyInv f len (-zetas k) st t) + (fun x y z => (x + y, z * (y - x))) := + ⟨fun _ _ _ _ => rfl, fun _ _ _ _ _ _ => blockN_add _ _ _ _ _ _ _, + fun f _ _ _ _ hl ht hs _ hi => by rw [blockN_bflyInv_get' f hl ht hs hi, neg_mul_sub]⟩ + +/-! ## A block -/ + +theorem f_in {rs : List Region} {fP : Addr} (hw : pR fP ∈ rs) {j : Nat} (hj : j + 4 ≤ 256) : + InRegions rs (coeffAddr fP j) 16 := + ⟨_, hw, pR_contains fP hj⟩ + +theorem tab_in {rs : List Region} {sP : Addr} (hw : pR sP ∈ rs) {k : Nat} (hk : k + 4 ≤ 256) : + InRegions rs (coeffAddr sP k) 16 := + ⟨_, hw, pR_contains sP hk⟩ + +/-- The facts a block keeps. -/ +structure BInv (fP : Addr) (s₀ s : State) : Prop where + keep : Keep [.r8, .rcx, .rdx, .rax] s₀ s + frame : Frame [pR fP] s₀.mem s.mem + consts : VConsts s + mxcsr : s.mxcsr = s₀.mxcsr + +theorem BInv.trans {fP : Addr} {s₁ s₂ s₃ : State} (h₁ : BInv fP s₁ s₂) (h₂ : BInv fP s₂ s₃) : BInv fP s₁ s₃ := + ⟨(h₁.keep.trans h₂.keep).mono (by simp), h₁.frame.trans h₂.frame, h₂.consts, h₂.mxcsr.trans h₁.mxcsr⟩ + +/-! ## Four butterflies -/ + +section +variable {bf : List Instr} {op : Zq → Zq → Zq → Zq × Zq} (hbf : VBflyOk bf op) + {blk : Poly → Nat → Nat → Nat → Nat → Poly} (hblk : BlkOk blk op) +include hbf hblk + +/-- The body of the loop over the vectors of a block. -/ +abbrev vbody (bf : List Instr) (len : Nat) : List Instr := + [.movdquLoad .xmm0 (at_ .rdx 0), .movdquLoad .xmm1 (at_ .rdx (4 * len))] ++ bf ++ + [.movdquStore (at_ .rdx 0) .xmm0, .movdquStore (at_ .rdx (4 * len)) .xmm3, .alu .add .rdx (.imm 16)] ++ + [.alu .sub .rcx (.imm 1)] + +theorem vstep {fP : Addr} {len st u k : Nat} (hl : 0 < len) (hs : st + 2 * len ≤ 256) (hu : 4 * u + 4 ≤ len) + {G : Poly} {s : State} (hc : VConsts s) (hz : ZLanes (s.xmm .xmm13) (fun _ => zetas k)) + (ho : ZOdd (s.xmm .xmm13) (s.xmm .xmm12)) + (hdx : s.gpr .rdx = coeffAddr fP (st + 4 * u)) (hS : PolyIs s.mem fP (blk G len k st (4 * u))) + (hw : pR fP ∈ s.wr) : + WP isa (.block (vbody bf len)) s fun s' => + PolyIs s'.mem fP (blk G len k st (4 * (u + 1))) ∧ s'.gpr .rdx = coeffAddr fP (st + 4 * (u + 1)) ∧ + Frame [pR fP] s.mem s'.mem ∧ VConsts s' ∧ s'.xmm .xmm13 = s.xmm .xmm13 ∧ + s'.xmm .xmm12 = s.xmm .xmm12 ∧ Keep [.rdx, .rcx] s s' ∧ + s'.gpr .rcx = s.gpr .rcx - 1 ∧ s'.zf = some (s.gpr .rcx - 1 == 0) ∧ s'.mxcsr = s.mxcsr := by + have j0 : st + 4 * u + 4 ≤ 256 := by omega + have j1 : st + 4 * u + len + 4 ≤ 256 := by omega + have a1 : coeffAddr fP (st + 4 * u) + BitVec.ofNat 64 (4 * len) = coeffAddr fP (st + 4 * u + len) := + coeffAddr_add _ _ _ + have r0 : InRegions (s.rd ++ s.wr) (coeffAddr fP (st + 4 * u)) 16 := f_in (List.mem_append_right _ hw) j0 + have r1 : InRegions (s.rd ++ s.wr) (coeffAddr fP (st + 4 * u + len)) 16 := f_in (List.mem_append_right _ hw) j1 + have w0 := f_in hw j0 + have w1 := f_in hw j1 + rw [vbody, List.append_assoc, List.append_assoc, WP.block_append_iff] + vrund [hdx, a1, r0, r1] + rw [WP.block_append_iff] + have hx := dlanes_load hS j0 + have hy := dlanes_load hS j1 + refine WP.mono (hbf _ ((hc.setXmm (by decide) (by decide) _).setXmm (by decide) (by decide) _) + (fun e => (blk G len k st (4 * u))[st + 4 * u + e]!) (fun e => (blk G len k st (4 * u))[st + 4 * u + len + e]!) + (fun _ => zetas k) (by rw [xmm_setXmm, xmm_setXmm]; exact hx) (by rw [xmm_setXmm]; exact hy) + (by rw [xmm_setXmm, xmm_setXmm]; exact hz) (by rw [xmm_setXmm, xmm_setXmm, xmm_setXmm, xmm_setXmm]; exact ho)) + fun s2 ⟨l0, l3, o2⟩ => ?_ + have c2 := xonly_vconsts o2 ((hc.setXmm (by decide) (by decide) _).setXmm (by decide) (by decide) _) (by decide) + (by decide) + have g2 : s2.gpr = s.gpr := o2.gpr + have m2 : s2.mem = s.mem := o2.mem + have e2 : s2.rd = s.rd ∧ s2.wr = s.wr := ⟨o2.rd, o2.wr⟩ + have x2 : s2.mxcsr = s.mxcsr := o2.mxcsr + have z2 : s2.xmm .xmm13 = s.xmm .xmm13 := by rw [o2.xmm _ (by decide), xmm_setXmm, xmm_setXmm]; rfl + have z2' : s2.xmm .xmm12 = s.xmm .xmm12 := by rw [o2.xmm _ (by decide), xmm_setXmm, xmm_setXmm]; rfl + vrund [g2, m2, e2.1, e2.2, hdx, a1, w0, w1, x2] + refine ⟨?_, ?_, ?_, ⟨c2.q, c2.qinv⟩, z2, z2', ⟨fun r hr => ?_, rfl, rfl⟩⟩ + · refine polyIs_write2 hS j0 j1 (by omega) l0 l3 fun i hi => ?_ + rw [show 4 * (u + 1) = 4 * u + 4 by omega, hblk.add, hblk.get _ _ _ _ _ hl (by omega) + (by rw [n_eq]; omega) _ (by rw [n_eq]; exact hi)] + by_cases c1 : st + 4 * u ≤ i ∧ i < st + 4 * u + 4 + · rw [ite_eq_left_of_eq_true _ _ (eq_true c1), ite_eq_left_of_eq_true _ _ (eq_true c1), + show st + 4 * u + (i - (st + 4 * u)) = i by omega, + show st + 4 * u + len + (i - (st + 4 * u)) = i + len by omega] + · rw [ite_eq_right_of_eq_false _ _ (eq_false c1), ite_eq_right_of_eq_false _ _ (eq_false c1)] + by_cases c2 : st + 4 * u + len ≤ i ∧ i < st + 4 * u + len + 4 + · rw [ite_eq_left_of_eq_true _ _ (eq_true c2), ite_eq_left_of_eq_true _ _ (eq_true (by omega)), + show st + 4 * u + (i - (st + 4 * u + len)) = i - len by omega, + show st + 4 * u + len + (i - (st + 4 * u + len)) = i by omega] + · rw [ite_eq_right_of_eq_false _ _ (eq_false c2), ite_eq_right_of_eq_false _ _ (eq_false (by omega))] + · rw [show (16 : BitVec 64) = BitVec.ofNat 64 (4 * 4) from rfl, coeffAddr_add, + show st + 4 * u + 4 = st + 4 * (u + 1) by omega] + · exact frame_write2 (Frame.refl _ _) j0 j1 _ _ + · simp only [List.mem_cons, List.not_mem_nil, or_false, not_or] at hr + simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, hr.1, hr.2, ite_false] + +/-- The code of a block of a layer with `len ≥ 4`. -/ +abbrev vblk (bf : List Instr) (len : Nat) (dz : BitVec 32) : Prog isa := + .seq (.block (vzeta 0 ++ [.alu .add .r8 (.imm dz)])) + (.seq (VG.Impl.MlKem.X86_64.rcxLoop (len / 4) ([.movdquLoad .xmm0 (at_ .rdx 0), + .movdquLoad .xmm1 (at_ .rdx (4 * len))] ++ + bf ++ [.movdquStore (at_ .rdx 0) .xmm0, .movdquStore (at_ .rdx (4 * len)) .xmm3, + .alu .add .rdx (.imm 16)])) + (.block [.alu .add .rdx (.imm (BitVec.ofNat 32 (4 * len))), .alu .sub .rax (.imm 1)])) + +theorem vblock_ok {fP sP : Addr} {len st kz : Nat} (h4 : 4 ≤ len) (hl4 : len % 4 = 0) (hl : len ≤ 128) + (hs : st + 2 * len ≤ 256) (hkz : kz + 4 ≤ 256) (dz : BitVec 32) {G : Poly} {s : State} (hc : VConsts s) + (hdx : s.gpr .rdx = coeffAddr fP st) (h8r : s.gpr .r8 = coeffAddr sP kz) (hS : PolyIs s.mem fP G) + (hT : Tab zmTab s.mem sP 256) (hwf : pR fP ∈ s.wr) (hw : pR sP ∈ s.wr) : + WP isa (vblk bf len dz) s fun s' => PolyIs s'.mem fP (blk G len kz st len) ∧ + s'.gpr .rdx = coeffAddr fP (st + 2 * len) ∧ s'.gpr .r8 = s.gpr .r8 + BitVec.signExtend 64 dz ∧ + s'.gpr .rax = s.gpr .rax - 1 ∧ s'.zf = some (s.gpr .rax - 1 == 0) ∧ BInv fP s s' := by + -- the zeta + refine WP.seq ?_ + rw [WP.block_append_iff] + refine WP.mono (vzeta_ok 0 (k := kz) (fun j _ => by rw [sel_zero]; omega) h8r + (tab_in (List.mem_append_right _ hw) hkz) hT) fun s1 ⟨z1, zo1, o1⟩ => ?_ + have g1 : s1.gpr = s.gpr := o1.gpr + refine WP.mono (Q := fun (s2 : State) => s2.gpr .r8 = s.gpr .r8 + BitVec.signExtend 64 dz ∧ + GOnly [.r8] s1 s2) + (by vrund [g1]; gonlyd) + fun s2 ⟨h82, o2⟩ => ?_ + have c2 := gonly_vconsts o2 (xonly_vconsts o1 hc (by decide) (by decide)) + have z2 : ZLanes (s2.xmm .xmm13) (fun _ => zetas kz) := by + rw [o2.xmm]; intro i hi; rw [z1 i hi]; dsimp only; rw [sel_zero, Nat.add_zero] + have zo2 : ZOdd (s2.xmm .xmm13) (s2.xmm .xmm12) := by rw [o2.xmm]; exact zo1 + have dx2 : s2.gpr .rdx = coeffAddr fP st := by rw [o2.keep.gpr (by decide), g1, hdx] + have hw2 : pR fP ∈ s2.wr := by rw [o2.keep.2.2, o1.wr]; exact hwf + have m2 : s2.mem = s.mem := by rw [o2.mem, o1.mem] + refine WP.seq (WP.mono (wp_rcxLoop (N := len / 4) (by omega) (by omega) + (fun u w => PolyIs w.mem fP (blk G len kz st (4 * u)) ∧ w.gpr .rdx = coeffAddr fP (st + 4 * u) ∧ + VConsts w ∧ w.xmm .xmm13 = s2.xmm .xmm13 ∧ w.xmm .xmm12 = s2.xmm .xmm12 ∧ Keep [.rcx, .rdx] s2 w ∧ + Frame [pR fP] s2.mem w.mem ∧ w.mxcsr = s2.mxcsr) + (fun w o hc => ⟨by rw [hblk.zero, o.mem, m2]; exact hS, by rw [o.keep.gpr (by decide), dx2]; rfl, + gonly_vconsts o c2, by rw [o.xmm], by rw [o.xmm], o.keep.mono (by simp), by rw [o.mem]; exact Frame.refl _ _, + o.mxcsr⟩) + (fun u hu w ⟨hS', hdx', hc', hz', hzo', hk', hf', hx'⟩ => WP.mono (vstep hbf hblk (by omega) hs (by + have := Nat.div_mul_cancel (Nat.dvd_of_mod_eq_zero hl4); omega) hc' (by rw [hz']; exact z2) + (by rw [hz', hzo']; exact zo2) hdx' hS' (by rw [hk'.2.2]; exact hw2)) + fun w' ⟨hS'', hdx'', hf'', hc'', hz'', hzo'', hk'', hcx, hzf, hx''⟩ => + ⟨⟨hS'', hdx'', hc'', by rw [hz'', hz'], by rw [hzo'', hzo'], (hk'.trans hk'').mono (by simp), + hf'.trans hf'', by rw [hx'', hx']⟩, hcx, hzf⟩)) fun w ⟨hS3, hdx3, hc3, _, _, hk3, hf3, hx3⟩ => ?_) + rw [show 4 * (len / 4) = len from Nat.mul_div_cancel' (Nat.dvd_of_mod_eq_zero hl4)] at hS3 hdx3 + have hax : w.gpr .rax = s.gpr .rax := by rw [hk3.gpr (by decide), o2.keep.gpr (by decide), g1] + have h8w : w.gpr .r8 = s.gpr .r8 + BitVec.signExtend 64 dz := by rw [hk3.gpr (by decide), h82] + vrund [hdx3, sx_ofNat (show 4 * len < 2 ^ 31 by omega), hax, h8w] + refine ⟨hS3, by rw [coeffAddr_add, show st + len + len = st + 2 * len by omega], ?_⟩ + have k1 : Keep [.r8, .rcx, .rdx, .rax] s w := + (Keep.trans (⟨fun r _ => by rw [g1], o1.rd, o1.wr⟩ : Keep [] s s1) (o2.keep.trans hk3)).mono (by simp) + refine ⟨⟨fun r hr => ?_, k1.2.1, k1.2.2⟩, by rw [← m2]; exact hf3, + ⟨by simp only [RegUpd.xmm_setReg, RegUpd.xmm_setFlags]; exact hc3.q, + by simp only [RegUpd.xmm_setReg, RegUpd.xmm_setFlags]; exact hc3.qinv⟩, + by simp only [RegUpd.mxcsr_setReg, RegUpd.mxcsr_setFlags]; rw [hx3, o2.mxcsr, o1.mxcsr]⟩ + simp only [List.mem_cons, List.not_mem_nil, or_false, not_or] at hr + simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, hr, ite_false] + exact k1.gpr (by simp [hr]) + +/-! ## A layer -/ + +/-- The first `b` blocks of the layer with `len`, block `c` with the zeta of +index `zi c`. -/ +def layF (blk : Poly → Nat → Nat → Nat → Nat → Poly) (F : Poly) (len : Nat) (zi : Nat → Nat) (b : Nat) : + Poly := + (List.range b).foldl (fun f c => blk f len (zi c) (2 * len * c) len) F + +theorem vlay_ok {fP sP : Addr} {len k : Nat} (hlen : len ∈ [4, 8, 16, 32, 64, 128]) (dz : BitVec 32) + (zi : Nat → Nat) (hz0 : zi 0 = k) (hzi : ∀ c < 128 / len, zi c + 4 ≤ 256) + (hstep : ∀ c < 128 / len, coeffAddr sP (zi c) + BitVec.signExtend 64 dz = coeffAddr sP (zi (c + 1))) + {F : Poly} {s : State} (hc : VConsts s) (hdi : s.gpr .rdi = fP) (hsi : s.gpr .rsi = sP) + (hS : PolyIs s.mem fP F) (hT : Tab zmTab s.mem sP 256) (hwf : pR fP ∈ s.wr) (hw : pR sP ∈ s.wr) + (hd : (pR sP).Disjoint (pR fP)) : + WP isa (vlay bf len k dz) s fun s' => PolyIs s'.mem fP (layF blk F len zi (128 / len)) ∧ + BInv fP s s' := by + have hl : 4 ≤ len ∧ len % 4 = 0 ∧ len ≤ 128 ∧ 2 * len * (128 / len) = 256 ∧ 0 < 128 / len ∧ + 128 / len ≤ 32 := by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hlen + rcases hlen with rfl | rfl | rfl | rfl | rfl | rfl <;> decide + obtain ⟨h4, hl4, hl128, hcov, hpos, h32⟩ := hl + have hk : k + 4 ≤ 256 := hz0 ▸ hzi 0 hpos + refine WP.seq (WP.mono (Q := fun (w : State) => w.gpr .rdx = fP ∧ w.gpr .r8 = coeffAddr sP k ∧ + w.gpr .rax = BitVec.ofNat 64 (128 / len) ∧ GOnly [.rdx, .r8, .rax] s w) + (by + simp only [leaR] + vrund [sx_ofNat (show 4 * k < 2 ^ 31 by omega), hsi, hdi] + refine ⟨?_, by gonlyd⟩ + apply BitVec.eq_of_toNat_eq + rw [BitVec.toNat_setWidth, BitVec.toNat_ofNat, BitVec.toNat_ofNat] + omega) fun w ⟨hdx, h8r, hax, o⟩ => ?_) + have hwf' : pR fP ∈ w.wr := by rw [o.keep.2.2]; exact hwf + have hw' : pR sP ∈ w.wr := by rw [o.keep.2.2]; exact hw + refine WP.mono (wp_countdown (cnt := .rax) (N := 128 / len) (by omega) hpos + (fun c u => PolyIs u.mem fP (layF blk F len zi c) ∧ u.gpr .rdx = coeffAddr fP (2 * len * c) ∧ + u.gpr .r8 = coeffAddr sP (zi c) ∧ BInv fP w u ∧ Tab zmTab u.mem sP 256) + (fun c hc u ⟨hS', hdx', h8', hb', hT'⟩ _ => ?_) (fun u h => h) + ⟨by rw [o.mem]; exact hS, by rw [hdx, Nat.mul_zero, coeffAddr, Nat.mul_zero, add_ofNat_zero], + by rw [h8r, hz0], ⟨Keep.refl _ _, Frame.refl _ _, gonly_vconsts o hc, rfl⟩, by rw [o.mem]; exact hT⟩ hax) + fun u ⟨hS', _, _, hb', _⟩ => ⟨hS', ⟨(o.keep.trans hb'.keep).mono (by simp), + by rw [← o.mem]; exact hb'.frame, hb'.consts, by rw [hb'.mxcsr, o.mxcsr]⟩⟩ + have hs : 2 * len * c + 2 * len ≤ 256 := by + have : 2 * len * (c + 1) ≤ 2 * len * (128 / len) := Nat.mul_le_mul_left _ (by omega) + rw [Nat.mul_succ] at this; omega + refine WP.mono (vblock_ok hbf hblk h4 hl4 hl128 hs (hzi c hc) dz hb'.consts hdx' h8' hS' hT' + (by rw [hb'.keep.2.2]; exact hwf') (by rw [hb'.keep.2.2]; exact hw')) + fun u' ⟨hS'', hdx'', h8'', hax'', hzf'', hb''⟩ => + ⟨⟨by rw [layF, foldl_range_succ]; exact hS'', by rw [hdx'', Nat.mul_succ], + by rw [h8'', h8', hstep c hc], hb'.trans hb'', + hT'.frame hb''.frame (fun r hr => by rw [List.mem_singleton.mp hr]; exact hd) (by decide)⟩, hax'', hzf''⟩ + +end + +end VG.Proof.MlDsa.X86_64.Arith diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VLay21.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VLay21.lean new file mode 100644 index 000000000..288058286 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VLay21.lean @@ -0,0 +1,458 @@ +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.VLay + +/-! +# ML-DSA on x86-64: the layers of the NTT and its inverse with `len` = 2 and 1 + +Untrusted: everything here is checked by Lean. The layer with `len = 2` runs +two blocks at a time (`vstep2`): the lower halves of their coefficients +gathered into `xmm0` and the upper ones into `xmm1` by `punpcklqdq` and +`punpckhqdq`, and back. The layer with `len = 1` runs four blocks at a time +(`vstep1`): their coefficients gathered by `pshufd` and `punpck{l,h}qdq`, +and interleaved back by `punpck{l,h}dq`. +-/ + +namespace VG.Proof.MlDsa.X86_64.Arith + +open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Arith +open VG.Proof.MlDsa.Arith +open VG.Proof.MlKem.X86_64 (XOnly xmm_setXmm mxcsr_setXmm ifp ifn sel sel_lt add_ofNat_zero Keep GOnly + wp_rcxLoop sx32) +open VG.Impl.MlKem.X86_64 (xb xmov) +open VG.Spec.MlDsa (q n Poly Zq coeffAt polyAt Reduced PolyIs zetas) + +theorem dword_punpcklqdq (a b : BitVec 128) {j : Nat} (hj : j < 4) : + dword (XBinOp.eval .punpcklqdq a b) j = if j < 2 then dword a j else dword b (j - 2) := by + rw [punpcklqdq_eq] + rcases cases4 hj with rfl | rfl | rfl | rfl <;> + simp (disch := decide) only [Nat.reduceLT, Nat.reduceSub, ite_true, ite_false, dword_ofDwords_0, + dword_ofDwords_1, dword_ofDwords_2, dword_ofDwords_3] + +theorem dword_punpckhqdq (a b : BitVec 128) {j : Nat} (hj : j < 4) : + dword (XBinOp.eval .punpckhqdq a b) j = if j < 2 then dword a (2 + j) else dword b j := by + rw [punpckhqdq_eq] + rcases cases4 hj with rfl | rfl | rfl | rfl <;> + simp (disch := decide) only [Nat.reduceLT, Nat.reduceAdd, ite_true, ite_false, dword_ofDwords_0, + dword_ofDwords_1, dword_ofDwords_2, dword_ofDwords_3] + +theorem dword_punpckldq' (a b : BitVec 128) {j : Nat} (hj : j < 4) : + dword (XBinOp.eval .punpckldq a b) j = if j % 2 = 0 then dword a (j / 2) else dword b (j / 2) := by + rw [dword_punpckldq] + rcases cases4 hj with rfl | rfl | rfl | rfl <;> simp + +theorem dword_punpckhdq' (a b : BitVec 128) {j : Nat} (hj : j < 4) : + dword (XBinOp.eval .punpckhdq a b) j = if j % 2 = 0 then dword a (2 + j / 2) else dword b (2 + j / 2) := by + rw [dword_punpckhdq] + rcases cases4 hj with rfl | rfl | rfl | rfl <;> simp + +/-- The doublewords of `x` that `pshufd` with `0xD8` puts in place `e`: the +even ones in the lower half, the odd ones in the upper half. -/ +theorem dword_d8 (x : BitVec 128) {e : Nat} (he : e < 4) : + dword (shufDwords x 0xD8) e = dword x (if e < 2 then 2 * e else 2 * (e - 2) + 1) := by + rw [dword_shufDwords _ _ he] + rcases cases4 he with rfl | rfl | rfl | rfl <;> rfl + +/-- The general-purpose registers but `rs`, memory, the permissions and +MXCSR are as they were. -/ +structure GKeep (rs : List Reg) (s s' : State) : Prop where + keep : Keep rs s s' + mem : s'.mem = s.mem + mxcsr : s'.mxcsr = s.mxcsr + +/-! ## The layer with `len = 2` -/ + +section +variable {bf : List Instr} {op : Zq → Zq → Zq → Zq × Zq} (hbf : VBflyOk bf op) + {blk : Poly → Nat → Nat → Nat → Nat → Poly} (hblk : BlkOk blk op) +include hbf hblk + +/-- The loads, the zetas and the gathering of the lower and upper halves. -/ +abbrev pre2 (o : BitVec 8) (dz : BitVec 32) : List Instr := + [.movdquLoad .xmm0 (at_ .rdx 0), .movdquLoad .xmm1 (at_ .rdx 16)] ++ vzeta o ++ + [.alu .add .r8 (.imm dz), xmov .xmm2 .xmm0, xb .punpcklqdq .xmm0 .xmm1, xb .punpckhqdq .xmm2 .xmm1, + xmov .xmm1 .xmm2] + +/-- The interleaving back, the stores and the counts. -/ +abbrev post2 : List Instr := + [xmov .xmm1 .xmm0, xb .punpcklqdq .xmm0 .xmm3, xb .punpckhqdq .xmm1 .xmm3, + .movdquStore (at_ .rdx 0) .xmm0, .movdquStore (at_ .rdx 16) .xmm1, .alu .add .rdx (.imm 32), + .alu .sub .rcx (.imm 1)] + +theorem vstep2 {fP sP : Addr} {i kz : Nat} (hi : i < 32) (o : BitVec 8) (dz : BitVec 32) (zi : Nat → Nat) + (hk : kz + 4 ≤ 256) (hsel : ∀ e < 4, kz + sel o e = zi (2 * i + e / 2)) + {F : Poly} {s : State} (hc : VConsts s) (hdx : s.gpr .rdx = coeffAddr fP (8 * i)) + (h8 : s.gpr .r8 = coeffAddr sP kz) (hS : PolyIs s.mem fP (layF blk F 2 zi (2 * i))) + (hT : Tab zmTab s.mem sP 256) (hwf : pR fP ∈ s.wr) (hw : pR sP ∈ s.wr) : + WP isa (.block (pre2 o dz ++ (bf ++ post2))) s fun s' => + PolyIs s'.mem fP (layF blk F 2 zi (2 * (i + 1))) ∧ s'.gpr .rdx = coeffAddr fP (8 * (i + 1)) ∧ + s'.gpr .r8 = s.gpr .r8 + BitVec.signExtend 64 dz ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ + s'.zf = some (s.gpr .rcx - 1 == 0) ∧ BInv fP s s' := by + have j0 : 8 * i + 4 ≤ 256 := by omega + have j1 : 8 * i + 4 + 4 ≤ 256 := by omega + have a1 : coeffAddr fP (8 * i) + BitVec.ofNat 64 16 = coeffAddr fP (8 * i + 4) := coeffAddr_add _ _ 4 + have r0 := f_in (List.mem_append_right s.rd hwf) j0 + have r1 := f_in (List.mem_append_right s.rd hwf) j1 + have hk' : ∀ j < 4, kz + sel o j < 256 := fun j _ => by have := sel_lt o j; omega + generalize hG : layF blk F 2 zi (2 * i) = G at hS + have lx := dlanes_load hS j0 + have ly := dlanes_load hS j1 + rw [WP.block_append_iff, show pre2 o dz = [.movdquLoad .xmm0 (at_ .rdx 0), .movdquLoad .xmm1 (at_ .rdx 16)] ++ + (vzeta o ++ [.alu .add .r8 (.imm dz), xmov .xmm2 .xmm0, xb .punpcklqdq .xmm0 .xmm1, + xb .punpckhqdq .xmm2 .xmm1, xmov .xmm1 .xmm2]) by simp, WP.block_append_iff] + vrund [hdx, a1, r0, r1] + rw [WP.block_append_iff] + refine WP.mono (vzeta_ok o hk' (by simp only [RegUpd.gpr_setXmm]; exact h8) + (by simp only [RegUpd.rd_setXmm, RegUpd.wr_setXmm]; exact tab_in (List.mem_append_right _ hw) hk) + (by simp only [RegUpd.mem_setXmm]; exact hT)) fun s1 ⟨z1, zo1, o1⟩ => ?_ + refine WP.mono (Q := fun (s1' : State) => DLanes (s1'.xmm .xmm0) (fun e => G[8 * i + e + 2 * (e / 2)]!) ∧ + DLanes (s1'.xmm .xmm1) (fun e => G[8 * i + 2 + e + 2 * (e / 2)]!) ∧ + ZLanes (s1'.xmm .xmm13) (fun e => zetas (zi (2 * i + e / 2))) ∧ ZOdd (s1'.xmm .xmm13) (s1'.xmm .xmm12) ∧ + VConsts s1' ∧ s1'.gpr .r8 = s.gpr .r8 + BitVec.signExtend 64 dz ∧ GKeep [.r8] s s1') ?_ + fun s1' ⟨l0, l1, l13, l12, c1, h81, o1'⟩ => ?_ + · have g1 : s1.gpr = s.gpr := by rw [o1.gpr]; rfl + have m1 : s1.mem = s.mem := by rw [o1.mem]; rfl + have e1 : s1.rd = s.rd ∧ s1.wr = s.wr := ⟨by rw [o1.rd]; rfl, by rw [o1.wr]; rfl⟩ + have x1 : s1.mxcsr = s.mxcsr := by rw [o1.mxcsr]; rfl + have x0 : s1.xmm .xmm0 = s.mem.readW (coeffAddr fP (8 * i)) 128 := by + rw [o1.xmm _ (by decide)]; simp only [xmm_setXmm]; rfl + have x1' : s1.xmm .xmm1 = s.mem.readW (coeffAddr fP (8 * i + 4)) 128 := by + rw [o1.xmm _ (by decide)]; simp only [xmm_setXmm]; rfl + have c1 : VConsts s1 := xonly_vconsts o1 ((hc.setXmm (by decide) (by decide) _).setXmm (by decide) + (by decide) _) (by decide) (by decide) + vrund [g1, m1, e1.1, e1.2, x1, eval_movdqa] + refine ⟨?_, ?_, ?_, ?_, ⟨?_, ?_⟩, ⟨⟨fun r hr => ?_, ?_, ?_⟩, ?_, ?_⟩⟩ + · intro e he + rw [dword_punpcklqdq _ _ he, x0, x1'] + split + · rw [lx e he]; dsimp only; rw [show 8 * i + e + 2 * (e / 2) = 8 * i + e by omega] + · rw [ly (e - 2) (by omega)]; dsimp only + rw [show 8 * i + e + 2 * (e / 2) = 8 * i + 4 + (e - 2) by omega] + · intro e he + rw [dword_punpckhqdq _ _ he, x0, x1'] + split + · rw [lx (2 + e) (by omega)]; dsimp only + rw [show 8 * i + 2 + e + 2 * (e / 2) = 8 * i + (2 + e) by omega] + · rw [ly e he]; dsimp only; rw [show 8 * i + 2 + e + 2 * (e / 2) = 8 * i + 4 + e by omega] + · intro e he + rw [z1 e he]; dsimp only; rw [hsel e he] + · exact zo1 + · exact c1.q + · exact c1.qinv + · simp only [RegUpd.gpr_setXmm, RegUpd.gpr_setReg, RegUpd.gpr_setFlags, g1] + rw [ifn (by simpa using hr)] + all_goals simp only [RegUpd.rd_setXmm, RegUpd.wr_setXmm, RegUpd.mem_setXmm, mxcsr_setXmm, + RegUpd.rd_setReg, RegUpd.wr_setReg, RegUpd.mem_setReg, RegUpd.mxcsr_setReg, RegUpd.rd_setFlags, + RegUpd.wr_setFlags, RegUpd.mem_setFlags, RegUpd.mxcsr_setFlags, e1.1, e1.2, m1, x1] + rw [WP.block_append_iff] + refine WP.mono (hbf _ c1 _ _ _ l0 l1 l13 l12) fun s2 ⟨a0, a3, o2⟩ => ?_ + have g2 : s2.gpr .rdx = s.gpr .rdx := by rw [o2.gpr, o1'.keep.gpr (by decide)] + have e2 : s2.rd = s.rd ∧ s2.wr = s.wr := ⟨by rw [o2.rd, o1'.keep.2.1], by rw [o2.wr, o1'.keep.2.2]⟩ + have m2 : s2.mem = s.mem := by rw [o2.mem, o1'.mem] + have w0 := f_in hwf j0 + have w1 := f_in hwf j1 + have c2 := xonly_vconsts o2 c1 (by decide) (by decide) + simp only [post2, xmov, xb] + vrund [g2, e2.1, e2.2, m2, hdx, a1, w0, w1, sx32] + have r81 : s2.gpr .r8 = s.gpr .r8 + BitVec.signExtend 64 dz := by rw [o2.gpr, h81] + have rc1 : s2.gpr .rcx = s.gpr .rcx := by rw [o2.gpr, o1'.keep.gpr (by decide)] + refine ⟨?_, by rw [show (32 : BitVec 64) = BitVec.ofNat 64 (4 * 8) from rfl, coeffAddr_add, Nat.mul_succ], r81, + by rw [rc1], by rw [rc1], ?_⟩ + · -- the coefficients stored + refine polyIs_write2 hS j0 j1 (by omega) + (a := fun e => if e < 2 then (op G[8 * i + e]! G[8 * i + 2 + e]! (zetas (zi (2 * i)))).1 + else (op G[8 * i + (e - 2)]! G[8 * i + 2 + (e - 2)]! (zetas (zi (2 * i)))).2) + (b := fun e => if e < 2 then (op G[8 * i + 4 + e]! G[8 * i + 6 + e]! (zetas (zi (2 * i + 1)))).1 + else (op G[8 * i + 4 + (e - 2)]! G[8 * i + 6 + (e - 2)]! (zetas (zi (2 * i + 1)))).2) + (fun e he => ?_) (fun e he => ?_) (fun j hj => ?_) + · rw [dword_punpcklqdq _ _ he] + split + · rw [a0 e he]; dsimp only + rw [ite_eq_left (by omega), show 8 * i + e + 2 * (e / 2) = 8 * i + e by omega, + show 8 * i + 2 + e + 2 * (e / 2) = 8 * i + 2 + e by omega, show 2 * i + e / 2 = 2 * i by omega] + · rw [a3 (e - 2) (by omega)]; dsimp only + rw [ite_eq_right (by omega), show 8 * i + (e - 2) + 2 * ((e - 2) / 2) = 8 * i + (e - 2) by omega, + show 8 * i + 2 + (e - 2) + 2 * ((e - 2) / 2) = 8 * i + 2 + (e - 2) by omega, + show 2 * i + (e - 2) / 2 = 2 * i by omega] + · rw [dword_punpckhqdq _ _ he, eval_movdqa] + split + · rw [a0 (2 + e) (by omega)]; dsimp only + rw [ite_eq_left (by omega), show 8 * i + (2 + e) + 2 * ((2 + e) / 2) = 8 * i + 4 + e by omega, + show 8 * i + 2 + (2 + e) + 2 * ((2 + e) / 2) = 8 * i + 6 + e by omega, + show 2 * i + (2 + e) / 2 = 2 * i + 1 by omega] + · rw [a3 e he]; dsimp only + rw [ite_eq_right (by omega), show 8 * i + e + 2 * (e / 2) = 8 * i + 4 + (e - 2) by omega, + show 8 * i + 2 + e + 2 * (e / 2) = 8 * i + 6 + (e - 2) by omega, + show 2 * i + e / 2 = 2 * i + 1 by omega] + · -- the specification: two blocks + rw [← hG, show 2 * (i + 1) = 2 * i + 1 + 1 by omega, layF, foldl_range_succ, foldl_range_succ, ← layF, + hG, show 2 * 2 * (2 * i) = 8 * i by omega, show 2 * 2 * (2 * i + 1) = 8 * i + 4 by omega] + have hn : ∀ j, j < 256 → j < n := fun j h => by rw [n_eq]; exact h + rw [hblk.get _ _ _ _ _ (by decide) (by decide) (by rw [n_eq]; omega) _ (hn j hj)] + have p2 := fun j (h : j < 256) => hblk.get G 2 (zi (2 * i)) (8 * i) 2 (by decide) (by decide) + (by rw [n_eq]; omega) j (hn j h) + rcases (by omega : j < 8 * i ∨ (8 * i ≤ j ∧ j < 8 * i + 2) ∨ (8 * i + 2 ≤ j ∧ j < 8 * i + 4) ∨ + (8 * i + 4 ≤ j ∧ j < 8 * i + 6) ∨ (8 * i + 6 ≤ j ∧ j < 8 * i + 8) ∨ 8 * i + 8 ≤ j) with + h | h | h | h | h | h + · simp (disch := omega) only [ite_eq_left, ite_eq_right, p2 j hj] + · simp (disch := omega) only [ite_eq_left, ite_eq_right, p2 j hj] + rw [show 8 * i + (j - 8 * i) = j by omega, show 8 * i + 2 + (j - 8 * i) = j + 2 by omega] + · simp (disch := omega) only [ite_eq_left, ite_eq_right, p2 j hj] + rw [show 8 * i + (j - 8 * i - 2) = j - 2 by omega, show 8 * i + 2 + (j - 8 * i - 2) = j by omega] + · simp (disch := omega) only [ite_eq_left, ite_eq_right, p2 j hj, p2 (j + 2) (by omega)] + rw [show 8 * i + 4 + (j - (8 * i + 4)) = j by omega, + show 8 * i + 6 + (j - (8 * i + 4)) = j + 2 by omega] + · simp (disch := omega) only [ite_eq_left, ite_eq_right, p2 j hj, p2 (j - 2) (by omega)] + rw [show 8 * i + 4 + (j - (8 * i + 4) - 2) = j - 2 by omega, + show 8 * i + 6 + (j - (8 * i + 4) - 2) = j by omega] + · simp (disch := omega) only [ite_eq_left, ite_eq_right, p2 j hj] + · -- what the step keeps + refine ⟨⟨fun r hr => ?_, by simp only [RegUpd.rd_setReg, RegUpd.rd_setFlags], + by simp only [RegUpd.wr_setReg, RegUpd.wr_setFlags]⟩, frame_write2 (Frame.refl _ _) j0 j1 _ _, ⟨?_, ?_⟩, + by simp only [RegUpd.mxcsr_setReg, RegUpd.mxcsr_setFlags]; rw [o2.mxcsr, o1'.mxcsr]⟩ + · simp only [List.mem_cons, List.not_mem_nil, or_false, not_or] at hr + simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, hr, ite_false] + rw [o2.gpr, o1'.keep.gpr (by simp [hr])] + · simp only [RegUpd.xmm_setReg, RegUpd.xmm_setFlags, xmm_setXmm, reduceCtorEq, ite_false] + exact c2.q + · simp only [RegUpd.xmm_setReg, RegUpd.xmm_setFlags, xmm_setXmm, reduceCtorEq, ite_false] + exact c2.qinv + +omit hbf hblk in +/-- The prologue of the layers with `len` = 2 and 1. -/ +theorem vpre21 {fP sP : Addr} (k : Nat) (hk : k + 4 ≤ 256) {s : State} (hdi : s.gpr .rdi = fP) + (hsi : s.gpr .rsi = sP) : + WP isa (.block (([.mov .rdx (.reg .rdi)] : List Instr) ++ leaR .r8 .rsi (4 * k))) s fun w => + w.gpr .rdx = fP ∧ w.gpr .r8 = coeffAddr sP k ∧ GOnly [.rdx, .r8] s w := by + simp only [leaR] + vrund [sx_ofNat (show 4 * k < 2 ^ 31 by omega), hsi, hdi] + gonlyd + +theorem vlay2_ok {fP sP : Addr} (k : Nat) (o : BitVec 8) (dz : BitVec 32) (zi kz : Nat → Nat) (hkz0 : kz 0 = k) + (hk : ∀ i < 32, kz i + 4 ≤ 256) (hsel : ∀ i < 32, ∀ e < 4, kz i + sel o e = zi (2 * i + e / 2)) + (hstep : ∀ i < 32, coeffAddr sP (kz i) + BitVec.signExtend 64 dz = coeffAddr sP (kz (i + 1))) + {F : Poly} {s : State} (hc : VConsts s) (hdi : s.gpr .rdi = fP) (hsi : s.gpr .rsi = sP) + (hS : PolyIs s.mem fP F) (hT : Tab zmTab s.mem sP 256) (hwf : pR fP ∈ s.wr) (hw : pR sP ∈ s.wr) + (hd : (pR sP).Disjoint (pR fP)) : + WP isa (vlay2 bf k o dz) s fun s' => PolyIs s'.mem fP (layF blk F 2 zi 64) ∧ BInv fP s s' := by + have hk0 : k + 4 ≤ 256 := by have := hk 0 (by decide); omega + refine WP.seq (WP.mono (vpre21 k hk0 hdi hsi) fun w ⟨hdx, h8, og⟩ => ?_) + refine WP.mono (wp_rcxLoop (N := 32) (by decide) (by decide) + (fun i u => PolyIs u.mem fP (layF blk F 2 zi (2 * i)) ∧ u.gpr .rdx = coeffAddr fP (8 * i) ∧ + u.gpr .r8 = coeffAddr sP (kz i) ∧ BInv fP w u) + (fun u ou _ => ⟨by rw [ou.mem, og.mem]; exact hS, + by rw [ou.keep.gpr (by decide), hdx, Nat.mul_zero, coeffAddr, Nat.mul_zero, add_ofNat_zero], + by rw [ou.keep.gpr (by decide), h8, hkz0], ⟨ou.keep.mono (by simp), by rw [ou.mem]; exact Frame.refl _ _, + gonly_vconsts ou (gonly_vconsts og hc), ou.mxcsr⟩⟩) + (fun i hi u ⟨hS', hdx', h8', hb'⟩ => ?_)) fun u ⟨hS', _, _, hb'⟩ => + ⟨hS', ⟨(og.keep.trans hb'.keep).mono (by simp), by rw [← og.mem]; exact hb'.frame, hb'.consts, + by rw [hb'.mxcsr, og.mxcsr]⟩⟩ + have hT' : Tab zmTab u.mem sP 256 := (by rw [og.mem]; exact hT : Tab zmTab w.mem sP 256).frame hb'.frame + (fun r hr => by rw [List.mem_singleton.mp hr]; exact hd) (by decide) + have hwf' : pR fP ∈ u.wr := by rw [hb'.keep.2.2, og.keep.2.2]; exact hwf + have hw' : pR sP ∈ u.wr := by rw [hb'.keep.2.2, og.keep.2.2]; exact hw + rw [show [Instr.movdquLoad .xmm0 (at_ .rdx 0), .movdquLoad .xmm1 (at_ .rdx 16)] ++ vzeta o ++ + [.alu .add .r8 (.imm dz), xmov .xmm2 .xmm0, xb .punpcklqdq .xmm0 .xmm1, xb .punpckhqdq .xmm2 .xmm1, + xmov .xmm1 .xmm2] ++ bf ++ [xmov .xmm1 .xmm0, xb .punpcklqdq .xmm0 .xmm3, xb .punpckhqdq .xmm1 .xmm3, + .movdquStore (at_ .rdx 0) .xmm0, .movdquStore (at_ .rdx 16) .xmm1, .alu .add .rdx (.imm 32)] ++ + [.alu .sub .rcx (.imm 1)] = pre2 o dz ++ (bf ++ post2) by simp [List.append_assoc]] + exact WP.mono (vstep2 hbf hblk hi o dz zi (hk i hi) (hsel i hi) hb'.consts hdx' h8' hS' hT' hwf' hw') + fun u' ⟨hS'', hdx'', h8'', hcx, hzf, hb''⟩ => ⟨⟨hS'', hdx'', by rw [h8'', h8', hstep i hi], + hb'.trans hb''⟩, hcx, hzf⟩ + +/-! ## The layer with `len = 1` -/ + +/-- The loads, the zetas and the gathering of the coefficients. -/ +abbrev pre1 (o : BitVec 8) (dz : BitVec 32) : List Instr := + [.movdquLoad .xmm0 (at_ .rdx 0), .movdquLoad .xmm2 (at_ .rdx 16)] ++ vzeta o ++ + [.alu .add .r8 (.imm dz), .xop (.pshufd .xmm0 .xmm0 0xD8), .xop (.pshufd .xmm2 .xmm2 0xD8), + xmov .xmm1 .xmm0, xb .punpcklqdq .xmm0 .xmm2, xb .punpckhqdq .xmm1 .xmm2] + +/-- The interleaving back, the stores and the counts. -/ +abbrev post1 : List Instr := + [xmov .xmm1 .xmm0, xb .punpckldq .xmm0 .xmm3, xb .punpckhdq .xmm1 .xmm3, + .movdquStore (at_ .rdx 0) .xmm0, .movdquStore (at_ .rdx 16) .xmm1, .alu .add .rdx (.imm 32), + .alu .sub .rcx (.imm 1)] + +omit hbf in +/-- Each coefficient after the first `b` blocks of the layer with `len = 1`. -/ +theorem layF1_get (F : Poly) (zi : Nat → Nat) {b : Nat} (hb : b ≤ 128) {j : Nat} (hj : j < 256) : + (layF blk F 1 zi b)[j]! = if j < 2 * b then + (if j % 2 = 0 then (op F[j]! F[j + 1]! (zetas (zi (j / 2)))).1 + else (op F[j - 1]! F[j]! (zetas (zi (j / 2)))).2) else F[j]! := by + induction b generalizing j with + | zero => rw [ite_eq_right (by omega)]; rfl + | succ b ih => + rw [layF, foldl_range_succ, ← layF, + hblk.get _ 1 _ _ 1 (by decide) (by decide) (by rw [n_eq]; omega) j (by rw [n_eq]; exact hj)] + by_cases h1 : 2 * 1 * b ≤ j ∧ j < 2 * 1 * b + 1 + · rw [ite_eq_left h1, ih (by omega) hj, ih (by omega) (by omega), show j / 2 = b by omega] + simp (disch := omega) only [ite_eq_left, ite_eq_right] + · rw [ite_eq_right h1] + by_cases h2 : 2 * 1 * b + 1 ≤ j ∧ j < 2 * 1 * b + 1 + 1 + · rw [ite_eq_left h2, ih (by omega) (by omega), ih (by omega) hj, show j / 2 = b by omega] + simp (disch := omega) only [ite_eq_left, ite_eq_right] + · rw [ite_eq_right h2, ih (by omega) hj] + by_cases h3 : j < 2 * b <;> simp (disch := omega) only [ite_eq_left, ite_eq_right] + +theorem vstep1 {fP sP : Addr} {i kz : Nat} (hi : i < 32) (o : BitVec 8) (dz : BitVec 32) (zi : Nat → Nat) + (hk : kz + 4 ≤ 256) (hsel : ∀ e < 4, kz + sel o e = zi (4 * i + e)) + {F : Poly} {s : State} (hc : VConsts s) (hdx : s.gpr .rdx = coeffAddr fP (8 * i)) + (h8 : s.gpr .r8 = coeffAddr sP kz) (hS : PolyIs s.mem fP (layF blk F 1 zi (4 * i))) + (hT : Tab zmTab s.mem sP 256) (hwf : pR fP ∈ s.wr) (hw : pR sP ∈ s.wr) : + WP isa (.block (pre1 o dz ++ (bf ++ post1))) s fun s' => + PolyIs s'.mem fP (layF blk F 1 zi (4 * (i + 1))) ∧ s'.gpr .rdx = coeffAddr fP (8 * (i + 1)) ∧ + s'.gpr .r8 = s.gpr .r8 + BitVec.signExtend 64 dz ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ + s'.zf = some (s.gpr .rcx - 1 == 0) ∧ BInv fP s s' := by + have j0 : 8 * i + 4 ≤ 256 := by omega + have j1 : 8 * i + 4 + 4 ≤ 256 := by omega + have a1 : coeffAddr fP (8 * i) + BitVec.ofNat 64 16 = coeffAddr fP (8 * i + 4) := coeffAddr_add _ _ 4 + have r0 := f_in (List.mem_append_right s.rd hwf) j0 + have r1 := f_in (List.mem_append_right s.rd hwf) j1 + have hk' : ∀ j < 4, kz + sel o j < 256 := fun j _ => by have := sel_lt o j; omega + generalize hG : layF blk F 1 zi (4 * i) = G at hS + have lx := dlanes_load hS j0 + have ly := dlanes_load hS j1 + rw [WP.block_append_iff, show pre1 o dz = [.movdquLoad .xmm0 (at_ .rdx 0), .movdquLoad .xmm2 (at_ .rdx 16)] ++ + (vzeta o ++ [.alu .add .r8 (.imm dz), .xop (.pshufd .xmm0 .xmm0 0xD8), .xop (.pshufd .xmm2 .xmm2 0xD8), + xmov .xmm1 .xmm0, xb .punpcklqdq .xmm0 .xmm2, xb .punpckhqdq .xmm1 .xmm2]) by simp, WP.block_append_iff] + vrund [hdx, a1, r0, r1] + rw [WP.block_append_iff] + refine WP.mono (vzeta_ok o hk' (by simp only [RegUpd.gpr_setXmm]; exact h8) + (by simp only [RegUpd.rd_setXmm, RegUpd.wr_setXmm]; exact tab_in (List.mem_append_right _ hw) hk) + (by simp only [RegUpd.mem_setXmm]; exact hT)) fun s1 ⟨z1, zo1, o1⟩ => ?_ + refine WP.mono (Q := fun (s1' : State) => DLanes (s1'.xmm .xmm0) (fun e => G[8 * i + 2 * e]!) ∧ + DLanes (s1'.xmm .xmm1) (fun e => G[8 * i + 2 * e + 1]!) ∧ + ZLanes (s1'.xmm .xmm13) (fun e => zetas (zi (4 * i + e))) ∧ ZOdd (s1'.xmm .xmm13) (s1'.xmm .xmm12) ∧ + VConsts s1' ∧ s1'.gpr .r8 = s.gpr .r8 + BitVec.signExtend 64 dz ∧ GKeep [.r8] s s1') ?_ + fun s1' ⟨l0, l1, l13, l12, c1, h81, o1'⟩ => ?_ + · have g1 : s1.gpr = s.gpr := by rw [o1.gpr]; rfl + have m1 : s1.mem = s.mem := by rw [o1.mem]; rfl + have e1 : s1.rd = s.rd ∧ s1.wr = s.wr := ⟨by rw [o1.rd]; rfl, by rw [o1.wr]; rfl⟩ + have x1 : s1.mxcsr = s.mxcsr := by rw [o1.mxcsr]; rfl + have x0 : s1.xmm .xmm0 = s.mem.readW (coeffAddr fP (8 * i)) 128 := by + rw [o1.xmm _ (by decide)]; simp only [xmm_setXmm]; rfl + have x2 : s1.xmm .xmm2 = s.mem.readW (coeffAddr fP (8 * i + 4)) 128 := by + rw [o1.xmm _ (by decide)]; simp only [xmm_setXmm]; rfl + have c1 : VConsts s1 := xonly_vconsts o1 ((hc.setXmm (by decide) (by decide) _).setXmm (by decide) + (by decide) _) (by decide) (by decide) + vrund [g1, m1, e1.1, e1.2, x1, eval_movdqa] + refine ⟨?_, ?_, ?_, ?_, ⟨?_, ?_⟩, ⟨⟨fun r hr => ?_, ?_, ?_⟩, ?_, ?_⟩⟩ + · intro e he + rw [dword_punpcklqdq _ _ he, x0, x2] + split + · rw [dword_d8 _ he, ite_eq_left (by omega), lx _ (by omega)] + · rw [dword_d8 _ (by omega), ite_eq_left (by omega), ly _ (by omega)]; dsimp only + rw [show 8 * i + 4 + 2 * (e - 2) = 8 * i + 2 * e by omega] + · intro e he + rw [dword_punpckhqdq _ _ he, x0, x2] + split + · rw [dword_d8 _ (by omega), ite_eq_right (by omega), lx _ (by omega)]; dsimp only + rw [show 8 * i + (2 * (2 + e - 2) + 1) = 8 * i + 2 * e + 1 by omega] + · rw [dword_d8 _ he, ite_eq_right (by omega), ly _ (by omega)]; dsimp only + rw [show 8 * i + 4 + (2 * (e - 2) + 1) = 8 * i + 2 * e + 1 by omega] + · intro e he + rw [z1 e he]; dsimp only; rw [hsel e he] + · exact zo1 + · exact c1.q + · exact c1.qinv + · simp only [RegUpd.gpr_setXmm, RegUpd.gpr_setReg, RegUpd.gpr_setFlags, g1] + rw [ifn (by simpa using hr)] + all_goals simp only [RegUpd.rd_setXmm, RegUpd.wr_setXmm, RegUpd.mem_setXmm, mxcsr_setXmm, + RegUpd.rd_setReg, RegUpd.wr_setReg, RegUpd.mem_setReg, RegUpd.mxcsr_setReg, RegUpd.rd_setFlags, + RegUpd.wr_setFlags, RegUpd.mem_setFlags, RegUpd.mxcsr_setFlags, e1.1, e1.2, m1, x1] + rw [WP.block_append_iff] + refine WP.mono (hbf _ c1 _ _ _ l0 l1 l13 l12) fun s2 ⟨a0, a3, o2⟩ => ?_ + have g2 : s2.gpr .rdx = s.gpr .rdx := by rw [o2.gpr, o1'.keep.gpr (by decide)] + have e2 : s2.rd = s.rd ∧ s2.wr = s.wr := ⟨by rw [o2.rd, o1'.keep.2.1], by rw [o2.wr, o1'.keep.2.2]⟩ + have m2 : s2.mem = s.mem := by rw [o2.mem, o1'.mem] + have w0 := f_in hwf j0 + have w1 := f_in hwf j1 + have c2 := xonly_vconsts o2 c1 (by decide) (by decide) + simp only [post1, xmov, xb] + vrund [g2, e2.1, e2.2, m2, hdx, a1, w0, w1, sx32] + have r81 : s2.gpr .r8 = s.gpr .r8 + BitVec.signExtend 64 dz := by rw [o2.gpr, h81] + have rc1 : s2.gpr .rcx = s.gpr .rcx := by rw [o2.gpr, o1'.keep.gpr (by decide)] + have lg : ∀ b, b ≤ 128 → ∀ j, j < 256 → _ := fun b hb j hj => layF1_get hblk F zi (b := b) hb (j := j) hj + refine ⟨?_, by rw [show (32 : BitVec 64) = BitVec.ofNat 64 (4 * 8) from rfl, coeffAddr_add, Nat.mul_succ], r81, + by rw [rc1], by rw [rc1], ?_⟩ + · -- the coefficients stored + refine polyIs_write2 hS j0 j1 (by omega) + (a := fun e => (layF blk F 1 zi (4 * (i + 1)))[8 * i + e]!) + (b := fun e => (layF blk F 1 zi (4 * (i + 1)))[8 * i + 4 + e]!) + (fun e he => ?_) (fun e he => ?_) (fun j hj => ?_) + · rw [dword_punpckldq' _ _ he] + split + · rw [a0 (e / 2) (by omega)]; dsimp only + rw [show 8 * i + 2 * (e / 2) = 8 * i + e by omega, + show 4 * i + e / 2 = (8 * i + e) / 2 by omega, ← hG] + simp (disch := omega) only [lg, ite_eq_left, ite_eq_right] + · rw [a3 (e / 2) (by omega)]; dsimp only + rw [show 8 * i + 2 * (e / 2) = 8 * i + e - 1 by omega, show 8 * i + e - 1 + 1 = 8 * i + e by omega, + show 4 * i + e / 2 = (8 * i + e) / 2 by omega, ← hG] + simp (disch := omega) only [lg, ite_eq_left, ite_eq_right] + · rw [dword_punpckhdq' _ _ he, eval_movdqa] + split + · rw [a0 (2 + e / 2) (by omega)]; dsimp only + rw [show 8 * i + 2 * (2 + e / 2) = 8 * i + 4 + e by omega, + show 4 * i + (2 + e / 2) = (8 * i + 4 + e) / 2 by omega, ← hG] + simp (disch := omega) only [lg, ite_eq_left, ite_eq_right] + · rw [a3 (2 + e / 2) (by omega)]; dsimp only + rw [show 8 * i + 2 * (2 + e / 2) = 8 * i + 4 + e - 1 by omega, + show 8 * i + 4 + e - 1 + 1 = 8 * i + 4 + e by omega, + show 4 * i + (2 + e / 2) = (8 * i + 4 + e) / 2 by omega, ← hG] + simp (disch := omega) only [lg, ite_eq_left, ite_eq_right] + · rcases (by omega : (8 * i ≤ j ∧ j < 8 * i + 4) ∨ (8 * i + 4 ≤ j ∧ j < 8 * i + 8) ∨ + j < 8 * i ∨ 8 * i + 8 ≤ j) with h | h | h | h + · rw [ite_eq_left h, show 8 * i + (j - 8 * i) = j by omega] + · rw [ite_eq_right (by omega), ite_eq_left h, show 8 * i + 4 + (j - (8 * i + 4)) = j by omega] + · rw [ite_eq_right (by omega), ite_eq_right (by omega), ← hG] + simp (disch := omega) only [lg, ite_eq_left] + · rw [ite_eq_right (by omega), ite_eq_right (by omega), ← hG] + simp (disch := omega) only [lg, ite_eq_left, ite_eq_right] + · -- what the step keeps + refine ⟨⟨fun r hr => ?_, by simp only [RegUpd.rd_setReg, RegUpd.rd_setFlags], + by simp only [RegUpd.wr_setReg, RegUpd.wr_setFlags]⟩, frame_write2 (Frame.refl _ _) j0 j1 _ _, ⟨?_, ?_⟩, + by simp only [RegUpd.mxcsr_setReg, RegUpd.mxcsr_setFlags]; rw [o2.mxcsr, o1'.mxcsr]⟩ + · simp only [List.mem_cons, List.not_mem_nil, or_false, not_or] at hr + simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, hr, ite_false] + rw [o2.gpr, o1'.keep.gpr (by simp [hr])] + · simp only [RegUpd.xmm_setReg, RegUpd.xmm_setFlags, xmm_setXmm, reduceCtorEq, ite_false] + exact c2.q + · simp only [RegUpd.xmm_setReg, RegUpd.xmm_setFlags, xmm_setXmm, reduceCtorEq, ite_false] + exact c2.qinv + +theorem vlay1_ok {fP sP : Addr} (k : Nat) (o : BitVec 8) (dz : BitVec 32) (zi kz : Nat → Nat) (hkz0 : kz 0 = k) + (hk : ∀ i < 32, kz i + 4 ≤ 256) (hsel : ∀ i < 32, ∀ e < 4, kz i + sel o e = zi (4 * i + e)) + (hstep : ∀ i < 32, coeffAddr sP (kz i) + BitVec.signExtend 64 dz = coeffAddr sP (kz (i + 1))) + {F : Poly} {s : State} (hc : VConsts s) (hdi : s.gpr .rdi = fP) (hsi : s.gpr .rsi = sP) + (hS : PolyIs s.mem fP F) (hT : Tab zmTab s.mem sP 256) (hwf : pR fP ∈ s.wr) (hw : pR sP ∈ s.wr) + (hd : (pR sP).Disjoint (pR fP)) : + WP isa (vlay1 bf k o dz) s fun s' => PolyIs s'.mem fP (layF blk F 1 zi 128) ∧ BInv fP s s' := by + have hk0 : k + 4 ≤ 256 := by have := hk 0 (by decide); omega + refine WP.seq (WP.mono (vpre21 k hk0 hdi hsi) fun w ⟨hdx, h8, og⟩ => ?_) + refine WP.mono (wp_rcxLoop (N := 32) (by decide) (by decide) + (fun i u => PolyIs u.mem fP (layF blk F 1 zi (4 * i)) ∧ u.gpr .rdx = coeffAddr fP (8 * i) ∧ + u.gpr .r8 = coeffAddr sP (kz i) ∧ BInv fP w u) + (fun u ou _ => ⟨by rw [ou.mem, og.mem]; exact hS, + by rw [ou.keep.gpr (by decide), hdx, Nat.mul_zero, coeffAddr, Nat.mul_zero, add_ofNat_zero], + by rw [ou.keep.gpr (by decide), h8, hkz0], ⟨ou.keep.mono (by simp), by rw [ou.mem]; exact Frame.refl _ _, + gonly_vconsts ou (gonly_vconsts og hc), ou.mxcsr⟩⟩) + (fun i hi u ⟨hS', hdx', h8', hb'⟩ => ?_)) fun u ⟨hS', _, _, hb'⟩ => + ⟨hS', ⟨(og.keep.trans hb'.keep).mono (by simp), by rw [← og.mem]; exact hb'.frame, hb'.consts, + by rw [hb'.mxcsr, og.mxcsr]⟩⟩ + have hT' : Tab zmTab u.mem sP 256 := (by rw [og.mem]; exact hT : Tab zmTab w.mem sP 256).frame hb'.frame + (fun r hr => by rw [List.mem_singleton.mp hr]; exact hd) (by decide) + have hwf' : pR fP ∈ u.wr := by rw [hb'.keep.2.2, og.keep.2.2]; exact hwf + have hw' : pR sP ∈ u.wr := by rw [hb'.keep.2.2, og.keep.2.2]; exact hw + rw [show [Instr.movdquLoad .xmm0 (at_ .rdx 0), .movdquLoad .xmm2 (at_ .rdx 16)] ++ vzeta o ++ + [.alu .add .r8 (.imm dz), .xop (.pshufd .xmm0 .xmm0 0xD8), .xop (.pshufd .xmm2 .xmm2 0xD8), + xmov .xmm1 .xmm0, xb .punpcklqdq .xmm0 .xmm2, xb .punpckhqdq .xmm1 .xmm2] ++ bf ++ + [xmov .xmm1 .xmm0, xb .punpckldq .xmm0 .xmm3, xb .punpckhdq .xmm1 .xmm3, + .movdquStore (at_ .rdx 0) .xmm0, .movdquStore (at_ .rdx 16) .xmm1, .alu .add .rdx (.imm 32)] ++ + [.alu .sub .rcx (.imm 1)] = pre1 o dz ++ (bf ++ post1) by simp [List.append_assoc]] + exact WP.mono (vstep1 hbf hblk hi o dz zi (hk i hi) (hsel i hi) hb'.consts hdx' h8' hS' hT' hwf' hw') + fun u' ⟨hS'', hdx'', h8'', hcx, hzf, hb''⟩ => ⟨⟨hS'', hdx'', by rw [h8'', h8', hstep i hi], + hb'.trans hb''⟩, hcx, hzf⟩ + +end + +end VG.Proof.MlDsa.X86_64.Arith diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VMem.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VMem.lean new file mode 100644 index 000000000..aff2c0409 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/VMem.lean @@ -0,0 +1,104 @@ +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.VLanes +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Table +import VerifiedGarbage.Proof.MlKem.X86_64.VMem +import VerifiedGarbage.Proof.MlDsa.Arith.Ntt + +/-! +# ML-DSA on x86-64: four coefficients at a time in memory + +Untrusted: everything here is checked by Lean. 16-byte loads of four +coefficients of a stored polynomial (`dlanes_load`) and stores of them +(`polyIs_write2`), and the table of the zetas in Montgomery form +(`Tab zmTab`), from which `vzeta` loads the zetas of up to four blocks +(`vzeta_ok`). +-/ + +namespace VG.Proof.MlDsa.X86_64.Arith + +open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Arith +open VG.Proof.MlDsa.Arith +open VG.Proof.MlKem.X86_64 (XOnly xmm_setXmm ifp ifn sel sel_lt add_ofNat_zero) +open VG.Spec.MlDsa (q n Poly Zq coeffAt polyAt Reduced PolyIs zetas) + +theorem dlanes_load {m : Mem} {p : Addr} {F : Poly} (h : PolyIs m p F) {j : Nat} (hj : j + 4 ≤ 256) : + DLanes (m.readW (coeffAddr p j) 128) (fun e => F[j + e]!) := fun e he => by + rw [dword_readW _ _ he, coeffAddr_add, ← coeffAt_eq] + exact polyIs_toNat h (by rw [n_eq]; omega) + +/-- Coefficient `i` after storing `x` at coefficient `j`. -/ +theorem coeffAt_write128 (m : Mem) (p : Addr) {j : Nat} (hj : j + 4 ≤ 256) (x : BitVec 128) {i : Nat} + (hi : i < 256) : + coeffAt (m.writeW (coeffAddr p j) x) p i = if j ≤ i ∧ i < j + 4 then dword x (i - j) else coeffAt m p i := by + split + · rename_i h + rw [coeffAt_eq, show coeffAddr p i = coeffAddr p j + BitVec.ofNat 64 (4 * (i - j)) by + rw [coeffAddr_add, show j + (i - j) = i by omega]] + exact readW_writeW128 _ _ _ (by omega) + · exact Mem.readW_writeW_sep (Offset.sep p (by omega) (by omega) (by omega)) (by decide) + +/-- Two vectors stored into a polynomial, with the lanes `a` and `b`. -/ +theorem polyIs_write2 {m : Mem} {p : Addr} {P R : Poly} (hP : PolyIs m p P) {j j' : Nat} + (hj : j + 4 ≤ 256) (hj' : j' + 4 ≤ 256) (hsep : j + 4 ≤ j' ∨ j' + 4 ≤ j) {x y : BitVec 128} + {a b : Nat → Zq} (hx : DLanes x a) (hy : DLanes y b) + (hR : ∀ i < 256, R[i]! = if j ≤ i ∧ i < j + 4 then a (i - j) + else if j' ≤ i ∧ i < j' + 4 then b (i - j') else P[i]!) : + PolyIs ((m.writeW (coeffAddr p j) x).writeW (coeffAddr p j') y) p R := polyIs_of_toNat fun i hi => by + rw [n_eq] at hi + rw [coeffAt_write128 _ _ hj' _ hi, coeffAt_write128 _ _ hj _ hi, hR i hi] + by_cases h1 : j' ≤ i ∧ i < j' + 4 + · rw [ite_eq_left_of_eq_true _ _ (eq_true h1), ite_eq_right_of_eq_false _ _ (eq_false (by omega)), + ite_eq_left_of_eq_true _ _ (eq_true h1)] + exact hy _ (by omega) + · rw [ite_eq_right_of_eq_false _ _ (eq_false h1)] + by_cases h2 : j ≤ i ∧ i < j + 4 + · rw [ite_eq_left_of_eq_true _ _ (eq_true h2), ite_eq_left_of_eq_true _ _ (eq_true h2)] + exact hx _ (by omega) + · rw [ite_eq_right_of_eq_false _ _ (eq_false h2), ite_eq_right_of_eq_false _ _ (eq_false h2), + ite_eq_right_of_eq_false _ _ (eq_false h1)] + exact polyIs_toNat hP (by rw [n_eq]; exact hi) + +theorem pR_contains (p : Addr) {j : Nat} (hj : j + 4 ≤ 256) : (pR p).Contains (coeffAddr p j) 16 := + Offset.contains_base p (by omega) (by omega) + +theorem frame_write2 {m m' : Mem} {p : Addr} (hf : Frame [pR p] m m') {j j' : Nat} (hj : j + 4 ≤ 256) + (hj' : j' + 4 ≤ 256) (x y : BitVec 128) : + Frame [pR p] m ((m'.writeW (coeffAddr p j) x).writeW (coeffAddr p j') y) := + (hf.writeW (List.mem_singleton_self _) x (pR_contains p hj)).writeW (List.mem_singleton_self _) y + (pR_contains p hj') + +/-! ## The table of zetas -/ + +theorem zmTab_lt (k : Nat) : zmTab k < q := Nat.mod_lt _ (by decide) + +theorem zmTab_eq (k : Nat) : zmTab k = (zetas k).val * 2 ^ 32 % q := by + rw [zmTab, ← zetaNat_eq, zetaNat, Nat.mod_mul_mod] + +/-- The zeta at index `k` of the table. -/ +theorem tab_zeta {m : Mem} {zP : Addr} (ht : Tab zmTab m zP 256) {k : Nat} (hk : k < 256) : + (m.readW (coeffAddr zP k) 32).toNat = (zetas k).val * 2 ^ 32 % q := by + rw [← coeffAt_eq, ht k hk, BitVec.toNat_ofNat, Nat.mod_eq_of_lt (Nat.lt_trans (zmTab_lt k) (by decide)), + zmTab_eq] + +theorem dword_shufDwords_sel (a : BitVec 128) (o : BitVec 8) {i : Nat} (hi : i < 4) : + dword (shufDwords a o) i = dword a (sel o i) := dword_shufDwords a o hi + +theorem vzeta_ok (o : BitVec 8) {zP : Addr} {k : Nat} (hk : ∀ j < 4, k + sel o j < 256) {s : State} + (h8 : s.gpr .r8 = coeffAddr zP k) (hin : InRegions (s.rd ++ s.wr) (coeffAddr zP k) 16) + (ht : Tab zmTab s.mem zP 256) : + WP isa (.block (vzeta o)) s fun s' => + ZLanes (s'.xmm .xmm13) (fun i => zetas (k + sel o i)) ∧ ZOdd (s'.xmm .xmm13) (s'.xmm .xmm12) ∧ + XOnly [.xmm13, .xmm12] s s' := by + simp only [vzeta] + apply WP.of_runBlock + simp only [runBlock_cons, runStep_some, runBlock_nil, exec, XOp.exec, State.load128, ea_atD, + add_ofNat_zero, h8, hin, ite_true, Option.map_some, Option.some.injEq, exists_eq_left'] + refine ⟨fun i hi => ?_, fun j hj => ?_, by xonly⟩ + · simp only [xmm_setXmm, ite_true, ite_false, reduceCtorEq] + have hs := sel_lt o i + rw [dword_shufDwords_sel _ _ hi, dword_readW _ _ hs, coeffAddr_add] + exact tab_zeta ht (hk i hi) + · simp only [xmm_setXmm, ite_true, ite_false, reduceCtorEq] + rw [dword_shufDwords _ _ (by omega)] + rcases (by omega : j = 0 ∨ j = 1) with rfl | rfl <;> rfl + +end VG.Proof.MlDsa.X86_64.Arith diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Correct.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Correct.lean index 115c6bc54..6b95ef645 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Correct.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Correct.lean @@ -1,4 +1,5 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Sign.PhaseO +import VerifiedGarbage.Proof.Framework.X86_64.Mxcsr /-! # ML-DSA signing on x86-64: correctness @@ -153,7 +154,7 @@ theorem entry_st {P : Prims} {D : Nat} (hP : PrimsOk P D) {p : Params} (h3 : Ok3 entry_bytes hf d6 (by omega) (ht.regs (.r13, .rdx) (by decide))⟩ theorem sign_correct {P : Prims} {D : Nat} (hP : PrimsOk P D) {p : Params} (h3 : Ok3 p) - (hmx : (Impl.MlDsa.X86_64.Sign.sign P p).allInstrs (fun i => !loadsMxcsr i) = true) (σ : State) + (hmx : ctlOk (Impl.MlDsa.X86_64.Sign.sign P p) = true) (σ : State) (hpre : (signK p D).pre σ) : ∃ t s', Exec isa (Impl.MlDsa.X86_64.Sign.sign P p) σ t s' ∧ abiPreserved σ s' ∧ (signK p D).post σ s' := by have hc := allChk_ok h3 @@ -182,7 +183,7 @@ theorem sign_correct {P : Prims} {D : Nat} (hP : PrimsOk P D) {p : Params} (h3 : ⟨hg, s₄, h₄, hr, hm⟩ obtain ⟨t, s', he, hF⟩ := main obtain ⟨hg, s₄, h₄, hr, hm⟩ := hF - refine ⟨t, s', he, abiPreserved_of_exec hmx he hg, ?_⟩ + refine ⟨t, s', he, abiPreserved_of_ctl hmx he hg, ?_⟩ have e14 : pa s₄ (.r14, 0) = σ.gpr .rcx := by rw [pa, h₄.st.top.regs (.r14, .rcx) (by decide), VG.Proof.MlKem.X86_64.add_ofNat_zero] show Outcome _ _ _ diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean index de1fbb3cd..089924f6d 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean @@ -63,7 +63,7 @@ theorem signK_implies {p : Params} (h3 : Ok3 p) : · sig_implies_sat [signContractT, signSig, X86_64.abi, X86_64.argRegs] [signSat] using signSat mlDsa87 theorem sign_verified {p : Params} (h3 : Ok3 p) - (hmx : (Impl.MlDsa.X86_64.Sign.sign prims p).allInstrs (fun i => !loadsMxcsr i) = true) : + (hmx : ctlOk (Impl.MlDsa.X86_64.Sign.sign prims p) = true) : Verified X86_64.target (Impl.MlDsa.X86_64.Sign.sign prims p) (signContractT p X86_64.abi signStack) := Verified.of_correct (sign_correct prims_ok h3 hmx) (sign_ct prims_ok h3) (signK_implies h3) diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Correct.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Correct.lean index 934c283a8..94176d623 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Correct.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Correct.lean @@ -128,6 +128,6 @@ theorem verify_correct {P : Prims} (C : PrimsOk P) {p : Params} (hp : p ∈ para rcases hr with ⟨e, hb⟩ | ⟨e, hb⟩ · exact .inl ⟨by rw [hres, e]; rfl, hb⟩ · exact .inr ⟨by rw [hres, e]; rfl, hb⟩⟩ : gprPreserved σ s₃ ∧ (verifyK p).post σ s₃))) - exact ⟨t, s', he, abiPreserved_of_exec (verify_mxcsr C hp) he hF.1, hF.2⟩ + exact ⟨t, s', he, abiPreserved_of_ctl (verify_ctl C hp) he hF.1, hF.2⟩ end VG.Proof.MlDsa.X86_64.Verify diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Entry.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Entry.lean index e00400b78..430f0bb51 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Entry.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Entry.lean @@ -1,6 +1,7 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Verify.Call import VerifiedGarbage.Proof.MlKem.X86_64.KCall import VerifiedGarbage.Proof.Framework.X86_64.Abi +import VerifiedGarbage.Proof.Framework.X86_64.Mxcsr /-! # ML-DSA verification on x86-64: entry to a callee @@ -25,14 +26,14 @@ open VG.Spec.Sha3 (bytesAt) /-- A callee: correct and constant time under the contract `k` (a shared contract with 16 bytes of stack), not writing `rsp`, calling at most two -deep, and never loading MXCSR or writing the stack pointer (which its -callers' artifacts check). -/ +deep, loading MXCSR only to restore it (`ctlOk`), and never writing the +stack pointer (which its callers' artifacts check). -/ structure CalleeOk (c : Prog isa) (k : Contract isa) : Prop where correct : ∀ s, k.pre s → ∃ t s', Exec isa c s t s' ∧ abiPreserved s s' ∧ k.post s s' ct : ConstantTime isa k.pre k.pub c nosp : NoSp c depth : c.depth ≤ 2 - mxcsr : c.allInstrs (fun i => !loadsMxcsr i) = true + ctl : ctlOk c = true spSafe : c.all (fun i => !isa.writesSp i) = true /-- A callee verified against its shared contract with at most 16 bytes of stack. -/ @@ -40,7 +41,7 @@ theorem CalleeOk.of_verified {c : Prog isa} {sig : Sig} {pre : Curry (sig.words {post : sig.Post X86_64.abi.ptrBits} {wa : Bool} {leak : Option (Curry (sig.words X86_64.abi.ptrBits) (Mem → List Nat))} {n : Nat} (h : Verified X86_64.target c (sig.contract X86_64.abi pre post wa n leak)) (hn : n ≤ 16) - (hsp : NoSp c) (hd : c.depth ≤ 2) (hmx : c.allInstrs (fun i => !loadsMxcsr i) = true) + (hsp : NoSp c) (hd : c.depth ≤ 2) (hmx : ctlOk c = true) (hss : c.all (fun i => !isa.writesSp i) = true) : CalleeOk c (sig.contract X86_64.abi pre post wa 16 leak) := ⟨fun s hs => h.1 s (pre_stack hn hs), diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Instrs.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Instrs.lean index 090c626e4..ce8c82e43 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Instrs.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Instrs.lean @@ -8,8 +8,9 @@ Untrusted: everything here is checked by Lean. A property `q` of every instruction of `verify P p` (`Code.allInstrs q`) holds if it holds of every instruction of the primitives `P` and of `verify P0 p`, the same code with the primitives empty (`verify_q`), which the kernel evaluates. So it never -loads MXCSR (`verify_mxcsr`) or writes the stack pointer (`verify_spSafe`) -if the primitives do not. +writes the stack pointer (`verify_spSafe`) if the primitives do not. Likewise +for `ctlC` (`verify_c`): it loads MXCSR only to restore it (`verify_ctl`) if +the primitives do (`ctlOk`). -/ namespace VG.Proof.MlDsa.X86_64.Verify @@ -107,19 +108,87 @@ theorem Code.allInstrs_of_all {I C : Type} {q : I → Bool} {c : Code I C} (h : | block is => induction is <;> simp_all [Code.all, Code.allInstrs] | _ => simp_all [Code.all, Code.allInstrs] -theorem verify0_mxcsr : ∀ p ∈ params, (verify P0 p).allInstrs (fun i => !loadsMxcsr i) = true := by +theorem verify0_sp : ∀ p ∈ params, (verify P0 p).allInstrs (fun i => !isa.writesSp i) = true := by decide +kernel -theorem verify0_sp : ∀ p ∈ params, (verify P0 p).allInstrs (fun i => !isa.writesSp i) = true := by +/-! ## MXCSR -/ + +/-- Every primitive of `P` loads MXCSR only to restore it. -/ +structure PrimsC (P : Prims) : Prop where + ntt : ctlOk P.ntt = true + invNtt : ctlOk P.invNtt = true + mul : ctlOk P.mul = true + mulAdd : ctlOk P.mulAdd = true + sub : ctlOk P.sub = true + rejNtt : ctlOk P.rejNtt = true + ball : ctlOk P.ball = true + useHint : ctlOk P.useHint = true + simpleBitPack : ctlOk P.simpleBitPack = true + bitUnpack : ctlOk P.bitUnpack = true + unpackT1 : ctlOk P.unpackT1 = true + hintUnpack : ctlOk P.hintUnpack = true + normLt : ctlOk P.normLt = true + +/-- `ctlC` holds of `c` exactly when it does of `c'`. -/ +def SameC (c c' : Prog isa) : Prop := ctlC c = ctlC c' + +theorem SameC.seq {a a' b b' : Prog isa} (ha : SameC a a') (hb : SameC b b') : SameC (.seq a b) (.seq a' b') := by + show (ctlC a && ctlC b) = (ctlC a' && ctlC b') + rw [show ctlC a = ctlC a' from ha, show ctlC b = ctlC b' from hb] + +theorem SameC.call {c : Prog isa} (hc : ctlOk c = true) (n : String) (as : List (Reg × Arg)) : + SameC (callAt n c as) (callAt n (.block []) as) := by + show (_ && ctlOk c) = (_ && true) + rw [hc] + +theorem SameC.seqR {f g : Nat → Prog isa} (h : ∀ k, SameC (f k) (g k)) : ∀ a n, SameC (seqR f a n) (seqR g a n) + | _, 0 => rfl + | a, n + 1 => (h a).seq (SameC.seqR h (a + 1) n) + +theorem SameC.ifOk {c c' : Prog isa} (h : SameC c c') : SameC (ifOk c) (ifOk c') := by + refine SameC.seq rfl ?_ + show (ctlC c && _) = (ctlC c' && _) + rw [show ctlC c = ctlC c' from h] + +theorem SameC.sampled {c c' : Prog isa} (h : SameC c c') (a : Ptr) : SameC (sampled c a) (sampled c' a) := + h.seq rfl + +section +variable {P : Prims} (hP : PrimsC P) (p : Params) +include hP + +theorem verify_c : SameC (verify P p) (verify P0 p) := by + have aOne : ∀ e, SameC (aOne P e) (aOne P0 e) := fun e => + SameC.seq rfl (SameC.sampled (SameC.call hP.rejNtt _ _) _) + have dot : ∀ r, SameC (dot P p r) (dot P0 p r) := fun r => + (SameC.call hP.mul _ _).seq (SameC.seqR (fun _ => SameC.call hP.mulAdd _ _) _ _) + have row : ∀ r, SameC (row P p r) (row P0 p r) := fun r => + (dot r).seq ((SameC.call hP.unpackT1 _ _).seq ((SameC.call hP.ntt _ _).seq ((SameC.call hP.mul _ _).seq + ((SameC.call hP.sub _ _).seq ((SameC.call hP.invNtt _ _).seq ((SameC.call hP.useHint _ _).seq + (SameC.call hP.simpleBitPack _ _))))))) + have samples : SameC (samples P p) (samples P0 p) := + SameC.seq rfl ((SameC.seqR (fun r => SameC.seqR aOne _ _) _ _).seq + (SameC.sampled (SameC.call hP.ball _ _) _)) + have compute : SameC (compute P p) (compute P0 p) := + (SameC.seqR (fun _ => SameC.call hP.ntt _ _) _ _).seq ((SameC.call hP.ntt _ _).seq + ((SameC.seqR row _ _).seq rfl)) + have zOne : ∀ i, SameC (zOne P p i) (zOne P0 p i) := fun _ => + (SameC.call hP.bitUnpack _ _).seq ((SameC.call hP.normLt _ _).seq rfl) + exact SameC.seq rfl ((((SameC.call hP.hintUnpack _ _).seq rfl).seq (SameC.ifOk ((SameC.seqR zOne _ _).seq + (SameC.ifOk (samples.seq compute))))).seq rfl) + +end + +theorem verify0_ctlC : ∀ p ∈ params, ctlC (verify P0 p) = true := by decide +kernel variable {P : Prims} (C : PrimsOk P) {p : Params} (hp : p ∈ params) include C hp -theorem verify_mxcsr : (verify P p).allInstrs (fun i => !loadsMxcsr i) = true := - (verify_q ⟨C.ntt.mxcsr, C.invNtt.mxcsr, C.mul.mxcsr, C.mulAdd.mxcsr, C.sub.mxcsr, C.rejNtt.mxcsr, C.ball.mxcsr, - C.useHint.mxcsr, C.simpleBitPack.mxcsr, C.bitUnpack.mxcsr, C.unpackT1.mxcsr, C.hintUnpack.mxcsr, - C.normLt.mxcsr⟩ p).trans (verify0_mxcsr p hp) +theorem verify_ctl : ctlOk (verify P p) = true := + ctlOk_of_ctlC ((verify_c ⟨C.ntt.ctl, C.invNtt.ctl, C.mul.ctl, C.mulAdd.ctl, C.sub.ctl, C.rejNtt.ctl, C.ball.ctl, + C.useHint.ctl, C.simpleBitPack.ctl, C.bitUnpack.ctl, C.unpackT1.ctl, C.hintUnpack.ctl, + C.normLt.ctl⟩ p).trans (verify0_ctlC p hp)) theorem verify_spSafe : (verify P p).all (fun i => !isa.writesSp i) = true := Code.all_of_allInstrs ((verify_q ⟨Code.allInstrs_of_all C.ntt.spSafe, Code.allInstrs_of_all C.invNtt.spSafe, diff --git a/src/asm/x86_64/mldsa.rs b/src/asm/x86_64/mldsa.rs index 8b5809e4c..59b6ad833 100644 --- a/src/asm/x86_64/mldsa.rs +++ b/src/asm/x86_64/mldsa.rs @@ -6,7 +6,7 @@ /// /// Contract: `VG.Spec.MlDsa.nttContract`. Constant time: only the pointers may affect timing, not the data. /// -/// The function stores a table of the 256 zetas in `scratch`. +/// The function computes on four coefficients at a time in SSE2 registers, with a table of the 256 zetas that it stores in `scratch`. It sets MXCSR to `0x1FBF` around its multiplications (Intel's mitigation of MXCSR-configuration-dependent timing) and loads the caller's MXCSR back before returning. /// /// # Safety /// @@ -19,850 +19,691 @@ #[unsafe(naked)] pub(crate) unsafe extern "sysv64" fn vg_mldsa_ntt(f: *mut [u32; 256], scratch: *mut [u64; 128]) { core::arch::naked_asm!( - "mov r9, rsi", - "mov eax, 1", - "mov DWORD PTR [r9], eax", - "mov eax, 4808194", - "mov DWORD PTR [r9+4], eax", - "mov eax, 3765607", - "mov DWORD PTR [r9+8], eax", - "mov eax, 3761513", - "mov DWORD PTR [r9+12], eax", - "mov eax, 5178923", - "mov DWORD PTR [r9+16], eax", - "mov eax, 5496691", - "mov DWORD PTR [r9+20], eax", - "mov eax, 5234739", - "mov DWORD PTR [r9+24], eax", - "mov eax, 5178987", - "mov DWORD PTR [r9+28], eax", - "mov eax, 7778734", - "mov DWORD PTR [r9+32], eax", - "mov eax, 3542485", - "mov DWORD PTR [r9+36], eax", - "mov eax, 2682288", - "mov DWORD PTR [r9+40], eax", - "mov eax, 2129892", - "mov DWORD PTR [r9+44], eax", - "mov eax, 3764867", - "mov DWORD PTR [r9+48], eax", - "mov eax, 7375178", - "mov DWORD PTR [r9+52], eax", - "mov eax, 557458", - "mov DWORD PTR [r9+56], eax", - "mov eax, 7159240", - "mov DWORD PTR [r9+60], eax", - "mov eax, 5010068", - "mov DWORD PTR [r9+64], eax", - "mov eax, 4317364", - "mov DWORD PTR [r9+68], eax", - "mov eax, 2663378", - "mov DWORD PTR [r9+72], eax", - "mov eax, 6705802", - "mov DWORD PTR [r9+76], eax", - "mov eax, 4855975", - "mov DWORD PTR [r9+80], eax", - "mov eax, 7946292", - "mov DWORD PTR [r9+84], eax", - "mov eax, 676590", - "mov DWORD PTR [r9+88], eax", - "mov eax, 7044481", - "mov DWORD PTR [r9+92], eax", - "mov eax, 5152541", - "mov DWORD PTR [r9+96], eax", - "mov eax, 1714295", - "mov DWORD PTR [r9+100], eax", - "mov eax, 2453983", - "mov DWORD PTR [r9+104], eax", - "mov eax, 1460718", - "mov DWORD PTR [r9+108], eax", - "mov eax, 7737789", - "mov DWORD PTR [r9+112], eax", - "mov eax, 4795319", - "mov DWORD PTR [r9+116], eax", - "mov eax, 2815639", - "mov DWORD PTR [r9+120], eax", - "mov eax, 2283733", - "mov DWORD PTR [r9+124], eax", - "mov eax, 3602218", - "mov DWORD PTR [r9+128], eax", - "mov eax, 3182878", - "mov DWORD PTR [r9+132], eax", - "mov eax, 2740543", - "mov DWORD PTR [r9+136], eax", - "mov eax, 4793971", - "mov DWORD PTR [r9+140], eax", - "mov eax, 5269599", - "mov DWORD PTR [r9+144], eax", - "mov eax, 2101410", - "mov DWORD PTR [r9+148], eax", - "mov eax, 3704823", - "mov DWORD PTR [r9+152], eax", - "mov eax, 1159875", - "mov DWORD PTR [r9+156], eax", - "mov eax, 394148", - "mov DWORD PTR [r9+160], eax", - "mov eax, 928749", - "mov DWORD PTR [r9+164], eax", - "mov eax, 1095468", - "mov DWORD PTR [r9+168], eax", - "mov eax, 4874037", - "mov DWORD PTR [r9+172], eax", - "mov eax, 2071829", - "mov DWORD PTR [r9+176], eax", - "mov eax, 4361428", - "mov DWORD PTR [r9+180], eax", - "mov eax, 3241972", - "mov DWORD PTR [r9+184], eax", - "mov eax, 2156050", - "mov DWORD PTR [r9+188], eax", - "mov eax, 3415069", - "mov DWORD PTR [r9+192], eax", - "mov eax, 1759347", - "mov DWORD PTR [r9+196], eax", - "mov eax, 7562881", - "mov DWORD PTR [r9+200], eax", - "mov eax, 4805951", - "mov DWORD PTR [r9+204], eax", - "mov eax, 3756790", - "mov DWORD PTR [r9+208], eax", - "mov eax, 6444618", - "mov DWORD PTR [r9+212], eax", - "mov eax, 6663429", - "mov DWORD PTR [r9+216], eax", - "mov eax, 4430364", - "mov DWORD PTR [r9+220], eax", - "mov eax, 5483103", - "mov DWORD PTR [r9+224], eax", - "mov eax, 3192354", - "mov DWORD PTR [r9+228], eax", - "mov eax, 556856", - "mov DWORD PTR [r9+232], eax", - "mov eax, 3870317", - "mov DWORD PTR [r9+236], eax", - "mov eax, 2917338", - "mov DWORD PTR [r9+240], eax", - "mov eax, 1853806", - "mov DWORD PTR [r9+244], eax", - "mov eax, 3345963", - "mov DWORD PTR [r9+248], eax", - "mov eax, 1858416", - "mov DWORD PTR [r9+252], eax", - "mov eax, 3073009", - "mov DWORD PTR [r9+256], eax", - "mov eax, 1277625", - "mov DWORD PTR [r9+260], eax", - "mov eax, 5744944", - "mov DWORD PTR [r9+264], eax", - "mov eax, 3852015", - "mov DWORD PTR [r9+268], eax", - "mov eax, 4183372", - "mov DWORD PTR [r9+272], eax", - "mov eax, 5157610", - "mov DWORD PTR [r9+276], eax", - "mov eax, 5258977", - "mov DWORD PTR [r9+280], eax", - "mov eax, 8106357", - "mov DWORD PTR [r9+284], eax", - "mov eax, 2508980", - "mov DWORD PTR [r9+288], eax", - "mov eax, 2028118", - "mov DWORD PTR [r9+292], eax", - "mov eax, 1937570", - "mov DWORD PTR [r9+296], eax", - "mov eax, 4564692", - "mov DWORD PTR [r9+300], eax", - "mov eax, 2811291", - "mov DWORD PTR [r9+304], eax", - "mov eax, 5396636", - "mov DWORD PTR [r9+308], eax", - "mov eax, 7270901", - "mov DWORD PTR [r9+312], eax", - "mov eax, 4158088", - "mov DWORD PTR [r9+316], eax", - "mov eax, 1528066", - "mov DWORD PTR [r9+320], eax", - "mov eax, 482649", - "mov DWORD PTR [r9+324], eax", - "mov eax, 1148858", - "mov DWORD PTR [r9+328], eax", - "mov eax, 5418153", - "mov DWORD PTR [r9+332], eax", - "mov eax, 7814814", - "mov DWORD PTR [r9+336], eax", - "mov eax, 169688", - "mov DWORD PTR [r9+340], eax", - "mov eax, 2462444", - "mov DWORD PTR [r9+344], eax", - "mov eax, 5046034", - "mov DWORD PTR [r9+348], eax", - "mov eax, 4213992", - "mov DWORD PTR [r9+352], eax", - "mov eax, 4892034", - "mov DWORD PTR [r9+356], eax", - "mov eax, 1987814", - "mov DWORD PTR [r9+360], eax", - "mov eax, 5183169", - "mov DWORD PTR [r9+364], eax", - "mov eax, 1736313", - "mov DWORD PTR [r9+368], eax", - "mov eax, 235407", - "mov DWORD PTR [r9+372], eax", - "mov eax, 5130263", - "mov DWORD PTR [r9+376], eax", - "mov eax, 3258457", - "mov DWORD PTR [r9+380], eax", - "mov eax, 5801164", - "mov DWORD PTR [r9+384], eax", - "mov eax, 1787943", - "mov DWORD PTR [r9+388], eax", - "mov eax, 5989328", - "mov DWORD PTR [r9+392], eax", - "mov eax, 6125690", - "mov DWORD PTR [r9+396], eax", - "mov eax, 3482206", - "mov DWORD PTR [r9+400], eax", - "mov eax, 4197502", - "mov DWORD PTR [r9+404], eax", - "mov eax, 7080401", - "mov DWORD PTR [r9+408], eax", - "mov eax, 6018354", - "mov DWORD PTR [r9+412], eax", - "mov eax, 7062739", - "mov DWORD PTR [r9+416], eax", - "mov eax, 2461387", - "mov DWORD PTR [r9+420], eax", - "mov eax, 3035980", - "mov DWORD PTR [r9+424], eax", - "mov eax, 621164", - "mov DWORD PTR [r9+428], eax", - "mov eax, 3901472", - "mov DWORD PTR [r9+432], eax", - "mov eax, 7153756", - "mov DWORD PTR [r9+436], eax", - "mov eax, 2925816", - "mov DWORD PTR [r9+440], eax", - "mov eax, 3374250", - "mov DWORD PTR [r9+444], eax", - "mov eax, 1356448", - "mov DWORD PTR [r9+448], eax", - "mov eax, 5604662", - "mov DWORD PTR [r9+452], eax", - "mov eax, 2683270", - "mov DWORD PTR [r9+456], eax", - "mov eax, 5601629", - "mov DWORD PTR [r9+460], eax", - "mov eax, 4912752", - "mov DWORD PTR [r9+464], eax", - "mov eax, 2312838", - "mov DWORD PTR [r9+468], eax", - "mov eax, 7727142", - "mov DWORD PTR [r9+472], eax", - "mov eax, 7921254", - "mov DWORD PTR [r9+476], eax", - "mov eax, 348812", - "mov DWORD PTR [r9+480], eax", - "mov eax, 8052569", - "mov DWORD PTR [r9+484], eax", - "mov eax, 1011223", - "mov DWORD PTR [r9+488], eax", - "mov eax, 6026202", - "mov DWORD PTR [r9+492], eax", - "mov eax, 4561790", - "mov DWORD PTR [r9+496], eax", - "mov eax, 6458164", - "mov DWORD PTR [r9+500], eax", - "mov eax, 6143691", - "mov DWORD PTR [r9+504], eax", - "mov eax, 1744507", - "mov DWORD PTR [r9+508], eax", - "mov eax, 1753", - "mov DWORD PTR [r9+512], eax", - "mov eax, 6444997", - "mov DWORD PTR [r9+516], eax", - "mov eax, 5720892", - "mov DWORD PTR [r9+520], eax", - "mov eax, 6924527", - "mov DWORD PTR [r9+524], eax", - "mov eax, 2660408", - "mov DWORD PTR [r9+528], eax", - "mov eax, 6600190", - "mov DWORD PTR [r9+532], eax", - "mov eax, 8321269", - "mov DWORD PTR [r9+536], eax", - "mov eax, 2772600", - "mov DWORD PTR [r9+540], eax", - "mov eax, 1182243", - "mov DWORD PTR [r9+544], eax", - "mov eax, 87208", - "mov DWORD PTR [r9+548], eax", - "mov eax, 636927", - "mov DWORD PTR [r9+552], eax", - "mov eax, 4415111", - "mov DWORD PTR [r9+556], eax", - "mov eax, 4423672", - "mov DWORD PTR [r9+560], eax", - "mov eax, 6084020", - "mov DWORD PTR [r9+564], eax", - "mov eax, 5095502", - "mov DWORD PTR [r9+568], eax", - "mov eax, 4663471", - "mov DWORD PTR [r9+572], eax", - "mov eax, 8352605", - "mov DWORD PTR [r9+576], eax", - "mov eax, 822541", - "mov DWORD PTR [r9+580], eax", - "mov eax, 1009365", - "mov DWORD PTR [r9+584], eax", - "mov eax, 5926272", - "mov DWORD PTR [r9+588], eax", - "mov eax, 6400920", - "mov DWORD PTR [r9+592], eax", - "mov eax, 1596822", - "mov DWORD PTR [r9+596], eax", - "mov eax, 4423473", - "mov DWORD PTR [r9+600], eax", - "mov eax, 4620952", - "mov DWORD PTR [r9+604], eax", - "mov eax, 6695264", - "mov DWORD PTR [r9+608], eax", - "mov eax, 4969849", - "mov DWORD PTR [r9+612], eax", - "mov eax, 2678278", - "mov DWORD PTR [r9+616], eax", - "mov eax, 4611469", - "mov DWORD PTR [r9+620], eax", - "mov eax, 4829411", - "mov DWORD PTR [r9+624], eax", - "mov eax, 635956", - "mov DWORD PTR [r9+628], eax", - "mov eax, 8129971", - "mov DWORD PTR [r9+632], eax", - "mov eax, 5925040", - "mov DWORD PTR [r9+636], eax", - "mov eax, 4234153", - "mov DWORD PTR [r9+640], eax", - "mov eax, 6607829", - "mov DWORD PTR [r9+644], eax", - "mov eax, 2192938", - "mov DWORD PTR [r9+648], eax", - "mov eax, 6653329", - "mov DWORD PTR [r9+652], eax", - "mov eax, 2387513", - "mov DWORD PTR [r9+656], eax", - "mov eax, 4768667", - "mov DWORD PTR [r9+660], eax", - "mov eax, 8111961", - "mov DWORD PTR [r9+664], eax", - "mov eax, 5199961", - "mov DWORD PTR [r9+668], eax", - "mov eax, 3747250", - "mov DWORD PTR [r9+672], eax", - "mov eax, 2296099", - "mov DWORD PTR [r9+676], eax", - "mov eax, 1239911", - "mov DWORD PTR [r9+680], eax", - "mov eax, 4541938", - "mov DWORD PTR [r9+684], eax", - "mov eax, 3195676", - "mov DWORD PTR [r9+688], eax", - "mov eax, 2642980", - "mov DWORD PTR [r9+692], eax", - "mov eax, 1254190", - "mov DWORD PTR [r9+696], eax", - "mov eax, 8368000", - "mov DWORD PTR [r9+700], eax", - "mov eax, 2998219", - "mov DWORD PTR [r9+704], eax", - "mov eax, 141835", - "mov DWORD PTR [r9+708], eax", - "mov eax, 8291116", - "mov DWORD PTR [r9+712], eax", - "mov eax, 2513018", - "mov DWORD PTR [r9+716], eax", - "mov eax, 7025525", - "mov DWORD PTR [r9+720], eax", - "mov eax, 613238", - "mov DWORD PTR [r9+724], eax", - "mov eax, 7070156", - "mov DWORD PTR [r9+728], eax", - "mov eax, 6161950", - "mov DWORD PTR [r9+732], eax", - "mov eax, 7921677", - "mov DWORD PTR [r9+736], eax", - "mov eax, 6458423", - "mov DWORD PTR [r9+740], eax", - "mov eax, 4040196", - "mov DWORD PTR [r9+744], eax", - "mov eax, 4908348", - "mov DWORD PTR [r9+748], eax", - "mov eax, 2039144", - "mov DWORD PTR [r9+752], eax", - "mov eax, 6500539", - "mov DWORD PTR [r9+756], eax", - "mov eax, 7561656", - "mov DWORD PTR [r9+760], eax", - "mov eax, 6201452", - "mov DWORD PTR [r9+764], eax", - "mov eax, 6757063", - "mov DWORD PTR [r9+768], eax", - "mov eax, 2105286", - "mov DWORD PTR [r9+772], eax", - "mov eax, 6006015", - "mov DWORD PTR [r9+776], eax", - "mov eax, 6346610", - "mov DWORD PTR [r9+780], eax", - "mov eax, 586241", - "mov DWORD PTR [r9+784], eax", - "mov eax, 7200804", - "mov DWORD PTR [r9+788], eax", - "mov eax, 527981", - "mov DWORD PTR [r9+792], eax", - "mov eax, 5637006", - "mov DWORD PTR [r9+796], eax", - "mov eax, 6903432", - "mov DWORD PTR [r9+800], eax", - "mov eax, 1994046", - "mov DWORD PTR [r9+804], eax", - "mov eax, 2491325", - "mov DWORD PTR [r9+808], eax", - "mov eax, 6987258", - "mov DWORD PTR [r9+812], eax", - "mov eax, 507927", - "mov DWORD PTR [r9+816], eax", - "mov eax, 7192532", - "mov DWORD PTR [r9+820], eax", - "mov eax, 7655613", - "mov DWORD PTR [r9+824], eax", - "mov eax, 6545891", - "mov DWORD PTR [r9+828], eax", - "mov eax, 5346675", - "mov DWORD PTR [r9+832], eax", - "mov eax, 8041997", - "mov DWORD PTR [r9+836], eax", - "mov eax, 2647994", - "mov DWORD PTR [r9+840], eax", - "mov eax, 3009748", - "mov DWORD PTR [r9+844], eax", - "mov eax, 5767564", - "mov DWORD PTR [r9+848], eax", - "mov eax, 4148469", - "mov DWORD PTR [r9+852], eax", - "mov eax, 749577", - "mov DWORD PTR [r9+856], eax", - "mov eax, 4357667", - "mov DWORD PTR [r9+860], eax", - "mov eax, 3980599", - "mov DWORD PTR [r9+864], eax", - "mov eax, 2569011", - "mov DWORD PTR [r9+868], eax", - "mov eax, 6764887", - "mov DWORD PTR [r9+872], eax", - "mov eax, 1723229", - "mov DWORD PTR [r9+876], eax", - "mov eax, 1665318", - "mov DWORD PTR [r9+880], eax", - "mov eax, 2028038", - "mov DWORD PTR [r9+884], eax", - "mov eax, 1163598", - "mov DWORD PTR [r9+888], eax", - "mov eax, 5011144", - "mov DWORD PTR [r9+892], eax", - "mov eax, 3994671", - "mov DWORD PTR [r9+896], eax", - "mov eax, 8368538", - "mov DWORD PTR [r9+900], eax", - "mov eax, 7009900", - "mov DWORD PTR [r9+904], eax", - "mov eax, 3020393", - "mov DWORD PTR [r9+908], eax", - "mov eax, 3363542", - "mov DWORD PTR [r9+912], eax", - "mov eax, 214880", - "mov DWORD PTR [r9+916], eax", - "mov eax, 545376", - "mov DWORD PTR [r9+920], eax", - "mov eax, 7609976", - "mov DWORD PTR [r9+924], eax", - "mov eax, 3105558", - "mov DWORD PTR [r9+928], eax", - "mov eax, 7277073", - "mov DWORD PTR [r9+932], eax", - "mov eax, 508145", - "mov DWORD PTR [r9+936], eax", - "mov eax, 7826699", - "mov DWORD PTR [r9+940], eax", - "mov eax, 860144", - "mov DWORD PTR [r9+944], eax", - "mov eax, 3430436", - "mov DWORD PTR [r9+948], eax", - "mov eax, 140244", - "mov DWORD PTR [r9+952], eax", - "mov eax, 6866265", - "mov DWORD PTR [r9+956], eax", - "mov eax, 6195333", - "mov DWORD PTR [r9+960], eax", - "mov eax, 3123762", - "mov DWORD PTR [r9+964], eax", - "mov eax, 2358373", - "mov DWORD PTR [r9+968], eax", - "mov eax, 6187330", - "mov DWORD PTR [r9+972], eax", - "mov eax, 5365997", - "mov DWORD PTR [r9+976], eax", - "mov eax, 6663603", - "mov DWORD PTR [r9+980], eax", - "mov eax, 2926054", - "mov DWORD PTR [r9+984], eax", - "mov eax, 7987710", - "mov DWORD PTR [r9+988], eax", - "mov eax, 8077412", - "mov DWORD PTR [r9+992], eax", - "mov eax, 3531229", - "mov DWORD PTR [r9+996], eax", - "mov eax, 4405932", - "mov DWORD PTR [r9+1000], eax", - "mov eax, 4606686", - "mov DWORD PTR [r9+1004], eax", - "mov eax, 1900052", - "mov DWORD PTR [r9+1008], eax", - "mov eax, 7598542", - "mov DWORD PTR [r9+1012], eax", - "mov eax, 1054478", - "mov DWORD PTR [r9+1016], eax", - "mov eax, 7648983", - "mov DWORD PTR [r9+1020], eax", - "mov rsi, rdi", - "mov r8, r9", + "stmxcsr DWORD PTR [rsi+768]", + "mov r11d, DWORD PTR [rsi+768]", + "and r11d, 65535", + "mov eax, 8127", + "mov DWORD PTR [rsi+772], eax", + "ldmxcsr DWORD PTR [rsi+772]", + "lfence", + "movabs r9, 111012023893504", + "mov QWORD PTR [rsi], r9", + "movabs r9, 33764919763013891", + "mov QWORD PTR [rsi+8], r9", + "movabs r9, 32652304184483396", + "mov QWORD PTR [rsi+16], r9", + "movabs r9, 2003464812134697", + "mov QWORD PTR [rsi+24], r9", + "movabs r9, 10107995079564843", + "mov QWORD PTR [rsi+32], r9", + "movabs r9, 27008953388524718", + "mov QWORD PTR [rsi+40], r9", + "movabs r9, 23603259066260085", + "mov QWORD PTR [rsi+48], r9", + "movabs r9, 11510954738022985", + "mov QWORD PTR [rsi+56], r9", + "movabs r9, 4398527550166616", + "mov QWORD PTR [rsi+64], r9", + "movabs r9, 15401443493111205", + "mov QWORD PTR [rsi+72], r9", + "movabs r9, 31185040284548497", + "mov QWORD PTR [rsi+80], r9", + "movabs r9, 26937467947448680", + "mov QWORD PTR [rsi+88], r9", + "movabs r9, 19416172761943511", + "mov QWORD PTR [rsi+96], r9", + "movabs r9, 21916122901808376", + "mov QWORD PTR [rsi+104], r9", + "movabs r9, 35910200088776757", + "mov QWORD PTR [rsi+112], r9", + "movabs r9, 1202612321726977", + "mov QWORD PTR [rsi+120], r9", + "movabs r9, 411354790447719", + "mov QWORD PTR [rsi+128], r9", + "movabs r9, 15163111458665677", + "mov QWORD PTR [rsi+136], r9", + "movabs r9, 20565458766169348", + "mov QWORD PTR [rsi+144], r9", + "movabs r9, 16816682460325845", + "mov QWORD PTR [rsi+152], r9", + "movabs r9, 22920956268049798", + "mov QWORD PTR [rsi+160], r9", + "movabs r9, 23677166863944342", + "mov QWORD PTR [rsi+168], r9", + "movabs r9, 34703121006855168", + "mov QWORD PTR [rsi+176], r9", + "movabs r9, 33677345376425628", + "mov QWORD PTR [rsi+184], r9", + "movabs r9, 28933472397947454", + "mov QWORD PTR [rsi+192], r9", + "movabs r9, 19579390106369566", + "mov QWORD PTR [rsi+200], r9", + "movabs r9, 26799599498134591", + "mov QWORD PTR [rsi+208], r9", + "movabs r9, 15889643835192413", + "mov QWORD PTR [rsi+216], r9", + "movabs r9, 2282148053410728", + "mov QWORD PTR [rsi+224], r9", + "movabs r9, 16668952760323958", + "mov QWORD PTR [rsi+232], r9", + "movabs r9, 25011900965946676", + "mov QWORD PTR [rsi+240], r9", + "movabs r9, 23977247637478740", + "mov QWORD PTR [rsi+248], r9", + "movabs r9, 29427887555995366", + "mov QWORD PTR [rsi+256], r9", + "movabs r9, 22931526182748624", + "mov QWORD PTR [rsi+264], r9", + "movabs r9, 14929091579459166", + "mov QWORD PTR [rsi+272], r9", + "movabs r9, 29185144592086471", + "mov QWORD PTR [rsi+280], r9", + "movabs r9, 8329290213797750", + "mov QWORD PTR [rsi+288], r9", + "movabs r9, 31697782066745459", + "mov QWORD PTR [rsi+296], r9", + "movabs r9, 22432987854353025", + "mov QWORD PTR [rsi+304], r9", + "movabs r9, 545125843890401", + "mov QWORD PTR [rsi+312], r9", + "movabs r9, 31769864501989618", + "mov QWORD PTR [rsi+320], r9", + "movabs r9, 11662103226140216", + "mov QWORD PTR [rsi+328], r9", + "movabs r9, 20130185304250276", + "mov QWORD PTR [rsi+336], r9", + "movabs r9, 25354781094156910", + "mov QWORD PTR [rsi+344], r9", + "movabs r9, 30717142252233347", + "mov QWORD PTR [rsi+352], r9", + "movabs r9, 30375073877558844", + "mov QWORD PTR [rsi+360], r9", + "movabs r9, 5794237307816926", + "mov QWORD PTR [rsi+368], r9", + "movabs r9, 29849966874477923", + "mov QWORD PTR [rsi+376], r9", + "movabs r9, 1137925820308458", + "mov QWORD PTR [rsi+384], r9", + "movabs r9, 13305774323778583", + "mov QWORD PTR [rsi+392], r9", + "movabs r9, 31268732009491712", + "mov QWORD PTR [rsi+400], r9", + "movabs r9, 17002134848261444", + "mov QWORD PTR [rsi+408], r9", + "movabs r9, 35956774717033419", + "mov QWORD PTR [rsi+416], r9", + "movabs r9, 22036141462600008", + "mov QWORD PTR [rsi+424], r9", + "movabs r9, 35087477629023596", + "mov QWORD PTR [rsi+432], r9", + "movabs r9, 30337763489061025", + "mov QWORD PTR [rsi+440], r9", + "movabs r9, 20732429908239468", + "mov QWORD PTR [rsi+448], r9", + "movabs r9, 28041905903253186", + "mov QWORD PTR [rsi+456], r9", + "movabs r9, 35231517950811284", + "mov QWORD PTR [rsi+464], r9", + "movabs r9, 5760968484459269", + "mov QWORD PTR [rsi+472], r9", + "movabs r9, 29186403016613413", + "mov QWORD PTR [rsi+480], r9", + "movabs r9, 29809972144732485", + "mov QWORD PTR [rsi+488], r9", + "movabs r9, 19324591173389987", + "mov QWORD PTR [rsi+496], r9", + "movabs r9, 16492506917666904", + "mov QWORD PTR [rsi+504], r9", + "movabs r9, 14635985826474643", + "mov QWORD PTR [rsi+512], r9", + "movabs r9, 16398082059229396", + "mov QWORD PTR [rsi+520], r9", + "movabs r9, 9638297459285875", + "mov QWORD PTR [rsi+528], r9", + "movabs r9, 20692959164533664", + "mov QWORD PTR [rsi+536], r9", + "movabs r9, 10455835889373941", + "mov QWORD PTR [rsi+544], r9", + "movabs r9, 15088997507073265", + "mov QWORD PTR [rsi+552], r9", + "movabs r9, 19847271512942753", + "mov QWORD PTR [rsi+560], r9", + "movabs r9, 22278162875259735", + "mov QWORD PTR [rsi+568], r9", + "movabs r9, 7984765110959710", + "mov QWORD PTR [rsi+576], r9", + "movabs r9, 3517724245221606", + "mov QWORD PTR [rsi+584], r9", + "movabs r9, 29065087369580419", + "mov QWORD PTR [rsi+592], r9", + "movabs r9, 33749496538019589", + "mov QWORD PTR [rsi+600], r9", + "movabs r9, 22582830675910690", + "mov QWORD PTR [rsi+608], r9", + "movabs r9, 13774157688799364", + "mov QWORD PTR [rsi+616], r9", + "movabs r9, 33738338209470846", + "mov QWORD PTR [rsi+624], r9", + "movabs r9, 20549610337740179", + "mov QWORD PTR [rsi+632], r9", + "movabs r9, 1232604074686745", + "mov QWORD PTR [rsi+640], r9", + "movabs r9, 17645078572608834", + "mov QWORD PTR [rsi+648], r9", + "movabs r9, 21638646536106727", + "mov QWORD PTR [rsi+656], r9", + "movabs r9, 872067341384903", + "mov QWORD PTR [rsi+664], r9", + "movabs r9, 11559822875647717", + "mov QWORD PTR [rsi+672], r9", + "movabs r9, 5433172289935931", + "mov QWORD PTR [rsi+680], r9", + "movabs r9, 5358487101890844", + "mov QWORD PTR [rsi+688], r9", + "movabs r9, 6854656137752657", + "mov QWORD PTR [rsi+696], r9", + "movabs r9, 5370830838457625", + "mov QWORD PTR [rsi+704], r9", + "movabs r9, 20753904747165841", + "mov QWORD PTR [rsi+712], r9", + "movabs r9, 8027804982718602", + "mov QWORD PTR [rsi+720], r9", + "movabs r9, 31479735164668747", + "mov QWORD PTR [rsi+728], r9", + "movabs r9, 5314055668205759", + "mov QWORD PTR [rsi+736], r9", + "movabs r9, 29850847345983039", + "mov QWORD PTR [rsi+744], r9", + "movabs r9, 5636951310400997", + "mov QWORD PTR [rsi+752], r9", + "movabs r9, 27564133741392515", + "mov QWORD PTR [rsi+760], r9", + "movabs r9, 8233800205883732", + "mov QWORD PTR [rsi+768], r9", + "movabs r9, 30088883024233849", + "mov QWORD PTR [rsi+776], r9", + "movabs r9, 3338009929245701", + "mov QWORD PTR [rsi+784], r9", + "movabs r9, 14628791756398056", + "mov QWORD PTR [rsi+792], r9", + "movabs r9, 23830870862829877", + "mov QWORD PTR [rsi+800], r9", + "movabs r9, 28061014216302585", + "mov QWORD PTR [rsi+808], r9", + "movabs r9, 19997999096164636", + "mov QWORD PTR [rsi+816], r9", + "movabs r9, 19771555530215640", + "mov QWORD PTR [rsi+824], r9", + "movabs r9, 10447056982320729", + "mov QWORD PTR [rsi+832], r9", + "movabs r9, 35286145636332471", + "mov QWORD PTR [rsi+840], r9", + "movabs r9, 14470225858518424", + "mov QWORD PTR [rsi+848], r9", + "movabs r9, 30807937853347003", + "mov QWORD PTR [rsi+856], r9", + "movabs r9, 699409659546815", + "mov QWORD PTR [rsi+864], r9", + "movabs r9, 12945035726727688", + "mov QWORD PTR [rsi+872], r9", + "movabs r9, 7098008983067813", + "mov QWORD PTR [rsi+880], r9", + "movabs r9, 28266511219523944", + "mov QWORD PTR [rsi+888], r9", + "movabs r9, 15135022374814013", + "mov QWORD PTR [rsi+896], r9", + "movabs r9, 1158610381635861", + "mov QWORD PTR [rsi+904], r9", + "movabs r9, 31802227079365879", + "mov QWORD PTR [rsi+912], r9", + "movabs r9, 2027559572878823", + "mov QWORD PTR [rsi+920], r9", + "movabs r9, 7402805639339334", + "mov QWORD PTR [rsi+928], r9", + "movabs r9, 8205002449640623", + "mov QWORD PTR [rsi+936], r9", + "movabs r9, 31250542829661849", + "mov QWORD PTR [rsi+944], r9", + "movabs r9, 19527171898598875", + "mov QWORD PTR [rsi+952], r9", + "movabs r9, 26390134497937253", + "mov QWORD PTR [rsi+960], r9", + "movabs r9, 26173917256840158", + "mov QWORD PTR [rsi+968], r9", + "movabs r9, 31797902045269139", + "mov QWORD PTR [rsi+976], r9", + "movabs r9, 20765007236602922", + "mov QWORD PTR [rsi+984], r9", + "movabs r9, 16834811519265361", + "mov QWORD PTR [rsi+992], r9", + "movabs r9, 30142973844857679", + "mov QWORD PTR [rsi+1000], r9", + "movabs r9, 6014775284471242", + "mov QWORD PTR [rsi+1008], r9", + "movabs r9, 8490214048855735", + "mov QWORD PTR [rsi+1016], r9", + "mov eax, 8380417", + "movq xmm15, rax", + "pshufd xmm15, xmm15, 0", + "mov eax, -58728449", + "movq xmm14, rax", + "pshufd xmm14, xmm14, 0", + "mov rdx, rdi", + "mov r8, rsi", "add r8, 4", - "mov edi, 1", + "mov eax, 1", "20:", - "mov r9d, DWORD PTR [r8]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 0", + "pshufd xmm12, xmm13, 245", "add r8, 4", - "mov ecx, 128", + "mov ecx, 32", "21:", - "mov eax, DWORD PTR [rsi+512]", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov eax, DWORD PTR [rsi]", - "mov edx, eax", - "add edx, 8380417", - "sub edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi+512], edx", - "add eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mov DWORD PTR [rsi], eax", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm1, XMMWORD PTR [rdx+512]", + "pshufd xmm4, xmm1, 245", + "pmuludq xmm1, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm1", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm1, xmm2", + "psrlq xmm1, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm1, xmm4", + "psubd xmm1, xmm15", + "movdqa xmm2, xmm1", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm1, xmm2", + "movdqa xmm3, xmm0", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "psubd xmm3, xmm1", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+512], xmm3", + "add rdx, 16", "sub rcx, 1", "jne 21b", - "add rsi, 512", - "sub rdi, 1", + "add rdx, 512", + "sub rax, 1", "jne 20b", - "sub rsi, 1024", - "mov edi, 2", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 8", + "mov eax, 2", "22:", - "mov r9d, DWORD PTR [r8]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 0", + "pshufd xmm12, xmm13, 245", "add r8, 4", - "mov ecx, 64", + "mov ecx, 16", "23:", - "mov eax, DWORD PTR [rsi+256]", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov eax, DWORD PTR [rsi]", - "mov edx, eax", - "add edx, 8380417", - "sub edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi+256], edx", - "add eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mov DWORD PTR [rsi], eax", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm1, XMMWORD PTR [rdx+256]", + "pshufd xmm4, xmm1, 245", + "pmuludq xmm1, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm1", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm1, xmm2", + "psrlq xmm1, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm1, xmm4", + "psubd xmm1, xmm15", + "movdqa xmm2, xmm1", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm1, xmm2", + "movdqa xmm3, xmm0", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "psubd xmm3, xmm1", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+256], xmm3", + "add rdx, 16", "sub rcx, 1", "jne 23b", - "add rsi, 256", - "sub rdi, 1", + "add rdx, 256", + "sub rax, 1", "jne 22b", - "sub rsi, 1024", - "mov edi, 4", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 16", + "mov eax, 4", "24:", - "mov r9d, DWORD PTR [r8]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 0", + "pshufd xmm12, xmm13, 245", "add r8, 4", - "mov ecx, 32", + "mov ecx, 8", "25:", - "mov eax, DWORD PTR [rsi+128]", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov eax, DWORD PTR [rsi]", - "mov edx, eax", - "add edx, 8380417", - "sub edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi+128], edx", - "add eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mov DWORD PTR [rsi], eax", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm1, XMMWORD PTR [rdx+128]", + "pshufd xmm4, xmm1, 245", + "pmuludq xmm1, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm1", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm1, xmm2", + "psrlq xmm1, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm1, xmm4", + "psubd xmm1, xmm15", + "movdqa xmm2, xmm1", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm1, xmm2", + "movdqa xmm3, xmm0", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "psubd xmm3, xmm1", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+128], xmm3", + "add rdx, 16", "sub rcx, 1", "jne 25b", - "add rsi, 128", - "sub rdi, 1", + "add rdx, 128", + "sub rax, 1", "jne 24b", - "sub rsi, 1024", - "mov edi, 8", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 32", + "mov eax, 8", "26:", - "mov r9d, DWORD PTR [r8]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 0", + "pshufd xmm12, xmm13, 245", "add r8, 4", - "mov ecx, 16", + "mov ecx, 4", "27:", - "mov eax, DWORD PTR [rsi+64]", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov eax, DWORD PTR [rsi]", - "mov edx, eax", - "add edx, 8380417", - "sub edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi+64], edx", - "add eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mov DWORD PTR [rsi], eax", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm1, XMMWORD PTR [rdx+64]", + "pshufd xmm4, xmm1, 245", + "pmuludq xmm1, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm1", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm1, xmm2", + "psrlq xmm1, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm1, xmm4", + "psubd xmm1, xmm15", + "movdqa xmm2, xmm1", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm1, xmm2", + "movdqa xmm3, xmm0", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "psubd xmm3, xmm1", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+64], xmm3", + "add rdx, 16", "sub rcx, 1", "jne 27b", - "add rsi, 64", - "sub rdi, 1", + "add rdx, 64", + "sub rax, 1", "jne 26b", - "sub rsi, 1024", - "mov edi, 16", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 64", + "mov eax, 16", "28:", - "mov r9d, DWORD PTR [r8]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 0", + "pshufd xmm12, xmm13, 245", "add r8, 4", - "mov ecx, 8", + "mov ecx, 2", "29:", - "mov eax, DWORD PTR [rsi+32]", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov eax, DWORD PTR [rsi]", - "mov edx, eax", - "add edx, 8380417", - "sub edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi+32], edx", - "add eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mov DWORD PTR [rsi], eax", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm1, XMMWORD PTR [rdx+32]", + "pshufd xmm4, xmm1, 245", + "pmuludq xmm1, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm1", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm1, xmm2", + "psrlq xmm1, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm1, xmm4", + "psubd xmm1, xmm15", + "movdqa xmm2, xmm1", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm1, xmm2", + "movdqa xmm3, xmm0", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "psubd xmm3, xmm1", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+32], xmm3", + "add rdx, 16", "sub rcx, 1", "jne 29b", - "add rsi, 32", - "sub rdi, 1", + "add rdx, 32", + "sub rax, 1", "jne 28b", - "sub rsi, 1024", - "mov edi, 32", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 128", + "mov eax, 32", "210:", - "mov r9d, DWORD PTR [r8]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 0", + "pshufd xmm12, xmm13, 245", "add r8, 4", - "mov ecx, 4", + "mov ecx, 1", "211:", - "mov eax, DWORD PTR [rsi+16]", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov eax, DWORD PTR [rsi]", - "mov edx, eax", - "add edx, 8380417", - "sub edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi+16], edx", - "add eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mov DWORD PTR [rsi], eax", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm1, XMMWORD PTR [rdx+16]", + "pshufd xmm4, xmm1, 245", + "pmuludq xmm1, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm1", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm1, xmm2", + "psrlq xmm1, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm1, xmm4", + "psubd xmm1, xmm15", + "movdqa xmm2, xmm1", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm1, xmm2", + "movdqa xmm3, xmm0", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "psubd xmm3, xmm1", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+16], xmm3", + "add rdx, 16", "sub rcx, 1", "jne 211b", - "add rsi, 16", - "sub rdi, 1", + "add rdx, 16", + "sub rax, 1", "jne 210b", - "sub rsi, 1024", - "mov edi, 64", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 256", + "mov ecx, 32", "212:", - "mov r9d, DWORD PTR [r8]", - "add r8, 4", - "mov ecx, 2", - "213:", - "mov eax, DWORD PTR [rsi+8]", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov eax, DWORD PTR [rsi]", - "mov edx, eax", - "add edx, 8380417", - "sub edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi+8], edx", - "add eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mov DWORD PTR [rsi], eax", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm1, XMMWORD PTR [rdx+16]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 80", + "pshufd xmm12, xmm13, 245", + "add r8, 8", + "movdqa xmm2, xmm0", + "punpcklqdq xmm0, xmm1", + "punpckhqdq xmm2, xmm1", + "movdqa xmm1, xmm2", + "pshufd xmm4, xmm1, 245", + "pmuludq xmm1, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm1", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm1, xmm2", + "psrlq xmm1, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm1, xmm4", + "psubd xmm1, xmm15", + "movdqa xmm2, xmm1", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm1, xmm2", + "movdqa xmm3, xmm0", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "psubd xmm3, xmm1", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqa xmm1, xmm0", + "punpcklqdq xmm0, xmm3", + "punpckhqdq xmm1, xmm3", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+16], xmm1", + "add rdx, 32", "sub rcx, 1", - "jne 213b", - "add rsi, 8", - "sub rdi, 1", "jne 212b", - "sub rsi, 1024", - "mov edi, 128", - "214:", - "mov r9d, DWORD PTR [r8]", - "add r8, 4", - "mov ecx, 1", - "215:", - "mov eax, DWORD PTR [rsi+4]", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov eax, DWORD PTR [rsi]", - "mov edx, eax", - "add edx, 8380417", - "sub edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi+4], edx", - "add eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mov DWORD PTR [rsi], eax", - "add rsi, 4", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 512", + "mov ecx, 32", + "213:", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm2, XMMWORD PTR [rdx+16]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 228", + "pshufd xmm12, xmm13, 245", + "add r8, 16", + "pshufd xmm0, xmm0, 216", + "pshufd xmm2, xmm2, 216", + "movdqa xmm1, xmm0", + "punpcklqdq xmm0, xmm2", + "punpckhqdq xmm1, xmm2", + "pshufd xmm4, xmm1, 245", + "pmuludq xmm1, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm1", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm1, xmm2", + "psrlq xmm1, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm1, xmm4", + "psubd xmm1, xmm15", + "movdqa xmm2, xmm1", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm1, xmm2", + "movdqa xmm3, xmm0", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "psubd xmm3, xmm1", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqa xmm1, xmm0", + "punpckldq xmm0, xmm3", + "punpckhdq xmm1, xmm3", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+16], xmm1", + "add rdx, 32", "sub rcx, 1", - "jne 215b", - "add rsi, 4", - "sub rdi, 1", - "jne 214b", - "sub rsi, 1024", + "jne 213b", + "lfence", + "mov DWORD PTR [rsi+768], r11d", + "ldmxcsr DWORD PTR [rsi+768]", "ret", ) } @@ -871,7 +712,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_ntt(f: *mut [u32; 256], scratch: * /// /// Contract: `VG.Spec.MlDsa.nttInvContract`. Constant time: only the pointers may affect timing, not the data. /// -/// The function stores a table of the 256 negated zetas in `scratch`. +/// The function computes on four coefficients at a time in SSE2 registers, with a table of the 256 zetas that it stores in `scratch`. It sets MXCSR to `0x1FBF` around its multiplications (Intel's mitigation of MXCSR-configuration-dependent timing) and loads the caller's MXCSR back before returning. /// /// # Safety /// @@ -884,870 +725,697 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_ntt(f: *mut [u32; 256], scratch: * #[unsafe(naked)] pub(crate) unsafe extern "sysv64" fn vg_mldsa_inv_ntt(f: *mut [u32; 256], scratch: *mut [u64; 128]) { core::arch::naked_asm!( - "mov r9, rsi", - "mov eax, 8380416", - "mov DWORD PTR [r9], eax", - "mov eax, 3572223", - "mov DWORD PTR [r9+4], eax", - "mov eax, 4614810", - "mov DWORD PTR [r9+8], eax", - "mov eax, 4618904", - "mov DWORD PTR [r9+12], eax", - "mov eax, 3201494", - "mov DWORD PTR [r9+16], eax", - "mov eax, 2883726", - "mov DWORD PTR [r9+20], eax", - "mov eax, 3145678", - "mov DWORD PTR [r9+24], eax", - "mov eax, 3201430", - "mov DWORD PTR [r9+28], eax", - "mov eax, 601683", - "mov DWORD PTR [r9+32], eax", - "mov eax, 4837932", - "mov DWORD PTR [r9+36], eax", - "mov eax, 5698129", - "mov DWORD PTR [r9+40], eax", - "mov eax, 6250525", - "mov DWORD PTR [r9+44], eax", - "mov eax, 4615550", - "mov DWORD PTR [r9+48], eax", - "mov eax, 1005239", - "mov DWORD PTR [r9+52], eax", - "mov eax, 7822959", - "mov DWORD PTR [r9+56], eax", - "mov eax, 1221177", - "mov DWORD PTR [r9+60], eax", - "mov eax, 3370349", - "mov DWORD PTR [r9+64], eax", - "mov eax, 4063053", - "mov DWORD PTR [r9+68], eax", - "mov eax, 5717039", - "mov DWORD PTR [r9+72], eax", - "mov eax, 1674615", - "mov DWORD PTR [r9+76], eax", - "mov eax, 3524442", - "mov DWORD PTR [r9+80], eax", - "mov eax, 434125", - "mov DWORD PTR [r9+84], eax", - "mov eax, 7703827", - "mov DWORD PTR [r9+88], eax", - "mov eax, 1335936", - "mov DWORD PTR [r9+92], eax", - "mov eax, 3227876", - "mov DWORD PTR [r9+96], eax", - "mov eax, 6666122", - "mov DWORD PTR [r9+100], eax", - "mov eax, 5926434", - "mov DWORD PTR [r9+104], eax", - "mov eax, 6919699", - "mov DWORD PTR [r9+108], eax", - "mov eax, 642628", - "mov DWORD PTR [r9+112], eax", - "mov eax, 3585098", - "mov DWORD PTR [r9+116], eax", - "mov eax, 5564778", - "mov DWORD PTR [r9+120], eax", - "mov eax, 6096684", - "mov DWORD PTR [r9+124], eax", - "mov eax, 4778199", - "mov DWORD PTR [r9+128], eax", - "mov eax, 5197539", - "mov DWORD PTR [r9+132], eax", - "mov eax, 5639874", - "mov DWORD PTR [r9+136], eax", - "mov eax, 3586446", - "mov DWORD PTR [r9+140], eax", - "mov eax, 3110818", - "mov DWORD PTR [r9+144], eax", - "mov eax, 6279007", - "mov DWORD PTR [r9+148], eax", - "mov eax, 4675594", - "mov DWORD PTR [r9+152], eax", - "mov eax, 7220542", - "mov DWORD PTR [r9+156], eax", - "mov eax, 7986269", - "mov DWORD PTR [r9+160], eax", - "mov eax, 7451668", - "mov DWORD PTR [r9+164], eax", - "mov eax, 7284949", - "mov DWORD PTR [r9+168], eax", - "mov eax, 3506380", - "mov DWORD PTR [r9+172], eax", - "mov eax, 6308588", - "mov DWORD PTR [r9+176], eax", - "mov eax, 4018989", - "mov DWORD PTR [r9+180], eax", - "mov eax, 5138445", - "mov DWORD PTR [r9+184], eax", - "mov eax, 6224367", - "mov DWORD PTR [r9+188], eax", - "mov eax, 4965348", - "mov DWORD PTR [r9+192], eax", - "mov eax, 6621070", - "mov DWORD PTR [r9+196], eax", - "mov eax, 817536", - "mov DWORD PTR [r9+200], eax", - "mov eax, 3574466", - "mov DWORD PTR [r9+204], eax", - "mov eax, 4623627", - "mov DWORD PTR [r9+208], eax", - "mov eax, 1935799", - "mov DWORD PTR [r9+212], eax", - "mov eax, 1716988", - "mov DWORD PTR [r9+216], eax", - "mov eax, 3950053", - "mov DWORD PTR [r9+220], eax", - "mov eax, 2897314", - "mov DWORD PTR [r9+224], eax", - "mov eax, 5188063", - "mov DWORD PTR [r9+228], eax", - "mov eax, 7823561", - "mov DWORD PTR [r9+232], eax", - "mov eax, 4510100", - "mov DWORD PTR [r9+236], eax", - "mov eax, 5463079", - "mov DWORD PTR [r9+240], eax", - "mov eax, 6526611", - "mov DWORD PTR [r9+244], eax", - "mov eax, 5034454", - "mov DWORD PTR [r9+248], eax", - "mov eax, 6522001", - "mov DWORD PTR [r9+252], eax", - "mov eax, 5307408", - "mov DWORD PTR [r9+256], eax", - "mov eax, 7102792", - "mov DWORD PTR [r9+260], eax", - "mov eax, 2635473", - "mov DWORD PTR [r9+264], eax", - "mov eax, 4528402", - "mov DWORD PTR [r9+268], eax", - "mov eax, 4197045", - "mov DWORD PTR [r9+272], eax", - "mov eax, 3222807", - "mov DWORD PTR [r9+276], eax", - "mov eax, 3121440", - "mov DWORD PTR [r9+280], eax", - "mov eax, 274060", - "mov DWORD PTR [r9+284], eax", - "mov eax, 5871437", - "mov DWORD PTR [r9+288], eax", - "mov eax, 6352299", - "mov DWORD PTR [r9+292], eax", - "mov eax, 6442847", - "mov DWORD PTR [r9+296], eax", - "mov eax, 3815725", - "mov DWORD PTR [r9+300], eax", - "mov eax, 5569126", - "mov DWORD PTR [r9+304], eax", - "mov eax, 2983781", - "mov DWORD PTR [r9+308], eax", - "mov eax, 1109516", - "mov DWORD PTR [r9+312], eax", - "mov eax, 4222329", - "mov DWORD PTR [r9+316], eax", - "mov eax, 6852351", - "mov DWORD PTR [r9+320], eax", - "mov eax, 7897768", - "mov DWORD PTR [r9+324], eax", - "mov eax, 7231559", - "mov DWORD PTR [r9+328], eax", - "mov eax, 2962264", - "mov DWORD PTR [r9+332], eax", - "mov eax, 565603", - "mov DWORD PTR [r9+336], eax", - "mov eax, 8210729", - "mov DWORD PTR [r9+340], eax", - "mov eax, 5917973", - "mov DWORD PTR [r9+344], eax", - "mov eax, 3334383", - "mov DWORD PTR [r9+348], eax", - "mov eax, 4166425", - "mov DWORD PTR [r9+352], eax", - "mov eax, 3488383", - "mov DWORD PTR [r9+356], eax", - "mov eax, 6392603", - "mov DWORD PTR [r9+360], eax", - "mov eax, 3197248", - "mov DWORD PTR [r9+364], eax", - "mov eax, 6644104", - "mov DWORD PTR [r9+368], eax", - "mov eax, 8145010", - "mov DWORD PTR [r9+372], eax", - "mov eax, 3250154", - "mov DWORD PTR [r9+376], eax", - "mov eax, 5121960", - "mov DWORD PTR [r9+380], eax", - "mov eax, 2579253", - "mov DWORD PTR [r9+384], eax", - "mov eax, 6592474", - "mov DWORD PTR [r9+388], eax", - "mov eax, 2391089", - "mov DWORD PTR [r9+392], eax", - "mov eax, 2254727", - "mov DWORD PTR [r9+396], eax", - "mov eax, 4898211", - "mov DWORD PTR [r9+400], eax", - "mov eax, 4182915", - "mov DWORD PTR [r9+404], eax", - "mov eax, 1300016", - "mov DWORD PTR [r9+408], eax", - "mov eax, 2362063", - "mov DWORD PTR [r9+412], eax", - "mov eax, 1317678", - "mov DWORD PTR [r9+416], eax", - "mov eax, 5919030", - "mov DWORD PTR [r9+420], eax", - "mov eax, 5344437", - "mov DWORD PTR [r9+424], eax", - "mov eax, 7759253", - "mov DWORD PTR [r9+428], eax", - "mov eax, 4478945", - "mov DWORD PTR [r9+432], eax", - "mov eax, 1226661", - "mov DWORD PTR [r9+436], eax", - "mov eax, 5454601", - "mov DWORD PTR [r9+440], eax", - "mov eax, 5006167", - "mov DWORD PTR [r9+444], eax", - "mov eax, 7023969", - "mov DWORD PTR [r9+448], eax", - "mov eax, 2775755", - "mov DWORD PTR [r9+452], eax", - "mov eax, 5697147", - "mov DWORD PTR [r9+456], eax", - "mov eax, 2778788", - "mov DWORD PTR [r9+460], eax", - "mov eax, 3467665", - "mov DWORD PTR [r9+464], eax", - "mov eax, 6067579", - "mov DWORD PTR [r9+468], eax", - "mov eax, 653275", - "mov DWORD PTR [r9+472], eax", - "mov eax, 459163", - "mov DWORD PTR [r9+476], eax", - "mov eax, 8031605", - "mov DWORD PTR [r9+480], eax", - "mov eax, 327848", - "mov DWORD PTR [r9+484], eax", - "mov eax, 7369194", - "mov DWORD PTR [r9+488], eax", - "mov eax, 2354215", - "mov DWORD PTR [r9+492], eax", - "mov eax, 3818627", - "mov DWORD PTR [r9+496], eax", - "mov eax, 1922253", - "mov DWORD PTR [r9+500], eax", - "mov eax, 2236726", - "mov DWORD PTR [r9+504], eax", - "mov eax, 6635910", - "mov DWORD PTR [r9+508], eax", - "mov eax, 8378664", - "mov DWORD PTR [r9+512], eax", - "mov eax, 1935420", - "mov DWORD PTR [r9+516], eax", - "mov eax, 2659525", - "mov DWORD PTR [r9+520], eax", - "mov eax, 1455890", - "mov DWORD PTR [r9+524], eax", - "mov eax, 5720009", - "mov DWORD PTR [r9+528], eax", - "mov eax, 1780227", - "mov DWORD PTR [r9+532], eax", - "mov eax, 59148", - "mov DWORD PTR [r9+536], eax", - "mov eax, 5607817", - "mov DWORD PTR [r9+540], eax", - "mov eax, 7198174", - "mov DWORD PTR [r9+544], eax", - "mov eax, 8293209", - "mov DWORD PTR [r9+548], eax", - "mov eax, 7743490", - "mov DWORD PTR [r9+552], eax", - "mov eax, 3965306", - "mov DWORD PTR [r9+556], eax", - "mov eax, 3956745", - "mov DWORD PTR [r9+560], eax", - "mov eax, 2296397", - "mov DWORD PTR [r9+564], eax", - "mov eax, 3284915", - "mov DWORD PTR [r9+568], eax", - "mov eax, 3716946", - "mov DWORD PTR [r9+572], eax", - "mov eax, 27812", - "mov DWORD PTR [r9+576], eax", - "mov eax, 7557876", - "mov DWORD PTR [r9+580], eax", - "mov eax, 7371052", - "mov DWORD PTR [r9+584], eax", - "mov eax, 2454145", - "mov DWORD PTR [r9+588], eax", - "mov eax, 1979497", - "mov DWORD PTR [r9+592], eax", - "mov eax, 6783595", - "mov DWORD PTR [r9+596], eax", - "mov eax, 3956944", - "mov DWORD PTR [r9+600], eax", - "mov eax, 3759465", - "mov DWORD PTR [r9+604], eax", - "mov eax, 1685153", - "mov DWORD PTR [r9+608], eax", - "mov eax, 3410568", - "mov DWORD PTR [r9+612], eax", - "mov eax, 5702139", - "mov DWORD PTR [r9+616], eax", - "mov eax, 3768948", - "mov DWORD PTR [r9+620], eax", - "mov eax, 3551006", - "mov DWORD PTR [r9+624], eax", - "mov eax, 7744461", - "mov DWORD PTR [r9+628], eax", - "mov eax, 250446", - "mov DWORD PTR [r9+632], eax", - "mov eax, 2455377", - "mov DWORD PTR [r9+636], eax", - "mov eax, 4146264", - "mov DWORD PTR [r9+640], eax", - "mov eax, 1772588", - "mov DWORD PTR [r9+644], eax", - "mov eax, 6187479", - "mov DWORD PTR [r9+648], eax", - "mov eax, 1727088", - "mov DWORD PTR [r9+652], eax", - "mov eax, 5992904", - "mov DWORD PTR [r9+656], eax", - "mov eax, 3611750", - "mov DWORD PTR [r9+660], eax", - "mov eax, 268456", - "mov DWORD PTR [r9+664], eax", - "mov eax, 3180456", - "mov DWORD PTR [r9+668], eax", - "mov eax, 4633167", - "mov DWORD PTR [r9+672], eax", - "mov eax, 6084318", - "mov DWORD PTR [r9+676], eax", - "mov eax, 7140506", - "mov DWORD PTR [r9+680], eax", - "mov eax, 3838479", - "mov DWORD PTR [r9+684], eax", - "mov eax, 5184741", - "mov DWORD PTR [r9+688], eax", - "mov eax, 5737437", - "mov DWORD PTR [r9+692], eax", - "mov eax, 7126227", - "mov DWORD PTR [r9+696], eax", - "mov eax, 12417", - "mov DWORD PTR [r9+700], eax", - "mov eax, 5382198", - "mov DWORD PTR [r9+704], eax", - "mov eax, 8238582", - "mov DWORD PTR [r9+708], eax", - "mov eax, 89301", - "mov DWORD PTR [r9+712], eax", - "mov eax, 5867399", - "mov DWORD PTR [r9+716], eax", - "mov eax, 1354892", - "mov DWORD PTR [r9+720], eax", - "mov eax, 7767179", - "mov DWORD PTR [r9+724], eax", - "mov eax, 1310261", - "mov DWORD PTR [r9+728], eax", - "mov eax, 2218467", - "mov DWORD PTR [r9+732], eax", - "mov eax, 458740", - "mov DWORD PTR [r9+736], eax", - "mov eax, 1921994", - "mov DWORD PTR [r9+740], eax", - "mov eax, 4340221", - "mov DWORD PTR [r9+744], eax", - "mov eax, 3472069", - "mov DWORD PTR [r9+748], eax", - "mov eax, 6341273", - "mov DWORD PTR [r9+752], eax", - "mov eax, 1879878", - "mov DWORD PTR [r9+756], eax", - "mov eax, 818761", - "mov DWORD PTR [r9+760], eax", - "mov eax, 2178965", - "mov DWORD PTR [r9+764], eax", - "mov eax, 1623354", - "mov DWORD PTR [r9+768], eax", - "mov eax, 6275131", - "mov DWORD PTR [r9+772], eax", - "mov eax, 2374402", - "mov DWORD PTR [r9+776], eax", - "mov eax, 2033807", - "mov DWORD PTR [r9+780], eax", - "mov eax, 7794176", - "mov DWORD PTR [r9+784], eax", - "mov eax, 1179613", - "mov DWORD PTR [r9+788], eax", - "mov eax, 7852436", - "mov DWORD PTR [r9+792], eax", - "mov eax, 2743411", - "mov DWORD PTR [r9+796], eax", - "mov eax, 1476985", - "mov DWORD PTR [r9+800], eax", - "mov eax, 6386371", - "mov DWORD PTR [r9+804], eax", - "mov eax, 5889092", - "mov DWORD PTR [r9+808], eax", - "mov eax, 1393159", - "mov DWORD PTR [r9+812], eax", - "mov eax, 7872490", - "mov DWORD PTR [r9+816], eax", - "mov eax, 1187885", - "mov DWORD PTR [r9+820], eax", - "mov eax, 724804", - "mov DWORD PTR [r9+824], eax", - "mov eax, 1834526", - "mov DWORD PTR [r9+828], eax", - "mov eax, 3033742", - "mov DWORD PTR [r9+832], eax", - "mov eax, 338420", - "mov DWORD PTR [r9+836], eax", - "mov eax, 5732423", - "mov DWORD PTR [r9+840], eax", - "mov eax, 5370669", - "mov DWORD PTR [r9+844], eax", - "mov eax, 2612853", - "mov DWORD PTR [r9+848], eax", - "mov eax, 4231948", - "mov DWORD PTR [r9+852], eax", - "mov eax, 7630840", - "mov DWORD PTR [r9+856], eax", - "mov eax, 4022750", - "mov DWORD PTR [r9+860], eax", - "mov eax, 4399818", - "mov DWORD PTR [r9+864], eax", - "mov eax, 5811406", - "mov DWORD PTR [r9+868], eax", - "mov eax, 1615530", - "mov DWORD PTR [r9+872], eax", - "mov eax, 6657188", - "mov DWORD PTR [r9+876], eax", - "mov eax, 6715099", - "mov DWORD PTR [r9+880], eax", - "mov eax, 6352379", - "mov DWORD PTR [r9+884], eax", - "mov eax, 7216819", - "mov DWORD PTR [r9+888], eax", - "mov eax, 3369273", - "mov DWORD PTR [r9+892], eax", - "mov eax, 4385746", - "mov DWORD PTR [r9+896], eax", - "mov eax, 11879", - "mov DWORD PTR [r9+900], eax", - "mov eax, 1370517", - "mov DWORD PTR [r9+904], eax", - "mov eax, 5360024", - "mov DWORD PTR [r9+908], eax", - "mov eax, 5016875", - "mov DWORD PTR [r9+912], eax", - "mov eax, 8165537", - "mov DWORD PTR [r9+916], eax", - "mov eax, 7835041", - "mov DWORD PTR [r9+920], eax", - "mov eax, 770441", - "mov DWORD PTR [r9+924], eax", - "mov eax, 5274859", - "mov DWORD PTR [r9+928], eax", - "mov eax, 1103344", - "mov DWORD PTR [r9+932], eax", - "mov eax, 7872272", - "mov DWORD PTR [r9+936], eax", - "mov eax, 553718", - "mov DWORD PTR [r9+940], eax", - "mov eax, 7520273", - "mov DWORD PTR [r9+944], eax", - "mov eax, 4949981", - "mov DWORD PTR [r9+948], eax", - "mov eax, 8240173", - "mov DWORD PTR [r9+952], eax", - "mov eax, 1514152", - "mov DWORD PTR [r9+956], eax", - "mov eax, 2185084", - "mov DWORD PTR [r9+960], eax", - "mov eax, 5256655", - "mov DWORD PTR [r9+964], eax", - "mov eax, 6022044", - "mov DWORD PTR [r9+968], eax", - "mov eax, 2193087", - "mov DWORD PTR [r9+972], eax", - "mov eax, 3014420", - "mov DWORD PTR [r9+976], eax", - "mov eax, 1716814", - "mov DWORD PTR [r9+980], eax", - "mov eax, 5454363", - "mov DWORD PTR [r9+984], eax", - "mov eax, 392707", - "mov DWORD PTR [r9+988], eax", - "mov eax, 303005", - "mov DWORD PTR [r9+992], eax", - "mov eax, 4849188", - "mov DWORD PTR [r9+996], eax", - "mov eax, 3974485", - "mov DWORD PTR [r9+1000], eax", - "mov eax, 3773731", - "mov DWORD PTR [r9+1004], eax", - "mov eax, 6480365", - "mov DWORD PTR [r9+1008], eax", - "mov eax, 781875", - "mov DWORD PTR [r9+1012], eax", - "mov eax, 7325939", - "mov DWORD PTR [r9+1016], eax", - "mov eax, 731434", - "mov DWORD PTR [r9+1020], eax", - "mov rsi, rdi", - "mov r8, r9", - "add r8, 1020", - "mov edi, 128", + "stmxcsr DWORD PTR [rsi+768]", + "mov r11d, DWORD PTR [rsi+768]", + "and r11d, 65535", + "mov eax, 8127", + "mov DWORD PTR [rsi+772], eax", + "ldmxcsr DWORD PTR [rsi+772]", + "lfence", + "movabs r9, 111012023893504", + "mov QWORD PTR [rsi], r9", + "movabs r9, 33764919763013891", + "mov QWORD PTR [rsi+8], r9", + "movabs r9, 32652304184483396", + "mov QWORD PTR [rsi+16], r9", + "movabs r9, 2003464812134697", + "mov QWORD PTR [rsi+24], r9", + "movabs r9, 10107995079564843", + "mov QWORD PTR [rsi+32], r9", + "movabs r9, 27008953388524718", + "mov QWORD PTR [rsi+40], r9", + "movabs r9, 23603259066260085", + "mov QWORD PTR [rsi+48], r9", + "movabs r9, 11510954738022985", + "mov QWORD PTR [rsi+56], r9", + "movabs r9, 4398527550166616", + "mov QWORD PTR [rsi+64], r9", + "movabs r9, 15401443493111205", + "mov QWORD PTR [rsi+72], r9", + "movabs r9, 31185040284548497", + "mov QWORD PTR [rsi+80], r9", + "movabs r9, 26937467947448680", + "mov QWORD PTR [rsi+88], r9", + "movabs r9, 19416172761943511", + "mov QWORD PTR [rsi+96], r9", + "movabs r9, 21916122901808376", + "mov QWORD PTR [rsi+104], r9", + "movabs r9, 35910200088776757", + "mov QWORD PTR [rsi+112], r9", + "movabs r9, 1202612321726977", + "mov QWORD PTR [rsi+120], r9", + "movabs r9, 411354790447719", + "mov QWORD PTR [rsi+128], r9", + "movabs r9, 15163111458665677", + "mov QWORD PTR [rsi+136], r9", + "movabs r9, 20565458766169348", + "mov QWORD PTR [rsi+144], r9", + "movabs r9, 16816682460325845", + "mov QWORD PTR [rsi+152], r9", + "movabs r9, 22920956268049798", + "mov QWORD PTR [rsi+160], r9", + "movabs r9, 23677166863944342", + "mov QWORD PTR [rsi+168], r9", + "movabs r9, 34703121006855168", + "mov QWORD PTR [rsi+176], r9", + "movabs r9, 33677345376425628", + "mov QWORD PTR [rsi+184], r9", + "movabs r9, 28933472397947454", + "mov QWORD PTR [rsi+192], r9", + "movabs r9, 19579390106369566", + "mov QWORD PTR [rsi+200], r9", + "movabs r9, 26799599498134591", + "mov QWORD PTR [rsi+208], r9", + "movabs r9, 15889643835192413", + "mov QWORD PTR [rsi+216], r9", + "movabs r9, 2282148053410728", + "mov QWORD PTR [rsi+224], r9", + "movabs r9, 16668952760323958", + "mov QWORD PTR [rsi+232], r9", + "movabs r9, 25011900965946676", + "mov QWORD PTR [rsi+240], r9", + "movabs r9, 23977247637478740", + "mov QWORD PTR [rsi+248], r9", + "movabs r9, 29427887555995366", + "mov QWORD PTR [rsi+256], r9", + "movabs r9, 22931526182748624", + "mov QWORD PTR [rsi+264], r9", + "movabs r9, 14929091579459166", + "mov QWORD PTR [rsi+272], r9", + "movabs r9, 29185144592086471", + "mov QWORD PTR [rsi+280], r9", + "movabs r9, 8329290213797750", + "mov QWORD PTR [rsi+288], r9", + "movabs r9, 31697782066745459", + "mov QWORD PTR [rsi+296], r9", + "movabs r9, 22432987854353025", + "mov QWORD PTR [rsi+304], r9", + "movabs r9, 545125843890401", + "mov QWORD PTR [rsi+312], r9", + "movabs r9, 31769864501989618", + "mov QWORD PTR [rsi+320], r9", + "movabs r9, 11662103226140216", + "mov QWORD PTR [rsi+328], r9", + "movabs r9, 20130185304250276", + "mov QWORD PTR [rsi+336], r9", + "movabs r9, 25354781094156910", + "mov QWORD PTR [rsi+344], r9", + "movabs r9, 30717142252233347", + "mov QWORD PTR [rsi+352], r9", + "movabs r9, 30375073877558844", + "mov QWORD PTR [rsi+360], r9", + "movabs r9, 5794237307816926", + "mov QWORD PTR [rsi+368], r9", + "movabs r9, 29849966874477923", + "mov QWORD PTR [rsi+376], r9", + "movabs r9, 1137925820308458", + "mov QWORD PTR [rsi+384], r9", + "movabs r9, 13305774323778583", + "mov QWORD PTR [rsi+392], r9", + "movabs r9, 31268732009491712", + "mov QWORD PTR [rsi+400], r9", + "movabs r9, 17002134848261444", + "mov QWORD PTR [rsi+408], r9", + "movabs r9, 35956774717033419", + "mov QWORD PTR [rsi+416], r9", + "movabs r9, 22036141462600008", + "mov QWORD PTR [rsi+424], r9", + "movabs r9, 35087477629023596", + "mov QWORD PTR [rsi+432], r9", + "movabs r9, 30337763489061025", + "mov QWORD PTR [rsi+440], r9", + "movabs r9, 20732429908239468", + "mov QWORD PTR [rsi+448], r9", + "movabs r9, 28041905903253186", + "mov QWORD PTR [rsi+456], r9", + "movabs r9, 35231517950811284", + "mov QWORD PTR [rsi+464], r9", + "movabs r9, 5760968484459269", + "mov QWORD PTR [rsi+472], r9", + "movabs r9, 29186403016613413", + "mov QWORD PTR [rsi+480], r9", + "movabs r9, 29809972144732485", + "mov QWORD PTR [rsi+488], r9", + "movabs r9, 19324591173389987", + "mov QWORD PTR [rsi+496], r9", + "movabs r9, 16492506917666904", + "mov QWORD PTR [rsi+504], r9", + "movabs r9, 14635985826474643", + "mov QWORD PTR [rsi+512], r9", + "movabs r9, 16398082059229396", + "mov QWORD PTR [rsi+520], r9", + "movabs r9, 9638297459285875", + "mov QWORD PTR [rsi+528], r9", + "movabs r9, 20692959164533664", + "mov QWORD PTR [rsi+536], r9", + "movabs r9, 10455835889373941", + "mov QWORD PTR [rsi+544], r9", + "movabs r9, 15088997507073265", + "mov QWORD PTR [rsi+552], r9", + "movabs r9, 19847271512942753", + "mov QWORD PTR [rsi+560], r9", + "movabs r9, 22278162875259735", + "mov QWORD PTR [rsi+568], r9", + "movabs r9, 7984765110959710", + "mov QWORD PTR [rsi+576], r9", + "movabs r9, 3517724245221606", + "mov QWORD PTR [rsi+584], r9", + "movabs r9, 29065087369580419", + "mov QWORD PTR [rsi+592], r9", + "movabs r9, 33749496538019589", + "mov QWORD PTR [rsi+600], r9", + "movabs r9, 22582830675910690", + "mov QWORD PTR [rsi+608], r9", + "movabs r9, 13774157688799364", + "mov QWORD PTR [rsi+616], r9", + "movabs r9, 33738338209470846", + "mov QWORD PTR [rsi+624], r9", + "movabs r9, 20549610337740179", + "mov QWORD PTR [rsi+632], r9", + "movabs r9, 1232604074686745", + "mov QWORD PTR [rsi+640], r9", + "movabs r9, 17645078572608834", + "mov QWORD PTR [rsi+648], r9", + "movabs r9, 21638646536106727", + "mov QWORD PTR [rsi+656], r9", + "movabs r9, 872067341384903", + "mov QWORD PTR [rsi+664], r9", + "movabs r9, 11559822875647717", + "mov QWORD PTR [rsi+672], r9", + "movabs r9, 5433172289935931", + "mov QWORD PTR [rsi+680], r9", + "movabs r9, 5358487101890844", + "mov QWORD PTR [rsi+688], r9", + "movabs r9, 6854656137752657", + "mov QWORD PTR [rsi+696], r9", + "movabs r9, 5370830838457625", + "mov QWORD PTR [rsi+704], r9", + "movabs r9, 20753904747165841", + "mov QWORD PTR [rsi+712], r9", + "movabs r9, 8027804982718602", + "mov QWORD PTR [rsi+720], r9", + "movabs r9, 31479735164668747", + "mov QWORD PTR [rsi+728], r9", + "movabs r9, 5314055668205759", + "mov QWORD PTR [rsi+736], r9", + "movabs r9, 29850847345983039", + "mov QWORD PTR [rsi+744], r9", + "movabs r9, 5636951310400997", + "mov QWORD PTR [rsi+752], r9", + "movabs r9, 27564133741392515", + "mov QWORD PTR [rsi+760], r9", + "movabs r9, 8233800205883732", + "mov QWORD PTR [rsi+768], r9", + "movabs r9, 30088883024233849", + "mov QWORD PTR [rsi+776], r9", + "movabs r9, 3338009929245701", + "mov QWORD PTR [rsi+784], r9", + "movabs r9, 14628791756398056", + "mov QWORD PTR [rsi+792], r9", + "movabs r9, 23830870862829877", + "mov QWORD PTR [rsi+800], r9", + "movabs r9, 28061014216302585", + "mov QWORD PTR [rsi+808], r9", + "movabs r9, 19997999096164636", + "mov QWORD PTR [rsi+816], r9", + "movabs r9, 19771555530215640", + "mov QWORD PTR [rsi+824], r9", + "movabs r9, 10447056982320729", + "mov QWORD PTR [rsi+832], r9", + "movabs r9, 35286145636332471", + "mov QWORD PTR [rsi+840], r9", + "movabs r9, 14470225858518424", + "mov QWORD PTR [rsi+848], r9", + "movabs r9, 30807937853347003", + "mov QWORD PTR [rsi+856], r9", + "movabs r9, 699409659546815", + "mov QWORD PTR [rsi+864], r9", + "movabs r9, 12945035726727688", + "mov QWORD PTR [rsi+872], r9", + "movabs r9, 7098008983067813", + "mov QWORD PTR [rsi+880], r9", + "movabs r9, 28266511219523944", + "mov QWORD PTR [rsi+888], r9", + "movabs r9, 15135022374814013", + "mov QWORD PTR [rsi+896], r9", + "movabs r9, 1158610381635861", + "mov QWORD PTR [rsi+904], r9", + "movabs r9, 31802227079365879", + "mov QWORD PTR [rsi+912], r9", + "movabs r9, 2027559572878823", + "mov QWORD PTR [rsi+920], r9", + "movabs r9, 7402805639339334", + "mov QWORD PTR [rsi+928], r9", + "movabs r9, 8205002449640623", + "mov QWORD PTR [rsi+936], r9", + "movabs r9, 31250542829661849", + "mov QWORD PTR [rsi+944], r9", + "movabs r9, 19527171898598875", + "mov QWORD PTR [rsi+952], r9", + "movabs r9, 26390134497937253", + "mov QWORD PTR [rsi+960], r9", + "movabs r9, 26173917256840158", + "mov QWORD PTR [rsi+968], r9", + "movabs r9, 31797902045269139", + "mov QWORD PTR [rsi+976], r9", + "movabs r9, 20765007236602922", + "mov QWORD PTR [rsi+984], r9", + "movabs r9, 16834811519265361", + "mov QWORD PTR [rsi+992], r9", + "movabs r9, 30142973844857679", + "mov QWORD PTR [rsi+1000], r9", + "movabs r9, 6014775284471242", + "mov QWORD PTR [rsi+1008], r9", + "movabs r9, 8490214048855735", + "mov QWORD PTR [rsi+1016], r9", + "mov eax, 8380417", + "movq xmm15, rax", + "pshufd xmm15, xmm15, 0", + "mov eax, -58728449", + "movq xmm14, rax", + "pshufd xmm14, xmm14, 0", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 1008", + "mov ecx, 32", "20:", - "mov r9d, DWORD PTR [r8]", - "add r8, -4", - "mov ecx, 1", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm2, XMMWORD PTR [rdx+16]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 27", + "pshufd xmm12, xmm13, 245", + "add r8, -16", + "pshufd xmm0, xmm0, 216", + "pshufd xmm2, xmm2, 216", + "movdqa xmm1, xmm0", + "punpcklqdq xmm0, xmm2", + "punpckhqdq xmm1, xmm2", + "movdqa xmm3, xmm1", + "psubd xmm3, xmm0", + "paddd xmm3, xmm15", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqa xmm1, xmm0", + "punpckldq xmm0, xmm3", + "punpckhdq xmm1, xmm3", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+16], xmm1", + "add rdx, 32", + "sub rcx, 1", + "jne 20b", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 504", + "mov ecx, 32", "21:", - "mov eax, DWORD PTR [rsi]", - "mov r10d, DWORD PTR [rsi+4]", - "mov edx, eax", - "add edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi], edx", - "add eax, 8380417", - "sub eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov DWORD PTR [rsi+4], r10d", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm1, XMMWORD PTR [rdx+16]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 5", + "pshufd xmm12, xmm13, 245", + "add r8, -8", + "movdqa xmm2, xmm0", + "punpcklqdq xmm0, xmm1", + "punpckhqdq xmm2, xmm1", + "movdqa xmm1, xmm2", + "movdqa xmm3, xmm1", + "psubd xmm3, xmm0", + "paddd xmm3, xmm15", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqa xmm1, xmm0", + "punpcklqdq xmm0, xmm3", + "punpckhqdq xmm1, xmm3", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+16], xmm1", + "add rdx, 32", "sub rcx, 1", "jne 21b", - "add rsi, 4", - "sub rdi, 1", - "jne 20b", - "sub rsi, 1024", - "mov edi, 64", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 252", + "mov eax, 32", "22:", - "mov r9d, DWORD PTR [r8]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 0", + "pshufd xmm12, xmm13, 245", "add r8, -4", - "mov ecx, 2", + "mov ecx, 1", "23:", - "mov eax, DWORD PTR [rsi]", - "mov r10d, DWORD PTR [rsi+8]", - "mov edx, eax", - "add edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi], edx", - "add eax, 8380417", - "sub eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov DWORD PTR [rsi+8], r10d", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm1, XMMWORD PTR [rdx+16]", + "movdqa xmm3, xmm1", + "psubd xmm3, xmm0", + "paddd xmm3, xmm15", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+16], xmm3", + "add rdx, 16", "sub rcx, 1", "jne 23b", - "add rsi, 8", - "sub rdi, 1", + "add rdx, 16", + "sub rax, 1", "jne 22b", - "sub rsi, 1024", - "mov edi, 32", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 124", + "mov eax, 16", "24:", - "mov r9d, DWORD PTR [r8]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 0", + "pshufd xmm12, xmm13, 245", "add r8, -4", - "mov ecx, 4", + "mov ecx, 2", "25:", - "mov eax, DWORD PTR [rsi]", - "mov r10d, DWORD PTR [rsi+16]", - "mov edx, eax", - "add edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi], edx", - "add eax, 8380417", - "sub eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov DWORD PTR [rsi+16], r10d", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm1, XMMWORD PTR [rdx+32]", + "movdqa xmm3, xmm1", + "psubd xmm3, xmm0", + "paddd xmm3, xmm15", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+32], xmm3", + "add rdx, 16", "sub rcx, 1", "jne 25b", - "add rsi, 16", - "sub rdi, 1", + "add rdx, 32", + "sub rax, 1", "jne 24b", - "sub rsi, 1024", - "mov edi, 16", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 60", + "mov eax, 8", "26:", - "mov r9d, DWORD PTR [r8]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 0", + "pshufd xmm12, xmm13, 245", "add r8, -4", - "mov ecx, 8", + "mov ecx, 4", "27:", - "mov eax, DWORD PTR [rsi]", - "mov r10d, DWORD PTR [rsi+32]", - "mov edx, eax", - "add edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi], edx", - "add eax, 8380417", - "sub eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov DWORD PTR [rsi+32], r10d", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm1, XMMWORD PTR [rdx+64]", + "movdqa xmm3, xmm1", + "psubd xmm3, xmm0", + "paddd xmm3, xmm15", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+64], xmm3", + "add rdx, 16", "sub rcx, 1", "jne 27b", - "add rsi, 32", - "sub rdi, 1", + "add rdx, 64", + "sub rax, 1", "jne 26b", - "sub rsi, 1024", - "mov edi, 8", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 28", + "mov eax, 4", "28:", - "mov r9d, DWORD PTR [r8]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 0", + "pshufd xmm12, xmm13, 245", "add r8, -4", - "mov ecx, 16", + "mov ecx, 8", "29:", - "mov eax, DWORD PTR [rsi]", - "mov r10d, DWORD PTR [rsi+64]", - "mov edx, eax", - "add edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi], edx", - "add eax, 8380417", - "sub eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov DWORD PTR [rsi+64], r10d", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm1, XMMWORD PTR [rdx+128]", + "movdqa xmm3, xmm1", + "psubd xmm3, xmm0", + "paddd xmm3, xmm15", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+128], xmm3", + "add rdx, 16", "sub rcx, 1", "jne 29b", - "add rsi, 64", - "sub rdi, 1", + "add rdx, 128", + "sub rax, 1", "jne 28b", - "sub rsi, 1024", - "mov edi, 4", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 12", + "mov eax, 2", "210:", - "mov r9d, DWORD PTR [r8]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 0", + "pshufd xmm12, xmm13, 245", "add r8, -4", - "mov ecx, 32", + "mov ecx, 16", "211:", - "mov eax, DWORD PTR [rsi]", - "mov r10d, DWORD PTR [rsi+128]", - "mov edx, eax", - "add edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi], edx", - "add eax, 8380417", - "sub eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov DWORD PTR [rsi+128], r10d", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm1, XMMWORD PTR [rdx+256]", + "movdqa xmm3, xmm1", + "psubd xmm3, xmm0", + "paddd xmm3, xmm15", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+256], xmm3", + "add rdx, 16", "sub rcx, 1", "jne 211b", - "add rsi, 128", - "sub rdi, 1", + "add rdx, 256", + "sub rax, 1", "jne 210b", - "sub rsi, 1024", - "mov edi, 2", + "mov rdx, rdi", + "mov r8, rsi", + "add r8, 4", + "mov eax, 1", "212:", - "mov r9d, DWORD PTR [r8]", + "movdqu xmm13, XMMWORD PTR [r8]", + "pshufd xmm13, xmm13, 0", + "pshufd xmm12, xmm13, 245", "add r8, -4", - "mov ecx, 64", + "mov ecx, 32", "213:", - "mov eax, DWORD PTR [rsi]", - "mov r10d, DWORD PTR [rsi+256]", - "mov edx, eax", - "add edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi], edx", - "add eax, 8380417", - "sub eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov DWORD PTR [rsi+256], r10d", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdx]", + "movdqu xmm1, XMMWORD PTR [rdx+512]", + "movdqa xmm3, xmm1", + "psubd xmm3, xmm0", + "paddd xmm3, xmm15", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdx], xmm0", + "movdqu XMMWORD PTR [rdx+512], xmm3", + "add rdx, 16", "sub rcx, 1", "jne 213b", - "add rsi, 256", - "sub rdi, 1", + "add rdx, 512", + "sub rax, 1", "jne 212b", - "sub rsi, 1024", - "mov edi, 1", + "mov rdx, rdi", + "mov eax, 16382", + "movq xmm13, rax", + "pshufd xmm13, xmm13, 0", + "movdqa xmm12, xmm13", + "mov ecx, 64", "214:", - "mov r9d, DWORD PTR [r8]", - "add r8, -4", - "mov ecx, 128", - "215:", - "mov eax, DWORD PTR [rsi]", - "mov r10d, DWORD PTR [rsi+512]", - "mov edx, eax", - "add edx, r10d", - "sub edx, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add edx, r11d", - "mov DWORD PTR [rsi], edx", - "add eax, 8380417", - "sub eax, r10d", - "sub eax, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add eax, r11d", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov DWORD PTR [rsi+512], r10d", - "add rsi, 4", + "movdqu xmm3, XMMWORD PTR [rdx]", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdx], xmm3", + "add rdx, 16", "sub rcx, 1", - "jne 215b", - "add rsi, 512", - "sub rdi, 1", "jne 214b", - "sub rsi, 1024", - "mov r9d, 8347681", - "mov ecx, 256", - "216:", - "mov eax, DWORD PTR [rsi]", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov DWORD PTR [rsi], r10d", - "add rsi, 4", - "sub rcx, 1", - "jne 216b", + "lfence", + "mov DWORD PTR [rsi+768], r11d", + "ldmxcsr DWORD PTR [rsi+768]", "ret", ) } From ae7eb79564f854750ddeecb53ab8f7e0bf4d17ed Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 14:36:01 +0000 Subject: [PATCH 2/2] ML-DSA on x86-64: polynomial arithmetic in SSE2 MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit vg_mldsa_multiply_ntt, vg_mldsa_multiply_add_ntt, vg_mldsa_add and vg_mldsa_sub now compute on four coefficients at a time in SSE2 registers, with the vector helpers of the NTT (Vec.lean). The products are two Montgomery multiplications: f·g·2⁻³², then by 2⁶⁴ mod q, which is f·g in [0, 2q), reduced with vcsub (for multiply_add, h is then added and the sum reduced). They use pmuludq, so they run inside ML-KEM's withMxcsr. These functions have no working space and use no stack, so MXCSR goes through the last 8 bytes of h, addressed through r8 = h. The last four coefficients of h are loaded into xmm6 first. The loop stores the first 252 coefficients. The last four are computed from registers before MXCSR is loaded back, and stored after it. Their callers' proofs are unchanged: the functions still need no stack and never write rsp. add and sub use paddd/psubd and vcsub/vcadd. Proofs: the lanes of a product (mul_lane, mulAdd_lane), the loop (Mul.step, Mul.loop_ok), the last block (Mul.last), the function (Mul.fn_ok), and withMxcsr through the end of a polynomial (withMxcsrH_ok). AddSub is reproven the same way. ML-DSA-65 instructions executed (callgrind, PR #464 -> this): sign -19%, verify -11%, keygen -2%. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01Ddof3szoTi7HB8iCsM2MCr --- README.md | 6 +- docs/algorithms/ml-dsa-44.toml | 2 +- docs/algorithms/ml-dsa-65.toml | 2 +- docs/algorithms/ml-dsa-87.toml | 2 +- .../Artifacts/MlDsaArith/X86_64.lean | 10 + .../Impl/MlDsa/X86_64/Arith/AddSub.lean | 36 +- .../Impl/MlDsa/X86_64/Arith/Mul.lean | 66 ++- .../Proof/MlDsa/X86_64/Arith/AddSub.lean | 307 +++++----- .../Proof/MlDsa/X86_64/Arith/Mul.lean | 548 +++++++++++------- .../Proof/MlDsa/X86_64/Arith/Mxcsr.lean | 73 +++ src/asm/x86_64/mldsa.rs | 306 ++++++++-- 11 files changed, 918 insertions(+), 440 deletions(-) create mode 100644 lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Mxcsr.lean diff --git a/README.md b/README.md index 00030ddd4..a14742701 100644 --- a/README.md +++ b/README.md @@ -774,7 +774,7 @@ yours to keep: ✅ -✅ SSE2 NTT +✅ SSE2 polynomial arithmetic ✅ SHA extensions @@ -790,7 +790,7 @@ yours to keep: ✅ -✅ SSE2 NTT +✅ SSE2 polynomial arithmetic ✅ SHA extensions @@ -806,7 +806,7 @@ yours to keep: ✅ -✅ SSE2 NTT +✅ SSE2 polynomial arithmetic ✅ SHA extensions diff --git a/docs/algorithms/ml-dsa-44.toml b/docs/algorithms/ml-dsa-44.toml index 7f316830d..854e6670e 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 NTT" } +optimized = { x86_64 = "SSE2 polynomial arithmetic" } diff --git a/docs/algorithms/ml-dsa-65.toml b/docs/algorithms/ml-dsa-65.toml index 26de285e2..e87725611 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 NTT" } +optimized = { x86_64 = "SSE2 polynomial arithmetic" } diff --git a/docs/algorithms/ml-dsa-87.toml b/docs/algorithms/ml-dsa-87.toml index 1f76a641b..9cfbc9f3a 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 NTT" } +optimized = { x86_64 = "SSE2 polynomial arithmetic" } diff --git a/lean/VerifiedGarbage/Artifacts/MlDsaArith/X86_64.lean b/lean/VerifiedGarbage/Artifacts/MlDsaArith/X86_64.lean index f8fdd427a..a3d95256b 100644 --- a/lean/VerifiedGarbage/Artifacts/MlDsaArith/X86_64.lean +++ b/lean/VerifiedGarbage/Artifacts/MlDsaArith/X86_64.lean @@ -46,6 +46,10 @@ def artifacts : List Artifact := [ { Spec.MlDsa.mulApi with target := X86_64.target doc := Spec.MlDsa.mulApi.doc + (notes := ["The function computes on four coefficients at a time in SSE2 registers. It sets MXCSR to \ + `0x1FBF` around its multiplications (Intel's mitigation of MXCSR-configuration-dependent timing), \ + through the last 8 bytes of `h`, which it stores last, and loads the caller's MXCSR back before \ + returning."]) code := Impl.MlDsa.X86_64.Arith.mul contract := Spec.MlDsa.mulContract X86_64.abi verified := Proof.MlDsa.X86_64.Arith.mul_verified @@ -53,6 +57,10 @@ def artifacts : List Artifact := [ { Spec.MlDsa.mulAddApi with target := X86_64.target doc := Spec.MlDsa.mulAddApi.doc + (notes := ["The function computes on four coefficients at a time in SSE2 registers. It sets MXCSR to \ + `0x1FBF` around its multiplications (Intel's mitigation of MXCSR-configuration-dependent timing), \ + through the last 8 bytes of `h`, which it stores last, and loads the caller's MXCSR back before \ + returning."]) code := Impl.MlDsa.X86_64.Arith.mulAdd contract := Spec.MlDsa.mulAddContract X86_64.abi verified := Proof.MlDsa.X86_64.Arith.mulAdd_verified @@ -60,6 +68,7 @@ def artifacts : List Artifact := [ { Spec.MlDsa.addApi with target := X86_64.target doc := Spec.MlDsa.addApi.doc + (notes := ["The function computes on four coefficients at a time in SSE2 registers."]) code := Impl.MlDsa.X86_64.Arith.add contract := Spec.MlDsa.addContract X86_64.abi verified := Proof.MlDsa.X86_64.Arith.add_verified @@ -67,6 +76,7 @@ def artifacts : List Artifact := [ { Spec.MlDsa.subApi with target := X86_64.target doc := Spec.MlDsa.subApi.doc + (notes := ["The function computes on four coefficients at a time in SSE2 registers."]) code := Impl.MlDsa.X86_64.Arith.sub contract := Spec.MlDsa.subContract X86_64.abi verified := Proof.MlDsa.X86_64.Arith.sub_verified diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/AddSub.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/AddSub.lean index 1ea537f64..3383b263c 100644 --- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/AddSub.lean +++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/AddSub.lean @@ -1,33 +1,39 @@ -import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Common +import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Vec /-! # ML-DSA on x86-64: `vg_mldsa_add` and `vg_mldsa_sub` -`add(f = rdi, g = rsi)` and `sub(f = rdi, g = rsi)` run over the 256 -coefficients with `rdi` and `rsi` pointing at coefficient `i` of `f` and -`g`, and `rcx = 256 - i` counting down: `f[i] + g[i]` (for `sub`, -`f[i] + q - g[i]`), less than `2q`, is reduced with `csubQ` and stored to -`f[i]`. Every address and branch depends only on the pointers. +`add(f = rdi, g = rsi)` and `sub(f = rdi, g = rsi)` compute on four +coefficients at a time, as doublewords of SSE registers (see `Vec.lean`), +with `rdi` and `rsi` pointing at coefficient `4i` of `f` and `g`, and +`rcx = 64 - i` counting down: `f + g`, less than `2q`, is reduced with +`vcsub` (for `sub`, `f - g`, in `(-q, q)`, with `vcadd`) and stored to `f`, +with `q` in the doublewords of `xmm15`. Every address and branch depends +only on the pointers. -/ namespace VG.Impl.MlDsa.X86_64.Arith open VG.X86_64 +open VG.Impl.MlKem.X86_64 (xb rcxLoop) -/-- Advance the two pointers and count down. -/ -def step2 : List Instr := - [.alu .add .rdi (.imm 4), .alu .add .rsi (.imm 4), .alu .sub .rcx (.imm 1)] +/-- `q` in the doublewords of `xmm15`, through `rax`. -/ +def qPro : List Instr := [.mov32 .rax (.imm 8380417), .xop (.movq .xmm15 .rax), .xop (.pshufd .xmm15 .xmm15 0)] + +/-- Store `xmm0` to `f` and advance the two pointers. -/ +def accTail : List Instr := + [.movdquStore (at_ .rdi 0) .xmm0, .alu .add .rdi (.imm 16), .alu .add .rsi (.imm 16)] def addBody : List Instr := - [.mov32 .rax (.mem (at_ .rdi 0)), .alu32 .add .rax (.mem (at_ .rsi 0))] ++ csubQ .rax .rdx ++ - [.store32 (at_ .rdi 0) .rax] ++ step2 + [.movdquLoad .xmm0 (at_ .rdi 0), .movdquLoad .xmm1 (at_ .rsi 0), xb .paddd .xmm0 .xmm1] ++ + vcsub .xmm0 .xmm2 ++ accTail def subBody : List Instr := - [.mov32 .rax (.mem (at_ .rdi 0)), .alu32 .add .rax (.imm qImm), .alu32 .sub .rax (.mem (at_ .rsi 0))] ++ - csubQ .rax .rdx ++ [.store32 (at_ .rdi 0) .rax] ++ step2 + [.movdquLoad .xmm0 (at_ .rdi 0), .movdquLoad .xmm1 (at_ .rsi 0), xb .psubd .xmm0 .xmm1] ++ + vcadd .xmm0 .xmm2 ++ accTail -def add : Prog isa := .seq (.block [.mov32 .rcx (.imm 256)]) (.loop (.block addBody) .ne) +def add : Prog isa := .seq (.block qPro) (rcxLoop 64 addBody) -def sub : Prog isa := .seq (.block [.mov32 .rcx (.imm 256)]) (.loop (.block subBody) .ne) +def sub : Prog isa := .seq (.block qPro) (rcxLoop 64 subBody) end VG.Impl.MlDsa.X86_64.Arith diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Mul.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Mul.lean index bd0679661..59f8d7b71 100644 --- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Mul.lean +++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Mul.lean @@ -1,41 +1,63 @@ -import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Common +import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Vec /-! # ML-DSA on x86-64: `vg_mldsa_multiply_ntt` and `vg_mldsa_multiply_add_ntt` `multiplyNTT(h = rdi, f = rsi, g = rdx)` and `multiplyAddNTT(h = rdi, -f = rsi, g = rdx)` run over the 256 coefficients with `rdi`, `rsi` and `r8` -(`g`, as `mul` writes `rdx`) pointing at coefficient `i` of `h`, `f` and -`g`, and `rcx = 256 - i` counting down: the product `f[i] · g[i]` (by -`mul`, less than `q²`; for `multiplyAddNTT`, plus `h[i]`) is reduced with -`reduce` and stored to `h[i]`. Every address and branch depends only on the -pointers. +f = rsi, g = rdx)` compute on four coefficients at a time, as doublewords of +SSE registers (see `Vec.lean`), with `rdi`, `rsi` and `rdx` pointing at +coefficient `4i` of `h`, `f` and `g`, and `rcx = 63 - i` counting down: +`vmont` of `f` by `g` is `f · g · 2⁻³²`, and `vmont` of that by +`2⁶⁴ mod q = 2365951` (in `xmm11`) is `f · g`, less than `2q`, which `vcsub` +reduces (for `multiplyAddNTT`, `h` is then added and the sum reduced), and +it is stored to `h`. + +The functions have no working space and use no stack, so `withMxcsr` saves +MXCSR through the last 8 bytes of `h` (`r8 = h`): the last four coefficients +of `h` are loaded to `xmm6` first, the loop stores the first 252, the last +four are computed from registers before MXCSR is loaded back, and stored +after it. Every address and branch depends only on the pointers. -/ namespace VG.Impl.MlDsa.X86_64.Arith open VG.X86_64 +open VG.Impl.MlKem.X86_64 (xb xmov withMxcsr rcxLoop) -/-- Advance the three pointers and count down. -/ -def step3 : List Instr := - [.alu .add .rdi (.imm 4), .alu .add .rsi (.imm 4), .alu .add .r8 (.imm 4), .alu .sub .rcx (.imm 1)] +/-- The constants and `2⁶⁴ mod q` in the doublewords of `xmm11`. -/ +def mulPro : List Instr := + vconsts ++ [.mov32 .rax (.imm 2365951), .xop (.movq .xmm11 .rax), .xop (.pshufd .xmm11 .xmm11 0)] -/-- `f[i] · g[i]`, in `rax`. -/ -def mulHead : List Instr := - [.mov32 .rax (.mem (at_ .rsi 0)), .mov32 .r9 (.mem (at_ .r8 0)), .mul .r9] +/-- The coefficients of `f`, `g` and `h` to `xmm3`, `xmm13` and `xmm5`. -/ +def mulLoads : List Instr := + [.movdquLoad .xmm3 (at_ .rsi 0), .movdquLoad .xmm13 (at_ .rdx 0), .movdquLoad .xmm5 (at_ .rdi 0)] -/-- `f[i] · g[i] + h[i]`, in `rax`. -/ -def mulAddHead : List Instr := - mulHead ++ [.mov32 .r9 (.mem (at_ .rdi 0)), .alu .add .rax (.reg .r9)] +/-- `f · g` for four coefficients, in `[0, q)`, in `xmm3`. -/ +def mulCore : List Instr := + [.xop (.pshufd .xmm12 .xmm13 0xF5)] ++ vmont .xmm3 .xmm13 .xmm12 .xmm2 .xmm4 ++ + vmont .xmm3 .xmm11 .xmm11 .xmm2 .xmm4 ++ vcsub .xmm3 .xmm2 -def mulBody : List Instr := mulHead ++ reduce ++ [.store32 (at_ .rdi 0) .r10] ++ step3 +/-- `h + f · g` for four coefficients, in `[0, q)`, in `xmm3`. -/ +def mulAddCore : List Instr := mulCore ++ [xb .paddd .xmm3 .xmm5] ++ vcsub .xmm3 .xmm2 -def mulAddBody : List Instr := mulAddHead ++ reduce ++ [.store32 (at_ .rdi 0) .r10] ++ step3 +/-- Store `xmm3` to `h` and advance the three pointers. -/ +def mulTail : List Instr := + [.movdquStore (at_ .rdi 0) .xmm3, .alu .add .rdi (.imm 16), .alu .add .rsi (.imm 16), + .alu .add .rdx (.imm 16)] -def mul : Prog isa := - .seq (.block [.mov .r8 (.reg .rdx)]) (.seq (.block [.mov32 .rcx (.imm 256)]) (.loop (.block mulBody) .ne)) +/-- The last four coefficients, with those of `h` in `xmm6`, to `xmm3`. -/ +def mulLast (core : List Instr) : List Instr := + [.movdquLoad .xmm3 (at_ .rsi 0), .movdquLoad .xmm13 (at_ .rdx 0), xmov .xmm5 .xmm6] ++ core -def mulAdd : Prog isa := - .seq (.block [.mov .r8 (.reg .rdx)]) (.seq (.block [.mov32 .rcx (.imm 256)]) (.loop (.block mulAddBody) .ne)) +/-- A function of `h`, `f` and `g` four coefficients at a time by `core`. -/ +def mulFn (core : List Instr) : Prog isa := + .seq (.block [.mov .r8 (.reg .rdi), .movdquLoad .xmm6 (at_ .rdi 1008)]) + (.seq (withMxcsr .r8 1016 + (.seq (.block mulPro) (.seq (rcxLoop 63 (mulLoads ++ core ++ mulTail)) (.block (mulLast core))))) + (.block [.movdquStore (at_ .rdi 0) .xmm3])) + +def mul : Prog isa := mulFn mulCore + +def mulAdd : Prog isa := mulFn mulAddCore end VG.Impl.MlDsa.X86_64.Arith diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/AddSub.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/AddSub.lean index 96b1bc978..bf574a3af 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/AddSub.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/AddSub.lean @@ -1,149 +1,194 @@ import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.AddSub +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.VLay import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Basic import VerifiedGarbage.Proof.Framework.X86_64.Abi /-! # ML-DSA on x86-64: `vg_mldsa_add` and `vg_mldsa_sub` -Untrusted: everything here is checked by Lean. +Untrusted: everything here is checked by Lean. Each iteration of the loop +stores four coefficients to `f` (`addBody_ok`, `subBody_ok`), each `csubL` +of the sum (`caddL` of the difference), whose value is `addD_toNat` +(`subD_toNat`); the loop leaves `f` with all 256 (`AddSub.fn_ok`). -/ namespace VG.Proof.MlDsa.X86_64.Arith open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Arith open VG.Proof.MlDsa.Arith -open VG.Proof.MlKem.X86_64 (Keep WP.keep writesOnly gprPreserved_of wp_countdown ifp ifn ptr_step) +open VG.Proof.MlKem.X86_64 (Keep WP.keep writesOnly gprPreserved_of ifp ifn ptr_step GOnly wp_rcxLoop xmm_setXmm + add_ofNat_zero) +open VG.Impl.MlKem.X86_64 (xb xmov) open VG.Spec.MlDsa (q n Poly Zq coeffAt polyAt Reduced PolyIs) -/-! ## One coefficient -/ - -theorem addBody_ok (s : State) (h1 : InRegions (s.rd ++ s.wr) (s.gpr .rdi) 4) - (h2 : InRegions (s.rd ++ s.wr) (s.gpr .rsi) 4) (h3 : InRegions s.wr (s.gpr .rdi) 4) : - WP isa (.block addBody) s fun s' => - (s'.mem = s.mem.writeW (s.gpr .rdi) (csubD (s.mem.readW (s.gpr .rdi) 32 + s.mem.readW (s.gpr .rsi) 32)) ∧ - s'.gpr .rdi = s.gpr .rdi + 4 ∧ s'.gpr .rsi = s.gpr .rsi + 4 ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ - s'.zf = some (s.gpr .rcx - 1 == 0)) ∧ Keep [.rax, .rdx, .rdi, .rsi, .rcx] s s' := by - refine WP.keep _ ?_ (by decide) - unfold addBody csubQ step2 - xrund [h1, h2, h3, List.cons_append, List.nil_append, csubD] - -theorem subBody_ok (s : State) (h1 : InRegions (s.rd ++ s.wr) (s.gpr .rdi) 4) - (h2 : InRegions (s.rd ++ s.wr) (s.gpr .rsi) 4) (h3 : InRegions s.wr (s.gpr .rdi) 4) : - WP isa (.block subBody) s fun s' => - (s'.mem = s.mem.writeW (s.gpr .rdi) - (csubD (s.mem.readW (s.gpr .rdi) 32 + qImm - s.mem.readW (s.gpr .rsi) 32)) ∧ - s'.gpr .rdi = s.gpr .rdi + 4 ∧ s'.gpr .rsi = s.gpr .rsi + 4 ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ - s'.zf = some (s.gpr .rcx - 1 == 0)) ∧ Keep [.rax, .rdx, .rdi, .rsi, .rcx] s s' := by - refine WP.keep _ ?_ (by decide) - unfold subBody csubQ step2 - xrund [h1, h2, h3, List.cons_append, List.nil_append, csubD] +/-! ## Four coefficients -/ + +/-- What an iteration of `add` stores. -/ +def addV (x y : BitVec 128) : BitVec 128 := csubV (XBinOp.eval .paddd x y) + +/-- What an iteration of `sub` stores. -/ +def subV (x y : BitVec 128) : BitVec 128 := caddV (XBinOp.eval .psubd x y) + +theorem dword_addV (x y : BitVec 128) {i : Nat} (hi : i < 4) : + dword (addV x y) i = csubL (dword x i + dword y i) := by + rw [addV, dword_csubV _ hi, dword_paddd _ _ hi] + +theorem dword_subV (x y : BitVec 128) {i : Nat} (hi : i < 4) : + dword (subV x y) i = caddL (dword x i - dword y i) := by + rw [subV, dword_caddV _ hi, dword_psubd _ _ hi] + +/-- The body of `add` or `sub`: the store of `F x y` of the vectors at `rdi` +and `rsi`, and the counts. -/ +theorem accBody_ok {op : XBinOp} {fix : List Instr} {F : BitVec 128 → BitVec 128 → BitVec 128} + (hF : ∀ (s : State), s.xmm .xmm15 = qV → + WP isa (.block (xb op .xmm0 .xmm1 :: fix)) s fun s' => + s'.xmm .xmm0 = F (s.xmm .xmm0) (s.xmm .xmm1) ∧ s'.gpr = s.gpr ∧ s'.mem = s.mem ∧ s'.rd = s.rd ∧ + s'.wr = s.wr ∧ s'.xmm .xmm15 = qV) + (s : State) (hq : s.xmm .xmm15 = qV) (h1 : InRegions (s.rd ++ s.wr) (s.gpr .rdi) 16) + (h2 : InRegions (s.rd ++ s.wr) (s.gpr .rsi) 16) (h3 : InRegions s.wr (s.gpr .rdi) 16) : + WP isa (.block (([.movdquLoad .xmm0 (at_ .rdi 0), .movdquLoad .xmm1 (at_ .rsi 0)] : List Instr) ++ + ((xb op .xmm0 .xmm1 :: fix) ++ (accTail ++ ([.alu .sub .rcx (.imm 1)] : List Instr))))) s fun s' => + s'.mem = s.mem.writeW (s.gpr .rdi) (F (s.mem.readW (s.gpr .rdi) 128) (s.mem.readW (s.gpr .rsi) 128)) ∧ + s'.gpr .rdi = s.gpr .rdi + 16 ∧ s'.gpr .rsi = s.gpr .rsi + 16 ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ + s'.zf = some (s.gpr .rcx - 1 == 0) ∧ s'.rd = s.rd ∧ s'.wr = s.wr ∧ s'.xmm .xmm15 = qV ∧ + Keep [.rdi, .rsi, .rcx] s s' := by + rw [WP.block_append_iff] + vrund [h1, h2] + rw [show Instr.xop (.bin op .xmm0 .xmm1) :: (fix ++ (accTail ++ [.alu .sub .rcx (.imm 1)])) = + (xb op .xmm0 .xmm1 :: fix) ++ (accTail ++ ([.alu .sub .rcx (.imm 1)] : List Instr)) from rfl, + WP.block_append_iff] + refine WP.mono (hF _ (by simp only [xmm_setXmm, reduceCtorEq, ite_false]; exact hq)) + fun s2 ⟨h0, g2, m2, r2, w2, q2⟩ => ?_ + simp only [accTail] + vrund [g2, m2, r2, w2, h3, h0, q2] + refine ⟨fun r hr => ?_, rfl, rfl⟩ + simp only [List.mem_cons, List.not_mem_nil, or_false, not_or] at hr + simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, hr.1, hr.2.1, hr.2.2, ite_false] + +theorem addFix_ok (s : State) (hq : s.xmm .xmm15 = qV) : + WP isa (.block (xb .paddd .xmm0 .xmm1 :: vcsub .xmm0 .xmm2)) s fun s' => + s'.xmm .xmm0 = addV (s.xmm .xmm0) (s.xmm .xmm1) ∧ s'.gpr = s.gpr ∧ s'.mem = s.mem ∧ s'.rd = s.rd ∧ + s'.wr = s.wr ∧ s'.xmm .xmm15 = qV := by + simp only [vcsub, vcadd, xmov, xb] + vrun [eval_movdqa] + rw [hq] + exact ⟨rfl, trivial, trivial, trivial, trivial, rfl⟩ + +theorem subFix_ok (s : State) (hq : s.xmm .xmm15 = qV) : + WP isa (.block (xb .psubd .xmm0 .xmm1 :: vcadd .xmm0 .xmm2)) s fun s' => + s'.xmm .xmm0 = subV (s.xmm .xmm0) (s.xmm .xmm1) ∧ s'.gpr = s.gpr ∧ s'.mem = s.mem ∧ s'.rd = s.rd ∧ + s'.wr = s.wr ∧ s'.xmm .xmm15 = qV := by + simp only [vcadd, xmov, xb] + vrun [eval_movdqa] + rw [hq] + exact ⟨rfl, trivial, trivial, trivial, trivial, rfl⟩ /-! ## The loop -/ namespace AddSub -/-- After `i` coefficients, each one `v k`. -/ +/-- After `i` vectors, each coefficient before `4i` is `v k`. -/ structure Inv (s₀ : State) (v : Nat → BitVec 32) (i : Nat) (s : State) : Prop where - rdi : s.gpr .rdi = s₀.gpr .rdi + BitVec.ofNat 64 (4 * i) - rsi : s.gpr .rsi = s₀.gpr .rsi + BitVec.ofNat 64 (4 * i) + rdi : s.gpr .rdi = s₀.gpr .rdi + BitVec.ofNat 64 (16 * i) + rsi : s.gpr .rsi = s₀.gpr .rsi + BitVec.ofNat 64 (16 * i) rd : s.rd = s₀.rd wr : s.wr = s₀.wr + q : s.xmm .xmm15 = qV frame : Frame [pR (s₀.gpr .rdi)] s₀.mem s.mem - coeff : ∀ k < 256, coeffAt s.mem (s₀.gpr .rdi) k = if k < i then v k else coeffAt s₀.mem (s₀.gpr .rdi) k + coeff : ∀ k < 256, coeffAt s.mem (s₀.gpr .rdi) k = if k < 4 * i then v k else coeffAt s₀.mem (s₀.gpr .rdi) k section variable {t : Poly → Poly → Poly} {s₀ : State} (hp : (accK t).pre s₀) include hp -theorem inF {i : Nat} (hi : i < 256) {s : State} (hrd : s.rd = s₀.rd) (hwr : s.wr = s₀.wr) : - InRegions (s.rd ++ s.wr) (coeffAddr (s₀.gpr .rdi) i) 4 ∧ InRegions s.wr (coeffAddr (s₀.gpr .rdi) i) 4 := by - rw [hrd, hwr, hp.1, hp.2.1] - exact ⟨⟨_, by simp, coeff_contains _ hi⟩, ⟨_, by simp, coeff_contains _ hi⟩⟩ - -theorem inG {i : Nat} (hi : i < 256) {s : State} (hrd : s.rd = s₀.rd) (hwr : s.wr = s₀.wr) : - InRegions (s.rd ++ s.wr) (coeffAddr (s₀.gpr .rsi) i) 4 := by - rw [hrd, hwr, hp.1, hp.2.1] - exact ⟨_, by simp, coeff_contains _ hi⟩ - -/-- `g` is not written. -/ -theorem coeffG {m : Mem} (hf : Frame [pR (s₀.gpr .rdi)] s₀.mem m) {i : Nat} (hi : i < 256) : - coeffAt m (s₀.gpr .rsi) i = coeffAt s₀.mem (s₀.gpr .rsi) i := - coeffAt_frame hf (by simpa using hp.2.2.1.symm) hi - -omit hp in -theorem inv_step {v : Nat → BitVec 32} {i : Nat} (hi : i < 256) {s s' : State} (hI : Inv s₀ v i s) - (hm : s'.mem = s.mem.writeW (s.gpr .rdi) (v i)) (hdi : s'.gpr .rdi = s.gpr .rdi + 4) - (hsi : s'.gpr .rsi = s.gpr .rsi + 4) (hrd : s'.rd = s.rd) (hwr : s'.wr = s.wr) : - Inv s₀ v (i + 1) s' where - rdi := by rw [hdi, hI.rdi]; exact ptr_step _ i 4 - rsi := by rw [hsi, hI.rsi]; exact ptr_step _ i 4 - rd := hrd.trans hI.rd - wr := hwr.trans hI.wr - frame := by - rw [hm, hI.rdi] - exact hI.frame.writeW (List.mem_singleton_self _) _ (coeff_contains _ hi) - coeff k hk := by - rw [hm, hI.rdi, ← coeffAddr, coeffAt_writeW _ _ hk hi, hI.coeff k hk] - by_cases e : i = k - · subst e; simp - · have : (k < i + 1) = (k < i) := propext (by omega) - simp only [e, this, ↓reduceIte] - -omit hp in -/-- The loop, from the prologue, with a body that stores `v i` to coefficient `i`. -/ -theorem loop_ok {body : List Instr} {v : Nat → BitVec 32} - (hbody : ∀ i < 256, ∀ s, Inv s₀ v i s → WP isa (.block body) s fun s' => - (s'.mem = s.mem.writeW (s.gpr .rdi) (v i) ∧ s'.gpr .rdi = s.gpr .rdi + 4 ∧ - s'.gpr .rsi = s.gpr .rsi + 4 ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ - s'.zf = some (s.gpr .rcx - 1 == 0)) ∧ Keep [.rax, .rdx, .rdi, .rsi, .rcx] s s') : - WP isa (.seq (.block [.mov32 .rcx (.imm 256)]) (.loop (.block body) .ne)) s₀ (Inv s₀ v 256) := by - refine WP.seq ?_ - xrund - refine wp_countdown (cnt := .rcx) (N := 256) (by decide) (by decide) (Inv s₀ v) (fun i hi s hI _ => ?_) - (fun _ h => h) ?_ (by simp [VG.Proof.MlKem.X86_64.setReg_gpr]) - · refine WP.mono (hbody i hi s hI) fun s' ⟨⟨hm, hdi, hsi, hcx, hz⟩, hk⟩ => ⟨?_, hcx, hz⟩ - exact inv_step hi hI hm hdi hsi hk.2.1 hk.2.2 - · refine ⟨by simp [VG.Proof.MlKem.X86_64.setReg_gpr], by simp [VG.Proof.MlKem.X86_64.setReg_gpr], rfl, rfl, - Frame.refl _ _, fun k _ => ?_⟩ - simp [VG.Proof.MlKem.X86_64.setReg_mem] - -/-- The coefficients the loop reads. -/ -theorem reads {v : Nat → BitVec 32} {i : Nat} (hi : i < 256) {s : State} (hI : Inv s₀ v i s) : - s.mem.readW (s.gpr .rdi) 32 = coeffAt s₀.mem (s₀.gpr .rdi) i ∧ - s.mem.readW (s.gpr .rsi) 32 = coeffAt s₀.mem (s₀.gpr .rsi) i ∧ - InRegions (s.rd ++ s.wr) (s.gpr .rdi) 4 ∧ InRegions (s.rd ++ s.wr) (s.gpr .rsi) 4 ∧ - InRegions s.wr (s.gpr .rdi) 4 := by - have hf := inF hp hi hI.rd hI.wr - refine ⟨?_, ?_, ?_, ?_, ?_⟩ - · rw [hI.rdi, ← coeffAddr, ← coeffAt_eq, hI.coeff i hi]; simp only [Nat.lt_irrefl, ↓reduceIte] - · rw [hI.rsi, ← coeffAddr, ← coeffAt_eq, coeffG hp hI.frame hi] - · rw [hI.rdi]; exact hf.1 - · rw [hI.rsi]; exact inG hp hi hI.rd hI.wr - · rw [hI.rdi]; exact hf.2 - -omit hp in -/-- The result, with the values of `add` or `sub`. -/ -theorem result {v : Nat → BitVec 32} {r : Poly} {s : State} (hI : Inv s₀ v 256 s) - (hv : ∀ i < 256, (v i).toNat = (r[i]!).val) : PolyIs s.mem (s₀.gpr .rdi) r := - polyIs_of_toNat fun i hi => by rw [hI.coeff i hi]; simp only [hi, ↓reduceIte]; exact hv i hi +/-- An iteration, which stores `F` of the vectors of `f` and `g`, whose +doublewords are `L` of theirs. -/ +theorem step {op : XBinOp} {fix : List Instr} {F : BitVec 128 → BitVec 128 → BitVec 128} + {L : BitVec 32 → BitVec 32 → BitVec 32} + (hF : ∀ (s : State), s.xmm .xmm15 = qV → + WP isa (.block (xb op .xmm0 .xmm1 :: fix)) s fun s' => + s'.xmm .xmm0 = F (s.xmm .xmm0) (s.xmm .xmm1) ∧ s'.gpr = s.gpr ∧ s'.mem = s.mem ∧ s'.rd = s.rd ∧ + s'.wr = s.wr ∧ s'.xmm .xmm15 = qV) + (hL : ∀ x y : BitVec 128, ∀ e < 4, dword (F x y) e = L (dword x e) (dword y e)) + {i : Nat} (hi : i < 64) {s : State} + (hI : Inv s₀ (fun k => L (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)) i s) : + WP isa (.block (([.movdquLoad .xmm0 (at_ .rdi 0), .movdquLoad .xmm1 (at_ .rsi 0)] : List Instr) ++ + ((xb op .xmm0 .xmm1 :: fix) ++ (accTail ++ ([.alu .sub .rcx (.imm 1)] : List Instr))))) s fun s' => + Inv s₀ (fun k => L (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)) (i + 1) s' ∧ + s'.gpr .rcx = s.gpr .rcx - 1 ∧ s'.zf = some (s.gpr .rcx - 1 == 0) := by + have j0 : 4 * i + 4 ≤ 256 := by omega + have hw : pR (s₀.gpr .rdi) ∈ s.wr := by rw [hI.wr, hp.2.1]; simp + have hr : pR (s₀.gpr .rsi) ∈ s.rd ++ s.wr := by rw [hI.rd, hI.wr, hp.1]; simp + have e1 : s.gpr .rdi = coeffAddr (s₀.gpr .rdi) (4 * i) := by rw [hI.rdi]; congr 2; omega + have e2 : s.gpr .rsi = coeffAddr (s₀.gpr .rsi) (4 * i) := by rw [hI.rsi]; congr 2; omega + refine WP.mono (accBody_ok hF s hI.q (by rw [e1]; exact f_in (List.mem_append_right _ hw) j0) + (by rw [e2]; exact ⟨_, hr, pR_contains _ j0⟩) (by rw [e1]; exact f_in hw j0)) + fun s' ⟨hm, hdi, hsi, hcx, hz, hrd, hwr, hq, _⟩ => ⟨?_, hcx, hz⟩ + refine ⟨by rw [hdi, hI.rdi]; exact ptr_step _ i 16, by rw [hsi, hI.rsi]; exact ptr_step _ i 16, + hrd.trans hI.rd, hwr.trans hI.wr, hq, ?_, fun k hk => ?_⟩ + · rw [hm, e1]; exact hI.frame.writeW (List.mem_singleton_self _) _ (pR_contains _ j0) + · rw [hm, e1, coeffAt_write128 _ _ j0 _ hk] + split + · rename_i h + rw [ifp (show k < 4 * (i + 1) by omega), hL _ _ _ (by omega), e2, dword_readW _ _ (by omega), + dword_readW _ _ (by omega), coeffAddr_add, coeffAddr_add, ← coeffAt_eq, ← coeffAt_eq, + show 4 * i + (k - 4 * i) = k by omega, hI.coeff k hk, ifn (by omega), + coeffAt_frame hI.frame (by simpa using hp.2.2.1.symm) (by rw [n_eq]; exact hk)] + · rename_i h + rw [hI.coeff k hk] + by_cases h' : k < 4 * i + · rw [ifp h', ifp (by omega)] + · rw [ifn h', ifn (by omega)] /-- The whole function, from its precondition. -/ -theorem fn_ok {body : List Instr} {v : Nat → BitVec 32} - (hbody : ∀ i < 256, ∀ s, Inv s₀ v i s → WP isa (.block body) s fun s' => - (s'.mem = s.mem.writeW (s.gpr .rdi) (v i) ∧ s'.gpr .rdi = s.gpr .rdi + 4 ∧ - s'.gpr .rsi = s.gpr .rsi + 4 ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ - s'.zf = some (s.gpr .rcx - 1 == 0)) ∧ Keep [.rax, .rdx, .rdi, .rsi, .rcx] s s') - (hv : ∀ i < 256, (v i).toNat = ((t (polyAt s₀.mem (s₀.gpr .rdi)) (polyAt s₀.mem (s₀.gpr .rsi)))[i]!).val) - (hc : writesOnly [.rax, .rdx, .rdi, .rsi, .rcx] - (.seq (.block [.mov32 .rcx (.imm 256)]) (.loop (.block body) .ne)) = true) - (hm : Code.allInstrs (fun i => !loadsMxcsr i) - (.seq (.block [.mov32 .rcx (.imm 256)]) (.loop (.block body) .ne) : Prog isa) = true) : - ∃ tr s', Exec isa (.seq (.block [.mov32 .rcx (.imm 256)]) (.loop (.block body) .ne)) s₀ tr s' ∧ - abiPreserved s₀ s' ∧ (accK t).post s₀ s' := by - obtain ⟨tr, s', he, hI, hk⟩ := WP.keep _ (loop_ok hbody) hc +theorem fn_ok {op : XBinOp} {fix : List Instr} {F : BitVec 128 → BitVec 128 → BitVec 128} + {L : BitVec 32 → BitVec 32 → BitVec 32} + (hF : ∀ (s : State), s.xmm .xmm15 = qV → + WP isa (.block (xb op .xmm0 .xmm1 :: fix)) s fun s' => + s'.xmm .xmm0 = F (s.xmm .xmm0) (s.xmm .xmm1) ∧ s'.gpr = s.gpr ∧ s'.mem = s.mem ∧ s'.rd = s.rd ∧ + s'.wr = s.wr ∧ s'.xmm .xmm15 = qV) + (hL : ∀ x y : BitVec 128, ∀ e < 4, dword (F x y) e = L (dword x e) (dword y e)) + (hv : ∀ k < 256, (L (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)).toNat = + ((t (polyAt s₀.mem (s₀.gpr .rdi)) (polyAt s₀.mem (s₀.gpr .rsi)))[k]!).val) + (hc : writesOnly [.rax, .rdi, .rsi, .rcx] (.seq (.block qPro) (VG.Impl.MlKem.X86_64.rcxLoop 64 + (([.movdquLoad .xmm0 (at_ .rdi 0), .movdquLoad .xmm1 (at_ .rsi 0)] : List Instr) ++ (xb op .xmm0 .xmm1 :: fix) ++ + accTail))) = true) + (hm : Code.allInstrs (fun i => !loadsMxcsr i) (.seq (.block qPro) (VG.Impl.MlKem.X86_64.rcxLoop 64 + (([.movdquLoad .xmm0 (at_ .rdi 0), .movdquLoad .xmm1 (at_ .rsi 0)] : List Instr) ++ (xb op .xmm0 .xmm1 :: fix) ++ + accTail)) : Prog isa) = true) : + ∃ tr s', Exec isa (.seq (.block qPro) (VG.Impl.MlKem.X86_64.rcxLoop 64 + (([.movdquLoad .xmm0 (at_ .rdi 0), .movdquLoad .xmm1 (at_ .rsi 0)] : List Instr) ++ (xb op .xmm0 .xmm1 :: fix) ++ + accTail))) s₀ tr s' ∧ abiPreserved s₀ s' ∧ (accK t).post s₀ s' := by + have hw : pR (s₀.gpr .rdi) ∈ s₀.wr := by rw [hp.2.1]; simp + have hW : WP isa (.seq (.block qPro) (VG.Impl.MlKem.X86_64.rcxLoop 64 + (([.movdquLoad .xmm0 (at_ .rdi 0), .movdquLoad .xmm1 (at_ .rsi 0)] : List Instr) ++ (xb op .xmm0 .xmm1 :: fix) ++ + accTail))) s₀ + (Inv s₀ (fun k => L (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)) 64) := by + refine WP.seq (WP.mono (Q := fun (w : State) => w.xmm .xmm15 = qV ∧ Keep [.rax] s₀ w ∧ w.mem = s₀.mem) + (by + simp only [qPro] + vrund + refine ⟨fun r hr => ?_, rfl, rfl⟩ + simp only [List.mem_singleton] at hr + simp only [RegUpd.gpr_setReg, RegUpd.gpr_setXmm, hr, ite_false]) fun w ⟨hq, k1, m1⟩ => ?_) + refine wp_rcxLoop (N := 64) (by decide) (by decide) _ (fun u o _ => ⟨?_, ?_, by rw [o.keep.2.1, k1.2.1], + by rw [o.keep.2.2, k1.2.2], by rw [o.xmm]; exact hq, by rw [o.mem, m1]; exact Frame.refl _ _, + fun k _ => by rw [o.mem, m1, ifn (by omega)]⟩) fun i hi u hI => ?_ + · rw [o.keep.gpr (by decide), k1.gpr (by decide), Nat.mul_zero, add_ofNat_zero] + · rw [o.keep.gpr (by decide), k1.gpr (by decide), Nat.mul_zero, add_ofNat_zero] + · rw [show [Instr.movdquLoad .xmm0 (at_ .rdi 0), .movdquLoad .xmm1 (at_ .rsi 0)] ++ + (xb op .xmm0 .xmm1 :: fix) ++ accTail ++ [.alu .sub .rcx (.imm 1)] = + [.movdquLoad .xmm0 (at_ .rdi 0), .movdquLoad .xmm1 (at_ .rsi 0)] ++ + ((xb op .xmm0 .xmm1 :: fix) ++ (accTail ++ ([.alu .sub .rcx (.imm 1)] : List Instr))) by + simp only [List.append_assoc]] + exact step hp hF hL hi hI + obtain ⟨tr, s', he, hI, hk⟩ := WP.keep _ hW hc refine ⟨tr, s', he, abiPreserved_of_exec hm he (gprPreserved_of hk (by decide) hI.frame ?_), - result hI hv⟩ - simpa using hp.2.2.2.1 + polyIs_of_toNat fun k hk => ?_⟩ + · simpa using hp.2.2.2.1 + · rw [n_eq] at hk + rw [hI.coeff k hk, ifp (by omega)] + exact hv k hk end @@ -153,27 +198,25 @@ end AddSub theorem add_correct (s : State) (hs : (accK Spec.MlDsa.add).pre s) : ∃ t s', Exec isa Impl.MlDsa.X86_64.Arith.add s t s' ∧ abiPreserved s s' ∧ (accK Spec.MlDsa.add).post s s' := - AddSub.fn_ok hs (v := fun i => csubD (coeffAt s.mem (s.gpr .rdi) i + coeffAt s.mem (s.gpr .rsi) i)) - (fun i hi s' hI => by - obtain ⟨e1, e2, h1, h2, h3⟩ := AddSub.reads hs hi hI - have := addBody_ok s' h1 h2 h3 - rwa [e1, e2] at this) - (fun i hi => by - rw [add_get _ _ hi] - exact csubD_add_val (polyAt_val hs.2.2.2.2.2.1 hi).symm (polyAt_val hs.2.2.2.2.2.2 hi).symm) - (by decide) (by decide) + AddSub.fn_ok hs (op := .paddd) (fix := vcsub .xmm0 .xmm2) (L := fun a b => csubL (a + b)) addFix_ok + (fun x y e he => dword_addV x y he) + (fun k hk => by + have hk' : k < n := by rw [n_eq]; exact hk + rw [add_get _ _ hk', addD_toNat (by rw [← polyAt_val hs.2.2.2.2.2.1 hk']; exact val_lt _) + (by rw [← polyAt_val hs.2.2.2.2.2.2 hk']; exact val_lt _), + ← polyAt_val hs.2.2.2.2.2.1 hk', ← polyAt_val hs.2.2.2.2.2.2 hk', val_add]) + (by decide +kernel) (by decide +kernel) theorem sub_correct (s : State) (hs : (accK Spec.MlDsa.sub).pre s) : ∃ t s', Exec isa Impl.MlDsa.X86_64.Arith.sub s t s' ∧ abiPreserved s s' ∧ (accK Spec.MlDsa.sub).post s s' := - AddSub.fn_ok hs (v := fun i => csubD (coeffAt s.mem (s.gpr .rdi) i + qImm - coeffAt s.mem (s.gpr .rsi) i)) - (fun i hi s' hI => by - obtain ⟨e1, e2, h1, h2, h3⟩ := AddSub.reads hs hi hI - have := subBody_ok s' h1 h2 h3 - rwa [e1, e2] at this) - (fun i hi => by - rw [sub_get _ _ hi] - exact csubD_sub_val (polyAt_val hs.2.2.2.2.2.1 hi).symm (polyAt_val hs.2.2.2.2.2.2 hi).symm) - (by decide) (by decide) + AddSub.fn_ok hs (op := .psubd) (fix := vcadd .xmm0 .xmm2) (L := fun a b => caddL (a - b)) subFix_ok + (fun x y e he => dword_subV x y he) + (fun k hk => by + have hk' : k < n := by rw [n_eq]; exact hk + rw [sub_get _ _ hk', subD_toNat (by rw [← polyAt_val hs.2.2.2.2.2.1 hk']; exact val_lt _) + (by rw [← polyAt_val hs.2.2.2.2.2.2 hk']; exact val_lt _), + ← polyAt_val hs.2.2.2.2.2.1 hk', ← polyAt_val hs.2.2.2.2.2.2 hk', val_sub]) + (by decide +kernel) (by decide +kernel) /-- The pointers and `rsp` are public. -/ def accτ : X86_64.Taint.T := X86_64.Taint.ofRegs [.rdi, .rsi, .rsp] diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Mul.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Mul.lean index 3939bb475..7cdf5d3b0 100644 --- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Mul.lean +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Mul.lean @@ -1,205 +1,353 @@ import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Mul -import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Basic -import VerifiedGarbage.Proof.Framework.X86_64.Abi +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Mxcsr +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.AddSub +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Ntt +import VerifiedGarbage.Proof.Framework.X86_64.Mxcsr /-! # ML-DSA on x86-64: `vg_mldsa_multiply_ntt` and `vg_mldsa_multiply_add_ntt` -Untrusted: everything here is checked by Lean. Both functions are one loop -over the coefficients (`Mul.fn_ok`), whose body stores the reduced product -(`mulBody_ok`), or the reduced sum of the product and the coefficient of -`h` (`mulAddBody_ok`). +Untrusted: everything here is checked by Lean. A doubleword of `mulV x y` +is the product of those of `x` and `y` (`mul_lane`): `mont` of `mont` of +their product by `2⁶⁴ mod q` (`mont_mont_R2`). The loop stores four of them +at a time to the first 252 coefficients of `h` (`Mul.step`, `Mul.loop_ok`), +inside `withMxcsr` through the last 8 bytes of `h`; the last four are +computed from the coefficients of `h` loaded before (`Mul.last`), and stored +after MXCSR is loaded back (`Mul.fn_ok`). -/ namespace VG.Proof.MlDsa.X86_64.Arith open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Arith open VG.Proof.MlDsa.Arith -open VG.Proof.MlKem.X86_64 (Keep WP.keep writesOnly gprPreserved_of wp_counted ifp ifn ptr_step - toNat_setWidth64) -open VG.Spec.MlDsa (q n Poly Zq coeffAt polyAt Reduced PolyIs) - -/-! ## One coefficient -/ - -/-- The product of two words, as `mul` leaves it. -/ -abbrev prod32 (a b : BitVec 32) : BitVec 64 := - BitVec.ofNat 64 ((BitVec.setWidth 64 a).toNat * (BitVec.setWidth 64 b).toNat) - -theorem mulHead_ok (s : State) (h1 : InRegions (s.rd ++ s.wr) (s.gpr .rsi) 4) - (h2 : InRegions (s.rd ++ s.wr) (s.gpr .r8) 4) : - WP isa (.block mulHead) s fun s' => - (s'.gpr .rax = prod32 (s.mem.readW (s.gpr .rsi) 32) (s.mem.readW (s.gpr .r8) 32) ∧ s'.mem = s.mem) ∧ - Keep [.rax, .rdx, .r9] s s' := by - refine WP.keep _ ?_ (by decide) - unfold mulHead - xrund [h1, h2] - -theorem mulAddHead_ok (s : State) (h1 : InRegions (s.rd ++ s.wr) (s.gpr .rsi) 4) - (h2 : InRegions (s.rd ++ s.wr) (s.gpr .r8) 4) (h3 : InRegions (s.rd ++ s.wr) (s.gpr .rdi) 4) : - WP isa (.block mulAddHead) s fun s' => - (s'.gpr .rax = prod32 (s.mem.readW (s.gpr .rsi) 32) (s.mem.readW (s.gpr .r8) 32) + - BitVec.setWidth 64 (s.mem.readW (s.gpr .rdi) 32) ∧ s'.mem = s.mem) ∧ Keep [.rax, .rdx, .r9] s s' := by - refine WP.keep _ ?_ (by decide) - unfold mulAddHead mulHead - xrund [h1, h2, h3, List.cons_append, List.nil_append] - -theorem mulTail_ok (s : State) (hw : InRegions s.wr (s.gpr .rdi) 4) : - WP isa (.block (([.store32 (at_ .rdi 0) .r10] : List Instr) ++ step3)) s fun s' => - (s'.mem = s.mem.writeW (s.gpr .rdi) (BitVec.setWidth 32 (s.gpr .r10)) ∧ s'.gpr .rdi = s.gpr .rdi + 4 ∧ - s'.gpr .rsi = s.gpr .rsi + 4 ∧ s'.gpr .r8 = s.gpr .r8 + 4 ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ - s'.zf = some (s.gpr .rcx - 1 == 0)) ∧ Keep [.rdi, .rsi, .r8, .rcx] s s' := by - refine WP.keep _ ?_ (by decide) - unfold step3 - xrund [hw, List.cons_append, List.nil_append] - -/-- A body: a head that leaves `x` in `rax`, `reduce`, and the store. -/ -theorem body_ok {head : List Instr} {x : BitVec 64} (s : State) - (hh : WP isa (.block head) s fun s' => (s'.gpr .rax = x ∧ s'.mem = s.mem) ∧ Keep [.rax, .rdx, .r9] s s') - (hw : InRegions s.wr (s.gpr .rdi) 4) : - WP isa (.block (head ++ reduce ++ ([.store32 (at_ .rdi 0) .r10] : List Instr) ++ step3)) s fun s' => - (s'.mem = s.mem.writeW (s.gpr .rdi) (BitVec.setWidth 32 (redD x)) ∧ s'.gpr .rdi = s.gpr .rdi + 4 ∧ - s'.gpr .rsi = s.gpr .rsi + 4 ∧ s'.gpr .r8 = s.gpr .r8 + 4 ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ - s'.zf = some (s.gpr .rcx - 1 == 0)) ∧ - Keep [.rax, .rdx, .r9, .r10, .r11, .rdi, .rsi, .r8, .rcx] s s' := by - rw [List.append_assoc, List.append_assoc, WP.block_append_iff] - refine WP.mono hh fun s1 ⟨⟨ha, hm1⟩, k1⟩ => ?_ - rw [WP.block_append_iff] - refine WP.mono (reduce_ok s1) fun s2 ⟨⟨hr, hm2⟩, k2⟩ => ?_ - have k12 := k1.trans k2 - refine WP.mono (mulTail_ok s2 (by rw [k12.2.2, k12.gpr (by decide)]; exact hw)) - fun s3 ⟨⟨hm3, hdi, hsi, h8, hcx, hz⟩, k3⟩ => ⟨?_, (k12.trans k3).mono (by decide)⟩ - rw [hm3, hdi, hsi, h8, hcx, hz, hr, hm2, ha, hm1, k12.gpr (r := .rdi) (by decide), k12.gpr (r := .rsi) (by decide), - k12.gpr (r := .r8) (by decide), k12.gpr (r := .rcx) (by decide)] - exact ⟨rfl, rfl, rfl, rfl, rfl, rfl⟩ - -theorem prod32_val {a b : BitVec 32} {x y : Zq} (ha : a.toNat = x.val) (hb : b.toNat = y.val) : - (BitVec.setWidth 32 (redD (prod32 a b))).toNat = (x * y).val := by - have e : (prod32 a b).toNat = x.val * y.val := by - rw [prod32, BitVec.toNat_ofNat, toNat_setWidth64, toNat_setWidth64, ha, hb] - exact Nat.mod_eq_of_lt (Nat.lt_of_lt_of_le (mul_lt_q2 x.isLt y.isLt) (by decide)) - rw [redD32_toNat, e, val_mul] - -theorem prod32_add_val {a b c : BitVec 32} {x y z : Zq} (ha : a.toNat = x.val) (hb : b.toNat = y.val) - (hc : c.toNat = z.val) : - (BitVec.setWidth 32 (redD (prod32 a b + BitVec.setWidth 64 c))).toNat = (z + x * y).val := by - have e : (prod32 a b + BitVec.setWidth 64 c).toNat = x.val * y.val + z.val := by - have := mul_lt_q2 x.isLt y.isLt - have := val_lt z - rw [BitVec.toNat_add, prod32, BitVec.toNat_ofNat, toNat_setWidth64, toNat_setWidth64, toNat_setWidth64, ha, - hb, hc] - omega - rw [redD32_toNat, e, val_add', val_mul, Nat.add_comm, Nat.add_mod_mod] +open VG.Proof.MlKem.X86_64 (Keep XOnly WP.keep writesOnly gprPreserved_of ifn xmm_setXmm wp_rcxLoop + add_ofNat_zero) +open VG.Impl.MlKem.X86_64 (xb xmov rcxLoop) +open VG.Spec.MlDsa (q n Poly Zq coeffAt polyAt Reduced) + +/-! ## Four products -/ + +/-- `2⁶⁴ mod q` in each doubleword, as the prologue leaves it in `xmm11`. -/ +def r2V : BitVec 128 := shufDwords ((0 : BitVec 64) ++ BitVec.setWidth 64 2365951#32) 0 + +theorem dword_r2V {i : Nat} (hi : i < 4) : dword r2V i = 2365951#32 := by + rcases cases4 hi with rfl | rfl | rfl | rfl <;> decide + +/-- What `mulCore` leaves in `xmm3` of the vectors of `f` and `g`. -/ +def mulV (x y : BitVec 128) : BitVec 128 := csubV (montV (montV x y (shufDwords y 0xF5)) r2V r2V) + +/-- What `mulAddCore` leaves in `xmm3`, with the vector of `h`. -/ +def mulAddV (x y z : BitVec 128) : BitVec 128 := csubV (XBinOp.eval .paddd (mulV x y) z) + +theorem mul_lane {x y : BitVec 128} {a b : Nat → Zq} (hx : DLanes x a) (hy : DLanes y b) {i : Nat} + (hi : i < 4) : (dword (mulV x y) i).toNat = (a i * b i).val := by + have hzo : ZOdd y (shufDwords y 0xF5) := fun j hj => by + rw [dword_shufDwords _ _ (by omega)] + rcases (by omega : j = 0 ∨ j = 1) with rfl | rfl <;> rfl + have hb1 : ∀ i < 4, (dword x i).toNat * (dword y i).toNat < q * 2 ^ 32 := fun i hi => by + rw [hx i hi, hy i hi] + exact Nat.lt_of_lt_of_le (Nat.mul_lt_mul_of_lt_of_le (val_lt (a i)) (Nat.le_of_lt (val_lt (b i))) + (by decide)) (by decide) + have m1 := fun i (hi : i < 4) => dword_montV hzo hb1 hi + have hb2 : ∀ i < 4, (dword (montV x y (shufDwords y 0xF5)) i).toNat * (dword r2V i).toNat < q * 2 ^ 32 := + fun i hi => by + rw [m1 i hi, dword_r2V hi] + have := mont_lt (hb1 i hi) + rw [q_eq] at this ⊢ + exact Nat.lt_of_lt_of_le (Nat.mul_lt_mul_of_lt_of_le this (Nat.le_refl 2365951) (by decide)) (by decide) + have hzo2 : ZOdd r2V r2V := fun j hj => by rw [dword_r2V (by omega), dword_r2V (by omega)] + rw [mulV, dword_csubV _ hi, csubL_toNat (by rw [dword_montV hzo2 hb2 hi]; exact mont_lt (hb2 i hi)), + dword_montV hzo2 hb2 hi, condSub_mont (hb2 i hi), m1 i hi, dword_r2V hi, + show (2365951#32).toNat = 2 ^ 64 % q from rfl, mont_mont_R2, hx i hi, hy i hi, val_mul] + +theorem mulAdd_lane {x y z : BitVec 128} {a b c : Nat → Zq} (hx : DLanes x a) (hy : DLanes y b) + (hz : DLanes z c) {i : Nat} (hi : i < 4) : (dword (mulAddV x y z) i).toNat = (c i + a i * b i).val := by + have hm := mul_lane hx hy hi + rw [mulAddV, dword_csubV _ hi, dword_paddd _ _ hi, addD_toNat (by rw [hm]; exact val_lt _) + (by rw [hz i hi]; exact val_lt _), hm, hz i hi, Nat.add_comm, ← val_add] + +theorem mulCore_ok {s : State} (hc : VConsts s) (h11 : s.xmm .xmm11 = r2V) : + WP isa (.block mulCore) s fun s' => + s'.xmm .xmm3 = mulV (s.xmm .xmm3) (s.xmm .xmm13) ∧ XOnly [.xmm12, .xmm3, .xmm2, .xmm4] s s' := by + simp only [mulCore, vmont, vredc, vcsub, vcadd, xmov, xb, List.cons_append, List.nil_append] + vrun [eval_movdqa] + rw [hc.q, hc.qinv, h11] + exact ⟨rfl, by xonly⟩ + +theorem mulAddCore_ok {s : State} (hc : VConsts s) (h11 : s.xmm .xmm11 = r2V) : + WP isa (.block mulAddCore) s fun s' => + s'.xmm .xmm3 = mulAddV (s.xmm .xmm3) (s.xmm .xmm13) (s.xmm .xmm5) ∧ + XOnly [.xmm12, .xmm3, .xmm2, .xmm4] s s' := by + simp only [mulAddCore, mulCore, vmont, vredc, vcsub, vcadd, xmov, xb, List.cons_append, List.nil_append] + vrun [eval_movdqa] + rw [hc.q, hc.qinv, h11] + exact ⟨rfl, by xonly⟩ + +theorem mulPro_ok (s : State) : + WP isa (.block mulPro) s fun s' => VConsts s' ∧ s'.xmm .xmm11 = r2V ∧ s'.xmm .xmm6 = s.xmm .xmm6 ∧ + Keep [.rax] s s' ∧ s'.mem = s.mem := by + simp only [mulPro, vconsts, List.cons_append, List.nil_append] + vrund + refine ⟨⟨?_, ?_⟩, ⟨fun r hr => ?_, rfl, rfl⟩⟩ + · simp only [RegUpd.xmm_setReg, xmm_setXmm, ite_true, ite_false, reduceCtorEq]; decide + · simp only [RegUpd.xmm_setReg, xmm_setXmm, ite_true, ite_false, reduceCtorEq]; decide + · simp only [List.mem_singleton] at hr + simp only [RegUpd.gpr_setReg, RegUpd.gpr_setXmm, hr, ite_false] /-! ## The loop -/ namespace Mul -/-- After `i` coefficients, each one `v k`. -/ -structure Inv (s₀ : State) (v : Nat → BitVec 32) (i : Nat) (s : State) : Prop where - rdi : s.gpr .rdi = s₀.gpr .rdi + BitVec.ofNat 64 (4 * i) - rsi : s.gpr .rsi = s₀.gpr .rsi + BitVec.ofNat 64 (4 * i) - r8 : s.gpr .r8 = s₀.gpr .rdx + BitVec.ofNat 64 (4 * i) +theorem ptr16 (p : Addr) (i : Nat) : coeffAddr p (4 * i) + 16 = p + BitVec.ofNat 64 (16 * (i + 1)) := by + rw [coeffAddr, BitVec.add_assoc, show (16 : BitVec 64) = BitVec.ofNat 64 16 from rfl, ← BitVec.ofNat_add] + congr 2; omega + +/-- After `i` iterations: the first `4i` coefficients of `h` are `R`'s, the +others as they were in `s₀`. -/ +structure Inv (h f g : Addr) (s₀ : State) (R : Poly) (i : Nat) (s : State) : Prop where + rdi : s.gpr .rdi = h + BitVec.ofNat 64 (16 * i) + rsi : s.gpr .rsi = f + BitVec.ofNat 64 (16 * i) + rdx : s.gpr .rdx = g + BitVec.ofNat 64 (16 * i) rd : s.rd = s₀.rd wr : s.wr = s₀.wr - frame : Frame [pR (s₀.gpr .rdi)] s₀.mem s.mem - coeff : ∀ k < 256, coeffAt s.mem (s₀.gpr .rdi) k = if k < i then v k else coeffAt s₀.mem (s₀.gpr .rdi) k + c : VConsts s + r2 : s.xmm .xmm11 = r2V + x6 : s.xmm .xmm6 = s₀.xmm .xmm6 + frame : Frame [pR h] s₀.mem s.mem + done : ∀ k < 4 * i, (coeffAt s.mem h k).toNat = (R[k]!).val + rest : ∀ k < 256, 4 * i ≤ k → coeffAt s.mem h k = coeffAt s₀.mem h k section -variable {t : Poly → Poly → Poly → Poly} {r : Mem → Addr → Prop} {s₀ : State} (hp : (mulK t r).pre s₀) -include hp - -theorem inR {p : Addr} (hp' : p = s₀.gpr .rsi ∨ p = s₀.gpr .rdx ∨ p = s₀.gpr .rdi) {k : Nat} (hk : k < 256) - {s : State} (hrd : s.rd = s₀.rd) (hwr : s.wr = s₀.wr) : InRegions (s.rd ++ s.wr) (coeffAddr p k) 4 := by - rw [hrd, hwr, hp.1, hp.2.1] - rcases hp' with rfl | rfl | rfl - · exact ⟨_, by simp, coeff_contains _ hk⟩ - · exact ⟨_, by simp, coeff_contains _ hk⟩ - · exact ⟨_, by simp, coeff_contains _ hk⟩ - -theorem inW {k : Nat} (hk : k < 256) {s : State} (hwr : s.wr = s₀.wr) : - InRegions s.wr (coeffAddr (s₀.gpr .rdi) k) 4 := by - rw [hwr, hp.2.1]; exact ⟨_, by simp, coeff_contains _ hk⟩ - -/-- `f` and `g` are not written. -/ -theorem coeffF {m : Mem} (hf : Frame [pR (s₀.gpr .rdi)] s₀.mem m) {k : Nat} (hk : k < 256) : - coeffAt m (s₀.gpr .rsi) k = coeffAt s₀.mem (s₀.gpr .rsi) k := - coeffAt_frame hf (by simpa using hp.2.2.1.symm) hk - -theorem coeffG {m : Mem} (hf : Frame [pR (s₀.gpr .rdi)] s₀.mem m) {k : Nat} (hk : k < 256) : - coeffAt m (s₀.gpr .rdx) k = coeffAt s₀.mem (s₀.gpr .rdx) k := - coeffAt_frame hf (by simpa using hp.2.2.2.1.symm) hk - -omit hp in -theorem inv_step {v : Nat → BitVec 32} {i : Nat} (hi : i < 256) {s s' : State} (hI : Inv s₀ v i s) - (hm : s'.mem = s.mem.writeW (s.gpr .rdi) (v i)) (hdi : s'.gpr .rdi = s.gpr .rdi + 4) - (hsi : s'.gpr .rsi = s.gpr .rsi + 4) (h8 : s'.gpr .r8 = s.gpr .r8 + 4) (hrd : s'.rd = s.rd) - (hwr : s'.wr = s.wr) : Inv s₀ v (i + 1) s' where - rdi := by rw [hdi, hI.rdi]; exact ptr_step _ i 4 - rsi := by rw [hsi, hI.rsi]; exact ptr_step _ i 4 - r8 := by rw [h8, hI.r8]; exact ptr_step _ i 4 - rd := hrd.trans hI.rd - wr := hwr.trans hI.wr - frame := by - rw [hm, hI.rdi] - exact hI.frame.writeW (List.mem_singleton_self _) _ (coeff_contains _ hi) - coeff k hk := by - rw [hm, hI.rdi, ← coeffAddr, coeffAt_writeW _ _ hk hi, hI.coeff k hk] - by_cases e : i = k - · subst e; simp - · have : (k < i + 1) = (k < i) := propext (by omega) - simp only [e, this, ↓reduceIte] - -/-- What a body reads at coefficient `i`, and where. -/ -theorem reads {v : Nat → BitVec 32} {i : Nat} (hi : i < 256) {s : State} (hI : Inv s₀ v i s) : - s.mem.readW (s.gpr .rsi) 32 = coeffAt s₀.mem (s₀.gpr .rsi) i ∧ - s.mem.readW (s.gpr .r8) 32 = coeffAt s₀.mem (s₀.gpr .rdx) i ∧ - s.mem.readW (s.gpr .rdi) 32 = coeffAt s₀.mem (s₀.gpr .rdi) i ∧ - InRegions (s.rd ++ s.wr) (s.gpr .rsi) 4 ∧ InRegions (s.rd ++ s.wr) (s.gpr .r8) 4 ∧ - InRegions (s.rd ++ s.wr) (s.gpr .rdi) 4 ∧ InRegions s.wr (s.gpr .rdi) 4 := by - refine ⟨?_, ?_, ?_, ?_, ?_, ?_, ?_⟩ - · rw [hI.rsi, ← coeffAddr, ← coeffAt_eq, coeffF hp hI.frame hi] - · rw [hI.r8, ← coeffAddr, ← coeffAt_eq, coeffG hp hI.frame hi] - · rw [hI.rdi, ← coeffAddr, ← coeffAt_eq, hI.coeff i hi]; simp only [Nat.lt_irrefl, ↓reduceIte] - · rw [hI.rsi]; exact inR hp (.inl rfl) hi hI.rd hI.wr - · rw [hI.r8]; exact inR hp (.inr (.inl rfl)) hi hI.rd hI.wr - · rw [hI.rdi]; exact inR hp (.inr (.inr rfl)) hi hI.rd hI.wr - · rw [hI.rdi]; exact inW hp hi hI.wr - -/-- The whole function, from its precondition, with a body that stores -`v i` to coefficient `i`. -/ -theorem fn_ok {body : List Instr} {v : Nat → BitVec 32} - (hbody : ∀ i < 256, ∀ s, Inv s₀ v i s → WP isa (.block body) s fun s' => - (s'.mem = s.mem.writeW (s.gpr .rdi) (v i) ∧ s'.gpr .rdi = s.gpr .rdi + 4 ∧ - s'.gpr .rsi = s.gpr .rsi + 4 ∧ s'.gpr .r8 = s.gpr .r8 + 4 ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ - s'.zf = some (s.gpr .rcx - 1 == 0)) ∧ Keep [.rax, .rdx, .r9, .r10, .r11, .rdi, .rsi, .r8, .rcx] s s') - (hv : ∀ i < 256, (v i).toNat = - ((t (polyAt s₀.mem (s₀.gpr .rdi)) (polyAt s₀.mem (s₀.gpr .rsi)) (polyAt s₀.mem (s₀.gpr .rdx)))[i]!).val) - (hc : writesOnly [.rax, .rdx, .r9, .r10, .r11, .rdi, .rsi, .r8, .rcx] - (.seq (.block [.mov .r8 (.reg .rdx)]) (.seq (.block [.mov32 .rcx (.imm 256)]) (.loop (.block body) .ne))) = - true) - (hm : Code.allInstrs (fun i => !loadsMxcsr i) - (.seq (.block [.mov .r8 (.reg .rdx)]) (.seq (.block [.mov32 .rcx (.imm 256)]) (.loop (.block body) .ne)) : - Prog isa) = true) : - ∃ tr s', Exec isa (.seq (.block [.mov .r8 (.reg .rdx)]) - (.seq (.block [.mov32 .rcx (.imm 256)]) (.loop (.block body) .ne))) s₀ tr s' ∧ - abiPreserved s₀ s' ∧ (mulK t r).post s₀ s' := by - obtain ⟨tr, s', he, hI, hk⟩ := WP.keep [.rax, .rdx, .r9, .r10, .r11, .rdi, .rsi, .r8, .rcx] - (WP.seq (WP.mono (WP.keep [.r8] (Q := fun s => s.mem = s₀.mem ∧ s.gpr .r8 = s₀.gpr .rdx) (by xrund) - (by decide)) fun s1 ⟨⟨hm1, h81⟩, k1⟩ => - wp_counted (s₀ := s1) (N := 256) (v := 256) rfl (by decide) (Inv s₀ v) - (fun s hm hk' => ⟨by rw [hk'.gpr (by decide), k1.gpr (by decide)]; simp, - by rw [hk'.gpr (by decide), k1.gpr (by decide)]; simp, - by rw [hk'.gpr (by decide), h81]; simp, hk'.2.1.trans k1.2.1, hk'.2.2.trans k1.2.2, - by rw [hm, hm1]; exact Frame.refl _ _, fun k _ => by simp [hm, hm1]⟩) - fun i hi s hI => WP.mono (hbody i hi s hI) fun s' ⟨⟨hm, hdi, hsi, h8, hcx, hz⟩, hk⟩ => - ⟨inv_step hi hI hm hdi hsi h8 hk.2.1 hk.2.2, hcx, hz⟩)) hc - refine ⟨tr, s', he, abiPreserved_of_exec hm he (gprPreserved_of hk (by decide) hI.frame - (by simpa using hp.2.2.2.2.1)), polyIs_of_toNat fun i hi => ?_⟩ - rw [hI.coeff i hi, ifp hi] - exact hv i hi +variable {h f g : Addr} {s₀ : State} (hwh : pR h ∈ s₀.wr) (hrf : pR f ∈ s₀.rd ++ s₀.wr) + (hrg : pR g ∈ s₀.rd ++ s₀.wr) (hdf : (pR h).Disjoint (pR f)) (hdg : (pR h).Disjoint (pR g)) + {F G : Poly} (hF : ∀ k < 256, (coeffAt s₀.mem f k).toNat = (F[k]!).val) + (hG : ∀ k < 256, (coeffAt s₀.mem g k).toNat = (G[k]!).val) + {core : List Instr} {Fv : BitVec 128 → BitVec 128 → BitVec 128 → BitVec 128} + (hcore : ∀ s : State, VConsts s → s.xmm .xmm11 = r2V → WP isa (.block core) s fun s' => + s'.xmm .xmm3 = Fv (s.xmm .xmm3) (s.xmm .xmm13) (s.xmm .xmm5) ∧ XOnly [.xmm12, .xmm3, .xmm2, .xmm4] s s') + {R : Poly} {H : Nat → BitVec 32} (hH : ∀ k < 252, coeffAt s₀.mem h k = H k) + (hlane : ∀ i < 64, ∀ x y z : BitVec 128, (∀ e < 4, (dword x e).toNat = (F[4 * i + e]!).val) → + (∀ e < 4, (dword y e).toNat = (G[4 * i + e]!).val) → (∀ e < 4, dword z e = H (4 * i + e)) → + ∀ e < 4, (dword (Fv x y z) e).toNat = (R[4 * i + e]!).val) +include hwh hrf hrg hdf hdg hF hG hcore hH hlane + +theorem step {i : Nat} (hi : i < 63) {s : State} (hI : Inv h f g s₀ R i s) : + WP isa (.block (mulLoads ++ core ++ mulTail ++ ([.alu .sub .rcx (.imm 1)] : List Instr))) s fun s' => + Inv h f g s₀ R (i + 1) s' ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ s'.zf = some (s.gpr .rcx - 1 == 0) := by + have j0 : 4 * i + 4 ≤ 256 := by omega + have e1 : s.gpr .rdi = coeffAddr h (4 * i) := by rw [hI.rdi]; congr 2; omega + have e2 : s.gpr .rsi = coeffAddr f (4 * i) := by rw [hI.rsi]; congr 2; omega + have e3 : s.gpr .rdx = coeffAddr g (4 * i) := by rw [hI.rdx]; congr 2; omega + have rf : InRegions (s.rd ++ s.wr) (coeffAddr f (4 * i)) 16 := by + rw [hI.rd, hI.wr]; exact ⟨_, hrf, pR_contains f j0⟩ + have rg : InRegions (s.rd ++ s.wr) (coeffAddr g (4 * i)) 16 := by + rw [hI.rd, hI.wr]; exact ⟨_, hrg, pR_contains g j0⟩ + have rh : InRegions (s.rd ++ s.wr) (coeffAddr h (4 * i)) 16 := by + rw [hI.rd, hI.wr]; exact f_in (List.mem_append_right _ hwh) j0 + have wh : InRegions s.wr (coeffAddr h (4 * i)) 16 := by rw [hI.wr]; exact f_in hwh j0 + have mf : ∀ k < 256, coeffAt s.mem f k = coeffAt s₀.mem f k := fun k hk => + coeffAt_frame hI.frame (by simpa using hdf.symm) (by rw [n_eq]; exact hk) + have mg : ∀ k < 256, coeffAt s.mem g k = coeffAt s₀.mem g k := fun k hk => + coeffAt_frame hI.frame (by simpa using hdg.symm) (by rw [n_eq]; exact hk) + rw [show mulLoads ++ core ++ mulTail ++ [.alu .sub .rcx (.imm 1)] = + mulLoads ++ (core ++ (mulTail ++ ([.alu .sub .rcx (.imm 1)] : List Instr))) by simp [List.append_assoc], + WP.block_append_iff] + simp only [mulLoads] + vrund [e1, e2, e3, rf, rg, rh] + rw [WP.block_append_iff] + refine WP.mono (hcore _ (((hI.c.setXmm (by decide) (by decide) _).setXmm (by decide) (by decide) _).setXmm + (by decide) (by decide) _) (by simp only [xmm_setXmm, reduceCtorEq, ite_false]; exact hI.r2)) + fun s2 ⟨h3, o2⟩ => ?_ + have c2 := xonly_vconsts o2 (((hI.c.setXmm (by decide) (by decide) _).setXmm (by decide) (by decide) _).setXmm + (by decide) (by decide) _) (by decide) (by decide) + have g2 : s2.gpr = s.gpr := o2.gpr + have m2 : s2.mem = s.mem := o2.mem + have r2 : s2.rd = s.rd := o2.rd + have w2 : s2.wr = s.wr := o2.wr + have x11 : s2.xmm .xmm11 = r2V := by rw [o2.xmm _ (by decide)]; simp only [xmm_setXmm, reduceCtorEq, ite_false]; exact hI.r2 + simp only [xmm_setXmm, ite_true, reduceCtorEq, ite_false] at h3 + generalize hV : Fv (s.mem.readW (coeffAddr f (4 * i)) 128) (s.mem.readW (coeffAddr g (4 * i)) 128) + (s.mem.readW (coeffAddr h (4 * i)) 128) = V at h3 + simp only [mulTail] + vrund [g2, m2, r2, w2, e1, e2, e3, wh, h3] + have lx : ∀ e < 4, (dword (s.mem.readW (coeffAddr f (4 * i)) 128) e).toNat = (F[4 * i + e]!).val := + fun e he => by rw [dword_readW _ _ he, coeffAddr_add, ← coeffAt_eq, mf _ (by omega), hF _ (by omega)] + have ly : ∀ e < 4, (dword (s.mem.readW (coeffAddr g (4 * i)) 128) e).toNat = (G[4 * i + e]!).val := + fun e he => by rw [dword_readW _ _ he, coeffAddr_add, ← coeffAt_eq, mg _ (by omega), hG _ (by omega)] + have lz : ∀ e < 4, dword (s.mem.readW (coeffAddr h (4 * i)) 128) e = H (4 * i + e) := + fun e he => by rw [dword_readW _ _ he, coeffAddr_add, ← coeffAt_eq, hI.rest _ (by omega) (by omega), + hH _ (by omega)] + have hl := hlane i (by omega) _ _ _ lx ly lz + rw [hV] at hl + refine ⟨?_, ?_, ?_, ?_, ?_, ⟨?_, ?_⟩, ?_, ?_, ?_, ?_, ?_⟩ <;> + simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, RegUpd.mem_setReg, RegUpd.mem_setFlags, + RegUpd.rd_setReg, RegUpd.rd_setFlags, RegUpd.wr_setReg, RegUpd.wr_setFlags, RegUpd.xmm_setReg, + RegUpd.xmm_setFlags, ite_true, ite_false, reduceCtorEq] + · exact ptr16 h i + · exact ptr16 f i + · exact ptr16 g i + · exact hI.rd + · exact hI.wr + · exact c2.q + · exact c2.qinv + · exact x11 + · rw [o2.xmm _ (by decide)]; simp only [xmm_setXmm, reduceCtorEq, ite_false]; exact hI.x6 + · exact hI.frame.writeW (List.mem_singleton_self _) _ (pR_contains h j0) + · intro k hk + rw [coeffAt_write128 _ _ j0 _ (by omega)] + split + · rw [hl _ (by omega), show 4 * i + (k - 4 * i) = k by omega] + · exact hI.done k (by omega) + · intro k hk hk' + rw [coeffAt_write128 _ _ j0 _ hk, ifn (by omega)] + exact hI.rest k hk (by omega) + +theorem loop_ok (hdi : s₀.gpr .rdi = h) (hsi : s₀.gpr .rsi = f) (hdx : s₀.gpr .rdx = g) (hc : VConsts s₀) + (h11 : s₀.xmm .xmm11 = r2V) : WP isa (rcxLoop 63 (mulLoads ++ core ++ mulTail)) s₀ (Inv h f g s₀ R 63) := + wp_rcxLoop (N := 63) (by decide) (by decide) _ (fun u o _ => + ⟨by rw [o.keep.gpr (by decide), hdi, Nat.mul_zero, add_ofNat_zero], + by rw [o.keep.gpr (by decide), hsi, Nat.mul_zero, add_ofNat_zero], + by rw [o.keep.gpr (by decide), hdx, Nat.mul_zero, add_ofNat_zero], o.keep.2.1, o.keep.2.2, + ⟨by rw [o.xmm]; exact hc.q, by rw [o.xmm]; exact hc.qinv⟩, by rw [o.xmm]; exact h11, by rw [o.xmm], + by rw [o.mem]; exact Frame.refl _ _, fun k hk => absurd hk (by omega), fun k _ _ => by rw [o.mem]⟩) + fun i hi u hI => step hwh hrf hrg hdf hdg hF hG hcore hH hlane hi hI + + +omit hwh hH in +/-- The last four coefficients, with those of `h` in `xmm6`. -/ +theorem last {s : State} (hI : Inv h f g s₀ R 63 s) (hz : ∀ e < 4, dword (s₀.xmm .xmm6) e = H (252 + e)) : + WP isa (.block (mulLast core)) s fun s' => + (∀ e < 4, (dword (s'.xmm .xmm3) e).toNat = (R[252 + e]!).val) ∧ s'.gpr = s.gpr ∧ s'.mem = s.mem := by + have j0 : 4 * 63 + 4 ≤ 256 := by decide + have e2 : s.gpr .rsi = coeffAddr f (4 * 63) := hI.rsi + have e3 : s.gpr .rdx = coeffAddr g (4 * 63) := hI.rdx + have rf : InRegions (s.rd ++ s.wr) (coeffAddr f (4 * 63)) 16 := by + rw [hI.rd, hI.wr]; exact ⟨_, hrf, pR_contains f j0⟩ + have rg : InRegions (s.rd ++ s.wr) (coeffAddr g (4 * 63)) 16 := by + rw [hI.rd, hI.wr]; exact ⟨_, hrg, pR_contains g j0⟩ + have mf : ∀ k < 256, coeffAt s.mem f k = coeffAt s₀.mem f k := fun k hk => + coeffAt_frame hI.frame (by simpa using hdf.symm) (by rw [n_eq]; exact hk) + have mg : ∀ k < 256, coeffAt s.mem g k = coeffAt s₀.mem g k := fun k hk => + coeffAt_frame hI.frame (by simpa using hdg.symm) (by rw [n_eq]; exact hk) + rw [mulLast, WP.block_append_iff] + vrund [e2, e3, rf, rg, eval_movdqa] + refine WP.mono (hcore _ (((hI.c.setXmm (by decide) (by decide) _).setXmm (by decide) (by decide) _).setXmm + (by decide) (by decide) _) (by simp only [xmm_setXmm, reduceCtorEq, ite_false]; exact hI.r2)) + fun s2 ⟨h3, o2⟩ => ⟨fun e he => ?_, o2.gpr, o2.mem⟩ + simp only [xmm_setXmm, ite_true, reduceCtorEq, ite_false, hI.x6] at h3 + rw [h3] + refine hlane 63 (by decide) _ _ _ (fun e he => ?_) (fun e he => ?_) hz e he + · rw [dword_readW _ _ he, coeffAddr_add, ← coeffAt_eq, mf _ (by omega), hF _ (by omega)] + · rw [dword_readW _ _ he, coeffAddr_add, ← coeffAt_eq, mg _ (by omega), hG _ (by omega)] end + +theorem coeffAt_mxH {m m' : Mem} {p : Addr} (hf : Frame [mxH p] m m') {k : Nat} (hk : k < 254) : + coeffAt m' p k = coeffAt m p k := + hf.readW (r := ⟨coeffAddr p k, 4⟩) (Region.contains_self _ _) (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact Offset.disjoint p (by omega) (by omega) (by omega)) (by decide) + +/-- The whole function, from its precondition, with a `core` whose lanes are +those of `t`. -/ +theorem fn_ok {t : Poly → Poly → Poly → Poly} {hPre : Mem → Addr → Prop} {σ : State} (hp : (mulK t hPre).pre σ) + {core : List Instr} {Fv : BitVec 128 → BitVec 128 → BitVec 128 → BitVec 128} + (hcore : ∀ s : State, VConsts s → s.xmm .xmm11 = r2V → WP isa (.block core) s fun s' => + s'.xmm .xmm3 = Fv (s.xmm .xmm3) (s.xmm .xmm13) (s.xmm .xmm5) ∧ XOnly [.xmm12, .xmm3, .xmm2, .xmm4] s s') + (hlane : ∀ i < 64, ∀ x y z : BitVec 128, + (∀ e < 4, (dword x e).toNat = ((polyAt σ.mem (σ.gpr .rsi))[4 * i + e]!).val) → + (∀ e < 4, (dword y e).toNat = ((polyAt σ.mem (σ.gpr .rdx))[4 * i + e]!).val) → + (∀ e < 4, dword z e = coeffAt σ.mem (σ.gpr .rdi) (4 * i + e)) → + ∀ e < 4, (dword (Fv x y z) e).toNat = ((t (polyAt σ.mem (σ.gpr .rdi)) (polyAt σ.mem (σ.gpr .rsi)) + (polyAt σ.mem (σ.gpr .rdx)))[4 * i + e]!).val) + (hk : writesOnly [.rax, .rdi, .rsi, .rdx, .rcx] + (.seq (.block mulPro) (.seq (rcxLoop 63 (mulLoads ++ core ++ mulTail)) (.block (mulLast core)))) = true) : + WP isa (mulFn core) σ fun s' => Keep [.r8, .rax, .r11, .rax, .rdi, .rsi, .rdx, .rcx] σ s' ∧ + Frame [pR (σ.gpr .rdi)] σ.mem s'.mem ∧ (mulK t hPre).post σ s' := by + obtain ⟨hrd, hwr, hdf, hdg, -, -, -, -, redf, redg⟩ := hp + generalize eh : σ.gpr .rdi = h at * + generalize ef : σ.gpr .rsi = f at * + generalize eg : σ.gpr .rdx = g at * + let R := t (polyAt σ.mem h) (polyAt σ.mem f) (polyAt σ.mem g) + have hwh : pR h ∈ σ.wr := by rw [hwr]; exact List.mem_singleton_self _ + have r6 : InRegions (σ.rd ++ σ.wr) (coeffAddr h 252) 16 := + ⟨_, List.mem_append_right _ hwh, Offset.contains_base h (by omega) (by omega)⟩ + simp only [mulFn] + refine WP.seq (WP.mono (Q := fun (s1 : State) => s1.gpr .r8 = h ∧ + s1.xmm .xmm6 = σ.mem.readW (coeffAddr h 252) 128 ∧ Keep [.r8] σ s1 ∧ s1.mem = σ.mem) (by + vrund [eh, r6] + exact ⟨fun r hr => by + simp only [List.mem_singleton] at hr + simp only [RegUpd.gpr_setXmm, RegUpd.gpr_setReg, hr, ite_false], rfl, rfl⟩) fun s1 ⟨h8, h6, k1, m1⟩ => ?_) + have hw1 : pR h ∈ s1.wr := by rw [k1.2.2]; exact hwh + refine WP.seq (WP.mono (withMxcsrH_ok (r := .r8) ⟨by decide, by decide⟩ [.rax, .rdi, .rsi, .rdx, .rcx] + ⟨by decide, by decide⟩ h8 hw1 hk (Q := fun (s3 : State) => Keep [.rax, .r11, .rax, .rdi, .rsi, .rdx, .rcx] s1 s3 ∧ + s3.gpr .rdi = coeffAddr h 252 ∧ Frame [pR h] σ.mem s3.mem ∧ + (∀ k < 252, (coeffAt s3.mem h k).toNat = (R[k]!).val) ∧ + ∀ e < 4, (dword (s3.xmm .xmm3) e).toNat = (R[252 + e]!).val) + fun s2 k2 f2 x2 => ?_) fun s4 ⟨s3, ⟨kk, hdi, fr, dn, ln⟩, f4, k4, x4⟩ => ?_) + · have fσ2 : Frame [mxH h] σ.mem s2.mem := by rw [← m1]; exact f2 + refine WP.mono (WP.keep [.rax, .rdi, .rsi, .rdx, .rcx] (WP.seq (WP.mono (mulPro_ok s2) + fun w ⟨cw, xw, x6, kw, mw⟩ => ?_) (Q := fun (s3 : State) => s3.gpr .rdi = coeffAddr h 252 ∧ + Frame [pR h] σ.mem s3.mem ∧ (∀ k < 252, (coeffAt s3.mem h k).toNat = (R[k]!).val) ∧ + ∀ e < 4, (dword (s3.xmm .xmm3) e).toNat = (R[252 + e]!).val)) hk) + fun s3 ⟨q3, kk⟩ => ⟨k2.trans kk, q3⟩ + have gw : ∀ r, r ≠ .rax → r ≠ .r11 → r ≠ .r8 → w.gpr r = σ.gpr r := fun r h1 h2 h3 => by + rw [kw.gpr (by simpa using h1), k2.gpr (by simp [h1, h2]), k1.gpr (by simpa using h3)] + have rw' : w.rd = σ.rd ∧ w.wr = σ.wr := + ⟨kw.2.1.trans (k2.2.1.trans k1.2.1), kw.2.2.trans (k2.2.2.trans k1.2.2)⟩ + have fσw : Frame [mxH h] σ.mem w.mem := by rw [mw]; exact fσ2 + have fσw' : Frame [pR h] σ.mem w.mem := + Frame.sub fσw fun r hr => ⟨pR h, List.mem_singleton_self _, by + simp only [List.mem_singleton] at hr; subst hr; exact mxH_sub h⟩ + refine WP.seq (WP.mono (loop_ok (s₀ := w) (h := h) (f := f) (g := g) (R := R) (F := polyAt σ.mem f) + (G := polyAt σ.mem g) (H := coeffAt σ.mem h) + (by rw [rw'.2]; exact hwh) (by rw [rw'.1, hrd]; simp) (by rw [rw'.1, hrd]; simp) hdf hdg + (fun k hk => by + rw [coeffAt_frame fσw' (by simpa using hdf.symm) (by rw [n_eq]; exact hk), + polyAt_val redf (by rw [n_eq]; exact hk)]) + (fun k hk => by + rw [coeffAt_frame fσw' (by simpa using hdg.symm) (by rw [n_eq]; exact hk), + polyAt_val redg (by rw [n_eq]; exact hk)]) + hcore (fun k hk => coeffAt_mxH fσw (by omega)) hlane + (by rw [gw _ (by decide) (by decide) (by decide), eh]) (by rw [gw _ (by decide) (by decide) (by decide), ef]) + (by rw [gw _ (by decide) (by decide) (by decide), eg]) cw xw) fun s hI => ?_) + refine WP.mono (last (s₀ := w) (h := h) (f := f) (g := g) (R := R) (F := polyAt σ.mem f) + (G := polyAt σ.mem g) (H := coeffAt σ.mem h) + (by rw [rw'.1, hrd]; simp) (by rw [rw'.1, hrd]; simp) hdf hdg + (fun k hk => by + rw [coeffAt_frame fσw' (by simpa using hdf.symm) (by rw [n_eq]; exact hk), + polyAt_val redf (by rw [n_eq]; exact hk)]) + (fun k hk => by + rw [coeffAt_frame fσw' (by simpa using hdg.symm) (by rw [n_eq]; exact hk), + polyAt_val redg (by rw [n_eq]; exact hk)]) + hcore hlane hI (fun e he => by + rw [x6, x2, h6, dword_readW _ _ he, coeffAddr_add, ← coeffAt_eq])) fun s3 ⟨ln, g3, m3⟩ => ?_ + refine ⟨by rw [g3, hI.rdi], ?_, fun k hk => by rw [m3]; exact hI.done k (by omega), ln⟩ + rw [m3] + exact Frame.trans fσw' hI.frame + · have k14 := (k1.trans kk).trans k4 + have hdi4 : s4.gpr .rdi = coeffAddr h 252 := by rw [k4.gpr (by simp), hdi] + have wh4 : InRegions s4.wr (coeffAddr h 252) 16 := by + rw [k14.2.2]; exact ⟨_, hwh, Offset.contains_base h (by omega) (by omega)⟩ + vrund [hdi4, wh4] + have f4' : Frame [pR h] s3.mem s4.mem := Frame.sub f4 fun r hr => ⟨pR h, List.mem_singleton_self _, by + simp only [List.mem_singleton] at hr; subst hr; exact mxH_sub h⟩ + refine ⟨k14.mono (by simp), (fr.trans f4').writeW (List.mem_singleton_self _) _ + (Offset.contains_base h (by omega) (by omega)), ?_⟩ + dsimp only [mulK] + rw [eh, ef, eg] + refine polyIs_of_toNat fun k hk => ?_ + rw [n_eq] at hk + rw [coeffAt_write128 _ _ (j := 252) (by decide) _ hk] + split + · rw [x4] + have := ln (k - 252) (by omega) + rwa [show 252 + (k - 252) = k by omega] at this + · rw [coeffAt_mxH f4 (by omega)] + exact dn k (by omega) + end Mul /-! ## The functions -/ @@ -211,31 +359,27 @@ abbrev mulK' : Contract isa := mulK (fun _ f g => Spec.MlDsa.multiplyNTT f g) fu abbrev mulAddK : Contract isa := mulK (fun h f g => Spec.MlDsa.add h (Spec.MlDsa.multiplyNTT f g)) Reduced theorem mul_correct (s : State) (hs : mulK'.pre s) : - ∃ t s', Exec isa Impl.MlDsa.X86_64.Arith.mul s t s' ∧ abiPreserved s s' ∧ mulK'.post s s' := - Mul.fn_ok hs (v := fun i => BitVec.setWidth 32 (redD (prod32 (coeffAt s.mem (s.gpr .rsi) i) - (coeffAt s.mem (s.gpr .rdx) i)))) - (fun i hi s' hI => by - obtain ⟨e1, e2, _, h1, h2, _, hw⟩ := Mul.reads hs hi hI - have := body_ok s' (mulHead_ok s' h1 h2) hw - rwa [e1, e2] at this) - (fun i hi => by - rw [mul_get _ _ hi] - exact prod32_val (polyAt_val hs.2.2.2.2.2.2.2.2.1 hi).symm (polyAt_val hs.2.2.2.2.2.2.2.2.2 hi).symm) - (by decide) (by decide) + ∃ t s', Exec isa Impl.MlDsa.X86_64.Arith.mul s t s' ∧ abiPreserved s s' ∧ mulK'.post s s' := by + obtain ⟨t, s', he, hk, hf, hq⟩ := Mul.fn_ok hs (core := mulCore) (Fv := fun x y _ => mulV x y) + (fun s hc h11 => mulCore_ok hc h11) + (fun i hi x y z hx hy _ e he => by + rw [mul_get _ _ (by rw [n_eq]; omega)] + exact mul_lane hx hy he) + (by decide +kernel) + exact ⟨t, s', he, abiPreserved_of_ctl (by decide +kernel) he (gprPreserved_of hk (by decide) hf + (by simpa using hs.2.2.2.2.1)), hq⟩ theorem mulAdd_correct (s : State) (hs : mulAddK.pre s) : - ∃ t s', Exec isa Impl.MlDsa.X86_64.Arith.mulAdd s t s' ∧ abiPreserved s s' ∧ mulAddK.post s s' := - Mul.fn_ok hs (v := fun i => BitVec.setWidth 32 (redD (prod32 (coeffAt s.mem (s.gpr .rsi) i) - (coeffAt s.mem (s.gpr .rdx) i) + BitVec.setWidth 64 (coeffAt s.mem (s.gpr .rdi) i)))) - (fun i hi s' hI => by - obtain ⟨e1, e2, e3, h1, h2, h3, hw⟩ := Mul.reads hs hi hI - have := body_ok s' (mulAddHead_ok s' h1 h2 h3) hw - rwa [e1, e2, e3] at this) - (fun i hi => by - rw [add_get _ _ hi, mul_get _ _ hi] - exact prod32_add_val (polyAt_val hs.2.2.2.2.2.2.2.2.1 hi).symm - (polyAt_val hs.2.2.2.2.2.2.2.2.2 hi).symm (polyAt_val hs.2.2.2.2.2.2.2.1 hi).symm) - (by decide) (by decide) + ∃ t s', Exec isa Impl.MlDsa.X86_64.Arith.mulAdd s t s' ∧ abiPreserved s s' ∧ mulAddK.post s s' := by + obtain ⟨t, s', he, hk, hf, hq⟩ := Mul.fn_ok hs (core := mulAddCore) (Fv := mulAddV) + (fun s hc h11 => mulAddCore_ok hc h11) + (fun i hi x y z hx hy hz e he => by + rw [add_get _ _ (by rw [n_eq]; omega), mul_get _ _ (by rw [n_eq]; omega)] + exact mulAdd_lane hx hy (c := fun e => (polyAt s.mem (s.gpr .rdi))[4 * i + e]!) + (fun e he => by rw [hz e he, polyAt_val hs.2.2.2.2.2.2.2.1 (by rw [n_eq]; omega)]) he) + (by decide +kernel) + exact ⟨t, s', he, abiPreserved_of_ctl (by decide +kernel) he (gprPreserved_of hk (by decide) hf + (by simpa using hs.2.2.2.2.1)), hq⟩ /-- The pointers and `rsp` are public. -/ def mulτ : X86_64.Taint.T := X86_64.Taint.ofRegs [.rdi, .rsi, .rdx, .rsp] diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Mxcsr.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Mxcsr.lean new file mode 100644 index 000000000..d87f9e044 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Mxcsr.lean @@ -0,0 +1,73 @@ +import VerifiedGarbage.Proof.MlKem.X86_64.VMxcsr +import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Basic + +/-! +# ML-DSA on x86-64: MXCSR through the end of a polynomial + +Untrusted: everything here is checked by Lean. `withMxcsr r 1016 c` +(ML-KEM's, see `Impl/MlKem/X86_64/Vec.lean`) through the last 8 bytes `mxH` +of a writable polynomial at `r` (`withMxcsrH_ok`), as ML-KEM's +`withMxcsr_ok` through `scratch + 768`. +-/ + +namespace VG.Proof.MlDsa.X86_64.Arith + +open VG VG.X86_64 VG.Impl.MlKem.X86_64 +open VG.Proof.MlKem.X86_64 (Keep WP.keep writesOnly ldmxcsr_ok) + +/-- The last 8 bytes of the polynomial at `p`. -/ +abbrev mxH (p : Addr) : Region := ⟨p + BitVec.ofNat 64 1016, 8⟩ + +theorem mxH_in {p : Addr} {rs : List Region} (hw : pR p ∈ rs) (d : Nat) (hd : 1016 ≤ d ∧ d ≤ 1020) : + InRegions rs (p + BitVec.ofNat 64 d) 4 := + ⟨_, hw, Offset.contains_base p (by omega) (by omega)⟩ + +theorem mxH_sub (p : Addr) : Region.Sub (mxH p) (pR p) := Offset.sub_base p (by decide) + +/-- `withMxcsr` through `mxH`: `c` runs from `s` but for `rax`, `r11` and +those bytes, and nothing more than they and MXCSR change after it. -/ +theorem withMxcsrH_ok {c : Prog isa} {r : Reg} (hr : r ≠ .r11 ∧ r ≠ .rax) (rs : List Reg) + (hrs : r ∉ rs ∧ Reg.r11 ∉ rs) {p : Addr} {s : State} {Q : State → Prop} + (hsi : s.gpr r = p) (hw : pR p ∈ s.wr) (hk : writesOnly rs c = true) + (hc : ∀ s1, Keep [.rax, .r11] s s1 → Frame [mxH p] s.mem s1.mem → s1.xmm = s.xmm → WP isa c s1 Q) : + WP isa (withMxcsr r 1016 c) s fun s' => + ∃ s2, Q s2 ∧ Frame [mxH p] s2.mem s'.mem ∧ Keep [] s2 s' ∧ s'.xmm = s2.xmm := by + have h0 := mxH_in hw 1016 (by decide) + have h0' := mxH_in (List.mem_append_right s.rd hw) 1016 (by decide) + have h4 := mxH_in hw 1020 (by decide) + simp only [withMxcsr] + refine WP.seq (WP.mono (Q := fun (s1 : State) => s1.gpr .r11 = (s.mxcsr &&& 0xFFFF).setWidth 64 ∧ Keep [.r11] s s1 ∧ + Frame [mxH p] s.mem s1.mem ∧ s1.xmm = s.xmm) (by + vrunm [hsi, h0, h0', Mem.readW_writeW_self32, hr.1] + refine ⟨by rw [BitVec.setWidth_setWidth_of_le _ (by decide), BitVec.setWidth_eq], + ⟨fun r hr => ?_, rfl, rfl⟩, (Frame.refl _ _).writeW (List.mem_singleton_self _) _ + (Offset.contains p (by decide) (by decide) (by decide))⟩ + simp only [List.mem_singleton] at hr + simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, hr, ite_false]) fun s1 ⟨h11, k1, f1, x1⟩ => ?_) + have hsi1 : s1.gpr r = p := by rw [k1.gpr (by simpa using hr.1), hsi] + have h4' : InRegions s1.wr (p + BitVec.ofNat 64 (1016 + 4)) 4 := by rw [k1.2.2]; exact h4 + have h4'' : InRegions (s1.rd ++ s1.wr) (p + BitVec.ofNat 64 (1016 + 4)) 4 := + let ⟨r, hr, hc⟩ := h4'; ⟨r, List.mem_append_right _ hr, hc⟩ + refine WP.seq (WP.seq (WP.mono (Q := fun (s2 : State) => Keep [.rax] s1 s2 ∧ Frame [mxH p] s1.mem s2.mem ∧ + s2.xmm = s1.xmm) + (by + vrunm [hsi1, h4', h4'', Mem.readW_writeW_self32, hr.2] + refine ⟨⟨fun r hr => ?_, rfl, rfl⟩, (Frame.refl _ _).writeW (List.mem_singleton_self _) _ + (Offset.contains p (by decide) (by decide) (by decide))⟩ + simp only [List.mem_singleton] at hr + simp only [RegUpd.gpr_setReg, hr, ite_false]) fun s2 ⟨k2, f2, x2⟩ => ?_)) + refine WP.seq (WP.mono (WP.keep _ (hc s2 ((k1.trans k2).mono (by simp)) (f1.trans f2) (x2.trans x1)) hk) + fun s3 ⟨hq, k3⟩ => ?_) + have k23 := k2.trans k3 + have hsi3 : s3.gpr r = p := by rw [k23.gpr (by simp [hr.2, hrs.1]), hsi1] + have h113 : s3.gpr .r11 = BitVec.setWidth 64 (s.mxcsr &&& 65535) := by rw [k23.gpr (by simp [hrs.2]), h11] + have h03 : InRegions s3.wr (p + BitVec.ofNat 64 1016) 4 := by rw [k23.2.2, k1.2.2]; exact h0 + have h03' : InRegions (s3.rd ++ s3.wr) (p + BitVec.ofNat 64 1016) 4 := + let ⟨r, hr, hc⟩ := h03; ⟨r, List.mem_append_right _ hr, hc⟩ + refine WP.mono (Q := fun s4 => s4 = s3) (by vrunm) fun s4 h4 => ?_ + subst h4 + vrunm [hsi3, h113, h03, h03', Mem.readW_writeW_self32, ldmxcsr_ok] + exact ⟨_, hq, (Frame.refl _ _).writeW (List.mem_singleton_self _) _ + (Offset.contains p (by decide) (by decide) (by decide)), ⟨fun _ _ => rfl, rfl, rfl⟩, rfl⟩ + +end VG.Proof.MlDsa.X86_64.Arith diff --git a/src/asm/x86_64/mldsa.rs b/src/asm/x86_64/mldsa.rs index 59b6ad833..01e804f36 100644 --- a/src/asm/x86_64/mldsa.rs +++ b/src/asm/x86_64/mldsa.rs @@ -1424,6 +1424,8 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_inv_ntt(f: *mut [u32; 256], scratc /// /// Contract: `VG.Spec.MlDsa.mulContract`. Constant time: only the pointers may affect timing, not the data. /// +/// The function computes on four coefficients at a time in SSE2 registers. It sets MXCSR to `0x1FBF` around its multiplications (Intel's mitigation of MXCSR-configuration-dependent timing), through the last 8 bytes of `h`, which it stores last, and loads the caller's MXCSR back before returning. +/// /// # Safety /// /// * `h` must be valid for reads and writes of 1024 bytes. @@ -1436,29 +1438,106 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_inv_ntt(f: *mut [u32; 256], scratc #[unsafe(naked)] pub(crate) unsafe extern "sysv64" fn vg_mldsa_multiply_ntt(h: *mut [u32; 256], f: *const [u32; 256], g: *const [u32; 256]) { core::arch::naked_asm!( - "mov r8, rdx", - "mov ecx, 256", + "mov r8, rdi", + "movdqu xmm6, XMMWORD PTR [rdi+1008]", + "stmxcsr DWORD PTR [r8+1016]", + "mov r11d, DWORD PTR [r8+1016]", + "and r11d, 65535", + "mov eax, 8127", + "mov DWORD PTR [r8+1020], eax", + "ldmxcsr DWORD PTR [r8+1020]", + "lfence", + "mov eax, 8380417", + "movq xmm15, rax", + "pshufd xmm15, xmm15, 0", + "mov eax, -58728449", + "movq xmm14, rax", + "pshufd xmm14, xmm14, 0", + "mov eax, 2365951", + "movq xmm11, rax", + "pshufd xmm11, xmm11, 0", + "mov ecx, 63", "20:", - "mov eax, DWORD PTR [rsi]", - "mov r9d, DWORD PTR [r8]", - "mul r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov DWORD PTR [rdi], r10d", - "add rdi, 4", - "add rsi, 4", - "add r8, 4", + "movdqu xmm3, XMMWORD PTR [rsi]", + "movdqu xmm13, XMMWORD PTR [rdx]", + "movdqu xmm5, XMMWORD PTR [rdi]", + "pshufd xmm12, xmm13, 245", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm11", + "pmuludq xmm4, xmm11", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdi], xmm3", + "add rdi, 16", + "add rsi, 16", + "add rdx, 16", "sub rcx, 1", "jne 20b", + "movdqu xmm3, XMMWORD PTR [rsi]", + "movdqu xmm13, XMMWORD PTR [rdx]", + "movdqa xmm5, xmm6", + "pshufd xmm12, xmm13, 245", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm11", + "pmuludq xmm4, xmm11", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "lfence", + "mov DWORD PTR [r8+1016], r11d", + "ldmxcsr DWORD PTR [r8+1016]", + "movdqu XMMWORD PTR [rdi], xmm3", "ret", ) } @@ -1467,6 +1546,8 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_multiply_ntt(h: *mut [u32; 256], f /// /// Contract: `VG.Spec.MlDsa.mulAddContract`. Constant time: only the pointers may affect timing, not the data. /// +/// The function computes on four coefficients at a time in SSE2 registers. It sets MXCSR to `0x1FBF` around its multiplications (Intel's mitigation of MXCSR-configuration-dependent timing), through the last 8 bytes of `h`, which it stores last, and loads the caller's MXCSR back before returning. +/// /// # Safety /// /// * `h` must be valid for reads and writes of 1024 bytes. @@ -1480,31 +1561,118 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_multiply_ntt(h: *mut [u32; 256], f #[unsafe(naked)] pub(crate) unsafe extern "sysv64" fn vg_mldsa_multiply_add_ntt(h: *mut [u32; 256], f: *const [u32; 256], g: *const [u32; 256]) { core::arch::naked_asm!( - "mov r8, rdx", - "mov ecx, 256", + "mov r8, rdi", + "movdqu xmm6, XMMWORD PTR [rdi+1008]", + "stmxcsr DWORD PTR [r8+1016]", + "mov r11d, DWORD PTR [r8+1016]", + "and r11d, 65535", + "mov eax, 8127", + "mov DWORD PTR [r8+1020], eax", + "ldmxcsr DWORD PTR [r8+1020]", + "lfence", + "mov eax, 8380417", + "movq xmm15, rax", + "pshufd xmm15, xmm15, 0", + "mov eax, -58728449", + "movq xmm14, rax", + "pshufd xmm14, xmm14, 0", + "mov eax, 2365951", + "movq xmm11, rax", + "pshufd xmm11, xmm11, 0", + "mov ecx, 63", "20:", - "mov eax, DWORD PTR [rsi]", - "mov r9d, DWORD PTR [r8]", - "mul r9", - "mov r9d, DWORD PTR [rdi]", - "add rax, r9", - "mov r10, rax", - "movabs r11, 2201172575745", - "mul r11", - "mov rax, rdx", - "mov r11, 8380417", - "mul r11", - "sub r10, rax", - "sub r10d, 8380417", - "sbb r11d, r11d", - "and r11d, 8380417", - "add r10d, r11d", - "mov DWORD PTR [rdi], r10d", - "add rdi, 4", - "add rsi, 4", - "add r8, 4", + "movdqu xmm3, XMMWORD PTR [rsi]", + "movdqu xmm13, XMMWORD PTR [rdx]", + "movdqu xmm5, XMMWORD PTR [rdi]", + "pshufd xmm12, xmm13, 245", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm11", + "pmuludq xmm4, xmm11", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "paddd xmm3, xmm5", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "movdqu XMMWORD PTR [rdi], xmm3", + "add rdi, 16", + "add rsi, 16", + "add rdx, 16", "sub rcx, 1", "jne 20b", + "movdqu xmm3, XMMWORD PTR [rsi]", + "movdqu xmm13, XMMWORD PTR [rdx]", + "movdqa xmm5, xmm6", + "pshufd xmm12, xmm13, 245", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm13", + "pmuludq xmm4, xmm12", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "pshufd xmm4, xmm3, 245", + "pmuludq xmm3, xmm11", + "pmuludq xmm4, xmm11", + "movdqa xmm2, xmm3", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm3, xmm2", + "psrlq xmm3, 32", + "movdqa xmm2, xmm4", + "pmuludq xmm2, xmm14", + "pmuludq xmm2, xmm15", + "paddq xmm4, xmm2", + "por xmm3, xmm4", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "paddd xmm3, xmm5", + "psubd xmm3, xmm15", + "movdqa xmm2, xmm3", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm3, xmm2", + "lfence", + "mov DWORD PTR [r8+1016], r11d", + "ldmxcsr DWORD PTR [r8+1016]", + "movdqu XMMWORD PTR [rdi], xmm3", "ret", ) } @@ -1513,6 +1681,8 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_multiply_add_ntt(h: *mut [u32; 256 /// /// Contract: `VG.Spec.MlDsa.addContract`. Constant time: only the pointers may affect timing, not the data. /// +/// The function computes on four coefficients at a time in SSE2 registers. +/// /// # Safety /// /// * `f` must be valid for reads and writes of 1024 bytes. @@ -1524,17 +1694,22 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_multiply_add_ntt(h: *mut [u32; 256 #[unsafe(naked)] pub(crate) unsafe extern "sysv64" fn vg_mldsa_add(f: *mut [u32; 256], g: *const [u32; 256]) { core::arch::naked_asm!( - "mov ecx, 256", + "mov eax, 8380417", + "movq xmm15, rax", + "pshufd xmm15, xmm15, 0", + "mov ecx, 64", "20:", - "mov eax, DWORD PTR [rdi]", - "add eax, DWORD PTR [rsi]", - "sub eax, 8380417", - "sbb edx, edx", - "and edx, 8380417", - "add eax, edx", - "mov DWORD PTR [rdi], eax", - "add rdi, 4", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdi]", + "movdqu xmm1, XMMWORD PTR [rsi]", + "paddd xmm0, xmm1", + "psubd xmm0, xmm15", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "movdqu XMMWORD PTR [rdi], xmm0", + "add rdi, 16", + "add rsi, 16", "sub rcx, 1", "jne 20b", "ret", @@ -1545,6 +1720,8 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_add(f: *mut [u32; 256], g: *const /// /// Contract: `VG.Spec.MlDsa.subContract`. Constant time: only the pointers may affect timing, not the data. /// +/// The function computes on four coefficients at a time in SSE2 registers. +/// /// # Safety /// /// * `f` must be valid for reads and writes of 1024 bytes. @@ -1556,18 +1733,21 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_add(f: *mut [u32; 256], g: *const #[unsafe(naked)] pub(crate) unsafe extern "sysv64" fn vg_mldsa_sub(f: *mut [u32; 256], g: *const [u32; 256]) { core::arch::naked_asm!( - "mov ecx, 256", + "mov eax, 8380417", + "movq xmm15, rax", + "pshufd xmm15, xmm15, 0", + "mov ecx, 64", "20:", - "mov eax, DWORD PTR [rdi]", - "add eax, 8380417", - "sub eax, DWORD PTR [rsi]", - "sub eax, 8380417", - "sbb edx, edx", - "and edx, 8380417", - "add eax, edx", - "mov DWORD PTR [rdi], eax", - "add rdi, 4", - "add rsi, 4", + "movdqu xmm0, XMMWORD PTR [rdi]", + "movdqu xmm1, XMMWORD PTR [rsi]", + "psubd xmm0, xmm1", + "movdqa xmm2, xmm0", + "psrad xmm2, 31", + "pand xmm2, xmm15", + "paddd xmm0, xmm2", + "movdqu XMMWORD PTR [rdi], xmm0", + "add rdi, 16", + "add rsi, 16", "sub rcx, 1", "jne 20b", "ret",