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