diff --git a/README.md b/README.md index 8e8048a9b..49aa2153d 100644 --- a/README.md +++ b/README.md @@ -827,7 +827,7 @@ yours to keep: ✅ -✅ SSE2 NTT +✅ SSE2 polynomial arithmetic ✅ SHA extensions @@ -843,7 +843,7 @@ yours to keep: ✅ -✅ SSE2 NTT +✅ SSE2 polynomial arithmetic ✅ SHA extensions @@ -859,7 +859,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",