Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
23 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
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
2c38b92
Merge branch 'claude/fervent-einstein-ukl7t7-rej4' 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
3d168fc
Merge branch 'claude/fervent-einstein-ukl7t7-rej4' into claude/ferven…
claude Oct 1, 2026
8cc896c
Merge main into claude/fervent-einstein-ukl7t7-rej4
claude Oct 2, 2026
850ee16
Merge branch 'claude/fervent-einstein-ukl7t7-rej4' into claude/ferven…
claude Oct 2, 2026
3c23c2a
Merge main into claude/fervent-einstein-ukl7t7-rej4kg
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 @@ -891,7 +891,7 @@ yours to keep:

<td>✅</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>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)</td>

<td>✅ SHA extensions</td>

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

<td>✅</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>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)</td>

<td>✅ SHA extensions</td>

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

<td>✅</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>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and 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; matrix sampled four entries at a time in verification (four SHAKE128 instances at once with AVX2)" }
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and 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; matrix sampled four entries at a time in verification (four SHAKE128 instances at once with AVX2)" }
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and 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; matrix sampled four entries at a time in verification (four SHAKE128 instances at once with AVX2)" }
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" }
2 changes: 2 additions & 0 deletions lean/VerifiedGarbage/Impl/MlDsa/X86_64/KeyGen/Inst.lean
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@ def prims : Prims where
power2Round := Round.power2Round
simpleBitPack := Pack.simpleBitPack
bitPack := Pack.bitPack
rej4 := Sample.Rej4.rejNTT4

/-- The primitives, with the polynomial arithmetic of `B`. -/
def primsWith (B : Arith.Backend) : Prims :=
Expand All @@ -38,6 +39,7 @@ def primsWith (B : Arith.Backend) : Prims :=
mul := B.mul
mulAdd := B.mulAdd
add := B.add
rej4 := B.rej4
sfx := B.sfx }

end VG.Impl.MlDsa.X86_64.KeyGen
50 changes: 38 additions & 12 deletions lean/VerifiedGarbage/Impl/MlDsa/X86_64/KeyGen/KeyGen.lean
Original file line number Diff line number Diff line change
Expand Up @@ -17,18 +17,22 @@ The layout of `scratch` (in bytes): the Keccak state at 0 and the sponge
functions' working space at 200, the saved registers at 840 (as ML-KEM's);
`k` and `ℓ` at 896 (`oKL`); `(ρ, ρ′, K)` at 1024 (`oHX`, 128 bytes); the
seed of `RejNTTPoly` at 1152 (`oSA`, 34 bytes) and of `RejBoundedPoly` at
1216 (`oSB`, 66 bytes); the working space of the primitives at 2048
1216 (`oSB`, 66 bytes); the four seeds of `vg_mldsa_rej_ntt_poly4` at 1408
(`oSA4`, 136 bytes); the working space of the primitives at 2048
(2048 bytes); and polynomials of 1024 bytes from 4096 (`oP j`): `Â[r, s]`
is polynomial `rℓ + s`, `s₁[j]` (then `ŝ₁[j]`) polynomial `kℓ + j`,
`s₂[i]` polynomial `kℓ + ℓ + i`, and `t`, `t₁` and `t₀` the three after
them.
them; then the working space of `vg_mldsa_rej_ntt_poly4` (8 KiB, `oR4`).

1. `(ρ, ρ′, K) = H(ξ ‖ k ‖ ℓ, 128)`, and `ρ` and `ρ′` to the seeds.
2. `Â[r, s] = RejNTTPoly(ρ ‖ s ‖ r)` and `s₁ ‖ s₂ = RejBoundedPoly(ρ′ ‖ r ‖ 0)`
(`ExpandA`, `ExpandS`). After each, `r15 ← r15 ∧ result`, and the
polynomial is ANDed with `-result` (`mask`): it is zero if the sampler
failed, so that every polynomial is reduced, and small, whatever the
samplers return, without a branch.
2. `Â[r, s] = RejNTTPoly(ρ ‖ s ‖ r)`, four consecutive entries at a time
(`vg_mldsa_rej_ntt_poly4`, from the seeds at `oSA4`, each `ρ` with its
entry's indices) and the last `kℓ mod 4` one at a time, and
`s₁ ‖ s₂ = RejBoundedPoly(ρ′ ‖ r ‖ 0)` (`ExpandA`, `ExpandS`). After each
call, `r15 ← r15 ∧ result`, and the polynomials it sampled are ANDed with
`-result` (`mask`): they are zero if the sampler failed, so that every
polynomial is reduced, and small, whatever the samplers return, without a
branch.
3. `ρ` and `K` to `sk`; `s₁` and `s₂`, `BitPack`ed, to `sk`; `ŝ₁ = NTT(s₁)`.
4. For each row `i`: `t = NTT⁻¹(Σⱼ Â[i, j] ŝ₁[j]) + s₂[i]`, `Power2Round`, and
`t₁` `SimpleBitPack`ed to `pk`, `t₀` `BitPack`ed to `sk`; `ρ` to `pk`.
Expand All @@ -50,6 +54,7 @@ def oKL : Nat := 896
def oHX : Nat := 1024
def oSA : Nat := 1152
def oSB : Nat := 1216
def oSA4 : Nat := 1408
/-- Polynomial `j`. -/
def oP (j : Nat) : Nat := 4096 + 1024 * j

Expand All @@ -60,6 +65,8 @@ abbrev sP (p : Params) (r : Nat) : Ptr := sc (oP (p.k * p.ℓ + r))
abbrev tP (p : Params) : Ptr := sc (oP (p.k * p.ℓ + p.ℓ + p.k))
abbrev t1P (p : Params) : Ptr := sc (oP (p.k * p.ℓ + p.ℓ + p.k + 1))
abbrev t0P (p : Params) : Ptr := sc (oP (p.k * p.ℓ + p.ℓ + p.k + 2))
/-- The working space of `vg_mldsa_rej_ntt_poly4`. -/
def oR4 (p : Params) : Nat := oP (p.k * p.ℓ + p.ℓ + p.k + 3)

/-- The length of a packed polynomial of `s₁` or `s₂`, `32 · bitlen (2η)`. -/
def lenS (p : Params) : Nat := 32 * bitlen (2 * p.η)
Expand Down Expand Up @@ -89,6 +96,9 @@ def addAt (sfx : String) (c : Prog isa) (f g : Ptr) : Prog isa :=
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)

def rej4At (c : Prog isa) (sfx : String) (a w : Ptr) : Prog isa :=
.seq (.block (lea .rdi (sc oSA4) ++ lea .rsi a ++ lea .rdx w)) (.call ("vg_mldsa_rej_ntt_poly4" ++ sfx) c)

def rejBoundedAt (c : Prog isa) (seed : Ptr) (eta : Nat) (a : Ptr) : Prog isa :=
.seq (.block (lea .rdi seed ++ imm .rsi eta ++ lea .rdx a ++ lea .rcx (sc oSS)))
(.call "vg_mldsa_rej_bounded_poly" c)
Expand All @@ -104,10 +114,11 @@ def bitPackAt (c : Prog isa) (f : Ptr) (a b : Nat) (out : Ptr) (len : Nat) : Pro
.seq (.block (lea .rdi f ++ imm .rsi a ++ imm .rdx b ++ lea .rcx out ++ imm .r8 len))
(.call "vg_mldsa_bit_pack" c)

/-- `r15 ← r15 ∧ eax`, and the polynomial at `a` ANDed with `-eax` (`eax` is 0 or 1). -/
def mask (a : Ptr) : Prog isa :=
/-- `r15 ← r15 ∧ eax`, and the polynomial at `a` (or the `N` coefficients from `a`) ANDed with `-eax`
(`eax` is 0 or 1). -/
def mask (a : Ptr) (N : Nat := 256) : Prog isa :=
.seq (.block ([.alu32 .and .r15 (.reg .rax), .mov32 .r8 (.imm 0), .alu32 .sub .r8 (.reg .rax)] ++
lea .rdi a ++ imm .rcx 256))
lea .rdi a ++ imm .rcx N))
(.loop (.block [.mov32 .rax (.mem (at_ .rdi 0)), .alu32 .and .rax (.reg .r8), .store32 (at_ .rdi 0) .rax,
.alu .add .rdi (.imm 4), .alu .sub .rcx (.imm 1)]) .ne)

Expand All @@ -120,13 +131,28 @@ def pro : List Instr := topPro .rcx [(.rbp, .rdi), (.r12, .rsi), (.r13, .rdx)]
def seeds (p : Params) : Prog isa :=
.seq (.block (setB (sc oKL) p.k ++ setB (sc (oKL + 1)) p.ℓ))
(.seq (hashAt [((.rbp, 0), 32), (sc oKL, 2)] 136 0x1f (sc oHX) 128)
(.seq (copy (sc oSA) (sc oHX) 32) (.seq (copy (sc oSB) (sc (oHX + 32)) 64) (.block (setB (sc (oSB + 65)) 0)))))
(.seq (copy (sc oSA) (sc oHX) 32) (.seq (copy (sc oSB) (sc (oHX + 32)) 64) (.seq (.block (setB (sc (oSB + 65)) 0))
(.seq (copy (sc oSA4) (sc oHX) 32) (.seq (copy (sc (oSA4 + 34)) (sc oHX) 32)
(.seq (copy (sc (oSA4 + 68)) (sc oHX) 32) (copy (sc (oSA4 + 102)) (sc oHX) 32))))))))

/-- `Â[e / ℓ, e % ℓ] = RejNTTPoly(ρ ‖ e % ℓ ‖ e / ℓ)`. -/
def expA (P : Prims) (p : Params) (e : Nat) : Prog isa :=
.seq (.block (setB (sc (oSA + 32)) (e % p.ℓ) ++ setB (sc (oSA + 33)) (e / p.ℓ)))
(.seq (rejNttAt P.rejNtt (sc oSA) (aP e)) (mask (aP e)))

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

/-- Entries `4g, …, 4g + 3` of `Â`. -/
def expA4 (P : Prims) (p : Params) (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)) (.seq (rej4At P.rej4 P.sfx (aP (4 * g)) (sc (oR4 p))) (mask (aP (4 * g)) 1024)))))

/-- The entries of `Â`: four at a time, then the last `kℓ mod 4` one at a time. -/
def expAll (P : Prims) (p : Params) : Prog isa :=
.seq (seqR (expA4 P p) 0 (p.k * p.ℓ / 4)) (seqR (expA P p) (4 * (p.k * p.ℓ / 4)) (p.k * p.ℓ % 4))

/-- Entry `r` of `s₁ ‖ s₂`: `RejBoundedPoly(ρ′ ‖ r ‖ 0)`. -/
def expS (P : Prims) (p : Params) (r : Nat) : Prog isa :=
.seq (.block (setB (sc (oSB + 64)) r))
Expand Down Expand Up @@ -162,7 +188,7 @@ def rest (P : Prims) (p : Params) : Prog isa :=

/-- `vg_mldsa*_keygen` for the parameter set `p`, calling the primitives `P`. -/
def keyGen (P : Prims) (p : Params) : Prog isa :=
.seq (.block pro) (.seq (seeds p) (.seq (seqR (expA P p) 0 (p.k * p.ℓ))
.seq (.block pro) (.seq (seeds p) (.seq (expAll P p)
(.seq (seqR (expS P p) 0 (p.ℓ + p.k)) (.seq (rest P p) (.block topEpi)))))

end VG.Impl.MlDsa.X86_64.KeyGen
2 changes: 2 additions & 0 deletions lean/VerifiedGarbage/Impl/MlDsa/X86_64/KeyGen/Prims.lean
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,8 @@ structure Prims where
simpleBitPack : Prog isa
/-- `vg_mldsa_bit_pack` -/
bitPack : 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 Down
10 changes: 9 additions & 1 deletion lean/VerifiedGarbage/Proof/MlDsa/X86_64/KeyGen/Base.lean
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,13 @@ structure Callee (c : Prog isa) (k : Nat → Contract isa) : Prop where
nosp : NoSp c
depth : c.depth ≤ 2

/-- `vg_mldsa_rej_ntt_poly4`: verified with 24 bytes of stack, never writing `rsp` but by calls nested at
most three deep. -/
structure Callee4 (c : Prog isa) : Prop where
verified : Verified X86_64.target c (Spec.MlDsa.rejNTT4Contract X86_64.abi 24)
nosp : NoSp c
depth : c.depth ≤ 3

/-- Verified implementations of the primitives key generation calls. -/
structure PrimsOk (P : Prims) : Prop where
ntt : Callee P.ntt (fun stk => Spec.MlDsa.nttContract X86_64.abi stk)
Expand All @@ -44,6 +51,7 @@ structure PrimsOk (P : Prims) : Prop where
power2Round : Callee P.power2Round (fun stk => Spec.MlDsa.power2RoundContract X86_64.abi stk)
simpleBitPack : Callee P.simpleBitPack (fun stk => Spec.MlDsa.simpleBitPackContract X86_64.abi stk)
bitPack : Callee P.bitPack (fun stk => Spec.MlDsa.bitPackContract X86_64.abi stk)
rej4 : Callee4 P.rej4

/-! ## MXCSR -/

Expand Down Expand Up @@ -118,7 +126,7 @@ theorem WP.callMx {n : String} {c : Prog isa} {k : Contract isa}
/-- `glueCall_ok`, keeping MXCSR's control bits. -/
theorem glueCallMx_ok {glue : List Instr} {n : String} {c : Prog isa} {k : Contract isa}
(hv : ∀ s, k.pre s → ∃ t s', Exec isa c s t s' ∧ abiPreserved s s' ∧ k.post s s')
(hsp : NoSp c) (hd : c.depth ≤ 2) (hgl : ∀ i ∈ glue, loadsMxcsr i = false) {s : State}
(hsp : NoSp c) (hd : c.depth ≤ 3) (hgl : ∀ i ∈ glue, loadsMxcsr i = false) {s : State}
{V : State → Prop}
(hg : WP isa (.block glue) s fun s1 => (V s1 ∧ s1.mem = s.mem) ∧ Keep MlKem.X86_64.argRegs s s1)
{rd wr : List Region} (hpre : ∀ s1, V s1 → s1.mem = s.mem → Keep MlKem.X86_64.argRegs s s1 →
Expand Down
99 changes: 97 additions & 2 deletions lean/VerifiedGarbage/Proof/MlDsa/X86_64/KeyGen/Call.lean
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,19 @@ theorem ceWf (h24 : 24 ≤ (s.gpr .rsp).toNat) {n : Nat} (hn : n + 1 ≤ 16) : n
have := (s.gpr .rsp).isLt
omega

omit L in
theorem ceWf24 (h32 : 32 ≤ (s.gpr .rsp).toNat) : 24 ≤ (s1.gpr .rsp - 8).toNat := by
rw [hsp, BitVec.toNat_sub]
have : (8 : BitVec 64).toNat = 8 := rfl
rw [this]
have := (s.gpr .rsp).isLt
omega

theorem ceD24 {q : Ptr} {l : Nat} (h : inB (kgB p) q l = true) :
Region.Disjoint (below (s1.gpr .rsp - 8) 24) ⟨pa s q, l⟩ := by
have := stk_disj24 s1 (R := ⟨pa s q, l⟩) (by rw [hsp]; exact L.stkD h)
simpa only [State.callEntry_rsp] using this

end

/-! ## The memory of a call's entry -/
Expand Down Expand Up @@ -93,7 +106,7 @@ theorem primOk {c : Prog isa} {kk : Nat → Contract isa} (hc : Callee c kk) {gl
∃ s₂ : State, s₂.mem = s'.mem ∧ (∀ r, r ≠ .rsp → s₂.gpr r = s'.gpr r) ∧
(kk stk).post (s1.callEntry.withRegions rd wr) s₂ := by
obtain ⟨stk, hs, hver⟩ := hc.verified
exact WP.mono (glueCallMx_ok hver.1 hc.nosp hc.depth hgl hg (hpre stk hs) hcov hw)
exact WP.mono (glueCallMx_ok hver.1 hc.nosp (Nat.le_succ_of_le hc.depth) hgl hg (hpre stk hs) hcov hw)
fun s' ⟨h1, h2, h3⟩ => ⟨h1, h2, stk, hs, h3⟩

/-- A block of moves leaks nothing. -/
Expand Down Expand Up @@ -188,7 +201,9 @@ theorem covers2 {s : State} {p : Params} (L : Lay kgR (kgW p) s) {a b : Ptr} {la
/-- A state of the function, where a call can be made. -/
structure Site (p : Params) (s : State) : Prop where
lay : Lay kgR (kgW p) s
h24 : 24 ≤ (s.gpr .rsp).toNat
h32 : 32 ≤ (s.gpr .rsp).toNat

theorem Site.h24 {p : Params} {s : State} (S : Site p s) : 24 ≤ (s.gpr .rsp).toNat := by have := S.h32; omega

/-- The registers that hold the pointers of the layout. -/
abbrev kgRegs : List Reg := [.rbx, .rbp, .r12, .r13]
Expand Down Expand Up @@ -594,6 +609,86 @@ theorem rejNttAt_tr {c : Prog isa} (hc : Callee c fun stk => Spec.MlDsa.rejNTTCo

end

/-! ## `RejNTTPoly` four times -/

section
variable {p : Params} {a w : Ptr} (ha : PtrOk a) (hw : PtrOk w)
(h1 : sepB (kgB p) (sc oSA4) 136 a 4096 = true) (h2 : sepB (kgB p) (sc oSA4) 136 w 8192 = true)
(h3 : sepB (kgB p) a 4096 w 8192 = true)
(w1 : inB (kgW p) a 4096 = true) (w2 : inB (kgW p) w 8192 = true)

include h1 h2 h3 in
theorem rej4_pre {s s1 : State} (S : Site p s)
(hv : s1.gpr .rdi = pa s (sc oSA4) ∧ s1.gpr .rsi = pa s a ∧ s1.gpr .rdx = pa s w)
(k : Keep MlKem.X86_64.argRegs s s1) :
(Spec.MlDsa.rejNTT4Contract X86_64.abi 24).pre
(s1.callEntry.withRegions [⟨pa s (sc oSA4), 136⟩] [⟨pa s a, 4096⟩, ⟨pa s w, 8192⟩]) := by
obtain ⟨i1, i2, _⟩ := sepB_spec h1
obtain ⟨_, i3, _⟩ := sepB_spec h3
have L := S.lay
have hsp := keep_rsp k
sig_pre [Spec.MlDsa.rejNTT4Contract, Spec.MlDsa.rejNTT4Sig, X86_64.abi, VG.X86_64.argRegs]
simp only [hv.1, hv.2.1, hv.2.2]
exact ⟨ceWf24 hsp S.h32, trivial, trivial, L.disj h1, L.disj h2, L.disj h3, ceD1 L hsp i1, ceD1 L hsp i2,
ceD1 L hsp i3, ceD24 L hsp i1, ceD24 L hsp i2, ceD24 L hsp i3, L.nwp i1, L.nwp i2, L.nwp i3⟩

include ha hw h1 h2 h3 w1 w2 in
theorem rej4At_ok {c : Prog isa} {sfx : String} (hc : Callee4 c) {s : State} (S : Site p s) :
WP isa (rej4At c sfx a w) s fun s' => Post s s' [⟨pa s a, 4096⟩, ⟨pa s w, 8192⟩] ∧ MX s' = MX s ∧
((s'.gpr .rax).setWidth 32 = 1 → ∀ k < 4, Spec.MlDsa.Reduced s'.mem (Spec.MlDsa.poly4 (pa s a) k)) ∧
(((s'.gpr .rax).setWidth 32 = 1 ∧ ∀ k < 4, ∃ b : Spec.MlDsa.Bounds,
Spec.MlDsa.rejNTTPoly b.rejNTT (Spec.MlDsa.seed4 s.mem (pa s (sc oSA4)) k) =
some (Spec.MlDsa.polyAt s'.mem (Spec.MlDsa.poly4 (pa s a) k))) ∨
((s'.gpr .rax).setWidth 32 = 0 ∧ ∃ k < 4,
Spec.MlDsa.rejNTTPoly Spec.MlDsa.minBounds.rejNTT (Spec.MlDsa.seed4 s.mem (pa s (sc oSA4)) k) = none)) := by
obtain ⟨i1, _, _⟩ := sepB_spec h1
have L := S.lay
refine WP.mono (glueCallMx_ok hc.verified.1 hc.nosp hc.depth
(noLd_append (noLd_append (lea_noLd _ _) (lea_noLd _ _)) (lea_noLd _ _))
(glue3_ok (sc_ok oSA4 (by decide)) ha hw s) (fun s1 hv _ k => rej4_pre h1 h2 h3 S hv k)
(covers_rww L i1 w1 w2) (covers2 L w1 w2))
fun s' ⟨hP, hx, s1, hv, hm, k, s₂, hm₂, hg₂, hpost⟩ => ⟨hP, hx, ?_⟩
have hsp := keep_rsp k
sig_post [Spec.MlDsa.rejNTT4Contract, Spec.MlDsa.rejNTT4Sig, X86_64.abi, VG.X86_64.argRegs] at hpost
simp only [hv.1, hv.2.1, hm₂, hg₂ .rax (by decide)] at hpost
have hseed : ∀ k < 4, Spec.MlDsa.seed4 (ceM s1) (pa s (sc oSA4)) k = Spec.MlDsa.seed4 s.mem (pa s (sc oSA4)) k :=
fun k hk => by
unfold Spec.MlDsa.seed4
rw [← hm]
refine Proof.MlKem.bytesAt_congr fun i hi => ?_
rw [BitVec.add_assoc, ← BitVec.ofNat_add]
exact callEntry_bytes s1 (R := ⟨pa s (sc oSA4), 136⟩) (k16 s1 (by rw [hsp]; exact L.stkD i1))
(show 136 ≤ 2 ^ 64 by decide) (show 34 * k + i < 136 by omega)
obtain ⟨hr, ho⟩ := hpost
refine ⟨hr, ?_⟩
rcases ho with ⟨h1', hb⟩ | ⟨h0, k, hk, hn⟩
· exact .inl ⟨h1', fun k hk => by rw [← hseed k hk]; exact hb k hk⟩
· exact .inr ⟨h0, k, hk, by rw [← hseed k hk]; exact hn⟩

include ha hw h1 h2 h3 w1 w2 in
theorem rej4At_tr {c : Prog isa} {sfx : String} (hc : Callee4 c) (hba : a.1 ∈ kgRegs) (hbw : w.1 ∈ kgRegs) :
RelCT isa (fun x y => Two p x y ∧ bytesAt x.mem (pa x (sc oSA4)) 136 = bytesAt y.mem (pa y (sc oSA4)) 136)
(rej4At c sfx a w) fun _ _ => True := by
obtain ⟨i1, _, _⟩ := sepB_spec h1
have g := fun x => glue3_ok (sc_ok oSA4 (by decide)) ha hw x
refine glueCall_tr hc.verified.1 hc.verified.2.1
(moves_tr (nomem_append (nomem_append (lea_nomem _ _) (lea_nomem _ _)) (lea_nomem _ _)))
(fun x y _ => ⟨g x, g y⟩)
fun x y x1 y1 ⟨T, e⟩ ⟨⟨hv1, hm1⟩, k1⟩ ⟨⟨hv2, hm2⟩, k2⟩ =>
⟨_, _, _, _, rej4_pre h1 h2 h3 T.sx hv1 k1, rej4_pre h1 h2 h3 T.sy hv2 k2, ?_,
by rw [k1.2.1, k1.2.2]; exact covers_rww T.sx.lay i1 w1 w2, by rw [k1.2.2]; exact covers2 T.sx.lay w1 w2,
by rw [k2.2.1, k2.2.2]; exact covers_rww T.sy.lay i1 w1 w2, by rw [k2.2.2]; exact covers2 T.sy.lay w1 w2,
by rw [keep_rsp k1, keep_rsp k2, T.rsp]⟩
sig_pub [Spec.MlDsa.rejNTT4Contract, Spec.MlDsa.rejNTT4Sig, X86_64.abi, VG.X86_64.argRegs]
simp only [hv1.1, hv1.2.1, hv1.2.2, hv2.1, hv2.2.1, hv2.2.2, T.pa hba, T.pa hbw,
T.pa (q := sc oSA4) (by decide), and_true]
refine ⟨by rw [keep_rsp k1, keep_rsp k2, T.rsp], ?_⟩
rw [T.pa (q := sc oSA4) (by decide)] at e
rw [ce_bytesAt' (by decide) (by rw [keep_rsp k1, ← T.pa (q := sc oSA4) (by decide)]; exact T.sx.lay.stkD i1),
ce_bytesAt' (by decide) (by rw [keep_rsp k2]; exact T.sy.lay.stkD i1), hm1, hm2, e]

end

/-! ## Moves with immediates -/

theorem w32_toNat {v : Nat} (h : v < 2 ^ 32) : (BitVec.setWidth 32 (BitVec.ofNat 64 v)).toNat = v := by
Expand Down
Loading
Loading