diff --git a/README.md b/README.md index 035f37374..aa190785e 100644 --- a/README.md +++ b/README.md @@ -219,7 +219,7 @@ yours to keep: ✅ AES, PMULL -❌ +✅ ❌ diff --git a/bench/benches/primitives/cmac_aes.rs b/bench/benches/primitives/cmac_aes.rs index 51a8dbf1c..e7bb7a381 100644 --- a/bench/benches/primitives/cmac_aes.rs +++ b/bench/benches/primitives/cmac_aes.rs @@ -8,7 +8,7 @@ pub const USES: &[&str] = &["cmac_aes", "aes"]; /// The MAC of a message with a 16-byte key (setup included), computed and /// verified. -#[cfg(any(target_arch = "x86_64", target_arch = "aarch64"))] +#[cfg(any(target_arch = "x86_64", target_arch = "aarch64", target_arch = "arm"))] pub fn bench(c: &mut Criterion) { use std::hint::black_box; @@ -63,5 +63,5 @@ pub fn bench(c: &mut Criterion) { g.finish(); } -#[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))] +#[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64", target_arch = "arm")))] pub fn bench(_: &mut Criterion) {} diff --git a/lean/VerifiedGarbage/Artifacts/CmacAes/Arm.lean b/lean/VerifiedGarbage/Artifacts/CmacAes/Arm.lean new file mode 100644 index 000000000..269ba4ca0 --- /dev/null +++ b/lean/VerifiedGarbage/Artifacts/CmacAes/Arm.lean @@ -0,0 +1,53 @@ +import VerifiedGarbage.TCB.Arm.Target +import VerifiedGarbage.Proof.CmacAes.Arm.Verified + +/-! +# AES-CMAC (NIST SP 800-38B) on ARMv7 + +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`. An +artifact made from a function's `Api` (in `Spec/`, reviewed with the +contract) takes them from there, 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. + +Each function calls `vg_aes_ctr32` in a frame that pushes its two stack +arguments, so uses 8 bytes of stack. +-/ + +namespace VG.Artifacts.CmacAes.Arm + +open VG.Proof.CmacAes.Arm + +/-- How the functions encrypt a block. -/ +def ctrNote : String := "This implementation encrypts each block with `vg_aes_ctr32`." + +def artifacts : List Artifact := [ + { Spec.Cmac.aesSubkeysApi with + target := Arm.target + doc := Spec.Cmac.aesSubkeysApi.doc (notes := [ctrNote]) + code := Impl.CmacAes.Arm.subkeys + contract := Spec.Cmac.aesSubkeysContract Arm.abi 8 + stack := 8 + verified := subkeys_verified + spSafe := Code.all_of_forall (fun _ => rfl) _ }, + { Spec.Cmac.aesUpdateApi with + target := Arm.target + doc := Spec.Cmac.aesUpdateApi.doc (notes := [ctrNote]) + code := Impl.CmacAes.Arm.update + contract := Spec.Cmac.aesUpdateContract Arm.abi 8 + stack := 8 + verified := update_verified + spSafe := Code.all_of_forall (fun _ => rfl) _ }, + { Spec.Cmac.aesFinalizeApi with + target := Arm.target + doc := Spec.Cmac.aesFinalizeApi.doc (notes := [ctrNote]) + code := Impl.CmacAes.Arm.finalize + contract := Spec.Cmac.aesFinalizeContract Arm.abi 8 + stack := 8 + verified := finalize_verified + spSafe := Code.all_of_forall (fun _ => rfl) _ }] + +end VG.Artifacts.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Impl/CmacAes/Arm.lean b/lean/VerifiedGarbage/Impl/CmacAes/Arm.lean new file mode 100644 index 000000000..d1dd98781 --- /dev/null +++ b/lean/VerifiedGarbage/Impl/CmacAes/Arm.lean @@ -0,0 +1,183 @@ +import VerifiedGarbage.Impl.Aes.Arm.Ctr32 + +/-! +# AES-CMAC: 32-bit ARM implementation + +`vg_cmac_aes_subkeys(schedule = r0, rounds = r1, subkeys = r2, scratch = r3)`, +`vg_cmac_aes_update(schedule = r0, rounds = r1, state = r2, data = r3, n = [sp], scratch = [sp + 4])` +and `vg_cmac_aes_finalize(key = r0, rounds = r1, state = r2, last = r3, last_len = [sp], scratch = [sp + 4])` +(see `VG.Spec.Cmac.aesSubkeysContract` and the others), composed of calls of +the verified `vg_aes_ctr32`, one block at a time: with a counter block `X` +and a zero data block, it leaves `CIPH_K(X)` in the data block. + +`vg_aes_ctr32(schedule, rounds, counter, data, n, scratch)` takes `n` and +`scratch` on the stack: a frame pushes them (`push {rA, rB}`, `n = 1` in +`rA` at `[sp]`) around each call, and its pop loads `rA` back. So the +functions use 8 bytes of stack. + +The scratch buffer (2176 bytes): `[0, 2048)` is the working space of +`vg_aes_ctr32`, `[2048, 2064)` the counter block, and `[2064, 2096)` our +caller's callee-saved registers and our return address `lr`. + +* `subkeys` computes `L = CIPH_K(0)` into the first block of `subkeys`, and + doubles it there (`K1`) and into the second block (`K2`): the block as a + big-endian 128-bit integer in `r0:r1:r2:r3`, shifted left by one bit, and + XORed with `0x87` masked by the bit shifted out. `r6` holds `subkeys` and + `r5` the scratch buffer across the call. +* `update` keeps its arguments in `r4` (schedule), `r5` (rounds), `r6` + (state), `r7` (data), `r8` (blocks left) and `r10` (scratch) across the + calls; each block, the counter block is `C ⊕ Mᵢ` and the state, zeroed, + receives `CIPH_K(C ⊕ Mᵢ)`. +* `finalize` forms `Mₙ` in the counter block: `Mₙ* ⊕ K1` for a complete + block, else `Mₙ*` copied a byte at a time onto zeros, `0x80` after it, and + XORed with `K2`. It XORs in the chaining value and calls `vg_aes_ctr32` + last, keeping only the scratch buffer (in `r5`) across the call. + +The model has no register-offset addressing: the last bytes are copied +through advancing pointers, counting down with `subs`. Only the pointers, +`rounds`, `n` and `last_len` can affect timing: the branches are on `n` and +`last_len`, and the doubling is masked. +-/ + +namespace VG.Impl.CmacAes.Arm + +open VG.Arm + +/-- `mov d, n`. -/ +def mov (d n : Reg) : Instr := .mov d (.reg n) + +/-- The offset of the counter block in the scratch buffer. -/ +def cOff : Nat := 2048 + +/-- The call of `vg_aes_ctr32`, with `n` (1) in `ra` and the scratch buffer in +`rb` pushed as its stack arguments, and `ra` popped. -/ +def ctrCall (ra rb : Reg) : Prog isa := + .frame (.push [ra, rb]) (.call "vg_aes_ctr32" Impl.Aes.Arm.ctr32) (.pop ra 8) + +/-! ## `vg_cmac_aes_subkeys` -/ + +/-- Saves `r4`–`r6` and `lr`, keeps `subkeys` in `r6` and the scratch buffer +in `r5`, zeroes the counter block and the first block of `subkeys`, and sets +up the arguments of `vg_aes_ctr32`. -/ +def subkeysPre : List Instr := + [.str .r4 .r3 2064, .str .r5 .r3 2068, .str .r6 .r3 2072, .str .lr .r3 2076, mov .r6 .r2, mov .r5 .r3, + .mov .r12 (.imm 0), .str .r12 .r3 cOff, .str .r12 .r3 (cOff + 4), .str .r12 .r3 (cOff + 8), + .str .r12 .r3 (cOff + 12), .str .r12 .r2 0, .str .r12 .r2 4, .str .r12 .r2 8, .str .r12 .r2 12, + .dp .add .r2 .r5 (.imm (BitVec.ofNat 32 cOff)), mov .r3 .r6, .mov .r4 (.imm 1)] + +/-- The block at `r6 + src`, doubled (`VG.Spec.Cmac.dbl 16`), to `r6 + dst`. -/ +def dbl (src dst : Nat) : List Instr := + [.ldr .r0 .r6 src, .ldr .r1 .r6 (src + 4), .ldr .r2 .r6 (src + 8), .ldr .r3 .r6 (src + 12), + .rev .r0 .r0, .rev .r1 .r1, .rev .r2 .r2, .rev .r3 .r3, + .mov .r12 (.shifted .r0 .lsr 31), .mov .r4 (.imm 0), .dp .sub .r12 .r4 (.reg .r12), + .dp .and .r12 .r12 (.imm 0x87), + .mov .r0 (.shifted .r0 .lsl 1), .dp .orr .r0 .r0 (.shifted .r1 .lsr 31), + .mov .r1 (.shifted .r1 .lsl 1), .dp .orr .r1 .r1 (.shifted .r2 .lsr 31), + .mov .r2 (.shifted .r2 .lsl 1), .dp .orr .r2 .r2 (.shifted .r3 .lsr 31), + .mov .r3 (.shifted .r3 .lsl 1), .dp .eor .r3 .r3 (.reg .r12), + .rev .r0 .r0, .rev .r1 .r1, .rev .r2 .r2, .rev .r3 .r3, + .str .r0 .r6 dst, .str .r1 .r6 (dst + 4), .str .r2 .r6 (dst + 8), .str .r3 .r6 (dst + 12)] + +/-- `K1` over `L`, `K2` after it, and the saved registers restored. -/ +def subkeysPost : List Instr := + dbl 0 0 ++ dbl 0 16 ++ + [.ldr .r4 .r5 2064, .ldr .r6 .r5 2072, .ldr .lr .r5 2076, .ldr .r5 .r5 2068] + +def subkeys : Prog isa := + .seq (.block subkeysPre) (.seq (ctrCall .r4 .r5) (.block subkeysPost)) + +/-! ## `vg_cmac_aes_update` -/ + +/-- The registers saved in the scratch buffer, and where (`r10`, the base of +the restore, last). -/ +def saved : List (Reg × Nat) := + [(.r4, 2064), (.r5, 2068), (.r6, 2072), (.r7, 2076), (.r8, 2080), (.r9, 2084), (.lr, 2092), + (.r10, 2088)] + +/-- Saves them, with the scratch buffer (the second stack argument) in `r12`. -/ +def save : List Instr := .ldrSp .r12 4 :: saved.map fun (r, d) => .str r .r12 d + +/-- Restores them, with `r10` (restored last) the scratch buffer. -/ +def restore : List Instr := saved.map fun (r, d) => .ldr r .r10 d + +/-- The arguments to their registers; Z is set if there are no blocks. -/ +def setup : List Instr := + [mov .r4 .r0, mov .r5 .r1, mov .r6 .r2, mov .r7 .r3, .ldrSp .r8 0, mov .r10 .r12, .cmp .r8 (.imm 0)] + +/-- The counter block `C ⊕ Mᵢ` (the state at `r6`, the block at `r7`), and +the state zeroed. -/ +def chainIn : List Instr := + [.ldr .r0 .r6 0, .ldr .r1 .r7 0, .dp .eor .r0 .r0 (.reg .r1), .str .r0 .r10 cOff, + .ldr .r0 .r6 4, .ldr .r1 .r7 4, .dp .eor .r0 .r0 (.reg .r1), .str .r0 .r10 (cOff + 4), + .ldr .r0 .r6 8, .ldr .r1 .r7 8, .dp .eor .r0 .r0 (.reg .r1), .str .r0 .r10 (cOff + 8), + .ldr .r0 .r6 12, .ldr .r1 .r7 12, .dp .eor .r0 .r0 (.reg .r1), .str .r0 .r10 (cOff + 12), + .mov .r0 (.imm 0), .str .r0 .r6 0, .str .r0 .r6 4, .str .r0 .r6 8, .str .r0 .r6 12] + +/-- The arguments of `vg_aes_ctr32` for the block. -/ +def updArgs : List Instr := + [mov .r0 .r4, mov .r1 .r5, .dp .add .r2 .r10 (.imm (BitVec.ofNat 32 cOff)), mov .r3 .r6, .mov .r9 (.imm 1)] + +/-- On to the next block (Z is set when none are left). -/ +def advance : List Instr := [.dp .add .r7 .r7 (.imm 16), .subs .r8 .r8 (.imm 1)] + +/-- One block. -/ +def body : Prog isa := + .seq (.block (chainIn ++ updArgs)) (.seq (ctrCall .r9 .r10) (.block advance)) + +def update : Prog isa := + .seq (.block (save ++ setup)) (.seq (.ite .eq (.block []) (.loop body .ne)) (.block restore)) + +/-! ## `vg_cmac_aes_finalize` -/ + +/-- The four words at `pb + pd` and `qb + qd` XORed into `cb + cd`, with +`r12` and `lr`. -/ +def xor4 (pb qb cb : Reg) (pd qd cd : Nat) : List Instr := + (List.range 4).flatMap fun i => + [.ldr .r12 pb (pd + 4 * i), .ldr .lr qb (qd + 4 * i), .dp .eor .r12 .r12 (.reg .lr), + .str .r12 cb (cd + 4 * i)] + +/-- Saves `r4`, `r5` and `lr`, keeps the scratch buffer in `r5` and `last_len` +in `r4`; Z is set if `last_len` is 16. -/ +def finSave : List Instr := + [.ldrSp .r12 4, .str .r4 .r12 2064, .str .r5 .r12 2068, .str .lr .r12 2072, mov .r5 .r12, + .ldrSp .r4 0, .cmp .r4 (.imm 16)] + +/-- `Mₙ = Mₙ* ⊕ K1` (`K1` at `r0 + 240`), for a complete last block. -/ +def full : List Instr := xor4 .r3 .r0 .r5 0 240 cOff + +/-- The counter block zeroed, with `lr` pointing at it; Z is set if +`last_len` is 0. -/ +def zero : List Instr := + [.mov .r12 (.imm 0), .str .r12 .r5 cOff, .str .r12 .r5 (cOff + 4), .str .r12 .r5 (cOff + 8), + .str .r12 .r5 (cOff + 12), .dp .add .lr .r5 (.imm (BitVec.ofNat 32 cOff)), .cmp .r4 (.imm 0)] + +/-- The `r4` (nonzero) bytes at `r3` copied to `lr`, advancing both. -/ +def copy : Prog isa := + .loop (.block [.ldrb .r12 .r3 0, .strb .r12 .lr 0, .dp .add .r3 .r3 (.imm 1), + .dp .add .lr .lr (.imm 1), .subs .r4 .r4 (.imm 1)]) .ne + +/-- `0x80` after the bytes (at `lr`), and the block XORed with `K2` (at +`r0 + 256`). -/ +def padK2 : List Instr := + [.mov .r12 (.imm 0x80), .strb .r12 .lr 0] ++ xor4 .r5 .r0 .r5 cOff 256 cOff + +/-- `Mₙ = K2 ⊕ (Mₙ* ‖ 10ʲ)`, for a partial last block (`last_len < 16`). -/ +def partialBlock : Prog isa := + .seq (.block zero) (.seq (.ite .eq (.block []) copy) (.block padK2)) + +/-- The counter block `C ⊕ Mₙ` (the state at `r2`), the state zeroed, and the +arguments of `vg_aes_ctr32` but the schedule (`r0`) and the rounds (`r1`), +which are ours. -/ +def finArgs : List Instr := + xor4 .r5 .r2 .r5 cOff 0 cOff ++ + [.mov .r12 (.imm 0), .str .r12 .r2 0, .str .r12 .r2 4, .str .r12 .r2 8, .str .r12 .r2 12, + mov .r3 .r2, .dp .add .r2 .r5 (.imm (BitVec.ofNat 32 cOff)), .mov .r4 (.imm 1)] + +/-- Everything before the call. -/ +def finPre : Prog isa := + .seq (.block finSave) (.seq (.ite .eq (.block full) partialBlock) (.block finArgs)) + +def finalize : Prog isa := + .seq finPre (.seq (ctrCall .r4 .r5) (.block [.ldr .r4 .r5 2064, .ldr .lr .r5 2072, .ldr .r5 .r5 2068])) + +end VG.Impl.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/Cmac/Block32.lean b/lean/VerifiedGarbage/Proof/Cmac/Block32.lean new file mode 100644 index 000000000..9c99c0f53 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/Cmac/Block32.lean @@ -0,0 +1,104 @@ +import VerifiedGarbage.Proof.Cmac.Mem32 +import VerifiedGarbage.Proof.Cmac.Block + +/-! +# CMAC: blocks formed a 32-bit word at a time + +Untrusted: everything here is checked by Lean. What the 32-bit targets' +stores leave: the XOR of two blocks stored a word at a time (`xor4Mem`; the +block written may be one of those read, as long as no word written is read +afterwards, `Sep4`), a zeroed block (`zero4`), and a counter block `C = P ⊕ Q` +with `P` zeroed (`chainMem4`). +-/ + +namespace VG.Proof.Cmac + +open VG + +/-- No word written at `c` is read at `p` after it: the word `i` written is +disjoint from every word `j > i` of `p`. -/ +def Sep4 (c p : Addr) : Prop := + ∀ i < 4, ∀ j < 4, i < j → (⟨c + BitVec.ofNat 64 (4 * i), 4⟩ : Region).Disjoint ⟨p + BitVec.ofNat 64 (4 * j), 4⟩ + +theorem Sep4.self (c : Addr) : Sep4 c c := fun i hi j hj hij => + Offset.disjoint c (by omega) (by omega) (by omega) + +theorem Sep4.of_disjoint {c p : Addr} (h : (⟨c, 16⟩ : Region).Disjoint ⟨p, 16⟩) : Sep4 c p := + fun i hi j hj _ => (h.sub_left (Offset.sub_base c (by omega))).sub_right (Offset.sub_base p (by omega)) + +theorem readW_writeW_disj {m : Mem} {a b : Addr} (v : BitVec 32) (h : (⟨a, 4⟩ : Region).Disjoint ⟨b, 4⟩) : + (m.writeW a v).readW b 32 = m.readW b 32 := + Mem.readW_writeW_sep (h.symm.sep (Region.contains_self _ _) (Region.contains_self _ _)) (by decide) + +/-- The memory after storing at `c` the XOR of the blocks at `p` and `q`, a +word at a time. -/ +def xor4Mem (m : Mem) (c p q : Addr) : Mem := + let m₁ := m.writeW c (m.readW p 32 ^^^ m.readW q 32) + let m₂ := m₁.writeW (c + BitVec.ofNat 64 4) + (m₁.readW (p + BitVec.ofNat 64 4) 32 ^^^ m₁.readW (q + BitVec.ofNat 64 4) 32) + let m₃ := m₂.writeW (c + BitVec.ofNat 64 8) + (m₂.readW (p + BitVec.ofNat 64 8) 32 ^^^ m₂.readW (q + BitVec.ofNat 64 8) 32) + m₃.writeW (c + BitVec.ofNat 64 12) (m₃.readW (p + BitVec.ofNat 64 12) 32 ^^^ m₃.readW (q + BitVec.ofNat 64 12) 32) + +theorem xor4Mem_eq (m : Mem) {c p q : Addr} (hp : Sep4 c p) (hq : Sep4 c q) : + xor4Mem m c p q = store4 m c (m.readW p 32 ^^^ m.readW q 32) + (m.readW (p + BitVec.ofNat 64 4) 32 ^^^ m.readW (q + BitVec.ofNat 64 4) 32) + (m.readW (p + BitVec.ofNat 64 8) 32 ^^^ m.readW (q + BitVec.ofNat 64 8) 32) + (m.readW (p + BitVec.ofNat 64 12) 32 ^^^ m.readW (q + BitVec.ofNat 64 12) 32) := by + have e (x : Addr) (h : Sep4 c x) (i j : Nat) (hi : i < 4) (hj : j < 4) (hij : i < j) : + (⟨c + BitVec.ofNat 64 (4 * i), 4⟩ : Region).Disjoint ⟨x + BitVec.ofNat 64 (4 * j), 4⟩ := h i hi j hj hij + have p01 : (⟨c, 4⟩ : Region).Disjoint ⟨p + BitVec.ofNat 64 4, 4⟩ := by simpa using e p hp 0 1 (by decide) (by decide) (by decide) + have q01 : (⟨c, 4⟩ : Region).Disjoint ⟨q + BitVec.ofNat 64 4, 4⟩ := by simpa using e q hq 0 1 (by decide) (by decide) (by decide) + have p02 : (⟨c, 4⟩ : Region).Disjoint ⟨p + BitVec.ofNat 64 8, 4⟩ := by simpa using e p hp 0 2 (by decide) (by decide) (by decide) + have q02 : (⟨c, 4⟩ : Region).Disjoint ⟨q + BitVec.ofNat 64 8, 4⟩ := by simpa using e q hq 0 2 (by decide) (by decide) (by decide) + have p03 : (⟨c, 4⟩ : Region).Disjoint ⟨p + BitVec.ofNat 64 12, 4⟩ := by simpa using e p hp 0 3 (by decide) (by decide) (by decide) + have q03 : (⟨c, 4⟩ : Region).Disjoint ⟨q + BitVec.ofNat 64 12, 4⟩ := by simpa using e q hq 0 3 (by decide) (by decide) (by decide) + have p12 : (⟨c + BitVec.ofNat 64 4, 4⟩ : Region).Disjoint ⟨p + BitVec.ofNat 64 8, 4⟩ := e p hp 1 2 (by decide) (by decide) (by decide) + have q12 : (⟨c + BitVec.ofNat 64 4, 4⟩ : Region).Disjoint ⟨q + BitVec.ofNat 64 8, 4⟩ := e q hq 1 2 (by decide) (by decide) (by decide) + have p13 : (⟨c + BitVec.ofNat 64 4, 4⟩ : Region).Disjoint ⟨p + BitVec.ofNat 64 12, 4⟩ := e p hp 1 3 (by decide) (by decide) (by decide) + have q13 : (⟨c + BitVec.ofNat 64 4, 4⟩ : Region).Disjoint ⟨q + BitVec.ofNat 64 12, 4⟩ := e q hq 1 3 (by decide) (by decide) (by decide) + have p23 : (⟨c + BitVec.ofNat 64 8, 4⟩ : Region).Disjoint ⟨p + BitVec.ofNat 64 12, 4⟩ := e p hp 2 3 (by decide) (by decide) (by decide) + have q23 : (⟨c + BitVec.ofNat 64 8, 4⟩ : Region).Disjoint ⟨q + BitVec.ofNat 64 12, 4⟩ := e q hq 2 3 (by decide) (by decide) (by decide) + simp only [xor4Mem, store4, readW_writeW_disj _ p01, readW_writeW_disj _ q01, readW_writeW_disj _ p02, + readW_writeW_disj _ q02, readW_writeW_disj _ p03, readW_writeW_disj _ q03, readW_writeW_disj _ p12, + readW_writeW_disj _ q12, readW_writeW_disj _ p13, readW_writeW_disj _ q13, readW_writeW_disj _ p23, + readW_writeW_disj _ q23] + +theorem xor4Mem_frame (m : Mem) (c p q : Addr) : Frame [⟨c, 16⟩] m (xor4Mem m c p q) := by + have k (d : Nat) (h : d + 4 ≤ 16) : (⟨c, 16⟩ : Region).Contains (c + BitVec.ofNat 64 d) 4 := + Offset.contains_base c h (by omega) + have k0 : (⟨c, 16⟩ : Region).Contains c 4 := by simpa using k 0 (by decide) + exact ((((Frame.refl _ _).writeW (List.mem_singleton_self _) _ k0).writeW (List.mem_singleton_self _) _ + (k 4 (by decide))).writeW (List.mem_singleton_self _) _ (k 8 (by decide))).writeW + (List.mem_singleton_self _) _ (k 12 (by decide)) + +theorem xor4Mem_bytes (m : Mem) {c p q : Addr} (hp : Sep4 c p) (hq : Sep4 c q) : + Spec.Aes.bytesAt (xor4Mem m c p q) c 16 = + Spec.Cmac.xor (Spec.Aes.bytesAt m p 16) (Spec.Aes.bytesAt m q 16) := by + rw [xor4Mem_eq m hp hq, bytesAt_store4, xor_words4] + +/-- The memory after zeroing the block at `c`, a word at a time. -/ +def zero4 (m : Mem) (c : Addr) : Mem := store4 m c 0 0 0 0 + +theorem zero4_bytes (m : Mem) (c : Addr) : Spec.Aes.bytesAt (zero4 m c) c 16 = Spec.Cmac.zeros 16 := by + rw [zero4, bytesAt_store4, le4_zero]; decide + +/-- The memory after forming a counter block: the block at `c` is the block +at `p` XORed with the block at `q`, and the block at `p` is zeroed. -/ +def chainMem4 (m : Mem) (c p q : Addr) : Mem := zero4 (xor4Mem m c p q) p + +theorem chainMem4_frame (m : Mem) (C P Q : Addr) : Frame [⟨C, 16⟩, ⟨P, 16⟩] m (chainMem4 m C P Q) := + ((xor4Mem_frame m C P Q).mono (fun r hr => by simp only [List.mem_singleton] at hr; simp [hr])).trans + ((frame_store4 P 0 0 0 0).mono (fun r hr => by simp only [List.mem_singleton] at hr; simp [hr])) + +theorem chainMem4_state (m : Mem) (C P Q : Addr) : + Spec.Aes.bytesAt (chainMem4 m C P Q) P 16 = Spec.Cmac.zeros 16 := zero4_bytes _ _ + +theorem chainMem4_counter (m : Mem) {C P Q : Addr} (hcp : (⟨C, 16⟩ : Region).Disjoint ⟨P, 16⟩) + (hcq : (⟨C, 16⟩ : Region).Disjoint ⟨Q, 16⟩) : + Spec.Aes.bytesAt (chainMem4 m C P Q) C 16 = + Spec.Cmac.xor (Spec.Aes.bytesAt m P 16) (Spec.Aes.bytesAt m Q 16) := by + rw [chainMem4, zero4, bytesAt_frame16 (frame_store4 P 0 0 0 0) (by simpa using hcp), + xor4Mem_bytes m (Sep4.of_disjoint hcp) (Sep4.of_disjoint hcq)] + +end VG.Proof.Cmac diff --git a/lean/VerifiedGarbage/Proof/Cmac/Dbl32.lean b/lean/VerifiedGarbage/Proof/Cmac/Dbl32.lean new file mode 100644 index 000000000..de25ce1a1 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/Cmac/Dbl32.lean @@ -0,0 +1,208 @@ +import VerifiedGarbage.Proof.Cmac.Dbl +import VerifiedGarbage.Proof.Cmac.Mem32 +import VerifiedGarbage.Proof.Framework.Bswap +import VerifiedGarbage.Proof.Gcm.Bits + +/-! +# CMAC: doubling a block in four 32-bit words + +Untrusted: everything here is checked by Lean. The 32-bit targets load a +block as four byte-reversed words (`byteRev32`), the block as a big-endian +128-bit integer `b₀ ++ b₁ ++ b₂ ++ b₃` (`ofBytes_rev4`), double it a word at a +time (`dbl_words4`), and store the words byte-reversed again (`le4_rev4`). +-/ + +namespace VG.Proof.Cmac + +open VG Spec.Cmac + +theorem ofBytes_toBytes (x : Spec.Gcm.Block) : Spec.Gcm.ofBytes (Spec.Gcm.toBytes x) = x := by + apply BitVec.eq_of_getLsbD_eq + intro p hp + obtain ⟨i, j, hi, hj, rfl⟩ : ∃ i j, i < 16 ∧ j < 8 ∧ p = 8 * (15 - i) + j := + ⟨15 - p / 8, p % 8, by omega, by omega, by omega⟩ + rw [ofBytes_bit (toBytes_length x) hi hj, Proof.Aes.toBytes_getD _ hi, BitVec.getLsbD_extractLsb'] + simp [hj] + +theorem getLsbD_byteRev32 (a : BitVec 32) {i : Nat} (hi : i < 32) : + (byteRev32 a).getLsbD i = a.getLsbD (8 * (3 - i / 8) + i % 8) := by + simp only [byteRev32] + rw [VG.getLsbD_cat4] + rcases (by omega : i < 8 ∨ (8 ≤ i ∧ i < 16) ∨ (16 ≤ i ∧ i < 24) ∨ 24 ≤ i) with h | h | h | h + · simp only [h, ite_true, BitVec.getLsbD_extractLsb', decide_true, Bool.true_and]; congr 1; omega + · simp only [show ¬ i < 8 by omega, h.2, ite_true, ite_false, BitVec.getLsbD_extractLsb', + show i - 8 < 8 by omega, decide_true, Bool.true_and]; congr 1; omega + · simp only [show ¬ i < 8 by omega, show ¬ i < 16 by omega, h.2, ite_true, ite_false, + BitVec.getLsbD_extractLsb', show i - 16 < 8 by omega, decide_true, Bool.true_and]; congr 1; omega + · simp only [show ¬ i < 8 by omega, show ¬ i < 16 by omega, show ¬ i < 24 by omega, ite_false, + BitVec.getLsbD_extractLsb', show i - 24 < 8 by omega, decide_true, Bool.true_and]; congr 1; omega + +theorem getD_le4_append4 (a b c d : BitVec 32) {k : Nat} (hk : k < 16) : + (le4 a ++ le4 b ++ le4 c ++ le4 d).getD k 0 = + (if k < 4 then a else if k < 8 then b else if k < 12 then c else d).extractLsb' (8 * (k % 4)) 8 := by + simp only [List.append_assoc, List.getD_eq_getElem?_getD] + rcases (by omega : k < 4 ∨ (4 ≤ k ∧ k < 8) ∨ (8 ≤ k ∧ k < 12) ∨ 12 ≤ k) with h | h | h | h + · rw [List.getElem?_append_left (by rw [length_le4]; omega), ← List.getD_eq_getElem?_getD, getD_le4 _ h] + simp [h, Nat.mod_eq_of_lt h] + · rw [List.getElem?_append_right (by rw [length_le4]; omega), length_le4, + List.getElem?_append_left (by rw [length_le4]; omega), ← List.getD_eq_getElem?_getD, + getD_le4 _ (by omega)] + simp [show ¬ k < 4 by omega, h.2, show k % 4 = k - 4 by omega] + · rw [List.getElem?_append_right (by rw [length_le4]; omega), length_le4, + List.getElem?_append_right (by rw [length_le4]; omega), length_le4, + List.getElem?_append_left (by rw [length_le4]; omega), ← List.getD_eq_getElem?_getD, + getD_le4 _ (by omega)] + simp [show ¬ k < 4 by omega, show ¬ k < 8 by omega, h.2, show k % 4 = k - 4 - 4 by omega] + · rw [List.getElem?_append_right (by rw [length_le4]; omega), length_le4, + List.getElem?_append_right (by rw [length_le4]; omega), length_le4, + List.getElem?_append_right (by rw [length_le4]; omega), length_le4, ← List.getD_eq_getElem?_getD, + getD_le4 _ (by omega)] + simp [show ¬ k < 4 by omega, show ¬ k < 8 by omega, show ¬ k < 12 by omega, show k % 4 = k - 4 - 4 - 4 by omega] + +/-- Storing the byte-reversed words of `a ++ b ++ c ++ d` stores its bytes, +big-endian. -/ +theorem le4_rev4 (a b c d : BitVec 32) : + le4 (byteRev32 a) ++ le4 (byteRev32 b) ++ le4 (byteRev32 c) ++ le4 (byteRev32 d) = + Spec.Gcm.toBytes (a ++ b ++ c ++ d) := by + refine ext16 (by simp [length_le4]) (toBytes_length _) fun k hk => ?_ + rw [Proof.Aes.toBytes_getD _ hk, getD_le4_append4 _ _ _ _ hk] + apply BitVec.eq_of_getLsbD_eq + intro j hj + rw [BitVec.getLsbD_extractLsb', BitVec.getLsbD_extractLsb'] + simp only [hj, decide_true, Bool.true_and] + rcases (by omega : k < 4 ∨ (4 ≤ k ∧ k < 8) ∨ (8 ≤ k ∧ k < 12) ∨ 12 ≤ k) with h | h | h | h + · simp only [h, ite_true, getLsbD_byteRev32 _ (show 8 * (k % 4) + j < 32 by omega)] + simp only [BitVec.getLsbD_append] + split_ifs <;> first | omega | (congr 1; omega) + · simp only [show ¬ k < 4 by omega, h.2, ite_true, ite_false, + getLsbD_byteRev32 _ (show 8 * (k % 4) + j < 32 by omega)] + simp only [BitVec.getLsbD_append] + split_ifs <;> first | omega | (congr 1; omega) + · simp only [show ¬ k < 4 by omega, show ¬ k < 8 by omega, h.2, ite_true, ite_false, + getLsbD_byteRev32 _ (show 8 * (k % 4) + j < 32 by omega)] + simp only [BitVec.getLsbD_append] + split_ifs <;> first | omega | (congr 1; omega) + · simp only [show ¬ k < 4 by omega, show ¬ k < 8 by omega, show ¬ k < 12 by omega, ite_false, + getLsbD_byteRev32 _ (show 8 * (k % 4) + j < 32 by omega)] + simp only [BitVec.getLsbD_append] + split_ifs <;> first | omega | (congr 1; omega) + +theorem mask_eq32 (hi : BitVec 32) : + ((0 : BitVec 32) - (hi >>> 31)) &&& 0x87 = if hi.msb then 0x87 else 0 := by + have h : hi >>> 31 = if hi.msb then 1 else 0 := by + apply BitVec.eq_of_toNat_eq + rw [BitVec.toNat_ushiftRight, Nat.shiftRight_eq_div_pow, BitVec.msb_eq_decide] + have := hi.isLt + by_cases hm : 2 ^ (32 - 1) ≤ hi.toNat + · rw [decide_eq_true hm]; simp; omega + · rw [decide_eq_false hm]; simp; omega + rw [h] + split <;> decide + +theorem bit135_32 : ∀ p < 32, (135 : BitVec 32).getLsbD p = (135 : BitVec 128).getLsbD p := by decide + +/-- The words the 32-bit targets store, from the big-endian words `b₀ … b₃` +of a block: the block doubled. -/ +def dblW0 (b₀ b₁ : BitVec 32) : BitVec 32 := (b₀ <<< 1) ||| (b₁ >>> 31) +def dblW3 (b₀ b₃ : BitVec 32) : BitVec 32 := (b₃ <<< 1) ^^^ (((0 : BitVec 32) - (b₀ >>> 31)) &&& 0x87) + +theorem getLsbD_cat4w (b₀ b₁ b₂ b₃ : BitVec 32) {i : Nat} (hi : i < 128) : + (b₀ ++ b₁ ++ b₂ ++ b₃).getLsbD i = if i < 32 then b₃.getLsbD i else if i < 64 then b₂.getLsbD (i - 32) + else if i < 96 then b₁.getLsbD (i - 64) else b₀.getLsbD (i - 96) := by + simp only [BitVec.getLsbD_append] + split_ifs <;> first | omega | (congr 1; omega) | rfl + +theorem getLsbD_shl1 (x : BitVec 32) {i : Nat} (hi : i < 32) : + (x <<< 1).getLsbD i = (decide (1 ≤ i) && x.getLsbD (i - 1)) := by + rw [BitVec.getLsbD_shiftLeft]; simp only [hi, decide_true, Bool.true_and] + by_cases h : i < 1 <;> simp [h] <;> omega + +theorem getLsbD_shr31 (x : BitVec 32) (i : Nat) : (x >>> 31).getLsbD i = (decide (i = 0) && x.getLsbD 31) := by + rw [BitVec.getLsbD_ushiftRight] + by_cases h : i = 0 + · subst h; simp + · simp only [h, decide_false, Bool.false_and]; exact BitVec.getLsbD_of_ge x _ (by omega) + +theorem getLsbD_shl1_128 (x : BitVec 128) {i : Nat} (hi : i < 128) : + (x <<< 1).getLsbD i = (decide (1 ≤ i) && x.getLsbD (i - 1)) := by + rw [BitVec.getLsbD_shiftLeft]; simp only [hi, decide_true, Bool.true_and] + by_cases h : i < 1 <;> simp [h] <;> omega + +theorem dbl_words4 (b₀ b₁ b₂ b₃ : BitVec 32) : + dblW0 b₀ b₁ ++ dblW0 b₁ b₂ ++ dblW0 b₂ b₃ ++ dblW3 b₀ b₃ = dbl128 (b₀ ++ b₁ ++ b₂ ++ b₃) := by + rw [dblW0, dblW0, dblW0, dblW3, mask_eq32, dbl128, BitVec.msb_append, BitVec.msb_append, BitVec.msb_append] + have h0 : ((32 : Nat) = 0) = False := by simp + have h64 : ((64 : Nat) = 0) = False := by simp + have h96 : ((96 : Nat) = 0) = False := by simp + simp only [h0, h64, h96, ite_false] + apply BitVec.eq_of_getLsbD_eq + intro p hp + rw [getLsbD_cat4w _ _ _ _ hp] + rw [BitVec.getLsbD_xor (x := (b₀ ++ b₁ ++ b₂ ++ b₃) <<< 1), getLsbD_shl1_128 _ hp] + have m87 : (if b₀.msb = true then (135 : BitVec 128) else 0).getLsbD p = + (decide (p < 32) && (if b₀.msb = true then (135 : BitVec 32) else 0).getLsbD p) := by + split + · by_cases h : p < 32 + · simp only [h, decide_true, Bool.true_and]; exact (bit135_32 p h).symm + · simp only [h, decide_false, Bool.false_and]; exact Proof.Cmac.high_0x87 (by omega) + · simp + rw [m87] + rcases (by omega : p = 0 ∨ (1 ≤ p ∧ p < 32) ∨ p = 32 ∨ (33 ≤ p ∧ p < 64) ∨ p = 64 ∨ (65 ≤ p ∧ p < 96) ∨ + p = 96 ∨ (97 ≤ p ∧ p < 128)) with h | h | h | h | h | h | h | h + · subst h + simp + · simp only [show p < 32 from h.2, ite_true, BitVec.getLsbD_xor, getLsbD_shl1 _ h.2, show 1 ≤ p from h.1, + decide_true, Bool.true_and, getLsbD_cat4w _ _ _ _ (show p - 1 < 128 by omega), show p - 1 < 32 by omega] + · subst h + simp only [show ¬ 32 < 32 by decide, show 32 < 64 by decide, ite_false, ite_true, Nat.sub_self, + BitVec.getLsbD_or, getLsbD_shl1 _ (by decide : 0 < 32), getLsbD_shr31, + getLsbD_cat4w _ _ _ _ (by decide : 32 - 1 < 128)] + simp + · simp only [show ¬ p < 32 by omega, show p < 64 from h.2, ite_true, ite_false, BitVec.getLsbD_or, + getLsbD_shl1 _ (show p - 32 < 32 by omega), getLsbD_shr31, show ¬ p - 32 = 0 by omega, + show 1 ≤ p - 32 by omega, show 1 ≤ p by omega, decide_true, decide_false, Bool.true_and, + Bool.false_and, Bool.or_false, Bool.xor_false, getLsbD_cat4w _ _ _ _ (show p - 1 < 128 by omega), + show ¬ p - 1 < 32 by omega, show p - 1 < 64 by omega] + congr 1 + · subst h + simp only [show ¬ 64 < 32 by decide, show ¬ 64 < 64 by decide, show 64 < 96 by decide, ite_false, + ite_true, Nat.sub_self, BitVec.getLsbD_or, getLsbD_shl1 _ (by decide : 0 < 32), getLsbD_shr31, + getLsbD_cat4w _ _ _ _ (by decide : 64 - 1 < 128)] + simp + · simp only [show ¬ p < 32 by omega, show ¬ p < 64 by omega, show p < 96 from h.2, ite_true, ite_false, + BitVec.getLsbD_or, getLsbD_shl1 _ (show p - 64 < 32 by omega), getLsbD_shr31, + show ¬ p - 64 = 0 by omega, show 1 ≤ p - 64 by omega, show 1 ≤ p by omega, decide_true, decide_false, + Bool.true_and, Bool.false_and, Bool.or_false, Bool.xor_false, + getLsbD_cat4w _ _ _ _ (show p - 1 < 128 by omega), show ¬ p - 1 < 32 by omega, + show ¬ p - 1 < 64 by omega, show p - 1 < 96 by omega] + congr 1 + · subst h + simp only [show ¬ 96 < 32 by decide, show ¬ 96 < 64 by decide, show ¬ 96 < 96 by decide, ite_false, + Nat.sub_self, BitVec.getLsbD_or, getLsbD_shl1 _ (by decide : 0 < 32), getLsbD_shr31, + getLsbD_cat4w _ _ _ _ (by decide : 96 - 1 < 128)] + simp + · simp only [show ¬ p < 32 by omega, show ¬ p < 64 by omega, show ¬ p < 96 by omega, ite_false, + BitVec.getLsbD_or, getLsbD_shl1 _ (show p - 96 < 32 by omega), getLsbD_shr31, + show ¬ p - 96 = 0 by omega, show 1 ≤ p - 96 by omega, show 1 ≤ p by omega, decide_true, decide_false, + Bool.true_and, Bool.false_and, Bool.or_false, Bool.xor_false, + getLsbD_cat4w _ _ _ _ (show p - 1 < 128 by omega), show ¬ p - 1 < 32 by omega, + show ¬ p - 1 < 64 by omega, show ¬ p - 1 < 96 by omega] + congr 1 + +theorem byteRev32_byteRev32 (a : BitVec 32) : byteRev32 (byteRev32 a) = a := by + apply BitVec.eq_of_getLsbD_eq + intro i hi + rw [getLsbD_byteRev32 _ hi, getLsbD_byteRev32 _ (by omega)] + congr 1; omega + +/-- The block at `p` as a big-endian integer, from its four byte-reversed words. -/ +theorem ofBytes_rev4 (m : Mem) (p : Addr) : + Spec.Gcm.ofBytes (Spec.Aes.bytesAt m p 16) = + byteRev32 (m.readW p 32) ++ byteRev32 (m.readW (p + BitVec.ofNat 64 4) 32) ++ + byteRev32 (m.readW (p + BitVec.ofNat 64 8) 32) ++ byteRev32 (m.readW (p + BitVec.ofNat 64 12) 32) := by + have h := le4_rev4 (byteRev32 (m.readW p 32)) (byteRev32 (m.readW (p + BitVec.ofNat 64 4) 32)) + (byteRev32 (m.readW (p + BitVec.ofNat 64 8) 32)) (byteRev32 (m.readW (p + BitVec.ofNat 64 12) 32)) + rw [byteRev32_byteRev32, byteRev32_byteRev32, byteRev32_byteRev32, byteRev32_byteRev32] at h + rw [bytesAt_split4, ← le4_readW, ← le4_readW, ← le4_readW, ← le4_readW, h, ofBytes_toBytes] + +end VG.Proof.Cmac diff --git a/lean/VerifiedGarbage/Proof/Cmac/Mem32.lean b/lean/VerifiedGarbage/Proof/Cmac/Mem32.lean new file mode 100644 index 000000000..45d12b865 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/Cmac/Mem32.lean @@ -0,0 +1,117 @@ +import VerifiedGarbage.Proof.Cmac.Frame + +/-! +# CMAC: blocks in memory as 32-bit words + +Untrusted: everything here is checked by Lean. + +On the 32-bit targets a 16-byte block is loaded and stored as four +little-endian 32-bit words: `le4 w` is the bytes of the word `w`, and storing +`w₀ … w₃` at `p`, `p + 4`, `p + 8` and `p + 12` (`store4`) leaves +`le4 w₀ ++ … ++ le4 w₃` there. `xor4Mem` stores the XOR of two blocks a +word at a time (the block written may be one of those read), `zero4` zeroes +a block, and `chainMem4` forms a counter block `C = P ⊕ Q` and zeroes `P`. +-/ + +namespace VG.Proof.Cmac + +open VG + +/-- The bytes of a 32-bit word, least significant first. -/ +def le4 (w : BitVec 32) : List Byte := (List.range 4).map fun i => w.extractLsb' (8 * i) 8 + +theorem length_le4 (w : BitVec 32) : (le4 w).length = 4 := by simp [le4] + +theorem getD_le4 (w : BitVec 32) {k : Nat} (hk : k < 4) : (le4 w).getD k 0 = w.extractLsb' (8 * k) 8 := by + simp [le4, List.getD_eq_getElem?_getD, hk] + +theorem le4_readW (m : Mem) (a : Addr) : le4 (m.readW a 32) = Spec.Aes.bytesAt m a 4 := by + apply List.ext_getElem (by simp [le4, Spec.Aes.bytesAt]) + intro k h₁ h₂ + have hk : k < 4 := by simpa [le4] using h₁ + simp only [le4, Spec.Aes.bytesAt, List.getElem_map, List.getElem_range] + rw [← Mem.extractLsb'_read m a (n := 4) hk] + simp only [Mem.readW] + rfl + +theorem le4_xor (a b : BitVec 32) : le4 (a ^^^ b) = Spec.Cmac.xor (le4 a) (le4 b) := by + apply List.ext_getElem (by simp [le4, Spec.Cmac.xor]) + intro k h₁ h₂ + have hk : k < 4 := by simpa [le4] using h₁ + simp only [le4, Spec.Cmac.xor, List.getElem_map, List.getElem_range, List.getElem_zipWith] + ext j hj + simp + +theorem le4_zero : le4 0 = Spec.Cmac.zeros 4 := by decide + +theorem bytesAt_add (m : Mem) (p : Addr) (a b : Nat) : + Spec.Aes.bytesAt m p (a + b) = Spec.Aes.bytesAt m p a ++ Spec.Aes.bytesAt m (p + BitVec.ofNat 64 a) b := by + simp only [Spec.Aes.bytesAt] + rw [List.range_add, List.map_append, List.map_map] + congr 1 + apply List.map_congr_left + intro i _ + simp only [Function.comp, BitVec.add_assoc] + congr 1 + rw [BitVec.ofNat_add] + +/-- A block as its four words' bytes. -/ +theorem bytesAt_split4 (m : Mem) (p : Addr) : + Spec.Aes.bytesAt m p 16 = Spec.Aes.bytesAt m p 4 ++ Spec.Aes.bytesAt m (p + BitVec.ofNat 64 4) 4 ++ + Spec.Aes.bytesAt m (p + BitVec.ofNat 64 8) 4 ++ Spec.Aes.bytesAt m (p + BitVec.ofNat 64 12) 4 := by + rw [show (16 : Nat) = 4 + (4 + (4 + 4)) from rfl, bytesAt_add, bytesAt_add, bytesAt_add] + simp only [Offset.add_add, List.append_assoc] + +/-- The four words `w₀ … w₃` stored at `p`. -/ +def store4 (m : Mem) (p : Addr) (w₀ w₁ w₂ w₃ : BitVec 32) : Mem := + (((m.writeW p w₀).writeW (p + BitVec.ofNat 64 4) w₁).writeW (p + BitVec.ofNat 64 8) w₂).writeW + (p + BitVec.ofNat 64 12) w₃ + +theorem frame_store4 {m : Mem} (p : Addr) (w₀ w₁ w₂ w₃ : BitVec 32) : + Frame [⟨p, 16⟩] m (store4 m p w₀ w₁ w₂ w₃) := by + have c (d : Nat) (h : d + 4 ≤ 16) : (⟨p, 16⟩ : Region).Contains (p + BitVec.ofNat 64 d) 4 := + Offset.contains_base p h (by omega) + have c0 : (⟨p, 16⟩ : Region).Contains p 4 := by simpa using c 0 (by decide) + exact ((((Frame.refl _ _).writeW (List.mem_singleton_self _) _ c0).writeW (List.mem_singleton_self _) _ + (c 4 (by decide))).writeW (List.mem_singleton_self _) _ (c 8 (by decide))).writeW + (List.mem_singleton_self _) _ (c 12 (by decide)) + +theorem readW_store4_of_sep {m : Mem} {p a : Addr} (w₀ w₁ w₂ w₃ : BitVec 32) (h : (⟨p, 16⟩ : Region).Disjoint ⟨a, 4⟩) : + (store4 m p w₀ w₁ w₂ w₃).readW a 32 = m.readW a 32 := + (frame_store4 (m := m) p w₀ w₁ w₂ w₃).readW (r := ⟨a, 4⟩) (Region.contains_self _ _) + (fun r hr => by simp only [List.mem_singleton] at hr; subst hr; exact h.symm) (by decide) + +/-- The bytes of a block after storing its four words. -/ +theorem bytesAt_store4 (m : Mem) (p : Addr) (w₀ w₁ w₂ w₃ : BitVec 32) : + Spec.Aes.bytesAt (store4 m p w₀ w₁ w₂ w₃) p 16 = le4 w₀ ++ le4 w₁ ++ le4 w₂ ++ le4 w₃ := by + have sep (d e : Nat) (h : d + 4 ≤ e ∨ e + 4 ≤ d) (he : e + 4 ≤ 16) (hd : d + 4 ≤ 16) : + Mem.Sep (p + BitVec.ofNat 64 d) (32 / 8) (p + BitVec.ofNat 64 e) (32 / 8) := + Offset.sep p h (by omega) (by omega) + rw [bytesAt_split4, ← le4_readW, ← le4_readW, ← le4_readW, ← le4_readW, store4] + rw [Mem.readW_writeW_self32, Mem.readW_writeW_sep (sep 8 12 (by decide) (by decide) (by decide)) (by decide), + Mem.readW_writeW_self32, Mem.readW_writeW_sep (sep 4 12 (by decide) (by decide) (by decide)) (by decide), + Mem.readW_writeW_sep (sep 4 8 (by decide) (by decide) (by decide)) (by decide), Mem.readW_writeW_self32] + rw [Mem.readW_writeW_sep (by simpa using sep 0 12 (by decide) (by decide) (by decide)) (by decide), + Mem.readW_writeW_sep (by simpa using sep 0 8 (by decide) (by decide) (by decide)) (by decide), + Mem.readW_writeW_sep (by simpa using sep 0 4 (by decide) (by decide) (by decide)) (by decide), + Mem.readW_writeW_self32] + +theorem xor_append4 {a b c d a' b' c' d' : List Byte} (ha : a.length = a'.length) (hb : b.length = b'.length) + (hc : c.length = c'.length) : + Spec.Cmac.xor (a ++ b ++ c ++ d) (a' ++ b' ++ c' ++ d') = + Spec.Cmac.xor a a' ++ Spec.Cmac.xor b b' ++ Spec.Cmac.xor c c' ++ Spec.Cmac.xor d d' := by + simp only [List.append_assoc] + rw [xor_append ha, xor_append hb, xor_append hc] + +/-- The XOR of two blocks, a word at a time. -/ +theorem xor_words4 (m : Mem) (p q : Addr) : + le4 (m.readW p 32 ^^^ m.readW q 32) ++ + le4 (m.readW (p + BitVec.ofNat 64 4) 32 ^^^ m.readW (q + BitVec.ofNat 64 4) 32) ++ + le4 (m.readW (p + BitVec.ofNat 64 8) 32 ^^^ m.readW (q + BitVec.ofNat 64 8) 32) ++ + le4 (m.readW (p + BitVec.ofNat 64 12) 32 ^^^ m.readW (q + BitVec.ofNat 64 12) 32) = + Spec.Cmac.xor (Spec.Aes.bytesAt m p 16) (Spec.Aes.bytesAt m q 16) := by + rw [le4_xor, le4_xor, le4_xor, le4_xor, le4_readW, le4_readW, le4_readW, le4_readW, le4_readW, + le4_readW, le4_readW, le4_readW, bytesAt_split4 m p, bytesAt_split4 m q, + xor_append4 (by simp [Spec.Aes.bytesAt]) (by simp [Spec.Aes.bytesAt]) (by simp [Spec.Aes.bytesAt])] + +end VG.Proof.Cmac diff --git a/lean/VerifiedGarbage/Proof/CmacAes/Arm/Call.lean b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Call.lean new file mode 100644 index 000000000..d2bdf9639 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Call.lean @@ -0,0 +1,281 @@ +import VerifiedGarbage.Proof.Aes.Arm.Ctr32 +import VerifiedGarbage.Proof.Cmac.Frame +import VerifiedGarbage.Proof.Framework.Arm.Frame +import VerifiedGarbage.Proof.Framework.Arm.RelCT +import VerifiedGarbage.Proof.Framework.Arm.RegUpd +import VerifiedGarbage.Impl.CmacAes.Arm + +/-! +# AES-CMAC on ARMv7: calling `vg_aes_ctr32` on one block + +Untrusted: everything here is checked by Lean. + +`ctr_call`: the frame that pushes `vg_aes_ctr32`'s stack arguments (`n = 1` +in `ra` and the working space `S` in `rb`) around its call, with the counter +block `C` and one data block `D` holding zeros: `D` then holds `CIPH_K(C)`, as +bytes (`Cmac.aesWith`), and only `C`, `D`, `S` and the 8 bytes below the stack +pointer change in memory. +-/ + +namespace VG.Proof.CmacAes.Arm + +open VG VG.Arm VG.Impl.CmacAes.Arm + +theorem ofBytes_zeros : Spec.Gcm.ofBytes (Spec.Cmac.zeros 16) = 0 := by decide + +theorem toNat_rounds {R : Nat} (hR : R = 10 ∨ R = 12 ∨ R = 14) : (BitVec.ofNat 32 R).toNat = R := by + rw [BitVec.toNat_ofNat]; exact Nat.mod_eq_of_lt (by omega) + +theorem storeWords_two (m : Mem) (a : BitVec 32) (x y : BitVec 32) : + storeWords m a [x, y] = (m.writeW (State.addr a) x).writeW (State.addr (a + 4)) y := rfl + +/-- Addresses below a pointer do not wrap. -/ +theorem addr_sub {a : BitVec 32} {k : Nat} (h : k ≤ a.toNat) : + State.addr (a - BitVec.ofNat 32 k) = State.addr a - BitVec.ofNat 64 k := by + simp only [State.addr] + apply BitVec.eq_of_toNat_eq + have := a.isLt + simp only [BitVec.toNat_setWidth, BitVec.toNat_sub, BitVec.toNat_ofNat] + rw [Nat.mod_eq_of_lt (a := k) (by omega), Nat.mod_eq_of_lt (a := k) (by omega), + Nat.mod_eq_of_lt (a := a.toNat) (by omega)] + omega + +/-- The 8 bytes below the stack pointer, where the frame pushes the stack arguments. -/ +abbrev below (s : State) : Region := ⟨State.addr s.sp - 8, 8⟩ + +/-- What a call of `vg_aes_ctr32` on one block needs. -/ +structure CallPre (s : State) (W C D S : BitVec 32) (R : Nat) (ra rb : Reg) : Prop where + r0 : s.gpr .r0 = W + r1 : s.gpr .r1 = BitVec.ofNat 32 R + r2 : s.gpr .r2 = C + r3 : s.gpr .r3 = D + hra : s.gpr ra = 1 + hrb : s.gpr rb = S + regs : regList [ra, rb] = true + rounds : R = 10 ∨ R = 12 ∨ R = 14 + hsp : 8 ≤ s.sp.toNat + wc : (⟨State.addr W, 240⟩ : Region).Disjoint ⟨State.addr C, 16⟩ + wd : (⟨State.addr W, 240⟩ : Region).Disjoint ⟨State.addr D, 16⟩ + ws : (⟨State.addr W, 240⟩ : Region).Disjoint ⟨State.addr S, 2048⟩ + cd : (⟨State.addr C, 16⟩ : Region).Disjoint ⟨State.addr D, 16⟩ + cs : (⟨State.addr C, 16⟩ : Region).Disjoint ⟨State.addr S, 2048⟩ + ds : (⟨State.addr D, 16⟩ : Region).Disjoint ⟨State.addr S, 2048⟩ + bw : (below s).Disjoint ⟨State.addr W, 240⟩ + bc : (below s).Disjoint ⟨State.addr C, 16⟩ + bd : (below s).Disjoint ⟨State.addr D, 16⟩ + bs : (below s).Disjoint ⟨State.addr S, 2048⟩ + hW : W.toNat + 240 ≤ 2 ^ 32 + hC : C.toNat + 16 ≤ 2 ^ 32 + hD : D.toNat + 16 ≤ 2 ^ 32 + hS : S.toNat + 2048 ≤ 2 ^ 32 + reads : Covers [⟨State.addr W, 240⟩] (s.rd ++ s.wr) + writes : Covers [⟨State.addr C, 16⟩, ⟨State.addr D, 16⟩, ⟨State.addr S, 2048⟩] s.wr + zero : Spec.Aes.bytesAt s.mem (State.addr D) 16 = Spec.Cmac.zeros 16 + +/-- What a call of `vg_aes_ctr32` on one block leaves. -/ +structure CallPost (s : State) (W C D S : BitVec 32) (R : Nat) (s' : State) : Prop where + rd : s'.rd = s.rd + wr : s'.wr = s.wr + sp : s'.sp = s.sp + saved : ∀ r ∈ preserved, r ≠ .lr → s'.gpr r = s.gpr r + frame : Frame [⟨State.addr C, 16⟩, ⟨State.addr D, 16⟩, ⟨State.addr S, 2048⟩, below s] s.mem s'.mem + out : Spec.Aes.bytesAt s'.mem (State.addr D) 16 = + Spec.Cmac.aesWith R (Spec.Aes.bytesAt s.mem (State.addr W) (16 * (R + 1))) + (Spec.Aes.bytesAt s.mem (State.addr C) 16) + +theorem e8 (ra rb : Reg) : BitVec.ofNat 32 (4 * [ra, rb].length) = 8 := rfl + +/-- The regions `vg_aes_ctr32` is called with. -/ +abbrev ctrRd (s : State) (W : BitVec 32) : List Region := [⟨State.addr W, 240⟩, below s] +abbrev ctrWr (C D S : BitVec 32) : List Region := + [⟨State.addr C, 16⟩, ⟨State.addr D, 16⟩, ⟨State.addr S, 2048⟩] + +/-- The state `vg_aes_ctr32` runs from, with the permissions it is given. -/ +abbrev ctrView (s : State) (ra rb : Reg) (W C D S : BitVec 32) : State := + (pushed [ra, rb] s).callEntry.withRegions (ctrRd s W) (ctrWr C D S) + +namespace CallPre +variable {s : State} {W C D S : BitVec 32} {R : Nat} {ra rb : Reg} (h : CallPre s W C D S R ra rb) +include h + +theorem hA : State.addr (s.sp - 8) = State.addr s.sp - 8 := addr_sub h.hsp + +theorem hspA : (s.sp - 8).toNat = s.sp.toNat - 8 := + BitVec.toNat_sub_of_le (by rw [BitVec.le_def]; exact h.hsp) + +theorem hA4 : State.addr (s.sp - 8 + BitVec.ofNat 32 4) = State.addr s.sp - 8 + 4 := by + have := s.sp.isLt + rw [addr_add (by rw [h.hspA]; omega), h.hA]; rfl + +theorem amem : (pushed [ra, rb] s).mem = + (s.mem.writeW (State.addr s.sp - 8) (1 : BitVec 32)).writeW (State.addr s.sp - 8 + 4) S := by + show storeWords s.mem (s.sp - BitVec.ofNat 32 (4 * [ra, rb].length)) [s.gpr ra, s.gpr rb] = _ + rw [e8, storeWords_two, h.hA, show (s.sp - 8 + 4 : BitVec 32) = s.sp - 8 + BitVec.ofNat 32 4 from rfl, h.hA4, + h.hra, h.hrb] + +theorem fA : Frame [below s] s.mem (pushed [ra, rb] s).mem := by + rw [h.amem] + refine ((Frame.refl _ _).writeW (List.mem_singleton_self _) _ ?_).writeW (List.mem_singleton_self _) _ ?_ + · simp only [Region.Contains, BitVec.sub_self, BitVec.toNat_zero]; omega + · simp only [Region.Contains] + rw [Offset.add_sub_cancel_left]; decide + +omit h in +theorem sp_view : (ctrView s ra rb W C D S).sp = s.sp - 8 := by + simp only [State.withRegions_sp, State.callEntry_sp, pushed_sp, e8] + +theorem arg0 : stackArg (ctrView s ra rb W C D S) 0 = 1 := by + have := h.hsp + rw [stackArg, show stackArgAddr (ctrView s ra rb W C D S) 0 = State.addr s.sp - 8 by + unfold stackArgAddr; rw [sp_view, show s.sp - 8 + BitVec.ofNat 32 (4 * 0) = s.sp - 8 from + BitVec.add_zero _, h.hA], + State.withRegions_mem, State.callEntry_mem, h.amem, Mem.readW_writeW_sep + (Offset.sep_base (State.addr s.sp - 8) (n := 4) (e := 4) (k := 4) (by decide) (by decide)) (by decide), + Mem.readW_writeW_self32] + +theorem arg1 : stackArg (ctrView s ra rb W C D S) 1 = S := by + rw [stackArg, show stackArgAddr (ctrView s ra rb W C D S) 1 = State.addr s.sp - 8 + 4 by + unfold stackArgAddr; rw [sp_view]; exact h.hA4, + State.withRegions_mem, State.callEntry_mem, h.amem, Mem.readW_writeW_self32] + +omit h in +theorem view_gpr (r : Reg) (hr : r ∉ linkRegs) : (ctrView s ra rb W C D S).gpr r = s.gpr r := by + rw [State.withRegions_gpr, State.callEntry_gpr _ hr, pushed_gpr] + +theorem pre : Proof.Aes.ctr32Arm.pre (ctrView s ra rb W C D S) := by + have hR := toNat_rounds h.rounds + have hsa : stackArgAddr (ctrView s ra rb W C D S) 0 = State.addr s.sp - 8 := by + unfold stackArgAddr; rw [sp_view, show s.sp - 8 + BitVec.ofNat 32 (4 * 0) = s.sp - 8 from + BitVec.add_zero _, h.hA] + simp only [Proof.Aes.ctr32Arm, h.arg0, h.arg1, hsa, view_gpr .r0 (by decide), view_gpr .r1 (by decide), + view_gpr .r2 (by decide), view_gpr .r3 (by decide), h.r0, h.r1, h.r2, h.r3, hR, State.withRegions_rd, + State.withRegions_wr, sp_view, h.hspA, show (1 : BitVec 32).toNat = 1 from rfl, Nat.mul_one] + refine ⟨trivial, trivial, h.wc, ?_, h.ws, ?_, h.cs, ?_, h.bc.symm, ?_, h.bs.symm, h.hW, h.hC, ?_, h.hS, ?_, + h.rounds⟩ + · simpa using h.wd + · simpa using h.cd + · simpa using h.ds + · simpa [below] using h.bd.symm + · simpa using h.hD + · have := s.sp.isLt; omega + +/-- The stack arguments are the frame. -/ +theorem cov : Covers (ctrRd s W ++ ctrWr C D S) ((pushed [ra, rb] s).rd ++ (pushed [ra, rb] s).wr) := by + intro x n' ⟨r, hr, hc⟩ + simp only [List.cons_append, List.nil_append, List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl | rfl + · obtain ⟨r', hr', hc'⟩ := h.reads x n' ⟨_, List.mem_singleton_self _, hc⟩ + refine ⟨r', ?_, hc'⟩ + rcases List.mem_append.mp hr' with h' | h' + · exact List.mem_append_left _ h' + · exact List.mem_append_right _ (by rw [pushed_wr]; exact List.mem_cons_of_mem _ h') + · refine ⟨_, List.mem_append_right _ (by rw [pushed_wr]; exact List.mem_cons_self ..), ?_⟩ + rwa [e8, h.hA] + all_goals + obtain ⟨r', hr', hc'⟩ := h.writes x n' ⟨_, by simp, hc⟩ + exact ⟨r', List.mem_append_right _ (by rw [pushed_wr]; exact List.mem_cons_of_mem _ hr'), hc'⟩ + +theorem covW : Covers (ctrWr C D S) (pushed [ra, rb] s).wr := by + intro x n' hi + obtain ⟨r', hr', hc'⟩ := h.writes x n' hi + exact ⟨r', by rw [pushed_wr]; exact List.mem_cons_of_mem _ hr', hc'⟩ + +end CallPre + +theorem ctr_noCalls : Impl.Aes.Arm.ctr32.noCalls = true := by decide +kernel + +theorem ctr_call {s : State} {W C D S : BitVec 32} {R : Nat} {ra rb : Reg} (h : CallPre s W C D S R ra rb) : + WP isa (ctrCall ra rb) s (CallPost s W C D S R) := by + refine WP.frame (rs := [ra, rb]) (r := ra) h.regs (by simpa using h.hsp) (by simp) ?_ + refine WP.call (k := Proof.Aes.ctr32Arm) Proof.Aes.Arm.ctr32_correct + (rd := ctrRd s W) (wr := ctrWr C D S) h.pre h.cov h.covW ?_ ctr_noCalls + intro s₂ hrd₂ hwr₂ hsp₂ hf hcs _ hpost + have hR := toNat_rounds h.rounds + have hR' : 16 * (R + 1) ≤ 240 := by rcases h.rounds with h' | h' | h' <;> omega + -- The memory of the frame. + have fBelow : Frame [⟨State.addr C, 16⟩, ⟨State.addr D, 16⟩, ⟨State.addr S, 2048⟩] + (pushed [ra, rb] s).mem s₂.mem := hf + have bytesW : Spec.Aes.bytesAt (pushed [ra, rb] s).mem (State.addr W) (16 * (R + 1)) = + Spec.Aes.bytesAt s.mem (State.addr W) (16 * (R + 1)) := + Proof.Cmac.bytesAt_frame h.fA (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact (h.bw.symm.sub_left (Region.sub_prefix hR')).symm.symm) (by omega) + have bytesC : Spec.Aes.bytesAt (pushed [ra, rb] s).mem (State.addr C) 16 = + Spec.Aes.bytesAt s.mem (State.addr C) 16 := + Proof.Cmac.bytesAt_frame16 h.fA (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact h.bc.symm) + have bytesD : Spec.Aes.bytesAt (pushed [ra, rb] s).mem (State.addr D) 16 = + Spec.Aes.bytesAt s.mem (State.addr D) 16 := + Proof.Cmac.bytesAt_frame16 h.fA (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact h.bd.symm) + obtain ⟨hdata, -⟩ := hpost + simp only [State.withRegions_mem, State.callEntry_mem, CallPre.view_gpr .r0 (by decide), + CallPre.view_gpr .r1 (by decide), CallPre.view_gpr .r2 (by decide), CallPre.view_gpr .r3 (by decide), h.r0, h.r1, h.r2, + h.r3, hR, h.arg0, show (1 : BitVec 32).toNat = 1 from rfl] at hdata + have one : ∀ m : Mem, Spec.Gcm.blocksAt m (State.addr D) 1 = [Spec.Gcm.blockAt m (State.addr D)] := + fun m => by simp [Spec.Gcm.blocksAt] + have bD : Spec.Gcm.blockAt (pushed [ra, rb] s).mem (State.addr D) = 0 := by + rw [Spec.Gcm.blockAt, bytesD, h.zero, ofBytes_zeros] + rw [one, one, bD, Proof.Cmac.ctr32_one, List.cons.injEq] at hdata + -- The register the pop loads is the one pushed. + have slot : s₂.mem.readW (State.addr (pushed [ra, rb] s).sp) 32 = 1 := by + rw [pushed_sp, e8, h.hA] + have := fBelow.readW (r := ⟨State.addr s.sp - 8, 4⟩) (a := State.addr s.sp - 8) (w := 32) + (Region.contains_self _ _) (fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl + · exact (h.bc.sub_left (Region.sub_prefix (by decide))).symm.symm + · exact (h.bd.sub_left (Region.sub_prefix (by decide))) + · exact (h.bs.sub_left (Region.sub_prefix (by decide)))) (by decide) + rw [this, h.amem, Mem.readW_writeW_sep + (Offset.sep_base (State.addr s.sp - 8) (n := 4) (e := 4) (k := 4) (by decide) (by decide)) (by decide), + Mem.readW_writeW_self32] + refine ⟨?_, ?_, ?_, fun r hr hlr => ?_, ?_, ?_⟩ + · rw [popped_rd, hrd₂, pushed_rd] + · rw [popped_wr, hwr₂, pushed_wr]; rfl + · rw [popped_sp, hsp₂, pushed_sp, e8]; exact BitVec.sub_add_cancel _ _ + · by_cases hra : r = ra + · subst hra + show (s₂.setReg r (s₂.mem.readW (State.addr s₂.sp) 32)).gpr r = _ + rw [VG.Arm.RegUpd.gpr_setReg_self, hsp₂, slot, h.hra] + · rw [popped_gpr hra, hcs r hr hlr, pushed_gpr] + · rw [popped_mem] + refine (h.fA.sub fun r hr => ⟨r, by simp at hr; simp [hr], fun _ h => h⟩).trans + (fBelow.sub fun r hr => ⟨r, by simp at hr; rcases hr with rfl | rfl | rfl <;> simp, fun _ h => h⟩) + · rw [popped_mem, Proof.Cmac.bytesAt_blockAt, hdata.1, Spec.Gcm.blockAt, bytesW, bytesC, + Proof.Cmac.aesWith_bytes _ _ (Proof.Cmac.bytesAt_length _ _ _)] + +/-- Calls of `vg_aes_ctr32` on one block, with the same arguments and stack +pointer in both runs, are constant time. -/ +theorem ctr_rel {W C D S sp₀ : BitVec 32} {R : Nat} {ra rb : Reg} {P : State → State → Prop} + (h : ∀ s₁ s₂, P s₁ s₂ → CallPre s₁ W C D S R ra rb ∧ CallPre s₂ W C D S R ra rb ∧ s₁.sp = sp₀ ∧ + s₂.sp = sp₀) : + RelCT isa P (ctrCall ra rb) fun _ _ => True := by + refine RelCT.frame (fun s₁ s₂ hp => by obtain ⟨_, _, h₁, h₂⟩ := h _ _ hp; rw [h₁, h₂]) ?_ + refine RelCT.call Proof.Aes.Arm.ctr32_correct Proof.Aes.Arm.ctr32_ct + [⟨State.addr W, 240⟩, ⟨State.addr sp₀ - 8, 8⟩] (ctrWr C D S) fun a b ⟨s₁, s₂, hp, pa, pb⟩ => ?_ + obtain ⟨h₁, h₂, sp₁, sp₂⟩ := h _ _ hp + rw [push_pushed h₁.regs (by simpa using h₁.hsp), Option.some.injEq] at pa + rw [push_pushed h₂.regs (by simpa using h₂.hsp), Option.some.injEq] at pb + subst pa pb + have e₁ : ctrRd s₁ W = [⟨State.addr W, 240⟩, ⟨State.addr sp₀ - 8, 8⟩] := by rw [ctrRd, below, sp₁] + have e₂ : ctrRd s₂ W = [⟨State.addr W, 240⟩, ⟨State.addr sp₀ - 8, 8⟩] := by rw [ctrRd, below, sp₂] + have p₁ := h₁.pre (C := C) (D := D) (S := S) + have p₂ := h₂.pre (C := C) (D := D) (S := S) + rw [ctrView, e₁] at p₁ + rw [ctrView, e₂] at p₂ + refine ⟨p₁, p₂, ?_, e₁ ▸ h₁.cov, h₁.covW, e₂ ▸ h₂.cov, h₂.covW⟩ + have a₁ := h₁.arg0 (C := C) (D := D) (S := S) + have b₁ := h₁.arg1 (C := C) (D := D) (S := S) + have a₂ := h₂.arg0 (C := C) (D := D) (S := S) + have b₂ := h₂.arg1 (C := C) (D := D) (S := S) + rw [ctrView, e₁] at a₁ b₁ + rw [ctrView, e₂] at a₂ b₂ + simp only [Proof.Aes.ctr32Arm, a₁, b₁, a₂, b₂, State.withRegions_gpr, State.withRegions_sp, + State.callEntry_sp, pushed_sp, State.callEntry_gpr _ (by decide : Reg.r0 ∉ linkRegs), + State.callEntry_gpr _ (by decide : Reg.r1 ∉ linkRegs), State.callEntry_gpr _ (by decide : Reg.r2 ∉ linkRegs), + State.callEntry_gpr _ (by decide : Reg.r3 ∉ linkRegs), pushed_gpr, h₁.r0, h₁.r1, h₁.r2, h₁.r3, h₂.r0, h₂.r1, + h₂.r2, h₂.r3, sp₁, sp₂] + exact ⟨trivial, trivial, trivial, trivial, trivial, trivial, trivial⟩ + +end VG.Proof.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/CmacAes/Arm/Contract.lean b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Contract.lean new file mode 100644 index 000000000..e1d8a13a1 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Contract.lean @@ -0,0 +1,97 @@ +import VerifiedGarbage.Proof.CmacAes.Arm.Call + +/-! +# AES-CMAC on ARMv7: the contracts the proofs are written against + +Untrusted: everything here is checked by Lean. The artifacts' contracts are +the shared ones of `Spec/Cmac/Contract.lean`, which imply these +(`Verified.lean`). Each function pushes `vg_aes_ctr32`'s two stack arguments +in the 8 bytes below the stack pointer, which may not overlap any buffer. +-/ + +namespace VG.Proof.CmacAes.Arm + +open VG VG.Arm + +/-- `CIPH_K` for AES with the key schedule at `w` for `R` rounds, in `m`. -/ +abbrev ciphAt (m : Mem) (w : Addr) (R : Nat) : Spec.Cmac.Cipher := + Spec.Cmac.aesWith R (Spec.Aes.bytesAt m w (16 * (R + 1))) + +/-- `vg_cmac_aes_update(schedule = r0, rounds = r1, state = r2, data = r3, n = [sp], scratch = [sp + 4])`. -/ +def updateArm : Contract isa where + pre s := + let sched : Region := ⟨State.addr (s.gpr .r0), 240⟩ + let state : Region := ⟨State.addr (s.gpr .r2), 16⟩ + let data : Region := ⟨State.addr (s.gpr .r3), 16 * (stackArg s 0).toNat⟩ + let scr : Region := ⟨State.addr (stackArg s 1), 2176⟩ + let args : Region := ⟨stackArgAddr s 0, 8⟩ + let below : Region := ⟨State.addr s.sp - BitVec.ofNat 64 8, 8⟩ + s.rd = [sched, data, args] ∧ s.wr = [state, scr] ∧ + sched.Disjoint state ∧ sched.Disjoint scr ∧ data.Disjoint state ∧ data.Disjoint scr ∧ + state.Disjoint scr ∧ state.Disjoint args ∧ scr.Disjoint args ∧ + below.Disjoint sched ∧ below.Disjoint data ∧ below.Disjoint state ∧ below.Disjoint scr ∧ + (s.gpr .r0).toNat + 240 ≤ 2 ^ 32 ∧ (s.gpr .r2).toNat + 16 ≤ 2 ^ 32 ∧ + (s.gpr .r3).toNat + 16 * (stackArg s 0).toNat ≤ 2 ^ 32 ∧ (stackArg s 1).toNat + 2176 ≤ 2 ^ 32 ∧ + 8 ≤ s.sp.toNat ∧ s.sp.toNat + 8 ≤ 2 ^ 32 ∧ + ((s.gpr .r1).toNat = 10 ∨ (s.gpr .r1).toNat = 12 ∨ (s.gpr .r1).toNat = 14) + post s s' := + Spec.Aes.bytesAt s'.mem (State.addr (s.gpr .r2)) 16 = + Spec.Cmac.chain (ciphAt s.mem (State.addr (s.gpr .r0)) (s.gpr .r1).toNat) + (Spec.Aes.bytesAt s.mem (State.addr (s.gpr .r2)) 16) + (Spec.Cmac.blocksAt s.mem (State.addr (s.gpr .r3)) 16 (stackArg s 0).toNat) + pub s₁ s₂ := + s₁.sp = s₂.sp ∧ s₁.gpr .r0 = s₂.gpr .r0 ∧ s₁.gpr .r1 = s₂.gpr .r1 ∧ s₁.gpr .r2 = s₂.gpr .r2 ∧ + s₁.gpr .r3 = s₂.gpr .r3 ∧ stackArg s₁ 0 = stackArg s₂ 0 ∧ stackArg s₁ 1 = stackArg s₂ 1 + +/-- `vg_cmac_aes_subkeys(schedule = r0, rounds = r1, subkeys = r2, scratch = r3)`. -/ +def subkeysArm : Contract isa where + pre s := + let sched : Region := ⟨State.addr (s.gpr .r0), 240⟩ + let subk : Region := ⟨State.addr (s.gpr .r2), 32⟩ + let scr : Region := ⟨State.addr (s.gpr .r3), 2176⟩ + let below : Region := ⟨State.addr s.sp - BitVec.ofNat 64 8, 8⟩ + s.rd = [sched] ∧ s.wr = [subk, scr] ∧ + sched.Disjoint subk ∧ sched.Disjoint scr ∧ subk.Disjoint scr ∧ + below.Disjoint sched ∧ below.Disjoint subk ∧ below.Disjoint scr ∧ + (s.gpr .r0).toNat + 240 ≤ 2 ^ 32 ∧ (s.gpr .r2).toNat + 32 ≤ 2 ^ 32 ∧ + (s.gpr .r3).toNat + 2176 ≤ 2 ^ 32 ∧ 8 ≤ s.sp.toNat ∧ + ((s.gpr .r1).toNat = 10 ∨ (s.gpr .r1).toNat = 12 ∨ (s.gpr .r1).toNat = 14) + post s s' := + let ks := Spec.Cmac.subkeys (ciphAt s.mem (State.addr (s.gpr .r0)) (s.gpr .r1).toNat) 16 + Spec.Aes.bytesAt s'.mem (State.addr (s.gpr .r2)) 32 = ks.1 ++ ks.2 + pub s₁ s₂ := + s₁.sp = s₂.sp ∧ s₁.gpr .r0 = s₂.gpr .r0 ∧ s₁.gpr .r1 = s₂.gpr .r1 ∧ s₁.gpr .r2 = s₂.gpr .r2 ∧ + s₁.gpr .r3 = s₂.gpr .r3 + +/-- `vg_cmac_aes_finalize(key = r0, rounds = r1, state = r2, last = r3, last_len = [sp], scratch = [sp + 4])`. -/ +def finalizeArm : Contract isa where + pre s := + let key : Region := ⟨State.addr (s.gpr .r0), 272⟩ + let state : Region := ⟨State.addr (s.gpr .r2), 16⟩ + let last : Region := ⟨State.addr (s.gpr .r3), (stackArg s 0).toNat⟩ + let scr : Region := ⟨State.addr (stackArg s 1), 2176⟩ + let args : Region := ⟨stackArgAddr s 0, 8⟩ + let below : Region := ⟨State.addr s.sp - BitVec.ofNat 64 8, 8⟩ + s.rd = [key, last, args] ∧ s.wr = [state, scr] ∧ + key.Disjoint state ∧ key.Disjoint scr ∧ last.Disjoint state ∧ last.Disjoint scr ∧ + state.Disjoint scr ∧ state.Disjoint args ∧ scr.Disjoint args ∧ + below.Disjoint key ∧ below.Disjoint last ∧ below.Disjoint state ∧ below.Disjoint scr ∧ + (s.gpr .r0).toNat + 272 ≤ 2 ^ 32 ∧ (s.gpr .r2).toNat + 16 ≤ 2 ^ 32 ∧ + (s.gpr .r3).toNat + (stackArg s 0).toNat ≤ 2 ^ 32 ∧ (stackArg s 1).toNat + 2176 ≤ 2 ^ 32 ∧ + 8 ≤ s.sp.toNat ∧ s.sp.toNat + 8 ≤ 2 ^ 32 ∧ + ((s.gpr .r1).toNat = 10 ∨ (s.gpr .r1).toNat = 12 ∨ (s.gpr .r1).toNat = 14) ∧ + (stackArg s 0).toNat ≤ 16 + post s s' := + let ciph := ciphAt s.mem (State.addr (s.gpr .r0)) (s.gpr .r1).toNat + let ks := Spec.Cmac.subkeys ciph 16 + Spec.Aes.bytesAt s.mem (State.addr (s.gpr .r0) + 240) 32 = ks.1 ++ ks.2 → + ∀ msg : List Byte, msg.length % 16 = 0 → (msg = [] ∨ 0 < (stackArg s 0).toNat) → + Spec.Aes.bytesAt s.mem (State.addr (s.gpr .r2)) 16 = + Spec.Cmac.chain ciph (Spec.Cmac.zeros 16) (Spec.Cmac.blocks 16 msg) → + Spec.Aes.bytesAt s'.mem (State.addr (s.gpr .r2)) 16 = + Spec.Cmac.macFull ciph 16 (msg ++ Spec.Aes.bytesAt s.mem (State.addr (s.gpr .r3)) (stackArg s 0).toNat) + pub s₁ s₂ := + s₁.sp = s₂.sp ∧ s₁.gpr .r0 = s₂.gpr .r0 ∧ s₁.gpr .r1 = s₂.gpr .r1 ∧ s₁.gpr .r2 = s₂.gpr .r2 ∧ + s₁.gpr .r3 = s₂.gpr .r3 ∧ stackArg s₁ 0 = stackArg s₂ 0 ∧ stackArg s₁ 1 = stackArg s₂ 1 + +end VG.Proof.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/CmacAes/Arm/Dbl.lean b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Dbl.lean new file mode 100644 index 000000000..71fdbaa56 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Dbl.lean @@ -0,0 +1,166 @@ +import VerifiedGarbage.Proof.CmacAes.Arm.Words +import VerifiedGarbage.Proof.Cmac.Dbl32 +import VerifiedGarbage.Proof.Cmac.Dbl + +/-! +# AES-CMAC on ARMv7: doubling a block in four 32-bit words + +Untrusted: everything here is checked by Lean. `dbl src dst` loads a block +as four byte-reversed words (`rev`), the block as a big-endian integer +(`Cmac.ofBytes_rev4`), doubles the integer a word at a time +(`Cmac.dbl_words4`), and stores the words byte-reversed again +(`Cmac.le4_rev4`). +-/ + +namespace VG.Proof.CmacAes.Arm + +open VG VG.Arm VG.Impl.CmacAes.Arm +open VG.Proof.MdStream.Arm (Upd Mupd op2_imm op2_reg op2_lsr op2_lsl wp_mov wp_sub wp_and wp_orr wp_rev wp_ldr wp_str) + +theorem rev_eq (a : BitVec 32) : rev a = byteRev32 a := rfl + +/-- The memory after `dbl src dst`, with `r6` pointing at `A`. -/ +def dblMem (m : Mem) (A : Addr) (src dst : Nat) : Mem := + let P := A + BitVec.ofNat 64 src + let b₀ := byteRev32 (m.readW P 32) + let b₁ := byteRev32 (m.readW (P + BitVec.ofNat 64 4) 32) + let b₂ := byteRev32 (m.readW (P + BitVec.ofNat 64 8) 32) + let b₃ := byteRev32 (m.readW (P + BitVec.ofNat 64 12) 32) + Proof.Cmac.store4 m (A + BitVec.ofNat 64 dst) (byteRev32 (Proof.Cmac.dblW0 b₀ b₁)) + (byteRev32 (Proof.Cmac.dblW0 b₁ b₂)) (byteRev32 (Proof.Cmac.dblW0 b₂ b₃)) (byteRev32 (Proof.Cmac.dblW3 b₀ b₃)) + +theorem dblMem_frame (m : Mem) (A : Addr) (src dst : Nat) : + Frame [⟨A + BitVec.ofNat 64 dst, 16⟩] m (dblMem m A src dst) := + Proof.Cmac.frame_store4 _ _ _ _ _ + +theorem dblMem_bytes (m : Mem) (A : Addr) (src dst : Nat) : + Spec.Aes.bytesAt (dblMem m A src dst) (A + BitVec.ofNat 64 dst) 16 = + Spec.Cmac.dbl 16 (Spec.Aes.bytesAt m (A + BitVec.ofNat 64 src) 16) := by + simp only [dblMem] + rw [Proof.Cmac.bytesAt_store4, Proof.Cmac.le4_rev4, Proof.Cmac.dbl_words4, + Proof.Cmac.dbl_eq (Proof.Cmac.bytesAt_length _ _ _), Proof.Cmac.ofBytes_rev4] + +/-- `dbl src dst`, with `r6` pointing at `K`. -/ +theorem dbl_wp {is : List Instr} {s : State} {Q : State → Prop} {K : BitVec 32} {src dst : Nat} + (h6 : s.gpr .r6 = K) (hs : src + 12 < 4096) (hd : dst + 12 < 4096) + (fs : K.toNat + src + 16 ≤ 2 ^ 32) (fd : K.toNat + dst + 16 ≤ 2 ^ 32) + (rS : Covers [⟨State.addr K + BitVec.ofNat 64 src, 16⟩] (s.rd ++ s.wr)) + (wD : Covers [⟨State.addr K + BitVec.ofNat 64 dst, 16⟩] s.wr) + (k : ∀ s', (∀ r, r ≠ .r0 → r ≠ .r1 → r ≠ .r2 → r ≠ .r3 → r ≠ .r4 → r ≠ .r12 → s'.gpr r = s.gpr r) → + s'.mem = dblMem s.mem (State.addr K) src dst → s'.rd = s.rd → s'.wr = s.wr → s'.sp = s.sp → + WP isa (.block is) s' Q) : + WP isa (.block (dbl src dst ++ is)) s Q := by + simp only [dbl, List.cons_append, List.nil_append] + refine wp_ldr (a := State.addr K + BitVec.ofNat 64 src) (by omega) (by rw [h6]; exact addr_add (by omega)) + (in_word0 rS) fun s₁ u₁ => ?_ + refine wp_ldr (a := State.addr K + BitVec.ofNat 64 src + BitVec.ofNat 64 4) (by omega) + (by rw [u₁.other _ (by decide), h6]; exact addr_word 4 fs (by decide)) + (by rw [u₁.rd, u₁.wr]; exact in_word rS (by decide)) fun s₂ u₂ => ?_ + refine wp_ldr (a := State.addr K + BitVec.ofNat 64 src + BitVec.ofNat 64 8) (by omega) + (by rw [u₂.other _ (by decide), u₁.other _ (by decide), h6]; exact addr_word 8 fs (by decide)) + (by rw [u₂.rd, u₂.wr, u₁.rd, u₁.wr]; exact in_word rS (by decide)) fun s₃ u₃ => ?_ + refine wp_ldr (a := State.addr K + BitVec.ofNat 64 src + BitVec.ofNat 64 12) (by omega) + (by rw [u₃.other _ (by decide), u₂.other _ (by decide), u₁.other _ (by decide), h6] + exact addr_word 12 fs (by decide)) + (by rw [u₃.rd, u₃.wr, u₂.rd, u₂.wr, u₁.rd, u₁.wr]; exact in_word rS (by decide)) fun s₄ u₄ => ?_ + refine wp_rev fun s₅ u₅ => wp_rev fun s₆ u₆ => wp_rev fun s₇ u₇ => wp_rev fun s₈ u₈ => ?_ + refine wp_mov (op2_lsr (by decide)) fun s₉ u₉ => wp_mov (op2_imm (by decide)) fun s₁₀ u₁₀ => + wp_sub (op2_reg _ _) fun s₁₁ u₁₁ => wp_and (op2_imm (by decide)) fun s₁₂ u₁₂ => ?_ + refine wp_mov (op2_lsl (by decide)) fun s₁₃ u₁₃ => wp_orr (op2_lsr (by decide)) fun s₁₄ u₁₄ => + wp_mov (op2_lsl (by decide)) fun s₁₅ u₁₅ => wp_orr (op2_lsr (by decide)) fun s₁₆ u₁₆ => + wp_mov (op2_lsl (by decide)) fun s₁₇ u₁₇ => wp_orr (op2_lsr (by decide)) fun s₁₈ u₁₈ => + wp_mov (op2_lsl (by decide)) fun s₁₉ u₁₉ => wp_eor (op2_reg _ _) fun s₂₀ u₂₀ => ?_ + refine wp_rev fun s₂₁ u₂₁ => wp_rev fun s₂₂ u₂₂ => wp_rev fun s₂₃ u₂₃ => wp_rev fun s₂₄ u₂₄ => ?_ + have g : ∀ r, r ≠ .r0 → r ≠ .r1 → r ≠ .r2 → r ≠ .r3 → r ≠ .r4 → r ≠ .r12 → s₂₄.gpr r = s.gpr r := + fun r h0 h1 h2 h3 h4 h12 => by + rw [u₂₄.other _ h3, u₂₃.other _ h2, u₂₂.other _ h1, u₂₁.other _ h0, u₂₀.other _ h3, u₁₉.other _ h3, + u₁₈.other _ h2, u₁₇.other _ h2, u₁₆.other _ h1, u₁₅.other _ h1, u₁₄.other _ h0, u₁₃.other _ h0, + u₁₂.other _ h12, u₁₁.other _ h12, u₁₀.other _ h4, u₉.other _ h12, u₈.other _ h3, u₇.other _ h2, + u₆.other _ h1, u₅.other _ h0, u₄.other _ h3, u₃.other _ h2, u₂.other _ h1, u₁.other _ h0] + have g6 : s₂₄.gpr .r6 = K := by + rw [g _ (by decide) (by decide) (by decide) (by decide) (by decide) (by decide), h6] + have m24 : s₂₄.mem = s.mem := by + rw [u₂₄.mem, u₂₃.mem, u₂₂.mem, u₂₁.mem, u₂₀.mem, u₁₉.mem, u₁₈.mem, u₁₇.mem, u₁₆.mem, u₁₅.mem, u₁₄.mem, + u₁₃.mem, u₁₂.mem, u₁₁.mem, u₁₀.mem, u₉.mem, u₈.mem, u₇.mem, u₆.mem, u₅.mem, u₄.mem, u₃.mem, u₂.mem, + u₁.mem] + have rd24 : s₂₄.rd = s.rd := by + rw [u₂₄.rd, u₂₃.rd, u₂₂.rd, u₂₁.rd, u₂₀.rd, u₁₉.rd, u₁₈.rd, u₁₇.rd, u₁₆.rd, u₁₅.rd, u₁₄.rd, u₁₃.rd, + u₁₂.rd, u₁₁.rd, u₁₀.rd, u₉.rd, u₈.rd, u₇.rd, u₆.rd, u₅.rd, u₄.rd, u₃.rd, u₂.rd, u₁.rd] + have wr24 : s₂₄.wr = s.wr := by + rw [u₂₄.wr, u₂₃.wr, u₂₂.wr, u₂₁.wr, u₂₀.wr, u₁₉.wr, u₁₈.wr, u₁₇.wr, u₁₆.wr, u₁₅.wr, u₁₄.wr, u₁₃.wr, + u₁₂.wr, u₁₁.wr, u₁₀.wr, u₉.wr, u₈.wr, u₇.wr, u₆.wr, u₅.wr, u₄.wr, u₃.wr, u₂.wr, u₁.wr] + have sp24 : s₂₄.sp = s.sp := by + rw [u₂₄.sp, u₂₃.sp, u₂₂.sp, u₂₁.sp, u₂₀.sp, u₁₉.sp, u₁₈.sp, u₁₇.sp, u₁₆.sp, u₁₅.sp, u₁₄.sp, u₁₃.sp, + u₁₂.sp, u₁₁.sp, u₁₀.sp, u₉.sp, u₈.sp, u₇.sp, u₆.sp, u₅.sp, u₄.sp, u₃.sp, u₂.sp, u₁.sp] + -- The four words. + have w₀ : s₁.gpr .r0 = s.mem.readW (State.addr K + BitVec.ofNat 64 src) 32 := u₁.gpr + have w₁ : s₂.gpr .r1 = s.mem.readW (State.addr K + BitVec.ofNat 64 src + BitVec.ofNat 64 4) 32 := by + rw [u₂.gpr, u₁.mem] + have w₂ : s₃.gpr .r2 = s.mem.readW (State.addr K + BitVec.ofNat 64 src + BitVec.ofNat 64 8) 32 := by + rw [u₃.gpr, u₂.mem, u₁.mem] + have w₃ : s₄.gpr .r3 = s.mem.readW (State.addr K + BitVec.ofNat 64 src + BitVec.ofNat 64 12) 32 := by + rw [u₄.gpr, u₃.mem, u₂.mem, u₁.mem] + have b₀ : s₈.gpr .r0 = byteRev32 (s.mem.readW (State.addr K + BitVec.ofNat 64 src) 32) := by + rw [u₈.other _ (by decide), u₇.other _ (by decide), u₆.other _ (by decide), u₅.gpr, u₄.other _ (by decide), + u₃.other _ (by decide), u₂.other _ (by decide), w₀, rev_eq] + have b₁ : s₈.gpr .r1 = byteRev32 (s.mem.readW (State.addr K + BitVec.ofNat 64 src + BitVec.ofNat 64 4) 32) := by + rw [u₈.other _ (by decide), u₇.other _ (by decide), u₆.gpr, u₅.other _ (by decide), u₄.other _ (by decide), + u₃.other _ (by decide), w₁, rev_eq] + have b₂ : s₈.gpr .r2 = byteRev32 (s.mem.readW (State.addr K + BitVec.ofNat 64 src + BitVec.ofNat 64 8) 32) := by + rw [u₈.other _ (by decide), u₇.gpr, u₆.other _ (by decide), u₅.other _ (by decide), u₄.other _ (by decide), + w₂, rev_eq] + have b₃ : s₈.gpr .r3 = byteRev32 (s.mem.readW (State.addr K + BitVec.ofNat 64 src + BitVec.ofNat 64 12) 32) := by + rw [u₈.gpr, u₇.other _ (by decide), u₆.other _ (by decide), u₅.other _ (by decide), w₃, rev_eq] + have mask : s₁₂.gpr .r12 = ((0 : BitVec 32) - (s₈.gpr .r0 >>> 31)) &&& 0x87 := by + rw [u₁₂.gpr, u₁₁.gpr, u₁₀.gpr, u₁₀.other _ (by decide), u₉.gpr] + have v₀ : s₂₄.gpr .r0 = byteRev32 (Proof.Cmac.dblW0 (s₈.gpr .r0) (s₈.gpr .r1)) := by + rw [u₂₄.other _ (by decide), u₂₃.other _ (by decide), u₂₂.other _ (by decide), u₂₁.gpr, rev_eq, + u₂₀.other _ (by decide), u₁₉.other _ (by decide), u₁₈.other _ (by decide), u₁₇.other _ (by decide), + u₁₆.other _ (by decide), u₁₅.other _ (by decide), u₁₄.gpr, u₁₃.gpr, u₁₃.other _ (by decide), + u₁₂.other _ (by decide), u₁₂.other _ (by decide), u₁₁.other _ (by decide), u₁₁.other _ (by decide), + u₁₀.other _ (by decide), u₁₀.other _ (by decide), u₉.other _ (by decide), u₉.other _ (by decide)] + rfl + have v₁ : s₂₄.gpr .r1 = byteRev32 (Proof.Cmac.dblW0 (s₈.gpr .r1) (s₈.gpr .r2)) := by + rw [u₂₄.other _ (by decide), u₂₃.other _ (by decide), u₂₂.gpr, rev_eq, u₂₁.other _ (by decide), + u₂₀.other _ (by decide), u₁₉.other _ (by decide), u₁₈.other _ (by decide), u₁₇.other _ (by decide), + u₁₆.gpr, u₁₅.gpr, u₁₅.other _ (by decide), u₁₄.other _ (by decide), u₁₄.other _ (by decide), + u₁₃.other _ (by decide), u₁₃.other _ (by decide), u₁₂.other _ (by decide), u₁₂.other _ (by decide), + u₁₁.other _ (by decide), u₁₁.other _ (by decide), u₁₀.other _ (by decide), u₁₀.other _ (by decide), + u₉.other _ (by decide), u₉.other _ (by decide)] + rfl + have v₂ : s₂₄.gpr .r2 = byteRev32 (Proof.Cmac.dblW0 (s₈.gpr .r2) (s₈.gpr .r3)) := by + rw [u₂₄.other _ (by decide), u₂₃.gpr, rev_eq, u₂₂.other _ (by decide), u₂₁.other _ (by decide), + u₂₀.other _ (by decide), u₁₉.other _ (by decide), u₁₈.gpr, u₁₇.gpr, u₁₇.other _ (by decide), + u₁₆.other _ (by decide), u₁₆.other _ (by decide), u₁₅.other _ (by decide), u₁₅.other _ (by decide), + u₁₄.other _ (by decide), u₁₄.other _ (by decide), u₁₃.other _ (by decide), u₁₃.other _ (by decide), + u₁₂.other _ (by decide), u₁₂.other _ (by decide), u₁₁.other _ (by decide), u₁₁.other _ (by decide), + u₁₀.other _ (by decide), u₁₀.other _ (by decide), u₉.other _ (by decide), u₉.other _ (by decide)] + rfl + have v₃ : s₂₄.gpr .r3 = byteRev32 (Proof.Cmac.dblW3 (s₈.gpr .r0) (s₈.gpr .r3)) := by + rw [u₂₄.gpr, rev_eq, u₂₃.other _ (by decide), u₂₂.other _ (by decide), u₂₁.other _ (by decide), u₂₀.gpr, + u₁₉.gpr, u₁₉.other _ (by decide), u₁₈.other _ (by decide), u₁₈.other _ (by decide), u₁₇.other _ (by decide), + u₁₇.other _ (by decide), u₁₆.other _ (by decide), u₁₆.other _ (by decide), u₁₅.other _ (by decide), + u₁₅.other _ (by decide), u₁₄.other _ (by decide), u₁₄.other _ (by decide), u₁₃.other _ (by decide), + u₁₃.other _ (by decide), mask, u₁₂.other _ (by decide), u₁₁.other _ (by decide), u₁₀.other _ (by decide), + u₉.other _ (by decide)] + rfl + refine wp_str (a := State.addr K + BitVec.ofNat 64 dst) (by omega) (by rw [g6]; exact addr_add (by omega)) + (by rw [wr24]; exact in_word0 wD) fun s₂₅ v₂₅ => ?_ + refine wp_str (a := State.addr K + BitVec.ofNat 64 dst + BitVec.ofNat 64 4) (by omega) + (by rw [v₂₅.gpr, g6]; exact addr_word 4 fd (by decide)) + (by rw [v₂₅.wr, wr24]; exact in_word wD (by decide)) fun s₂₆ v₂₆ => ?_ + refine wp_str (a := State.addr K + BitVec.ofNat 64 dst + BitVec.ofNat 64 8) (by omega) + (by rw [v₂₆.gpr, v₂₅.gpr, g6]; exact addr_word 8 fd (by decide)) + (by rw [v₂₆.wr, v₂₅.wr, wr24]; exact in_word wD (by decide)) fun s₂₇ v₂₇ => ?_ + refine wp_str (a := State.addr K + BitVec.ofNat 64 dst + BitVec.ofNat 64 12) (by omega) + (by rw [v₂₇.gpr, v₂₆.gpr, v₂₅.gpr, g6]; exact addr_word 12 fd (by decide)) + (by rw [v₂₇.wr, v₂₆.wr, v₂₅.wr, wr24]; exact in_word wD (by decide)) fun s₂₈ v₂₈ => k s₂₈ ?_ ?_ ?_ ?_ ?_ + · intro r h0 h1 h2 h3 h4 h12 + rw [v₂₈.gpr, v₂₇.gpr, v₂₆.gpr, v₂₅.gpr, g r h0 h1 h2 h3 h4 h12] + · rw [v₂₈.mem, v₂₇.mem, v₂₆.mem, v₂₅.mem, v₂₇.gpr, v₂₆.gpr, v₂₅.gpr, m24, v₀, v₁, v₂, v₃, b₀, b₁, b₂, b₃] + rfl + · rw [v₂₈.rd, v₂₇.rd, v₂₆.rd, v₂₅.rd, rd24] + · rw [v₂₈.wr, v₂₇.wr, v₂₆.wr, v₂₅.wr, wr24] + · rw [v₂₈.sp, v₂₇.sp, v₂₆.sp, v₂₅.sp, sp24] + +end VG.Proof.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/CmacAes/Arm/Finalize.lean b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Finalize.lean new file mode 100644 index 000000000..9a4ccf374 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Finalize.lean @@ -0,0 +1,444 @@ +import VerifiedGarbage.Proof.CmacAes.Arm.Subkeys +import VerifiedGarbage.Proof.Framework.WriteBytes + +/-! +# AES-CMAC on ARMv7: `vg_cmac_aes_finalize`, the last block + +Untrusted: everything here is checked by Lean. The steps that form the last +block `Mₙ` (§6.2 step 4) in the counter block, before the chaining value is +XORed in: `Mₙ* ⊕ K1` for a complete block (`full_wp`), else `Mₙ*` copied a +byte at a time onto zeros (`copy_wp`), `0x80` after it, and the block XORed +with `K2` (`partial_wp`). +-/ + +namespace VG.Proof.CmacAes.Arm + +open VG VG.Arm VG.Impl.CmacAes.Arm +open VG.Proof.MdStream.Arm (Upd Mupd Fupd op2_imm op2_reg wp_mov wp_add wp_subs wp_cmp wp_ldr wp_ldrb wp_strb + wp_ldrSp saveMem saveList_ok saveMem_frame readW_writeW_save eval_eq eval_ne sub_beq ofNat_beq_zero sub_ofNat) +open VG.WriteBytes (writeBytes writeBytes_nil writeBytes_snoc writeBytes_frame) + +section +variable (s₀ : State) + +/-- The key: the schedule and the subkeys `K1` and `K2` after it. -/ +abbrev keyR : Region := ⟨State.addr (W s₀), 272⟩ +/-- The last bytes `Mₙ*`. -/ +abbrev lastR : Region := ⟨State.addr (Dp s₀), N s₀⟩ + +/-- The last block `Mₙ` (§6.2 step 4), from the key and the last bytes. -/ +abbrev mn : List Byte := + Spec.Cmac.lastBlock 16 (Spec.Aes.bytesAt s₀.mem (State.addr (W s₀) + BitVec.ofNat 64 240) 16) + (Spec.Aes.bytesAt s₀.mem (State.addr (W s₀) + BitVec.ofNat 64 256) 16) + (Spec.Aes.bytesAt s₀.mem (State.addr (Dp s₀)) (N s₀)) + +/-- The counter block. -/ +abbrev Ca : Addr := State.addr (S s₀) + BitVec.ofNat 64 2048 + +end + +/-- The precondition, by name. -/ +structure FPre (s₀ : State) : Prop where + rd : s₀.rd = [keyR s₀, lastR s₀, argsR s₀] + wr : s₀.wr = [stR s₀, scrR s₀] + key_st : (keyR s₀).Disjoint (stR s₀) + key_scr : (keyR s₀).Disjoint (scrR s₀) + last_st : (lastR s₀).Disjoint (stR s₀) + last_scr : (lastR s₀).Disjoint (scrR s₀) + st_scr : (stR s₀).Disjoint (scrR s₀) + st_args : (stR s₀).Disjoint (argsR s₀) + scr_args : (scrR s₀).Disjoint (argsR s₀) + b_key : (belowR s₀).Disjoint (keyR s₀) + b_last : (belowR s₀).Disjoint (lastR s₀) + b_st : (belowR s₀).Disjoint (stR s₀) + b_scr : (belowR s₀).Disjoint (scrR s₀) + key_fit : (W s₀).toNat + 272 ≤ 2 ^ 32 + st_fit : (St s₀).toNat + 16 ≤ 2 ^ 32 + last_fit : (Dp s₀).toNat + N s₀ ≤ 2 ^ 32 + scr_fit : (S s₀).toNat + 2176 ≤ 2 ^ 32 + sp8 : 8 ≤ s₀.sp.toNat + sp_fit : s₀.sp.toNat + 8 ≤ 2 ^ 32 + rounds : R s₀ = 10 ∨ R s₀ = 12 ∨ R s₀ = 14 + len : N s₀ ≤ 16 + +theorem FPre.of {s₀ : State} (h : finalizeArm.pre s₀) : FPre s₀ := + let ⟨a, b, c, d, e, f, g, h, i, j, k, l, m, n, o, p, q, r, t, u, v⟩ := h + ⟨a, b, c, d, e, f, g, h, i, j, k, l, m, n, o, p, q, r, t, u, v⟩ + +section +variable {s₀ : State} (hp : FPre s₀) +include hp + +theorem FPre.arg1 : stackArgAddr s₀ 1 = stackArgAddr s₀ 0 + BitVec.ofNat 64 4 := by + have := hp.sp_fit + simp only [stackArgAddr] + rw [addr_add (by omega), addr_add (by omega)] + simp + +theorem FPre.arg_in {k : Nat} (hk : k < 2) : InRegions (s₀.rd ++ s₀.wr) (stackArgAddr s₀ k) 4 := by + refine ⟨argsR s₀, by simp [hp.rd], ?_⟩ + rcases (by omega : k = 0 ∨ k = 1) with rfl | rfl + · simpa using Offset.contains_base (stackArgAddr s₀ 0) (d := 0) (n := 4) (k := 8) (by decide) (by decide) + · rw [hp.arg1]; exact Offset.contains_base _ (by decide) (by decide) + +theorem FPre.cS {d n : Nat} (h : d + n ≤ 2176) : + Covers [⟨State.addr (S s₀) + BitVec.ofNat 64 d, n⟩] s₀.wr := by + rw [hp.wr] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact ⟨scrR s₀, by simp, d, rfl, h⟩ + +theorem FPre.cKey {d n : Nat} (h : d + n ≤ 272) : + Covers [⟨State.addr (W s₀) + BitVec.ofNat 64 d, n⟩] (s₀.rd ++ s₀.wr) := by + rw [hp.rd] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact ⟨keyR s₀, by simp, d, rfl, h⟩ + +theorem FPre.cLast {d n : Nat} (h : d + n ≤ N s₀) : + Covers [⟨State.addr (Dp s₀) + BitVec.ofNat 64 d, n⟩] (s₀.rd ++ s₀.wr) := by + rw [hp.rd] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact ⟨lastR s₀, by simp, d, rfl, h⟩ + +theorem FPre.ca_key {d n : Nat} (h : d + n ≤ 272) : + (⟨Ca s₀, 16⟩ : Region).Disjoint ⟨State.addr (W s₀) + BitVec.ofNat 64 d, n⟩ := + (hp.key_scr.symm.sub_left (Offset.sub_base _ (by decide))).sub_right (Offset.sub_base _ h) + +theorem FPre.ca_last : (⟨Ca s₀, 16⟩ : Region).Disjoint (lastR s₀) := + hp.last_scr.symm.sub_left (Offset.sub_base _ (by decide)) + +end + +theorem in_of_cov {rs : List Region} {a : Addr} {n : Nat} (h : Covers [⟨a, n⟩] rs) : InRegions rs a n := + h _ _ ⟨_, List.mem_singleton_self _, Region.contains_self _ _⟩ + +theorem cov_mono {rs rs' : List Region} {r : Region} (h : Covers [r] rs) (e : rs = rs') : Covers [r] rs' := e ▸ h + +/-! ## Saving the registers -/ + +/-- The registers `finalize` saves, and where. -/ +def fsaved : List (Reg × Nat) := [(.r4, 2064), (.r5, 2068), (.lr, 2072)] + +/-- The memory after saving them. -/ +def fsMem (s₀ : State) : Mem := saveMem s₀.mem (State.addr (S s₀)) s₀.gpr fsaved + +theorem finSave_eq : finSave = .ldrSp .r12 4 :: (fsaved.map (fun p => Instr.str p.1 .r12 p.2) ++ + ([.mov .r5 (.reg .r12), .ldrSp .r4 0, .cmp .r4 (.imm 16)] : List Instr)) := rfl + +theorem fsMem_frame (s₀ : State) : Frame [scrR s₀] s₀.mem (fsMem s₀) := + saveMem_frame _ _ _ (by decide) fsaved (by decide) + +set_option simprocs false in +theorem fsMem_slot (s₀ : State) {r : Reg} {d : Nat} (h : (r, d) ∈ fsaved) : + (fsMem s₀).readW (State.addr (S s₀) + BitVec.ofNat 64 d) 32 = s₀.gpr r := by + simp only [fsaved, List.mem_cons, List.not_mem_nil, or_false, Prod.mk.injEq] at h + rcases h with ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ <;> + simp (disch := decide) only [fsMem, fsaved, saveMem, Mem.readW_writeW_self32, readW_writeW_save] + +/-- What `finSave` leaves. -/ +structure FS (s₀ s : State) : Prop where + keep : ∀ r, r ≠ .r4 → r ≠ .r5 → r ≠ .r12 → s.gpr r = s₀.gpr r + r4 : s.gpr .r4 = BitVec.ofNat 32 (N s₀) + r5 : s.gpr .r5 = S s₀ + z : s.z = decide (N s₀ = 16) + mem : s.mem = fsMem s₀ + sp : s.sp = s₀.sp + rd : s.rd = s₀.rd + wr : s.wr = s₀.wr + +theorem finSave_wp {s₀ : State} (hp : FPre s₀) : WP isa (.block finSave) s₀ (FS s₀) := by + have hsc := hp.scr_fit + rw [finSave_eq] + refine wp_ldrSp (a := stackArgAddr s₀ 1) (by decide) rfl (hp.arg_in (by decide)) fun s₁ u₁ => ?_ + have h12 : s₁.gpr .r12 = S s₀ := u₁.gpr + refine saveList_ok fsaved s₁ _ (fun p hp' => ?_) fun s₂ g₂ rd₂ wr₂ sp₂ m₂ => ?_ + · have hb : 2064 ≤ p.2 ∧ p.2 + 4 ≤ 2076 := by + simp only [fsaved, List.mem_cons, List.not_mem_nil, or_false] at hp' + rcases hp' with rfl | rfl | rfl <;> decide + rw [h12, u₁.wr] + exact ⟨by omega, by omega, in_of_cov (hp.cS (d := p.2) (n := 4) (by omega))⟩ + have hm₂ : s₂.mem = fsMem s₀ := by + rw [m₂, u₁.mem, h12, fsMem] + exact saveMem_congr _ _ _ fun p hp' => u₁.other _ (by + simp only [fsaved, List.mem_cons, List.not_mem_nil, or_false] at hp' + rcases hp' with rfl | rfl | rfl <;> decide) + have harg : s₂.mem.readW (stackArgAddr s₀ 0) 32 = stackArg s₀ 0 := by + rw [hm₂] + exact (fsMem_frame s₀).readW (Region.contains_self _ _) (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact hp.scr_args.symm.sub_left UPre.arg_sub) (by decide) + refine wp_mov (op2_reg _ _) fun s₃ u₃ => ?_ + refine wp_ldrSp (a := stackArgAddr s₀ 0) (by decide) (by rw [u₃.sp, sp₂, u₁.sp]; rfl) + (by rw [u₃.rd, u₃.wr, rd₂, wr₂, u₁.rd, u₁.wr]; exact hp.arg_in (by decide)) fun s₄ u₄ => ?_ + refine wp_cmp (op2_imm (by decide)) fun s₅ f₅ z₅ => WP.block_nil ?_ + have a0 : stackArg s₀ 0 = BitVec.ofNat 32 (N s₀) := by simp [N] + have r4 : s₅.gpr .r4 = BitVec.ofNat 32 (N s₀) := by rw [f₅.gpr, u₄.gpr, u₃.mem, harg, a0] + refine ⟨fun r h4 h5 h12' => ?_, r4, ?_, ?_, ?_, ?_, ?_, ?_⟩ + · rw [f₅.gpr, u₄.other _ h4, u₃.other _ h5, g₂, u₁.other _ h12'] + · rw [f₅.gpr, u₄.other _ (by decide), u₃.gpr, g₂, h12] + · rw [z₅, show s₄.gpr .r4 = s₅.gpr .r4 from (congrFun f₅.gpr _).symm, r4, + show (16 : BitVec 32) = BitVec.ofNat 32 16 from rfl, sub_beq (by have := hp.len; omega) (by decide)] + · rw [f₅.mem, u₄.mem, u₃.mem, hm₂] + · rw [f₅.sp, u₄.sp, u₃.sp, sp₂, u₁.sp] + · rw [f₅.rd, u₄.rd, u₃.rd, rd₂, u₁.rd] + · rw [f₅.wr, u₄.wr, u₃.wr, wr₂, u₁.wr] + +/-! ## The last block -/ + +/-- What the branch on the length leaves: `Mₙ` in the counter block. -/ +structure BPost (s₀ s : State) : Prop where + keep : ∀ r, r ≠ .r3 → r ≠ .r4 → r ≠ .r5 → r ≠ .r12 → r ≠ .lr → s.gpr r = s₀.gpr r + r5 : s.gpr .r5 = S s₀ + sp : s.sp = s₀.sp + rd : s.rd = s₀.rd + wr : s.wr = s₀.wr + frame : Frame [⟨Ca s₀, 16⟩] (fsMem s₀) s.mem + blk : Spec.Aes.bytesAt s.mem (Ca s₀) 16 = mn s₀ + +section +variable {s₀ : State} (hp : FPre s₀) +include hp + +theorem FPre.key_bytes {d : Nat} (h : d + 16 ≤ 272) : + Spec.Aes.bytesAt (fsMem s₀) (State.addr (W s₀) + BitVec.ofNat 64 d) 16 = + Spec.Aes.bytesAt s₀.mem (State.addr (W s₀) + BitVec.ofNat 64 d) 16 := + Proof.Cmac.bytesAt_frame16 (fsMem_frame s₀) fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact hp.key_scr.sub_left (Offset.sub_base _ h) + +theorem FPre.last_bytes : Spec.Aes.bytesAt (fsMem s₀) (State.addr (Dp s₀)) (N s₀) = + Spec.Aes.bytesAt s₀.mem (State.addr (Dp s₀)) (N s₀) := + Proof.Cmac.bytesAt_frame (fsMem_frame s₀) (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact hp.last_scr) (by have := hp.len; omega) + +end + +theorem full_eq : full = xorBlk .r12 .lr .r3 .r0 .r5 0 240 2048 ++ [] := rfl + +theorem full_wp {s₀ : State} (hp : FPre s₀) (hL : N s₀ = 16) {s : State} (h : FS s₀ s) : + WP isa (.block full) s (BPost s₀) := by + have sf := hp.scr_fit + have kf := hp.key_fit + have lf := hp.last_fit + have hrw : s.rd ++ s.wr = s₀.rd ++ s₀.wr := by rw [h.rd, h.wr] + have r3 : s.gpr .r3 = Dp s₀ := h.keep _ (by decide) (by decide) (by decide) + have r0 : s.gpr .r0 = W s₀ := h.keep _ (by decide) (by decide) (by decide) + rw [full_eq] + refine xorBlk_ok (by decide) (by decide) (by decide) (by decide) (by decide) (by decide) (by decide) + (by decide) (by decide) (by decide) (by rw [r3]; omega) (by rw [r0]; omega) (by rw [h.r5]; omega) + (by rw [r3, add0, hrw]; exact hp.cLast (d := 0) (n := 16) (by omega) |> fun c => by simpa using c) + (by rw [r0, hrw]; exact hp.cKey (by decide)) (by rw [h.r5, h.wr]; exact hp.cS (by decide)) + fun s' g' => WP.block_nil ?_ + refine ⟨fun r h3 h4 h5 h12 hlr => by rw [g'.gpr r h12 hlr, h.keep r h4 h5 h12], by rw [g'.gpr _ (by decide) + (by decide), h.r5], by rw [g'.sp, h.sp], by rw [g'.rd, h.rd], by rw [g'.wr, h.wr], ?_, ?_⟩ + · rw [g'.mem, h.mem, h.r5]; exact Proof.Cmac.xor4Mem_frame _ _ _ _ + · rw [g'.mem, h.mem, h.r5, r3, r0, add0, Proof.Cmac.xor4Mem_bytes _ + (Proof.Cmac.Sep4.of_disjoint (hp.ca_last.sub_right (Region.sub_prefix (by omega)))) + (Proof.Cmac.Sep4.of_disjoint (hp.ca_key (by decide))), hp.key_bytes (by decide)] + have lb := hp.last_bytes + rw [hL] at lb + rw [lb] + simp only [mn, Spec.Cmac.lastBlock, Proof.Cmac.bytesAt_length, hL, ite_true] + exact Proof.Cmac.xor_comm _ _ + +/-! ## Copying the last bytes -/ + +theorem byte_rt32 (b : BitVec 8) : (b.setWidth 32).setWidth 8 = b := by + apply BitVec.eq_of_toNat_eq + have := b.isLt + simp only [BitVec.toNat_setWidth] + omega + +theorem copy_wp {s : State} {p c : BitVec 32} {L : Nat} (hL₀ : 0 < L) (hL : L ≤ 16) + (h3 : s.gpr .r3 = p) (hlr : s.gpr .lr = c) (h4 : s.gpr .r4 = BitVec.ofNat 32 L) + (fp : p.toNat + L ≤ 2 ^ 32) (fc : c.toNat + 16 ≤ 2 ^ 32) + (hr : Covers [⟨State.addr p, L⟩] (s.rd ++ s.wr)) (hw : Covers [⟨State.addr c, 16⟩] s.wr) + (hd : (⟨State.addr p, L⟩ : Region).Disjoint ⟨State.addr c, 16⟩) : + WP isa copy s fun s' => + s'.mem = writeBytes s.mem (State.addr c) (Spec.Aes.bytesAt s.mem (State.addr p) L) ∧ + s'.gpr .lr = c + BitVec.ofNat 32 L ∧ + (∀ r, r ≠ .r3 → r ≠ .r4 → r ≠ .r12 → r ≠ .lr → s'.gpr r = s.gpr r) ∧ + s'.sp = s.sp ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine WP.loop (M := isa) (body := .block [.ldrb .r12 .r3 0, .strb .r12 .lr 0, .dp .add .r3 .r3 (.imm 1), + .dp .add .lr .lr (.imm 1), .subs .r4 .r4 (.imm 1)]) (c := .ne) + (fun (n : Nat) (t : State) => ∃ i, n = L - i ∧ i < L ∧ t.gpr .r3 = p + BitVec.ofNat 32 i ∧ + t.gpr .lr = c + BitVec.ofNat 32 i ∧ t.gpr .r4 = BitVec.ofNat 32 (L - i) ∧ + t.mem = writeBytes s.mem (State.addr c) (Spec.Aes.bytesAt s.mem (State.addr p) i) ∧ + (∀ r, r ≠ .r3 → r ≠ .r4 → r ≠ .r12 → r ≠ .lr → t.gpr r = s.gpr r) ∧ + t.sp = s.sp ∧ t.rd = s.rd ∧ t.wr = s.wr) ?_ (L - 0) _ + ⟨0, rfl, hL₀, by rw [h3]; exact (BitVec.add_zero p).symm, by rw [hlr]; exact (BitVec.add_zero c).symm, by rw [h4, Nat.sub_zero], + by simp [Spec.Aes.bytesAt, writeBytes_nil], fun _ _ _ _ _ => rfl, rfl, rfl, rfl⟩ + rintro n t ⟨i, rfl, hi, x3, xlr, x4, mem, g, sp, rd, wr⟩ + have aP : State.addr (p + BitVec.ofNat 32 i) = State.addr p + BitVec.ofNat 64 i := addr_add (by omega) + have aC : State.addr (c + BitVec.ofNat 32 i) = State.addr c + BitVec.ofNat 64 i := addr_add (by omega) + refine wp_ldrb (a := State.addr p + BitVec.ofNat 64 i) (by decide) (by rw [x3, BitVec.add_zero, aP]) + (by rw [rd, wr]; exact hr _ _ ⟨_, List.mem_singleton_self _, Offset.contains_base _ (by omega) (by omega)⟩) + fun t₁ u₁ => ?_ + refine wp_strb (a := State.addr c + BitVec.ofNat 64 i) (by decide) + (by rw [u₁.other _ (by decide), xlr, BitVec.add_zero, aC]) + (by rw [u₁.wr, wr]; exact hw _ _ ⟨_, List.mem_singleton_self _, Offset.contains_base _ (by omega) (by omega)⟩) + fun t₂ v₂ => ?_ + refine wp_add (op2_imm (by decide)) fun t₃ u₃ => wp_add (op2_imm (by decide)) fun t₄ u₄ => + wp_subs (op2_imm (by decide)) fun t₅ u₅ z₅ => WP.block_nil ?_ + have hlen : (Spec.Aes.bytesAt s.mem (State.addr p) i).length = i := Proof.Cmac.bytesAt_length _ _ _ + have hx : writeBytes s.mem (State.addr c) (Spec.Aes.bytesAt s.mem (State.addr p) i) (State.addr p + BitVec.ofNat 64 i) = + s.mem (State.addr p + BitVec.ofNat 64 i) := + (writeBytes_frame s.mem (State.addr c) _ (R := ⟨State.addr c, i⟩) (by rw [hlen]; exact Region.contains_self _ _)) _ + fun r hr hcon => by + simp only [List.mem_singleton] at hr; subst hr + exact hd _ (Offset.contains_base _ (by omega) (by omega)) (Region.sub_prefix (by omega) _ hcon) + have hmem : t₅.mem = writeBytes s.mem (State.addr c) (Spec.Aes.bytesAt s.mem (State.addr p) (i + 1)) := by + rw [u₅.mem, u₄.mem, u₃.mem, v₂.mem, u₁.gpr, u₁.mem, mem, byte_rt32, hx, Proof.Cmac.bytesAt_succ, + writeBytes_snoc s.mem _ _ _ (by rw [hlen]; omega), hlen] + have x4' : t₅.gpr .r4 = BitVec.ofNat 32 (L - (i + 1)) := by + rw [u₅.gpr, u₄.other _ (by decide), u₃.other _ (by decide), v₂.gpr, u₁.other _ (by decide), x4, + show (1 : BitVec 32) = BitVec.ofNat 32 1 from rfl, sub_ofNat (by omega)]; rfl + have ev : isa.eval .ne t₅ = some !decide (L - (i + 1) = 0) := by + show VG.Arm.eval .ne t₅ = _ + rw [eval_ne, z₅, u₄.other _ (by decide), u₃.other _ (by decide), v₂.gpr, u₁.other _ (by decide), x4, + show (1 : BitVec 32) = BitVec.ofNat 32 1 from rfl, sub_ofNat (by omega), Nat.sub_sub, + ofNat_beq_zero (by omega)] + have gg : ∀ r, r ≠ .r3 → r ≠ .r4 → r ≠ .r12 → r ≠ .lr → t₅.gpr r = s.gpr r := fun r h₃ h₄ h₁₂ hl => by + rw [u₅.other _ h₄, u₄.other _ hl, u₃.other _ h₃, v₂.gpr, u₁.other _ h₁₂, g r h₃ h₄ h₁₂ hl] + have xlr' : t₅.gpr .lr = c + BitVec.ofNat 32 (i + 1) := by + rw [u₅.other _ (by decide), u₄.gpr, u₃.other _ (by decide), v₂.gpr, u₁.other _ (by decide), xlr, + show (1 : BitVec 32) = BitVec.ofNat 32 1 from rfl, Offset.add_add] + have x3' : t₅.gpr .r3 = p + BitVec.ofNat 32 (i + 1) := by + rw [u₅.other _ (by decide), u₄.other _ (by decide), u₃.gpr, v₂.gpr, u₁.other _ (by decide), x3, + show (1 : BitVec 32) = BitVec.ofNat 32 1 from rfl, Offset.add_add] + have sp' : t₅.sp = s.sp := by rw [u₅.sp, u₄.sp, u₃.sp, v₂.sp, u₁.sp, sp] + have rd' : t₅.rd = s.rd := by rw [u₅.rd, u₄.rd, u₃.rd, v₂.rd, u₁.rd, rd] + have wr' : t₅.wr = s.wr := by rw [u₅.wr, u₄.wr, u₃.wr, v₂.wr, u₁.wr, wr] + by_cases he : i + 1 = L + · left + exact ⟨by rw [ev]; simp [he], by rw [hmem, he], by rw [xlr', he], gg, sp', rd', wr'⟩ + · right + exact ⟨by rw [ev]; simp; omega, L - (i + 1), by omega, i + 1, rfl, by omega, x3', xlr', x4', hmem, gg, + sp', rd', wr'⟩ + +/-! ## A partial last block -/ + +theorem zero_eq : zero = .mov .r12 (.imm 0) :: (zeroBlk .r12 .r5 2048 ++ + ([.dp .add .lr .r5 (.imm (BitVec.ofNat 32 2048)), .cmp .r4 (.imm 0)] : List Instr)) := rfl + +theorem padK2_eq : padK2 = .mov .r12 (.imm 0x80) :: .strb .r12 .lr 0 :: + (xorBlk .r12 .lr .r5 .r0 .r5 2048 256 2048 ++ []) := rfl + +theorem b80 : ((0x80 : BitVec 32).setWidth 8 : Byte) = 0x80 := by decide + +theorem partial_wp {s₀ : State} (hp : FPre s₀) (hL : N s₀ < 16) {s : State} (h : FS s₀ s) : + WP isa partialBlock s (BPost s₀) := by + have sf := hp.scr_fit + have kf := hp.key_fit + have lf := hp.last_fit + have hrw : s.rd ++ s.wr = s₀.rd ++ s₀.wr := by rw [h.rd, h.wr] + have cA : State.addr (S s₀ + BitVec.ofNat 32 2048) = Ca s₀ := addr_add (by omega) + -- Zero the counter block. + refine WP.seq ?_ + rw [zero_eq] + refine wp_mov (op2_imm (by decide)) fun s₁ u₁ => ?_ + refine Proof.CmacAes.Arm.zeroBlk_ok u₁.gpr (by decide) (by rw [u₁.other _ (by decide), h.r5]; omega) + (by rw [u₁.other _ (by decide), h.r5, u₁.wr, h.wr]; exact hp.cS (by decide)) fun s₂ G₂ m₂ rd₂ wr₂ sp₂ => ?_ + refine wp_add (op2_imm (by decide)) fun s₃ u₃ => wp_cmp (op2_imm (by decide)) fun s₄ f₄ z₄ => WP.block_nil ?_ + have r5₂ : s₂.gpr .r5 = S s₀ := by rw [G₂, u₁.other _ (by decide), h.r5] + have mem₄ : s₄.mem = Proof.Cmac.zero4 (fsMem s₀) (Ca s₀) := by + rw [f₄.mem, u₃.mem, m₂, u₁.other _ (by decide), h.r5, u₁.mem, h.mem] + have k₄ : ∀ r, r ≠ .r12 → r ≠ .lr → s₄.gpr r = s.gpr r := fun r h12 hlr => by + rw [f₄.gpr, u₃.other _ hlr, G₂, u₁.other _ h12] + have lr₄ : s₄.gpr .lr = S s₀ + BitVec.ofNat 32 2048 := by rw [f₄.gpr, u₃.gpr, r5₂] + have r4₄ : s₄.gpr .r4 = BitVec.ofNat 32 (N s₀) := by + rw [k₄ _ (by decide) (by decide), h.r4] + have sp₄ : s₄.sp = s.sp := by rw [f₄.sp, u₃.sp, sp₂, u₁.sp] + have rd₄ : s₄.rd = s.rd := by rw [f₄.rd, u₃.rd, rd₂, u₁.rd] + have wr₄ : s₄.wr = s.wr := by rw [f₄.wr, u₃.wr, wr₂, u₁.wr] + have ev : isa.eval .eq s₄ = some (decide (N s₀ = 0)) := by + show VG.Arm.eval .eq s₄ = _ + rw [eval_eq, z₄, show s₃.gpr .r4 = s₄.gpr .r4 from (congrFun f₄.gpr _).symm, r4₄, + MdStream.Arm.cmp0 (by omega)] + have fz : Frame [⟨Ca s₀, 16⟩] (fsMem s₀) (Proof.Cmac.zero4 (fsMem s₀) (Ca s₀)) := + Proof.Cmac.frame_store4 _ _ _ _ _ + have lastZ : Spec.Aes.bytesAt (Proof.Cmac.zero4 (fsMem s₀) (Ca s₀)) (State.addr (Dp s₀)) (N s₀) = + Spec.Aes.bytesAt s₀.mem (State.addr (Dp s₀)) (N s₀) := by + rw [Proof.Cmac.bytesAt_frame fz (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact hp.ca_last.symm) (by omega), hp.last_bytes] + have keyZ : Spec.Aes.bytesAt (Proof.Cmac.zero4 (fsMem s₀) (Ca s₀)) (State.addr (W s₀) + BitVec.ofNat 64 256) 16 = + Spec.Aes.bytesAt s₀.mem (State.addr (W s₀) + BitVec.ofNat 64 256) 16 := by + rw [Proof.Cmac.bytesAt_frame16 fz (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact (hp.ca_key (by decide)).symm), hp.key_bytes (by decide)] + -- Copy the last bytes. + refine WP.seq (WP.mono (Q := fun (s₅ : State) => + s₅.mem = writeBytes (Proof.Cmac.zero4 (fsMem s₀) (Ca s₀)) (Ca s₀) + (Spec.Aes.bytesAt s₀.mem (State.addr (Dp s₀)) (N s₀)) ∧ + s₅.gpr .lr = S s₀ + BitVec.ofNat 32 (2048 + N s₀) ∧ + (∀ r, r ≠ .r3 → r ≠ .r4 → r ≠ .r12 → r ≠ .lr → s₅.gpr r = s.gpr r) ∧ + s₅.sp = s.sp ∧ s₅.rd = s.rd ∧ s₅.wr = s.wr) ?_ fun s₅ h₅ => ?_) + · by_cases hL0 : N s₀ = 0 + · refine WP.ite true (by rw [ev]; simp [hL0]) (fun _ => WP.block_nil ?_) (fun h => by cases h) + refine ⟨by rw [mem₄, hL0]; simp [Spec.Aes.bytesAt, writeBytes_nil], by rw [lr₄, hL0], fun r h₃ h₄ h₁₂ hl => + k₄ r h₁₂ hl, sp₄, rd₄, wr₄⟩ + · refine WP.ite false (by rw [ev]; simp [hL0]) (fun h => by cases h) fun _ => ?_ + refine WP.mono (copy_wp (p := Dp s₀) (c := S s₀ + BitVec.ofNat 32 2048) (L := N s₀) (by omega) (by omega) + (by rw [k₄ _ (by decide) (by decide), h.keep _ (by decide) (by decide) (by decide)]) lr₄ r4₄ lf + (by rw [BitVec.toNat_add, BitVec.toNat_ofNat, Nat.mod_eq_of_lt (a := 2048) (by decide), + Nat.mod_eq_of_lt (by omega)]; omega) + (by rw [rd₄, wr₄, hrw]; simpa using hp.cLast (d := 0) (n := N s₀) (by omega)) + (by rw [cA, wr₄, h.wr]; exact hp.cS (by decide)) (by rw [cA]; exact hp.ca_last.symm)) ?_ + rintro s₅ ⟨m₅, lr₅, g₅, sp₅, rd₅, wr₅⟩ + refine ⟨by rw [m₅, mem₄, cA, lastZ], by rw [lr₅, Offset.add_add], fun r h₃ h₄ h₁₂ hl => by + rw [g₅ r h₃ h₄ h₁₂ hl, k₄ r h₁₂ hl], by rw [sp₅, sp₄], by rw [rd₅, rd₄], by rw [wr₅, wr₄]⟩ + · obtain ⟨m₅, lr₅, g₅, sp₅, rd₅, wr₅⟩ := h₅ + rw [padK2_eq] + have aL : State.addr (S s₀ + BitVec.ofNat 32 (2048 + N s₀)) = Ca s₀ + BitVec.ofNat 64 (N s₀) := by + rw [addr_add (by omega), Offset.add_add] + refine wp_mov (op2_imm (by decide)) fun s₆ u₆ => ?_ + refine wp_strb (a := Ca s₀ + BitVec.ofNat 64 (N s₀)) (by decide) + (by rw [u₆.other _ (by decide), lr₅, BitVec.add_zero, aL]) + (by + rw [u₆.wr, wr₅, h.wr] + show InRegions s₀.wr (State.addr (S s₀) + BitVec.ofNat 64 2048 + BitVec.ofNat 64 (N s₀)) 1 + rw [Offset.add_add] + exact in_of_cov (hp.cS (d := 2048 + N s₀) (n := 1) (by omega))) fun s₇ v₇ => ?_ + have g₇ : ∀ r, r ≠ .r3 → r ≠ .r4 → r ≠ .r12 → r ≠ .lr → s₇.gpr r = s.gpr r := fun r h₃ h₄ h₁₂ hl => by + rw [v₇.gpr, u₆.other _ h₁₂, g₅ r h₃ h₄ h₁₂ hl] + have r5₇ : s₇.gpr .r5 = S s₀ := by + rw [g₇ _ (by decide) (by decide) (by decide) (by decide), h.r5] + have r0₇ : s₇.gpr .r0 = W s₀ := by + rw [g₇ _ (by decide) (by decide) (by decide) (by decide), h.keep _ (by decide) (by decide) (by decide)] + have hlen : (Spec.Aes.bytesAt s₀.mem (State.addr (Dp s₀)) (N s₀)).length = N s₀ := + Proof.Cmac.bytesAt_length _ _ _ + have m₇ : s₇.mem = (writeBytes (Proof.Cmac.zero4 (fsMem s₀) (Ca s₀)) (Ca s₀) + (Spec.Aes.bytesAt s₀.mem (State.addr (Dp s₀)) (N s₀))).writeW (Ca s₀ + BitVec.ofNat 64 (N s₀)) + (0x80 : Byte) := by + rw [v₇.mem, u₆.gpr, u₆.mem, m₅, b80] + refine xorBlk_ok (by decide) (by decide) (by decide) (by decide) (by decide) (by decide) (by decide) + (by decide) (by decide) (by decide) (by rw [r5₇]; omega) (by rw [r0₇]; omega) (by rw [r5₇]; omega) + (by rw [r5₇, v₇.rd, v₇.wr, u₆.rd, u₆.wr, rd₅, wr₅, hrw] + exact fun a n hi => (hp.cS (d := 2048) (n := 16) (by decide)) a n hi |> + fun ⟨r, hr, hc⟩ => ⟨r, List.mem_append_right _ hr, hc⟩) + (by rw [r0₇, v₇.rd, v₇.wr, u₆.rd, u₆.wr, rd₅, wr₅, hrw]; exact hp.cKey (by decide)) + (by rw [r5₇, v₇.wr, u₆.wr, wr₅, h.wr]; exact hp.cS (by decide)) fun s₈ g₈ => WP.block_nil ?_ + have fW : Frame [⟨Ca s₀, 16⟩] (fsMem s₀) s₇.mem := by + rw [m₇] + refine (fz.trans (writeBytes_frame _ _ _ ?_)).trans + ((Frame.refl _ _).writeW (List.mem_singleton_self _) _ (Offset.contains_base _ (by omega) (by omega))) + rw [hlen]; simpa using Offset.contains_base (Ca s₀) (d := 0) (n := N s₀) (k := 16) (by omega) (by decide) + have pad : Spec.Aes.bytesAt s₇.mem (Ca s₀) 16 = + Spec.Aes.bytesAt s₀.mem (State.addr (Dp s₀)) (N s₀) ++ [0x80] ++ Spec.Cmac.zeros (16 - N s₀ - 1) := by + have := Proof.Cmac.padded_bytes (Proof.Cmac.zero4 (fsMem s₀) (Ca s₀)) (Ca s₀) + (Spec.Aes.bytesAt s₀.mem (State.addr (Dp s₀)) (N s₀)) (by rw [hlen]; exact hL) + (Proof.Cmac.zero4_bytes _ _) + rw [hlen] at this + rw [m₇]; exact this + have k2 : Spec.Aes.bytesAt s₇.mem (State.addr (W s₀) + BitVec.ofNat 64 256) 16 = + Spec.Aes.bytesAt s₀.mem (State.addr (W s₀) + BitVec.ofNat 64 256) 16 := by + rw [Proof.Cmac.bytesAt_frame16 fW (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact (hp.ca_key (by decide)).symm), hp.key_bytes (by decide)] + refine ⟨fun r h₃ h₄ h₅ h₁₂ hl => by rw [g₈.gpr r h₁₂ hl, g₇ r h₃ h₄ h₁₂ hl, h.keep r h₄ h₅ h₁₂], + by rw [g₈.gpr _ (by decide) (by decide), r5₇], by rw [g₈.sp, v₇.sp, u₆.sp, sp₅, h.sp], + by rw [g₈.rd, v₇.rd, u₆.rd, rd₅, h.rd], by rw [g₈.wr, v₇.wr, u₆.wr, wr₅, h.wr], ?_, ?_⟩ + · rw [g₈.mem, r5₇]; exact fW.trans (Proof.Cmac.xor4Mem_frame _ _ _ _) + · rw [g₈.mem, r5₇, r0₇, Proof.Cmac.xor4Mem_bytes _ (Proof.Cmac.Sep4.self _) + (Proof.Cmac.Sep4.of_disjoint (hp.ca_key (by decide))), pad, k2] + simp only [mn, Spec.Cmac.lastBlock, hlen, show N s₀ ≠ 16 by omega, ite_false] + exact Proof.Cmac.xor_comm _ _ + +end VG.Proof.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/CmacAes/Arm/FinalizeCT.lean b/lean/VerifiedGarbage/Proof/CmacAes/Arm/FinalizeCT.lean new file mode 100644 index 000000000..9735fdfed --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/Arm/FinalizeCT.lean @@ -0,0 +1,72 @@ +import VerifiedGarbage.Proof.CmacAes.Arm.FinalizeCorrect +import VerifiedGarbage.Proof.CmacAes.Arm.UpdateCT + +/-! +# AES-CMAC on ARMv7: `vg_cmac_aes_finalize` is constant time + +Untrusted: everything here is checked by Lean. The code before the call is +checked by the taint analysis, from the public arguments (its branches and +the copy loop depend only on `last_len`), the call of `vg_aes_ctr32`, in its +frame, is constant time by its own proof (`ctr_rel`), its arguments pinned +by the correctness proof (`FMid`), and the restore after it by the taint +analysis again, from `r5` (the scratch buffer). +-/ + +namespace VG.Proof.CmacAes.Arm + +open VG VG.Arm VG.Impl.CmacAes.Arm + +/-- The restore after the call. -/ +abbrev finEnd : List Instr := [.ldr .r4 .r5 2064, .ldr .lr .r5 2072, .ldr .r5 .r5 2068] + +theorem finalize_rel {s₀ s₀' : State} (h0 : finalizeArm.pre s₀) (h0' : finalizeArm.pre s₀') + (hq : finalizeArm.pub s₀ s₀') : + RelCT isa (fun a b => a = s₀ ∧ b = s₀') finalize fun _ _ => True := by + obtain ⟨q₀, q₁, q₂, q₃, q₄, q₅, q₆⟩ := hq + have hp := FPre.of h0 + have hp' := FPre.of h0' + have eW : W s₀' = W s₀ := q₁.symm + have eR : R s₀' = R s₀ := by rw [R, R, q₂] + have eSt : St s₀' = St s₀ := q₃.symm + have eS : S s₀' = S s₀ := q₆.symm + obtain ⟨_, hA⟩ : ∃ h, (taint.check (argTaint [.r0, .r1, .r2, .r3] 8) finPre h).isSome = true := + ⟨_, by taint_decide⟩ + obtain ⟨_, hB⟩ : ∃ h, (taint.check (Taint.ofRegs [.r5]) (.block finEnd) h).isSome = true := + ⟨_, by taint_decide⟩ + have wfA : ∀ {s : State}, FPre s → + s.sp.toNat + 8 ≤ 2 ^ 32 ∧ ∀ r ∈ s.wr, Region.Disjoint ⟨State.addr s.sp, 8⟩ r := fun {s} h => by + have e : (⟨State.addr s.sp, 8⟩ : Region) = argsR s := by simp [stackArgAddr] + refine ⟨h.sp_fit, ?_⟩ + rw [e, h.wr] + simp only [List.mem_cons, List.not_mem_nil, or_false] + rintro r (rfl | rfl) + · exact h.st_args.symm + · exact h.scr_args.symm + have a := rel_agree (F := fun s => s = s₀) (F' := fun s => s = s₀') (G := FMid s₀) (G' := FMid s₀') + (argTaint [.r0, .r1, .r2, .r3] 8) + (fun s s' e e' => by + subst e e' + refine agree_argTaint (fun r hr => ?_) q₀ (wfA hp) (wfA hp') + (argMem_of (j := 2) q₀ hp.sp_fit fun i hi => ?_) + · simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl <;> assumption + · rcases (by omega : i = 0 ∨ i = 1) with rfl | rfl + · exact q₅ + · exact q₆) ⟨_, hA⟩ + (fun s e => by rw [e]; exact finPre_wp hp) (fun s e => by rw [e]; exact finPre_wp hp') + have c := rel_wp (F := FMid s₀) (F' := FMid s₀') (G := fun s => s.gpr .r5 = S s₀) + (G' := fun s => s.gpr .r5 = S s₀') + (ctr_rel (sp₀ := s₀.sp) fun s₁ s₂ h => + ⟨h.1.pre, by have := h.2.pre; rwa [eW, eS, eSt, eR] at this, h.1.sp, h.2.sp.trans q₀.symm⟩) + (fun _ h => WP.mono (ctr_call h.pre) fun _ hc => by rw [hc.saved .r5 (by simp [preserved]) (by decide), h.r5]) + (fun _ h => WP.mono (ctr_call h.pre) fun _ hc => by rw [hc.saved .r5 (by simp [preserved]) (by decide), h.r5]) + have b := RelCT.taint (A := taint) (P := fun s₁ s₂ => s₁.gpr .r5 = S s₀ ∧ s₂.gpr .r5 = S s₀') + (Taint.ofRegs [.r5]) (fun s₁ s₂ h => Taint.agree_ofRegs fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + subst hr; rw [h.1, h.2, eS]) hB + exact a.seq (c.seq b) + +theorem finalize_ct : ConstantTime isa finalizeArm.pre finalizeArm.pub finalize := + fun _ _ _ _ _ _ h₁ h₂ hq e₁ e₂ => (finalize_rel h₁ h₂ hq _ _ _ _ _ _ ⟨rfl, rfl⟩ e₁ e₂).1 + +end VG.Proof.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/CmacAes/Arm/FinalizeCorrect.lean b/lean/VerifiedGarbage/Proof/CmacAes/Arm/FinalizeCorrect.lean new file mode 100644 index 000000000..2785d25db --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/Arm/FinalizeCorrect.lean @@ -0,0 +1,228 @@ +import VerifiedGarbage.Proof.CmacAes.Arm.Finalize + +/-! +# AES-CMAC on ARMv7: `vg_cmac_aes_finalize` is correct + +Untrusted: everything here is checked by Lean. Before the call, the +counter block holds `Mₙ ⊕ C`, for the last block `Mₙ` of §6.2 step 4 and the +chaining value `C` at `state`, and the state is zeroed; the call leaves +`CIPH_K(C ⊕ Mₙ)` there, the MAC (`Cmac.macFull_split`). +-/ + +namespace VG.Proof.CmacAes.Arm + +open VG VG.Arm VG.Impl.CmacAes.Arm +open VG.Proof.MdStream.Arm (Upd Mupd op2_imm op2_reg wp_mov wp_add wp_ldr eval_eq) + +/-! ## Up to the call -/ + +/-- What the code before the call leaves. -/ +structure FMid (s₀ s : State) : Prop where + pre : CallPre s (W s₀) (S s₀ + BitVec.ofNat 32 2048) (St s₀) (S s₀) (R s₀) .r4 .r5 + blk : Spec.Aes.bytesAt s.mem (Ca s₀) 16 = + Spec.Cmac.xor (mn s₀) (Spec.Aes.bytesAt s₀.mem (State.addr (St s₀)) 16) + frame : Frame [⟨Ca s₀, 16⟩, stR s₀] (fsMem s₀) s.mem + r5 : s.gpr .r5 = S s₀ + keep : ∀ r ∈ preserved, r ≠ .r4 → r ≠ .r5 → r ≠ .lr → s.gpr r = s₀.gpr r + sp : s.sp = s₀.sp + rd : s.rd = s₀.rd + wr : s.wr = s₀.wr + +theorem finArgs_eq : finArgs = xorBlk .r12 .lr .r5 .r2 .r5 2048 0 2048 ++ + (.mov .r12 (.imm 0) :: (zeroBlk .r12 .r2 0 ++ + ([.mov .r3 (.reg .r2), .dp .add .r2 .r5 (.imm (BitVec.ofNat 32 2048)), .mov .r4 (.imm 1)] : + List Instr))) := rfl + +theorem preserved_ne {r : Reg} (hr : r ∈ preserved) : r ≠ .r0 ∧ r ≠ .r1 ∧ r ≠ .r2 ∧ r ≠ .r3 ∧ r ≠ .r12 := by + simp only [preserved, List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl <;> decide + +theorem finArgs_wp {s₀ : State} (hp : FPre s₀) {s : State} (h : BPost s₀ s) : + WP isa (.block finArgs) s (FMid s₀) := by + have sf := hp.scr_fit + have tf := hp.st_fit + have hR := hp.rounds + have hrw : s.rd ++ s.wr = s₀.rd ++ s₀.wr := by rw [h.rd, h.wr] + have r2 : s.gpr .r2 = St s₀ := h.keep _ (by decide) (by decide) (by decide) (by decide) (by decide) + have cA : State.addr (S s₀ + BitVec.ofNat 32 2048) = Ca s₀ := addr_add (by omega) + have cSt : (⟨Ca s₀, 16⟩ : Region).Disjoint (stR s₀) := hp.st_scr.symm.sub_left (Offset.sub_base _ (by decide)) + have wSt : Covers [stR s₀] s₀.wr := by + rw [hp.wr] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact ⟨stR s₀, by simp, 0, by simp, by simp⟩ + rw [finArgs_eq] + refine xorBlk_ok (by decide) (by decide) (by decide) (by decide) (by decide) (by decide) (by decide) + (by decide) (by decide) (by decide) (by rw [h.r5]; omega) (by rw [r2]; omega) (by rw [h.r5]; omega) + (by + rw [h.r5, hrw] + exact fun a n hi => (hp.cS (d := 2048) (n := 16) (by decide)) a n hi |> + fun ⟨r, hr, hc⟩ => ⟨r, List.mem_append_right _ hr, hc⟩) + (by + rw [r2, add0, hrw] + exact fun a n hi => wSt a n hi |> fun ⟨r, hr, hc⟩ => ⟨r, List.mem_append_right _ hr, hc⟩) + (by rw [h.r5, h.wr]; exact hp.cS (by decide)) fun s₁ g₁ => ?_ + refine wp_mov (op2_imm (by decide)) fun s₂ u₂ => ?_ + have r2₂ : s₂.gpr .r2 = St s₀ := by rw [u₂.other _ (by decide), g₁.gpr _ (by decide) (by decide), r2] + refine Proof.CmacAes.Arm.zeroBlk_ok u₂.gpr (by decide) (by rw [r2₂]; omega) + (by rw [r2₂, add0, u₂.wr, g₁.wr, h.wr]; exact wSt) fun s₃ G₃ m₃ rd₃ wr₃ sp₃ => ?_ + refine wp_mov (op2_reg _ _) fun s₄ u₄ => wp_add (op2_imm (by decide)) fun s₅ u₅ => + wp_mov (op2_imm (by decide)) fun s₆ u₆ => WP.block_nil ?_ + have keep : ∀ r, r ≠ .r2 → r ≠ .r3 → r ≠ .r4 → r ≠ .r12 → r ≠ .lr → s₆.gpr r = s.gpr r := + fun r h2 h3 h4 h12 hlr => by + rw [u₆.other _ h4, u₅.other _ h2, u₄.other _ h3, G₃, u₂.other _ h12, g₁.gpr _ h12 hlr] + have r5₆ : s₆.gpr .r5 = S s₀ := by rw [keep _ (by decide) (by decide) (by decide) (by decide) (by decide), h.r5] + have sp₆ : s₆.sp = s₀.sp := by rw [u₆.sp, u₅.sp, u₄.sp, sp₃, u₂.sp, g₁.sp, h.sp] + have rd₆ : s₆.rd = s₀.rd := by rw [u₆.rd, u₅.rd, u₄.rd, rd₃, u₂.rd, g₁.rd, h.rd] + have wr₆ : s₆.wr = s₀.wr := by rw [u₆.wr, u₅.wr, u₄.wr, wr₃, u₂.wr, g₁.wr, h.wr] + have mem₆ : s₆.mem = Proof.Cmac.zero4 (Proof.Cmac.xor4Mem s.mem (Ca s₀) (Ca s₀) (State.addr (St s₀))) + (State.addr (St s₀)) := by + rw [u₆.mem, u₅.mem, u₄.mem, m₃, r2₂, add0, u₂.mem, g₁.mem, h.r5, r2, add0] + have hb : below s₆ = belowR s₀ := by rw [below, sp₆]; rfl + have stS : Spec.Aes.bytesAt s.mem (State.addr (St s₀)) 16 = Spec.Aes.bytesAt s₀.mem (State.addr (St s₀)) 16 := + Proof.Cmac.bytesAt_frame16 ((fsMem_frame s₀).trans (h.frame.sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact ⟨scrR s₀, List.mem_singleton_self _, Offset.sub_base _ (by decide)⟩)) fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact hp.st_scr + refine ⟨?_, ?_, ?_, r5₆, fun r hr h4 h5 hlr => ?_, sp₆, rd₆, wr₆⟩ + · exact + { r0 := by rw [keep _ (by decide) (by decide) (by decide) (by decide) (by decide), + h.keep _ (by decide) (by decide) (by decide) (by decide) (by decide)] + r1 := by + rw [keep _ (by decide) (by decide) (by decide) (by decide) (by decide), + h.keep _ (by decide) (by decide) (by decide) (by decide) (by decide)]; simp [R] + r2 := by rw [u₆.other _ (by decide), u₅.gpr, u₄.other _ (by decide), G₃, u₂.other _ (by decide), + g₁.gpr _ (by decide) (by decide), h.r5] + r3 := by rw [u₆.other _ (by decide), u₅.other _ (by decide), u₄.gpr, G₃, r2₂] + hra := u₆.gpr + hrb := r5₆ + regs := by decide + rounds := hR + hsp := by rw [sp₆]; exact hp.sp8 + wc := by rw [cA]; exact (hp.ca_key (d := 0) (n := 240) (by decide)).symm |> fun d => by simpa using d + wd := hp.key_st.sub_left (Region.sub_prefix (by decide)) + ws := (hp.key_scr.sub_left (Region.sub_prefix (by decide))).sub_right (Region.sub_prefix (by decide)) + cd := by rw [cA]; exact cSt + cs := by rw [cA]; exact Offset.disjoint_base _ (by decide) (by omega) + ds := hp.st_scr.sub_right (Region.sub_prefix (by decide)) + bw := by rw [hb]; exact hp.b_key.sub_right (Region.sub_prefix (by decide)) + bc := by rw [hb, cA]; exact hp.b_scr.sub_right (Offset.sub_base _ (by decide)) + bd := by rw [hb]; exact hp.b_st + bs := by rw [hb]; exact hp.b_scr.sub_right (Region.sub_prefix (by decide)) + hW := by have := hp.key_fit; omega + hC := by + rw [BitVec.toNat_add, BitVec.toNat_ofNat, Nat.mod_eq_of_lt (a := 2048) (by decide), + Nat.mod_eq_of_lt (by omega)]; omega + hD := tf + hS := by omega + reads := by + rw [rd₆, wr₆, hp.rd] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact ⟨keyR s₀, by simp, 0, by simp, by simp⟩ + writes := by + rw [wr₆, hp.wr, cA] + refine Covers.of_sub fun r hr => ?_ + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl + · exact ⟨scrR s₀, by simp, 2048, rfl, by simp⟩ + · exact ⟨stR s₀, by simp, 0, by simp, by simp⟩ + · exact ⟨scrR s₀, by simp, 0, by simp, by simp⟩ + zero := by rw [mem₆]; exact Proof.Cmac.zero4_bytes _ _ } + · rw [mem₆, Proof.Cmac.zero4, Proof.Cmac.bytesAt_frame16 (Proof.Cmac.frame_store4 _ _ _ _ _) (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact cSt), + Proof.Cmac.xor4Mem_bytes _ (Proof.Cmac.Sep4.self _) (Proof.Cmac.Sep4.of_disjoint cSt), h.blk, stS] + · rw [mem₆] + refine ((h.frame.trans (Proof.Cmac.xor4Mem_frame _ _ _ _)).mono (by simp)).trans + ((Proof.Cmac.frame_store4 _ _ _ _ _).mono (by simp)) + · have hne := preserved_ne hr + rw [keep r hne.2.2.1 hne.2.2.2.1 h4 hne.2.2.2.2 hlr, h.keep r hne.2.2.2.1 h4 h5 hne.2.2.2.2 hlr] + +theorem finPre_wp {s₀ : State} (hp : FPre s₀) : WP isa finPre s₀ (FMid s₀) := by + refine WP.seq (WP.mono (finSave_wp hp) fun s₁ h₁ => ?_) + refine WP.seq (WP.mono (Q := BPost s₀) ?_ fun _ h => finArgs_wp hp h) + have ev : isa.eval .eq s₁ = some (decide (N s₀ = 16)) := by + show VG.Arm.eval .eq s₁ = _; rw [eval_eq, h₁.z] + by_cases hL : N s₀ = 16 + · exact WP.ite true (by rw [ev]; simp [hL]) (fun _ => full_wp hp hL h₁) (fun h => by cases h) + · exact WP.ite false (by rw [ev]; simp [hL]) (fun h => by cases h) + (fun _ => partial_wp hp (by have := hp.len; omega) h₁) + +/-! ## The whole function -/ + +theorem finalize_wp {s₀ : State} (h0 : finalizeArm.pre s₀) : + WP isa finalize s₀ fun s' => abiPreserved s₀ s' ∧ finalizeArm.post s₀ s' := by + have hp := FPre.of h0 + have hR := hp.rounds + have hRb : 16 * (R s₀ + 1) ≤ 240 := by rcases hR with h | h | h <;> omega + have sf := hp.scr_fit + have cA : State.addr (S s₀ + BitVec.ofNat 32 2048) = Ca s₀ := addr_add (by omega) + refine WP.seq (WP.mono (finPre_wp hp) fun s₁ h₁ => ?_) + refine WP.seq (WP.mono (ctr_call h₁.pre) fun s₂ h₂ => ?_) + have r5₂ : s₂.gpr .r5 = S s₀ := by rw [h₂.saved .r5 (by simp [preserved]) (by decide), h₁.r5] + have rdwr₂ : s₂.rd ++ s₂.wr = s₀.rd ++ s₀.wr := by rw [h₂.rd, h₂.wr, h₁.rd, h₁.wr] + have inS : ∀ d, d + 4 ≤ 2176 → InRegions (s₂.rd ++ s₂.wr) (State.addr (S s₀) + BitVec.ofNat 64 d) 4 := + fun d hd => by + rw [rdwr₂] + obtain ⟨r, hr, hc⟩ := in_of_cov (hp.cS (d := d) (n := 4) hd) + exact ⟨r, List.mem_append_right _ hr, hc⟩ + rw [show ([.ldr .r4 .r5 2064, .ldr .lr .r5 2072, .ldr .r5 .r5 2068] : List Instr) = + [(.r4, 2064), (.lr, 2072)].map (fun (p : Reg × Nat) => Instr.ldr p.1 .r5 p.2) ++ [.ldr .r5 .r5 2068] from rfl] + refine restoreB_ok [(.r4, 2064), (.lr, 2072)] s₂ _ (by decide) (fun p hp' => ?_) + fun s₃ ld₃ ho₃ m₃ rd₃ wr₃ sp₃ => ?_ + · have hb : 2064 ≤ p.2 ∧ p.2 + 4 ≤ 2076 ∧ p.1 ≠ .r5 := by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hp' + rcases hp' with rfl | rfl <;> decide + exact ⟨hb.2.2, by omega, by rw [r5₂]; omega, by rw [r5₂]; exact inS _ (by omega)⟩ + refine wp_ldr (a := State.addr (S s₀) + BitVec.ofNat 64 2068) (by decide) + (by rw [ho₃ _ (by decide), r5₂]; exact addr_add (by omega)) + (by rw [rd₃, wr₃]; exact inS _ (by decide)) fun s₄ u₄ => WP.block_nil ?_ + -- The slots, which nothing after the save writes. + have slot : ∀ r d, (r, d) ∈ fsaved → s₂.mem.readW (State.addr (S s₀) + BitVec.ofNat 64 d) 32 = s₀.gpr r := by + intro r d hrd + have hd : 2064 ≤ d ∧ d + 4 ≤ 2076 := by + simp only [fsaved, List.mem_cons, List.not_mem_nil, or_false, Prod.mk.injEq] at hrd + omega + rw [h₂.frame.readW (r := ⟨State.addr (S s₀) + BitVec.ofNat 64 d, 4⟩) (Region.contains_self _ _) (fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl + · rw [cA]; exact Offset.disjoint _ (by omega) (by omega) (by omega) + · exact hp.st_scr.symm.sub_left (Offset.sub_base _ (by omega)) + · exact Offset.disjoint_base _ (by omega) (by omega) + · rw [below, h₁.sp]; exact hp.b_scr.symm.sub_left (Offset.sub_base _ (by omega))) (by decide), + h₁.frame.readW (r := ⟨State.addr (S s₀) + BitVec.ofNat 64 d, 4⟩) (Region.contains_self _ _) (fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl + · exact Offset.disjoint _ (by omega) (by omega) (by omega) + · exact hp.st_scr.symm.sub_left (Offset.sub_base _ (by omega))) (by decide), + fsMem_slot s₀ hrd] + have sch : Spec.Aes.bytesAt s₁.mem (State.addr (W s₀)) (16 * (R s₀ + 1)) = + Spec.Aes.bytesAt s₀.mem (State.addr (W s₀)) (16 * (R s₀ + 1)) := by + have f : Frame [scrR s₀, stR s₀] s₀.mem s₁.mem := + (((fsMem_frame s₀).mono (by simp)).trans (h₁.frame.sub fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl + · exact ⟨scrR s₀, by simp, Offset.sub_base _ (by decide)⟩ + · exact ⟨stR s₀, by simp, fun _ h => h⟩)) + exact Proof.Cmac.bytesAt_frame f (fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl + · exact hp.key_scr.sub_left (Region.sub_prefix (by omega)) + · exact hp.key_st.sub_left (Region.sub_prefix (by omega))) (by omega) + refine ⟨⟨fun r hr => ?_, by rw [u₄.sp, sp₃, h₂.sp, h₁.sp]⟩, ?_⟩ + · by_cases h4 : r = .r4 + · subst h4; rw [u₄.other _ (by decide), ld₃ (.r4, 2064) (by simp), r5₂, slot .r4 2064 (by decide)] + by_cases h5 : r = .r5 + · subst h5; rw [u₄.gpr, m₃, slot .r5 2068 (by decide)] + by_cases hlr : r = .lr + · subst hlr; rw [u₄.other _ (by decide), ld₃ (.lr, 2072) (by simp), r5₂, slot .lr 2072 (by decide)] + rw [u₄.other _ h5, ho₃ _ (by simp [h4, hlr]), h₂.saved r hr hlr, h₁.keep r hr h4 h5 hlr] + · intro hk msg hm hne hst + have hk' : Spec.Aes.bytesAt s₀.mem (State.addr (W s₀) + BitVec.ofNat 64 240) 32 = + (Spec.Cmac.subkeys (ciph s₀) 16).1 ++ (Spec.Cmac.subkeys (ciph s₀) 16).2 := hk + obtain ⟨e1, e2⟩ := Proof.Cmac.k1k2 (Proof.Cmac.subkeys_aes_length _ _) hk' + show Spec.Aes.bytesAt s₄.mem (State.addr (St s₀)) 16 = _ + rw [u₄.mem, m₃, h₂.out, sch, cA, h₁.blk, mn, e1, e2, hst, + Proof.Cmac.macFull_split _ hm (by rw [Proof.Cmac.bytesAt_length]; exact hp.len) + (by rw [Proof.Cmac.bytesAt_length]; exact hne), Proof.Cmac.xor_comm] + +end VG.Proof.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/CmacAes/Arm/Subkeys.lean b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Subkeys.lean new file mode 100644 index 000000000..4787e900f --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Subkeys.lean @@ -0,0 +1,331 @@ +import VerifiedGarbage.Proof.CmacAes.Arm.Dbl +import VerifiedGarbage.Proof.CmacAes.Arm.UpdateCorrect + +/-! +# AES-CMAC on ARMv7: `vg_cmac_aes_subkeys` + +Untrusted: everything here is checked by Lean. `L = CIPH_K(0)` is computed +into the first block of the subkeys (a zero counter block and a zero data +block), then doubled there (`K1`) and into the second block (`K2`). +-/ + +namespace VG.Proof.CmacAes.Arm + +open VG VG.Arm VG.Impl.CmacAes.Arm +open VG.Proof.MdStream.Arm (Upd Mupd op2_imm op2_reg wp_mov wp_add wp_ldr saveMem saveList_ok saveMem_frame + readW_writeW_save) + +section +variable (s₀ : State) + +/-- The subkeys. -/ +abbrev Kb : BitVec 32 := s₀.gpr .r2 +/-- The scratch buffer. -/ +abbrev Sc : BitVec 32 := s₀.gpr .r3 + +abbrev kR : Region := ⟨State.addr (Kb s₀), 32⟩ +abbrev scR : Region := ⟨State.addr (Sc s₀), 2176⟩ + +end + +/-- The precondition, by name. -/ +structure SPre (s₀ : State) : Prop where + rd : s₀.rd = [schR s₀] + wr : s₀.wr = [kR s₀, scR s₀] + sch_k : (schR s₀).Disjoint (kR s₀) + sch_scr : (schR s₀).Disjoint (scR s₀) + k_scr : (kR s₀).Disjoint (scR s₀) + b_sch : (belowR s₀).Disjoint (schR s₀) + b_k : (belowR s₀).Disjoint (kR s₀) + b_scr : (belowR s₀).Disjoint (scR s₀) + sch_fit : (W s₀).toNat + 240 ≤ 2 ^ 32 + k_fit : (Kb s₀).toNat + 32 ≤ 2 ^ 32 + scr_fit : (Sc s₀).toNat + 2176 ≤ 2 ^ 32 + sp8 : 8 ≤ s₀.sp.toNat + rounds : R s₀ = 10 ∨ R s₀ = 12 ∨ R s₀ = 14 + +theorem SPre.of {s₀ : State} (h : subkeysArm.pre s₀) : SPre s₀ := + let ⟨a, b, c, d, e, f, g, h, i, j, k, l, m⟩ := h + ⟨a, b, c, d, e, f, g, h, i, j, k, l, m⟩ + +/-- The registers `subkeys` saves, and where. -/ +def saved4 : List (Reg × Nat) := [(.r4, 2064), (.r5, 2068), (.r6, 2072), (.lr, 2076)] + +theorem subkeysPre_eq : subkeysPre = saved4.map (fun p => Instr.str p.1 .r3 p.2) ++ + (.mov .r6 (.reg .r2) :: .mov .r5 (.reg .r3) :: .mov .r12 (.imm 0) :: (zeroBlk .r12 .r3 2048 ++ + (zeroBlk .r12 .r2 0 ++ ([.dp .add .r2 .r5 (.imm (BitVec.ofNat 32 2048)), .mov .r3 (.reg .r6), + .mov .r4 (.imm 1)] : List Instr)))) := rfl + +theorem subkeysPost_eq : subkeysPost = dbl 0 0 ++ (dbl 0 16 ++ + ([(.r4, 2064), (.r6, 2072), (.lr, 2076)].map (fun (p : Reg × Nat) => Instr.ldr p.1 .r5 p.2) ++ + ([.ldr .r5 .r5 2068] : List Instr))) := rfl + +set_option simprocs false in +theorem saved4_slot (m : Mem) (B : Addr) (g : Reg → BitVec 32) {r : Reg} {d : Nat} (h : (r, d) ∈ saved4) : + (saveMem m B g saved4).readW (B + BitVec.ofNat 64 d) 32 = g r := by + simp only [saved4, List.mem_cons, List.not_mem_nil, or_false, Prod.mk.injEq] at h + rcases h with ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ <;> + simp (disch := decide) only [saved4, saveMem, Mem.readW_writeW_self32, readW_writeW_save] + +/-- The memory before the call. -/ +def preMem (s₀ : State) : Mem := + Proof.Cmac.zero4 (Proof.Cmac.zero4 (saveMem s₀.mem (State.addr (Sc s₀)) s₀.gpr saved4) + (State.addr (Sc s₀) + BitVec.ofNat 64 2048)) (State.addr (Kb s₀)) + +/-- What the code before the call leaves. -/ +structure SAfter (s₀ s : State) : Prop where + pre : CallPre s (W s₀) (Sc s₀ + BitVec.ofNat 32 2048) (Kb s₀) (Sc s₀) (R s₀) .r4 .r5 + keep : ∀ r, r ≠ .r2 → r ≠ .r3 → r ≠ .r4 → r ≠ .r5 → r ≠ .r6 → r ≠ .r12 → s.gpr r = s₀.gpr r + r5 : s.gpr .r5 = Sc s₀ + r6 : s.gpr .r6 = Kb s₀ + sp : s.sp = s₀.sp + rd : s.rd = s₀.rd + wr : s.wr = s₀.wr + mem : s.mem = preMem s₀ + +theorem pre_wp {s₀ : State} (hp : SPre s₀) : WP isa (.block subkeysPre) s₀ (SAfter s₀) := by + have kf := hp.k_fit + have sf := hp.scr_fit + have sf' : (s₀.gpr .r3).toNat + 2176 ≤ 2 ^ 32 := sf + have hR := hp.rounds + have hRb : 16 * (R s₀ + 1) ≤ 240 := by rcases hR with h | h | h <;> omega + have cK : ∀ d n, d + n ≤ 32 → Covers [⟨State.addr (Kb s₀) + BitVec.ofNat 64 d, n⟩] s₀.wr := fun d n h => by + rw [hp.wr] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact ⟨kR s₀, by simp, d, rfl, h⟩ + have cS : ∀ d n, d + n ≤ 2176 → Covers [⟨State.addr (Sc s₀) + BitVec.ofNat 64 d, n⟩] s₀.wr := fun d n h => by + rw [hp.wr] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact ⟨scR s₀, by simp, d, rfl, h⟩ + have rw' : ∀ {rs a n}, Covers rs s₀.wr → InRegions rs a n → InRegions (s₀.rd ++ s₀.wr) a n := + fun h hi => by obtain ⟨r, hr, hc⟩ := h _ _ hi; exact ⟨r, List.mem_append_right _ hr, hc⟩ + have cKr : ∀ d n, d + n ≤ 32 → Covers [⟨State.addr (Kb s₀) + BitVec.ofNat 64 d, n⟩] (s₀.rd ++ s₀.wr) := + fun d n h a k hi => rw' (cK d n h) hi + have aS : ∀ {d}, d < 2176 → State.addr (Sc s₀ + BitVec.ofNat 32 d) = State.addr (Sc s₀) + BitVec.ofNat 64 d := + fun _ => addr_add (by omega) + -- Before the call. + rw [subkeysPre_eq] + refine saveList_ok saved4 s₀ _ (fun p hp' => ?_) fun s₁ g₁ rd₁ wr₁ sp₁ m₁ => ?_ + · have hb : 2064 ≤ p.2 ∧ p.2 + 4 ≤ 2080 := by + simp only [saved4, List.mem_cons, List.not_mem_nil, or_false] at hp' + rcases hp' with rfl | rfl | rfl | rfl <;> decide + exact ⟨by omega, by omega, cS p.2 4 (by omega) _ _ ⟨_, List.mem_singleton_self _, Region.contains_self _ _⟩⟩ + refine wp_mov (op2_reg _ _) fun s₂ u₂ => wp_mov (op2_reg _ _) fun s₃ u₃ => + wp_mov (op2_imm (by decide)) fun s₄ u₄ => ?_ + have r3₄ : s₄.gpr .r3 = Sc s₀ := by rw [u₄.other _ (by decide), u₃.other _ (by decide), u₂.other _ (by decide), g₁] + refine Proof.CmacAes.Arm.zeroBlk_ok u₄.gpr (by decide) (by rw [r3₄]; omega) + (by rw [r3₄, u₄.wr, u₃.wr, u₂.wr, wr₁]; exact cS 2048 16 (by decide)) fun s₅ G₅ m₅ rd₅ wr₅ sp₅ => ?_ + have r2₅ : s₅.gpr .r2 = Kb s₀ := by + rw [G₅, u₄.other _ (by decide), u₃.other _ (by decide), u₂.other _ (by decide), g₁] + refine Proof.CmacAes.Arm.zeroBlk_ok (by rw [G₅, u₄.gpr]) (by decide) (by rw [r2₅]; omega) + (by rw [r2₅, add0, wr₅, u₄.wr, u₃.wr, u₂.wr, wr₁]; simpa using cK 0 16 (by decide)) + fun s₆ G₆ m₆ rd₆ wr₆ sp₆ => ?_ + refine wp_add (op2_imm (by decide)) fun s₇ u₇ => wp_mov (op2_reg _ _) fun s₈ u₈ => + wp_mov (op2_imm (by decide)) fun s₉ u₉ => WP.block_nil ?_ + have keep₉ : ∀ r, r ≠ .r2 → r ≠ .r3 → r ≠ .r4 → r ≠ .r5 → r ≠ .r6 → r ≠ .r12 → s₉.gpr r = s₀.gpr r := + fun r h2 h3 h4 h5 h6 h12 => by + rw [u₉.other _ h4, u₈.other _ h3, u₇.other _ h2, G₆, G₅, u₄.other _ h12, u₃.other _ h5, u₂.other _ h6, g₁] + have r5₉ : s₉.gpr .r5 = Sc s₀ := by + rw [u₉.other _ (by decide), u₈.other _ (by decide), u₇.other _ (by decide), G₆, G₅, u₄.other _ (by decide), + u₃.gpr, u₂.other _ (by decide), g₁] + have r6₉ : s₉.gpr .r6 = Kb s₀ := by + rw [u₉.other _ (by decide), u₈.other _ (by decide), u₇.other _ (by decide), G₆, G₅, u₄.other _ (by decide), + u₃.other _ (by decide), u₂.gpr, g₁] + have sp₉ : s₉.sp = s₀.sp := by rw [u₉.sp, u₈.sp, u₇.sp, sp₆, sp₅, u₄.sp, u₃.sp, u₂.sp, sp₁] + have rd₉ : s₉.rd = s₀.rd := by rw [u₉.rd, u₈.rd, u₇.rd, rd₆, rd₅, u₄.rd, u₃.rd, u₂.rd, rd₁] + have wr₉ : s₉.wr = s₀.wr := by rw [u₉.wr, u₈.wr, u₇.wr, wr₆, wr₅, u₄.wr, u₃.wr, u₂.wr, wr₁] + have mem₉ : s₉.mem = preMem s₀ := by + rw [u₉.mem, u₈.mem, u₇.mem, m₆, r2₅, add0, m₅, r3₄, u₄.mem, u₃.mem, u₂.mem, m₁]; rfl + -- The memory before the call. + have cA : State.addr (Sc s₀ + BitVec.ofNat 32 2048) = State.addr (Sc s₀) + BitVec.ofNat 64 2048 := aS (by decide) + have kC : (⟨State.addr (Kb s₀), 16⟩ : Region).Disjoint ⟨State.addr (Sc s₀) + BitVec.ofNat 64 2048, 16⟩ := + (hp.k_scr.sub_left (Region.sub_prefix (by decide))).sub_right (Offset.sub_base _ (by decide)) + have hb : below s₉ = belowR s₀ := by rw [below, sp₉]; rfl + have pre : CallPre s₉ (W s₀) (Sc s₀ + BitVec.ofNat 32 2048) (Kb s₀) (Sc s₀) (R s₀) .r4 .r5 := + { r0 := keep₉ _ (by decide) (by decide) (by decide) (by decide) (by decide) (by decide) + r1 := by + rw [keep₉ _ (by decide) (by decide) (by decide) (by decide) (by decide) (by decide)]; simp [R] + r2 := by + rw [u₉.other _ (by decide), u₈.other _ (by decide), u₇.gpr, G₆, G₅, u₄.other _ (by decide), u₃.gpr, + u₂.other _ (by decide), g₁] + r3 := by rw [u₉.other _ (by decide), u₈.gpr, u₇.other _ (by decide), G₆, G₅, u₄.other _ (by decide), + u₃.other _ (by decide), u₂.gpr, g₁] + hra := u₉.gpr + hrb := r5₉ + regs := by decide + rounds := hR + hsp := by rw [sp₉]; exact hp.sp8 + wc := by rw [cA]; exact hp.sch_scr.sub_right (Offset.sub_base _ (by decide)) + wd := hp.sch_k.sub_right (Region.sub_prefix (by decide)) + ws := hp.sch_scr.sub_right (Region.sub_prefix (by decide)) + cd := by rw [cA]; exact kC.symm + cs := by rw [cA]; exact Offset.disjoint_base _ (by decide) (by omega) + ds := (hp.k_scr.sub_left (Region.sub_prefix (by decide))).sub_right (Region.sub_prefix (by decide)) + bw := by rw [hb]; exact hp.b_sch + bc := by rw [hb, cA]; exact hp.b_scr.sub_right (Offset.sub_base _ (by decide)) + bd := by rw [hb]; exact hp.b_k.sub_right (Region.sub_prefix (by decide)) + bs := by rw [hb]; exact hp.b_scr.sub_right (Region.sub_prefix (by decide)) + hW := hp.sch_fit + hC := by + rw [BitVec.toNat_add, BitVec.toNat_ofNat, Nat.mod_eq_of_lt (a := 2048) (by decide), + Nat.mod_eq_of_lt (by omega)]; omega + hD := by omega + hS := by omega + reads := by + rw [rd₉, wr₉, hp.rd] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact ⟨schR s₀, by simp, 0, by simp, by simp⟩ + writes := by + rw [wr₉, hp.wr, cA] + refine Covers.of_sub fun r hr => ?_ + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl + · exact ⟨scR s₀, by simp, 2048, rfl, by simp⟩ + · exact ⟨kR s₀, by simp, 0, by simp, by simp⟩ + · exact ⟨scR s₀, by simp, 0, by simp, by simp⟩ + zero := by rw [mem₉, preMem]; exact Proof.Cmac.zero4_bytes _ _ } + exact ⟨pre, keep₉, r5₉, r6₉, sp₉, rd₉, wr₉, mem₉⟩ + +theorem subkeys_wp {s₀ : State} (h0 : subkeysArm.pre s₀) : + WP isa subkeys s₀ fun s' => abiPreserved s₀ s' ∧ subkeysArm.post s₀ s' := by + have hp := SPre.of h0 + have kf := hp.k_fit + have sf := hp.scr_fit + have sf' : (s₀.gpr .r3).toNat + 2176 ≤ 2 ^ 32 := sf + have hR := hp.rounds + have hRb : 16 * (R s₀ + 1) ≤ 240 := by rcases hR with h | h | h <;> omega + have cK : ∀ d n, d + n ≤ 32 → Covers [⟨State.addr (Kb s₀) + BitVec.ofNat 64 d, n⟩] s₀.wr := fun d n h => by + rw [hp.wr] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact ⟨kR s₀, by simp, d, rfl, h⟩ + have cS : ∀ d n, d + n ≤ 2176 → Covers [⟨State.addr (Sc s₀) + BitVec.ofNat 64 d, n⟩] s₀.wr := fun d n h => by + rw [hp.wr] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact ⟨scR s₀, by simp, d, rfl, h⟩ + have rw' : ∀ {rs a n}, Covers rs s₀.wr → InRegions rs a n → InRegions (s₀.rd ++ s₀.wr) a n := + fun h hi => by obtain ⟨r, hr, hc⟩ := h _ _ hi; exact ⟨r, List.mem_append_right _ hr, hc⟩ + have cKr : ∀ d n, d + n ≤ 32 → Covers [⟨State.addr (Kb s₀) + BitVec.ofNat 64 d, n⟩] (s₀.rd ++ s₀.wr) := + fun d n h a k hi => rw' (cK d n h) hi + have aS : ∀ {d}, d < 2176 → State.addr (Sc s₀ + BitVec.ofNat 32 d) = State.addr (Sc s₀) + BitVec.ofNat 64 d := + fun _ => addr_add (by omega) + unfold subkeys + refine WP.seq (WP.mono (pre_wp hp) fun s₉ a => ?_) + obtain ⟨pre, keep₉, r5₉, r6₉, sp₉, rd₉, wr₉, mem₉⟩ := a + -- The memory before the call. + have cA : State.addr (Sc s₀ + BitVec.ofNat 32 2048) = State.addr (Sc s₀) + BitVec.ofNat 64 2048 := aS (by decide) + have kC : (⟨State.addr (Kb s₀), 16⟩ : Region).Disjoint ⟨State.addr (Sc s₀) + BitVec.ofNat 64 2048, 16⟩ := + (hp.k_scr.sub_left (Region.sub_prefix (by decide))).sub_right (Offset.sub_base _ (by decide)) + have f₉ : Frame [scR s₀, kR s₀] s₀.mem s₉.mem := by + rw [mem₉, preMem] + refine (((saveMem_frame _ _ _ (L := 2176) (by decide) saved4 (by decide)).sub fun r hr => ?_).trans + ((Proof.Cmac.frame_store4 _ _ _ _ _).sub fun r hr => ?_)).trans + ((Proof.Cmac.frame_store4 _ _ _ _ _).sub fun r hr => ?_) <;> + simp only [List.mem_singleton] at hr <;> subst hr + · exact ⟨scR s₀, by simp, fun _ h => h⟩ + · exact ⟨scR s₀, by simp, Offset.sub_base _ (by decide)⟩ + · exact ⟨kR s₀, by simp, Region.sub_prefix (by decide)⟩ + have zC : Spec.Aes.bytesAt s₉.mem (State.addr (Sc s₀) + BitVec.ofNat 64 2048) 16 = Spec.Cmac.zeros 16 := by + rw [mem₉, preMem, Proof.Cmac.zero4, Proof.Cmac.bytesAt_frame16 (Proof.Cmac.frame_store4 _ _ _ _ _) (by + intro r hr; simp only [List.mem_singleton] at hr; subst hr; exact kC.symm)] + exact Proof.Cmac.zero4_bytes _ _ + have hb : below s₉ = belowR s₀ := by rw [below, sp₉]; rfl + -- The call. + refine WP.seq (WP.mono (ctr_call pre) fun s₁₀ h₁₀ => ?_) + have sv (r : Reg) (hr : r ∈ preserved) (hlr : r ≠ .lr) : s₁₀.gpr r = s₉.gpr r := h₁₀.saved r hr hlr + have r6₁₀ : s₁₀.gpr .r6 = Kb s₀ := by rw [sv .r6 (by simp [preserved]) (by decide), r6₉] + have rdwr₁₀ : s₁₀.rd ++ s₁₀.wr = s₀.rd ++ s₀.wr := by rw [h₁₀.rd, h₁₀.wr, rd₉, wr₉] + have wr₁₀ : s₁₀.wr = s₀.wr := by rw [h₁₀.wr, wr₉] + have L : Spec.Aes.bytesAt s₁₀.mem (State.addr (Kb s₀)) 16 = ciph s₀ (Spec.Cmac.zeros 16) := by + rw [h₁₀.out, cA, zC, Proof.Cmac.bytesAt_frame f₉ (fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl + · exact hp.sch_scr.sub_left (Region.sub_prefix hRb) + · exact hp.sch_k.sub_left (Region.sub_prefix hRb)) (by omega)] + -- After the call. + rw [subkeysPost_eq] + refine dbl_wp r6₁₀ (by decide) (by decide) (by omega) (by omega) + (by rw [rdwr₁₀]; exact cKr 0 16 (by decide)) (by rw [wr₁₀]; exact cK 0 16 (by decide)) + fun s₁₁ g₁₁ m₁₁ rd₁₁ wr₁₁ sp₁₁ => ?_ + have r6₁₁ : s₁₁.gpr .r6 = Kb s₀ := by + rw [g₁₁ _ (by decide) (by decide) (by decide) (by decide) (by decide) (by decide), r6₁₀] + refine dbl_wp r6₁₁ (by decide) (by decide) (by omega) (by omega) + (by rw [rd₁₁, wr₁₁, rdwr₁₀]; exact cKr 0 16 (by decide)) (by rw [wr₁₁, wr₁₀]; exact cK 16 16 (by decide)) + fun s₁₂ g₁₂ m₁₂ rd₁₂ wr₁₂ sp₁₂ => ?_ + have k₁₂ : ∀ r, r ≠ .r0 → r ≠ .r1 → r ≠ .r2 → r ≠ .r3 → r ≠ .r4 → r ≠ .r12 → s₁₂.gpr r = s₁₀.gpr r := + fun r h0 h1 h2 h3 h4 h12 => by rw [g₁₂ r h0 h1 h2 h3 h4 h12, g₁₁ r h0 h1 h2 h3 h4 h12] + have r5₁₂ : s₁₂.gpr .r5 = Sc s₀ := by + rw [k₁₂ _ (by decide) (by decide) (by decide) (by decide) (by decide) (by decide), + sv .r5 (by simp [preserved]) (by decide), r5₉] + have rdwr₁₂ : s₁₂.rd ++ s₁₂.wr = s₀.rd ++ s₀.wr := by rw [rd₁₂, wr₁₂, rd₁₁, wr₁₁, rdwr₁₀] + have inS : ∀ d, d + 4 ≤ 2176 → InRegions (s₁₂.rd ++ s₁₂.wr) (State.addr (Sc s₀) + BitVec.ofNat 64 d) 4 := + fun d hd => by + rw [rdwr₁₂] + exact rw' (cS d 4 hd) ⟨_, List.mem_singleton_self _, Region.contains_self _ _⟩ + refine restoreB_ok [(.r4, 2064), (.r6, 2072), (.lr, 2076)] s₁₂ _ (by decide) (fun p hp' => ?_) + fun s₁₃ ld₁₃ ho₁₃ m₁₃ rd₁₃ wr₁₃ sp₁₃ => ?_ + · have hb : 2064 ≤ p.2 ∧ p.2 + 4 ≤ 2080 ∧ p.1 ≠ .r5 := by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hp' + rcases hp' with rfl | rfl | rfl <;> decide + exact ⟨hb.2.2, by omega, by rw [r5₁₂]; omega, by rw [r5₁₂]; exact inS _ (by omega)⟩ + refine wp_ldr (a := State.addr (Sc s₀) + BitVec.ofNat 64 2068) (by decide) + (by rw [ho₁₃ _ (by decide), r5₁₂]; exact aS (by decide)) + (by rw [rd₁₃, wr₁₃]; exact inS _ (by decide)) fun s₁₄ u₁₄ => WP.block_nil ?_ + -- The slots. + have slotD : ∀ d, 2064 ≤ d → d + 4 ≤ 2080 → + ∀ r ∈ [⟨State.addr (Sc s₀ + BitVec.ofNat 32 2048), 16⟩, ⟨State.addr (Kb s₀), 16⟩, + ⟨State.addr (Sc s₀), 2048⟩, below s₉], (⟨State.addr (Sc s₀) + BitVec.ofNat 64 d, 4⟩ : Region).Disjoint r := by + intro d h₁ h₂ r hr + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl + · rw [cA]; exact Offset.disjoint _ (by omega) (by omega) (by omega) + · exact (hp.k_scr.symm.sub_left (Offset.sub_base _ (by omega))).sub_right (Region.sub_prefix (by decide)) + · exact Offset.disjoint_base _ (by omega) (by omega) + · rw [hb]; exact hp.b_scr.symm.sub_left (Offset.sub_base _ (by omega)) + have slotK : ∀ d, 2064 ≤ d → d + 4 ≤ 2080 → ∀ e, e ≤ 16 → + (⟨State.addr (Sc s₀) + BitVec.ofNat 64 d, 4⟩ : Region).Disjoint ⟨State.addr (Kb s₀) + BitVec.ofNat 64 e, 16⟩ := + fun d h₁ h₂ e he => (hp.k_scr.symm.sub_left (Offset.sub_base _ (by omega))).sub_right (Offset.sub_base _ (by omega)) + have slot : ∀ r d, (r, d) ∈ saved4 → s₁₂.mem.readW (State.addr (Sc s₀) + BitVec.ofNat 64 d) 32 = s₀.gpr r := by + intro r d hrd + have hd : 2064 ≤ d ∧ d + 4 ≤ 2080 := by + simp only [saved4, List.mem_cons, List.not_mem_nil, or_false, Prod.mk.injEq] at hrd + omega + rw [m₁₂, dblMem_frame _ _ _ _ |>.readW (Region.contains_self _ _) (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact slotK d hd.1 hd.2 16 (by decide)) (by decide), + m₁₁, dblMem_frame _ _ _ _ |>.readW (Region.contains_self _ _) (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact slotK d hd.1 hd.2 0 (by decide)) (by decide), + h₁₀.frame.readW (Region.contains_self _ _) (slotD d hd.1 hd.2) (by decide), mem₉, preMem, + Proof.Cmac.zero4, Proof.Cmac.readW_store4_of_sep _ _ _ _ + ((hp.k_scr.sub_left (Region.sub_prefix (by decide))).sub_right (Offset.sub_base _ (by omega))), + Proof.Cmac.zero4, Proof.Cmac.readW_store4_of_sep _ _ _ _ (Offset.disjoint _ (by omega) (by omega) (by omega)), + saved4_slot _ _ _ hrd] + refine ⟨⟨fun r hr => ?_, by rw [u₁₄.sp, sp₁₃, sp₁₂, sp₁₁, h₁₀.sp, sp₉]⟩, ?_⟩ + · simp only [preserved, List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl + · rw [u₁₄.other _ (by decide), ld₁₃ (.r4, 2064) (by simp), r5₁₂, slot .r4 2064 (by decide)] + · rw [u₁₄.gpr, m₁₃, slot .r5 2068 (by decide)] + · rw [u₁₄.other _ (by decide), ld₁₃ (.r6, 2072) (by simp), r5₁₂, slot .r6 2072 (by decide)] + all_goals first + | rw [u₁₄.other _ (by decide), ld₁₃ (.lr, 2076) (by simp), r5₁₂, slot .lr 2076 (by decide)] + | rw [u₁₄.other _ (by decide), ho₁₃ _ (by decide), + k₁₂ _ (by decide) (by decide) (by decide) (by decide) (by decide) (by decide), + sv _ (by simp [preserved]) (by decide), + keep₉ _ (by decide) (by decide) (by decide) (by decide) (by decide) (by decide)] + · show Spec.Aes.bytesAt s₁₄.mem (State.addr (Kb s₀)) 32 = _ + have b₁₁ : Spec.Aes.bytesAt s₁₁.mem (State.addr (Kb s₀)) 16 = + Spec.Cmac.dbl 16 (Spec.Aes.bytesAt s₁₀.mem (State.addr (Kb s₀)) 16) := by + have := dblMem_bytes s₁₀.mem (State.addr (Kb s₀)) 0 0 + rw [add0] at this; rw [m₁₁, this] + have lo : Spec.Aes.bytesAt s₁₂.mem (State.addr (Kb s₀)) 16 = Spec.Aes.bytesAt s₁₁.mem (State.addr (Kb s₀)) 16 := by + rw [m₁₂] + exact Proof.Cmac.bytesAt_frame16 (dblMem_frame _ _ _ _) fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact (Offset.disjoint_base _ (by decide) (by omega)).symm + have hi : Spec.Aes.bytesAt s₁₂.mem (State.addr (Kb s₀) + BitVec.ofNat 64 16) 16 = + Spec.Cmac.dbl 16 (Spec.Aes.bytesAt s₁₁.mem (State.addr (Kb s₀)) 16) := by + have := dblMem_bytes s₁₁.mem (State.addr (Kb s₀)) 0 16 + rw [add0] at this; rw [m₁₂, this] + rw [u₁₄.mem, m₁₃, Proof.Cmac.bytesAt_32, lo, hi, b₁₁, L] + rfl + +end VG.Proof.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/CmacAes/Arm/SubkeysCT.lean b/lean/VerifiedGarbage/Proof/CmacAes/Arm/SubkeysCT.lean new file mode 100644 index 000000000..8e24074a0 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/Arm/SubkeysCT.lean @@ -0,0 +1,65 @@ +import VerifiedGarbage.Proof.CmacAes.Arm.Subkeys +import VerifiedGarbage.Proof.CmacAes.Arm.UpdateCT + +/-! +# AES-CMAC on ARMv7: `vg_cmac_aes_subkeys` is constant time + +Untrusted: everything here is checked by Lean. The code before the call +and after it is checked by the taint analysis, from the registers the +correctness proof pins (the arguments, then `r5` and `r6`); the call of +`vg_aes_ctr32`, in its frame, is constant time by its own proof +(`ctr_rel`). +-/ + +namespace VG.Proof.CmacAes.Arm + +open VG VG.Arm VG.Impl.CmacAes.Arm + +/-- What is known after the call. -/ +structure SPost (s₀ : State) (s : State) : Prop where + r5 : s.gpr .r5 = Sc s₀ + r6 : s.gpr .r6 = Kb s₀ + +theorem spost_wp {s₀ s : State} (h : SAfter s₀ s) : WP isa (ctrCall .r4 .r5) s (SPost s₀) := + WP.mono (ctr_call h.pre) fun _ hc => + ⟨by rw [hc.saved .r5 (by simp [preserved]) (by decide), h.r5], + by rw [hc.saved .r6 (by simp [preserved]) (by decide), h.r6]⟩ + +theorem subkeys_rel {s₀ s₀' : State} (h0 : subkeysArm.pre s₀) (h0' : subkeysArm.pre s₀') + (hq : subkeysArm.pub s₀ s₀') : + RelCT isa (fun a b => a = s₀ ∧ b = s₀') subkeys fun _ _ => True := by + obtain ⟨q₀, q₁, q₂, q₃, q₄⟩ := hq + have hp := SPre.of h0 + have hp' := SPre.of h0' + have eW : W s₀' = W s₀ := q₁.symm + have eR : R s₀' = R s₀ := by rw [R, R, q₂] + have eK : Kb s₀' = Kb s₀ := q₃.symm + have eS : Sc s₀' = Sc s₀ := q₄.symm + obtain ⟨_, hA⟩ : ∃ h, (taint.check (Taint.ofRegs [.r0, .r1, .r2, .r3]) (.block subkeysPre) h).isSome = + true := ⟨_, by taint_decide⟩ + obtain ⟨_, hB⟩ : ∃ h, (taint.check (Taint.ofRegs [.r5, .r6]) (.block subkeysPost) h).isSome = true := + ⟨_, by taint_decide⟩ + have a := rel_agree (F := fun s => s = s₀) (F' := fun s => s = s₀') (G := SAfter s₀) (G' := SAfter s₀') + (Taint.ofRegs [.r0, .r1, .r2, .r3]) + (fun s s' e e' => by + subst e e' + refine Taint.agree_ofRegs fun r hr => ?_ + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl <;> assumption) ⟨_, hA⟩ + (fun s e => by rw [e]; exact pre_wp hp) (fun s e => by rw [e]; exact pre_wp hp') + have c := rel_wp (F := SAfter s₀) (F' := SAfter s₀') (G := SPost s₀) (G' := SPost s₀') + (ctr_rel (sp₀ := s₀.sp) fun s₁ s₂ h => + ⟨h.1.pre, by have := h.2.pre; rwa [eW, eS, eK, eR] at this, h.1.sp, h.2.sp.trans q₀.symm⟩) + (fun _ h => spost_wp h) (fun _ h => spost_wp h) + have b := RelCT.taint (A := taint) (P := fun s₁ s₂ => SPost s₀ s₁ ∧ SPost s₀' s₂) (Taint.ofRegs [.r5, .r6]) + (fun s₁ s₂ h => Taint.agree_ofRegs fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl + · rw [h.1.r5, h.2.r5, eS] + · rw [h.1.r6, h.2.r6, eK]) hB + exact a.seq (c.seq b) + +theorem subkeys_ct : ConstantTime isa subkeysArm.pre subkeysArm.pub subkeys := + fun _ _ _ _ _ _ h₁ h₂ hq e₁ e₂ => (subkeys_rel h₁ h₂ hq _ _ _ _ _ _ ⟨rfl, rfl⟩ e₁ e₂).1 + +end VG.Proof.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/CmacAes/Arm/Update.lean b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Update.lean new file mode 100644 index 000000000..36ad26ce2 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Update.lean @@ -0,0 +1,228 @@ +import VerifiedGarbage.Proof.CmacAes.Arm.Contract +import VerifiedGarbage.Proof.CmacAes.Arm.Words + +/-! +# AES-CMAC on ARMv7: `vg_cmac_aes_update`, the blocks before and in the loop + +Untrusted: everything here is checked by Lean. The invariant after `k` +blocks (`LInv`): the registers hold the arguments (`r7` the next block, `r8` +the blocks left), only the state, the first 2064 bytes of the scratch buffer +and the 8 bytes below the stack pointer have changed since the registers were +saved, and the state is the chaining value after the first `k` blocks. +-/ + +namespace VG.Proof.CmacAes.Arm + +open VG VG.Arm VG.Impl.CmacAes.Arm +open VG.Proof.MdStream.Arm (Upd Mupd Fupd op2_imm op2_reg wp_mov wp_add wp_subs wp_cmp wp_ldrSp saveMem + saveList_ok saveMem_frame readW_writeW_save cmp0 ofNat_beq_zero sub_ofNat) + +section +variable (s₀ : State) + +abbrev W : BitVec 32 := s₀.gpr .r0 +abbrev R : Nat := (s₀.gpr .r1).toNat +abbrev St : BitVec 32 := s₀.gpr .r2 +abbrev Dp : BitVec 32 := s₀.gpr .r3 +abbrev N : Nat := (stackArg s₀ 0).toNat +abbrev S : BitVec 32 := stackArg s₀ 1 + +abbrev schR : Region := ⟨State.addr (W s₀), 240⟩ +abbrev stR : Region := ⟨State.addr (St s₀), 16⟩ +abbrev dataR : Region := ⟨State.addr (Dp s₀), 16 * N s₀⟩ +abbrev scrR : Region := ⟨State.addr (S s₀), 2176⟩ +abbrev argsR : Region := ⟨stackArgAddr s₀ 0, 8⟩ +abbrev belowR : Region := ⟨State.addr s₀.sp - BitVec.ofNat 64 8, 8⟩ + +/-- The cipher. -/ +abbrev ciph : Spec.Cmac.Cipher := ciphAt s₀.mem (State.addr (W s₀)) (R s₀) + +/-- The message blocks. -/ +abbrev blks : List (List Byte) := Spec.Cmac.blocksAt s₀.mem (State.addr (Dp s₀)) 16 (N s₀) + +/-- The memory after saving the registers in the scratch buffer. -/ +def savedMem : Mem := saveMem s₀.mem (State.addr (S s₀)) s₀.gpr saved + +end + +/-- The precondition, by name. -/ +structure UPre (s₀ : State) : Prop where + rd : s₀.rd = [schR s₀, dataR s₀, argsR s₀] + wr : s₀.wr = [stR s₀, scrR s₀] + sch_st : (schR s₀).Disjoint (stR s₀) + sch_scr : (schR s₀).Disjoint (scrR s₀) + data_st : (dataR s₀).Disjoint (stR s₀) + data_scr : (dataR s₀).Disjoint (scrR s₀) + st_scr : (stR s₀).Disjoint (scrR s₀) + st_args : (stR s₀).Disjoint (argsR s₀) + scr_args : (scrR s₀).Disjoint (argsR s₀) + b_sch : (belowR s₀).Disjoint (schR s₀) + b_data : (belowR s₀).Disjoint (dataR s₀) + b_st : (belowR s₀).Disjoint (stR s₀) + b_scr : (belowR s₀).Disjoint (scrR s₀) + sch_fit : (W s₀).toNat + 240 ≤ 2 ^ 32 + st_fit : (St s₀).toNat + 16 ≤ 2 ^ 32 + data_fit : (Dp s₀).toNat + 16 * N s₀ ≤ 2 ^ 32 + scr_fit : (S s₀).toNat + 2176 ≤ 2 ^ 32 + sp8 : 8 ≤ s₀.sp.toNat + sp_fit : s₀.sp.toNat + 8 ≤ 2 ^ 32 + rounds : R s₀ = 10 ∨ R s₀ = 12 ∨ R s₀ = 14 + +theorem UPre.of {s₀ : State} (h : updateArm.pre s₀) : UPre s₀ := + let ⟨a, b, c, d, e, f, g, h, i, j, k, l, m, n, o, p, q, r, t, u⟩ := h + ⟨a, b, c, d, e, f, g, h, i, j, k, l, m, n, o, p, q, r, t, u⟩ + +/-- The loop invariant, after `k` blocks. -/ +structure LInv (s₀ : State) (k : Nat) (s : State) : Prop where + r4 : s.gpr .r4 = W s₀ + r5 : s.gpr .r5 = s₀.gpr .r1 + r6 : s.gpr .r6 = St s₀ + r7 : s.gpr .r7 = Dp s₀ + BitVec.ofNat 32 (16 * k) + r8 : s.gpr .r8 = BitVec.ofNat 32 (N s₀ - k) + r10 : s.gpr .r10 = S s₀ + r11 : s.gpr .r11 = s₀.gpr .r11 + sp : s.sp = s₀.sp + rd : s.rd = s₀.rd + wr : s.wr = s₀.wr + frame : Frame [stR s₀, ⟨State.addr (S s₀), 2064⟩, belowR s₀] (savedMem s₀) s.mem + state : Spec.Aes.bytesAt s.mem (State.addr (St s₀)) 16 = + Spec.Cmac.chain (ciph s₀) (Spec.Aes.bytesAt s₀.mem (State.addr (St s₀)) 16) ((blks s₀).take k) + +/-! ## Addresses and regions -/ + +theorem add0 (p : Addr) : p + BitVec.ofNat 64 0 = p := BitVec.add_zero p + +section +variable {s₀ : State} (hp : UPre s₀) +include hp + +theorem UPre.scrA {d : Nat} (hd : d < 2176) : + State.addr (S s₀ + BitVec.ofNat 32 d) = State.addr (S s₀) + BitVec.ofNat 64 d := + addr_add (by have := hp.scr_fit; have := (S s₀).isLt; omega) + +theorem UPre.dataA {k : Nat} (hk : k < N s₀) : + State.addr (Dp s₀ + BitVec.ofNat 32 (16 * k)) = State.addr (Dp s₀) + BitVec.ofNat 64 (16 * k) := + addr_add (by have := hp.data_fit; have := (Dp s₀).isLt; omega) + +theorem UPre.dataN {k : Nat} (hk : k < N s₀) : + (Dp s₀ + BitVec.ofNat 32 (16 * k)).toNat = (Dp s₀).toNat + 16 * k := by + have := hp.data_fit + rw [BitVec.toNat_add, BitVec.toNat_ofNat, Nat.mod_eq_of_lt (a := 16 * k) (by omega), + Nat.mod_eq_of_lt (by omega)] + +theorem UPre.scrN {d : Nat} (hd : d < 2176) : (S s₀ + BitVec.ofNat 32 d).toNat = (S s₀).toNat + d := by + have := hp.scr_fit + rw [BitVec.toNat_add, BitVec.toNat_ofNat, Nat.mod_eq_of_lt (a := d) (by omega), Nat.mod_eq_of_lt (by omega)] + +omit hp in +theorem UPre.scr_sub {d n : Nat} (h : d + n ≤ 2176) : + Region.Sub ⟨State.addr (S s₀) + BitVec.ofNat 64 d, n⟩ (scrR s₀) := + Offset.sub_base _ h + +omit hp in +theorem UPre.data_sub {k : Nat} (hk : k < N s₀) : + Region.Sub ⟨State.addr (Dp s₀) + BitVec.ofNat 64 (16 * k), 16⟩ (dataR s₀) := + Offset.sub_base _ (by omega) + +theorem UPre.arg1 : stackArgAddr s₀ 1 = stackArgAddr s₀ 0 + BitVec.ofNat 64 4 := by + have := hp.sp_fit + simp only [stackArgAddr] + rw [addr_add (by omega), addr_add (by omega)] + simp + +theorem UPre.arg_in {k : Nat} (hk : k < 2) : InRegions (s₀.rd ++ s₀.wr) (stackArgAddr s₀ k) 4 := by + refine ⟨argsR s₀, by simp [hp.rd], ?_⟩ + rcases (by omega : k = 0 ∨ k = 1) with rfl | rfl + · simpa using Offset.contains_base (stackArgAddr s₀ 0) (d := 0) (n := 4) (k := 8) (by decide) (by decide) + · rw [hp.arg1]; exact Offset.contains_base _ (by decide) (by decide) + +omit hp in +theorem UPre.arg_sub : Region.Sub ⟨stackArgAddr s₀ 0, 4⟩ (argsR s₀) := Region.sub_prefix (by decide) + +end + +/-! ## Saving the registers -/ + +theorem saved_bound : ∀ p ∈ saved, 2064 ≤ p.2 ∧ p.2 + 4 ≤ 2096 := by decide + +theorem saved_ne_r12 : ∀ p ∈ saved, p.1 ≠ .r12 := by decide + +theorem saveMem_congr (m : Mem) (B : Addr) {g g' : Reg → BitVec 32} : + ∀ (l : List (Reg × Nat)), (∀ p ∈ l, g p.1 = g' p.1) → saveMem m B g l = saveMem m B g' l := by + intro l + induction l generalizing m with + | nil => intro _; rfl + | cons p l ih => + intro h + simp only [saveMem] + rw [h p (List.mem_cons_self ..)] + exact ih _ fun q hq => h q (List.mem_cons_of_mem _ hq) + +theorem savedMem_frame (s₀ : State) : Frame [⟨State.addr (S s₀), 2096⟩] s₀.mem (savedMem s₀) := + saveMem_frame _ _ _ (by decide) saved fun p hp => (saved_bound p hp).2 + +set_option simprocs false in +/-- Each slot holds the register saved there. -/ +theorem savedMem_slot (s₀ : State) {r : Reg} {d : Nat} (h : (r, d) ∈ saved) : + (savedMem s₀).readW (State.addr (S s₀) + BitVec.ofNat 64 d) 32 = s₀.gpr r := by + simp only [saved, List.mem_cons, List.not_mem_nil, or_false, Prod.mk.injEq] at h + rcases h with ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | + ⟨rfl, rfl⟩ <;> + simp (disch := decide) only [savedMem, saved, saveMem, Mem.readW_writeW_self32, readW_writeW_save] + +/-! ## The prologue -/ + +theorem prologue_wp {s₀ : State} (hp : UPre s₀) : + WP isa (.block (save ++ setup)) s₀ fun s => LInv s₀ 0 s ∧ s.z = decide (N s₀ = 0) := by + have hsc := hp.scr_fit + rw [show save ++ setup = .ldrSp .r12 4 :: (saved.map (fun p => Instr.str p.1 .r12 p.2) ++ setup) from rfl] + refine wp_ldrSp (a := stackArgAddr s₀ 1) (by decide) rfl (hp.arg_in (by decide)) fun s₁ u₁ => ?_ + have h12 : s₁.gpr .r12 = S s₀ := u₁.gpr + refine saveList_ok saved s₁ _ (fun p hp' => ?_) fun s₂ g₂ rd₂ wr₂ sp₂ m₂ => ?_ + · have hb := saved_bound p hp' + rw [h12, u₁.wr, hp.wr] + exact ⟨by omega, by omega, ⟨scrR s₀, by simp, Offset.contains_base _ (by omega) (by omega)⟩⟩ + have hm₂ : s₂.mem = savedMem s₀ := by + rw [m₂, u₁.mem, h12, savedMem] + exact saveMem_congr _ _ _ fun p hp' => u₁.other _ (saved_ne_r12 p hp') + have harg : s₂.mem.readW (stackArgAddr s₀ 0) 32 = stackArg s₀ 0 := by + rw [hm₂] + exact (savedMem_frame s₀).readW (Region.contains_self _ _) (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact (hp.scr_args.symm.sub_left UPre.arg_sub).sub_right (Region.sub_prefix (by decide))) (by decide) + simp only [setup, mov] + refine wp_mov (op2_reg _ _) fun s₃ u₃ => wp_mov (op2_reg _ _) fun s₄ u₄ => wp_mov (op2_reg _ _) fun s₅ u₅ => + wp_mov (op2_reg _ _) fun s₆ u₆ => ?_ + refine wp_ldrSp (a := stackArgAddr s₀ 0) (by decide) (by rw [u₆.sp, u₅.sp, u₄.sp, u₃.sp, sp₂, u₁.sp]; rfl) + (by rw [u₆.rd, u₆.wr, u₅.rd, u₅.wr, u₄.rd, u₄.wr, u₃.rd, u₃.wr, rd₂, wr₂, u₁.rd, u₁.wr] + exact hp.arg_in (by decide)) fun s₇ u₇ => ?_ + refine wp_mov (op2_reg _ _) fun s₈ u₈ => wp_cmp (op2_imm (by decide)) fun s₉ f₉ z₉ => WP.block_nil ?_ + have r8 : s₉.gpr .r8 = stackArg s₀ 0 := by + rw [f₉.gpr, u₈.other _ (by decide), u₇.gpr, u₆.mem, u₅.mem, u₄.mem, u₃.mem, harg] + have mm : s₉.mem = savedMem s₀ := by + rw [f₉.mem, u₈.mem, u₇.mem, u₆.mem, u₅.mem, u₄.mem, u₃.mem, hm₂] + have a0 : stackArg s₀ 0 = BitVec.ofNat 32 (N s₀) := by simp [N] + have stS : Spec.Aes.bytesAt (savedMem s₀) (State.addr (St s₀)) 16 = + Spec.Aes.bytesAt s₀.mem (State.addr (St s₀)) 16 := + Proof.Cmac.bytesAt_frame16 (savedMem_frame s₀) fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact hp.st_scr.sub_right (Region.sub_prefix (by decide)) + refine ⟨⟨?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_⟩, ?_⟩ + · simp (disch := decide) only [f₉.gpr, u₈.other, u₇.other, u₆.other, u₅.other, u₄.other, u₃.gpr, g₂, u₁.other] + · simp (disch := decide) only [f₉.gpr, u₈.other, u₇.other, u₆.other, u₅.other, u₄.gpr, u₃.other, g₂, u₁.other] + · simp (disch := decide) only [f₉.gpr, u₈.other, u₇.other, u₆.other, u₅.gpr, u₄.other, u₃.other, g₂, u₁.other] + · simp (disch := decide) only [f₉.gpr, u₈.other, u₇.other, u₆.gpr, u₅.other, u₄.other, u₃.other, g₂, u₁.other] + exact (BitVec.add_zero _).symm + · rw [r8, a0]; rfl + · simp (disch := decide) only [f₉.gpr, u₈.gpr, u₇.other, u₆.other, u₅.other, u₄.other, u₃.other, g₂, h12] + · simp (disch := decide) only [f₉.gpr, u₈.other, u₇.other, u₆.other, u₅.other, u₄.other, u₃.other, g₂, + u₁.other] + · rw [f₉.sp, u₈.sp, u₇.sp, u₆.sp, u₅.sp, u₄.sp, u₃.sp, sp₂, u₁.sp] + · rw [f₉.rd, u₈.rd, u₇.rd, u₆.rd, u₅.rd, u₄.rd, u₃.rd, rd₂, u₁.rd] + · rw [f₉.wr, u₈.wr, u₇.wr, u₆.wr, u₅.wr, u₄.wr, u₃.wr, wr₂, u₁.wr] + · rw [mm]; exact Frame.refl _ _ + · rw [mm, stS]; rfl + · rw [z₉, show s₈.gpr .r8 = s₉.gpr .r8 from (congrFun f₉.gpr _).symm, r8, a0] + exact cmp0 (stackArg s₀ 0).isLt + +end VG.Proof.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/CmacAes/Arm/UpdateCT.lean b/lean/VerifiedGarbage/Proof/CmacAes/Arm/UpdateCT.lean new file mode 100644 index 000000000..07297c51a --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/Arm/UpdateCT.lean @@ -0,0 +1,211 @@ +import VerifiedGarbage.Proof.CmacAes.Arm.UpdateCorrect +import VerifiedGarbage.Proof.Framework.Arm.ArgTaint + +/-! +# AES-CMAC on ARMv7: `vg_cmac_aes_update` is constant time + +Untrusted: everything here is checked by Lean. The taint analysis does not +analyse frames, so two runs from states that agree on the public arguments +are related piece by piece (`RelCT`): the taint analysis covers the code +between the calls, from the public arguments for the prologue and from the +registers the correctness proof pins to them (`LInv`) afterwards, and each +call of `vg_aes_ctr32`, in its frame, is constant time by its own proof +(`ctr_rel`). +-/ + +namespace VG.Proof.CmacAes.Arm + +open VG VG.Arm VG.Impl.CmacAes.Arm +open VG.Proof.MdStream.Arm (eval_eq eval_ne) + +/-- The registers holding our variables in the loop. -/ +abbrev vars : List Reg := [.r4, .r5, .r6, .r7, .r8, .r10] + +section +variable {s₀ s₀' : State} (hq : updateArm.pub s₀ s₀') +include hq + +theorem pub_sp : s₀.sp = s₀'.sp := hq.1 +theorem pub_W : W s₀ = W s₀' := hq.2.1 +theorem pub_r1 : s₀.gpr .r1 = s₀'.gpr .r1 := hq.2.2.1 +theorem pub_R : R s₀ = R s₀' := by rw [R, R, pub_r1 hq] +theorem pub_St : St s₀ = St s₀' := hq.2.2.2.1 +theorem pub_Dp : Dp s₀ = Dp s₀' := hq.2.2.2.2.1 +theorem pub_N : N s₀ = N s₀' := by rw [N, N, hq.2.2.2.2.2.1] +theorem pub_S : S s₀ = S s₀' := hq.2.2.2.2.2.2 +theorem pub_Cb : Cb s₀ = Cb s₀' := by rw [Cb, Cb, pub_S hq] + +/-- The registers the invariant pins agree in both runs. -/ +theorem LInv.agree {k : Nat} {s₁ s₂ : State} (h₁ : LInv s₀ k s₁) (h₂ : LInv s₀' k s₂) : + ∀ r ∈ vars, s₁.gpr r = s₂.gpr r := by + intro r hr + simp only [vars, List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl | rfl | rfl + · rw [h₁.r4, h₂.r4, pub_W hq] + · rw [h₁.r5, h₂.r5, pub_r1 hq] + · rw [h₁.r6, h₂.r6, pub_St hq] + · rw [h₁.r7, h₂.r7, pub_Dp hq] + · rw [h₁.r8, h₂.r8, pub_N hq] + · rw [h₁.r10, h₂.r10, pub_S hq] + +end + +/-! ## One block -/ + +/-- What is known between the code before the call and the call. -/ +structure Mid (s₀ : State) (k : Nat) (s : State) : Prop where + pre : CallPre s (W s₀) (Cb s₀) (St s₀) (S s₀) (R s₀) .r9 .r10 + r7 : s.gpr .r7 = Dp s₀ + BitVec.ofNat 32 (16 * k) + r8 : s.gpr .r8 = BitVec.ofNat 32 (N s₀ - k) + sp : s.sp = s₀.sp + +theorem bodyMid_wp {s₀ : State} (hp : UPre s₀) {k : Nat} (hk : k < N s₀) {s : State} (h : LInv s₀ k s) : + WP isa (.block (chainIn ++ updArgs)) s (Mid s₀ k) := + WP.mono (bodyA_wp hp hk h) fun _ hb => + ⟨hb.pre, by rw [hb.keep .r7 (by decide) (by decide) (by decide) (by decide) (by decide), h.r7], + by rw [hb.keep .r8 (by decide) (by decide) (by decide) (by decide) (by decide), h.r8], + by rw [hb.sp, h.sp]⟩ + +/-- What is known after the call. -/ +structure After (s₀ : State) (k : Nat) (s : State) : Prop where + r7 : s.gpr .r7 = Dp s₀ + BitVec.ofNat 32 (16 * k) + r8 : s.gpr .r8 = BitVec.ofNat 32 (N s₀ - k) + +theorem call_after {s₀ : State} {k : Nat} {s : State} (h : Mid s₀ k s) : + WP isa (ctrCall .r9 .r10) s (After s₀ k) := + WP.mono (ctr_call h.pre) fun _ hc => + ⟨by rw [hc.saved .r7 (by simp [preserved]) (by decide), h.r7], + by rw [hc.saved .r8 (by simp [preserved]) (by decide), h.r8]⟩ + +/-- The relation before a block, in two runs. -/ +def BRel (s₀ s₀' : State) (k : Nat) (s₁ s₂ : State) : Prop := + (k < N s₀ ∧ LInv s₀ k s₁) ∧ (k < N s₀' ∧ LInv s₀' k s₂) + +theorem body_ct {s₀ s₀' : State} (hp : UPre s₀) (hp' : UPre s₀') (hq : updateArm.pub s₀ s₀') (k : Nat) : + RelCT isa (BRel s₀ s₀' k) body fun _ _ => True := by + obtain ⟨_, hB⟩ : ∃ h, (taint.check (Taint.ofRegs [.r7, .r8]) (.block advance) h).isSome = true := + ⟨_, by taint_decide⟩ + have a := rel_agree (F := fun s => k < N s₀ ∧ LInv s₀ k s) (F' := fun s => k < N s₀' ∧ LInv s₀' k s) + (G := Mid s₀ k) (G' := Mid s₀' k) (Taint.ofRegs vars) + (fun _ _ h h' => Taint.agree_ofRegs (LInv.agree hq h.2 h'.2)) ⟨_, by taint_decide⟩ + (fun _ h => bodyMid_wp hp h.1 h.2) (fun _ h => bodyMid_wp hp' h.1 h.2) + have c := rel_wp (F := Mid s₀ k) (F' := Mid s₀' k) (G := After s₀ k) (G' := After s₀' k) + (ctr_rel (sp₀ := s₀.sp) fun s₁ s₂ h => + ⟨h.1.pre, by rw [pub_W hq, pub_Cb hq, pub_St hq, pub_S hq, pub_R hq]; exact h.2.pre, h.1.sp, + h.2.sp.trans (pub_sp hq).symm⟩) + (fun _ h => call_after h) (fun _ h => call_after h) + have b := RelCT.taint (A := taint) (P := fun s₁ s₂ => After s₀ k s₁ ∧ After s₀' k s₂) (Taint.ofRegs [.r7, .r8]) + (fun s₁ s₂ h => Taint.agree_ofRegs fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl + · rw [h.1.r7, h.2.r7, pub_Dp hq] + · rw [h.1.r8, h.2.r8, pub_N hq]) hB + exact a.seq (c.seq b) + +/-! ## The loop -/ + +/-- The loop's relation, with the number of iterations left. -/ +def LRel (s₀ s₀' : State) (n : Nat) (s₁ s₂ : State) : Prop := + ∃ k, n = N s₀ - k ∧ BRel s₀ s₀' k s₁ s₂ + +theorem loop_ct {s₀ s₀' : State} (hp : UPre s₀) (hp' : UPre s₀') (hq : updateArm.pub s₀ s₀') (n : Nat) : + RelCT isa (LRel s₀ s₀' n) (.loop body .ne) fun s₁ s₂ => LInv s₀ (N s₀) s₁ ∧ LInv s₀' (N s₀') s₂ := by + refine RelCT.loop (M := isa) (LRel s₀ s₀') (fun n => ?_) n + have hN := pub_N hq + refine (RelCT.exists_ fun k => ?_).mono (fun s₁ s₂ (h : LRel s₀ s₀' n s₁ s₂) => h) fun _ _ h => h + by_cases hn : n = N s₀ - k + swap + · exact RelCT.of_false fun _ _ h => hn h.1 + subst hn + have ct := (body_ct hp hp' hq k).wp + (F₁ := fun (s : State) => (LInv s₀ (k + 1) s ∧ s.z = decide (N s₀ - (k + 1) = 0)) ∧ k < N s₀) + (F₂ := fun (s : State) => LInv s₀' (k + 1) s ∧ s.z = decide (N s₀' - (k + 1) = 0)) + fun _ _ h => ⟨WP.mono (body_ok hp h.1.1 h.1.2) fun _ r => ⟨r, h.1.1⟩, body_ok hp' h.2.1 h.2.2⟩ + refine ct.mono (fun _ _ h => h.2) fun s₁ s₂ ⟨_, ⟨⟨l₁, z₁⟩, hk⟩, ⟨l₂, z₂⟩⟩ => ?_ + have e₁ : isa.eval .ne s₁ = some !decide (N s₀ - (k + 1) = 0) := by + show VG.Arm.eval .ne s₁ = _; rw [eval_ne, z₁] + have e₂ : isa.eval .ne s₂ = some !decide (N s₀ - (k + 1) = 0) := by + show VG.Arm.eval .ne s₂ = _; rw [eval_ne, z₂, ← hN] + refine ⟨by rw [e₁, e₂], fun hf => ?_, fun ht => ?_⟩ + · rw [e₁] at hf + have h0 : N s₀ = k + 1 := by + have : N s₀ - (k + 1) = 0 := by simpa using hf + omega + exact ⟨h0 ▸ l₁, by rw [← hN, h0]; exact l₂⟩ + · rw [e₁] at ht + have h0 : N s₀ - (k + 1) ≠ 0 := by simpa using ht + exact ⟨N s₀ - (k + 1), by omega, k + 1, rfl, ⟨by omega, l₁⟩, ⟨by omega, l₂⟩⟩ + +/-! ## The whole function -/ + +theorem update_rel {s₀ s₀' : State} (h0 : updateArm.pre s₀) (h0' : updateArm.pre s₀') + (hq : updateArm.pub s₀ s₀') : + RelCT isa (fun a b => a = s₀ ∧ b = s₀') update fun _ _ => True := by + have hp := UPre.of h0 + have hp' := UPre.of h0' + obtain ⟨_, hpro⟩ : ∃ h, (taint.check (argTaint [.r0, .r1, .r2, .r3] 8) (.block (save ++ setup)) h).isSome = + true := ⟨_, by taint_decide⟩ + obtain ⟨_, hepi⟩ : ∃ h, (taint.check (Taint.ofRegs [.r10]) (.block restore) h).isSome = true := + ⟨_, by taint_decide⟩ + obtain ⟨_, hnil⟩ : ∃ h, (taint.check (Taint.ofRegs []) (.block []) h).isSome = true := + ⟨_, by taint_decide⟩ + have hN := pub_N hq + have hsp := pub_sp hq + have wfA : ∀ {s : State}, UPre s → + s.sp.toNat + 8 ≤ 2 ^ 32 ∧ ∀ r ∈ s.wr, Region.Disjoint ⟨State.addr s.sp, 8⟩ r := fun {s} h => by + have e : (⟨State.addr s.sp, 8⟩ : Region) = argsR s := by simp [stackArgAddr] + refine ⟨h.sp_fit, ?_⟩ + rw [e, h.wr] + simp only [List.mem_cons, List.not_mem_nil, or_false] + rintro r (rfl | rfl) + · exact h.st_args.symm + · exact h.scr_args.symm + have pro := rel_agree (F := fun s => s = s₀) (F' := fun s => s = s₀') + (G := fun s => LInv s₀ 0 s ∧ s.z = decide (N s₀ = 0)) (G' := fun s => LInv s₀' 0 s ∧ s.z = decide (N s₀' = 0)) + (argTaint [.r0, .r1, .r2, .r3] 8) + (fun s s' e e' => by + subst e e' + refine agree_argTaint (fun r hr => ?_) hsp (wfA hp) (wfA hp') + (argMem_of (j := 2) hsp hp.sp_fit fun i hi => ?_) + · simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl + · exact hq.2.1 + · exact hq.2.2.1 + · exact hq.2.2.2.1 + · exact hq.2.2.2.2.1 + · rcases (by omega : i = 0 ∨ i = 1) with rfl | rfl + · exact hq.2.2.2.2.2.1 + · exact hq.2.2.2.2.2.2) ⟨_, hpro⟩ + (fun s e => by rw [e]; exact prologue_wp hp) (fun s e => by rw [e]; exact prologue_wp hp') + have ev {s : State} (h : s.z = decide (N s₀ = 0)) : isa.eval .eq s = some (decide (N s₀ = 0)) := by + show VG.Arm.eval .eq s = _; rw [eval_eq, h] + have ev' {s : State} (h : s.z = decide (N s₀' = 0)) : isa.eval .eq s = some (decide (N s₀ = 0)) := by + show VG.Arm.eval .eq s = _; rw [eval_eq, h, hN] + have nil := RelCT.taint (A := taint) + (P := fun a b => ((LInv s₀ 0 a ∧ a.z = decide (N s₀ = 0)) ∧ (LInv s₀' 0 b ∧ b.z = decide (N s₀' = 0))) ∧ + isa.eval .eq a = some true) (Taint.ofRegs []) (fun _ _ _ => Taint.agree_ofRegs fun r hr => by simp at hr) + hnil + have mid : RelCT isa (fun a b => (LInv s₀ 0 a ∧ a.z = decide (N s₀ = 0)) ∧ (LInv s₀' 0 b ∧ b.z = decide (N s₀' = 0))) + (.ite .eq (.block []) (.loop body .ne)) (fun a b => LInv s₀ (N s₀) a ∧ LInv s₀' (N s₀') b) := by + refine RelCT.ite (fun a b h => by rw [ev h.1.2, ev' h.2.2]) ?_ ?_ + · refine (nil.wp (F₁ := LInv s₀ (N s₀)) (F₂ := LInv s₀' (N s₀')) fun a b h => ?_).mono + (fun _ _ h => h) fun _ _ h => h.2 + have h0 : N s₀ = 0 := by + have := h.2; rw [ev h.1.1.2] at this; simpa using this + exact ⟨WP.block_nil (h0 ▸ h.1.1.1), WP.block_nil (by rw [← hN, h0]; exact h.1.2.1)⟩ + · refine (loop_ct hp hp' hq (N s₀ - 0)).mono (fun a b h => ⟨0, rfl, ⟨?_, h.1.1.1⟩, ⟨?_, h.1.2.1⟩⟩) + fun _ _ h => h + all_goals + have := h.2; rw [ev h.1.1.2] at this + have : N s₀ ≠ 0 := by simpa using this + omega + have epi := RelCT.taint (A := taint) (P := fun a b => LInv s₀ (N s₀) a ∧ LInv s₀' (N s₀') b) + (Taint.ofRegs [.r10]) (fun a b h => Taint.agree_ofRegs fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + subst hr; rw [h.1.r10, h.2.r10, pub_S hq]) hepi + exact (pro.mono (fun _ _ h => h) fun _ _ h => h).seq (mid.seq epi) + +theorem update_ct : ConstantTime isa updateArm.pre updateArm.pub update := + fun _ _ _ _ _ _ h₁ h₂ hq e₁ e₂ => (update_rel h₁ h₂ hq _ _ _ _ _ _ ⟨rfl, rfl⟩ e₁ e₂).1 + +end VG.Proof.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/CmacAes/Arm/UpdateCorrect.lean b/lean/VerifiedGarbage/Proof/CmacAes/Arm/UpdateCorrect.lean new file mode 100644 index 000000000..0a2880100 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/Arm/UpdateCorrect.lean @@ -0,0 +1,112 @@ +import VerifiedGarbage.Proof.CmacAes.Arm.UpdateLoop + +/-! +# AES-CMAC on ARMv7: `vg_cmac_aes_update` is correct + +Untrusted: everything here is checked by Lean. +-/ + +namespace VG.Proof.CmacAes.Arm + +open VG VG.Arm VG.Impl.CmacAes.Arm +open VG.Proof.MdStream.Arm (Upd wp_ldr eval_eq) + +/-! ## Restoring the registers -/ + +/-- Loads of the registers `l` from `b + offset`, none of them `b`. -/ +theorem restoreB_ok {b : Reg} {rest : List Instr} (l : List (Reg × Nat)) : + ∀ (s : State) (Q : State → Prop), (l.map Prod.fst).Nodup → + (∀ p ∈ l, p.1 ≠ b ∧ p.2 < 4096 ∧ (s.gpr b).toNat + p.2 < 2 ^ 32 ∧ + InRegions (s.rd ++ s.wr) (State.addr (s.gpr b) + BitVec.ofNat 64 p.2) 4) → + (∀ s', (∀ p ∈ l, s'.gpr p.1 = s.mem.readW (State.addr (s.gpr b) + BitVec.ofNat 64 p.2) 32) → + (∀ r, r ∉ l.map Prod.fst → s'.gpr r = s.gpr r) → s'.mem = s.mem → s'.rd = s.rd → s'.wr = s.wr → + s'.sp = s.sp → WP isa (.block rest) s' Q) → + WP isa (.block (l.map (fun p => Instr.ldr p.1 b p.2) ++ rest)) s Q := by + induction l with + | nil => intro s Q _ _ k; exact k s (fun _ h => by cases h) (fun _ _ => rfl) rfl rfl rfl rfl + | cons p l ih => + intro s Q hnd hl k + obtain ⟨h0, h1, h2, h3⟩ := hl p (by simp) + simp only [List.map_cons, List.nodup_cons] at hnd + refine wp_ldr h1 (addr_add h2) h3 fun s₁ u₁ => ?_ + have eb : s₁.gpr b = s.gpr b := u₁.other _ (Ne.symm h0) + refine ih s₁ Q hnd.2 (fun q hq => ?_) fun s' hl' ho hm hrd hwr hsp => k s' (fun q hq => ?_) + (fun r hr => ?_) (hm.trans u₁.mem) (hrd.trans u₁.rd) (hwr.trans u₁.wr) (hsp.trans u₁.sp) + · rw [eb, u₁.rd, u₁.wr]; exact hl q (List.mem_cons_of_mem _ hq) + · rcases List.mem_cons.mp hq with rfl | hq + · rw [ho _ hnd.1, u₁.gpr] + · rw [hl' q hq, u₁.mem, eb] + · simp only [List.map_cons, List.mem_cons, not_or] at hr + rw [ho r hr.2, u₁.other r hr.1] + +theorem restore_eq : restore = (saved.take 7).map (fun p => Instr.ldr p.1 .r10 p.2) ++ + ([.ldr .r10 .r10 2088] : List Instr) := rfl + +theorem take7_ne : ∀ p ∈ saved.take 7, p.1 ≠ .r10 := by decide + +theorem slot_read {s₀ : State} (hp : UPre s₀) {m : Mem} + (hf : Frame [stR s₀, ⟨State.addr (S s₀), 2064⟩, belowR s₀] (savedMem s₀) m) {d : Nat} (h₁ : 2064 ≤ d) + (h₂ : d + 4 ≤ 2096) : + m.readW (State.addr (S s₀) + BitVec.ofNat 64 d) 32 = (savedMem s₀).readW (State.addr (S s₀) + BitVec.ofNat 64 d) 32 := + hf.readW (r := ⟨State.addr (S s₀) + BitVec.ofNat 64 d, 4⟩) (Region.contains_self _ _) (fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl + · exact hp.st_scr.symm.sub_left (UPre.scr_sub (by omega)) + · exact Offset.disjoint_base _ h₁ (by omega) + · exact hp.b_scr.symm.sub_left (UPre.scr_sub (by omega))) (by decide) + +theorem epilogue_wp {s₀ : State} (hp : UPre s₀) {s : State} (h : LInv s₀ (N s₀) s) : + WP isa (.block restore) s fun s' => abiPreserved s₀ s' ∧ updateArm.post s₀ s' := by + have hsc := hp.scr_fit + have rdwr : s.rd ++ s.wr = [schR s₀, dataR s₀, argsR s₀, stR s₀, scrR s₀] := by + rw [h.rd, h.wr, hp.rd, hp.wr]; rfl + have inS : ∀ d, d + 4 ≤ 2176 → InRegions (s.rd ++ s.wr) (State.addr (S s₀) + BitVec.ofNat 64 d) 4 := + fun d hd => by rw [rdwr]; exact ⟨scrR s₀, by simp, Offset.contains_base _ hd (by omega)⟩ + have sl : ∀ r d, (r, d) ∈ saved → s.mem.readW (State.addr (S s₀) + BitVec.ofNat 64 d) 32 = s₀.gpr r := + fun r d hrd => by + have hb := saved_bound _ hrd + rw [slot_read hp h.frame hb.1 hb.2, savedMem_slot s₀ hrd] + rw [restore_eq] + refine restoreB_ok (saved.take 7) s _ (by decide) (fun p hp' => ?_) fun s₁ ld₁ ho₁ m₁ rd₁ wr₁ sp₁ => ?_ + · have hb := saved_bound p (List.mem_of_mem_take hp') + exact ⟨take7_ne p hp', by omega, by rw [h.r10]; omega, by rw [h.r10]; exact inS _ (by omega)⟩ + refine wp_ldr (a := State.addr (S s₀) + BitVec.ofNat 64 2088) (by decide) + (by rw [ho₁ _ (by decide), h.r10]; exact hp.scrA (by decide)) + (by rw [rd₁, wr₁]; exact inS _ (by decide)) fun s₂ u₂ => WP.block_nil ⟨⟨fun r hr => ?_, ?_⟩, ?_⟩ + · have ld : ∀ r d, (r, d) ∈ saved.take 7 → s₂.gpr r = s₀.gpr r := fun r d hrd => by + have hne : r ≠ .r10 := take7_ne _ hrd + rw [u₂.other _ hne, ld₁ _ hrd, h.r10, sl r d (List.mem_of_mem_take hrd)] + simp only [preserved, List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl + · exact ld _ 2064 (by decide) + · exact ld _ 2068 (by decide) + · exact ld _ 2072 (by decide) + · exact ld _ 2076 (by decide) + · exact ld _ 2080 (by decide) + · exact ld _ 2084 (by decide) + · rw [u₂.gpr, m₁, sl .r10 2088 (by decide)] + · rw [u₂.other _ (by decide), ho₁ _ (by decide), h.r11] + · exact ld _ 2092 (by decide) + · rw [u₂.sp, sp₁, h.sp] + · show Spec.Aes.bytesAt s₂.mem (State.addr (St s₀)) 16 = Spec.Cmac.chain (ciph s₀) _ (blks s₀) + rw [u₂.mem, m₁, h.state, List.take_of_length_le (by simp [Spec.Cmac.blocksAt])] + +/-! ## The whole function -/ + +theorem mid_wp {s₀ : State} (hp : UPre s₀) {s₁ : State} (h : LInv s₀ 0 s₁) (hz : s₁.z = decide (N s₀ = 0)) : + WP isa (.ite .eq (.block []) (.loop body .ne)) s₁ (LInv s₀ (N s₀)) := by + have ev : isa.eval .eq s₁ = some (decide (N s₀ = 0)) := by + show VG.Arm.eval .eq s₁ = _; rw [eval_eq, hz] + by_cases hn : N s₀ = 0 + · refine WP.ite true (by rw [ev]; simp [hn]) (fun _ => WP.block_nil ?_) (fun h => by cases h) + rw [hn]; exact h + · refine WP.ite false (by rw [ev]; simp [hn]) (fun h => by cases h) fun _ => ?_ + exact loop_ok hp (by omega) h + +theorem update_wp {s₀ : State} (h0 : updateArm.pre s₀) : + WP isa update s₀ fun s' => abiPreserved s₀ s' ∧ updateArm.post s₀ s' := by + have hp := UPre.of h0 + exact WP.seq (WP.mono (prologue_wp hp) fun s₁ ⟨h₁, hz⟩ => + WP.seq (WP.mono (mid_wp hp h₁ hz) fun _ h₂ => epilogue_wp hp h₂)) + +end VG.Proof.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/CmacAes/Arm/UpdateLoop.lean b/lean/VerifiedGarbage/Proof/CmacAes/Arm/UpdateLoop.lean new file mode 100644 index 000000000..81c8098b5 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/Arm/UpdateLoop.lean @@ -0,0 +1,259 @@ +import VerifiedGarbage.Proof.CmacAes.Arm.Update + +/-! +# AES-CMAC on ARMv7: the loop of `vg_cmac_aes_update` + +Untrusted: everything here is checked by Lean. One block keeps the loop +invariant (`body_ok`): the counter block is `C ⊕ Mᵢ` and the state is +zeroed (`Cmac.chainMem4`), and the call of `vg_aes_ctr32` leaves +`CIPH_K(C ⊕ Mᵢ)` in the state. +-/ + +namespace VG.Proof.CmacAes.Arm + +open VG VG.Arm VG.Impl.CmacAes.Arm +open VG.Proof.MdStream.Arm (Upd Mupd Fupd op2_imm op2_reg wp_mov wp_add wp_subs eval_ne ofNat_beq_zero sub_ofNat) + +theorem take_succ_blks (s₀ : State) {k : Nat} (hk : k < N s₀) : + (blks s₀).take (k + 1) = + (blks s₀).take k ++ [Spec.Aes.bytesAt s₀.mem (State.addr (Dp s₀) + BitVec.ofNat 64 (16 * k)) 16] := by + rw [List.take_add_one, List.getElem?_eq_getElem (by simp [Spec.Cmac.blocksAt]; omega)] + simp [Spec.Cmac.blocksAt] + +/-! ## Memory outside the writable regions -/ + +/-- The regions the function writes, with the stack below it. -/ +abbrev Big (s₀ : State) : List Region := [stR s₀, scrR s₀, belowR s₀] + +section +variable {s₀ : State} (hp : UPre s₀) +include hp + +theorem UPre.sched_bytes {m : Mem} (hf : Frame (Big s₀) s₀.mem m) : + Spec.Aes.bytesAt m (State.addr (W s₀)) (16 * (R s₀ + 1)) = + Spec.Aes.bytesAt s₀.mem (State.addr (W s₀)) (16 * (R s₀ + 1)) := by + have hR : 16 * (R s₀ + 1) ≤ 240 := by rcases hp.rounds with h | h | h <;> omega + refine Proof.Cmac.bytesAt_frame hf (fun r hr => ?_) (by omega) + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl + · exact hp.sch_st.sub_left (Region.sub_prefix hR) + · exact hp.sch_scr.sub_left (Region.sub_prefix hR) + · exact hp.b_sch.symm.sub_left (Region.sub_prefix hR) + +theorem UPre.block_bytes {m : Mem} (hf : Frame (Big s₀) s₀.mem m) {k : Nat} (hk : k < N s₀) : + Spec.Aes.bytesAt m (State.addr (Dp s₀) + BitVec.ofNat 64 (16 * k)) 16 = + Spec.Aes.bytesAt s₀.mem (State.addr (Dp s₀) + BitVec.ofNat 64 (16 * k)) 16 := by + refine Proof.Cmac.bytesAt_frame hf (fun r hr => ?_) (by decide) + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl + · exact hp.data_st.sub_left (UPre.data_sub hk) + · exact hp.data_scr.sub_left (UPre.data_sub hk) + · exact hp.b_data.symm.sub_left (UPre.data_sub hk) + +omit hp in +theorem UPre.big_of {m : Mem} (hf : Frame [stR s₀, ⟨State.addr (S s₀), 2064⟩, belowR s₀] (savedMem s₀) m) : + Frame (Big s₀) s₀.mem m := by + have f₀ : Frame (Big s₀) s₀.mem (savedMem s₀) := + (savedMem_frame s₀).sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact ⟨scrR s₀, by simp, Region.sub_prefix (by decide)⟩ + exact f₀.trans (hf.sub fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl + · exact ⟨stR s₀, by simp, fun _ h => h⟩ + · exact ⟨scrR s₀, by simp, Region.sub_prefix (by decide)⟩ + · exact ⟨belowR s₀, by simp, fun _ h => h⟩) + +end + +/-! ## One block -/ + +/-- The counter block's address. -/ +abbrev Cb (s₀ : State) : BitVec 32 := S s₀ + BitVec.ofNat 32 2048 + +/-- What the code before the call leaves. -/ +structure BodyA (s₀ : State) (k : Nat) (s s₁ : State) : Prop where + pre : CallPre s₁ (W s₀) (Cb s₀) (St s₀) (S s₀) (R s₀) .r9 .r10 + keep : ∀ r, r ≠ .r0 → r ≠ .r1 → r ≠ .r2 → r ≠ .r3 → r ≠ .r9 → s₁.gpr r = s.gpr r + sp : s₁.sp = s.sp + mem : s₁.mem = Proof.Cmac.chainMem4 s.mem (State.addr (S s₀) + BitVec.ofNat 64 2048) (State.addr (St s₀)) + (State.addr (Dp s₀) + BitVec.ofNat 64 (16 * k)) + rd : s₁.rd = s.rd + wr : s₁.wr = s.wr + +theorem chainIn_eq : chainIn ++ updArgs = xorBlk .r0 .r1 .r6 .r7 .r10 0 0 2048 ++ + (.mov .r0 (.imm 0) :: (zeroBlk .r0 .r6 0 ++ + ([.mov .r0 (.reg .r4), .mov .r1 (.reg .r5), .dp .add .r2 .r10 (.imm (BitVec.ofNat 32 2048)), + .mov .r3 (.reg .r6), .mov .r9 (.imm 1)] : List Instr))) := rfl + +theorem bodyA_wp {s₀ : State} (hp : UPre s₀) {k : Nat} (hk : k < N s₀) {s : State} (h : LInv s₀ k s) : + WP isa (.block (chainIn ++ updArgs)) s (BodyA s₀ k s) := by + have hRegs : s.rd ++ s.wr = [schR s₀, dataR s₀, argsR s₀, stR s₀, scrR s₀] := by + rw [h.rd, h.wr, hp.rd, hp.wr]; rfl + have hW : s.wr = [stR s₀, scrR s₀] := by rw [h.wr, hp.wr] + have hsc := hp.scr_fit + have hst := hp.st_fit + have hdf := hp.data_fit + have qN := hp.dataN hk + rw [chainIn_eq] + refine xorBlk_ok (by decide) (by decide) (by decide) (by decide) (by decide) (by decide) (by decide) + (by decide) (by decide) (by decide) (by rw [h.r6]; omega) (by rw [h.r7, qN]; omega) + (by rw [h.r10]; omega) ?_ ?_ ?_ fun s₁ g₁ => ?_ + · rw [h.r6, add0, hRegs] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact ⟨stR s₀, by simp, 0, by simp, by simp⟩ + · rw [h.r7, add0, hp.dataA hk, hRegs] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact ⟨dataR s₀, by simp, 16 * k, rfl, by simp; omega⟩ + · rw [h.r10, hW] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact ⟨scrR s₀, by simp, 2048, rfl, by simp⟩ + have e₁ : ∀ r, r ≠ .r0 → r ≠ .r1 → s₁.gpr r = s.gpr r := g₁.gpr + refine wp_mov (op2_imm (by decide)) fun s₂ u₂ => ?_ + have r6₂ : s₂.gpr .r6 = St s₀ := by rw [u₂.other _ (by decide), e₁ _ (by decide) (by decide), h.r6] + refine Proof.CmacAes.Arm.zeroBlk_ok u₂.gpr (by decide) (by rw [r6₂]; omega) ?_ fun s₃ G₃ m₃ rd₃ wr₃ sp₃ => ?_ + · rw [r6₂, add0, u₂.wr, g₁.wr, hW] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact ⟨stR s₀, by simp, 0, by simp, by simp⟩ + refine wp_mov (op2_reg _ _) fun s₄ u₄ => wp_mov (op2_reg _ _) fun s₅ u₅ => + wp_add (op2_imm (by decide)) fun s₆ u₆ => wp_mov (op2_reg _ _) fun s₇ u₇ => + wp_mov (op2_imm (by decide)) fun s₈ u₈ => WP.block_nil ?_ + have keep : ∀ r, r ≠ .r0 → r ≠ .r1 → r ≠ .r2 → r ≠ .r3 → r ≠ .r9 → s₈.gpr r = s.gpr r := + fun r h0 h1 h2 h3 h9 => by + rw [u₈.other _ h9, u₇.other _ h3, u₆.other _ h2, u₅.other _ h1, u₄.other _ h0, G₃, u₂.other _ h0, + e₁ _ h0 h1] + have sp₈ : s₈.sp = s₀.sp := by + rw [u₈.sp, u₇.sp, u₆.sp, u₅.sp, u₄.sp, sp₃, u₂.sp, g₁.sp, h.sp] + have rd₈ : s₈.rd = s.rd := by rw [u₈.rd, u₇.rd, u₆.rd, u₅.rd, u₄.rd, rd₃, u₂.rd, g₁.rd] + have wr₈ : s₈.wr = s.wr := by rw [u₈.wr, u₇.wr, u₆.wr, u₅.wr, u₄.wr, wr₃, u₂.wr, g₁.wr] + have mem₈ : s₈.mem = Proof.Cmac.chainMem4 s.mem (State.addr (S s₀) + BitVec.ofNat 64 2048) + (State.addr (St s₀)) (State.addr (Dp s₀) + BitVec.ofNat 64 (16 * k)) := by + rw [u₈.mem, u₇.mem, u₆.mem, u₅.mem, u₄.mem, m₃, r6₂, u₂.mem, g₁.mem, h.r6, h.r7, h.r10, add0, add0, + hp.dataA hk] + rfl + have hb : below s₈ = belowR s₀ := by rw [below, sp₈]; rfl + have cA : State.addr (Cb s₀) = State.addr (S s₀) + BitVec.ofNat 64 2048 := hp.scrA (by decide) + have cSt : (⟨State.addr (Cb s₀), 16⟩ : Region).Disjoint (stR s₀) := by + rw [cA]; exact hp.st_scr.symm.sub_left (UPre.scr_sub (by decide)) + refine ⟨⟨?_, ?_, ?_, ?_, u₈.gpr, ?_, by decide, hp.rounds, by rw [sp₈]; exact hp.sp8, ?_, hp.sch_st, ?_, cSt, + ?_, ?_, by rw [hb]; exact hp.b_sch, ?_, by rw [hb]; exact hp.b_st, ?_, hp.sch_fit, ?_, hp.st_fit, ?_, ?_, ?_, + ?_⟩, keep, by rw [sp₈, h.sp], mem₈, rd₈, wr₈⟩ + · rw [u₈.other _ (by decide), u₇.other _ (by decide), u₆.other _ (by decide), u₅.other _ (by decide), u₄.gpr, + G₃, u₂.other _ (by decide), e₁ _ (by decide) (by decide), h.r4] + · rw [u₈.other _ (by decide), u₇.other _ (by decide), u₆.other _ (by decide), u₅.gpr, u₄.other _ (by decide), + G₃, u₂.other _ (by decide), e₁ _ (by decide) (by decide), h.r5] + simp [R] + · rw [u₈.other _ (by decide), u₇.other _ (by decide), u₆.gpr, u₅.other _ (by decide), u₄.other _ (by decide), + G₃, u₂.other _ (by decide), e₁ _ (by decide) (by decide), h.r10] + · rw [u₈.other _ (by decide), u₇.gpr, u₆.other _ (by decide), u₅.other _ (by decide), u₄.other _ (by decide), + G₃, r6₂] + · rw [keep _ (by decide) (by decide) (by decide) (by decide) (by decide), h.r10] + · rw [cA]; exact hp.sch_scr.sub_right (UPre.scr_sub (by decide)) + · exact hp.sch_scr.sub_right (Region.sub_prefix (by decide)) + · rw [cA]; exact Offset.disjoint_base _ (by decide) (by omega) + · exact hp.st_scr.sub_right (Region.sub_prefix (by decide)) + · rw [hb, cA]; exact hp.b_scr.sub_right (UPre.scr_sub (by decide)) + · rw [hb]; exact hp.b_scr.sub_right (Region.sub_prefix (by decide)) + · rw [hp.scrN (by decide)]; omega + · omega + · rw [rd₈, wr₈, hRegs] + exact Covers.of_sub fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact ⟨schR s₀, by simp, 0, by simp, by simp⟩ + · rw [wr₈, hW, cA] + refine Covers.of_sub fun r hr => ?_ + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl + · exact ⟨scrR s₀, by simp, 2048, rfl, by simp⟩ + · exact ⟨stR s₀, by simp, 0, by simp, by simp⟩ + · exact ⟨scrR s₀, by simp, 0, by simp, by simp⟩ + · rw [mem₈]; exact Proof.Cmac.chainMem4_state _ _ _ _ + +theorem body_ok {s₀ : State} (hp : UPre s₀) {k : Nat} (hk : k < N s₀) {s : State} (h : LInv s₀ k s) : + WP isa body s fun s' => LInv s₀ (k + 1) s' ∧ s'.z = decide (N s₀ - (k + 1) = 0) := by + have hdf := hp.data_fit + have hN := (stackArg s₀ 0).isLt + refine WP.seq (WP.mono (bodyA_wp hp hk h) fun s₁ a => ?_) + refine WP.seq (WP.mono (ctr_call a.pre) fun s₂ h₂ => ?_) + refine wp_add (op2_imm (by decide)) fun s₃ u₃ => wp_subs (op2_imm (by decide)) fun s₄ u₄ z₄ => WP.block_nil ?_ + have g (r : Reg) (hr : r ∈ preserved) (hlr : r ≠ .lr) (h7 : r ≠ .r7) (h8 : r ≠ .r8) (h9 : r ≠ .r9) : + s₄.gpr r = s.gpr r := by + have : r ≠ .r0 ∧ r ≠ .r1 ∧ r ≠ .r2 ∧ r ≠ .r3 := by + simp only [preserved, List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl <;> decide + rw [u₄.other _ h8, u₃.other _ h7, h₂.saved r hr hlr, a.keep r this.1 this.2.1 this.2.2.1 this.2.2.2 h9] + have r7₂ : s₂.gpr .r7 = Dp s₀ + BitVec.ofNat 32 (16 * k) := by + rw [h₂.saved .r7 (by simp [preserved]) (by decide), a.keep _ (by decide) (by decide) (by decide) (by decide) + (by decide), h.r7] + have r8₃ : s₃.gpr .r8 = BitVec.ofNat 32 (N s₀ - k) := by + rw [u₃.other _ (by decide), h₂.saved .r8 (by simp [preserved]) (by decide), + a.keep _ (by decide) (by decide) (by decide) (by decide) (by decide), h.r8] + have dec : BitVec.ofNat 32 (N s₀ - k) - 1 = BitVec.ofNat 32 (N s₀ - (k + 1)) := by + rw [show (1 : BitVec 32) = BitVec.ofNat 32 1 from rfl, sub_ofNat (by omega)]; rfl + -- Memory. + have bigS := UPre.big_of h.frame + have cA : State.addr (Cb s₀) = State.addr (S s₀) + BitVec.ofNat 64 2048 := hp.scrA (by decide) + have f₁ : Frame [⟨State.addr (S s₀) + BitVec.ofNat 64 2048, 16⟩, stR s₀] s.mem s₁.mem := by + rw [a.mem]; exact Proof.Cmac.chainMem4_frame _ _ _ _ + have big₁ : Frame (Big s₀) s₀.mem s₁.mem := (UPre.big_of h.frame).trans (f₁.sub fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl + · exact ⟨scrR s₀, by simp, UPre.scr_sub (by decide)⟩ + · exact ⟨stR s₀, by simp, fun _ h => h⟩) + have cst : (⟨State.addr (S s₀) + BitVec.ofNat 64 2048, 16⟩ : Region).Disjoint (stR s₀) := + hp.st_scr.symm.sub_left (UPre.scr_sub (by decide)) + have cq : (⟨State.addr (S s₀) + BitVec.ofNat 64 2048, 16⟩ : Region).Disjoint + ⟨State.addr (Dp s₀) + BitVec.ofNat 64 (16 * k), 16⟩ := + (hp.data_scr.symm.sub_left (UPre.scr_sub (by decide))).sub_right (UPre.data_sub hk) + have out := h₂.out + rw [UPre.sched_bytes hp big₁, cA, a.mem, Proof.Cmac.chainMem4_counter _ cst cq, h.state, + UPre.block_bytes hp bigS hk] at out + refine ⟨⟨?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_⟩, ?_⟩ + · rw [g .r4 (by simp [preserved]) (by decide) (by decide) (by decide) (by decide), h.r4] + · rw [g .r5 (by simp [preserved]) (by decide) (by decide) (by decide) (by decide), h.r5] + · rw [g .r6 (by simp [preserved]) (by decide) (by decide) (by decide) (by decide), h.r6] + · rw [u₄.other _ (by decide), u₃.gpr, r7₂, show (16 : BitVec 32) = BitVec.ofNat 32 16 from rfl, + Offset.add_add_eq _ (c := 16 * (k + 1)) (by omega)] + · rw [u₄.gpr, r8₃, dec] + · rw [g .r10 (by simp [preserved]) (by decide) (by decide) (by decide) (by decide), h.r10] + · rw [g .r11 (by simp [preserved]) (by decide) (by decide) (by decide) (by decide), h.r11] + · rw [u₄.sp, u₃.sp, h₂.sp, a.sp, h.sp] + · rw [u₄.rd, u₃.rd, h₂.rd, a.rd, h.rd] + · rw [u₄.wr, u₃.wr, h₂.wr, a.wr, h.wr] + · have hb : below s₁ = belowR s₀ := by rw [below, a.sp, h.sp]; rfl + rw [u₄.mem, u₃.mem] + refine h.frame.trans ((f₁.sub fun r hr => ?_).trans (h₂.frame.sub fun r hr => ?_)) + · simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl + · exact ⟨⟨State.addr (S s₀), 2064⟩, by simp, Offset.sub_base _ (by decide)⟩ + · exact ⟨stR s₀, by simp, fun _ h => h⟩ + · simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl + · exact ⟨⟨State.addr (S s₀), 2064⟩, by simp, by rw [cA]; exact Offset.sub_base _ (by decide)⟩ + · exact ⟨stR s₀, by simp, fun _ h => h⟩ + · exact ⟨⟨State.addr (S s₀), 2064⟩, by simp, Region.sub_prefix (by decide)⟩ + · exact ⟨belowR s₀, by simp, by rw [hb]; exact fun _ h => h⟩ + · rw [u₄.mem, u₃.mem, out, take_succ_blks s₀ hk, Proof.Cmac.chain_append, Proof.Cmac.chain_single] + · rw [z₄, r8₃, dec]; exact ofNat_beq_zero (by omega) + +theorem loop_ok {s₀ : State} (hp : UPre s₀) {k : Nat} (hk : k < N s₀) {s : State} (h : LInv s₀ k s) : + WP isa (.loop body .ne) s (LInv s₀ (N s₀)) := by + refine WP.loop (M := isa) (body := body) (c := .ne) (Q := LInv s₀ (N s₀)) + (fun (n : Nat) (t : State) => ∃ j, n = N s₀ - j ∧ j < N s₀ ∧ LInv s₀ j t) ?_ (N s₀ - k) s + ⟨k, rfl, hk, h⟩ + rintro n s ⟨k, rfl, hk, h⟩ + refine WP.mono (body_ok hp hk h) fun s' ⟨h', hz⟩ => ?_ + have ev : isa.eval .ne s' = some !decide (N s₀ - (k + 1) = 0) := by + show VG.Arm.eval .ne s' = _; rw [eval_ne, hz] + by_cases hz' : N s₀ - (k + 1) = 0 + · left + refine ⟨by rw [ev]; simp [hz'], ?_⟩ + rwa [show N s₀ = k + 1 by omega] + · right + refine ⟨by rw [ev]; simp [hz'], N s₀ - (k + 1), by omega, k + 1, rfl, by omega, h'⟩ + +end VG.Proof.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/CmacAes/Arm/Verified.lean b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Verified.lean new file mode 100644 index 000000000..0a1a2f5f7 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Verified.lean @@ -0,0 +1,78 @@ +import VerifiedGarbage.Proof.CmacAes.Arm.UpdateCT +import VerifiedGarbage.Proof.CmacAes.Arm.SubkeysCT +import VerifiedGarbage.Proof.CmacAes.Arm.FinalizeCT +import VerifiedGarbage.Proof.Framework.Contract +import VerifiedGarbage.Spec.Cmac.Contract + +/-! +# AES-CMAC on ARMv7: `Verified` + +Untrusted: everything here is checked by Lean. Correctness and constant +time, a state satisfying each precondition, and the shared contracts of +`Spec/Cmac/Contract.lean`, with 8 bytes of stack: each call of +`vg_aes_ctr32` pushes its two stack arguments. +-/ + +namespace VG.Proof.CmacAes.Arm + +open VG VG.Arm VG.Impl.CmacAes.Arm + +/-- A state satisfying `vg_cmac_aes_update`'s precondition (with no blocks, +and the scratch buffer at 0, which the zero stack arguments point at). -/ +def updSat : State where + gpr r := match r with + | .r0 => 0x1000 | .r1 => 10 | .r2 => 0x2000 | .r3 => 0x3000 | _ => 0 + sp := 0x8000 + n := false + z := false + c := false + v := false + mem _ := 0 + rd := [⟨0x1000, 240⟩, ⟨0x3000, 0⟩, ⟨0x8000, 8⟩] + wr := [⟨0x2000, 16⟩, ⟨0, 2176⟩] + +theorem update_verified : Verified Arm.target update (Spec.Cmac.aesUpdateContract Arm.abi 8) := + Verified.of_correct (fun _ hs => update_wp hs) update_ct (by + sig_implies [Spec.Cmac.aesUpdateContract, Spec.Cmac.aesUpdateSig, updateArm, Arm.abi, Arm.argRegs, + Arm.reduceClassify, Arm.Loc.val, Arm.State.addr] + [updSat, Arm.stackArg, Arm.stackArgAddr, Mem.readW, Mem.read] using updSat) + +/-- A state satisfying `vg_cmac_aes_subkeys`'s precondition. -/ +def subSat : State where + gpr r := match r with + | .r0 => 0x1000 | .r1 => 10 | .r2 => 0x2000 | .r3 => 0x3000 | _ => 0 + sp := 0x8000 + n := false + z := false + c := false + v := false + mem _ := 0 + rd := [⟨0x1000, 240⟩] + wr := [⟨0x2000, 32⟩, ⟨0x3000, 2176⟩] + +theorem subkeys_verified : Verified Arm.target subkeys (Spec.Cmac.aesSubkeysContract Arm.abi 8) := + Verified.of_correct (fun _ hs => subkeys_wp hs) subkeys_ct (by + sig_implies [Spec.Cmac.aesSubkeysContract, Spec.Cmac.aesSubkeysSig, subkeysArm, Arm.abi, Arm.argRegs, + Arm.reduceClassify, Arm.Loc.val, Arm.State.addr] [subSat] using subSat) + +/-- A state satisfying `vg_cmac_aes_finalize`'s precondition (with no last +bytes, and the scratch buffer at 0). -/ +def finSat : State where + gpr r := match r with + | .r0 => 0x1000 | .r1 => 10 | .r2 => 0x2000 | .r3 => 0x3000 | _ => 0 + sp := 0x8000 + n := false + z := false + c := false + v := false + mem _ := 0 + rd := [⟨0x1000, 272⟩, ⟨0x3000, 0⟩, ⟨0x8000, 8⟩] + wr := [⟨0x2000, 16⟩, ⟨0, 2176⟩] + +theorem finalize_verified : Verified Arm.target finalize (Spec.Cmac.aesFinalizeContract Arm.abi 8) := + Verified.of_correct (fun _ hs => finalize_wp hs) finalize_ct (by + sig_implies [Spec.Cmac.aesFinalizeContract, Spec.Cmac.aesFinalizeSig, finalizeArm, Arm.abi, Arm.argRegs, + Arm.reduceClassify, Arm.Loc.val, Arm.State.addr] + [finSat, Arm.stackArg, Arm.stackArgAddr, Mem.readW, Mem.read] using finSat) + +end VG.Proof.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/CmacAes/Arm/Words.lean b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Words.lean new file mode 100644 index 000000000..0c7a8ee60 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/Arm/Words.lean @@ -0,0 +1,156 @@ +import VerifiedGarbage.Proof.CmacAes.Arm.Call +import VerifiedGarbage.Proof.Cmac.Block32 +import VerifiedGarbage.Proof.MdStream.Arm.Common + +/-! +# AES-CMAC on ARMv7: blocks formed a word at a time + +Untrusted: everything here is checked by Lean. Weakest preconditions of the +instruction sequences the functions build blocks with: the XOR of the blocks +at `pb + pd` and `qb + qd` stored at `cb + cd` through two temporaries +(`xorBlk`, which leaves `Cmac.xor4Mem`), and four stores of a zero register +(`zeroBlk`, which leaves `Cmac.zero4`). +-/ + +namespace VG.Proof.CmacAes.Arm + +open VG VG.Arm +open VG.Proof.MdStream.Arm (Upd Mupd WP.cons op2_reg wp_ldr wp_str) + +/-- The words of the blocks at `pb + pd` and `qb + qd`, XORed through `t₁` +and `t₂` and stored at `cb + cd`. -/ +def xorBlk (t₁ t₂ pb qb cb : Reg) (pd qd cd : Nat) : List Instr := + [.ldr t₁ pb pd, .ldr t₂ qb qd, .dp .eor t₁ t₁ (.reg t₂), .str t₁ cb cd, + .ldr t₁ pb (pd + 4), .ldr t₂ qb (qd + 4), .dp .eor t₁ t₁ (.reg t₂), .str t₁ cb (cd + 4), + .ldr t₁ pb (pd + 8), .ldr t₂ qb (qd + 8), .dp .eor t₁ t₁ (.reg t₂), .str t₁ cb (cd + 8), + .ldr t₁ pb (pd + 12), .ldr t₂ qb (qd + 12), .dp .eor t₁ t₁ (.reg t₂), .str t₁ cb (cd + 12)] + +/-- `z` stored in the four words at `b + d`. -/ +def zeroBlk (z b : Reg) (d : Nat) : List Instr := + [.str z b d, .str z b (d + 4), .str z b (d + 8), .str z b (d + 12)] + +/-- `s'` is `s` with memory `m`, and `t₁` and `t₂` clobbered. -/ +structure Step (s s' : State) (t₁ t₂ : Reg) (m : Mem) : Prop where + gpr : ∀ r, r ≠ t₁ → r ≠ t₂ → s'.gpr r = s.gpr r + mem : s'.mem = m + rd : s'.rd = s.rd + wr : s'.wr = s.wr + sp : s'.sp = s.sp + +theorem wp_eor {is : List Instr} {s : State} {Q : State → Prop} {d n : Reg} {o : Op2} {y : BitVec 32} + (ho : o.eval s = some y) (k : ∀ s', Upd s s' d (s.gpr n ^^^ y) → WP isa (.block is) s' Q) : + WP isa (.block (.dp .eor d n o :: is)) s Q := + WP.cons (s' := s.setReg d (s.gpr n ^^^ y)) (by simp [exec, ho]) (k _ (MdStream.Arm.Upd.setReg _ _ _)) + +/-- One word. -/ +theorem xw_ok {t₁ t₂ pb qb cb : Reg} {pd qd cd : Nat} {is : List Instr} {s : State} {Q : State → Prop} + {P Q' C : Addr} (h12 : t₁ ≠ t₂) (hq : qb ≠ t₁) (hc₁ : cb ≠ t₁) (hc₂ : cb ≠ t₂) + (hpd : pd < 4096) (hqd : qd < 4096) (hcd : cd < 4096) + (hP : State.addr (s.gpr pb + BitVec.ofNat 32 pd) = P) (hQ : State.addr (s.gpr qb + BitVec.ofNat 32 qd) = Q') + (hC : State.addr (s.gpr cb + BitVec.ofNat 32 cd) = C) + (rP : InRegions (s.rd ++ s.wr) P 4) (rQ : InRegions (s.rd ++ s.wr) Q' 4) (wC : InRegions s.wr C 4) + (k : ∀ s', Step s s' t₁ t₂ (s.mem.writeW C (s.mem.readW P 32 ^^^ s.mem.readW Q' 32)) → + WP isa (.block is) s' Q) : + WP isa (.block (.ldr t₁ pb pd :: .ldr t₂ qb qd :: .dp .eor t₁ t₁ (.reg t₂) :: .str t₁ cb cd :: is)) s Q := by + subst hQ hC + refine wp_ldr hpd hP rP fun s₁ u₁ => ?_ + refine wp_ldr hqd (by rw [u₁.other _ hq]) (by rw [u₁.rd, u₁.wr]; exact rQ) fun s₂ u₂ => ?_ + refine wp_eor (op2_reg _ _) fun s₃ u₃ => ?_ + refine wp_str hcd (by rw [u₃.other _ hc₁, u₂.other _ hc₂, u₁.other _ hc₁]) + (by rw [u₃.wr, u₂.wr, u₁.wr]; exact wC) fun s₄ u₄ => k s₄ ⟨fun r h₁ h₂ => ?_, ?_, ?_, ?_, ?_⟩ + · rw [u₄.gpr, u₃.other _ h₁, u₂.other _ h₂, u₁.other _ h₁] + · rw [u₄.mem, u₃.gpr, u₂.other _ h12, u₂.gpr, u₁.gpr, u₃.mem, u₂.mem, u₁.mem] + · rw [u₄.rd, u₃.rd, u₂.rd, u₁.rd] + · rw [u₄.wr, u₃.wr, u₂.wr, u₁.wr] + · rw [u₄.sp, u₃.sp, u₂.sp, u₁.sp] + +/-- Word `i` of a block that does not wrap the 32-bit space. -/ +theorem addr_word {b : BitVec 32} {d : Nat} (i : Nat) (h : b.toNat + d + 16 ≤ 2 ^ 32) (hi : i ≤ 12) : + State.addr (b + BitVec.ofNat 32 (d + i)) = State.addr b + BitVec.ofNat 64 d + BitVec.ofNat 64 i := by + rw [addr_add (by omega), Offset.add_add] + +theorem in_word {rs : List Region} {P : Addr} (h : Covers [⟨P, 16⟩] rs) {i : Nat} (hi : i ≤ 12) : + InRegions rs (P + BitVec.ofNat 64 i) 4 := + h _ _ ⟨_, List.mem_singleton_self _, Offset.contains_base P (by omega) (by omega)⟩ + +theorem in_word0 {rs : List Region} {P : Addr} (h : Covers [⟨P, 16⟩] rs) : InRegions rs P 4 := by + have c := Offset.contains_base P (d := 0) (n := 4) (k := 16) (by decide) (by decide) + rw [show P + BitVec.ofNat 64 0 = P from BitVec.add_zero P] at c + exact h _ _ ⟨_, List.mem_singleton_self _, c⟩ + +/-- The XOR of the blocks at `pb + pd` and `qb + qd`, stored at `cb + cd`. -/ +theorem xorBlk_ok {t₁ t₂ pb qb cb : Reg} {pd qd cd : Nat} {is : List Instr} {s : State} {Q : State → Prop} + (h12 : t₁ ≠ t₂) (hp₁ : pb ≠ t₁) (hp₂ : pb ≠ t₂) (hq₁ : qb ≠ t₁) (hq₂ : qb ≠ t₂) (hc₁ : cb ≠ t₁) + (hc₂ : cb ≠ t₂) (hpd : pd + 12 < 4096) (hqd : qd + 12 < 4096) (hcd : cd + 12 < 4096) + (fp : (s.gpr pb).toNat + pd + 16 ≤ 2 ^ 32) (fq : (s.gpr qb).toNat + qd + 16 ≤ 2 ^ 32) + (fc : (s.gpr cb).toNat + cd + 16 ≤ 2 ^ 32) + (rP : Covers [⟨State.addr (s.gpr pb) + BitVec.ofNat 64 pd, 16⟩] (s.rd ++ s.wr)) + (rQ : Covers [⟨State.addr (s.gpr qb) + BitVec.ofNat 64 qd, 16⟩] (s.rd ++ s.wr)) + (wC : Covers [⟨State.addr (s.gpr cb) + BitVec.ofNat 64 cd, 16⟩] s.wr) + (k : ∀ s', Step s s' t₁ t₂ (Proof.Cmac.xor4Mem s.mem (State.addr (s.gpr cb) + BitVec.ofNat 64 cd) + (State.addr (s.gpr pb) + BitVec.ofNat 64 pd) (State.addr (s.gpr qb) + BitVec.ofNat 64 qd)) → + WP isa (.block is) s' Q) : + WP isa (.block (xorBlk t₁ t₂ pb qb cb pd qd cd ++ is)) s Q := by + simp only [xorBlk, List.cons_append, List.nil_append] + refine xw_ok h12 hq₁ hc₁ hc₂ (by omega) (by omega) (by omega) (addr_add (by omega)) (addr_add (by omega)) + (addr_add (by omega)) (in_word0 rP) (in_word0 rQ) (in_word0 wC) fun s₁ g₁ => ?_ + have e₁ : ∀ r, r ≠ t₁ → r ≠ t₂ → s₁.gpr r = s.gpr r := g₁.gpr + refine xw_ok (P := State.addr (s.gpr pb) + BitVec.ofNat 64 pd + BitVec.ofNat 64 4) + (Q' := State.addr (s.gpr qb) + BitVec.ofNat 64 qd + BitVec.ofNat 64 4) + (C := State.addr (s.gpr cb) + BitVec.ofNat 64 cd + BitVec.ofNat 64 4) h12 hq₁ hc₁ hc₂ (by omega) (by omega) (by omega) + (by rw [e₁ _ hp₁ hp₂]; exact addr_word 4 fp (by decide)) + (by rw [e₁ _ hq₁ hq₂]; exact addr_word 4 fq (by decide)) + (by rw [e₁ _ hc₁ hc₂]; exact addr_word 4 fc (by decide)) + (by rw [g₁.rd, g₁.wr]; exact in_word rP (by decide)) (by rw [g₁.rd, g₁.wr]; exact in_word rQ (by decide)) + (by rw [g₁.wr]; exact in_word wC (by decide)) fun s₂ g₂ => ?_ + have e₂ : ∀ r, r ≠ t₁ → r ≠ t₂ → s₂.gpr r = s.gpr r := fun r h₁ h₂ => by rw [g₂.gpr r h₁ h₂, e₁ r h₁ h₂] + refine xw_ok (P := State.addr (s.gpr pb) + BitVec.ofNat 64 pd + BitVec.ofNat 64 8) + (Q' := State.addr (s.gpr qb) + BitVec.ofNat 64 qd + BitVec.ofNat 64 8) + (C := State.addr (s.gpr cb) + BitVec.ofNat 64 cd + BitVec.ofNat 64 8) h12 hq₁ hc₁ hc₂ (by omega) (by omega) (by omega) + (by rw [e₂ _ hp₁ hp₂]; exact addr_word 8 fp (by decide)) + (by rw [e₂ _ hq₁ hq₂]; exact addr_word 8 fq (by decide)) + (by rw [e₂ _ hc₁ hc₂]; exact addr_word 8 fc (by decide)) + (by rw [g₂.rd, g₂.wr, g₁.rd, g₁.wr]; exact in_word rP (by decide)) + (by rw [g₂.rd, g₂.wr, g₁.rd, g₁.wr]; exact in_word rQ (by decide)) + (by rw [g₂.wr, g₁.wr]; exact in_word wC (by decide)) fun s₃ g₃ => ?_ + have e₃ : ∀ r, r ≠ t₁ → r ≠ t₂ → s₃.gpr r = s.gpr r := fun r h₁ h₂ => by rw [g₃.gpr r h₁ h₂, e₂ r h₁ h₂] + refine xw_ok (P := State.addr (s.gpr pb) + BitVec.ofNat 64 pd + BitVec.ofNat 64 12) + (Q' := State.addr (s.gpr qb) + BitVec.ofNat 64 qd + BitVec.ofNat 64 12) + (C := State.addr (s.gpr cb) + BitVec.ofNat 64 cd + BitVec.ofNat 64 12) h12 hq₁ hc₁ hc₂ (by omega) (by omega) (by omega) + (by rw [e₃ _ hp₁ hp₂]; exact addr_word 12 fp (by decide)) + (by rw [e₃ _ hq₁ hq₂]; exact addr_word 12 fq (by decide)) + (by rw [e₃ _ hc₁ hc₂]; exact addr_word 12 fc (by decide)) + (by rw [g₃.rd, g₃.wr, g₂.rd, g₂.wr, g₁.rd, g₁.wr]; exact in_word rP (by decide)) + (by rw [g₃.rd, g₃.wr, g₂.rd, g₂.wr, g₁.rd, g₁.wr]; exact in_word rQ (by decide)) + (by rw [g₃.wr, g₂.wr, g₁.wr]; exact in_word wC (by decide)) fun s₄ g₄ => k s₄ ⟨?_, ?_, ?_, ?_, ?_⟩ + · intro r h₁ h₂; rw [g₄.gpr r h₁ h₂, e₃ r h₁ h₂] + · rw [g₄.mem, g₃.mem, g₂.mem, g₁.mem]; rfl + · rw [g₄.rd, g₃.rd, g₂.rd, g₁.rd] + · rw [g₄.wr, g₃.wr, g₂.wr, g₁.wr] + · rw [g₄.sp, g₃.sp, g₂.sp, g₁.sp] + +/-- The block at `b + d` zeroed, from a register `z` holding zero. -/ +theorem zeroBlk_ok {z b : Reg} {d : Nat} {is : List Instr} {s : State} {Q : State → Prop} + (hz : s.gpr z = 0) (hd : d + 12 < 4096) (fb : (s.gpr b).toNat + d + 16 ≤ 2 ^ 32) + (wB : Covers [⟨State.addr (s.gpr b) + BitVec.ofNat 64 d, 16⟩] s.wr) + (k : ∀ s', s'.gpr = s.gpr → s'.mem = Proof.Cmac.zero4 s.mem (State.addr (s.gpr b) + BitVec.ofNat 64 d) → + s'.rd = s.rd → s'.wr = s.wr → s'.sp = s.sp → WP isa (.block is) s' Q) : + WP isa (.block (zeroBlk z b d ++ is)) s Q := by + simp only [zeroBlk, List.cons_append, List.nil_append] + refine wp_str (by omega) (addr_add (by omega)) (in_word0 wB) fun s₁ u₁ => ?_ + refine wp_str (a := State.addr (s.gpr b) + BitVec.ofNat 64 d + BitVec.ofNat 64 4) (by omega) + (by rw [u₁.gpr]; exact addr_word 4 fb (by decide)) + (by rw [u₁.wr]; exact in_word wB (by decide)) fun s₂ u₂ => ?_ + refine wp_str (a := State.addr (s.gpr b) + BitVec.ofNat 64 d + BitVec.ofNat 64 8) (by omega) + (by rw [u₂.gpr, u₁.gpr]; exact addr_word 8 fb (by decide)) + (by rw [u₂.wr, u₁.wr]; exact in_word wB (by decide)) fun s₃ u₃ => ?_ + refine wp_str (a := State.addr (s.gpr b) + BitVec.ofNat 64 d + BitVec.ofNat 64 12) (by omega) + (by rw [u₃.gpr, u₂.gpr, u₁.gpr]; exact addr_word 12 fb (by decide)) + (by rw [u₃.wr, u₂.wr, u₁.wr]; exact in_word wB (by decide)) fun s₄ u₄ => k s₄ ?_ ?_ ?_ ?_ ?_ + · rw [u₄.gpr, u₃.gpr, u₂.gpr, u₁.gpr] + · rw [u₄.mem, u₃.mem, u₂.mem, u₁.mem, u₃.gpr, u₂.gpr, u₁.gpr, hz]; rfl + · rw [u₄.rd, u₃.rd, u₂.rd, u₁.rd] + · rw [u₄.wr, u₃.wr, u₂.wr, u₁.wr] + · rw [u₄.sp, u₃.sp, u₂.sp, u₁.sp] + +end VG.Proof.CmacAes.Arm diff --git a/lean/VerifiedGarbage/Proof/Framework/Arm/ArgTaint.lean b/lean/VerifiedGarbage/Proof/Framework/Arm/ArgTaint.lean new file mode 100644 index 000000000..d751255ae --- /dev/null +++ b/lean/VerifiedGarbage/Proof/Framework/Arm/ArgTaint.lean @@ -0,0 +1,73 @@ +import VerifiedGarbage.Proof.Framework.Arm.Taint +import VerifiedGarbage.Proof.Framework.RelCT + +/-! +# A taint state for code reading its stack arguments (ARMv7) + +Untrusted: everything here is checked by Lean. + +As `Framework/X86/ArgTaint.lean`: pieces of code between calls (proved by +relating two runs, `RelCT`) may read their function's stack arguments: +`argTaint rs n` makes the registers `rs` and the first `n` bytes of stack +arguments public, which two runs agree on as long as the arguments are the +same and lie outside the writable regions (`agree_argTaint`). `rel_agree` +relates two runs of code the taint analysis checks, each described by `WP`. +-/ + +namespace VG.Arm + +/-- The taint in which the registers `rs` and the first `n` bytes of stack +arguments are public. -/ +def argTaint (rs : List Reg) (n : Nat) : VG.Arm.Taint.T := + { regs := RegSet.ofList rs, flags := false, argLen := n } + +/-- The first `4 j` bytes of stack arguments agree when their first `j` words do. -/ +theorem argMem_of {s₁ s₂ : State} {j : Nat} (hsp : s₁.sp = s₂.sp) (hf : s₁.sp.toNat + 4 * j ≤ 2 ^ 32) + (h : ∀ i < j, stackArg s₁ i = stackArg s₂ i) : + ∀ k < 4 * j, s₁.mem (VG.Arm.Taint.argByte s₁ k) = s₂.mem (VG.Arm.Taint.argByte s₂ k) := by + intro k hk + have e : ∀ s : State, s.sp.toNat + 4 * j ≤ 2 ^ 32 → + VG.Arm.Taint.argByte s k = stackArgAddr s (k / 4) + BitVec.ofNat 64 (k % 4) := fun s hs => by + simp only [VG.Arm.Taint.argByte, stackArgAddr] + rw [addr_add (by omega), BitVec.add_assoc, ← BitVec.ofNat_add] + congr 2; omega + rw [e s₁ hf, e s₂ (hsp ▸ hf), Mem.readW_byte s₁.mem _ (Nat.mod_lt _ (by omega)), + Mem.readW_byte s₂.mem _ (Nat.mod_lt _ (by omega))] + exact congrArg _ (h _ (by omega)) + +theorem agree_argTaint {rs : List Reg} {n : Nat} {s₁ s₂ : State} (h : ∀ r ∈ rs, s₁.gpr r = s₂.gpr r) + (hsp : s₁.sp = s₂.sp) + (hw₁ : s₁.sp.toNat + n ≤ 2 ^ 32 ∧ ∀ r ∈ s₁.wr, Region.Disjoint ⟨State.addr s₁.sp, n⟩ r) + (hw₂ : s₂.sp.toNat + n ≤ 2 ^ 32 ∧ ∀ r ∈ s₂.wr, Region.Disjoint ⟨State.addr s₂.sp, n⟩ r) + (hm : ∀ k < n, s₁.mem (VG.Arm.Taint.argByte s₁ k) = s₂.mem (VG.Arm.Taint.argByte s₂ k)) : + VG.Arm.Taint.Agree (argTaint rs n) s₁ s₂ where + rf := ⟨fun r hr => h r (RegSet.mem_ofList.mp hr), fun h => by cases h⟩ + wr h := absurd rfl h + wf₁ := ⟨fun h => absurd rfl h, fun _ h => (List.not_mem_nil h).elim, fun _ => hw₁, + fun _ h => (List.not_mem_nil h).elim⟩ + wf₂ := ⟨fun h => absurd rfl h, fun _ h => (List.not_mem_nil h).elim, fun _ => hw₂, + fun _ h => (List.not_mem_nil h).elim⟩ + ok _ h := (List.not_mem_nil h).elim + slots _ h := (List.not_mem_nil h).elim + sp _ := hsp + argMem := hm + +/-- Code the taint analysis checks from `τ`, in two runs whose single-run +facts `F` and `F'` make them agree on it. -/ +theorem rel_agree {F F' G G' : State → Prop} {c : Prog isa} (τ : VG.Arm.Taint.T) + (hag : ∀ s s', F s → F' s' → VG.Arm.Taint.Agree τ s s') + (hc : ∃ hc, (VG.Taint.check taint τ c hc).isSome = true) + (hw : ∀ s, F s → WP isa c s G) (hw' : ∀ s, F' s → WP isa c s G') : + RelCT isa (fun s s' => F s ∧ F' s') c fun s s' => G s ∧ G' s' := by + obtain ⟨_, hc⟩ := hc + exact ((RelCT.taint (A := taint) τ (fun s s' h => hag s s' h.1 h.2) hc).wp + fun s s' h => ⟨hw s h.1, hw' s' h.2⟩).mono (fun _ _ h => h) fun _ _ h => h.2 + +/-- Code constant time in two runs, each described by `WP`. -/ +theorem rel_wp {F F' G G' : State → Prop} {c : Prog isa} + (hct : RelCT isa (fun s s' => F s ∧ F' s') c fun _ _ => True) + (hw : ∀ s, F s → WP isa c s G) (hw' : ∀ s, F' s → WP isa c s G') : + RelCT isa (fun s s' => F s ∧ F' s') c fun s s' => G s ∧ G' s' := + (hct.wp fun s s' h => ⟨hw s h.1, hw' s' h.2⟩).mono (fun _ _ h => h) fun _ _ h => h.2 + +end VG.Arm diff --git a/src/asm/arm/cmac_aes.rs b/src/asm/arm/cmac_aes.rs new file mode 100644 index 000000000..6d7a7e846 --- /dev/null +++ b/src/asm/arm/cmac_aes.rs @@ -0,0 +1,310 @@ +// @generated from lean/VerifiedGarbage/Artifacts.lean by lean/Emit.lean. DO NOT EDIT. +//! Verified `cmac_aes` functions for `arm`. +#![allow(dead_code)] + +/// The CMAC subkey generation (NIST SP 800-38B §6.1) for AES: writes `K1 ‖ K2` to `*subkeys`, where `L = CIPH_K(0¹²⁸)`, `K1 = L << 1` (XORed with `R₁₂₈ = 0¹²⁰10000111` if the leftmost bit of `L` is 1) and `K2` is `K1` doubled the same way. `CIPH_K` is AES (FIPS 197) with `rounds` rounds and the key schedule in the first `16 * (rounds + 1)` bytes of `*schedule`, as `vg_aes_expand_key` writes it. +/// +/// Contract: `VG.Spec.Cmac.aesSubkeysContract`. Constant time: only the pointers and `rounds` may affect timing, not the key schedule or the subkeys. +/// +/// This implementation encrypts each block with `vg_aes_ctr32`. +/// +/// # Safety +/// +/// * `schedule` must be valid for reads of 240 bytes. +/// * `subkeys` must be valid for reads and writes of 32 bytes. +/// * `scratch` must be valid for reads and writes of 2176 bytes. +/// * `rounds` must be 10, 12 or 14. +/// * The contents of `scratch` on return are unspecified. +/// * `subkeys` and `scratch` must not overlap each other or `schedule` (distinct Rust objects never do). +/// * None of `schedule`, `subkeys` and `scratch` may overlap the 8 bytes of stack below the stack pointer, or wrap around the end of the address space (no Rust object does). +#[unsafe(naked)] +pub(crate) unsafe extern "C" fn vg_cmac_aes_subkeys(schedule: *const [u8; 240], rounds: usize, subkeys: *mut [u8; 32], scratch: *mut [u64; 272]) { + core::arch::naked_asm!( + "str r4, [r3, #2064]", + "str r5, [r3, #2068]", + "str r6, [r3, #2072]", + "str lr, [r3, #2076]", + "mov r6, r2", + "mov r5, r3", + "mov r12, #0", + "str r12, [r3, #2048]", + "str r12, [r3, #2052]", + "str r12, [r3, #2056]", + "str r12, [r3, #2060]", + "str r12, [r2, #0]", + "str r12, [r2, #4]", + "str r12, [r2, #8]", + "str r12, [r2, #12]", + "add r2, r5, #2048", + "mov r3, r6", + "mov r4, #1", + "push {{r4, r5}}", + "bl {vg_aes_ctr32}", + "ldr r4, [sp], #8", + "ldr r0, [r6, #0]", + "ldr r1, [r6, #4]", + "ldr r2, [r6, #8]", + "ldr r3, [r6, #12]", + "rev r0, r0", + "rev r1, r1", + "rev r2, r2", + "rev r3, r3", + "lsr r12, r0, #31", + "mov r4, #0", + "sub r12, r4, r12", + "and r12, r12, #135", + "lsl r0, r0, #1", + "orr r0, r0, r1, lsr #31", + "lsl r1, r1, #1", + "orr r1, r1, r2, lsr #31", + "lsl r2, r2, #1", + "orr r2, r2, r3, lsr #31", + "lsl r3, r3, #1", + "eor r3, r3, r12", + "rev r0, r0", + "rev r1, r1", + "rev r2, r2", + "rev r3, r3", + "str r0, [r6, #0]", + "str r1, [r6, #4]", + "str r2, [r6, #8]", + "str r3, [r6, #12]", + "ldr r0, [r6, #0]", + "ldr r1, [r6, #4]", + "ldr r2, [r6, #8]", + "ldr r3, [r6, #12]", + "rev r0, r0", + "rev r1, r1", + "rev r2, r2", + "rev r3, r3", + "lsr r12, r0, #31", + "mov r4, #0", + "sub r12, r4, r12", + "and r12, r12, #135", + "lsl r0, r0, #1", + "orr r0, r0, r1, lsr #31", + "lsl r1, r1, #1", + "orr r1, r1, r2, lsr #31", + "lsl r2, r2, #1", + "orr r2, r2, r3, lsr #31", + "lsl r3, r3, #1", + "eor r3, r3, r12", + "rev r0, r0", + "rev r1, r1", + "rev r2, r2", + "rev r3, r3", + "str r0, [r6, #16]", + "str r1, [r6, #20]", + "str r2, [r6, #24]", + "str r3, [r6, #28]", + "ldr r4, [r5, #2064]", + "ldr r6, [r5, #2072]", + "ldr lr, [r5, #2076]", + "ldr r5, [r5, #2068]", + "bx lr", + vg_aes_ctr32 = sym super::aes::vg_aes_ctr32, + ) +} + +/// CMAC's chaining (NIST SP 800-38B §6.2 step 6) for AES, over whole blocks: replaces the block `C₀` at `*state` with `Cₙ`, where `Cᵢ = CIPH_K(Cᵢ₋₁ ⊕ Mᵢ)` for the `n` 16-byte blocks `M₁ … Mₙ` starting at `data`. `CIPH_K` is AES (FIPS 197) with `rounds` rounds and the key schedule in the first `16 * (rounds + 1)` bytes of `*schedule`, as `vg_aes_expand_key` writes it. +/// +/// Contract: `VG.Spec.Cmac.aesUpdateContract`. Constant time: only the pointers, `rounds` and `n` may affect timing, not the key schedule, the chaining value or the data. +/// +/// This implementation encrypts each block with `vg_aes_ctr32`. +/// +/// # Safety +/// +/// * `schedule` must be valid for reads of 240 bytes. +/// * `state` must be valid for reads and writes of 16 bytes. +/// * `data` must be valid for reads of `16 * n` bytes. +/// * `scratch` must be valid for reads and writes of 2176 bytes. +/// * `rounds` must be 10, 12 or 14. +/// * The contents of `scratch` on return are unspecified. +/// * `state` and `scratch` must not overlap each other, `schedule`, `data` or the arguments on the stack (distinct Rust objects never do). +/// * None of `schedule`, `state`, `data` and `scratch` may overlap the 8 bytes of stack below the stack pointer, or wrap around the end of the address space (no Rust object does). +#[unsafe(naked)] +pub(crate) unsafe extern "C" fn vg_cmac_aes_update(schedule: *const [u8; 240], rounds: usize, state: *mut [u8; 16], data: *const [u8; 16], n: usize, scratch: *mut [u64; 272]) { + core::arch::naked_asm!( + "ldr r12, [sp, #4]", + "str r4, [r12, #2064]", + "str r5, [r12, #2068]", + "str r6, [r12, #2072]", + "str r7, [r12, #2076]", + "str r8, [r12, #2080]", + "str r9, [r12, #2084]", + "str lr, [r12, #2092]", + "str r10, [r12, #2088]", + "mov r4, r0", + "mov r5, r1", + "mov r6, r2", + "mov r7, r3", + "ldr r8, [sp, #0]", + "mov r10, r12", + "cmp r8, #0", + "beq 20f", + "22:", + "ldr r0, [r6, #0]", + "ldr r1, [r7, #0]", + "eor r0, r0, r1", + "str r0, [r10, #2048]", + "ldr r0, [r6, #4]", + "ldr r1, [r7, #4]", + "eor r0, r0, r1", + "str r0, [r10, #2052]", + "ldr r0, [r6, #8]", + "ldr r1, [r7, #8]", + "eor r0, r0, r1", + "str r0, [r10, #2056]", + "ldr r0, [r6, #12]", + "ldr r1, [r7, #12]", + "eor r0, r0, r1", + "str r0, [r10, #2060]", + "mov r0, #0", + "str r0, [r6, #0]", + "str r0, [r6, #4]", + "str r0, [r6, #8]", + "str r0, [r6, #12]", + "mov r0, r4", + "mov r1, r5", + "add r2, r10, #2048", + "mov r3, r6", + "mov r9, #1", + "push {{r9, r10}}", + "bl {vg_aes_ctr32}", + "ldr r9, [sp], #8", + "add r7, r7, #16", + "subs r8, r8, #1", + "bne 22b", + "b 21f", + "20:", + "21:", + "ldr r4, [r10, #2064]", + "ldr r5, [r10, #2068]", + "ldr r6, [r10, #2072]", + "ldr r7, [r10, #2076]", + "ldr r8, [r10, #2080]", + "ldr r9, [r10, #2084]", + "ldr lr, [r10, #2092]", + "ldr r10, [r10, #2088]", + "bx lr", + vg_aes_ctr32 = sym super::aes::vg_aes_ctr32, + ) +} + +/// Finishes an AES-CMAC computation (NIST SP 800-38B §6.2, with `Tlen = 128`): if the block at `*state` is the chaining value `Cₙ₋₁` of the message's blocks but the last (as `vg_cmac_aes_update` computes it from a zero block), and the `last_len` bytes at `last` are the message's last bytes `Mₙ*`, replaces it with the MAC `Cₙ = CIPH_K(Cₙ₋₁ ⊕ Mₙ)`, where `Mₙ = K1 ⊕ Mₙ*` if `last_len` is 16, and `Mₙ = K2 ⊕ (Mₙ* ‖ 10ʲ)` otherwise. `*key` is the 240 bytes `vg_aes_expand_key` writes the key schedule for `rounds` rounds to, followed by the subkeys `K1 ‖ K2` (as `vg_cmac_aes_subkeys` writes them). `last_len` is 0 only for the empty message. +/// +/// Contract: `VG.Spec.Cmac.aesFinalizeContract`. Constant time: only the pointers, `rounds` and `last_len` may affect timing, not the key schedule, the subkeys, the chaining value or the data. +/// +/// This implementation encrypts each block with `vg_aes_ctr32`. +/// +/// # Safety +/// +/// * `key` must be valid for reads of 272 bytes. +/// * `state` must be valid for reads and writes of 16 bytes. +/// * `last` must be valid for reads of `last_len` bytes. +/// * `scratch` must be valid for reads and writes of 2176 bytes. +/// * `rounds` must be 10, 12 or 14. +/// * `last_len` must be at most 16. +/// * The contents of `scratch` on return are unspecified. +/// * `state` and `scratch` must not overlap each other, `key`, `last` or the arguments on the stack (distinct Rust objects never do). +/// * None of `key`, `state`, `last` and `scratch` may overlap the 8 bytes of stack below the stack pointer, or wrap around the end of the address space (no Rust object does). +#[unsafe(naked)] +pub(crate) unsafe extern "C" fn vg_cmac_aes_finalize(key: *const [u8; 272], rounds: usize, state: *mut [u8; 16], last: *const u8, last_len: usize, scratch: *mut [u64; 272]) { + core::arch::naked_asm!( + "ldr r12, [sp, #4]", + "str r4, [r12, #2064]", + "str r5, [r12, #2068]", + "str lr, [r12, #2072]", + "mov r5, r12", + "ldr r4, [sp, #0]", + "cmp r4, #16", + "beq 20f", + "mov r12, #0", + "str r12, [r5, #2048]", + "str r12, [r5, #2052]", + "str r12, [r5, #2056]", + "str r12, [r5, #2060]", + "add lr, r5, #2048", + "cmp r4, #0", + "beq 22f", + "24:", + "ldrb r12, [r3, #0]", + "strb r12, [lr, #0]", + "add r3, r3, #1", + "add lr, lr, #1", + "subs r4, r4, #1", + "bne 24b", + "b 23f", + "22:", + "23:", + "mov r12, #128", + "strb r12, [lr, #0]", + "ldr r12, [r5, #2048]", + "ldr lr, [r0, #256]", + "eor r12, r12, lr", + "str r12, [r5, #2048]", + "ldr r12, [r5, #2052]", + "ldr lr, [r0, #260]", + "eor r12, r12, lr", + "str r12, [r5, #2052]", + "ldr r12, [r5, #2056]", + "ldr lr, [r0, #264]", + "eor r12, r12, lr", + "str r12, [r5, #2056]", + "ldr r12, [r5, #2060]", + "ldr lr, [r0, #268]", + "eor r12, r12, lr", + "str r12, [r5, #2060]", + "b 21f", + "20:", + "ldr r12, [r3, #0]", + "ldr lr, [r0, #240]", + "eor r12, r12, lr", + "str r12, [r5, #2048]", + "ldr r12, [r3, #4]", + "ldr lr, [r0, #244]", + "eor r12, r12, lr", + "str r12, [r5, #2052]", + "ldr r12, [r3, #8]", + "ldr lr, [r0, #248]", + "eor r12, r12, lr", + "str r12, [r5, #2056]", + "ldr r12, [r3, #12]", + "ldr lr, [r0, #252]", + "eor r12, r12, lr", + "str r12, [r5, #2060]", + "21:", + "ldr r12, [r5, #2048]", + "ldr lr, [r2, #0]", + "eor r12, r12, lr", + "str r12, [r5, #2048]", + "ldr r12, [r5, #2052]", + "ldr lr, [r2, #4]", + "eor r12, r12, lr", + "str r12, [r5, #2052]", + "ldr r12, [r5, #2056]", + "ldr lr, [r2, #8]", + "eor r12, r12, lr", + "str r12, [r5, #2056]", + "ldr r12, [r5, #2060]", + "ldr lr, [r2, #12]", + "eor r12, r12, lr", + "str r12, [r5, #2060]", + "mov r12, #0", + "str r12, [r2, #0]", + "str r12, [r2, #4]", + "str r12, [r2, #8]", + "str r12, [r2, #12]", + "mov r3, r2", + "add r2, r5, #2048", + "mov r4, #1", + "push {{r4, r5}}", + "bl {vg_aes_ctr32}", + "ldr r4, [sp], #8", + "ldr r4, [r5, #2064]", + "ldr lr, [r5, #2072]", + "ldr r5, [r5, #2068]", + "bx lr", + vg_aes_ctr32 = sym super::aes::vg_aes_ctr32, + ) +} diff --git a/src/asm/arm/mod.rs b/src/asm/arm/mod.rs index 93aa724f6..d157c1ca8 100644 --- a/src/asm/arm/mod.rs +++ b/src/asm/arm/mod.rs @@ -16,6 +16,9 @@ pub(crate) mod chacha20; #[rustfmt::skip] pub(crate) mod chacha20poly1305; +#[rustfmt::skip] +pub(crate) mod cmac_aes; + #[rustfmt::skip] pub(crate) mod ct; diff --git a/src/cmac/aes.rs b/src/cmac/aes.rs index 2c9cd0f2d..539826c16 100644 --- a/src/cmac/aes.rs +++ b/src/cmac/aes.rs @@ -15,9 +15,9 @@ //! same verified CMAC code, calling `vg_aes_ctr32_aesni` rather than //! `vg_aes_ctr32` to encrypt each block. On AArch64, CPUs with the AES //! extension run `vg_aes_expand_key_aes` and the `_aes` CMAC functions, -//! calling `vg_aes_ctr32_aes`. +//! calling `vg_aes_ctr32_aes`. ARMv7 has only the scalar implementation. -#![cfg(any(target_arch = "x86_64", target_arch = "aarch64"))] +#![cfg(any(target_arch = "x86_64", target_arch = "aarch64", target_arch = "arm"))] use super::{InvalidKeyLength, InvalidMac}; use crate::arch::aes::vg_aes_expand_key; @@ -91,6 +91,12 @@ impl Backend { Backend::Scalar } } + + /// The only implementation there is. + #[cfg(target_arch = "arm")] + fn select(_: Features) -> Backend { + Backend::Scalar + } } /// An incremental AES-CMAC computation. @@ -349,4 +355,11 @@ mod tests { assert_eq!(Backend::select(Features::of(&[])), Backend::Scalar); assert_eq!(Backend::select(Features::of(&["aes"])), Backend::ArmCrypto); } + + /// The scalar implementation is the only one. + #[cfg(target_arch = "arm")] + #[test] + fn select() { + assert_eq!(Backend::select(Features::of(&[])), Backend::Scalar); + } } diff --git a/tests/cavp/cmac_aes.rs b/tests/cavp/cmac_aes.rs index 490db371b..dfe7ba7d4 100644 --- a/tests/cavp/cmac_aes.rs +++ b/tests/cavp/cmac_aes.rs @@ -1,7 +1,7 @@ //! AES-CMAC: every vector of the CMAC generation and verification files, //! for 128-, 192- and 256-bit keys (the MAC truncated to `Tlen` bytes). -#![cfg(any(target_arch = "x86_64", target_arch = "aarch64"))] +#![cfg(any(target_arch = "x86_64", target_arch = "aarch64", target_arch = "arm"))] use verified_garbage::cmac::aes::AesCmac; diff --git a/tests/wycheproof/cmac_aes.rs b/tests/wycheproof/cmac_aes.rs index 7a07326cd..135e3cdb1 100644 --- a/tests/wycheproof/cmac_aes.rs +++ b/tests/wycheproof/cmac_aes.rs @@ -5,7 +5,7 @@ //! invalid one is either a modified tag, which `verify` must reject, or a //! key of a length AES does not take, which `new` must reject. -#![cfg(any(target_arch = "x86_64", target_arch = "aarch64"))] +#![cfg(any(target_arch = "x86_64", target_arch = "aarch64", target_arch = "arm"))] use serde::Deserialize; use verified_garbage::cmac::InvalidKeyLength;