Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
41 commits
Select commit Hold shift + click to select a range
c8a8986
ML-DSA on x86-64: NTT and NTT⁻¹ in SSE2
claude Oct 1, 2026
e49e9bb
Merge remote-tracking branch 'origin/main' into pr1-merge
claude Oct 1, 2026
ae7eb79
ML-DSA on x86-64: polynomial arithmetic in SSE2
claude Oct 1, 2026
04e1ac6
Merge remote-tracking branch 'origin/claude/fervent-einstein-ukl7t7' …
claude Oct 1, 2026
6703c7a
ML-DSA on x86-64: key generation, signing and verification generic ov…
claude Oct 1, 2026
3e2aa65
Merge origin/main into claude/fervent-einstein-ukl7t7-mul
claude Oct 1, 2026
55761ef
Merge claude/fervent-einstein-ukl7t7-mul (with origin/main) into clau…
claude Oct 1, 2026
791054f
ML-DSA on x86-64: polynomial arithmetic in AVX2
claude Oct 1, 2026
b818911
Spec: ML-DSA's RejNTTPoly four times (vg_mldsa_rej_ntt_poly4)
claude Oct 1, 2026
779ca16
Spec: give vg_mldsa_rej_ntt_poly4 8 KiB of scratch
claude Oct 1, 2026
e365dde
CI: test ML-DSA on the CPUs of the x86-64 feature matrix
claude Oct 1, 2026
e04e23b
Merge branch 'claude/fervent-einstein-ukl7t7-rej4spec' into claude/fe…
claude Oct 1, 2026
be2f80c
ML-DSA on x86-64: sample  four entries at a time in verification
claude Oct 1, 2026
9a8bde9
ML-DSA on x86-64: sample  four entries at a time in key generation
claude Oct 1, 2026
f8a40e9
ML-DSA on x86-64: sample  four entries at a time in signing
claude Oct 1, 2026
856f4bb
ML-DSA on x86-64: HighBits and LowBits with AVX2
claude Oct 1, 2026
4563920
Merge main into claude/fervent-einstein-ukl7t7-avx2
claude Oct 1, 2026
11c34b1
Merge branch 'claude/fervent-einstein-ukl7t7-avx2' into claude/ferven…
claude Oct 1, 2026
8160f69
Merge branch 'claude/fervent-einstein-ukl7t7-avx2' into claude/ferven…
claude Oct 1, 2026
2c38b92
Merge branch 'claude/fervent-einstein-ukl7t7-rej4' into claude/ferven…
claude Oct 1, 2026
1c58e17
Merge branch 'claude/fervent-einstein-ukl7t7-rej4kg' into claude/ferv…
claude Oct 1, 2026
6819894
Merge main into claude/fervent-einstein-ukl7t7-avx2
claude Oct 1, 2026
ac45439
Merge branch 'claude/fervent-einstein-ukl7t7-avx2' into claude/ferven…
claude Oct 1, 2026
3f51d68
Merge branch 'claude/fervent-einstein-ukl7t7-avx2' into claude/ferven…
claude Oct 1, 2026
3d168fc
Merge branch 'claude/fervent-einstein-ukl7t7-rej4' into claude/ferven…
claude Oct 1, 2026
5803118
Merge branch 'claude/fervent-einstein-ukl7t7-rej4kg' into claude/ferv…
claude Oct 1, 2026
8cc896c
Merge main into claude/fervent-einstein-ukl7t7-rej4
claude Oct 2, 2026
b482411
Merge main into claude/fervent-einstein-ukl7t7-ybits2
claude Oct 2, 2026
850ee16
Merge branch 'claude/fervent-einstein-ukl7t7-rej4' into claude/ferven…
claude Oct 2, 2026
17760ea
Merge branch 'claude/fervent-einstein-ukl7t7-rej4kg' into claude/ferv…
claude Oct 2, 2026
db72deb
Merge branch 'claude/fervent-einstein-ukl7t7-rej4' into claude/ferven…
claude Oct 2, 2026
c3926b2
Merge main into claude/fervent-einstein-ukl7t7-ybits2
claude Oct 2, 2026
3c23c2a
Merge main into claude/fervent-einstein-ukl7t7-rej4kg
claude Oct 2, 2026
65fc75b
Merge branch 'claude/fervent-einstein-ukl7t7-rej4kg' into claude/ferv…
claude Oct 2, 2026
199d9b3
Merge main into claude/fervent-einstein-ukl7t7-ybits2
claude Oct 2, 2026
88c4d9a
ML-DSA on x86-64: the norm check and MakeHint with AVX2 (sign −8% ins…
alex Oct 2, 2026
2e45fea
Merge #519 (already on the remote branch) with main
claude Oct 2, 2026
e6062e4
Merge main into claude/fervent-einstein-ukl7t7-rej4sg
claude Oct 2, 2026
201fb3b
Merge #513 into claude/fervent-einstein-ukl7t7-rej4sg
claude Oct 2, 2026
f0998e4
Merge main into claude/fervent-einstein-ukl7t7-rej4sg
claude Oct 2, 2026
80ac2dc
Merge main into claude/fervent-einstein-ukl7t7-rej4sg
claude Oct 2, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -923,7 +923,7 @@ yours to keep:

<td>✅</td>

<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)</td>
<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time (four SHAKE128 instances at once with AVX2)</td>

<td>✅ SHA extensions</td>

Expand All @@ -939,7 +939,7 @@ yours to keep:

<td>✅</td>

<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)</td>
<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time (four SHAKE128 instances at once with AVX2)</td>

<td>✅ SHA extensions</td>

Expand All @@ -955,7 +955,7 @@ yours to keep:

<td>✅</td>

<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)</td>
<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time (four SHAKE128 instances at once with AVX2)</td>

<td>✅ SHA extensions</td>

Expand Down
2 changes: 1 addition & 1 deletion docs/algorithms/ml-dsa-44.toml
Original file line number Diff line number Diff line change
Expand Up @@ -3,4 +3,4 @@ family = "Signatures"
specs = ["MlDsa"]
modules = ["src/mldsa44.rs"]
asm = ["mldsa44", "mldsa"]
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" }
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time (four SHAKE128 instances at once with AVX2)" }
2 changes: 1 addition & 1 deletion docs/algorithms/ml-dsa-65.toml
Original file line number Diff line number Diff line change
Expand Up @@ -3,4 +3,4 @@ family = "Signatures"
specs = ["MlDsa"]
modules = ["src/mldsa65.rs"]
asm = ["mldsa65", "mldsa"]
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" }
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time (four SHAKE128 instances at once with AVX2)" }
2 changes: 1 addition & 1 deletion docs/algorithms/ml-dsa-87.toml
Original file line number Diff line number Diff line change
Expand Up @@ -3,4 +3,4 @@ family = "Signatures"
specs = ["MlDsa"]
modules = ["src/mldsa87.rs"]
asm = ["mldsa87", "mldsa"]
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" }
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time (four SHAKE128 instances at once with AVX2)" }
14 changes: 7 additions & 7 deletions lean/VerifiedGarbage/Generic/MlDsaArith/X86_64/MlDsaSign.lean
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@ 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 \
["The function saves its caller's callee-saved registers in `scratch`; its calls use the 32 \
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 \
Expand All @@ -37,8 +37,8 @@ def artifacts (v : ArithImpl) : List Artifact := [
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
contract := Spec.MlDsa.signContract Spec.MlDsa.mlDsa44 X86_64.abi 32
stack := 32
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
Expand All @@ -47,8 +47,8 @@ def artifacts (v : ArithImpl) : List Artifact := [
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
contract := Spec.MlDsa.signContract Spec.MlDsa.mlDsa65 X86_64.abi 32
stack := 32
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
Expand All @@ -57,8 +57,8 @@ def artifacts (v : ArithImpl) : List Artifact := [
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
contract := Spec.MlDsa.signContract Spec.MlDsa.mlDsa87 X86_64.abi 32
stack := 32
verified := Proof.MlDsa.X86_64.Sign.sign_verified' v (.inr (.inr rfl))
spSafe := Proof.MlDsa.X86_64.Sign.sign_spSafe v (.inr (.inr rfl)) }]

Expand Down
12 changes: 11 additions & 1 deletion lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,8 @@ structure Prims where
bitPack : Prog isa
bitUnpack : Prog isa
hintBitPack : Prog isa
/-- `vg_mldsa_rej_ntt_poly4` -/
rej4 : Prog isa
/-- What the names of the polynomial arithmetic's functions end with (`Arith.Backend`). -/
sfx : String := ""

Expand All @@ -54,7 +56,8 @@ at 200 (640 bytes); the caller's callee-saved registers at 840 (48 bytes);
the iterations left (`CNT`), the counter `κ` (`KAP`) and the number of 1s
of the hint (`ONES`) at 888, 896 and 904 (8 bytes each); the seed of
`RejNTTPoly` (`RS`, 34 bytes) at 912; the seed of `ExpandMask` (`MS`,
`ρ″` and two bytes) at 960; `c̃` (`CT`, up to 64 bytes) at 1040; the
`ρ″` and two bytes) at 960; `c̃` (`CT`, up to 64 bytes) at 1040; the four
seeds of `vg_mldsa_rej_ntt_poly4` (`RS4`, 136 bytes) at 1152; the
encoding of `w₁` (`W1`, up to 1024 bytes) at 2048; the working space of
the primitives (`PS`, 2048 bytes) at 3072; and polynomials of 1024 bytes
from 5120 (`P i`). -/
Expand All @@ -66,6 +69,7 @@ def oONES : Nat := 904
def oRS : Nat := 912
def oMS : Nat := 960
def oCT : Nat := 1040
def oRS4 : Nat := 1152
def oW1 : Nat := 2048
def oPS : Nat := 3072
/-- Polynomial `i`. -/
Expand Down Expand Up @@ -184,6 +188,12 @@ def rejAt (a : Ptr) : Prog isa :=
.seq (callP "vg_mldsa_rej_ntt_poly" P.rejNTT [.ptr (sc oRS), .ptr a, .ptr (sc oPS)])
(.block [.alu32 .and .r15 (.reg .rax)])

/-- `RejNTTPoly` of the four seeds at `RS4` to the four polynomials from `a`, with the working
space `w`, and `r15 ← r15 ∧ result`. -/
def rej4At (a w : Ptr) : Prog isa :=
.seq (callP ("vg_mldsa_rej_ntt_poly4" ++ P.sfx) P.rej4 [.ptr (sc oRS4), .ptr a, .ptr w])
(.block [.alu32 .and .r15 (.reg .rax)])

/-- A polynomial of `ExpandMask` from the seed at `MS` to `a`. -/
def maskAt (gamma1 : Nat) (a : Ptr) : Prog isa :=
callP "vg_mldsa_expand_mask_poly" P.expandMask [.ptr (sc oMS), .imm gamma1, .ptr a, .ptr (sc oPS)]
Expand Down
37 changes: 30 additions & 7 deletions lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Sign.lean
Original file line number Diff line number Diff line change
Expand Up @@ -13,8 +13,11 @@ keeps `scratch` in `rbx`, `sk` in `rbp`, `mu` in `r12`, `rnd` in `r13` and
`scratch`.

1. `ρ` (the first 32 bytes of `sk`) to `RS`, the seed of `RejNTTPoly`, and
`Â[r, s] = RejNTTPoly(ρ ‖ s ‖ r)` for the `kℓ` entries, with `r15` the
AND of the results. If one failed (`r15 = 0`), it returns 0 at once.
to each of the four seeds at `RS4`, and `Â[r, s] = RejNTTPoly(ρ ‖ s ‖ r)`
for the `kℓ` entries: four consecutive entries at a time
(`vg_mldsa_rej_ntt_poly4`, each seed at `RS4` with its entry's indices),
then the last `kℓ mod 4` one at a time, with `r15` the AND of the
results. If one failed (`r15 = 0`), it returns 0 at once.
2. `ŝ₁`, `ŝ₂` and `t̂₀`: the `NTT` of the `BitUnpack` of their pieces of
`sk`; and `ρ″ = H(K ‖ rnd ‖ μ, 64)` to `MS`.
3. The rejection sampling loop, at most 814 iterations (`minBounds.sign`),
Expand All @@ -29,8 +32,8 @@ keeps `scratch` in `rbx`, `sk` in `rbp`, `mu` in `r12`, `rnd` in `r13` and
iterations.
4. If `r15 = 1`: `c̃`, the `BitPack` of `z` and `HintBitPack(h)` to `sig`.

Only the calls of `vg_mldsa_rej_ntt_poly` (whose seeds are `ρ` and two
indices), and the branch on their results, depend on `ρ`; only the calls of
Only the calls of `vg_mldsa_rej_ntt_poly` and `vg_mldsa_rej_ntt_poly4`
(whose seeds are `ρ` and two indices), and the branch on their results, depend on `ρ`; only the calls of
`vg_mldsa_sample_in_ball`, and the branch on their results, on `c̃`; only
the branch on the validity checks on whether they passed, and only the call
of `vg_mldsa_hint_bit_pack` on the hint of the signature. Every other
Expand Down Expand Up @@ -71,7 +74,8 @@ abbrev sigH : Nat := cLen p + zLen p * p.ℓ
/-! ## The polynomials of the working space

`ĉ` (0), four temporaries (1–4), then `h` (`k`), `y` (`ℓ`), `ŷ` (`ℓ`),
`w` (`k`), `ŝ₁` (`ℓ`), `ŝ₂` (`k`), `t̂₀` (`k`) and `Â` (`kℓ`, row by row). -/
`w` (`k`), `ŝ₁` (`ℓ`), `ŝ₂` (`k`), `t̂₀` (`k`) and `Â` (`kℓ`, row by row), then the
working space of `vg_mldsa_rej_ntt_poly4` (8 KiB). -/

abbrev pS (i : Nat) : Ptr := sc (oP i)
abbrev cP : Ptr := pS 0
Expand All @@ -87,6 +91,10 @@ abbrev s1P (r : Nat) : Ptr := pS (5 + 2 * p.k + 2 * p.ℓ + r)
abbrev s2P (i : Nat) : Ptr := pS (5 + 2 * p.k + 3 * p.ℓ + i)
abbrev t0P (i : Nat) : Ptr := pS (5 + 3 * p.k + 3 * p.ℓ + i)
abbrev aP (i j : Nat) : Ptr := pS (5 + 4 * p.k + 3 * p.ℓ + p.ℓ * i + j)
/-- Entry `e = ℓi + j` of `Â`. -/
abbrev aE (e : Nat) : Ptr := pS (5 + 4 * p.k + 3 * p.ℓ + e)
/-- The working space of `vg_mldsa_rej_ntt_poly4`. -/
abbrev r4P : Ptr := pS (5 + 4 * p.k + 3 * p.ℓ + p.k * p.ℓ)

end

Expand All @@ -99,8 +107,23 @@ variable (P : Prims) (p : Params)
def sampleE (e : Nat) : Prog isa :=
.seq (.block (setB (sc (oRS + 32)) (e % p.ℓ) ++ setB (sc (oRS + 33)) (e / p.ℓ))) (rejAt P (aP p (e / p.ℓ) (e % p.ℓ)))

/-- `ρ` to `RS`, and the `kℓ` entries of `Â`. -/
def expandA : Prog isa := .seq (copy (sc oRS) (.rbp, 0) 32) (seqR (sampleE P p) 0 (p.k * p.ℓ))
/-- `ρ` to seed `k` of `RS4`. -/
def cpR4 (k : Nat) : Prog isa := copy (sc (oRS4 + 34 * k)) (.rbp, 0) 32

/-- The indices of entry `e + k` of `Â` to seed `k` of `RS4`. -/
def setSR (e k : Nat) : List Instr :=
setB (sc (oRS4 + 34 * k + 32)) ((e + k) % p.ℓ) ++ setB (sc (oRS4 + 34 * k + 33)) ((e + k) / p.ℓ)

/-- Entries `4g, …, 4g + 3` of `Â`. -/
def sample4 (g : Nat) : Prog isa :=
.seq (.block (setSR p (4 * g) 0)) (.seq (.block (setSR p (4 * g) 1)) (.seq (.block (setSR p (4 * g) 2))
(.seq (.block (setSR p (4 * g) 3)) (rej4At P (aE p (4 * g)) (r4P p)))))

/-- `ρ` to `RS` and to the four seeds of `RS4`, and the `kℓ` entries of `Â`: four at a time, then
the last `kℓ mod 4` one at a time. -/
def expandA : Prog isa :=
.seq (copy (sc oRS) (.rbp, 0) 32) (.seq (seqR cpR4 0 4)
(.seq (seqR (sample4 P p) 0 (p.k * p.ℓ / 4)) (seqR (sampleE P p) (4 * (p.k * p.ℓ / 4)) (p.k * p.ℓ % 4))))

/-- `ŝ₁[r]`. -/
def decS1 (r : Nat) : Prog isa :=
Expand Down
6 changes: 5 additions & 1 deletion lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,10 @@ structure Rej4Ok (c : Prog isa) : Prop where
depth : c.depth ≤ 3
ctl : ctlOk c = true
sp : c.all (fun i => !isa.writesSp i) = true
/-- Its result is whether each seed has 256 coefficients in its first 1008 bytes of output, as
both implementations', which signing branches on. -/
ret : ∀ s t s', (Spec.MlDsa.rejNTT4Contract X86_64.abi 24).pre s → Exec isa c s t s' →
(s'.gpr .rax).setWidth 32 = Rej4.rej4Res s.mem (s.gpr .rdi)

/-- Each function of the backend `B` meets its contract, and is safe to call. -/
structure BackendOk (B : Backend) : Prop where
Expand Down Expand Up @@ -96,7 +100,7 @@ def ArithImpl.sse2 : ArithImpl where
makeHint := FnOk.of Round.makeHint_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
(by decide +kernel)
rej4 := ⟨Rej4.rejNTT4_verified, Proof.MlKem.X86_64.nosp_of (by decide +kernel), by decide +kernel,
by decide +kernel, Code.all_of_allInstrs (by decide +kernel)⟩ }
by decide +kernel, Code.all_of_allInstrs (by decide +kernel), fun _ _ _ => Rej4.rejNTT4_ret⟩ }
features := []

end VG.Proof.MlDsa.X86_64
Original file line number Diff line number Diff line change
Expand Up @@ -41,7 +41,7 @@ def ArithImpl.avx2 : ArithImpl where
makeHint := FnOk.of Round.makeHintY_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
(by decide +kernel)
rej4 := ⟨Rej4.rejNTT4Avx2_verified, Proof.MlKem.X86_64.nosp_of (by decide +kernel), by decide +kernel,
by decide +kernel, Code.all_of_allInstrs (by decide +kernel)⟩ }
by decide +kernel, Code.all_of_allInstrs (by decide +kernel), fun _ _ _ => Rej4.rejNTT4Avx2_ret⟩ }
features := ["avx", "avx2"]

end VG.Proof.MlDsa.X86_64
65 changes: 63 additions & 2 deletions lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sample/Rej4Verified.lean
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import VerifiedGarbage.Proof.MlDsa.X86_64.Sample.Rej4Scalar
import VerifiedGarbage.Proof.MlKem.X86_64.S4Verified
import VerifiedGarbage.Proof.MlDsa.KeyGen.Mono

/-!
# ML-DSA on x86-64: `vg_mldsa_rej_ntt_poly4` and `vg_mldsa_rej_ntt_poly4_avx2`, verified
Expand All @@ -17,6 +18,7 @@ open VG VG.X86_64
open VG.Proof.MlDsa.Sample (rnFold rejNTT_some rejNTT_none)
open VG.Proof.MlDsa.X86_64.Sample (leakBytes_inj)
open VG.Spec.MlDsa (G)
open VG.Spec.Sha3 (bytesAt)

theorem seed4_eq : Spec.MlDsa.seed4 = Spec.MlKem.seed4 := rfl
theorem poly4_eq : Spec.MlDsa.poly4 = Spec.MlKem.poly4 := rfl
Expand Down Expand Up @@ -45,11 +47,62 @@ theorem r4_post {s s' : State} (h : r4K.post s s') :
obtain ⟨k, hk, hk'⟩ := hall
exact ⟨k, hk, rejNTT_none (B := 1008) (by decide) (by decide) hk'⟩

theorem r4_pre (s : State) (h : (Spec.MlDsa.rejNTT4Contract X86_64.abi 24).pre s) : r4K.pre s := by
revert s h
sig_implies_pre [Spec.MlDsa.rejNTT4Contract, Spec.MlDsa.rejNTT4Sig, r4K, MlKem.X86_64.sample4K, X86_64.abi,
X86_64.argRegs]

/-- The result of either implementation: whether each seed has 256 coefficients in its first 1008
bytes of output. -/
def rej4Res (m : Mem) (a : Addr) : BitVec 32 :=
if (List.range 4).all (fun k => (rnFold [] (G (Spec.MlDsa.seed4 m a k) 1008)).length == 256) then 1 else 0

theorem seed4_of136 {m m' : Mem} {a a' : Addr} (h : bytesAt m a 136 = bytesAt m' a' 136) {k : Nat} (hk : k < 4) :
Spec.MlDsa.seed4 m a k = Spec.MlDsa.seed4 m' a' k := by
unfold Spec.MlDsa.seed4
rw [← Proof.MlKem.bytesAt_slice m a (show 34 * k + 34 ≤ 136 by omega),
← Proof.MlKem.bytesAt_slice m' a' (show 34 * k + 34 ≤ 136 by omega), h]

/-- The result depends only on the 136 bytes of the seeds. -/
theorem rej4Res_congr {m m' : Mem} {a a' : Addr} (h : bytesAt m a 136 = bytesAt m' a' 136) :
rej4Res m a = rej4Res m' a' := by
have e : ((List.range 4).all fun k => (rnFold [] (G (Spec.MlDsa.seed4 m a k) 1008)).length == 256) =
((List.range 4).all fun k => (rnFold [] (G (Spec.MlDsa.seed4 m' a' k) 1008)).length == 256) := by
rw [Bool.eq_iff_iff, List.all_eq_true, List.all_eq_true]
exact ⟨fun H k hk => by rw [← seed4_of136 h (List.mem_range.mp hk)]; exact H k hk,
fun H k hk => by rw [seed4_of136 h (List.mem_range.mp hk)]; exact H k hk⟩
simp only [rej4Res, e]

/-- The public data of two calls include their seeds. -/
theorem r4_pub {S : Nat} (s₁ s₂ : State) (h : (Spec.MlDsa.rejNTT4Contract X86_64.abi S).pub s₁ s₂) :
bytesAt s₁.mem (s₁.gpr .rdi) 136 = bytesAt s₂.mem (s₂.gpr .rdi) 136 := by
sig_pub [Spec.MlDsa.rejNTT4Contract, Spec.MlDsa.rejNTT4Sig, X86_64.abi, X86_64.argRegs] at h
obtain ⟨_, hb, _⟩ := h
exact leakBytes_inj hb

/-- Within the bound both implementations sample to, so within `maxBounds`'s. -/
theorem rej4Res_max {m : Mem} {a : Addr} (h : rej4Res m a = 1) {k : Nat} (hk : k < 4) {B : Nat} (hB : 1008 ≤ B) :
(Spec.MlDsa.rejNTTPoly B (Spec.MlDsa.seed4 m a k)).isSome := by
unfold rej4Res at h
by_cases hall : ((List.range 4).all fun k => (rnFold [] (G (Spec.MlDsa.seed4 m a k) 1008)).length == 256) = true
· have hs : (rnFold [] (G (Spec.MlDsa.seed4 m a k) 1008)).length = 256 := by
simpa using List.all_eq_true.mp hall k (List.mem_range.mpr hk)
rw [Proof.MlDsa.KeyGen.rejNTTPoly_mono hB (rejNTT_some hs)]; rfl
· rw [ite_eq_right hall] at h; exact absurd h (by decide)

/-- The result of code that meets `r4K`. -/
theorem rej4_ret {c : Prog isa}
(hc : ∀ σ, r4K.pre σ → ∃ t s', Exec isa c σ t s' ∧ abiPreserved σ s' ∧ r4K.post σ s') {s s' : State}
{t : List Leak} (h : (Spec.MlDsa.rejNTT4Contract X86_64.abi 24).pre s) (e : Exec isa c s t s') :
(s'.gpr .rax).setWidth 32 = rej4Res s.mem (s.gpr .rdi) := by
obtain ⟨_, _, e', _, hq⟩ := hc s (r4_pre s h)
obtain ⟨-, rfl⟩ := Exec.det e e'
exact hq.1

theorem rej4_verified (c : Prog isa) (hc : ∀ σ, r4K.pre σ → ∃ t s', Exec isa c σ t s' ∧ abiPreserved σ s' ∧ r4K.post σ s')
(ht : ConstantTime isa r4K.pre r4K.pub c) : Verified X86_64.target c (Spec.MlDsa.rejNTT4Contract X86_64.abi 24) :=
Verified.of_correct hc ht
{ pre := by sig_implies_pre [Spec.MlDsa.rejNTT4Contract, Spec.MlDsa.rejNTT4Sig, r4K, MlKem.X86_64.sample4K,
X86_64.abi, X86_64.argRegs]
{ pre := r4_pre
post := by
intro s s' _ h
sig_post [Spec.MlDsa.rejNTT4Contract, Spec.MlDsa.rejNTT4Sig, r4K, X86_64.abi, X86_64.argRegs]
Expand All @@ -76,4 +129,12 @@ theorem rejNTT4Avx2_verified : Verified X86_64.target Impl.MlDsa.X86_64.Sample.R
theorem rejNTT4_verified : Verified X86_64.target Impl.MlDsa.X86_64.Sample.Rej4.rejNTT4
(Spec.MlDsa.rejNTT4Contract X86_64.abi 24) := rej4_verified _ correct_scalar ct_scalar

theorem rejNTT4Avx2_ret {s s' : State} {t : List Leak} (h : (Spec.MlDsa.rejNTT4Contract X86_64.abi 24).pre s)
(e : Exec isa Impl.MlDsa.X86_64.Sample.Rej4.rejNTT4Avx2 s t s') :
(s'.gpr .rax).setWidth 32 = rej4Res s.mem (s.gpr .rdi) := rej4_ret correct h e

theorem rejNTT4_ret {s s' : State} {t : List Leak} (h : (Spec.MlDsa.rejNTT4Contract X86_64.abi 24).pre s)
(e : Exec isa Impl.MlDsa.X86_64.Sample.Rej4.rejNTT4 s t s') :
(s'.gpr .rax).setWidth 32 = rej4Res s.mem (s.gpr .rdi) := rej4_ret correct_scalar h e

end VG.Proof.MlDsa.X86_64.Rej4
Loading
Loading