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/Artifacts/MlDsaKeyGen/X86_64.lean b/lean/VerifiedGarbage/Artifacts/MlDsaKeyGen/X86_64.lean
deleted file mode 100644
index 5760b7b53..000000000
--- a/lean/VerifiedGarbage/Artifacts/MlDsaKeyGen/X86_64.lean
+++ /dev/null
@@ -1,53 +0,0 @@
-import VerifiedGarbage.TCB.X86_64.Target
-import VerifiedGarbage.Proof.MlDsa.X86_64.KeyGen.Inst
-
-/-!
-# ML-DSA (FIPS 204) key generation on x86-64
-
-A registration file (see `TCB/Emit.lean`): the artifacts it lists are
-emitted. **Review note**: `sig` and `doc` are trusted, as they tie the Rust
-caller to the contract; check them against the contract's `pre`/`post`. Each
-artifact is made from its function's `Api` (in `Spec/MlDsa/Contract.lean`,
-reviewed with the contract), and this file adds only notes on the
-implementation. The emitter adds the `# Safety` items that depend on the
-target (`Sig.layoutDoc`), from `stack` and `writeArgs`, which `ofSig` checks
-against the contract.
--/
-
-namespace VG.Artifacts.MlDsaKeyGen.X86_64
-
-/-- Notes on the implementation, the same for every parameter set. -/
-def notes : List String :=
- ["The function saves its caller's callee-saved registers in `scratch`; its calls use the 32 \
- bytes of stack below its return address.",
- "It samples every polynomial of `A` and of `s1` and `s2` whatever the samplers return, and \
- zeroes the polynomial of a sampler that fails rather than branching on it: its timing does not \
- depend on whether key generation fails."]
-
-def artifacts : List Artifact := [
- { Spec.MlDsa.keyGen44Api with
- target := X86_64.target
- doc := Spec.MlDsa.keyGen44Api.doc (notes := notes)
- code := Impl.MlDsa.X86_64.KeyGen.keyGen44
- contract := Spec.MlDsa.keyGenContract Spec.MlDsa.mlDsa44 X86_64.abi 32
- stack := 32
- verified := Proof.MlDsa.X86_64.KeyGen.keyGen44_verified
- spSafe := Code.all_of_allInstrs (by decide +kernel) },
- { Spec.MlDsa.keyGen65Api with
- target := X86_64.target
- doc := Spec.MlDsa.keyGen65Api.doc (notes := notes)
- code := Impl.MlDsa.X86_64.KeyGen.keyGen65
- contract := Spec.MlDsa.keyGenContract Spec.MlDsa.mlDsa65 X86_64.abi 32
- stack := 32
- verified := Proof.MlDsa.X86_64.KeyGen.keyGen65_verified
- spSafe := Code.all_of_allInstrs (by decide +kernel) },
- { Spec.MlDsa.keyGen87Api with
- target := X86_64.target
- doc := Spec.MlDsa.keyGen87Api.doc (notes := notes)
- code := Impl.MlDsa.X86_64.KeyGen.keyGen87
- contract := Spec.MlDsa.keyGenContract Spec.MlDsa.mlDsa87 X86_64.abi 32
- stack := 32
- verified := Proof.MlDsa.X86_64.KeyGen.keyGen87_verified
- spSafe := Code.all_of_allInstrs (by decide +kernel) }]
-
-end VG.Artifacts.MlDsaKeyGen.X86_64
diff --git a/lean/VerifiedGarbage/Artifacts/MlDsaSign/X86_64.lean b/lean/VerifiedGarbage/Artifacts/MlDsaSign/X86_64.lean
deleted file mode 100644
index c64c90d20..000000000
--- a/lean/VerifiedGarbage/Artifacts/MlDsaSign/X86_64.lean
+++ /dev/null
@@ -1,55 +0,0 @@
-import VerifiedGarbage.TCB.X86_64.Target
-import VerifiedGarbage.Proof.MlDsa.X86_64.Sign.Verified
-
-/-!
-# ML-DSA (FIPS 204) signing on x86-64
-
-A registration file (see `TCB/Emit.lean`): the artifacts it lists are
-emitted. **Review note**: `sig` and `doc` are trusted, as they tie the Rust
-caller to the contract; check them against the contract's `pre`/`post`. Each
-artifact is made from its function's `Api` (in `Spec/MlDsa/Contract.lean`,
-reviewed with the contract), and this file adds only notes on the
-implementation. The emitter adds the `# Safety` items that depend on the
-target (`Sig.layoutDoc`), from `stack` and `writeArgs`, which `ofSig` checks
-against the contract.
--/
-
-namespace VG.Artifacts.MlDsaSign.X86_64
-
-open VG.Proof.MlDsa.X86_64.Sign (prims)
-
-/-- Notes on the implementation, the same for every parameter set. -/
-def notes : List String :=
- ["The function saves its caller's callee-saved registers in `scratch`; its calls use the 24 \
- bytes of stack below its return address.",
- "The signing loop runs at most 814 iterations (FIPS 204 Appendix C). Each iteration computes \
- every validity check and combines them without branching: the one branch on their result \
- is the only place an iteration's outcome affects timing."]
-
-def artifacts : List Artifact := [
- { Spec.MlDsa.sign44Api with
- target := X86_64.target
- doc := Spec.MlDsa.sign44Api.doc (notes := notes)
- code := Impl.MlDsa.X86_64.Sign.sign prims Spec.MlDsa.mlDsa44
- contract := Spec.MlDsa.signContract Spec.MlDsa.mlDsa44 X86_64.abi 24
- stack := 24
- verified := Proof.MlDsa.X86_64.Sign.sign44_verified'
- spSafe := Proof.MlDsa.X86_64.Sign.sign44_spSafe },
- { Spec.MlDsa.sign65Api with
- target := X86_64.target
- doc := Spec.MlDsa.sign65Api.doc (notes := notes)
- code := Impl.MlDsa.X86_64.Sign.sign prims Spec.MlDsa.mlDsa65
- contract := Spec.MlDsa.signContract Spec.MlDsa.mlDsa65 X86_64.abi 24
- stack := 24
- verified := Proof.MlDsa.X86_64.Sign.sign65_verified'
- spSafe := Proof.MlDsa.X86_64.Sign.sign65_spSafe },
- { Spec.MlDsa.sign87Api with
- target := X86_64.target
- doc := Spec.MlDsa.sign87Api.doc (notes := notes)
- code := Impl.MlDsa.X86_64.Sign.sign prims Spec.MlDsa.mlDsa87
- contract := Spec.MlDsa.signContract Spec.MlDsa.mlDsa87 X86_64.abi 24
- stack := 24
- verified := Proof.MlDsa.X86_64.Sign.sign87_verified'
- spSafe := Proof.MlDsa.X86_64.Sign.sign87_spSafe }]
-
-end VG.Artifacts.MlDsaSign.X86_64
diff --git a/lean/VerifiedGarbage/Artifacts/MlDsaVerify/X86_64.lean b/lean/VerifiedGarbage/Artifacts/MlDsaVerify/X86_64.lean
deleted file mode 100644
index bb1ea07ef..000000000
--- a/lean/VerifiedGarbage/Artifacts/MlDsaVerify/X86_64.lean
+++ /dev/null
@@ -1,56 +0,0 @@
-import VerifiedGarbage.TCB.X86_64.Target
-import VerifiedGarbage.Proof.MlDsa.X86_64.Verify.Prims
-
-/-!
-# ML-DSA (FIPS 204) on x86-64: verification
-
-A registration file (see `TCB/Emit.lean`): the artifacts it lists are
-emitted. **Review note**: `sig` and `doc` are trusted, as they tie the Rust
-caller to the contract; check them against the contract's `pre`/`post`. Each
-artifact is made from its function's `Api` (in `Spec/MlDsa/Contract.lean`,
-reviewed with the contract), and this file adds only notes on the
-implementation. The emitter adds the `# Safety` items that depend on the
-target (`Sig.layoutDoc`), from `stack` and `writeArgs`, which `ofSig` checks
-against the contract.
-
-The stack is 24 bytes: the return address of a call of a primitive, and up
-to 16 bytes for its own calls.
--/
-
-namespace VG.Artifacts.MlDsaVerify.X86_64
-
-open VG
-open VG.Proof.MlDsa.X86_64.Verify (prims prims_ok verify_prims verify_spSafe)
-
-/-- What the documentation says of the implementation. -/
-def note : String :=
- "It calls the `vg_mldsa_*` primitives and the SHAKE256 sponge. The samplers' results are combined \
- without a branch, so the only branches depend on the public key and the signature."
-
-def artifacts : List Artifact := [
- { Spec.MlDsa.verify44Api with
- target := X86_64.target
- doc := Spec.MlDsa.verify44Api.doc (notes := [note])
- code := Impl.MlDsa.X86_64.Verify.verify prims Spec.MlDsa.mlDsa44
- contract := Spec.MlDsa.verifyContract Spec.MlDsa.mlDsa44 X86_64.abi 24
- stack := 24
- verified := verify_prims (List.mem_cons_self ..)
- spSafe := verify_spSafe prims_ok (List.mem_cons_self ..) },
- { Spec.MlDsa.verify65Api with
- target := X86_64.target
- doc := Spec.MlDsa.verify65Api.doc (notes := [note])
- code := Impl.MlDsa.X86_64.Verify.verify prims Spec.MlDsa.mlDsa65
- contract := Spec.MlDsa.verifyContract Spec.MlDsa.mlDsa65 X86_64.abi 24
- stack := 24
- verified := verify_prims (List.mem_cons_of_mem _ (List.mem_cons_self ..))
- spSafe := verify_spSafe prims_ok (List.mem_cons_of_mem _ (List.mem_cons_self ..)) },
- { Spec.MlDsa.verify87Api with
- target := X86_64.target
- doc := Spec.MlDsa.verify87Api.doc (notes := [note])
- code := Impl.MlDsa.X86_64.Verify.verify prims Spec.MlDsa.mlDsa87
- contract := Spec.MlDsa.verifyContract Spec.MlDsa.mlDsa87 X86_64.abi 24
- stack := 24
- verified := verify_prims (List.mem_cons_of_mem _ (List.mem_cons_of_mem _ (List.mem_cons_self ..)))
- spSafe := verify_spSafe prims_ok (List.mem_cons_of_mem _ (List.mem_cons_of_mem _ (List.mem_cons_self ..))) }]
-
-end VG.Artifacts.MlDsaVerify.X86_64
diff --git a/lean/VerifiedGarbage/Generic/MlDsaArith/X86_64/MlDsaKeyGen.lean b/lean/VerifiedGarbage/Generic/MlDsaArith/X86_64/MlDsaKeyGen.lean
new file mode 100644
index 000000000..bddd1b047
--- /dev/null
+++ b/lean/VerifiedGarbage/Generic/MlDsaArith/X86_64/MlDsaKeyGen.lean
@@ -0,0 +1,65 @@
+import VerifiedGarbage.TCB.X86_64.Target
+import VerifiedGarbage.Proof.MlDsa.X86_64.KeyGen.Inst
+
+/-!
+# ML-DSA (FIPS 204) on x86-64: key generation
+
+A generic file (see `TCB/Emit.lean`): the artifacts it lists, which call an
+implementation `v` of the polynomial arithmetic (`vg_mldsa_ntt`, …), are
+emitted once for each implementation (`Variants/MlDsaArith/X86_64/`), named
+with its suffix (e.g. `vg_mldsa44_keygen_avx2`), and need its CPU features.
+**Review note**: `sig` and `doc` are trusted, as they tie the Rust caller to
+the contract; check them against the contract's `pre`/`post`. Each artifact
+is made from its function's `Api` (in `Spec/MlDsa/Contract.lean`, reviewed
+with the contract), and this file adds only notes on the implementation. The
+emitter adds the `# Safety` items that depend on the target
+(`Sig.layoutDoc`), from `stack` and `writeArgs`, which `ofSig` checks
+against the contract.
+-/
+
+namespace VG.Generic.MlDsaArith.X86_64.MlDsaKeyGen
+
+open VG.Proof.MlDsa.X86_64 (ArithImpl)
+open VG.Impl.MlDsa.X86_64.KeyGen (keyGen primsWith)
+
+/-- Notes on the implementation, the same for every parameter set. -/
+def notes : List String :=
+ ["The function saves its caller's callee-saved registers in `scratch`; its calls use the 32 \
+ bytes of stack below its return address.",
+ "It samples every polynomial of `A` and of `s1` and `s2` whatever the samplers return, and \
+ zeroes the polynomial of a sampler that fails rather than branching on it: its timing does not \
+ depend on whether key generation fails."]
+
+def artifacts (v : ArithImpl) : List Artifact := [
+ { Spec.MlDsa.keyGen44Api with
+ name := Spec.MlDsa.keyGen44Api.name ++ v.code.sfx
+ features := v.features
+ target := X86_64.target
+ doc := Spec.MlDsa.keyGen44Api.doc (notes := notes)
+ code := keyGen (primsWith v.code) Spec.MlDsa.mlDsa44
+ contract := Spec.MlDsa.keyGenContract Spec.MlDsa.mlDsa44 X86_64.abi 32
+ stack := 32
+ verified := Proof.MlDsa.X86_64.KeyGen.keyGen_verifiedWith v (.inl rfl)
+ spSafe := Proof.MlDsa.X86_64.KeyGen.keyGen_spSafe v (.inl rfl) },
+ { Spec.MlDsa.keyGen65Api with
+ name := Spec.MlDsa.keyGen65Api.name ++ v.code.sfx
+ features := v.features
+ target := X86_64.target
+ doc := Spec.MlDsa.keyGen65Api.doc (notes := notes)
+ code := keyGen (primsWith v.code) Spec.MlDsa.mlDsa65
+ contract := Spec.MlDsa.keyGenContract Spec.MlDsa.mlDsa65 X86_64.abi 32
+ stack := 32
+ verified := Proof.MlDsa.X86_64.KeyGen.keyGen_verifiedWith v (.inr (.inl rfl))
+ spSafe := Proof.MlDsa.X86_64.KeyGen.keyGen_spSafe v (.inr (.inl rfl)) },
+ { Spec.MlDsa.keyGen87Api with
+ name := Spec.MlDsa.keyGen87Api.name ++ v.code.sfx
+ features := v.features
+ target := X86_64.target
+ doc := Spec.MlDsa.keyGen87Api.doc (notes := notes)
+ code := keyGen (primsWith v.code) Spec.MlDsa.mlDsa87
+ contract := Spec.MlDsa.keyGenContract Spec.MlDsa.mlDsa87 X86_64.abi 32
+ stack := 32
+ verified := Proof.MlDsa.X86_64.KeyGen.keyGen_verifiedWith v (.inr (.inr rfl))
+ spSafe := Proof.MlDsa.X86_64.KeyGen.keyGen_spSafe v (.inr (.inr rfl)) }]
+
+end VG.Generic.MlDsaArith.X86_64.MlDsaKeyGen
diff --git a/lean/VerifiedGarbage/Generic/MlDsaArith/X86_64/MlDsaSign.lean b/lean/VerifiedGarbage/Generic/MlDsaArith/X86_64/MlDsaSign.lean
new file mode 100644
index 000000000..2a231e969
--- /dev/null
+++ b/lean/VerifiedGarbage/Generic/MlDsaArith/X86_64/MlDsaSign.lean
@@ -0,0 +1,65 @@
+import VerifiedGarbage.TCB.X86_64.Target
+import VerifiedGarbage.Proof.MlDsa.X86_64.Sign.Verified
+
+/-!
+# ML-DSA (FIPS 204) on x86-64: signing
+
+A generic file (see `TCB/Emit.lean`): the artifacts it lists, which call an
+implementation `v` of the polynomial arithmetic (`vg_mldsa_ntt`, …), are
+emitted once for each implementation (`Variants/MlDsaArith/X86_64/`), named
+with its suffix (e.g. `vg_mldsa44_sign_avx2`), and need its CPU features.
+**Review note**: `sig` and `doc` are trusted, as they tie the Rust caller to
+the contract; check them against the contract's `pre`/`post`. Each artifact
+is made from its function's `Api` (in `Spec/MlDsa/Contract.lean`, reviewed
+with the contract), and this file adds only notes on the implementation. The
+emitter adds the `# Safety` items that depend on the target
+(`Sig.layoutDoc`), from `stack` and `writeArgs`, which `ofSig` checks
+against the contract.
+-/
+
+namespace VG.Generic.MlDsaArith.X86_64.MlDsaSign
+
+open VG.Proof.MlDsa.X86_64 (ArithImpl)
+open VG.Proof.MlDsa.X86_64.Sign (primsWith)
+
+/-- Notes on the implementation, the same for every parameter set. -/
+def notes : List String :=
+ ["The function saves its caller's callee-saved registers in `scratch`; its calls use the 24 \
+ bytes of stack below its return address.",
+ "The signing loop runs at most 814 iterations (FIPS 204 Appendix C). Each iteration computes \
+ every validity check and combines them without branching: the one branch on their result \
+ is the only place an iteration's outcome affects timing."]
+
+def artifacts (v : ArithImpl) : List Artifact := [
+ { Spec.MlDsa.sign44Api with
+ name := Spec.MlDsa.sign44Api.name ++ v.code.sfx
+ features := v.features
+ target := X86_64.target
+ doc := Spec.MlDsa.sign44Api.doc (notes := notes)
+ code := Impl.MlDsa.X86_64.Sign.sign (primsWith v.code) Spec.MlDsa.mlDsa44
+ contract := Spec.MlDsa.signContract Spec.MlDsa.mlDsa44 X86_64.abi 24
+ stack := 24
+ verified := Proof.MlDsa.X86_64.Sign.sign_verified' v (.inl rfl)
+ spSafe := Proof.MlDsa.X86_64.Sign.sign_spSafe v (.inl rfl) },
+ { Spec.MlDsa.sign65Api with
+ name := Spec.MlDsa.sign65Api.name ++ v.code.sfx
+ features := v.features
+ target := X86_64.target
+ doc := Spec.MlDsa.sign65Api.doc (notes := notes)
+ code := Impl.MlDsa.X86_64.Sign.sign (primsWith v.code) Spec.MlDsa.mlDsa65
+ contract := Spec.MlDsa.signContract Spec.MlDsa.mlDsa65 X86_64.abi 24
+ stack := 24
+ verified := Proof.MlDsa.X86_64.Sign.sign_verified' v (.inr (.inl rfl))
+ spSafe := Proof.MlDsa.X86_64.Sign.sign_spSafe v (.inr (.inl rfl)) },
+ { Spec.MlDsa.sign87Api with
+ name := Spec.MlDsa.sign87Api.name ++ v.code.sfx
+ features := v.features
+ target := X86_64.target
+ doc := Spec.MlDsa.sign87Api.doc (notes := notes)
+ code := Impl.MlDsa.X86_64.Sign.sign (primsWith v.code) Spec.MlDsa.mlDsa87
+ contract := Spec.MlDsa.signContract Spec.MlDsa.mlDsa87 X86_64.abi 24
+ stack := 24
+ verified := Proof.MlDsa.X86_64.Sign.sign_verified' v (.inr (.inr rfl))
+ spSafe := Proof.MlDsa.X86_64.Sign.sign_spSafe v (.inr (.inr rfl)) }]
+
+end VG.Generic.MlDsaArith.X86_64.MlDsaSign
diff --git a/lean/VerifiedGarbage/Generic/MlDsaArith/X86_64/MlDsaVerify.lean b/lean/VerifiedGarbage/Generic/MlDsaArith/X86_64/MlDsaVerify.lean
new file mode 100644
index 000000000..903d1f596
--- /dev/null
+++ b/lean/VerifiedGarbage/Generic/MlDsaArith/X86_64/MlDsaVerify.lean
@@ -0,0 +1,66 @@
+import VerifiedGarbage.TCB.X86_64.Target
+import VerifiedGarbage.Proof.MlDsa.X86_64.Verify.Prims
+
+/-!
+# ML-DSA (FIPS 204) on x86-64: verification
+
+A generic file (see `TCB/Emit.lean`): the artifacts it lists, which call an
+implementation `v` of the polynomial arithmetic (`vg_mldsa_ntt`, …), are
+emitted once for each implementation (`Variants/MlDsaArith/X86_64/`), named
+with its suffix (e.g. `vg_mldsa44_verify_avx2`), and need its CPU features.
+**Review note**: `sig` and `doc` are trusted, as they tie the Rust caller to
+the contract; check them against the contract's `pre`/`post`. Each artifact
+is made from its function's `Api` (in `Spec/MlDsa/Contract.lean`, reviewed
+with the contract), and this file adds only notes on the implementation. The
+emitter adds the `# Safety` items that depend on the target
+(`Sig.layoutDoc`), from `stack` and `writeArgs`, which `ofSig` checks
+against the contract.
+
+The stack is 24 bytes: the return address of a call of a primitive, and up
+to 16 bytes for its own calls.
+-/
+
+namespace VG.Generic.MlDsaArith.X86_64.MlDsaVerify
+
+open VG
+open VG.Proof.MlDsa.X86_64 (ArithImpl)
+open VG.Proof.MlDsa.X86_64.Verify (primsWith prims_okWith verify_prims verify_spSafe)
+
+/-- What the documentation says of the implementation. -/
+def note : String :=
+ "It calls the `vg_mldsa_*` primitives and the SHAKE256 sponge. The samplers' results are combined \
+ without a branch, so the only branches depend on the public key and the signature."
+
+def artifacts (v : ArithImpl) : List Artifact := [
+ { Spec.MlDsa.verify44Api with
+ name := Spec.MlDsa.verify44Api.name ++ v.code.sfx
+ features := v.features
+ target := X86_64.target
+ doc := Spec.MlDsa.verify44Api.doc (notes := [note])
+ code := Impl.MlDsa.X86_64.Verify.verify (primsWith v.code) Spec.MlDsa.mlDsa44
+ contract := Spec.MlDsa.verifyContract Spec.MlDsa.mlDsa44 X86_64.abi 24
+ stack := 24
+ verified := verify_prims v (List.mem_cons_self ..)
+ spSafe := verify_spSafe (prims_okWith v) (List.mem_cons_self ..) },
+ { Spec.MlDsa.verify65Api with
+ name := Spec.MlDsa.verify65Api.name ++ v.code.sfx
+ features := v.features
+ target := X86_64.target
+ doc := Spec.MlDsa.verify65Api.doc (notes := [note])
+ code := Impl.MlDsa.X86_64.Verify.verify (primsWith v.code) Spec.MlDsa.mlDsa65
+ contract := Spec.MlDsa.verifyContract Spec.MlDsa.mlDsa65 X86_64.abi 24
+ stack := 24
+ verified := verify_prims v (List.mem_cons_of_mem _ (List.mem_cons_self ..))
+ spSafe := verify_spSafe (prims_okWith v) (List.mem_cons_of_mem _ (List.mem_cons_self ..)) },
+ { Spec.MlDsa.verify87Api with
+ name := Spec.MlDsa.verify87Api.name ++ v.code.sfx
+ features := v.features
+ target := X86_64.target
+ doc := Spec.MlDsa.verify87Api.doc (notes := [note])
+ code := Impl.MlDsa.X86_64.Verify.verify (primsWith v.code) Spec.MlDsa.mlDsa87
+ contract := Spec.MlDsa.verifyContract Spec.MlDsa.mlDsa87 X86_64.abi 24
+ stack := 24
+ verified := verify_prims v (List.mem_cons_of_mem _ (List.mem_cons_of_mem _ (List.mem_cons_self ..)))
+ spSafe := verify_spSafe (prims_okWith v) (List.mem_cons_of_mem _ (List.mem_cons_of_mem _ (List.mem_cons_self ..))) }]
+
+end VG.Generic.MlDsaArith.X86_64.MlDsaVerify
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/Backend.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean
new file mode 100644
index 000000000..c6ef074df
--- /dev/null
+++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean
@@ -0,0 +1,41 @@
+import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Ntt
+import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Mul
+import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.AddSub
+
+/-!
+# ML-DSA on x86-64: implementations of the polynomial arithmetic
+
+Key generation, signing and verification call the polynomial arithmetic of
+one implementation, a `Backend`: the code of `vg_mldsa_ntt`,
+`vg_mldsa_inv_ntt`, `vg_mldsa_multiply_ntt`, `vg_mldsa_multiply_add_ntt`,
+`vg_mldsa_add` and `vg_mldsa_sub`, whose names end with `sfx` (e.g.
+`_avx2`; nothing for the SSE2 code, `sse2`). Each is a variant of the
+interface `MlDsaArith` on x86-64 (`Variants/MlDsaArith/X86_64/`), and the
+functions that call them are emitted once for each
+(`Generic/MlDsaArith/X86_64/`).
+-/
+
+namespace VG.Impl.MlDsa.X86_64.Arith
+
+open VG.X86_64
+
+/-- An implementation of the polynomial arithmetic. -/
+structure Backend where
+ ntt : Prog isa
+ invNtt : Prog isa
+ mul : Prog isa
+ mulAdd : Prog isa
+ add : Prog isa
+ sub : Prog isa
+ /-- What the names of its functions, and of those calling them, end with. -/
+ sfx : String
+
+/-- The SSE2 code. -/
+def Backend.sse2 : Backend := ⟨Arith.ntt, Arith.nttInv, Arith.mul, Arith.mulAdd, Arith.add, Arith.sub, ""⟩
+
+/-- Every function empty, which the proofs that the functions calling a
+backend never write `rsp` (and load MXCSR only to restore it) evaluate in
+its place. -/
+def Backend.empty : Backend := ⟨.block [], .block [], .block [], .block [], .block [], .block [], ""⟩
+
+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/Impl/MlDsa/X86_64/KeyGen/Inst.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/KeyGen/Inst.lean
index f4a1d0c87..5206623ad 100644
--- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/KeyGen/Inst.lean
+++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/KeyGen/Inst.lean
@@ -1,7 +1,5 @@
import VerifiedGarbage.Impl.MlDsa.X86_64.KeyGen.KeyGen
-import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Ntt
-import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Mul
-import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.AddSub
+import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Backend
import VerifiedGarbage.Impl.MlDsa.X86_64.Sample.RejNtt
import VerifiedGarbage.Impl.MlDsa.X86_64.Sample.RejBounded
import VerifiedGarbage.Impl.MlDsa.X86_64.Round.Round
@@ -11,7 +9,8 @@ import VerifiedGarbage.Impl.MlDsa.X86_64.Pack.Encode
# ML-DSA key generation on x86-64, with this library's primitives
`keyGen` (`KeyGen.lean`) called with the x86-64 implementations of the
-primitives it calls (`Arith/`, `Sample/`, `Round/`, `Pack/`).
+primitives it calls (`Arith/`, `Sample/`, `Round/`, `Pack/`), with the
+polynomial arithmetic of a `Backend` (`primsWith`).
-/
namespace VG.Impl.MlDsa.X86_64.KeyGen
@@ -31,11 +30,14 @@ def prims : Prims where
simpleBitPack := Pack.simpleBitPack
bitPack := Pack.bitPack
-/-- `vg_mldsa44_keygen` -/
-def keyGen44 : Prog isa := keyGen prims Spec.MlDsa.mlDsa44
-/-- `vg_mldsa65_keygen` -/
-def keyGen65 : Prog isa := keyGen prims Spec.MlDsa.mlDsa65
-/-- `vg_mldsa87_keygen` -/
-def keyGen87 : Prog isa := keyGen prims Spec.MlDsa.mlDsa87
+/-- The primitives, with the polynomial arithmetic of `B`. -/
+def primsWith (B : Arith.Backend) : Prims :=
+ { prims with
+ ntt := B.ntt
+ invNtt := B.invNtt
+ mul := B.mul
+ mulAdd := B.mulAdd
+ add := B.add
+ sfx := B.sfx }
end VG.Impl.MlDsa.X86_64.KeyGen
diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/KeyGen/KeyGen.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/KeyGen/KeyGen.lean
index 06126c30f..4f3d0ef4c 100644
--- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/KeyGen/KeyGen.lean
+++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/KeyGen/KeyGen.lean
@@ -71,20 +71,20 @@ def oT0 (p : Params) : Nat := 128 + lenS p * (p.ℓ + p.k)
/-- `r ← v`, a 32-bit immediate. -/
def imm (r : Reg) (v : Nat) : List Instr := [.mov32 r (.imm (BitVec.ofNat 32 v))]
-def nttAt (c : Prog isa) (f : Ptr) : Prog isa :=
- .seq (.block (lea .rdi f ++ lea .rsi (sc oSS))) (.call "vg_mldsa_ntt" c)
+def nttAt (sfx : String) (c : Prog isa) (f : Ptr) : Prog isa :=
+ .seq (.block (lea .rdi f ++ lea .rsi (sc oSS))) (.call ("vg_mldsa_ntt" ++ sfx) c)
-def invNttAt (c : Prog isa) (f : Ptr) : Prog isa :=
- .seq (.block (lea .rdi f ++ lea .rsi (sc oSS))) (.call "vg_mldsa_inv_ntt" c)
+def invNttAt (sfx : String) (c : Prog isa) (f : Ptr) : Prog isa :=
+ .seq (.block (lea .rdi f ++ lea .rsi (sc oSS))) (.call ("vg_mldsa_inv_ntt" ++ sfx) c)
-def mulAt (c : Prog isa) (h f g : Ptr) : Prog isa :=
- .seq (.block (lea .rdi h ++ lea .rsi f ++ lea .rdx g)) (.call "vg_mldsa_multiply_ntt" c)
+def mulAt (sfx : String) (c : Prog isa) (h f g : Ptr) : Prog isa :=
+ .seq (.block (lea .rdi h ++ lea .rsi f ++ lea .rdx g)) (.call ("vg_mldsa_multiply_ntt" ++ sfx) c)
-def mulAddAt (c : Prog isa) (h f g : Ptr) : Prog isa :=
- .seq (.block (lea .rdi h ++ lea .rsi f ++ lea .rdx g)) (.call "vg_mldsa_multiply_add_ntt" c)
+def mulAddAt (sfx : String) (c : Prog isa) (h f g : Ptr) : Prog isa :=
+ .seq (.block (lea .rdi h ++ lea .rsi f ++ lea .rdx g)) (.call ("vg_mldsa_multiply_add_ntt" ++ sfx) c)
-def addAt (c : Prog isa) (f g : Ptr) : Prog isa :=
- .seq (.block (lea .rdi f ++ lea .rsi g)) (.call "vg_mldsa_add" c)
+def addAt (sfx : String) (c : Prog isa) (f g : Ptr) : Prog isa :=
+ .seq (.block (lea .rdi f ++ lea .rsi g)) (.call ("vg_mldsa_add" ++ sfx) c)
def rejNttAt (c : Prog isa) (seed a : Ptr) : Prog isa :=
.seq (.block (lea .rdi seed ++ lea .rsi a ++ lea .rdx (sc oSS))) (.call "vg_mldsa_rej_ntt_poly" c)
@@ -141,13 +141,13 @@ def packS (P : Prims) (p : Params) (r : Nat) : Prog isa :=
bitPackAt P.bitPack (sP p r) p.η p.η (.r13, 128 + lenS p * r) (lenS p)
/-- `ŝ₁[j] = NTT(s₁[j])`. -/
-def nttS (P : Prims) (p : Params) (j : Nat) : Prog isa := nttAt P.ntt (sP p j)
+def nttS (P : Prims) (p : Params) (j : Nat) : Prog isa := nttAt P.sfx P.ntt (sP p j)
/-- Row `i`: `t = NTT⁻¹(Σⱼ Â[i, j] ŝ₁[j]) + s₂[i]`, and its `t₁` to `pk` and `t₀` to `sk`. -/
def row (P : Prims) (p : Params) (i : Nat) : Prog isa :=
- .seq (mulAt P.mul (tP p) (aP (p.ℓ * i)) (sP p 0))
- (.seq (seqR (fun j => mulAddAt P.mulAdd (tP p) (aP (p.ℓ * i + j)) (sP p j)) 1 (p.ℓ - 1))
- (.seq (invNttAt P.invNtt (tP p)) (.seq (addAt P.add (tP p) (sP p (p.ℓ + i)))
+ .seq (mulAt P.sfx P.mul (tP p) (aP (p.ℓ * i)) (sP p 0))
+ (.seq (seqR (fun j => mulAddAt P.sfx P.mulAdd (tP p) (aP (p.ℓ * i + j)) (sP p j)) 1 (p.ℓ - 1))
+ (.seq (invNttAt P.sfx P.invNtt (tP p)) (.seq (addAt P.sfx P.add (tP p) (sP p (p.ℓ + i)))
(.seq (power2RoundAt P.power2Round (tP p) (t1P p) (t0P p))
(.seq (simpleBitPackAt P.simpleBitPack (t1P p) 1023 (.r12, 32 + 320 * i) 320)
(bitPackAt P.bitPack (t0P p) 4095 4096 (.r13, oT0 p + 416 * i) 416))))))
diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/KeyGen/Prims.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/KeyGen/Prims.lean
index a3e13476d..a2ed604be 100644
--- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/KeyGen/Prims.lean
+++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/KeyGen/Prims.lean
@@ -36,5 +36,7 @@ structure Prims where
simpleBitPack : Prog isa
/-- `vg_mldsa_bit_pack` -/
bitPack : Prog isa
+ /-- What the names of the polynomial arithmetic's functions end with (`Arith.Backend`). -/
+ sfx : String := ""
end VG.Impl.MlDsa.X86_64.KeyGen
diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean
index 11dc35463..4fa71ab43 100644
--- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean
+++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean
@@ -44,6 +44,8 @@ structure Prims where
bitPack : Prog isa
bitUnpack : Prog isa
hintBitPack : Prog isa
+ /-- What the names of the polynomial arithmetic's functions end with (`Arith.Backend`). -/
+ sfx : String := ""
/-! ## The layout of the working space (in bytes)
@@ -165,17 +167,17 @@ Each takes its working space (if any) at `PS`. -/
section
variable (P : Prims)
-def nttAt (f : Ptr) : Prog isa := callP "vg_mldsa_ntt" P.ntt [.ptr f, .ptr (sc oPS)]
+def nttAt (f : Ptr) : Prog isa := callP ("vg_mldsa_ntt" ++ P.sfx) P.ntt [.ptr f, .ptr (sc oPS)]
-def invNttAt (f : Ptr) : Prog isa := callP "vg_mldsa_inv_ntt" P.invNtt [.ptr f, .ptr (sc oPS)]
+def invNttAt (f : Ptr) : Prog isa := callP ("vg_mldsa_inv_ntt" ++ P.sfx) P.invNtt [.ptr f, .ptr (sc oPS)]
-def mulAt (h f g : Ptr) : Prog isa := callP "vg_mldsa_multiply_ntt" P.mul [.ptr h, .ptr f, .ptr g]
+def mulAt (h f g : Ptr) : Prog isa := callP ("vg_mldsa_multiply_ntt" ++ P.sfx) P.mul [.ptr h, .ptr f, .ptr g]
-def mulAddAt (h f g : Ptr) : Prog isa := callP "vg_mldsa_multiply_add_ntt" P.mulAdd [.ptr h, .ptr f, .ptr g]
+def mulAddAt (h f g : Ptr) : Prog isa := callP ("vg_mldsa_multiply_add_ntt" ++ P.sfx) P.mulAdd [.ptr h, .ptr f, .ptr g]
-def addAt (f g : Ptr) : Prog isa := callP "vg_mldsa_add" P.add [.ptr f, .ptr g]
+def addAt (f g : Ptr) : Prog isa := callP ("vg_mldsa_add" ++ P.sfx) P.add [.ptr f, .ptr g]
-def subAt (f g : Ptr) : Prog isa := callP "vg_mldsa_sub" P.sub [.ptr f, .ptr g]
+def subAt (f g : Ptr) : Prog isa := callP ("vg_mldsa_sub" ++ P.sfx) P.sub [.ptr f, .ptr g]
/-- `RejNTTPoly` of the seed at `RS` to `a`, and `r15 ← r15 ∧ result`. -/
def rejAt (a : Ptr) : Prog isa :=
diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean
index cdeaf14d1..5b4e5e24f 100644
--- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean
+++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean
@@ -137,22 +137,24 @@ structure Prims where
unpackT1 : Prog isa
hintUnpack : Prog isa
normLt : Prog isa
+ /-- What the names of the polynomial arithmetic's functions end with (`Arith.Backend`). -/
+ sfx : String := ""
section
variable (P : Prims)
-def nttAt (f : Ptr) : Prog isa := callAt "vg_mldsa_ntt" P.ntt [(.rdi, .ptr f), (.rsi, .ptr (sc oSS))]
+def nttAt (f : Ptr) : Prog isa := callAt ("vg_mldsa_ntt" ++ P.sfx) P.ntt [(.rdi, .ptr f), (.rsi, .ptr (sc oSS))]
def invNttAt (f : Ptr) : Prog isa :=
- callAt "vg_mldsa_inv_ntt" P.invNtt [(.rdi, .ptr f), (.rsi, .ptr (sc oSS))]
+ callAt ("vg_mldsa_inv_ntt" ++ P.sfx) P.invNtt [(.rdi, .ptr f), (.rsi, .ptr (sc oSS))]
def mulAt (h f g : Ptr) : Prog isa :=
- callAt "vg_mldsa_multiply_ntt" P.mul [(.rdi, .ptr h), (.rsi, .ptr f), (.rdx, .ptr g)]
+ callAt ("vg_mldsa_multiply_ntt" ++ P.sfx) P.mul [(.rdi, .ptr h), (.rsi, .ptr f), (.rdx, .ptr g)]
def mulAddAt (h f g : Ptr) : Prog isa :=
- callAt "vg_mldsa_multiply_add_ntt" P.mulAdd [(.rdi, .ptr h), (.rsi, .ptr f), (.rdx, .ptr g)]
+ callAt ("vg_mldsa_multiply_add_ntt" ++ P.sfx) P.mulAdd [(.rdi, .ptr h), (.rsi, .ptr f), (.rdx, .ptr g)]
-def subAt (f g : Ptr) : Prog isa := callAt "vg_mldsa_sub" P.sub [(.rdi, .ptr f), (.rsi, .ptr g)]
+def subAt (f g : Ptr) : Prog isa := callAt ("vg_mldsa_sub" ++ P.sfx) P.sub [(.rdi, .ptr f), (.rsi, .ptr g)]
def rejNttAt (a : Ptr) : Prog isa :=
callAt "vg_mldsa_rej_ntt_poly" P.rejNtt [(.rdi, .ptr (sc oSB)), (.rsi, .ptr a), (.rdx, .ptr (sc oSS))]
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/Backend.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean
new file mode 100644
index 000000000..6a27d6c37
--- /dev/null
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean
@@ -0,0 +1,74 @@
+import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Backend
+import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Ntt
+import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.NttInv
+import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Mul
+import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.AddSub
+import VerifiedGarbage.Proof.MlKem.X86_64.ArithOk
+
+/-!
+# ML-DSA on x86-64: what the callers of the polynomial arithmetic need of it
+
+Untrusted: everything here is checked by Lean. An `ArithImpl` is an
+implementation of the polynomial arithmetic (`Impl.MlDsa.X86_64.Arith.Backend`)
+with what key generation, signing and verification need of each of its
+functions (`FnOk`): it meets its contract without using the stack, never
+writes `rsp`, calls no deeper than twice, and loads MXCSR only to restore
+it. Each is a variant of the interface `MlDsaArith` on x86-64
+(`Variants/MlDsaArith/X86_64/`), and the functions that call it are proven
+once for all of them (`Generic/MlDsaArith/X86_64/`).
+-/
+
+namespace VG.Proof.MlDsa.X86_64
+
+open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Arith
+
+/-- What a caller needs of a function with the contract `k` (given the
+stack its calls use) and the code `c`. -/
+structure FnOk (k : Nat → Contract isa) (c : Prog isa) : Prop where
+ ver : Verified X86_64.target c (k 0)
+ nosp : NoSp c
+ depth : c.depth ≤ 2
+ ctl : ctlOk c = true
+ sp : c.all (fun i => !isa.writesSp i) = true
+
+/-- Each function of the backend `B` meets its contract, and is safe to call. -/
+structure BackendOk (B : Backend) : Prop where
+ ntt : FnOk (fun S => Spec.MlDsa.nttContract X86_64.abi S) B.ntt
+ invNtt : FnOk (fun S => Spec.MlDsa.nttInvContract X86_64.abi S) B.invNtt
+ mul : FnOk (fun S => Spec.MlDsa.mulContract X86_64.abi S) B.mul
+ mulAdd : FnOk (fun S => Spec.MlDsa.mulAddContract X86_64.abi S) B.mulAdd
+ add : FnOk (fun S => Spec.MlDsa.addContract X86_64.abi S) B.add
+ sub : FnOk (fun S => Spec.MlDsa.subContract X86_64.abi S) B.sub
+
+/-- An implementation of the polynomial arithmetic on x86-64. -/
+structure ArithImpl where
+ code : Backend
+ ok : BackendOk code
+ /-- The CPU features its code requires, which its callers require too. -/
+ features : List String
+
+/-- `FnOk` of code verified without stack, from evaluating it. -/
+theorem FnOk.of {k : Nat → Contract isa} {c : Prog isa} (h : Verified X86_64.target c (k 0))
+ (hn : c.allInstrs (fun i => !Taint.clobbers i .rsp) = true) (hd : c.depth ≤ 2) (hc : ctlOk c = true)
+ (hs : c.allInstrs (fun i => !isa.writesSp i) = true) : FnOk k c :=
+ ⟨h, Proof.MlKem.X86_64.nosp_of hn, hd, hc, Code.all_of_allInstrs hs⟩
+
+/-- The SSE2 code. -/
+def ArithImpl.sse2 : ArithImpl where
+ code := .sse2
+ ok :=
+ { ntt := FnOk.of Arith.ntt_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
+ (by decide +kernel)
+ invNtt := FnOk.of Arith.nttInv_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
+ (by decide +kernel)
+ mul := FnOk.of Arith.mul_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
+ (by decide +kernel)
+ mulAdd := FnOk.of Arith.mulAdd_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
+ (by decide +kernel)
+ add := FnOk.of Arith.add_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
+ (by decide +kernel)
+ sub := FnOk.of Arith.sub_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
+ (by decide +kernel) }
+ features := []
+
+end VG.Proof.MlDsa.X86_64
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/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Same.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Same.lean
new file mode 100644
index 000000000..b08c4c356
--- /dev/null
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Same.lean
@@ -0,0 +1,100 @@
+import VerifiedGarbage.Proof.Framework.X86_64.Mxcsr
+import VerifiedGarbage.Impl.MlDsa.X86_64.Sign.Frag
+import VerifiedGarbage.Impl.MlDsa.X86_64.Verify.Frag
+import VerifiedGarbage.Impl.MlKem.X86_64.Frag
+
+/-!
+# ML-DSA on x86-64: checks of code, but for the functions it calls
+
+Untrusted: everything here is checked by Lean. A check of code that
+composes over its structure and looks at the code of each function called
+only through `mc` (`Comp m mc`: `ctlC`, with `ctlOk` of the functions
+called, and `Code.allInstrs q`) gives the same result on two programs that
+differ only in functions called, if it holds (`mc`) of those of the first
+and the second calls empty code instead (`Same`, proven by `same_tac` from
+the structure of the code). The top-level functions, written for any
+implementation of the polynomial arithmetic, are checked once with every
+function of it empty, so that the kernel evaluates no implementation of it.
+-/
+
+namespace VG.Proof.MlDsa.X86_64
+
+open VG VG.X86_64
+
+/-- `m` composes over the structure of code, and looks at the code of a
+function called only through `mc`. -/
+structure Comp (m mc : Prog isa → Bool) : Prop where
+ seq : ∀ a b, m (.seq a b) = (m a && m b)
+ ite : ∀ c t e, m (.ite c t e) = (m t && m e)
+ loop : ∀ b c, m (.loop b c) = m b
+ call : ∀ n b, m (.call n b) = mc b
+ nil : mc (.block []) = true
+
+theorem Comp.ctlC : Comp ctlC ctlOk :=
+ ⟨fun _ _ => rfl, fun _ _ _ => rfl, fun _ _ => rfl, fun _ _ => rfl, rfl⟩
+
+theorem Comp.all (q : Instr → Bool) : Comp (Code.allInstrs q) (Code.allInstrs q) :=
+ ⟨fun _ _ => rfl, fun _ _ _ => rfl, fun _ _ => rfl, fun _ _ => rfl, rfl⟩
+
+theorem Code.allInstrs_of_all {I C : Type} {q : I → Bool} {c : Code I C} (h : c.all q = true) :
+ c.allInstrs q = true := by
+ induction c with
+ | block is => induction is <;> simp_all [Code.all, Code.allInstrs]
+ | _ => simp_all [Code.all, Code.allInstrs]
+
+/-- `m` gives the same result on `c` and `c'`. -/
+def Same (m : Prog isa → Bool) (c c' : Prog isa) : Prop := m c = m c'
+
+section
+variable {m mc : Prog isa → Bool} (hm : Comp m mc)
+include hm
+
+theorem Same.seq {a a' b b' : Prog isa} (ha : Same m a a') (hb : Same m b b') : Same m (.seq a b) (.seq a' b') := by
+ unfold Same at *; rw [hm.seq, hm.seq, ha, hb]
+
+theorem Same.ite {c : isa.Cond} {t t' e e' : Prog isa} (ht : Same m t t') (he : Same m e e') :
+ Same m (.ite c t e) (.ite c t' e') := by
+ unfold Same at *; rw [hm.ite, hm.ite, ht, he]
+
+theorem Same.loop {b b' : Prog isa} {c : isa.Cond} (hb : Same m b b') : Same m (.loop b c) (.loop b' c) := by
+ unfold Same at *; rw [hm.loop, hm.loop, hb]
+
+theorem Same.call {c : Prog isa} (hc : mc c = true) (n n' : String) : Same m (.call n c) (.call n' (.block [])) := by
+ unfold Same; rw [hm.call, hm.call, hc, hm.nil]
+
+theorem Same.seqRS {f g : Nat → Prog isa} (h : ∀ k, Same m (f k) (g k)) :
+ ∀ a n, Same m (Impl.MlDsa.X86_64.Sign.seqR f a n) (Impl.MlDsa.X86_64.Sign.seqR g a n)
+ | _, 0 => rfl
+ | a, n + 1 => Same.seq hm (h a) (Same.seqRS h (a + 1) n)
+
+theorem Same.seqRV {f g : Nat → Prog isa} (h : ∀ k, Same m (f k) (g k)) :
+ ∀ a n, Same m (Impl.MlDsa.X86_64.Verify.seqR f a n) (Impl.MlDsa.X86_64.Verify.seqR g a n)
+ | _, 0 => rfl
+ | a, n + 1 => Same.seq hm (h a) (Same.seqRV h (a + 1) n)
+
+theorem Same.seqRK {f g : Nat → Prog isa} (h : ∀ k, Same m (f k) (g k)) :
+ ∀ a n, Same m (Impl.MlKem.X86_64.seqR f a n) (Impl.MlKem.X86_64.seqR g a n)
+ | _, 0 => rfl
+ | a, n + 1 => Same.seq hm (h a) (Same.seqRK h (a + 1) n)
+
+end
+
+/-- `m` of `a`, from that of code `b` that it gives the same result on. -/
+theorem Same.ok {m : Prog isa → Bool} {a b : Prog isa} (h : Same m a b) (hb : m b = true) : m a = true :=
+ Eq.trans h hb
+
+/-- `Same m c c'` for code `c` that calls functions whose `mc` the
+hypotheses state, and the same code `c'` but for empty functions in their
+place: from the structure of the code. -/
+macro "same_tac " hm:term : tactic =>
+ `(tactic| repeat' (first
+ | (apply Same.call $hm; assumption)
+ | apply Same.seq $hm
+ | apply Same.ite $hm
+ | apply Same.loop $hm
+ | (apply Same.seqRS $hm; intro)
+ | (apply Same.seqRV $hm; intro)
+ | (apply Same.seqRK $hm; intro)
+ | rfl))
+
+end VG.Proof.MlDsa.X86_64
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/KeyGen/Call.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/KeyGen/Call.lean
index dda7b3182..f8a03397c 100644
--- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/KeyGen/Call.lean
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/KeyGen/Call.lean
@@ -297,9 +297,9 @@ theorem covers_rw {s : State} {p : Params} (L : Lay kgR (kgW p) s) {a b : Ptr} {
covers_append (covers_cons (L.cR ha) covers_nil) (covers_cons (L.cR (inB_mono hb)) covers_nil)
include hf hg h1 w1 in
-theorem addAt_ok {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.addContract X86_64.abi stk) {s : State}
+theorem addAt_ok {sfx : String} {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.addContract X86_64.abi stk) {s : State}
(S : Site p s) (rf : Spec.MlDsa.Reduced s.mem (pa s f)) (rg : Spec.MlDsa.Reduced s.mem (pa s g)) :
- WP isa (addAt c f g) s fun s' => Post s s' [⟨pa s f, 1024⟩] ∧ MX s' = MX s ∧
+ WP isa (addAt sfx c f g) s fun s' => Post s s' [⟨pa s f, 1024⟩] ∧ MX s' = MX s ∧
Spec.MlDsa.PolyIs s'.mem (pa s f) (Spec.MlDsa.add (Spec.MlDsa.polyAt s.mem (pa s f))
(Spec.MlDsa.polyAt s.mem (pa s g))) := by
obtain ⟨i1, i2, _⟩ := sepB_spec h1
@@ -314,10 +314,10 @@ theorem addAt_ok {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.addContract
rwa [ce_polyAt (by rw [hsp]; exact L.stkD i1), ce_polyAt (by rw [hsp]; exact L.stkD i2), hm] at hpost
include hf hg h1 w1 in
-theorem addAt_tr {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.addContract X86_64.abi stk)
+theorem addAt_tr {sfx : String} {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.addContract X86_64.abi stk)
(hbf : f.1 ∈ kgRegs) (hbg : g.1 ∈ kgRegs) :
RelCT isa (fun x y => Two p x y ∧ (Spec.MlDsa.Reduced x.mem (pa x f) ∧ Spec.MlDsa.Reduced x.mem (pa x g)) ∧
- (Spec.MlDsa.Reduced y.mem (pa y f) ∧ Spec.MlDsa.Reduced y.mem (pa y g))) (addAt c f g) fun _ _ => True := by
+ (Spec.MlDsa.Reduced y.mem (pa y f) ∧ Spec.MlDsa.Reduced y.mem (pa y g))) (addAt sfx c f g) fun _ _ => True := by
obtain ⟨i1, i2, _⟩ := sepB_spec h1
refine primTr hc (nomem_append (lea_nomem _ _) (lea_nomem _ _)) (fun x y _ => ⟨glue2_ok hf hg x, glue2_ok hf hg y⟩)
fun stk hs x y x1 y1 ⟨T, rx, ry⟩ ⟨⟨hv1, hm1⟩, k1⟩ ⟨⟨hv2, hm2⟩, k2⟩ =>
@@ -362,9 +362,9 @@ theorem mul_pre {stk : Nat} (hstk : stk ≤ 16) {s s1 : State} (S : Site p s)
· exact (ce_reduced (by rw [hsp]; exact L.stkD i3)).mpr (hm ▸ rg)
include hh hf hg h1 h2 w1 in
-theorem mulAt_ok {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.mulContract X86_64.abi stk) {s : State}
+theorem mulAt_ok {sfx : String} {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.mulContract X86_64.abi stk) {s : State}
(S : Site p s) (rf : Spec.MlDsa.Reduced s.mem (pa s f)) (rg : Spec.MlDsa.Reduced s.mem (pa s g)) :
- WP isa (mulAt c h f g) s fun s' => Post s s' [⟨pa s h, 1024⟩] ∧ MX s' = MX s ∧
+ WP isa (mulAt sfx c h f g) s fun s' => Post s s' [⟨pa s h, 1024⟩] ∧ MX s' = MX s ∧
Spec.MlDsa.PolyIs s'.mem (pa s h) (Spec.MlDsa.multiplyNTT (Spec.MlDsa.polyAt s.mem (pa s f)) (Spec.MlDsa.polyAt s.mem (pa s g))) := by
obtain ⟨i1, i2, _⟩ := sepB_spec h1
obtain ⟨_, i3, _⟩ := sepB_spec h2
@@ -379,9 +379,9 @@ theorem mulAt_ok {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.mulContract
rwa [ce_polyAt (by rw [hsp]; exact L.stkD i2), ce_polyAt (by rw [hsp]; exact L.stkD i3), hm] at hpost
include hh hf hg h1 h2 w1 in
-theorem mulAt_tr {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.mulContract X86_64.abi stk)
+theorem mulAt_tr {sfx : String} {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.mulContract X86_64.abi stk)
(hbh : h.1 ∈ kgRegs) (hbf : f.1 ∈ kgRegs) (hbg : g.1 ∈ kgRegs) :
- RelCT isa (fun x y => Two p x y ∧ (Spec.MlDsa.Reduced x.mem (pa x f) ∧ Spec.MlDsa.Reduced x.mem (pa x g)) ∧ (Spec.MlDsa.Reduced y.mem (pa y f) ∧ Spec.MlDsa.Reduced y.mem (pa y g))) (mulAt c h f g) fun _ _ => True := by
+ RelCT isa (fun x y => Two p x y ∧ (Spec.MlDsa.Reduced x.mem (pa x f) ∧ Spec.MlDsa.Reduced x.mem (pa x g)) ∧ (Spec.MlDsa.Reduced y.mem (pa y f) ∧ Spec.MlDsa.Reduced y.mem (pa y g))) (mulAt sfx c h f g) fun _ _ => True := by
obtain ⟨i1, i2, _⟩ := sepB_spec h1
obtain ⟨_, i3, _⟩ := sepB_spec h2
refine primTr hc (nomem_append (nomem_append (lea_nomem _ _) (lea_nomem _ _)) (lea_nomem _ _))
@@ -425,9 +425,9 @@ theorem mulAdd_pre {stk : Nat} (hstk : stk ≤ 16) {s s1 : State} (S : Site p s)
· exact (ce_reduced (by rw [hsp]; exact L.stkD i3)).mpr (hm ▸ rg)
include hh hf hg h1 h2 w1 in
-theorem mulAddAt_ok {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.mulAddContract X86_64.abi stk) {s : State}
+theorem mulAddAt_ok {sfx : String} {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.mulAddContract X86_64.abi stk) {s : State}
(S : Site p s) (rh : Spec.MlDsa.Reduced s.mem (pa s h)) (rf : Spec.MlDsa.Reduced s.mem (pa s f)) (rg : Spec.MlDsa.Reduced s.mem (pa s g)) :
- WP isa (mulAddAt c h f g) s fun s' => Post s s' [⟨pa s h, 1024⟩] ∧ MX s' = MX s ∧
+ WP isa (mulAddAt sfx c h f g) s fun s' => Post s s' [⟨pa s h, 1024⟩] ∧ MX s' = MX s ∧
Spec.MlDsa.PolyIs s'.mem (pa s h) (Spec.MlDsa.add (Spec.MlDsa.polyAt s.mem (pa s h)) (Spec.MlDsa.multiplyNTT (Spec.MlDsa.polyAt s.mem (pa s f)) (Spec.MlDsa.polyAt s.mem (pa s g)))) := by
obtain ⟨i1, i2, _⟩ := sepB_spec h1
obtain ⟨_, i3, _⟩ := sepB_spec h2
@@ -442,9 +442,9 @@ theorem mulAddAt_ok {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.mulAddCo
rwa [ce_polyAt (by rw [hsp]; exact L.stkD i1), ce_polyAt (by rw [hsp]; exact L.stkD i2), ce_polyAt (by rw [hsp]; exact L.stkD i3), hm] at hpost
include hh hf hg h1 h2 w1 in
-theorem mulAddAt_tr {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.mulAddContract X86_64.abi stk)
+theorem mulAddAt_tr {sfx : String} {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.mulAddContract X86_64.abi stk)
(hbh : h.1 ∈ kgRegs) (hbf : f.1 ∈ kgRegs) (hbg : g.1 ∈ kgRegs) :
- RelCT isa (fun x y => Two p x y ∧ (Spec.MlDsa.Reduced x.mem (pa x h) ∧ Spec.MlDsa.Reduced x.mem (pa x f) ∧ Spec.MlDsa.Reduced x.mem (pa x g)) ∧ (Spec.MlDsa.Reduced y.mem (pa y h) ∧ Spec.MlDsa.Reduced y.mem (pa y f) ∧ Spec.MlDsa.Reduced y.mem (pa y g))) (mulAddAt c h f g) fun _ _ => True := by
+ RelCT isa (fun x y => Two p x y ∧ (Spec.MlDsa.Reduced x.mem (pa x h) ∧ Spec.MlDsa.Reduced x.mem (pa x f) ∧ Spec.MlDsa.Reduced x.mem (pa x g)) ∧ (Spec.MlDsa.Reduced y.mem (pa y h) ∧ Spec.MlDsa.Reduced y.mem (pa y f) ∧ Spec.MlDsa.Reduced y.mem (pa y g))) (mulAddAt sfx c h f g) fun _ _ => True := by
obtain ⟨i1, i2, _⟩ := sepB_spec h1
obtain ⟨_, i3, _⟩ := sepB_spec h2
refine primTr hc (nomem_append (nomem_append (lea_nomem _ _) (lea_nomem _ _)) (lea_nomem _ _))
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/KeyGen/Inst.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/KeyGen/Inst.lean
index 245ada812..0713c4074 100644
--- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/KeyGen/Inst.lean
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/KeyGen/Inst.lean
@@ -9,14 +9,17 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Sample.RejBoundedCT
import VerifiedGarbage.Proof.MlDsa.X86_64.Round.Power2Round
import VerifiedGarbage.Proof.MlDsa.X86_64.Pack.SimpleBitPack
import VerifiedGarbage.Proof.MlDsa.X86_64.Pack.BitPack
+import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Backend
+import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Same
/-!
# ML-DSA key generation on x86-64, with this library's primitives
Untrusted: everything here is checked by Lean. The x86-64 implementations of
the primitives (`prims`) are verified, use at most 16 bytes of stack, and
-never write `rsp` but by calls nested at most twice (`prims_ok`), so key
-generation with them is verified (`keyGen44_verified`, …).
+never write `rsp` but by calls nested at most twice, with any implementation
+`v` of the polynomial arithmetic (`prims_okWith`), so key generation with
+them is verified (`keyGen_verifiedWith`).
-/
namespace VG.Proof.MlDsa.X86_64.KeyGen
@@ -24,29 +27,56 @@ namespace VG.Proof.MlDsa.X86_64.KeyGen
open VG VG.X86_64 VG.Proof.MlKem.X86_64
open VG.Impl.MlDsa.X86_64.KeyGen
-theorem prims_ok : PrimsOk prims where
- ntt := ⟨⟨0, by decide, Arith.ntt_verified⟩, nosp_of (by decide +kernel), by decide +kernel⟩
- invNtt := ⟨⟨0, by decide, Arith.nttInv_verified⟩, nosp_of (by decide +kernel), by decide +kernel⟩
- mul := ⟨⟨0, by decide, Arith.mul_verified⟩, nosp_of (by decide +kernel), by decide +kernel⟩
- mulAdd := ⟨⟨0, by decide, Arith.mulAdd_verified⟩, nosp_of (by decide +kernel), by decide +kernel⟩
- add := ⟨⟨0, by decide, Arith.add_verified⟩, nosp_of (by decide +kernel), by decide +kernel⟩
- rejNtt := ⟨⟨16, by decide, Sample.rejNTT_verified⟩, nosp_of (by decide +kernel), by decide +kernel⟩
- rejBounded := ⟨⟨16, by decide, Sample.rejBounded_verified⟩, nosp_of (by decide +kernel), by decide +kernel⟩
- power2Round := ⟨⟨0, by decide, Round.power2Round_verified⟩, nosp_of (by decide +kernel), by decide +kernel⟩
- simpleBitPack := ⟨⟨0, by decide, Pack.simpleBitPack_verified⟩, nosp_of (by decide +kernel),
- by decide +kernel⟩
- bitPack := ⟨⟨0, by decide, Pack.bitPack_verified⟩, nosp_of (by decide +kernel), by decide +kernel⟩
-
-theorem keyGen44_verified :
- Verified X86_64.target keyGen44 (Spec.MlDsa.keyGenContract Spec.MlDsa.mlDsa44 X86_64.abi 32) :=
- keyGen_verified prims_ok _ (.inl rfl)
-
-theorem keyGen65_verified :
- Verified X86_64.target keyGen65 (Spec.MlDsa.keyGenContract Spec.MlDsa.mlDsa65 X86_64.abi 32) :=
- keyGen_verified prims_ok _ (.inr (.inl rfl))
-
-theorem keyGen87_verified :
- Verified X86_64.target keyGen87 (Spec.MlDsa.keyGenContract Spec.MlDsa.mlDsa87 X86_64.abi 32) :=
- keyGen_verified prims_ok _ (.inr (.inr rfl))
+open VG.Proof.MlDsa.X86_64 (FnOk ArithImpl Comp Same Same.ok Code.allInstrs_of_all)
+open VG.Impl.MlDsa.X86_64.Arith (Backend)
+
+/-- A function of the polynomial arithmetic satisfies what the proofs of key generation need of it. -/
+theorem calleeOf {k : Nat → Contract isa} {c : Prog isa} (h : FnOk k c) : Callee c k :=
+ ⟨⟨0, by decide, h.ver⟩, h.nosp, h.depth⟩
+
+theorem prims_okWith (v : ArithImpl) : PrimsOk (primsWith v.code) where
+ ntt := calleeOf v.ok.ntt
+ invNtt := calleeOf v.ok.invNtt
+ mul := calleeOf v.ok.mul
+ mulAdd := calleeOf v.ok.mulAdd
+ add := calleeOf v.ok.add
+ rejNtt := (⟨⟨16, by decide, Sample.rejNTT_verified⟩, nosp_of (by decide +kernel), by decide +kernel⟩ :
+ Callee prims.rejNtt _)
+ rejBounded := (⟨⟨16, by decide, Sample.rejBounded_verified⟩, nosp_of (by decide +kernel), by decide +kernel⟩ :
+ Callee prims.rejBounded _)
+ power2Round := (⟨⟨0, by decide, Round.power2Round_verified⟩, nosp_of (by decide +kernel), by decide +kernel⟩ :
+ Callee prims.power2Round _)
+ simpleBitPack := (⟨⟨0, by decide, Pack.simpleBitPack_verified⟩, nosp_of (by decide +kernel),
+ by decide +kernel⟩ : Callee prims.simpleBitPack _)
+ bitPack := (⟨⟨0, by decide, Pack.bitPack_verified⟩, nosp_of (by decide +kernel), by decide +kernel⟩ :
+ Callee prims.bitPack _)
+
+/-! For any implementation of the polynomial arithmetic, that key
+generation never writes `rsp` is checked by evaluating it with every
+function of it empty (`keyGen_same`, as `sign_same`). -/
+
+theorem keyGen_same {m mc : Prog isa → Bool} (hm : Comp m mc) {B : Backend} (h1 : mc B.ntt = true)
+ (h2 : mc B.invNtt = true) (h3 : mc B.mul = true) (h4 : mc B.mulAdd = true) (h5 : mc B.add = true)
+ (p : Spec.MlDsa.Params) : Same m (keyGen (primsWith B) p) (keyGen (primsWith .empty) p) := by
+ unfold keyGen
+ same_tac hm
+
+theorem keyGen0_sp {p : Spec.MlDsa.Params}
+ (hp : p = Spec.MlDsa.mlDsa44 ∨ p = Spec.MlDsa.mlDsa65 ∨ p = Spec.MlDsa.mlDsa87) :
+ (keyGen (primsWith .empty) p).allInstrs (fun i => !isa.writesSp i) = true := by
+ rcases hp with rfl | rfl | rfl <;> decide +kernel
+
+variable (v : ArithImpl) {p : Spec.MlDsa.Params}
+ (hp : p = Spec.MlDsa.mlDsa44 ∨ p = Spec.MlDsa.mlDsa65 ∨ p = Spec.MlDsa.mlDsa87)
+include hp
+
+theorem keyGen_spSafe : (keyGen (primsWith v.code) p).all (fun i => !isa.writesSp i) = true :=
+ Code.all_of_allInstrs (Same.ok (keyGen_same (Comp.all _) (Code.allInstrs_of_all v.ok.ntt.sp)
+ (Code.allInstrs_of_all v.ok.invNtt.sp) (Code.allInstrs_of_all v.ok.mul.sp)
+ (Code.allInstrs_of_all v.ok.mulAdd.sp) (Code.allInstrs_of_all v.ok.add.sp) p) (keyGen0_sp hp))
+
+theorem keyGen_verifiedWith :
+ Verified X86_64.target (keyGen (primsWith v.code) p) (Spec.MlDsa.keyGenContract p X86_64.abi 32) :=
+ keyGen_verified (prims_okWith v) p hp
end VG.Proof.MlDsa.X86_64.KeyGen
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/KeyGen/RestRow.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/KeyGen/RestRow.lean
index c1ff8245b..636a4b89d 100644
--- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/KeyGen/RestRow.lean
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/KeyGen/RestRow.lean
@@ -75,7 +75,7 @@ include hP hF hp hi
/-- `t = Â[i, 0] ŝ₁[0]`. -/
theorem mul_ok {A : Nat → Poly} {S : Nat → IPoly} {R : BitVec 64} {s : State} (h : KR p σ A S R (p.ℓ + p.k) p.ℓ i s) :
- WP isa (mulAt P.mul (tP p) (aP (p.ℓ * i)) (sP p 0)) s fun s' =>
+ WP isa (mulAt P.sfx P.mul (tP p) (aP (p.ℓ * i)) (sP p 0)) s fun s' =>
KR p σ A S R (p.ℓ + p.k) p.ℓ i s' ∧ tIs p (fun A S => dotK p A S i 1) A S s' := by
have hkl := hF.kl; have hl := hF.l; have hk := hF.k
have hx0 := idx_lt (j := 0) hi (by omega)
@@ -94,7 +94,7 @@ theorem mul_ok {A : Nat → Poly} {S : Nat → IPoly} {R : BitVec 64} {s : State
/-- `t = t + Â[i, j] ŝ₁[j]`. -/
theorem mulAdd_ok {j : Nat} (hj : j < p.ℓ) {A : Nat → Poly} {S : Nat → IPoly} {R : BitVec 64} {s : State}
(h : KR p σ A S R (p.ℓ + p.k) p.ℓ i s) (ht : tIs p (fun A S => dotK p A S i j) A S s) :
- WP isa (mulAddAt P.mulAdd (tP p) (aP (p.ℓ * i + j)) (sP p j)) s fun s' =>
+ WP isa (mulAddAt P.sfx P.mulAdd (tP p) (aP (p.ℓ * i + j)) (sP p j)) s fun s' =>
KR p σ A S R (p.ℓ + p.k) p.ℓ i s' ∧ tIs p (fun A S => dotK p A S i (j + 1)) A S s' := by
have hkl := hF.kl; have hl := hF.l; have hk := hF.k
dsimp only [tIs] at ht
@@ -112,7 +112,7 @@ theorem mulAdd_ok {j : Nat} (hj : j < p.ℓ) {A : Nat → Poly} {S : Nat → IPo
/-- `t = NTT⁻¹(t)`. -/
theorem inv_ok {A : Nat → Poly} {S : Nat → IPoly} {R : BitVec 64} {s : State}
(h : KR p σ A S R (p.ℓ + p.k) p.ℓ i s) (ht : tIs p (fun A S => dotK p A S i p.ℓ) A S s) :
- WP isa (invNttAt P.invNtt (tP p)) s fun s' =>
+ WP isa (invNttAt P.sfx P.invNtt (tP p)) s fun s' =>
KR p σ A S R (p.ℓ + p.k) p.ℓ i s' ∧ tIs p (fun A S => nttInv (dotK p A S i p.ℓ)) A S s' := by
have hkl := hF.kl; have hl := hF.l; have hk := hF.k
dsimp only [tIs] at ht
@@ -128,7 +128,7 @@ theorem inv_ok {A : Nat → Poly} {S : Nat → IPoly} {R : BitVec 64} {s : State
/-- `t = t + s₂[i]`. -/
theorem addS2_ok {A : Nat → Poly} {S : Nat → IPoly} {R : BitVec 64} {s : State}
(h : KR p σ A S R (p.ℓ + p.k) p.ℓ i s) (ht : tIs p (fun A S => nttInv (dotK p A S i p.ℓ)) A S s) :
- WP isa (addAt P.add (tP p) (sP p (p.ℓ + i))) s fun s' =>
+ WP isa (addAt P.sfx P.add (tP p) (sP p (p.ℓ + i))) s fun s' =>
KR p σ A S R (p.ℓ + p.k) p.ℓ i s' ∧ tIs p (fun A S => tK p A S i) A S s' := by
have hkl := hF.kl; have hl := hF.l; have hk := hF.k
dsimp only [tIs] at ht
@@ -226,7 +226,7 @@ variable {P : Prims} (hP : PrimsOk P) {p : Params} (hF : PFacts p) {i : Nat} (hi
include hP hF hi
theorem mul_piece : Piece p (KRx p (p.ℓ + p.k) p.ℓ i) (RowI p i (tIs p fun A S => dotK p A S i 1))
- (mulAt P.mul (tP p) (aP (p.ℓ * i)) (sP p 0)) := by
+ (mulAt P.sfx P.mul (tP p) (aP (p.ℓ * i)) (sP p 0)) := by
have hkl := hF.kl; have hl := hF.l; have hk := hF.k
have hx0 := idx_lt (j := 0) hi (by omega)
rw [Nat.add_zero] at hx0
@@ -244,7 +244,7 @@ theorem mul_piece : Piece p (KRx p (p.ℓ + p.k) p.ℓ i) (RowI p i (tIs p fun A
theorem mulAdd_piece {j : Nat} (hj : j < p.ℓ) :
Piece p (RowI p i (tIs p fun A S => dotK p A S i j)) (RowI p i (tIs p fun A S => dotK p A S i (j + 1)))
- (mulAddAt P.mulAdd (tP p) (aP (p.ℓ * i + j)) (sP p j)) := by
+ (mulAddAt P.sfx P.mulAdd (tP p) (aP (p.ℓ * i + j)) (sP p j)) := by
have hkl := hF.kl; have hl := hF.l; have hk := hF.k
have hx0 := idx_lt hi hj
refine ⟨fun _ _ hp ⟨A, S, R, h, ht⟩ => WP.mono (mulAdd_ok hP hF hp hi hj h ht) fun _ h => ⟨A, S, R, h⟩, ?_⟩
@@ -259,7 +259,7 @@ theorem mulAdd_piece {j : Nat} (hj : j < p.ℓ) :
(show Reg.rbx ∈ kgRegs by decide)
theorem inv_piece : Piece p (RowI p i (tIs p fun A S => dotK p A S i p.ℓ))
- (RowI p i (tIs p fun A S => nttInv (dotK p A S i p.ℓ))) (invNttAt P.invNtt (tP p)) := by
+ (RowI p i (tIs p fun A S => nttInv (dotK p A S i p.ℓ))) (invNttAt P.sfx P.invNtt (tP p)) := by
have hkl := hF.kl; have hl := hF.l; have hk := hF.k
refine ⟨fun _ _ hp ⟨A, S, R, h, ht⟩ => WP.mono (inv_ok hP hF hp hi h ht) fun _ h => ⟨A, S, R, h⟩, ?_⟩
refine rel_of (Q := fun x y => Two p x y ∧ Reduced x.mem (pa x (tP p)) ∧ Reduced y.mem (pa y (tP p))) ?_
@@ -268,7 +268,7 @@ theorem inv_piece : Piece p (RowI p i (tIs p fun A S => dotK p A S i p.ℓ))
exact ipAt_tr (tP_ok hF) (by lay) (by lay) (by lay) hP.invNtt (show Reg.rbx ∈ kgRegs by decide)
theorem addS2_piece : Piece p (RowI p i (tIs p fun A S => nttInv (dotK p A S i p.ℓ)))
- (RowI p i (tIs p fun A S => Proof.MlDsa.KeyGen.tK p A S i)) (addAt P.add (tP p) (sP p (p.ℓ + i))) := by
+ (RowI p i (tIs p fun A S => Proof.MlDsa.KeyGen.tK p A S i)) (addAt P.sfx P.add (tP p) (sP p (p.ℓ + i))) := by
have hkl := hF.kl; have hl := hF.l; have hk := hF.k
refine ⟨fun _ _ hp ⟨A, S, R, h, ht⟩ => WP.mono (addS2_ok hP hF hp hi h ht) fun _ h => ⟨A, S, R, h⟩, ?_⟩
refine rel_of (Q := fun x y => Two p x y ∧ (Reduced x.mem (pa x (tP p)) ∧ Reduced x.mem (pa x (sP p (p.ℓ + i)))) ∧
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Inst.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Inst.lean
index f88d2d5a0..0d9b4dcd2 100644
--- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Inst.lean
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Inst.lean
@@ -11,13 +11,15 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Pack.HintPack
import VerifiedGarbage.Proof.MlDsa.X86_64.Sample.RejNttCT
import VerifiedGarbage.Proof.MlDsa.X86_64.Sample.ExpandMask
import VerifiedGarbage.Proof.MlDsa.X86_64.Sample.BallCT
+import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Backend
/-!
# ML-DSA signing on x86-64: the primitives it calls
Untrusted: everything here is checked by Lean. The verified x86-64
implementations of the primitives (`prims`), and what the proofs of signing
-need of them (`prims_ok`), with 24 bytes of stack for each call: their
+need of them, with any implementation `v` of the polynomial arithmetic
+(`prims_okWith`), with 24 bytes of stack for each call: their
contracts, and, of the two samplers whose result signing branches on, that
it depends only on their public data and that they succeed only if the
algorithm finishes within `maxBounds` (from what their own proofs say they
@@ -30,6 +32,7 @@ open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Sign
open VG.Proof.MlDsa.Sign
open VG.Spec.MlDsa
open VG.Spec.Sha3 (bytesAt)
+open VG.Proof.MlDsa.X86_64 (FnOk ArithImpl)
/-- The x86-64 implementations of the primitives. -/
def prims : Prims where
@@ -51,6 +54,17 @@ def prims : Prims where
bitUnpack := Impl.MlDsa.X86_64.Pack.bitUnpack
hintBitPack := Impl.MlDsa.X86_64.Pack.hintBitPack
+/-- The primitives, with the polynomial arithmetic of `B`. -/
+def primsWith (B : Impl.MlDsa.X86_64.Arith.Backend) : Prims :=
+ { prims with
+ ntt := B.ntt
+ invNtt := B.invNtt
+ mul := B.mul
+ mulAdd := B.mulAdd
+ add := B.add
+ sub := B.sub
+ sfx := B.sfx }
+
theorem nosp_of {c : Prog isa} (h : c.allInstrs (fun i => !Taint.clobbers i .rsp) = true) : NoSp c := by
rw [Code.allInstrs_eq] at h
intro i hi
@@ -105,31 +119,47 @@ end
theorem one_ne_zero32 : (1 : BitVec 32) ≠ 0 := by decide
-/-- The primitives satisfy what the proofs of signing need of them. -/
-def prims_ok : PrimsOk prims signStack where
- ntt := ⟨0, by decide, Proof.MlDsa.X86_64.Arith.ntt_verified, nosp_of (by decide +kernel), by decide +kernel⟩
- invNtt := ⟨0, by decide, Proof.MlDsa.X86_64.Arith.nttInv_verified, nosp_of (by decide +kernel), by decide +kernel⟩
- mul := ⟨0, by decide, Proof.MlDsa.X86_64.Arith.mul_verified, nosp_of (by decide +kernel), by decide +kernel⟩
- mulAdd := ⟨0, by decide, Proof.MlDsa.X86_64.Arith.mulAdd_verified, nosp_of (by decide +kernel), by decide +kernel⟩
- add := ⟨0, by decide, Proof.MlDsa.X86_64.Arith.add_verified, nosp_of (by decide +kernel), by decide +kernel⟩
- sub := ⟨0, by decide, Proof.MlDsa.X86_64.Arith.sub_verified, nosp_of (by decide +kernel), by decide +kernel⟩
- rejNTT := ⟨16, by decide, Proof.MlDsa.X86_64.Sample.rejNTT_verified, nosp_of (by decide +kernel),
- by decide +kernel⟩
- expandMask := ⟨16, by decide, Proof.MlDsa.X86_64.Sample.expandMask_verified, nosp_of (by decide +kernel),
- by decide +kernel⟩
- ball := ⟨16, by decide, Proof.MlDsa.X86_64.Sample.sampleInBall_verified, nosp_of (by decide +kernel),
- by decide +kernel⟩
- highBits := ⟨0, by decide, Proof.MlDsa.X86_64.Round.highBits_verified, nosp_of (by decide +kernel), by decide +kernel⟩
- lowBits := ⟨0, by decide, Proof.MlDsa.X86_64.Round.lowBits_verified, nosp_of (by decide +kernel), by decide +kernel⟩
- normLt := ⟨0, by decide, Proof.MlDsa.X86_64.Round.normLt_verified, nosp_of (by decide +kernel), by decide +kernel⟩
- makeHint := ⟨0, by decide, Proof.MlDsa.X86_64.Round.makeHint_verified, nosp_of (by decide +kernel), by decide +kernel⟩
- simpleBitPack := ⟨0, by decide, Proof.MlDsa.X86_64.Pack.simpleBitPack_verified, nosp_of (by decide +kernel),
- by decide +kernel⟩
- bitPack := ⟨0, by decide, Proof.MlDsa.X86_64.Pack.bitPack_verified, nosp_of (by decide +kernel), by decide +kernel⟩
- bitUnpack := ⟨0, by decide, Proof.MlDsa.X86_64.Pack.bitUnpack_verified, nosp_of (by decide +kernel),
- by decide +kernel⟩
- hintBitPack := ⟨0, by decide, Proof.MlDsa.X86_64.Pack.hintBitPack_verified, nosp_of (by decide +kernel),
- by decide +kernel⟩
+/-- A function of the polynomial arithmetic satisfies what the proofs of signing need of it. -/
+def calleeOf {k : Nat → Contract isa} {c : Prog isa} (h : FnOk k c) : Callee k signStack c :=
+ ⟨0, by decide, h.ver, h.nosp, by have := h.depth; unfold signStack; omega⟩
+
+/-- The primitives, with the polynomial arithmetic of `v`, satisfy what the
+proofs of signing need of them. -/
+def prims_okWith (v : ArithImpl) : PrimsOk (primsWith v.code) signStack where
+ ntt := calleeOf v.ok.ntt
+ invNtt := calleeOf v.ok.invNtt
+ mul := calleeOf v.ok.mul
+ mulAdd := calleeOf v.ok.mulAdd
+ add := calleeOf v.ok.add
+ sub := calleeOf v.ok.sub
+ rejNTT := (⟨16, by decide, Proof.MlDsa.X86_64.Sample.rejNTT_verified, nosp_of (by decide +kernel),
+ by decide +kernel⟩ :
+ Callee _ signStack prims.rejNTT)
+ expandMask := (⟨16, by decide, Proof.MlDsa.X86_64.Sample.expandMask_verified, nosp_of (by decide +kernel),
+ by decide +kernel⟩ :
+ Callee _ signStack prims.expandMask)
+ ball := (⟨16, by decide, Proof.MlDsa.X86_64.Sample.sampleInBall_verified, nosp_of (by decide +kernel),
+ by decide +kernel⟩ :
+ Callee _ signStack prims.ball)
+ highBits := (⟨0, by decide, Proof.MlDsa.X86_64.Round.highBits_verified, nosp_of (by decide +kernel), by decide +kernel⟩ :
+ Callee _ signStack prims.highBits)
+ lowBits := (⟨0, by decide, Proof.MlDsa.X86_64.Round.lowBits_verified, nosp_of (by decide +kernel), by decide +kernel⟩ :
+ Callee _ signStack prims.lowBits)
+ normLt := (⟨0, by decide, Proof.MlDsa.X86_64.Round.normLt_verified, nosp_of (by decide +kernel), by decide +kernel⟩ :
+ Callee _ signStack prims.normLt)
+ makeHint := (⟨0, by decide, Proof.MlDsa.X86_64.Round.makeHint_verified, nosp_of (by decide +kernel), by decide +kernel⟩ :
+ Callee _ signStack prims.makeHint)
+ simpleBitPack := (⟨0, by decide, Proof.MlDsa.X86_64.Pack.simpleBitPack_verified, nosp_of (by decide +kernel),
+ by decide +kernel⟩ :
+ Callee _ signStack prims.simpleBitPack)
+ bitPack := (⟨0, by decide, Proof.MlDsa.X86_64.Pack.bitPack_verified, nosp_of (by decide +kernel), by decide +kernel⟩ :
+ Callee _ signStack prims.bitPack)
+ bitUnpack := (⟨0, by decide, Proof.MlDsa.X86_64.Pack.bitUnpack_verified, nosp_of (by decide +kernel),
+ by decide +kernel⟩ :
+ Callee _ signStack prims.bitUnpack)
+ hintBitPack := (⟨0, by decide, Proof.MlDsa.X86_64.Pack.hintBitPack_verified, nosp_of (by decide +kernel),
+ by decide +kernel⟩ :
+ Callee _ signStack prims.hintBitPack)
rejRet := fun s₁ s₂ t₁ t₂ s₁' s₂' ⟨h₁, h₂, hp⟩ e₁ e₂ =>
⟨Proof.MlDsa.X86_64.Sample.rejNTT_verified.2.1 s₁ s₂ t₁ t₂ s₁' s₂' h₁ h₂ hp e₁ e₂,
show _ = _ by rw [rn_ret h₁ e₁, rn_ret h₂ e₂, rn_pub s₁ s₂ hp]⟩
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean
index 089924f6d..3b2dc921e 100644
--- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean
@@ -1,10 +1,12 @@
import VerifiedGarbage.Proof.MlDsa.X86_64.Sign.Inst
+import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Same
/-!
# ML-DSA signing on x86-64: verified
Untrusted: everything here is checked by Lean. `vg_mldsa{44,65,87}_sign`
-(`sign prims p`) is verified against `signContractT`: `signContract` with
+(`sign (primsWith v.code) p`, for an implementation `v` of the polynomial
+arithmetic) is verified against `signContractT`: `signContract` with
`signLeakT` (`Proof/MlDsa/Sign/Leak.lean`) for `signLeak`, which tags what
each iteration of the loop leaks after its `c̃` with whether it was
rejected. The contract's `signLeak` tags the iterations the same way
@@ -18,6 +20,8 @@ open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Sign
open VG.Proof.MlDsa.Sign
open VG.Spec.MlDsa
open VG.Spec.Sha3 (bytesAt)
+open VG.Proof.MlDsa.X86_64 (Comp Same Same.ok ArithImpl Code.allInstrs_of_all)
+open VG.Impl.MlDsa.X86_64.Arith (Backend)
/-- `signContract`, with `signLeakT` for `signLeak`. -/
def signContractT (p : Params) {M : ISA} (A : Abi M) (stack : Nat := 0) : Contract M :=
@@ -62,52 +66,56 @@ theorem signK_implies {p : Params} (h3 : Ok3 p) :
· sig_implies_sat [signContractT, signSig, X86_64.abi, X86_64.argRegs] [signSat] using signSat mlDsa65
· sig_implies_sat [signContractT, signSig, X86_64.abi, X86_64.argRegs] [signSat] using signSat mlDsa87
-theorem sign_verified {p : Params} (h3 : Ok3 p)
- (hmx : ctlOk (Impl.MlDsa.X86_64.Sign.sign prims p) = true) :
- Verified X86_64.target (Impl.MlDsa.X86_64.Sign.sign prims p) (signContractT p X86_64.abi signStack) :=
- Verified.of_correct (sign_correct prims_ok h3 hmx) (sign_ct prims_ok h3) (signK_implies h3)
+/-! ## For any implementation of the polynomial arithmetic
-theorem sign44_verified :
- Verified X86_64.target (Impl.MlDsa.X86_64.Sign.sign prims mlDsa44) (signContractT mlDsa44 X86_64.abi signStack) :=
- sign_verified (.inl rfl) (by decide +kernel)
+A check that composes over the code (`Comp`) gives the same result on
+signing with the polynomial arithmetic `B` as with every function of it
+empty, if it holds of `B`'s functions (`sign_same`), so MXCSR (`sign_ctl`)
+and the stack pointer (`sign_spSafe`) are checked by evaluating signing
+with no implementation of it. -/
-theorem sign65_verified :
- Verified X86_64.target (Impl.MlDsa.X86_64.Sign.sign prims mlDsa65) (signContractT mlDsa65 X86_64.abi signStack) :=
- sign_verified (.inr (.inl rfl)) (by decide +kernel)
+theorem sign_same {m mc : Prog isa → Bool} (hm : Comp m mc) {B : Backend} (h1 : mc B.ntt = true)
+ (h2 : mc B.invNtt = true) (h3 : mc B.mul = true) (h4 : mc B.mulAdd = true) (h5 : mc B.add = true)
+ (h6 : mc B.sub = true) (p : Params) :
+ Same m (Impl.MlDsa.X86_64.Sign.sign (primsWith B) p) (Impl.MlDsa.X86_64.Sign.sign (primsWith .empty) p) := by
+ unfold Impl.MlDsa.X86_64.Sign.sign
+ same_tac hm
-theorem sign87_verified :
- Verified X86_64.target (Impl.MlDsa.X86_64.Sign.sign prims mlDsa87) (signContractT mlDsa87 X86_64.abi signStack) :=
- sign_verified (.inr (.inr rfl)) (by decide +kernel)
+theorem sign0_ctlC {p : Params} (h3 : Ok3 p) : ctlC (Impl.MlDsa.X86_64.Sign.sign (primsWith .empty) p) = true := by
+ rcases h3 with rfl | rfl | rfl <;> decide +kernel
-/-! Against the contract: `signContractT` is `signContract`, whose leakage
-tags each iteration as `signLeakT` does (`signLeakT_eq_signLeak`). -/
-
-theorem signContractT_eq (p : Params) {M : ISA} (A : Abi M) (stack : Nat) :
- signContractT p A stack = signContract p A stack := by
- unfold signContractT signContract
- simp only [Sign.signLeakT_eq_signLeak]
+theorem sign0_sp {p : Params} (h3 : Ok3 p) :
+ (Impl.MlDsa.X86_64.Sign.sign (primsWith .empty) p).allInstrs (fun i => !isa.writesSp i) = true := by
+ rcases h3 with rfl | rfl | rfl <;> decide +kernel
-theorem sign44_verified' :
- Verified X86_64.target (Impl.MlDsa.X86_64.Sign.sign prims mlDsa44) (signContract mlDsa44 X86_64.abi signStack) :=
- signContractT_eq mlDsa44 X86_64.abi signStack ▸ sign44_verified
+variable (v : ArithImpl) {p : Params} (h3 : Ok3 p)
+include h3
-theorem sign65_verified' :
- Verified X86_64.target (Impl.MlDsa.X86_64.Sign.sign prims mlDsa65) (signContract mlDsa65 X86_64.abi signStack) :=
- signContractT_eq mlDsa65 X86_64.abi signStack ▸ sign65_verified
+theorem sign_ctl : ctlOk (Impl.MlDsa.X86_64.Sign.sign (primsWith v.code) p) = true :=
+ ctlOk_of_ctlC (Same.ok (sign_same Comp.ctlC v.ok.ntt.ctl v.ok.invNtt.ctl v.ok.mul.ctl v.ok.mulAdd.ctl
+ v.ok.add.ctl v.ok.sub.ctl p) (sign0_ctlC h3))
-theorem sign87_verified' :
- Verified X86_64.target (Impl.MlDsa.X86_64.Sign.sign prims mlDsa87) (signContract mlDsa87 X86_64.abi signStack) :=
- signContractT_eq mlDsa87 X86_64.abi signStack ▸ sign87_verified
+theorem sign_spSafe : (Impl.MlDsa.X86_64.Sign.sign (primsWith v.code) p).all (fun i => !isa.writesSp i) = true :=
+ Code.all_of_allInstrs (Same.ok (sign_same (Comp.all _) (Code.allInstrs_of_all v.ok.ntt.sp)
+ (Code.allInstrs_of_all v.ok.invNtt.sp) (Code.allInstrs_of_all v.ok.mul.sp)
+ (Code.allInstrs_of_all v.ok.mulAdd.sp) (Code.allInstrs_of_all v.ok.add.sp)
+ (Code.allInstrs_of_all v.ok.sub.sp) p) (sign0_sp h3))
-/-! What registering them needs of the code besides: it never writes the stack pointer. -/
+theorem sign_verified :
+ Verified X86_64.target (Impl.MlDsa.X86_64.Sign.sign (primsWith v.code) p) (signContractT p X86_64.abi signStack) :=
+ Verified.of_correct (sign_correct (prims_okWith v) h3 (sign_ctl v h3)) (sign_ct (prims_okWith v) h3)
+ (signK_implies h3)
-theorem sign44_spSafe : (Impl.MlDsa.X86_64.Sign.sign prims mlDsa44).all (fun i => !X86_64.target.isa.writesSp i) = true :=
- Code.all_of_allInstrs (by decide +kernel)
-
-theorem sign65_spSafe : (Impl.MlDsa.X86_64.Sign.sign prims mlDsa65).all (fun i => !X86_64.target.isa.writesSp i) = true :=
- Code.all_of_allInstrs (by decide +kernel)
+omit h3 in
+/-- Against the contract: `signContractT` is `signContract`, whose leakage
+tags each iteration as `signLeakT` does (`signLeakT_eq_signLeak`). -/
+theorem signContractT_eq (p : Params) {M : ISA} (A : Abi M) (stack : Nat) :
+ signContractT p A stack = signContract p A stack := by
+ unfold signContractT signContract
+ simp only [Sign.signLeakT_eq_signLeak]
-theorem sign87_spSafe : (Impl.MlDsa.X86_64.Sign.sign prims mlDsa87).all (fun i => !X86_64.target.isa.writesSp i) = true :=
- Code.all_of_allInstrs (by decide +kernel)
+theorem sign_verified' :
+ Verified X86_64.target (Impl.MlDsa.X86_64.Sign.sign (primsWith v.code) p) (signContract p X86_64.abi signStack) :=
+ signContractT_eq p X86_64.abi signStack ▸ sign_verified v h3
end VG.Proof.MlDsa.X86_64.Sign
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Instrs.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Instrs.lean
index ce8c82e43..852e3e7a1 100644
--- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Instrs.lean
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Instrs.lean
@@ -1,5 +1,6 @@
import VerifiedGarbage.Proof.MlDsa.X86_64.Verify.PrimsOk
import VerifiedGarbage.Impl.MlDsa.X86_64.Verify.Verify
+import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Same
/-!
# ML-DSA verification on x86-64: properties of every instruction
@@ -20,7 +21,7 @@ open VG.Spec.MlDsa
/-- The primitives, each empty. -/
def P0 : Prims := ⟨.block [], .block [], .block [], .block [], .block [], .block [], .block [], .block [], .block [],
- .block [], .block [], .block [], .block []⟩
+ .block [], .block [], .block [], .block [], ""⟩
/-- `q` holds of every instruction of the primitives `P`. -/
structure PrimsQ (q : Instr → Bool) (P : Prims) : Prop where
@@ -49,8 +50,8 @@ theorem SameQ.seq {a a' b b' : Prog isa} (ha : SameQ q a a') (hb : SameQ q b b')
show (a.allInstrs q && b.allInstrs q) = (a'.allInstrs q && b'.allInstrs q)
rw [show a.allInstrs q = a'.allInstrs q from ha, show b.allInstrs q = b'.allInstrs q from hb]
-theorem SameQ.call {c : Prog isa} (hc : c.allInstrs q = true) (n : String) (as : List (Reg × Arg)) :
- SameQ q (callAt n c as) (callAt n (.block []) as) := by
+theorem SameQ.call {c : Prog isa} (hc : c.allInstrs q = true) {n n' : String} (as : List (Reg × Arg)) :
+ SameQ q (callAt n c as) (callAt n' (.block []) as) := by
show (_ && c.allInstrs q) = (_ && true)
rw [hc]
@@ -70,30 +71,30 @@ theorem SameQ.sampled {c c' : Prog isa} (h : SameQ q c c') (a : Ptr) : SameQ q (
variable {P : Prims} (hP : PrimsQ q P) (p : Params)
include hP
-theorem hint_q : SameQ q (hint P p) (hint P0 p) := (SameQ.call hP.hintUnpack _ _).seq rfl
+theorem hint_q : SameQ q (hint P p) (hint P0 p) := (SameQ.call hP.hintUnpack _).seq rfl
theorem zOne_q (i : Nat) : SameQ q (zOne P p i) (zOne P0 p i) :=
- (SameQ.call hP.bitUnpack _ _).seq ((SameQ.call hP.normLt _ _).seq rfl)
+ (SameQ.call hP.bitUnpack _).seq ((SameQ.call hP.normLt _).seq rfl)
omit p in
theorem aOne_q (e : Nat) : SameQ q (aOne P e) (aOne P0 e) :=
- SameQ.seq rfl (SameQ.sampled (SameQ.call hP.rejNtt _ _) _)
+ SameQ.seq rfl (SameQ.sampled (SameQ.call hP.rejNtt _) _)
theorem aRow_q (r : Nat) : SameQ q (aRow P p r) (aRow P0 p r) := SameQ.seqR (aOne_q hP) _ _
theorem samples_q : SameQ q (samples P p) (samples P0 p) :=
- SameQ.seq rfl ((SameQ.seqR (aRow_q hP p) _ _).seq (SameQ.sampled (SameQ.call hP.ball _ _) _))
+ SameQ.seq rfl ((SameQ.seqR (aRow_q hP p) _ _).seq (SameQ.sampled (SameQ.call hP.ball _) _))
theorem dot_q (r : Nat) : SameQ q (dot P p r) (dot P0 p r) :=
- (SameQ.call hP.mul _ _).seq (SameQ.seqR (fun _ => SameQ.call hP.mulAdd _ _) _ _)
+ (SameQ.call hP.mul _).seq (SameQ.seqR (fun _ => SameQ.call hP.mulAdd _) _ _)
theorem row_q (r : Nat) : SameQ q (row P p r) (row P0 p r) :=
- (dot_q hP p r).seq ((SameQ.call hP.unpackT1 _ _).seq ((SameQ.call hP.ntt _ _).seq ((SameQ.call hP.mul _ _).seq
- ((SameQ.call hP.sub _ _).seq ((SameQ.call hP.invNtt _ _).seq ((SameQ.call hP.useHint _ _).seq
- (SameQ.call hP.simpleBitPack _ _)))))))
+ (dot_q hP p r).seq ((SameQ.call hP.unpackT1 _).seq ((SameQ.call hP.ntt _).seq ((SameQ.call hP.mul _).seq
+ ((SameQ.call hP.sub _).seq ((SameQ.call hP.invNtt _).seq ((SameQ.call hP.useHint _).seq
+ (SameQ.call hP.simpleBitPack _)))))))
theorem compute_q : SameQ q (compute P p) (compute P0 p) :=
- (SameQ.seqR (fun _ => SameQ.call hP.ntt _ _) _ _).seq ((SameQ.call hP.ntt _ _).seq
+ (SameQ.seqR (fun _ => SameQ.call hP.ntt _) _ _).seq ((SameQ.call hP.ntt _).seq
((SameQ.seqR (row_q hP p) _ _).seq rfl))
theorem verify_q : SameQ q (verify P p) (verify P0 p) :=
@@ -102,12 +103,6 @@ theorem verify_q : SameQ q (verify P p) (verify P0 p) :=
end
-theorem Code.allInstrs_of_all {I C : Type} {q : I → Bool} {c : Code I C} (h : c.all q = true) :
- c.allInstrs q = true := by
- induction c with
- | block is => induction is <;> simp_all [Code.all, Code.allInstrs]
- | _ => simp_all [Code.all, Code.allInstrs]
-
theorem verify0_sp : ∀ p ∈ params, (verify P0 p).allInstrs (fun i => !isa.writesSp i) = true := by
decide +kernel
@@ -136,8 +131,8 @@ theorem SameC.seq {a a' b b' : Prog isa} (ha : SameC a a') (hb : SameC b b') : S
show (ctlC a && ctlC b) = (ctlC a' && ctlC b')
rw [show ctlC a = ctlC a' from ha, show ctlC b = ctlC b' from hb]
-theorem SameC.call {c : Prog isa} (hc : ctlOk c = true) (n : String) (as : List (Reg × Arg)) :
- SameC (callAt n c as) (callAt n (.block []) as) := by
+theorem SameC.call {c : Prog isa} (hc : ctlOk c = true) {n n' : String} (as : List (Reg × Arg)) :
+ SameC (callAt n c as) (callAt n' (.block []) as) := by
show (_ && ctlOk c) = (_ && true)
rw [hc]
@@ -159,22 +154,22 @@ include hP
theorem verify_c : SameC (verify P p) (verify P0 p) := by
have aOne : ∀ e, SameC (aOne P e) (aOne P0 e) := fun e =>
- SameC.seq rfl (SameC.sampled (SameC.call hP.rejNtt _ _) _)
+ SameC.seq rfl (SameC.sampled (SameC.call hP.rejNtt _) _)
have dot : ∀ r, SameC (dot P p r) (dot P0 p r) := fun r =>
- (SameC.call hP.mul _ _).seq (SameC.seqR (fun _ => SameC.call hP.mulAdd _ _) _ _)
+ (SameC.call hP.mul _).seq (SameC.seqR (fun _ => SameC.call hP.mulAdd _) _ _)
have row : ∀ r, SameC (row P p r) (row P0 p r) := fun r =>
- (dot r).seq ((SameC.call hP.unpackT1 _ _).seq ((SameC.call hP.ntt _ _).seq ((SameC.call hP.mul _ _).seq
- ((SameC.call hP.sub _ _).seq ((SameC.call hP.invNtt _ _).seq ((SameC.call hP.useHint _ _).seq
- (SameC.call hP.simpleBitPack _ _)))))))
+ (dot r).seq ((SameC.call hP.unpackT1 _).seq ((SameC.call hP.ntt _).seq ((SameC.call hP.mul _).seq
+ ((SameC.call hP.sub _).seq ((SameC.call hP.invNtt _).seq ((SameC.call hP.useHint _).seq
+ (SameC.call hP.simpleBitPack _)))))))
have samples : SameC (samples P p) (samples P0 p) :=
SameC.seq rfl ((SameC.seqR (fun r => SameC.seqR aOne _ _) _ _).seq
- (SameC.sampled (SameC.call hP.ball _ _) _))
+ (SameC.sampled (SameC.call hP.ball _) _))
have compute : SameC (compute P p) (compute P0 p) :=
- (SameC.seqR (fun _ => SameC.call hP.ntt _ _) _ _).seq ((SameC.call hP.ntt _ _).seq
+ (SameC.seqR (fun _ => SameC.call hP.ntt _) _ _).seq ((SameC.call hP.ntt _).seq
((SameC.seqR row _ _).seq rfl))
have zOne : ∀ i, SameC (zOne P p i) (zOne P0 p i) := fun _ =>
- (SameC.call hP.bitUnpack _ _).seq ((SameC.call hP.normLt _ _).seq rfl)
- exact SameC.seq rfl ((((SameC.call hP.hintUnpack _ _).seq rfl).seq (SameC.ifOk ((SameC.seqR zOne _ _).seq
+ (SameC.call hP.bitUnpack _).seq ((SameC.call hP.normLt _).seq rfl)
+ exact SameC.seq rfl ((((SameC.call hP.hintUnpack _).seq rfl).seq (SameC.ifOk ((SameC.seqR zOne _ _).seq
(SameC.ifOk (samples.seq compute))))).seq rfl)
end
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean
index 829f96384..8d87a8c4d 100644
--- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean
@@ -11,6 +11,7 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Round.UseHint
import VerifiedGarbage.Proof.MlDsa.X86_64.Pack.SimpleBitPack
import VerifiedGarbage.Proof.MlDsa.X86_64.Pack.Unpack
import VerifiedGarbage.Proof.MlDsa.X86_64.Pack.HintUnpack
+import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Backend
/-!
# ML-DSA verification on x86-64: the primitives it calls
@@ -18,13 +19,15 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Pack.HintUnpack
Untrusted: everything here is checked by Lean. The x86-64 implementations
of the primitives (`prims`) meet their contracts with at most 16 bytes of
stack, never write the stack pointer or load MXCSR, and call at most two
-deep (`prims_ok`), so `verify prims p` meets `verifyContract p`.
+deep, with any implementation `v` of the polynomial arithmetic
+(`prims_okWith`), so `verify (primsWith v.code) p` meets `verifyContract p`.
-/
namespace VG.Proof.MlDsa.X86_64.Verify
open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Verify
open VG.Impl.MlDsa.X86_64
+open VG.Proof.MlDsa.X86_64 (FnOk ArithImpl)
/-- The x86-64 implementations of the primitives. -/
def prims : Prims where
@@ -42,54 +45,72 @@ def prims : Prims where
hintUnpack := Pack.hintBitUnpack
normLt := Round.normLt
-theorem prims_ok : PrimsOk prims where
+/-- The primitives, with the polynomial arithmetic of `B`. -/
+def primsWith (B : Arith.Backend) : Prims :=
+ { prims with
+ ntt := B.ntt
+ invNtt := B.invNtt
+ mul := B.mul
+ mulAdd := B.mulAdd
+ sub := B.sub
+ sfx := B.sfx }
+
+/-- A function of the polynomial arithmetic satisfies what the proofs of verification need of it. -/
+theorem calleeOf {sig : Sig} {pre : Curry (sig.words X86_64.abi.ptrBits) (Mem → Prop)}
+ {post : sig.Post X86_64.abi.ptrBits} {wa : Bool} {c : Prog isa}
+ (h : FnOk (fun S => sig.contract X86_64.abi pre post wa S none) c) :
+ CalleeOk c (sig.contract X86_64.abi pre post wa 16 none) :=
+ CalleeOk.of_verified h.ver (by decide) h.nosp h.depth h.ctl h.sp
+
+theorem prims_okWith (v : ArithImpl) : PrimsOk (primsWith v.code) where
ntt := by
- have h := Proof.MlDsa.X86_64.Arith.ntt_verified
+ have h := v.ok.ntt
unfold Spec.MlDsa.nttContract Spec.MlDsa.inPlaceContract at h ⊢
- exact CalleeOk.of_verified h (by decide) (Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide)
- (by lit_decide) (Code.all_of_allInstrs (by lit_decide))
+ exact calleeOf h
invNtt := by
- have h := Proof.MlDsa.X86_64.Arith.nttInv_verified
+ have h := v.ok.invNtt
unfold Spec.MlDsa.nttInvContract Spec.MlDsa.inPlaceContract at h ⊢
- exact CalleeOk.of_verified h (by decide) (Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide)
- (by lit_decide) (Code.all_of_allInstrs (by lit_decide))
- mul := CalleeOk.of_verified Proof.MlDsa.X86_64.Arith.mul_verified (by decide)
- (Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide)
- (Code.all_of_allInstrs (by lit_decide))
- mulAdd := CalleeOk.of_verified Proof.MlDsa.X86_64.Arith.mulAdd_verified (by decide)
- (Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide)
- (Code.all_of_allInstrs (by lit_decide))
- sub := CalleeOk.of_verified Proof.MlDsa.X86_64.Arith.sub_verified (by decide)
- (Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide)
- (Code.all_of_allInstrs (by lit_decide))
- rejNtt := CalleeOk.of_verified Proof.MlDsa.X86_64.Sample.rejNTT_verified (by decide)
+ exact calleeOf h
+ mul := calleeOf v.ok.mul
+ mulAdd := calleeOf v.ok.mulAdd
+ sub := calleeOf v.ok.sub
+ rejNtt := (CalleeOk.of_verified Proof.MlDsa.X86_64.Sample.rejNTT_verified (by decide)
(Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide)
- (Code.all_of_allInstrs (by lit_decide))
- ball := CalleeOk.of_verified Proof.MlDsa.X86_64.Sample.sampleInBall_verified (by decide)
+ (Code.all_of_allInstrs (by lit_decide)) :
+ CalleeOk prims.rejNtt _)
+ ball := (CalleeOk.of_verified Proof.MlDsa.X86_64.Sample.sampleInBall_verified (by decide)
(Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide)
- (Code.all_of_allInstrs (by lit_decide))
- useHint := CalleeOk.of_verified Proof.MlDsa.X86_64.Round.useHint_verified (by decide)
+ (Code.all_of_allInstrs (by lit_decide)) :
+ CalleeOk prims.ball _)
+ useHint := (CalleeOk.of_verified Proof.MlDsa.X86_64.Round.useHint_verified (by decide)
(Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide)
- (Code.all_of_allInstrs (by lit_decide))
- simpleBitPack := CalleeOk.of_verified Proof.MlDsa.X86_64.Pack.simpleBitPack_verified (by decide)
+ (Code.all_of_allInstrs (by lit_decide)) :
+ CalleeOk prims.useHint _)
+ simpleBitPack := (CalleeOk.of_verified Proof.MlDsa.X86_64.Pack.simpleBitPack_verified (by decide)
(Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide)
- (Code.all_of_allInstrs (by lit_decide))
- bitUnpack := CalleeOk.of_verified Proof.MlDsa.X86_64.Pack.bitUnpack_verified (by decide)
+ (Code.all_of_allInstrs (by lit_decide)) :
+ CalleeOk prims.simpleBitPack _)
+ bitUnpack := (CalleeOk.of_verified Proof.MlDsa.X86_64.Pack.bitUnpack_verified (by decide)
(Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide)
- (Code.all_of_allInstrs (by lit_decide))
- unpackT1 := CalleeOk.of_verified Proof.MlDsa.X86_64.Pack.unpackT1_verified (by decide)
+ (Code.all_of_allInstrs (by lit_decide)) :
+ CalleeOk prims.bitUnpack _)
+ unpackT1 := (CalleeOk.of_verified Proof.MlDsa.X86_64.Pack.unpackT1_verified (by decide)
(Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide)
- (Code.all_of_allInstrs (by lit_decide))
- hintUnpack := CalleeOk.of_verified Proof.MlDsa.X86_64.Pack.hintBitUnpack_verified (by decide)
+ (Code.all_of_allInstrs (by lit_decide)) :
+ CalleeOk prims.unpackT1 _)
+ hintUnpack := (CalleeOk.of_verified Proof.MlDsa.X86_64.Pack.hintBitUnpack_verified (by decide)
(Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide)
- (Code.all_of_allInstrs (by lit_decide))
- normLt := CalleeOk.of_verified Proof.MlDsa.X86_64.Round.normLt_verified (by decide)
+ (Code.all_of_allInstrs (by lit_decide)) :
+ CalleeOk prims.hintUnpack _)
+ normLt := (CalleeOk.of_verified Proof.MlDsa.X86_64.Round.normLt_verified (by decide)
(Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide)
- (Code.all_of_allInstrs (by lit_decide))
+ (Code.all_of_allInstrs (by lit_decide)) :
+ CalleeOk prims.normLt _)
-/-- `vg_mldsa*_verify` for the parameter set `p`, calling the x86-64 primitives. -/
-theorem verify_prims {p : Spec.MlDsa.Params} (hp : p ∈ params) :
- Verified X86_64.target (verify prims p) (Spec.MlDsa.verifyContract p X86_64.abi 24) :=
- verify_verified prims_ok hp
+/-- `vg_mldsa*_verify` for the parameter set `p`, calling the x86-64
+primitives, with the polynomial arithmetic of `v`. -/
+theorem verify_prims (v : ArithImpl) {p : Spec.MlDsa.Params} (hp : p ∈ params) :
+ Verified X86_64.target (verify (primsWith v.code) p) (Spec.MlDsa.verifyContract p X86_64.abi 24) :=
+ verify_verified (prims_okWith v) hp
end VG.Proof.MlDsa.X86_64.Verify
diff --git a/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Sse2.lean b/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Sse2.lean
new file mode 100644
index 000000000..51dccc5dd
--- /dev/null
+++ b/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Sse2.lean
@@ -0,0 +1,16 @@
+import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Backend
+
+/-!
+# ML-DSA's polynomial arithmetic on x86-64: SSE2
+
+A variant of `MlDsaArith` on x86-64 (see `TCB/Emit.lean`): `vg_mldsa_ntt`,
+`vg_mldsa_inv_ntt`, `vg_mldsa_multiply_ntt`, `vg_mldsa_multiply_add_ntt`,
+`vg_mldsa_add` and `vg_mldsa_sub`, in the baseline ISA (SSE2), which key
+generation, signing and verification call.
+-/
+
+namespace VG.Variants.MlDsaArith.X86_64.Sse2
+
+def variant : Proof.MlDsa.X86_64.ArithImpl := .sse2
+
+end VG.Variants.MlDsaArith.X86_64.Sse2
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",