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",