Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
18 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
4563920
Merge main into claude/fervent-einstein-ukl7t7-avx2
claude Oct 1, 2026
8160f69
Merge branch 'claude/fervent-einstein-ukl7t7-avx2' into claude/ferven…
claude Oct 1, 2026
6819894
Merge main into claude/fervent-einstein-ukl7t7-avx2
claude Oct 1, 2026
3f51d68
Merge branch 'claude/fervent-einstein-ukl7t7-avx2' into claude/ferven…
claude Oct 1, 2026
8cc896c
Merge main into claude/fervent-einstein-ukl7t7-rej4
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 @@ -827,7 +827,7 @@ yours to keep:

<td>✅</td>

<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic</td>
<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in verification (four SHAKE128 instances at once with AVX2)</td>

<td>✅ SHA extensions</td>

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

<td>✅</td>

<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic</td>
<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in verification (four SHAKE128 instances at once with AVX2)</td>

<td>✅ SHA extensions</td>

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

<td>✅</td>

<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic</td>
<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in verification (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" }
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in verification (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" }
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in verification (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" }
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in verification (four SHAKE128 instances at once with AVX2)" }
23 changes: 23 additions & 0 deletions lean/VerifiedGarbage/Artifacts/MlDsaSample/X86_64.lean
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Sample.RejNttCT
import VerifiedGarbage.Proof.MlDsa.X86_64.Sample.RejBoundedCT
import VerifiedGarbage.Proof.MlDsa.X86_64.Sample.ExpandMask
import VerifiedGarbage.Proof.MlDsa.X86_64.Sample.BallCT
import VerifiedGarbage.Proof.MlDsa.X86_64.Sample.Rej4Verified

/-!
# ML-DSA (FIPS 204) on x86-64: the sampling primitives
Expand Down Expand Up @@ -32,6 +33,28 @@ def artifacts : List Artifact := [
stack := 16
verified := Proof.MlDsa.X86_64.Sample.rejNTT_verified
spSafe := Code.all_of_allInstrs (by lit_decide) },
{ Spec.MlDsa.rejNTT4Api with
target := X86_64.target
doc := Spec.MlDsa.rejNTT4Api.doc (notes := ["It calls `vg_mldsa_rej_ntt_poly` on each seed, and saves \
its caller's callee-saved registers in `scratch`."])
code := Impl.MlDsa.X86_64.Sample.Rej4.rejNTT4
contract := Spec.MlDsa.rejNTT4Contract X86_64.abi 24
stack := 24
verified := Proof.MlDsa.X86_64.Rej4.rejNTT4_verified
spSafe := Code.all_of_allInstrs (by decide +kernel) },
{ Spec.MlDsa.rejNTT4Api with
name := Spec.MlDsa.rejNTT4Api.name ++ "_avx2"
target := X86_64.target
doc := Spec.MlDsa.rejNTT4Api.doc (notes := ["It runs the four instances of SHAKE128 at once, in the four \
64-bit elements of AVX2 registers (as `vg_mlkem_sample_ntt4_avx2` does), squeezing three blocks of each \
twice, and runs the loop of `RejNTTPoly` over the 1008 bytes of each seed's output, as \
`vg_mldsa_rej_ntt_poly` does. It saves its caller's callee-saved registers in `scratch`."])
code := Impl.MlDsa.X86_64.Sample.Rej4.rejNTT4Avx2
contract := Spec.MlDsa.rejNTT4Contract X86_64.abi 24
stack := 24
verified := Proof.MlDsa.X86_64.Rej4.rejNTT4Avx2_verified
spSafe := Code.all_of_allInstrs (by decide +kernel)
features := ["avx", "avx2"] },
{ Spec.MlDsa.rejBoundedApi with
target := X86_64.target
doc := Spec.MlDsa.rejBoundedApi.doc (notes := ["It squeezes 544 bytes of SHAKE256 output (4 blocks) and \
Expand Down
16 changes: 8 additions & 8 deletions lean/VerifiedGarbage/Generic/MlDsaArith/X86_64/MlDsaVerify.lean
Original file line number Diff line number Diff line change
Expand Up @@ -16,8 +16,8 @@ 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.
The stack is 32 bytes: the return address of a call of a primitive, and up
to 24 bytes for its own calls (`vg_mldsa_rej_ntt_poly4`'s).
-/

namespace VG.Generic.MlDsaArith.X86_64.MlDsaVerify
Expand All @@ -38,8 +38,8 @@ def artifacts (v : ArithImpl) : List Artifact := [
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
contract := Spec.MlDsa.verifyContract Spec.MlDsa.mlDsa44 X86_64.abi 32
stack := 32
verified := verify_prims v (List.mem_cons_self ..)
spSafe := verify_spSafe (prims_okWith v) (List.mem_cons_self ..) },
{ Spec.MlDsa.verify65Api with
Expand All @@ -48,8 +48,8 @@ def artifacts (v : ArithImpl) : List Artifact := [
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
contract := Spec.MlDsa.verifyContract Spec.MlDsa.mlDsa65 X86_64.abi 32
stack := 32
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
Expand All @@ -58,8 +58,8 @@ def artifacts (v : ArithImpl) : List Artifact := [
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
contract := Spec.MlDsa.verifyContract Spec.MlDsa.mlDsa87 X86_64.abi 32
stack := 32
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 ..))) }]

Expand Down
3 changes: 2 additions & 1 deletion lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean
Original file line number Diff line number Diff line change
Expand Up @@ -198,6 +198,7 @@ def subAvx2 : Prog isa :=
.seq (.block (yconst .xmm15 8380417)) (.seq (rcxLoop 32 (yaccBody .psubd (vcadd .xmm0 .xmm2))) (.block yepi))

/-- The AVX2 code. -/
def Backend.avx2 : Backend := ⟨nttAvx2, nttInvAvx2, mulAvx2, mulAddAvx2, addAvx2, subAvx2, "_avx2"⟩
def Backend.avx2 : Backend :=
⟨nttAvx2, nttInvAvx2, mulAvx2, mulAddAvx2, addAvx2, subAvx2, Sample.Rej4.rejNTT4Avx2, "_avx2"⟩

end VG.Impl.MlDsa.X86_64.Arith
10 changes: 7 additions & 3 deletions lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean
Original file line number Diff line number Diff line change
@@ -1,14 +1,16 @@
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.Sample.RejNtt4

/-!
# 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.
`vg_mldsa_add` and `vg_mldsa_sub`, and of `vg_mldsa_rej_ntt_poly4`
(which samples four entries of the matrix `Â` at once), 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
Expand All @@ -27,15 +29,17 @@ structure Backend where
mulAdd : Prog isa
add : Prog isa
sub : Prog isa
rej4 : 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, ""⟩
def Backend.sse2 : Backend :=
⟨Arith.ntt, Arith.nttInv, Arith.mul, Arith.mulAdd, Arith.add, Arith.sub, Sample.Rej4.rejNTT4, ""⟩

/-- 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 [], ""⟩
def Backend.empty : Backend := ⟨.block [], .block [], .block [], .block [], .block [], .block [], .block [], ""⟩

end VG.Impl.MlDsa.X86_64.Arith
85 changes: 85 additions & 0 deletions lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sample/RejNtt4.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,85 @@
import VerifiedGarbage.Impl.MlDsa.X86_64.Sample.RejNtt
import VerifiedGarbage.Impl.MlKem.X86_64.Sample4

/-!
# ML-DSA on x86-64: `vg_mldsa_rej_ntt_poly4` and `vg_mldsa_rej_ntt_poly4_avx2`

`rejNTT4(seeds = rdi, a = rsi, scratch = rdx) -> eax` runs `RejNTTPoly` on
four seeds. The baseline implementation (`rejNTT4`) calls
`vg_mldsa_rej_ntt_poly` on each, with the prologue and epilogue of
`vg_mlkem_sample_ntt4` (`Impl/MlKem/X86_64/Sample4.lean`: `scratch` in
`rbx`, `seeds` in `r12`, `a` in `r13`, the AND of the results in `r14`, and
their caller's values and `rbp`'s saved in `scratch[4384..4424)`) and its
scratch space from byte 6144 of `scratch`.

The one for AVX2 (`rejNTT4Avx2`) runs the four SHAKE128 instances at once
in the four 64-bit elements of `ymm` registers, as `vg_mlkem_sample_ntt4_avx2`
does, with the same code to absorb the seeds and squeeze three blocks of
each (504 bytes, from byte `2368 + 504 k` of `scratch`). It samples from them
with the loop of `vg_mldsa_rej_ntt_poly` (`rnBody`, 168 iterations of 3
bytes), keeping the number `j` of coefficients of seed `k` at
`scratch[4424 + 8 k]`; squeezes three more blocks of each to the same
place, and runs 168 more iterations of the loop on them. That is the loop of
`vg_mldsa_rej_ntt_poly` on the same 1008 bytes of output, so seed `k` has
256 coefficients if and only if `vg_mldsa_rej_ntt_poly` would have them
(and if not, neither has `RejNTTPoly` within the least bound of FIPS 204
Appendix C, 894 bytes). It returns 1 if every seed has 256 coefficients
(the AND of `j >> 8`), and 0 otherwise.

The loops' branches and the addresses of their stores depend on the XOF
output, a function of the seeds, and on nothing else; every other address
and branch depends only on the pointers.
-/

namespace VG.Impl.MlDsa.X86_64.Sample.Rej4

open VG.X86_64
open VG.Impl.MlKem.X86_64 (at_)
open VG.Impl.MlKem.X86_64.Sample4 (absorb4 squeeze4 oRc oBuf oScalar)
open VG.Impl.Sha3.X86_64.X4 (rcTable)

/-- Where `j` of seed `k` is kept, at `scratch + oJ + 8 k`. -/
def oJ : Nat := 4424

/-- `j ← 0` for each seed. -/
def zeroJ : List Instr :=
.mov32 .rax (.imm 0) :: (List.range 4).map fun k => .store (at_ .rbx (oJ + 8 * k)) .rax

/-- 168 iterations of `RejNTTPoly`'s loop on the 504 bytes of the output of
seed `k`, to polynomial `k`, from `j` at `scratch + oJ + 8 k`. -/
def half (k : Nat) : Prog isa :=
.seq (.block [.mov .rsi (.reg .rbx), .alu .add .rsi (.imm (BitVec.ofNat 32 (oBuf + 504 * k))),
.mov .rbp (.reg .r13), .alu .add .rbp (.imm (BitVec.ofNat 32 (1024 * k))),
.mov .rdi (.mem (at_ .rbx (oJ + 8 * k))), .mov32 .rcx (.imm 168)])
(.loop rnBody .ne)

/-- The first half of seed `k`, and `j` kept. -/
def first (k : Nat) : Prog isa := .seq (half k) (.block [.store (at_ .rbx (oJ + 8 * k)) .rdi])

/-- The second half of seed `k`, and `r14 ← r14 ∧ (j >> 8)`. -/
def second (k : Nat) : Prog isa :=
.seq (half k) (.block [.mov .rax (.reg .rdi), .shift .shr .rax 8, .alu32 .and .r14 (.reg .rax)])

/-- `vg_mldsa_rej_ntt_poly4_avx2`. `vzeroupper` clears the upper halves of
the vector registers after the last permutation. -/
def rejNTT4Avx2 : Prog isa :=
.seq (.block (VG.Impl.MlKem.X86_64.Sample4.pro ++ rcTable .rbx (oRc / 32) ++ absorb4))
(.seq (squeeze4 0) (.seq (squeeze4 1) (.seq (squeeze4 2) (.seq (.block zeroJ)
(.seq (first 0) (.seq (first 1) (.seq (first 2) (.seq (first 3)
(.seq (squeeze4 0) (.seq (squeeze4 1) (.seq (squeeze4 2) (.seq (.block [.vop .vzeroupper])
(.seq (second 0) (.seq (second 1) (.seq (second 2) (.seq (second 3) (.block VG.Impl.MlKem.X86_64.Sample4.epi)))))))))))))))))

/-- `vg_mldsa_rej_ntt_poly` on seed `k`, and `r14 ← r14 ∧ result`. -/
def callK (k : Nat) : Prog isa :=
.seq (.block [.mov .rdi (.reg .r12), .alu .add .rdi (.imm (BitVec.ofNat 32 (34 * k))),
.mov .rsi (.reg .r13), .alu .add .rsi (.imm (BitVec.ofNat 32 (1024 * k))), .mov .rdx (.reg .rbx),
.alu .add .rdx (.imm (BitVec.ofNat 32 oScalar))])
(.seq (.call "vg_mldsa_rej_ntt_poly" rejNTT) (.block [.alu32 .and .r14 (.reg .rax)]))

/-- `vg_mldsa_rej_ntt_poly4`: `vg_mldsa_rej_ntt_poly` on each seed, with the
prologue and epilogue of `vg_mldsa_rej_ntt_poly4_avx2`. -/
def rejNTT4 : Prog isa :=
.seq (.block VG.Impl.MlKem.X86_64.Sample4.pro) (.seq (callK 0) (.seq (callK 1) (.seq (callK 2) (.seq (callK 3)
(.block VG.Impl.MlKem.X86_64.Sample4.epi)))))

end VG.Impl.MlDsa.X86_64.Sample.Rej4
26 changes: 19 additions & 7 deletions lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean
Original file line number Diff line number Diff line change
Expand Up @@ -21,11 +21,13 @@ The layout of `scratch` (in bytes): the Keccak state at 0 (200 bytes) and
the sponge functions' working space at 200 (640 bytes); the caller's
callee-saved registers at 840 (48 bytes); the seed of `RejNTTPoly` (`SB`,
34 bytes) at 896; `w1Encode(w′₁)` at 1024 (at most 1024 bytes); the
recomputed commitment hash `c̃′` at 2048 (at most 64 bytes); the working
recomputed commitment hash `c̃′` at 2048 (at most 64 bytes); the four seeds
of `vg_mldsa_rej_ntt_poly4` (`SB4`, 136 bytes) at 2560; the working
space of the primitives at 4096 (2048 bytes); and polynomials of 1024 bytes
from 8192 (`P j`): the hint `h` (polynomials 0 to 7, of which the first
`k`), `z` (8 to 14), `c` (15), two temporaries (16, 17), `w′` (18), `w′₁`
(19), and `Â[r, s]` (`20 + 8r + s`).
(19), and `Â[r, s]` (`20 + 8r + s`), then the working space of
`vg_mldsa_rej_ntt_poly4` (8 KiB, after the last row of `Â`).
-/

namespace VG.Impl.MlDsa.X86_64.Verify
Expand All @@ -42,6 +44,7 @@ def oSV : Nat := 840
def oSB : Nat := 896
def oB : Nat := 1024
def oCT : Nat := 2048
def oSB4 : Nat := 2560
def oSS : Nat := 4096
/-- Polynomial `j`. -/
def oP (j : Nat) : Nat := 8192 + 1024 * j
Expand Down Expand Up @@ -137,6 +140,7 @@ structure Prims where
unpackT1 : Prog isa
hintUnpack : Prog isa
normLt : Prog isa
rej4 : Prog isa
/-- What the names of the polynomial arithmetic's functions end with (`Arith.Backend`). -/
sfx : String := ""

Expand All @@ -159,6 +163,9 @@ def subAt (f g : Ptr) : Prog isa := callAt ("vg_mldsa_sub" ++ P.sfx) P.sub [(.rd
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))]

def rej4At (a w : Ptr) : Prog isa :=
callAt ("vg_mldsa_rej_ntt_poly4" ++ P.sfx) P.rej4 [(.rdi, .ptr (sc oSB4)), (.rsi, .ptr a), (.rdx, .ptr w)]

def ballAt (ct : Ptr) (len tau : Nat) (c : Ptr) : Prog isa :=
callAt "vg_mldsa_sample_in_ball" P.ball
[(.rdi, .ptr ct), (.rsi, .imm len), (.rdx, .imm tau), (.rcx, .ptr c), (.r8, .ptr (sc oSS))]
Expand Down Expand Up @@ -191,19 +198,24 @@ end
/-- `r15 ← r15 ∧ eax`. -/
def and15 : List Instr := [.alu32 .and .r15 (.reg .rax)]

/-- The polynomial at `a` masked by the result `eax` (0 or 1) of the
sampler that wrote it: unchanged if 1, and zero if 0, so that it is reduced
either way, without a branch. `edx ← -eax`, then each coefficient `∧ edx`. -/
def mask (a : Ptr) : Prog isa :=
/-- The polynomial at `a` (or the `N` coefficients from `a`) masked by the
result `eax` (0 or 1) of the sampler that wrote it: unchanged if 1, and zero
if 0, so that it is reduced either way, without a branch. `edx ← -eax`, then
each coefficient `∧ edx`. -/
def mask (a : Ptr) (N : Nat := 256) : Prog isa :=
.seq (.block ([.mov32 .rdx (.imm 0), .alu32 .sub .rdx (.reg .rax)] ++
glue [(.rdi, .ptr a), (.rcx, .imm 256)]))
glue [(.rdi, .ptr a), (.rcx, .imm N)]))
(.loop (.block [.mov32 .rax (.mem (at_ .rdi 0)), .alu32 .and .rax (.reg .rdx), .store32 (at_ .rdi 0) .rax,
.alu .add .rdi (.imm 4), .alu .sub .rcx (.imm 1)]) .ne)

/-- A sampler's call, its result ANDed into `r15`, and its output masked. -/
def sampled (call : Prog isa) (a : Ptr) : Prog isa :=
.seq call (.seq (.block and15) (mask a))

/-- The same for a call that samples the four polynomials from `a`. -/
def sampled4 (call : Prog isa) (a : Ptr) : Prog isa :=
.seq call (.seq (.block and15) (mask a 1024))

/-- `c` if `r15 ≠ 0`. -/
def ifOk (c : Prog isa) : Prog isa :=
.seq (.block [.alu32 .test .r15 (.reg .r15)]) (.ite .ne c (.block []))
Expand Down
32 changes: 26 additions & 6 deletions lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Verify.lean
Original file line number Diff line number Diff line change
Expand Up @@ -16,8 +16,11 @@ in `scratch`.
returns 0 at once if the hint is malformed.
2. `z[i] = BitUnpack` of the `i`-th piece of `σ` (polynomial `8 + i`), and
`r15 ← r15 ∧ (‖z[i]‖∞ < γ₁ - β)`; it returns 0 if one of them is not.
3. `ρ` (`pk[0 : 32]`) to `SB`, and `Â[r, s] = RejNTTPoly(ρ ‖ s ‖ r)`
(polynomial `20 + 8r + s`); `c = SampleInBall(c̃)` (polynomial 15). Each
3. `ρ` (`pk[0 : 32]`) to `SB` and to each seed of `SB4`, and
`Â[r, s] = RejNTTPoly(ρ ‖ s ‖ r)` (polynomial `20 + 8r + s`), four
entries of a row at a time (`vg_mldsa_rej_ntt_poly4`): those from 0, then
for `ℓ = 7` those from 3 (sampling entry 3 again), or for `ℓ = 5` entry 4
alone; `c = SampleInBall(c̃)` (polynomial 15). Each
sampler's result is ANDed into `r15`, and its output masked with it
(`sampled`): a sampler that fails leaves its output unspecified, and
masking makes it reduced (zero) without a branch on the result, which
Expand Down Expand Up @@ -73,12 +76,29 @@ def aOne (e : Nat) : Prog isa :=
.seq (.block (setB (sc (oSB + 32)) (e % 8) ++ setB (sc (oSB + 33)) (e / 8)))
(sampled (rejNttAt P (pS (20 + e))) (pS (20 + e)))

/-- The entries `8r + s` of row `r` of `Â`. -/
def aRow (r : Nat) : Prog isa := seqR (aOne P) (8 * r) p.ℓ
/-- Where `vg_mldsa_rej_ntt_poly4` works: after the last row of `Â`. -/
def oR4 : Nat := oP (20 + 8 * p.k)

/-- `ρ` to `SB`, `Â`, and `c`. -/
/-- The bytes `s ‖ r` of seed `k` of `SB4`. -/
def setSR (r s k : Nat) : List Instr := setB (sc (oSB4 + 34 * k + 32)) (s + k) ++ setB (sc (oSB4 + 34 * k + 33)) r

/-- `Â[r, s], …, Â[r, s + 3]`. -/
def aGrp (r s : Nat) : Prog isa :=
.seq (.block (setSR r s 0)) (.seq (.block (setSR r s 1)) (.seq (.block (setSR r s 2)) (.seq (.block (setSR r s 3))
(sampled4 (rej4At P (pA r s) (sc (oR4 p))) (pA r s)))))

/-- The entries of row `r` of `Â`: four from 0, then the last four if `ℓ = 7`, or the last one if `ℓ = 5`. -/
def aRow (r : Nat) : Prog isa :=
.seq (aGrp P p r 0) (if p.ℓ = 7 then aGrp P p r 3 else seqR (aOne P) (8 * r + 4) (p.ℓ - 4))

/-- `ρ` to `SB` and to the four seeds of `SB4`. -/
def rhos : Prog isa :=
.seq (copy (sc oSB) (.rbp, 0) 32) (.seq (copy (sc oSB4) (.rbp, 0) 32) (.seq (copy (sc (oSB4 + 34)) (.rbp, 0) 32)
(.seq (copy (sc (oSB4 + 68)) (.rbp, 0) 32) (copy (sc (oSB4 + 102)) (.rbp, 0) 32))))

/-- `ρ` to `SB` and `SB4`, `Â`, and `c`. -/
def samples : Prog isa :=
.seq (copy (sc oSB) (.rbp, 0) 32) (.seq (seqR (aRow P p) 0 p.k)
.seq rhos (.seq (seqR (aRow P p) 0 p.k)
(sampled (ballAt P (.r13, 0) p.ctildeLen p.τ pC) pC))

/-- `Σₛ Â[r, s] ẑ[s]` to `W`. -/
Expand Down
Loading
Loading