diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index cd358a843..77d025c99 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -407,7 +407,7 @@ jobs: fail-fast: false matrix: chip: [native, skx, hsw, p4p] - tests: [sha1 sha256 sha384 sha512 aes_gcm chacha20 mlkem x25519 ed25519 cpu] + tests: [sha1 sha256 sha384 sha512 aes_gcm cmac chacha20 mlkem x25519 ed25519 cpu] include: - chip: arl tests: sha384 sha512 ed25519 cpu diff --git a/README.md b/README.md index 25b339af4..272a85b9b 100644 --- a/README.md +++ b/README.md @@ -199,7 +199,7 @@ yours to keep: ✅ -❌ +✅ AES-NI ❌ diff --git a/bench/benches/primitives/cmac_aes.rs b/bench/benches/primitives/cmac_aes.rs new file mode 100644 index 000000000..3de7113fe --- /dev/null +++ b/bench/benches/primitives/cmac_aes.rs @@ -0,0 +1,67 @@ +//! AES-CMAC. + +use criterion::Criterion; + +/// The library modules whose code these benchmarks run (see +/// `ci/bench_arches.py`): this one and those it calls. +pub const USES: &[&str] = &["cmac_aes", "aes"]; + +/// The MAC of a message with a 16-byte key (setup included), computed and +/// verified. +#[cfg(target_arch = "x86_64")] +pub fn bench(c: &mut Criterion) { + use std::hint::black_box; + + use criterion::{BenchmarkId, Throughput}; + use openssl::pkey::PKey; + use openssl::sign::Signer; + use openssl::symm::Cipher; + use verified_garbage::cmac::aes::AesCmac; + + use crate::{OPENSSL, SIZES, VG}; + + let key = [0x42; 16]; + let pkey = PKey::cmac(&Cipher::aes_128_cbc(), &key).unwrap(); + let mut g = c.benchmark_group("aes-128-cmac"); + for size in SIZES { + g.throughput(Throughput::Bytes(size as u64)); + let data = vec![0x5a; size]; + g.bench_function(BenchmarkId::new(VG, size), |b| { + b.iter(|| AesCmac::mac(black_box(&key), black_box(&data)).unwrap()) + }); + let mut out = [0u8; 16]; + g.bench_function(BenchmarkId::new(OPENSSL, size), |b| { + b.iter(|| { + let mut s = Signer::new_without_digest(&pkey).unwrap(); + s.sign_oneshot(&mut out, black_box(&data)).unwrap() + }) + }); + } + g.finish(); + + let mut g = c.benchmark_group("aes-128-cmac-verify"); + for size in SIZES { + g.throughput(Throughput::Bytes(size as u64)); + let data = vec![0x5a; size]; + let mac = AesCmac::mac(&key, &data).unwrap(); + g.bench_function(BenchmarkId::new(VG, size), |b| { + b.iter(|| { + let mut m = AesCmac::new(black_box(&key)).unwrap(); + m.update(black_box(&data)); + m.verify(black_box(&mac)).unwrap() + }) + }); + let mut out = [0u8; 16]; + g.bench_function(BenchmarkId::new(OPENSSL, size), |b| { + b.iter(|| { + let mut s = Signer::new_without_digest(&pkey).unwrap(); + let n = s.sign_oneshot(&mut out, black_box(&data)).unwrap(); + assert!(openssl::memcmp::eq(&out[..n], black_box(&mac))) + }) + }); + } + g.finish(); +} + +#[cfg(not(target_arch = "x86_64"))] +pub fn bench(_: &mut Criterion) {} diff --git a/bench/benches/primitives/main.rs b/bench/benches/primitives/main.rs index 787c57cc9..d317bbf76 100644 --- a/bench/benches/primitives/main.rs +++ b/bench/benches/primitives/main.rs @@ -21,6 +21,7 @@ mod blake2b; mod blake2s; mod chacha20; mod chacha20poly1305; +mod cmac_aes; mod ed25519; mod hmac_md5; mod hmac_sha1; @@ -193,6 +194,7 @@ const BENCHES: &[Bench] = &[ (blake2s::USES, blake2s::bench), (chacha20::USES, chacha20::bench), (chacha20poly1305::USES, chacha20poly1305::bench), + (cmac_aes::USES, cmac_aes::bench), (hmac_md5::USES, hmac_md5::bench), (hmac_sha1::USES, hmac_sha1::bench), (hmac_sha256::USES, hmac_sha256::bench), diff --git a/lean/VerifiedGarbage/Generic/AesCtr32/X86_64/CmacAes.lean b/lean/VerifiedGarbage/Generic/AesCtr32/X86_64/CmacAes.lean new file mode 100644 index 000000000..43c8fddd7 --- /dev/null +++ b/lean/VerifiedGarbage/Generic/AesCtr32/X86_64/CmacAes.lean @@ -0,0 +1,62 @@ +import VerifiedGarbage.TCB.X86_64.Target +import VerifiedGarbage.Proof.CmacAes.X86_64.Verified + +/-! +# AES-CMAC (NIST SP 800-38B) on x86-64 + +A generic file (see `TCB/Emit.lean`): the artifacts it lists, calling an +implementation `v` of `vg_aes_ctr32`, are emitted once for each +implementation (`Variants/AesCtr32/X86_64/`), named with its suffix (e.g. +`vg_cmac_aes_update_aesni`), and need its CPU features. **Review note**: +`sig` and `doc` are trusted, as they tie the Rust caller to the contract; +check them against the contract's `pre`/`post`. 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. + +The stack is 8 bytes for every implementation: the return address of the +call of `vg_aes_ctr32`, which makes no calls. +-/ + +namespace VG.Generic.AesCtr32.X86_64.CmacAes + +open VG.Proof.CmacAes.X86_64 + +/-- Which implementation of `vg_aes_ctr32` an instance calls. -/ +def ctrNote (v : Proof.Aes.X86_64.Ctr32Impl) : String := + "This implementation encrypts each block with `" ++ v.callee.name ++ "`." + +def artifacts (v : Proof.Aes.X86_64.Ctr32Impl) : List Artifact := [ + { Spec.Cmac.aesSubkeysApi with + name := Spec.Cmac.aesSubkeysApi.name ++ v.suffix + target := X86_64.target + doc := Spec.Cmac.aesSubkeysApi.doc (notes := [ctrNote v]) + code := Impl.CmacAes.X86_64.subkeys v.callee + contract := Spec.Cmac.aesSubkeysContract X86_64.abi 8 + stack := 8 + verified := subkeys_verified v + spSafe := subkeys_spSafe v + features := v.features }, + { Spec.Cmac.aesUpdateApi with + name := Spec.Cmac.aesUpdateApi.name ++ v.suffix + target := X86_64.target + doc := Spec.Cmac.aesUpdateApi.doc (notes := [ctrNote v]) + code := Impl.CmacAes.X86_64.update v.callee + contract := Spec.Cmac.aesUpdateContract X86_64.abi 8 + stack := 8 + verified := update_verified v + spSafe := update_spSafe v + features := v.features }, + { Spec.Cmac.aesFinalizeApi with + name := Spec.Cmac.aesFinalizeApi.name ++ v.suffix + target := X86_64.target + doc := Spec.Cmac.aesFinalizeApi.doc (notes := [ctrNote v]) + code := Impl.CmacAes.X86_64.finalize v.callee + contract := Spec.Cmac.aesFinalizeContract X86_64.abi 8 + stack := 8 + verified := finalize_verified v + spSafe := finalize_spSafe v + features := v.features }] + +end VG.Generic.AesCtr32.X86_64.CmacAes diff --git a/lean/VerifiedGarbage/Impl/Aes/X86_64/Callee.lean b/lean/VerifiedGarbage/Impl/Aes/X86_64/Callee.lean new file mode 100644 index 000000000..2e85e2e4f --- /dev/null +++ b/lean/VerifiedGarbage/Impl/Aes/X86_64/Callee.lean @@ -0,0 +1,23 @@ +import VerifiedGarbage.Impl.Aes.X86_64.Ctr32 +import VerifiedGarbage.Impl.Aes.X86_64.AesNi + +/-! +# The implementations of `vg_aes_ctr32` on x86-64 + +A function that calls `vg_aes_ctr32` (AES-CMAC's) takes the implementation it +calls, a `Ctr32`, and is emitted once for each (`Generic/AesCtr32/X86_64/`). +-/ + +namespace VG.Impl.Aes.X86_64 + +open VG.X86_64 + +/-- An implementation of `vg_aes_ctr32` to call: its symbol and its code. -/ +structure Ctr32 where + name : String + code : Prog isa + +def Ctr32.scalar : Ctr32 := ⟨"vg_aes_ctr32", ctr32⟩ +def Ctr32.aesni : Ctr32 := ⟨"vg_aes_ctr32_aesni", AesNi.ctr32⟩ + +end VG.Impl.Aes.X86_64 diff --git a/lean/VerifiedGarbage/Impl/CmacAes/X86_64.lean b/lean/VerifiedGarbage/Impl/CmacAes/X86_64.lean new file mode 100644 index 000000000..28297cab6 --- /dev/null +++ b/lean/VerifiedGarbage/Impl/CmacAes/X86_64.lean @@ -0,0 +1,168 @@ +import VerifiedGarbage.Impl.Aes.X86_64.Callee + +/-! +# AES-CMAC: x86-64 implementation + +`vg_cmac_aes_subkeys(schedule = rdi, rounds = rsi, subkeys = rdx, scratch = rcx)`, +`vg_cmac_aes_update(schedule = rdi, rounds = rsi, state = rdx, data = rcx, n = r8, scratch = r9)` +and `vg_cmac_aes_finalize(key = rdi, rounds = rsi, state = rdx, last = rcx, last_len = r8, scratch = r9)` +(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. They are +generic over the implementation of `vg_aes_ctr32` they call (`Ctr32`): each +is emitted once for each implementation (e.g. `vg_cmac_aes_update` calls +`vg_aes_ctr32`, and `vg_cmac_aes_update_aesni` calls `vg_aes_ctr32_aesni`). + +The scratch buffer (2176 bytes): `[0, 2048)` is the working space of +`vg_aes_ctr32`, `[2048, 2064)` the counter block, and `[2064, 2112)` our +caller's callee-saved registers. + +* `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 `rax:rdx`, shifted left by one bit, and + XORed with `0x87` masked by the bit shifted out. `rbx` holds `subkeys` + and `rbp` the scratch buffer across the call. +* `update` keeps its arguments in `rbx` (schedule), `rbp` (rounds), `r12` + (state), `r13` (data), `r14` (blocks left) and `r15` (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, so it keeps nothing across the call. + +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.X86_64 + +open VG.X86_64 +open VG.Impl.Aes.X86_64 (Ctr32) + +def at_ (b : Reg) (d : Nat) : MemOp := { base := b, disp := d } + +/-- The offset of the counter block in the scratch buffer. -/ +def cOff : Nat := 2048 + +/-- The arguments of `vg_aes_ctr32` for one block, other than the schedule +and the rounds: the counter block at `scr + cOff`, the data block at `out`, +`n = 1`, and the working space at `scr`. -/ +def ctrArgs (scr out : Reg) : List Instr := + [.mov .r9 (.reg scr), .mov .rdx (.reg scr), .alu .add .rdx (.imm (BitVec.ofNat 32 cOff)), + .mov .rcx (.reg out), .mov32 .r8 (.imm 1)] + +/-! ## `vg_cmac_aes_subkeys` -/ + +/-- Saves `rbx` and `rbp`, keeps `subkeys` in `rbx` and the scratch buffer in +`rbp`, zeroes the counter block and the first block of `subkeys`, and sets up +the arguments of `vg_aes_ctr32`. -/ +def subkeysPre : List Instr := + [.store (at_ .rcx 2064) .rbx, .store (at_ .rcx 2072) .rbp, .mov .rbx (.reg .rdx), + .mov .rbp (.reg .rcx), .mov32 .rax (.imm 0), .store (at_ .rcx cOff) .rax, + .store (at_ .rcx (cOff + 8)) .rax, .store (at_ .rdx 0) .rax, .store (at_ .rdx 8) .rax] ++ + ctrArgs .rbp .rbx + +/-- The block at `rbx + src`, doubled (`VG.Spec.Cmac.dbl 16`), to `rbx + dst`. -/ +def dbl (src dst : Nat) : List Instr := + [.mov .rax (.mem (at_ .rbx src)), .bswap .rax, .mov .rdx (.mem (at_ .rbx (src + 8))), .bswap .rdx, + .mov .rcx (.reg .rax), .shift .shr .rcx 63, .mov32 .r8 (.imm 0), .alu .sub .r8 (.reg .rcx), + .alu .and .r8 (.imm 0x87), + .mov .rcx (.reg .rdx), .shift .shr .rcx 63, .alu .add .rax (.reg .rax), .alu .or .rax (.reg .rcx), + .alu .add .rdx (.reg .rdx), .alu .xor .rdx (.reg .r8), + .bswap .rax, .bswap .rdx, .store (at_ .rbx dst) .rax, .store (at_ .rbx (dst + 8)) .rdx] + +/-- `K1` over `L`, `K2` after it, and the saved registers restored. -/ +def subkeysPost : List Instr := + dbl 0 0 ++ dbl 0 16 ++ [.mov .rbx (.mem (at_ .rbp 2064)), .mov .rbp (.mem (at_ .rbp 2072))] + +def subkeys (c : Ctr32) : Prog isa := + .seq (.block subkeysPre) (.seq (.call c.name c.code) (.block subkeysPost)) + +/-! ## `vg_cmac_aes_update` -/ + +def saved : List (Reg × Nat) := + [(.rbx, 2064), (.rbp, 2072), (.r12, 2080), (.r13, 2088), (.r14, 2096), (.r15, 2104)] + +def save : List Instr := saved.map fun (r, d) => .store (at_ .r9 d) r + +/-- Restores the registers, with `r15` (restored last) the scratch buffer. -/ +def restore : List Instr := saved.map fun (r, d) => .mov r (.mem (at_ .r15 d)) + +/-- The arguments to their registers; ZF is set if there are no blocks. -/ +def setup : List Instr := + [.mov .rbx (.reg .rdi), .mov .rbp (.reg .rsi), .mov .r12 (.reg .rdx), .mov .r13 (.reg .rcx), + .mov .r14 (.reg .r8), .mov .r15 (.reg .r9), .alu .test .r14 (.reg .r14)] + +/-- The counter block `C ⊕ Mᵢ` (the state at `r12`, the block at `r13`), and +the state zeroed. -/ +def chainIn : List Instr := + [.mov .rax (.mem (at_ .r12 0)), .alu .xor .rax (.mem (at_ .r13 0)), .store (at_ .r15 cOff) .rax, + .mov .rax (.mem (at_ .r12 8)), .alu .xor .rax (.mem (at_ .r13 8)), .store (at_ .r15 (cOff + 8)) .rax, + .mov32 .rax (.imm 0), .store (at_ .r12 0) .rax, .store (at_ .r12 8) .rax] + +/-- The arguments of `vg_aes_ctr32` for the block. -/ +def updArgs : List Instr := [.mov .rdi (.reg .rbx), .mov .rsi (.reg .rbp)] ++ ctrArgs .r15 .r12 + +/-- On to the next block (ZF is set when none are left). -/ +def advance : List Instr := [.alu .add .r13 (.imm 16), .alu .sub .r14 (.imm 1)] + +/-- One block. -/ +def body (c : Ctr32) : Prog isa := + .seq (.block (chainIn ++ updArgs)) (.seq (.call c.name c.code) (.block advance)) + +def update (c : Ctr32) : Prog isa := + .seq (.block (save ++ setup)) + (.seq (.ite .e (.block []) (.loop (body c) .ne)) (.block restore)) + +/-! ## `vg_cmac_aes_finalize` -/ + +/-- `Mₙ = Mₙ* ⊕ K1` (`K1` at `rdi + 240`), for a complete last block. -/ +def full : List Instr := + [.mov .rax (.mem (at_ .rcx 0)), .alu .xor .rax (.mem (at_ .rdi 240)), .store (at_ .r9 cOff) .rax, + .mov .rax (.mem (at_ .rcx 8)), .alu .xor .rax (.mem (at_ .rdi 248)), .store (at_ .r9 (cOff + 8)) .rax] + +/-- `[rcx + r10]` and `[r9 + r10 + cOff]`. -/ +def lastByte : MemOp := { base := .rcx, index := some .r10 } +def padByte : MemOp := { base := .r9, index := some .r10, disp := cOff } + +/-- The counter block zeroed; ZF is set if `last_len` is 0. -/ +def zero : List Instr := + [.mov32 .rax (.imm 0), .store (at_ .r9 cOff) .rax, .store (at_ .r9 (cOff + 8)) .rax, + .alu .test .r8 (.reg .r8)] + +/-- The `last_len` (1 to 15) bytes at `rcx` copied to the counter block. -/ +def copy : Prog isa := + .seq (.block [.mov32 .r10 (.imm 0)]) + (.loop (.block [.movzx8 .rax lastByte, .store8 padByte .rax, .alu .add .r10 (.imm 1), + .alu .cmp .r10 (.reg .r8)]) .ne) + +/-- `0x80` after the bytes, and the block XORed with `K2` (at `rdi + 256`). -/ +def padK2 : List Instr := + [.mov32 .rax (.imm 0x80), .store8 { base := .r9, index := some .r8, disp := cOff } .rax, + .mov .rax (.mem (at_ .r9 cOff)), .alu .xor .rax (.mem (at_ .rdi 256)), .store (at_ .r9 cOff) .rax, + .mov .rax (.mem (at_ .r9 (cOff + 8))), .alu .xor .rax (.mem (at_ .rdi 264)), + .store (at_ .r9 (cOff + 8)) .rax] + +/-- `Mₙ = K2 ⊕ (Mₙ* ‖ 10ʲ)`, for a partial last block (`last_len < 16`). -/ +def partialBlock : Prog isa := + .seq (.block zero) (.seq (.ite .e (.block []) copy) (.block padK2)) + +/-- The counter block `C ⊕ Mₙ` (the state at `rdx`), the state zeroed, and +the arguments of `vg_aes_ctr32` but the schedule (`rdi`) and the rounds +(`rsi`), which are ours. -/ +def finArgs : List Instr := + [.mov .rax (.mem (at_ .r9 cOff)), .alu .xor .rax (.mem (at_ .rdx 0)), .store (at_ .r9 cOff) .rax, + .mov .rax (.mem (at_ .r9 (cOff + 8))), .alu .xor .rax (.mem (at_ .rdx 8)), + .store (at_ .r9 (cOff + 8)) .rax, + .mov32 .rax (.imm 0), .store (at_ .rdx 0) .rax, .store (at_ .rdx 8) .rax, + .mov .rcx (.reg .rdx), .mov .rdx (.reg .r9), .alu .add .rdx (.imm (BitVec.ofNat 32 cOff)), + .mov32 .r8 (.imm 1)] + +/-- Everything before the call. -/ +def finPre : Prog isa := + .seq (.block [.alu .cmp .r8 (.imm 16)]) (.seq (.ite .e (.block full) partialBlock) (.block finArgs)) + +def finalize (c : Ctr32) : Prog isa := .seq finPre (.call c.name c.code) + +end VG.Impl.CmacAes.X86_64 diff --git a/lean/VerifiedGarbage/Proof/Aes/X86_64/Variant.lean b/lean/VerifiedGarbage/Proof/Aes/X86_64/Variant.lean new file mode 100644 index 000000000..ac7239b8d --- /dev/null +++ b/lean/VerifiedGarbage/Proof/Aes/X86_64/Variant.lean @@ -0,0 +1,88 @@ +import VerifiedGarbage.Proof.Aes.X86_64.Ctr32 +import VerifiedGarbage.Proof.Aes.X86_64.AesNi.Ctr32 +import VerifiedGarbage.Impl.Aes.X86_64.Callee +import VerifiedGarbage.Proof.Framework.X86_64.Call + +/-! +# Implementations of `vg_aes_ctr32` on x86-64 + +Untrusted: everything here is checked by Lean. + +A `Ctr32Impl` is what a function that calls `vg_aes_ctr32` needs of it, so +that its proof holds for every implementation: each is a variant of the +interface `AesCtr32` on x86-64 (`Variants/AesCtr32/X86_64/`), and each caller +(in `Generic/AesCtr32/X86_64/`) is emitted once for each of them (see +`TCB/Emit.lean`). Every implementation is proven against the same contract, +`Proof.Aes.ctr32X86_64`, and makes no calls. +-/ + +namespace VG.Proof.Aes.X86_64 + +open VG.X86_64 + +/-- An implementation of `vg_aes_ctr32` on x86-64. -/ +structure Ctr32Impl where + /-- Its symbol and code. -/ + callee : Impl.Aes.X86_64.Ctr32 + /-- It makes no calls. -/ + depth : callee.code.depth = 0 + ok : ∀ s, Proof.Aes.ctr32X86_64.pre s → + ∃ t s', Exec isa callee.code s t s' ∧ abiPreserved s s' ∧ Proof.Aes.ctr32X86_64.post s s' + ct : ConstantTime isa Proof.Aes.ctr32X86_64.pre Proof.Aes.ctr32X86_64.pub callee.code + /-- It never writes the stack pointer. -/ + nosp : NoSp callee.code + /-- It never loads MXCSR. -/ + mxcsr : callee.code.allInstrs (fun i => !loadsMxcsr i) = true + spSafe : callee.code.all (fun i => !isa.writesSp i) = true + /-- What the names of its callers' instances end with (e.g. `_aesni`; + nothing for the baseline implementation). -/ + suffix : String + /-- The CPU features its code requires, which its callers require too. -/ + features : List String + +namespace Ctr32Impl + +theorem scalar_nosp : NoSp Impl.Aes.X86_64.Ctr32.scalar.code := by + have : ((instrs Impl.Aes.X86_64.Ctr32.scalar.code).all fun i => !Taint.clobbers i .rsp) = true := by + rw [← Code.allInstrs_eq]; lit_decide + exact fun i hi => by simpa using List.all_eq_true.mp this i hi + +/-- The bitsliced implementation, `vg_aes_ctr32`, in the baseline ISA. -/ +def scalar : Ctr32Impl where + callee := .scalar + depth := by lit_decide + ok := ctr32_correct + ct := ctr32_ct + nosp := scalar_nosp + mxcsr := by lit_decide + spSafe := Code.all_of_allInstrs (by lit_decide) + suffix := "" + features := [] + +theorem aesni_nosp : NoSp Impl.Aes.X86_64.Ctr32.aesni.code := by + have : ((instrs Impl.Aes.X86_64.Ctr32.aesni.code).all fun i => !Taint.clobbers i .rsp) = true := by + rw [← Code.allInstrs_eq]; lit_decide + exact fun i hi => by simpa using List.all_eq_true.mp this i hi + +/-- `vg_aes_ctr32_aesni`'s own contract is the same, but for `rsp`, which its +public data leaves out. -/ +theorem aesni_ct : ConstantTime isa Proof.Aes.ctr32X86_64.pre Proof.Aes.ctr32X86_64.pub + Impl.Aes.X86_64.Ctr32.aesni.code := + fun s₁ s₂ t₁ t₂ s₁' s₂' h₁ h₂ ⟨a, b, c, d, e, f, _⟩ e₁ e₂ => + AesNi.ctr32_ct s₁ s₂ t₁ t₂ s₁' s₂' h₁ h₂ ⟨a, b, c, d, e, f⟩ e₁ e₂ + +/-- The AES-NI implementation, `vg_aes_ctr32_aesni`. -/ +def aesni : Ctr32Impl where + callee := .aesni + depth := by lit_decide + ok := AesNi.ctr32_correct + ct := aesni_ct + nosp := aesni_nosp + mxcsr := by lit_decide + spSafe := Code.all_of_allInstrs (by lit_decide) + suffix := "_aesni" + features := ["aes", "ssse3"] + +end Ctr32Impl + +end VG.Proof.Aes.X86_64 diff --git a/lean/VerifiedGarbage/Proof/Cmac/Dbl.lean b/lean/VerifiedGarbage/Proof/Cmac/Dbl.lean new file mode 100644 index 000000000..323aa1664 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/Cmac/Dbl.lean @@ -0,0 +1,102 @@ +import VerifiedGarbage.Proof.Cmac.Spec + +/-! +# CMAC: doubling a 16-byte block as a 128-bit integer + +Untrusted: everything here is checked by Lean. + +`dbl_eq`: the doubling of §6.1 on 16 bytes (`Spec.Cmac.dbl 16`) is, on the +block as a big-endian 128-bit integer `x` (`Spec.Gcm.ofBytes`), the shift +`x << 1` XORed with `0x87` if the bit shifted out was 1. +-/ + +namespace VG.Proof.Cmac + +open VG Spec.Cmac + +/-- Bit `j` of byte `i` of a block, as an integer. -/ +theorem ofBytes_bit {L : List Byte} (hL : L.length = 16) {i j : Nat} (hi : i < 16) (hj : j < 8) : + (Spec.Gcm.ofBytes L).getLsbD (8 * (15 - i) + j) = (L.getD i 0).getLsbD j := by + rw [← Proof.Aes.toBytes_ofBytes hL hi, Proof.Aes.toBytes_getD _ hi, BitVec.getLsbD_extractLsb'] + simp [hj] + +/-- The 128-bit doubling. -/ +def dbl128 (x : BitVec 128) : BitVec 128 := (x <<< 1) ^^^ (if x.msb then 0x87 else 0) + +theorem getD_shiftLeft1 {L : List Byte} (hL : L.length = 16) {k : Nat} (hk : k < 16) : + (shiftLeft1 L).getD k 0 = (L.getD k 0 <<< 1) ||| (((L.drop 1 ++ [0]).getD k 0 : Byte) >>> 7) := by + simp [shiftLeft1, List.getD_eq_getElem?_getD, hL, hk] + +theorem next_getD {L : List Byte} (hL : L.length = 16) {k : Nat} (hk : k < 15) : + (L.drop 1 ++ [0]).getD k 0 = L.getD (k + 1) 0 := by + simp [List.getD_eq_getElem?_getD, List.getElem?_append, hL, hk, + List.getElem?_eq_getElem (show k + 1 < L.length by omega)] + +theorem next_getD15 {L : List Byte} (hL : L.length = 16) : (L.drop 1 ++ [0]).getD 15 0 = 0 := by + simp [List.getD_eq_getElem?_getD, hL] + +theorem getD_rb : ∀ k < 16, (rb 16).getD k 0 = if k = 15 then 0x87 else 0 := by decide + +theorem testBit_135 : ∀ p < 128, 8 ≤ p → Nat.testBit 135 p = false := by decide + +theorem high_0x87 {p : Nat} (hp : 8 ≤ p) : (0x87 : BitVec 128).getLsbD p = false := by + rw [show (0x87 : BitVec 128) = BitVec.ofNat 128 135 from rfl, BitVec.getLsbD_ofNat] + by_cases h : p < 128 + · rw [testBit_135 p h hp, Bool.and_false] + · simp [h] + +theorem bit_0x87 : ∀ j < 8, (0x87 : BitVec 128).getLsbD j = (0x87 : Byte).getLsbD j := by decide + +theorem dbl_eq {L : List Byte} (hL : L.length = 16) : + dbl 16 L = Spec.Gcm.toBytes (dbl128 (Spec.Gcm.ofBytes L)) := by + have hmsb : (Spec.Gcm.ofBytes L).msb = msb1 L := by + rw [BitVec.msb_eq_getLsbD_last, show 128 - 1 = 8 * (15 - 0) + 7 from rfl, + ofBytes_bit hL (by decide) (by decide), msb1, BitVec.msb_eq_getLsbD_last] + cases L with + | nil => simp at hL + | cons a _ => rfl + have hsl : (shiftLeft1 L).length = 16 := by simp [shiftLeft1, hL] + refine ext16 (by unfold dbl; split <;> simp [length_xor, hsl, rb, zeros]) (toBytes_length _) + fun k hk => ?_ + rw [Proof.Aes.toBytes_getD _ hk] + apply BitVec.eq_of_getLsbD_eq + intro j hj + rw [BitVec.getLsbD_extractLsb', dbl128] + simp only [hj, decide_true, Bool.true_and, BitVec.getLsbD_xor, hmsb, BitVec.getLsbD_shiftLeft] + have hdbl : (dbl 16 L).getD k 0 = (shiftLeft1 L).getD k 0 ^^^ (if msb1 L then (rb 16).getD k 0 else 0) := by + unfold dbl + split + · rw [getD_xor (by simp [hsl, rb, zeros]) (by rw [hsl]; exact hk)] + · simp + rw [hdbl, getD_shiftLeft1 hL hk, getD_rb k hk] + simp only [BitVec.getLsbD_xor, BitVec.getLsbD_or, BitVec.getLsbD_shiftLeft, BitVec.getLsbD_ushiftRight, + hj, decide_true, Bool.true_and] + have hm : ∀ p, 8 ≤ p → (if msb1 L = true then (135 : BitVec 128) else 0).getLsbD p = false := by + intro p hp; split + · exact high_0x87 hp + · simp + rcases Nat.lt_or_ge k 15 with hk15 | hk15 + · have hm0 : (if msb1 L = true then (if k = 15 then (135 : Byte) else 0) else 0).getLsbD j = false := by + simp [show k ≠ 15 by omega] + rw [hm0, hm _ (by omega), Bool.xor_false, Bool.xor_false, next_getD hL hk15] + rcases Nat.eq_zero_or_pos j with rfl | hj0 + · rw [show 8 * (15 - k) + 0 - 1 = 8 * (15 - (k + 1)) + 7 by omega, ofBytes_bit hL (by omega) (by decide)] + simp [show ¬ 8 * (15 - k) < 1 by omega, show 8 * (15 - k) < 128 by omega] + · rw [show 8 * (15 - k) + j - 1 = 8 * (15 - k) + (j - 1) by omega, ofBytes_bit hL hk (by omega), + BitVec.getLsbD_of_ge _ (7 + j) (by omega)] + simp [show ¬ j < 1 by omega, show ¬ 8 * (15 - k) + j < 1 by omega, show 8 * (15 - k) + j < 128 by omega] + · have hk' : k = 15 := by omega + subst hk' + rw [next_getD15 hL] + rcases Nat.eq_zero_or_pos j with rfl | hj0 + · cases hb : msb1 L + · simp + · simp + · rw [show 8 * (15 - 15) + j - 1 = 8 * (15 - 15) + (j - 1) by omega, + ofBytes_bit hL (i := 15) (by decide) (by omega)] + cases hb : msb1 L + · simp [show ¬ j < 1 by omega, show j < 128 by omega] + · simp [show ¬ j < 1 by omega, show j < 128 by omega] + rw [← BitVec.getLsbD_eq_getElem]; exact (bit_0x87 j hj).symm + +end VG.Proof.Cmac diff --git a/lean/VerifiedGarbage/Proof/Cmac/Mem.lean b/lean/VerifiedGarbage/Proof/Cmac/Mem.lean new file mode 100644 index 000000000..a4b374eb7 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/Cmac/Mem.lean @@ -0,0 +1,85 @@ +import VerifiedGarbage.Proof.Cmac.Spec +import VerifiedGarbage.Proof.Framework.Mem +import VerifiedGarbage.Proof.Framework.Offset + +/-! +# CMAC: blocks in memory as 64-bit words + +Untrusted: everything here is checked by Lean. + +A 16-byte block is loaded and stored as two little-endian 64-bit words: +`le8 w` is the bytes of the word `w`, so the bytes at `p` are +`le8 (readW p) ++ le8 (readW (p + BitVec.ofNat 64 8))`, and storing `w₀` at `p` and `w₁` at +`p + 8` leaves `le8 w₀ ++ le8 w₁` there. +-/ + +namespace VG.Proof.Cmac + +open VG Spec.Cmac + +/-- The bytes of a 64-bit word, least significant first. -/ +def le8 (w : BitVec 64) : List Byte := (List.range 8).map fun i => w.extractLsb' (8 * i) 8 + +theorem length_le8 (w : BitVec 64) : (le8 w).length = 8 := by simp [le8] + +theorem getD_le8 (w : BitVec 64) {k : Nat} (hk : k < 8) : (le8 w).getD k 0 = w.extractLsb' (8 * k) 8 := by + simp [le8, List.getD_eq_getElem?_getD, hk] + +theorem getD_bytesAt (m : Mem) (p : Addr) {n k : Nat} (hk : k < n) : + (Spec.Aes.bytesAt m p n).getD k 0 = m (p + BitVec.ofNat 64 k) := by + simp [Spec.Aes.bytesAt, List.getD_eq_getElem?_getD, hk] + +theorem le8_readW (m : Mem) (a : Addr) : le8 (m.readW a 64) = Spec.Aes.bytesAt m a 8 := by + apply List.ext_getElem (by simp [le8, Spec.Aes.bytesAt]) + intro k h₁ h₂ + have hk : k < 8 := by simpa [le8] using h₁ + simp only [le8, Spec.Aes.bytesAt, List.getElem_map, List.getElem_range] + rw [← Mem.extractLsb'_read m a (n := 8) hk] + simp only [Mem.readW] + rfl + +theorem le8_xor (a b : BitVec 64) : le8 (a ^^^ b) = Spec.Cmac.xor (le8 a) (le8 b) := by + apply List.ext_getElem (by simp [le8, Spec.Cmac.xor]) + intro k h₁ h₂ + have hk : k < 8 := by simpa [le8] using h₁ + simp only [le8, Spec.Cmac.xor, List.getElem_map, List.getElem_range, List.getElem_zipWith] + ext j hj + simp + +theorem le8_zero : le8 0 = zeros 8 := by decide + +theorem bytesAt_split (m : Mem) (p : Addr) : + Spec.Aes.bytesAt m p 16 = Spec.Aes.bytesAt m p 8 ++ Spec.Aes.bytesAt m (p + BitVec.ofNat 64 8) 8 := by + simp only [Spec.Aes.bytesAt] + rw [show (16 : Nat) = 8 + 8 from rfl, 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] + +/-- The bytes of a block after storing its two words. -/ +theorem bytesAt_store2 (m : Mem) (p : Addr) (w₀ w₁ : BitVec 64) : + Spec.Aes.bytesAt ((m.writeW p w₀).writeW (p + BitVec.ofNat 64 8) w₁) p 16 = le8 w₀ ++ le8 w₁ := by + have hs : Mem.Sep p (64 / 8) (p + BitVec.ofNat 64 8) (64 / 8) := by + have := Offset.sep p (d := 0) (n := 8) (e := 8) (k := 8) (by decide) (by decide) (by decide) + simpa using this + have h₀ : Spec.Aes.bytesAt ((m.writeW p w₀).writeW (p + BitVec.ofNat 64 8) w₁) p 8 = le8 w₀ := by + rw [← le8_readW, Mem.readW_writeW_sep hs (by decide), Mem.readW_writeW_self64] + have h₁ : Spec.Aes.bytesAt ((m.writeW p w₀).writeW (p + BitVec.ofNat 64 8) w₁) (p + BitVec.ofNat 64 8) 8 = le8 w₁ := by + rw [← le8_readW, Mem.readW_writeW_self64] + rw [bytesAt_split, h₀, h₁] + +theorem xor_append {a b c d : List Byte} (h : a.length = c.length) : + Spec.Cmac.xor (a ++ b) (c ++ d) = Spec.Cmac.xor a c ++ Spec.Cmac.xor b d := by + simp [Spec.Cmac.xor, List.zipWith_append h] + +/-- The XOR of two blocks, a word at a time. -/ +theorem xor_words (m : Mem) (p q : Addr) : + le8 (m.readW p 64 ^^^ m.readW q 64) ++ le8 (m.readW (p + BitVec.ofNat 64 8) 64 ^^^ m.readW (q + BitVec.ofNat 64 8) 64) = + Spec.Cmac.xor (Spec.Aes.bytesAt m p 16) (Spec.Aes.bytesAt m q 16) := by + rw [le8_xor, le8_xor, le8_readW, le8_readW, le8_readW, le8_readW, bytesAt_split m p, + bytesAt_split m q, xor_append (by simp [Spec.Aes.bytesAt])] + +end VG.Proof.Cmac diff --git a/lean/VerifiedGarbage/Proof/Cmac/Spec.lean b/lean/VerifiedGarbage/Proof/Cmac/Spec.lean new file mode 100644 index 000000000..2bfd30a77 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/Cmac/Spec.lean @@ -0,0 +1,114 @@ +import VerifiedGarbage.Spec.Cmac +import VerifiedGarbage.Proof.Aes.Blocks + +/-! +# CMAC: lemmas about the specification + +Untrusted: everything here is checked by Lean. + +* `macFull_split`: the MAC of a message of whole blocks followed by its last + bytes `Mₙ*` (at most a block, and some unless the message is empty) is the + cipher of the chaining value of the whole blocks XORed with `Mₙ` (§6.2 + steps 3–6), which is how `update` and `finalize` compute it. +* `ctr32_one`: counter mode on one zero block is the cipher of the counter + block, which is how the implementations call `vg_aes_ctr32`: as bytes, + `Spec.Gcm.toBytes (aesWith nr w (Spec.Gcm.ofBytes x)) = Cmac.aesWith nr w x`. +-/ + +namespace VG.Proof.Cmac + +open VG Spec.Cmac + +/-! ## Lists of bytes -/ + +theorem length_xor (x y : List Byte) : (Spec.Cmac.xor x y).length = min x.length y.length := by + simp [Spec.Cmac.xor] + +theorem length_zeros (n : Nat) : (zeros n).length = n := by simp [zeros] + +theorem getD_xor {x y : List Byte} (h : x.length = y.length) {k : Nat} (hk : k < x.length) : + (Spec.Cmac.xor x y).getD k 0 = x.getD k 0 ^^^ y.getD k 0 := by + simp only [Spec.Cmac.xor, List.getD_eq_getElem?_getD, List.getElem?_zipWith, + List.getElem?_eq_getElem hk, List.getElem?_eq_getElem (show k < y.length by omega)] + rfl + +/-- Lists of 16 bytes are equal if their bytes are. -/ +theorem ext16 {x y : List Byte} (hx : x.length = 16) (hy : y.length = 16) + (h : ∀ k < 16, x.getD k 0 = y.getD k 0) : x = y := by + apply List.ext_getElem (by rw [hx, hy]) + intro k h₁ h₂ + have := h k (by omega) + simpa [List.getD_eq_getElem?_getD, h₁, h₂] using this + +/-! ## Splitting the MAC -/ + +theorem chain_append (ciph : Cipher) (c : List Byte) (xs ys : List (List Byte)) : + chain ciph c (xs ++ ys) = chain ciph (chain ciph c xs) ys := by + simp [chain, List.foldl_append] + +theorem chain_single (ciph : Cipher) (c m : List Byte) : chain ciph c [m] = ciph (xor c m) := rfl + +theorem blocks_eq {msg : List Byte} (hm : msg.length % 16 = 0) (last : List Byte) (q : Nat) + (hq : msg.length = 16 * q) : + (List.range q).map (fun i => ((msg ++ last).drop (16 * i)).take 16) = blocks 16 msg := by + simp only [blocks, hq, Nat.mul_div_cancel_left _ (by decide : 0 < 16)] + apply List.map_congr_left + intro i hi + rw [List.mem_range] at hi + rw [List.drop_append_of_le_length (by omega), List.take_append_of_le_length (by simp; omega)] + +/-- §6.2 steps 3–6, for a message of whole blocks `msg` followed by `last`. -/ +theorem macFull_split (ciph : Cipher) {msg last : List Byte} (hm : msg.length % 16 = 0) + (hl : last.length ≤ 16) (hne : msg = [] ∨ 0 < last.length) : + macFull ciph 16 (msg ++ last) = + ciph (xor (chain ciph (zeros 16) (blocks 16 msg)) + (lastBlock 16 (subkeys ciph 16).1 (subkeys ciph 16).2 last)) := by + obtain ⟨q, hq⟩ : ∃ q, msg.length = 16 * q := ⟨msg.length / 16, by omega⟩ + have hn : (if (msg ++ last).length = 0 then 1 else ((msg ++ last).length + 16 - 1) / 16) = q + 1 := by + rw [List.length_append] + split + · omega + · have hl0 : 0 < last.length := by + rcases hne with h | h + · subst h; simp only [List.length_nil] at *; omega + · exact h + omega + simp only [macFull, hn, Nat.add_sub_cancel] + rw [blocks_eq hm last q hq, chain_append, chain_single, + List.drop_append_of_le_length (by omega), List.drop_eq_nil_of_le (by omega), List.nil_append] + +/-! ## One block of counter mode -/ + +theorem toBytes_length (x : Spec.Gcm.Block) : (Spec.Gcm.toBytes x).length = 16 := by simp [Spec.Gcm.toBytes] + +theorem toBytes_ofBytes {bs : List Byte} (h : bs.length = 16) : Spec.Gcm.toBytes (Spec.Gcm.ofBytes bs) = bs := + ext16 (toBytes_length _) h fun _ hk => Proof.Aes.toBytes_ofBytes h hk + +theorem ctr32_one (ciph : Spec.Gcm.Block → Spec.Gcm.Block) (icb : Spec.Gcm.Block) : + Spec.Gcm.ctr32 ciph icb [0] = [ciph icb] := by + simp [Spec.Gcm.ctr32, Spec.Gcm.keystream, Nat.repeat] + +theorem aesWith_bytes (nr : Nat) (w : List Byte) {x : List Byte} (h : x.length = 16) : + Spec.Gcm.toBytes (Spec.Gcm.aesWith nr w (Spec.Gcm.ofBytes x)) = aesWith nr w x := by + rw [Spec.Gcm.aesWith, toBytes_ofBytes (by simp), aesWith] + congr 2 + exact Vector.ext fun i hi => by simpa only [Vector.getElem_ofFn] using Proof.Aes.toBytes_ofBytes h hi + +theorem bytesAt_length (m : Mem) (p : Addr) (n : Nat) : (Spec.Aes.bytesAt m p n).length = n := by + simp [Spec.Aes.bytesAt] + +theorem bytesAt_blockAt (m : Mem) (p : Addr) : + Spec.Aes.bytesAt m p 16 = Spec.Gcm.toBytes (Spec.Gcm.blockAt m p) := + (toBytes_ofBytes (bytesAt_length m p 16)).symm + +/-! ## Lengths -/ + +theorem aesWith_length (nr : Nat) (w x : List Byte) : (aesWith nr w x).length = 16 := by simp [aesWith] + +theorem dbl_length {x : List Byte} (h : x.length = 16) : (dbl 16 x).length = 16 := by + unfold dbl; split <;> simp [shiftLeft1, length_xor, rb, zeros, h] + +theorem subkeys_aes_length (nr : Nat) (w : List Byte) : (subkeys (aesWith nr w) 16).1.length = 16 := + dbl_length (aesWith_length _ _ _) + +end VG.Proof.Cmac diff --git a/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Call.lean b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Call.lean new file mode 100644 index 000000000..e53fd2326 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Call.lean @@ -0,0 +1,141 @@ +import VerifiedGarbage.Proof.Aes.X86_64.Variant +import VerifiedGarbage.Proof.Cmac.Spec +import VerifiedGarbage.Proof.Framework.X86_64.RelCT + +/-! +# AES-CMAC on x86-64: calling `vg_aes_ctr32` on one block + +Untrusted: everything here is checked by Lean. + +`ctr_call`: a call of any implementation of `vg_aes_ctr32` with the counter +block `C`, one data block `D` holding zeros, and working space `S`, from its +contract (with `WP.call`): `D` then holds `CIPH_K(C)`, as bytes +(`Cmac.aesWith`), and only `C`, `D`, `S` and the return address change. +-/ + +namespace VG.Proof.CmacAes.X86_64 + +open VG VG.X86_64 +open VG.Proof.Aes.X86_64 (Ctr32Impl) + +/-- Bytes outside a frame are unchanged. -/ +theorem bytesAt_frame {rs : List Region} {m m' : Mem} (hf : Frame rs m m') {p : Addr} {n : Nat} + (hd : ∀ r ∈ rs, (⟨p, n⟩ : Region).Disjoint r) (hn : n ≤ 2 ^ 64) : + Spec.Aes.bytesAt m' p n = Spec.Aes.bytesAt m p n := by + simp only [Spec.Aes.bytesAt] + apply List.map_congr_left + intro i hi + exact hf.bytes (R := ⟨p, n⟩) hd hn (List.mem_range.mp hi) + +/-- The return address a call stores. -/ +theorem callEntry_frame (s : State) : Frame [below (s.gpr .rsp) 8] s.mem s.callEntry.mem := by + rw [State.callEntry_mem] + exact (Frame.refl _ _).writeW (List.mem_singleton_self _) _ (below_call _ (by decide) (by decide)) + +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 64 R).toNat = R := by + rw [BitVec.toNat_ofNat]; exact Nat.mod_eq_of_lt (by omega) + +theorem one_toNat : (1 : BitVec 64).toNat = 1 := rfl + +/-- What a call of `vg_aes_ctr32` on one block needs. -/ +structure CallPre (s : State) (W C D S : Addr) (R : Nat) : Prop where + rdi : s.gpr .rdi = W + rsi : s.gpr .rsi = BitVec.ofNat 64 R + rdx : s.gpr .rdx = C + rcx : s.gpr .rcx = D + r8 : s.gpr .r8 = 1 + r9 : s.gpr .r9 = S + rounds : R = 10 ∨ R = 12 ∨ R = 14 + wc : (⟨W, 240⟩ : Region).Disjoint ⟨C, 16⟩ + wd : (⟨W, 240⟩ : Region).Disjoint ⟨D, 16⟩ + ws : (⟨W, 240⟩ : Region).Disjoint ⟨S, 2048⟩ + cd : (⟨C, 16⟩ : Region).Disjoint ⟨D, 16⟩ + cs : (⟨C, 16⟩ : Region).Disjoint ⟨S, 2048⟩ + ds : (⟨D, 16⟩ : Region).Disjoint ⟨S, 2048⟩ + stkW : (below (s.gpr .rsp) 8).Disjoint ⟨W, 240⟩ + stkC : (below (s.gpr .rsp) 8).Disjoint ⟨C, 16⟩ + stkD : (below (s.gpr .rsp) 8).Disjoint ⟨D, 16⟩ + stkS : (below (s.gpr .rsp) 8).Disjoint ⟨S, 2048⟩ + wrap : D.toNat + 16 ≤ 2 ^ 64 + reads : Covers ([⟨W, 240⟩] ++ [⟨C, 16⟩, ⟨D, 16⟩, ⟨S, 2048⟩]) (s.rd ++ s.wr) + writes : Covers [⟨C, 16⟩, ⟨D, 16⟩, ⟨S, 2048⟩] s.wr + zero : Spec.Aes.bytesAt s.mem 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 : Addr) (R : Nat) (s' : State) : Prop where + rd : s'.rd = s.rd + wr : s'.wr = s.wr + saved : ∀ r ∈ calleeSaved, s'.gpr r = s.gpr r + frame : Frame [⟨C, 16⟩, ⟨D, 16⟩, ⟨S, 2048⟩, below (s.gpr .rsp) 8] s.mem s'.mem + out : Spec.Aes.bytesAt s'.mem D 16 = + Spec.Cmac.aesWith R (Spec.Aes.bytesAt s.mem W (16 * (R + 1))) (Spec.Aes.bytesAt s.mem C 16) + +/-- `vg_aes_ctr32`'s precondition, on entry to a call with the regions it is given. -/ +theorem CallPre.ctr_pre {s : State} {W C D S : Addr} {R : Nat} (h : CallPre s W C D S R) : + Proof.Aes.ctr32X86_64.pre + (s.callEntry.withRegions [⟨W, 240⟩] [⟨C, 16⟩, ⟨D, 16⟩, ⟨S, 2048⟩]) := by + have hR := toNat_rounds h.rounds + simp only [Proof.Aes.ctr32X86_64, State.withRegions_gpr, State.withRegions_rd, + State.withRegions_wr, State.callEntry_rsp, State.callEntry_gpr s (by decide : Reg.rdi ≠ .rsp), + State.callEntry_gpr s (by decide : Reg.rsi ≠ .rsp), State.callEntry_gpr s (by decide : Reg.rdx ≠ .rsp), + State.callEntry_gpr s (by decide : Reg.rcx ≠ .rsp), State.callEntry_gpr s (by decide : Reg.r8 ≠ .rsp), + State.callEntry_gpr s (by decide : Reg.r9 ≠ .rsp), h.rdi, h.rsi, h.rdx, h.rcx, h.r8, h.r9, hR, + one_toNat, Nat.mul_one] + exact ⟨trivial, trivial, h.wc, by simpa using h.wd, h.ws, by simpa using h.cd, h.cs, + by simpa using h.ds, h.stkC, by simpa using h.stkD, h.stkS, by simpa using h.wrap, h.rounds⟩ + +theorem ctr_call (v : Ctr32Impl) {s : State} {W C D S : Addr} {R : Nat} (h : CallPre s W C D S R) : + WP isa (.call v.callee.name v.callee.code) s (CallPost s W C D S R) := by + have hR := toNat_rounds h.rounds + refine WP.call (k := Proof.Aes.ctr32X86_64) v.ok v.nosp (by rw [v.depth]; decide) + (rd := [⟨W, 240⟩]) (wr := [⟨C, 16⟩, ⟨D, 16⟩, ⟨S, 2048⟩]) h.ctr_pre h.reads h.writes ?_ + intro s' hrd hwr hcs hf _ ⟨s₂, hm₂, _, hpost⟩ + rw [v.depth] at hf + refine ⟨hrd, hwr, hcs, ?_, ?_⟩ + · simpa using hf + · obtain ⟨hdata, -⟩ := hpost + simp only [State.withRegions_gpr, State.withRegions_mem, + State.callEntry_gpr s (by decide : Reg.rdi ≠ .rsp), State.callEntry_gpr s (by decide : Reg.rsi ≠ .rsp), + State.callEntry_gpr s (by decide : Reg.rdx ≠ .rsp), State.callEntry_gpr s (by decide : Reg.rcx ≠ .rsp), + State.callEntry_gpr s (by decide : Reg.r8 ≠ .rsp), h.rdi, h.rsi, h.rdx, h.rcx, h.r8, hR, + one_toNat] at hdata + have fE := callEntry_frame s + have hRb : 16 * (R + 1) ≤ 240 := by rcases h.rounds with rfl | rfl | rfl <;> decide + have eW := bytesAt_frame fE (p := W) (n := 16 * (R + 1)) + (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact (h.stkW.sub_right (Region.sub_prefix hRb)).symm) + (by omega) + have eC := bytesAt_frame fE (p := C) (n := 16) + (fun r hr => by simp only [List.mem_singleton] at hr; subst hr; exact h.stkC.symm) (by decide) + have eD := bytesAt_frame fE (p := D) (n := 16) + (fun r hr => by simp only [List.mem_singleton] at hr; subst hr; exact h.stkD.symm) (by decide) + have one : ∀ m : Mem, Spec.Gcm.blocksAt m D 1 = [Spec.Gcm.blockAt m D] := fun m => by + simp [Spec.Gcm.blocksAt] + have bD : Spec.Gcm.blockAt s.callEntry.mem D = 0 := by + rw [Spec.Gcm.blockAt, eD, h.zero, ofBytes_zeros] + have bC : Spec.Gcm.blockAt s.callEntry.mem C = Spec.Gcm.ofBytes (Spec.Aes.bytesAt s.mem C 16) := by + rw [Spec.Gcm.blockAt, eC] + rw [one, one, bD, bC, eW, Proof.Cmac.ctr32_one, List.cons.injEq] at hdata + rw [← hm₂, Proof.Cmac.bytesAt_blockAt, hdata.1, + Proof.Cmac.aesWith_bytes _ _ (Proof.Cmac.bytesAt_length _ _ _)] + +/-- Calls of `vg_aes_ctr32` on one block, with the same arguments in both +runs, are constant time. -/ +theorem ctr_rel (v : Ctr32Impl) {P : State → State → Prop} + (h : ∀ s₁ s₂, P s₁ s₂ → ∃ W C D S : Addr, ∃ R : Nat, + CallPre s₁ W C D S R ∧ CallPre s₂ W C D S R ∧ s₁.gpr .rsp = s₂.gpr .rsp) : + RelCT isa P (.call v.callee.name v.callee.code) fun _ _ => True := by + refine RelCT.callEx v.ok v.ct fun s₁ s₂ hp => ?_ + obtain ⟨W, C, D, S, R, h₁, h₂, hsp⟩ := h s₁ s₂ hp + refine ⟨_, _, _, _, h₁.ctr_pre, h₂.ctr_pre, ?_, h₁.reads, h₁.writes, h₂.reads, h₂.writes, hsp⟩ + simp only [Proof.Aes.ctr32X86_64, State.withRegions_gpr, State.callEntry_rsp, + State.callEntry_gpr _ (by decide : Reg.rdi ≠ .rsp), State.callEntry_gpr _ (by decide : Reg.rsi ≠ .rsp), + State.callEntry_gpr _ (by decide : Reg.rdx ≠ .rsp), State.callEntry_gpr _ (by decide : Reg.rcx ≠ .rsp), + State.callEntry_gpr _ (by decide : Reg.r8 ≠ .rsp), State.callEntry_gpr _ (by decide : Reg.r9 ≠ .rsp), + h₁.rdi, h₁.rsi, h₁.rdx, h₁.rcx, h₁.r8, h₁.r9, h₂.rdi, h₂.rsi, h₂.rdx, h₂.rcx, h₂.r8, h₂.r9, hsp] + exact ⟨trivial, trivial, trivial, trivial, trivial, trivial, trivial⟩ + +end VG.Proof.CmacAes.X86_64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Contract.lean b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Contract.lean new file mode 100644 index 000000000..bd5a40a15 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Contract.lean @@ -0,0 +1,97 @@ +import VerifiedGarbage.Proof.CmacAes.X86_64.Call +import VerifiedGarbage.Impl.CmacAes.X86_64 + +/-! +# AES-CMAC on x86-64: 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 calls `vg_aes_ctr32`, whose return address +is in the 8 bytes below the stack pointer, which may not overlap any buffer. +-/ + +namespace VG.Proof.CmacAes.X86_64 + +open VG VG.X86_64 + +/-- `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 = rdi, rounds = rsi, state = rdx, data = rcx, n = r8, scratch = r9)`. -/ +def updateX86_64 : Contract isa where + pre s := + let sched : Region := ⟨s.gpr .rdi, 240⟩ + let state : Region := ⟨s.gpr .rdx, 16⟩ + let data : Region := ⟨s.gpr .rcx, 16 * (s.gpr .r8).toNat⟩ + let scr : Region := ⟨s.gpr .r9, 2176⟩ + let ret : Region := ⟨s.gpr .rsp, 8⟩ + let stack := below (s.gpr .rsp) 8 + s.rd = [sched, data] ∧ s.wr = [state, scr] ∧ + sched.Disjoint state ∧ sched.Disjoint scr ∧ data.Disjoint state ∧ data.Disjoint scr ∧ + state.Disjoint scr ∧ ret.Disjoint state ∧ ret.Disjoint scr ∧ + stack.Disjoint sched ∧ stack.Disjoint data ∧ stack.Disjoint state ∧ stack.Disjoint scr ∧ + (s.gpr .rdx).toNat + 16 ≤ 2 ^ 64 ∧ (s.gpr .rcx).toNat + 16 * (s.gpr .r8).toNat ≤ 2 ^ 64 ∧ + (s.gpr .r9).toNat + 2176 ≤ 2 ^ 64 ∧ + ((s.gpr .rsi).toNat = 10 ∨ (s.gpr .rsi).toNat = 12 ∨ (s.gpr .rsi).toNat = 14) + post s s' := + Spec.Aes.bytesAt s'.mem (s.gpr .rdx) 16 = + Spec.Cmac.chain (ciphAt s.mem (s.gpr .rdi) (s.gpr .rsi).toNat) (Spec.Aes.bytesAt s.mem (s.gpr .rdx) 16) + (Spec.Cmac.blocksAt s.mem (s.gpr .rcx) 16 (s.gpr .r8).toNat) + pub s₁ s₂ := + s₁.gpr .rdi = s₂.gpr .rdi ∧ s₁.gpr .rsi = s₂.gpr .rsi ∧ s₁.gpr .rdx = s₂.gpr .rdx ∧ + s₁.gpr .rcx = s₂.gpr .rcx ∧ s₁.gpr .r8 = s₂.gpr .r8 ∧ s₁.gpr .r9 = s₂.gpr .r9 ∧ + s₁.gpr .rsp = s₂.gpr .rsp + +/-- `vg_cmac_aes_subkeys(schedule = rdi, rounds = rsi, subkeys = rdx, scratch = rcx)`. -/ +def subkeysX86_64 : Contract isa where + pre s := + let sched : Region := ⟨s.gpr .rdi, 240⟩ + let subk : Region := ⟨s.gpr .rdx, 32⟩ + let scr : Region := ⟨s.gpr .rcx, 2176⟩ + let ret : Region := ⟨s.gpr .rsp, 8⟩ + let stack := below (s.gpr .rsp) 8 + s.rd = [sched] ∧ s.wr = [subk, scr] ∧ + sched.Disjoint subk ∧ sched.Disjoint scr ∧ subk.Disjoint scr ∧ + ret.Disjoint subk ∧ ret.Disjoint scr ∧ + stack.Disjoint sched ∧ stack.Disjoint subk ∧ stack.Disjoint scr ∧ + (s.gpr .rdx).toNat + 32 ≤ 2 ^ 64 ∧ (s.gpr .rcx).toNat + 2176 ≤ 2 ^ 64 ∧ + ((s.gpr .rsi).toNat = 10 ∨ (s.gpr .rsi).toNat = 12 ∨ (s.gpr .rsi).toNat = 14) + post s s' := + let ks := Spec.Cmac.subkeys (ciphAt s.mem (s.gpr .rdi) (s.gpr .rsi).toNat) 16 + Spec.Aes.bytesAt s'.mem (s.gpr .rdx) 32 = ks.1 ++ ks.2 + pub s₁ s₂ := + s₁.gpr .rdi = s₂.gpr .rdi ∧ s₁.gpr .rsi = s₂.gpr .rsi ∧ s₁.gpr .rdx = s₂.gpr .rdx ∧ + s₁.gpr .rcx = s₂.gpr .rcx ∧ s₁.gpr .rsp = s₂.gpr .rsp + +/-- `vg_cmac_aes_finalize(key = rdi, rounds = rsi, state = rdx, last = rcx, last_len = r8, scratch = r9)`. -/ +def finalizeX86_64 : Contract isa where + pre s := + let key : Region := ⟨s.gpr .rdi, 272⟩ + let state : Region := ⟨s.gpr .rdx, 16⟩ + let last : Region := ⟨s.gpr .rcx, (s.gpr .r8).toNat⟩ + let scr : Region := ⟨s.gpr .r9, 2176⟩ + let ret : Region := ⟨s.gpr .rsp, 8⟩ + let stack := below (s.gpr .rsp) 8 + s.rd = [key, last] ∧ s.wr = [state, scr] ∧ + key.Disjoint state ∧ key.Disjoint scr ∧ last.Disjoint state ∧ last.Disjoint scr ∧ + state.Disjoint scr ∧ ret.Disjoint state ∧ ret.Disjoint scr ∧ + stack.Disjoint key ∧ stack.Disjoint last ∧ stack.Disjoint state ∧ stack.Disjoint scr ∧ + (s.gpr .rdi).toNat + 272 ≤ 2 ^ 64 ∧ (s.gpr .rdx).toNat + 16 ≤ 2 ^ 64 ∧ + (s.gpr .rcx).toNat + (s.gpr .r8).toNat ≤ 2 ^ 64 ∧ (s.gpr .r9).toNat + 2176 ≤ 2 ^ 64 ∧ + ((s.gpr .rsi).toNat = 10 ∨ (s.gpr .rsi).toNat = 12 ∨ (s.gpr .rsi).toNat = 14) ∧ + (s.gpr .r8).toNat ≤ 16 + post s s' := + let ciph := ciphAt s.mem (s.gpr .rdi) (s.gpr .rsi).toNat + let ks := Spec.Cmac.subkeys ciph 16 + Spec.Aes.bytesAt s.mem (s.gpr .rdi + 240) 32 = ks.1 ++ ks.2 → + ∀ msg : List Byte, msg.length % 16 = 0 → (msg = [] ∨ 0 < (s.gpr .r8).toNat) → + Spec.Aes.bytesAt s.mem (s.gpr .rdx) 16 = Spec.Cmac.chain ciph (Spec.Cmac.zeros 16) (Spec.Cmac.blocks 16 msg) → + Spec.Aes.bytesAt s'.mem (s.gpr .rdx) 16 = + Spec.Cmac.macFull ciph 16 (msg ++ Spec.Aes.bytesAt s.mem (s.gpr .rcx) (s.gpr .r8).toNat) + pub s₁ s₂ := + s₁.gpr .rdi = s₂.gpr .rdi ∧ s₁.gpr .rsi = s₂.gpr .rsi ∧ s₁.gpr .rdx = s₂.gpr .rdx ∧ + s₁.gpr .rcx = s₂.gpr .rcx ∧ s₁.gpr .r8 = s₂.gpr .r8 ∧ s₁.gpr .r9 = s₂.gpr .r9 ∧ + s₁.gpr .rsp = s₂.gpr .rsp + +end VG.Proof.CmacAes.X86_64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Dbl.lean b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Dbl.lean new file mode 100644 index 000000000..ade5a17fe --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Dbl.lean @@ -0,0 +1,102 @@ +import VerifiedGarbage.Proof.Cmac.Dbl +import VerifiedGarbage.Proof.Cmac.Mem +import VerifiedGarbage.Proof.Gcm.X86_64.Bits + +/-! +# AES-CMAC on x86-64: doubling a block in two 64-bit words + +Untrusted: everything here is checked by Lean. `subkeys` loads a block as +two byte-reversed words, the high and low halves of the block as a +big-endian integer (`Proof.Gcm.X86_64.blockAt_bswap`), doubles the integer a +word at a time (`dbl_words`), and stores the halves byte-reversed again +(`le8_bswap`). +-/ + +namespace VG.Proof.CmacAes.X86_64 + +open VG VG.X86_64 Proof.Cmac + +theorem getD_le8_append (a b : BitVec 64) {k : Nat} (hk : k < 16) : + (le8 a ++ le8 b).getD k 0 = if k < 8 then a.extractLsb' (8 * k) 8 else b.extractLsb' (8 * (k - 8)) 8 := by + rw [List.getD_eq_getElem?_getD] + split + · rw [List.getElem?_append_left (by rw [length_le8]; omega), ← List.getD_eq_getElem?_getD, getD_le8 _ ‹_›] + · rw [List.getElem?_append_right (by rw [length_le8]; omega), length_le8, ← List.getD_eq_getElem?_getD, + getD_le8 _ (by omega)] + +theorem getLsbD_bswap64 (x : BitVec 64) {p : Nat} (hp : p < 64) : + (bswap64 x).getLsbD p = x.getLsbD (8 * (7 - p / 8) + p % 8) := + Proof.Gcm.getLsbD_byteRev64 x p hp + +/-- Storing the byte-reversed halves of `h ++ l` stores its bytes, big-endian. -/ +theorem le8_bswap (h l : BitVec 64) : + le8 (bswap64 h) ++ le8 (bswap64 l) = Spec.Gcm.toBytes (h ++ l) := by + refine ext16 (by simp [length_le8]) (toBytes_length _) fun k hk => ?_ + rw [Proof.Aes.toBytes_getD _ hk] + apply BitVec.eq_of_getLsbD_eq + intro j hj + rw [BitVec.getLsbD_extractLsb', BitVec.getLsbD_append] + simp only [hj, decide_true, Bool.true_and] + rcases Nat.lt_or_ge k 8 with h8 | h8 + · rw [getD_le8_append _ _ hk] + simp only [h8, ↓reduceIte] + rw [BitVec.getLsbD_extractLsb', getLsbD_bswap64 _ (by omega)] + simp only [hj, decide_true, Bool.true_and, show ¬ 8 * (15 - k) + j < 64 by omega, ite_false] + congr 1; omega + · rw [getD_le8_append _ _ hk] + simp only [show ¬ k < 8 by omega, ↓reduceIte] + rw [BitVec.getLsbD_extractLsb', getLsbD_bswap64 _ (by omega)] + simp only [hj, decide_true, Bool.true_and, show 8 * (15 - k) + j < 64 by omega, ite_true] + congr 1; omega + +theorem add_self (x : BitVec 64) : x + x = x <<< 1 := by + apply BitVec.eq_of_toNat_eq + rw [BitVec.toNat_add, BitVec.toNat_shiftLeft, Nat.shiftLeft_eq] + omega + +theorem mask_eq (hi : BitVec 64) : + ((0 : BitVec 64) - (hi >>> 63)) &&& BitVec.signExtend 64 (0x87 : BitVec 32) = if hi.msb then 0x87 else 0 := by + have h : hi >>> 63 = 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 ^ (64 - 1) ≤ hi.toNat + · rw [decide_eq_true hm]; simp; omega + · rw [decide_eq_false hm]; simp; omega + rw [h] + split <;> decide + +theorem bit135 : ∀ p < 64, (135 : BitVec 64).getLsbD p = (135 : BitVec 128).getLsbD p := by decide + +/-- `subkeys`' doubling of `hi ++ lo`. -/ +theorem dbl_words (hi lo : BitVec 64) : + ((hi + hi) ||| (lo >>> 63)) ++ + ((lo + lo) ^^^ (((0 : BitVec 64) - (hi >>> 63)) &&& BitVec.signExtend 64 (0x87 : BitVec 32))) = + dbl128 (hi ++ lo) := by + rw [mask_eq, add_self, add_self, dbl128, BitVec.msb_append] + apply BitVec.eq_of_getLsbD_eq + intro p hp + rw [BitVec.getLsbD_append] + simp only [BitVec.getLsbD_xor, BitVec.getLsbD_or, BitVec.getLsbD_shiftLeft, + BitVec.getLsbD_ushiftRight, BitVec.getLsbD_append, hp, decide_true, Bool.true_and] + have h0 : ((64 : Nat) = 0) = False := by simp + simp only [h0, ite_false] + by_cases h64 : p < 64 + · simp only [h64, ↓reduceIte, show p - 1 < 64 by omega, decide_true, Bool.true_and] + congr 1 + split + · exact bit135 p h64 + · simp + · have hm : (if hi.msb = true then (135 : BitVec 128) else 0).getLsbD p = false := by + split + · exact Proof.Cmac.high_0x87 (by omega) + · simp + rw [hm, Bool.xor_false] + simp only [h64, ↓reduceIte, show p - 64 < 64 by omega, decide_true, Bool.true_and] + rcases Nat.eq_or_lt_of_le (show 64 ≤ p by omega) with rfl | hlt + · simp + · rw [BitVec.getLsbD_of_ge lo (63 + (p - 64)) (by omega)] + simp [show ¬ p - 64 < 1 by omega, show ¬ p < 1 by omega, show ¬ p - 1 < 64 by omega] + congr 1 + +end VG.Proof.CmacAes.X86_64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Finalize.lean b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Finalize.lean new file mode 100644 index 000000000..faee34500 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Finalize.lean @@ -0,0 +1,257 @@ +import VerifiedGarbage.Proof.CmacAes.X86_64.Subkeys +import VerifiedGarbage.Proof.Framework.WriteBytes + +/-! +# AES-CMAC on x86-64: `vg_cmac_aes_finalize`, the last block + +Untrusted: everything here is checked by Lean. The steps that form the +counter block `C ⊕ Mₙ` in the scratch buffer before the call. +-/ + +namespace VG.Proof.CmacAes.X86_64 + +open VG VG.X86_64 VG.X86_64.RegUpd VG.Impl.CmacAes.X86_64 + +/-! ## XORing two blocks, a word at a time -/ + +/-- The memory after storing at `c` the XOR of the blocks at `p` and `q`, +a word at a time. -/ +def xor2Mem (m : Mem) (c p q : Addr) : Mem := + let m₁ := m.writeW c (m.readW p 64 ^^^ m.readW q 64) + m₁.writeW (c + BitVec.ofNat 64 8) (m₁.readW (p + BitVec.ofNat 64 8) 64 ^^^ m₁.readW (q + BitVec.ofNat 64 8) 64) + +theorem xor2Mem_frame (m : Mem) (c p q : Addr) : Frame [⟨c, 16⟩] m (xor2Mem m c p q) := frame_store2 _ _ _ + +theorem xor2Mem_bytes (m : Mem) {c p q : Addr} + (hp : (⟨c, 8⟩ : Region).Disjoint ⟨p + BitVec.ofNat 64 8, 8⟩) + (hq : (⟨c, 8⟩ : Region).Disjoint ⟨q + BitVec.ofNat 64 8, 8⟩) : + Spec.Aes.bytesAt (xor2Mem m c p q) c 16 = + Spec.Cmac.xor (Spec.Aes.bytesAt m p 16) (Spec.Aes.bytesAt m q 16) := by + have g : Frame [⟨c, 8⟩] m (m.writeW c (m.readW p 64 ^^^ m.readW q 64)) := + (Frame.refl _ _).writeW (List.mem_singleton_self _) _ (Region.contains_self _ _) + rw [xor2Mem, Proof.Cmac.bytesAt_store2, + g.readW (r := ⟨p + BitVec.ofNat 64 8, 8⟩) (Region.contains_self _ _) + (fun r hr => by simp only [List.mem_singleton] at hr; subst hr; exact hp.symm) (by decide), + g.readW (r := ⟨q + BitVec.ofNat 64 8, 8⟩) (Region.contains_self _ _) + (fun r hr => by simp only [List.mem_singleton] at hr; subst hr; exact hq.symm) (by decide)] + exact Proof.Cmac.xor_words m p q + +theorem xor_comm (x y : List Byte) : Spec.Cmac.xor x y = Spec.Cmac.xor y x := by + simp only [Spec.Cmac.xor] + exact List.zipWith_comm_of_comm (fun a b => BitVec.xor_comm a b) + +/-! ## Copying the last bytes -/ + +/-- The copy loop's body. -/ +abbrev copyBody : List Instr := + [.movzx8 .rax lastByte, .store8 padByte .rax, .alu .add .r10 (.imm 1), .alu .cmp .r10 (.reg .r8)] + +theorem copyStep_ok (s : State) {P C : Addr} {i L : Nat} (hc : s.gpr .rcx = P) + (h9 : s.gpr .r9 + BitVec.ofNat 64 2048 = C) (hi : s.gpr .r10 = BitVec.ofNat 64 i) + (h8 : s.gpr .r8 = BitVec.ofNat 64 L) + (r : InRegions (s.rd ++ s.wr) (P + BitVec.ofNat 64 i) 1) (w : InRegions s.wr (C + BitVec.ofNat 64 i) 1) : + ∃ s', runBlock isa copyBody s = some s' ∧ + s'.mem = s.mem.writeW (C + BitVec.ofNat 64 i) (s.mem (P + BitVec.ofNat 64 i)) ∧ + s'.gpr .r10 = BitVec.ofNat 64 i + 1 ∧ + s'.zf = some (BitVec.ofNat 64 i + 1 - BitVec.ofNat 64 L == 0) ∧ + (∀ r, r ≠ .rax → r ≠ .r10 → s'.gpr r = s.gpr r) ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + have ea₁ : s.gpr .rcx + s.gpr .r10 * BitVec.ofNat 64 1 + BitVec.ofInt 64 0 = P + BitVec.ofNat 64 i := by + rw [hc, hi, BitVec.mul_one]; simp + have ea₂ : s.gpr .r9 + s.gpr .r10 * BitVec.ofNat 64 1 + BitVec.ofInt 64 (cOff : Int) = C + BitVec.ofNat 64 i := by + rw [hi, BitVec.mul_one, ← h9, BitVec.add_assoc, BitVec.add_assoc, BitVec.add_comm (BitVec.ofNat 64 i)] + rfl + refine ⟨_, by + simp (config := {decide := true}) only [copyBody, lastByte, padByte, runBlock_cons, runStep_some, + runBlock_nil, exec, readSrc, execAlu, State.load8, State.store8, State.ea, Option.bind_some, + Option.map_some, gpr_setReg, gpr_arithFlags, mem_setReg, rd_setReg, wr_setReg, ite_true, ite_false, + ea₁, ea₂, r, w] + rfl, ?_⟩ + refine ⟨?_, ?_, ?_, ?_, ?_, ?_⟩ + · simp only [mem_setReg, mem_arithFlags, BitVec.setWidth_setWidth_of_le _ (show 8 ≤ 64 by decide), + BitVec.setWidth_eq] + · simp [gpr_setReg, hi] + · simp [hi, h8] + · intro r h₁ h₂; simp [gpr_setReg, h₁, h₂] + · rfl + · rfl + +theorem succ_ofNat (i : Nat) : BitVec.ofNat 64 i + 1 = BitVec.ofNat 64 (i + 1) := (BitVec.ofNat_add i 1).symm + +open VG.WriteBytes in +theorem bytesAt_succ (m : Mem) (p : Addr) (i : Nat) : + Spec.Aes.bytesAt m p (i + 1) = Spec.Aes.bytesAt m p i ++ [m (p + BitVec.ofNat 64 i)] := by + simp [Spec.Aes.bytesAt, List.range_succ] + +open VG.WriteBytes in +theorem copy_ok (s : State) {P C : Addr} {L : Nat} (hL₀ : 0 < L) (hL : L < 16) (hc : s.gpr .rcx = P) + (h9 : s.gpr .r9 + BitVec.ofNat 64 2048 = C) (h8 : s.gpr .r8 = BitVec.ofNat 64 L) + (hr : ∀ i < L, InRegions (s.rd ++ s.wr) (P + BitVec.ofNat 64 i) 1) + (hw : ∀ i < 16, InRegions s.wr (C + BitVec.ofNat 64 i) 1) + (hd : (⟨P, L⟩ : Region).Disjoint ⟨C, 16⟩) : + WP isa copy s fun s' => s'.mem = writeBytes s.mem C (Spec.Aes.bytesAt s.mem P L) ∧ + (∀ r, r ≠ .rax → r ≠ .r10 → s'.gpr r = s.gpr r) ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + obtain ⟨s₁, run₁, r10₁, g₁⟩ : ∃ s₁, runBlock isa [.mov32 .r10 (.imm 0)] s = some s₁ ∧ + s₁.gpr .r10 = BitVec.ofNat 64 0 ∧ s₁ = s.setReg .r10 (BitVec.setWidth 64 (0 : BitVec 32)) := + ⟨_, by simp only [runBlock_cons, runStep_some, runBlock_nil, exec, readSrc32, Option.map_some, + State.setReg32], by simp [gpr_setReg], rfl⟩ + refine WP.seq (WP.of_runBlock ⟨s₁, run₁, ?_⟩) + subst g₁ + refine WP.loop (M := isa) (body := .block copyBody) (c := .ne) + (fun (n : Nat) (t : State) => ∃ i, n = L - i ∧ i < L ∧ t.gpr .r10 = BitVec.ofNat 64 i ∧ + t.mem = writeBytes s.mem C (Spec.Aes.bytesAt s.mem P i) ∧ + (∀ r, r ≠ .rax → r ≠ .r10 → t.gpr r = s.gpr r) ∧ t.rd = s.rd ∧ t.wr = s.wr) ?_ (L - 0) _ + ⟨0, rfl, hL₀, r10₁, by simp [Spec.Aes.bytesAt, writeBytes_nil, mem_setReg], + fun r h₁ h₂ => by simp [gpr_setReg, h₂], rfl, rfl⟩ + rintro n t ⟨i, rfl, hi, r10, mem, g, rd, wr⟩ + have tc : t.gpr .rcx = P := by rw [g _ (by decide) (by decide), hc] + have t9 : t.gpr .r9 + BitVec.ofNat 64 2048 = C := by rw [g _ (by decide) (by decide), h9] + have t8 : t.gpr .r8 = BitVec.ofNat 64 L := by rw [g _ (by decide) (by decide), h8] + obtain ⟨t', run', mem', r10', zf', g', rd', wr'⟩ := copyStep_ok t tc t9 r10 t8 + (by rw [rd, wr]; exact hr i hi) (by rw [wr]; exact hw i (by omega)) + refine WP.of_runBlock ⟨t', run', ?_⟩ + have hlen : (Spec.Aes.bytesAt s.mem P i).length = i := by simp [Spec.Aes.bytesAt] + have hx : writeBytes s.mem C (Spec.Aes.bytesAt s.mem P i) (P + BitVec.ofNat 64 i) = s.mem (P + BitVec.ofNat 64 i) := + (writeBytes_frame s.mem C _ (R := ⟨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 P (by omega) (by omega)) (Region.sub_prefix (by omega) _ hcon) + have hmem : t'.mem = writeBytes s.mem C (Spec.Aes.bytesAt s.mem P (i + 1)) := by + rw [mem', mem, hx, bytesAt_succ, writeBytes_snoc s.mem C (Spec.Aes.bytesAt s.mem P i) (s.mem (P + BitVec.ofNat 64 i)) + (by rw [hlen]; omega), hlen] + have hz : t'.zf = some (decide (i + 1 = L)) := by + rw [zf', succ_ofNat, Offset.ofNat_sub_ofNat_beq (by omega) (by omega)] + have gg : ∀ r, r ≠ .rax → r ≠ .r10 → t'.gpr r = s.gpr r := fun r h₁ h₂ => by rw [g' r h₁ h₂, g r h₁ h₂] + by_cases he : i + 1 = L + · left + refine ⟨by simp [eval, hz, he], by rw [hmem, he], gg, by rw [rd', rd], by rw [wr', wr]⟩ + · right + refine ⟨by simp [eval, hz, he], L - (i + 1), by omega, i + 1, rfl, by omega, by rw [r10', succ_ofNat], hmem, gg, + by rw [rd', rd], by rw [wr', wr]⟩ + +/-! ## The straight-line pieces -/ + +/-- Two words XORed from `[p]` and `[q]` into `[c]`, as `full`, `padK2` and +`finArgs` do. -/ +theorem xor2_ok (s : State) (pb qb cb : Reg) (pd qd cd : Nat) {P Q C : Addr} + (hp : s.gpr pb + BitVec.ofNat 64 pd = P) (hp8 : s.gpr pb + BitVec.ofNat 64 (pd + 8) = P + BitVec.ofNat 64 8) + (hq : s.gpr qb + BitVec.ofNat 64 qd = Q) (hq8 : s.gpr qb + BitVec.ofNat 64 (qd + 8) = Q + BitVec.ofNat 64 8) + (hc : s.gpr cb + BitVec.ofNat 64 cd = C) (hc8 : s.gpr cb + BitVec.ofNat 64 (cd + 8) = C + BitVec.ofNat 64 8) + (hrax : pb ≠ .rax ∧ qb ≠ .rax ∧ cb ≠ .rax) + (rp : InRegions (s.rd ++ s.wr) P 8) (rp8 : InRegions (s.rd ++ s.wr) (P + BitVec.ofNat 64 8) 8) + (rq : InRegions (s.rd ++ s.wr) Q 8) (rq8 : InRegions (s.rd ++ s.wr) (Q + BitVec.ofNat 64 8) 8) + (wc : InRegions s.wr C 8) (wc8 : InRegions s.wr (C + BitVec.ofNat 64 8) 8) : + ∃ s', runBlock isa [.mov .rax (.mem (at_ pb pd)), .alu .xor .rax (.mem (at_ qb qd)), .store (at_ cb cd) .rax, + .mov .rax (.mem (at_ pb (pd + 8))), .alu .xor .rax (.mem (at_ qb (qd + 8))), .store (at_ cb (cd + 8)) .rax] + s = some s' ∧ + s'.mem = xor2Mem s.mem C P Q ∧ (∀ r, r ≠ .rax → s'.gpr r = s.gpr r) ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + obtain ⟨h₁, h₂, h₃⟩ := hrax + refine ⟨_, by + simp (config := {decide := true}) only [runBlock_cons, runStep_some, runBlock_nil, at_, exec, readSrc, + execAlu, State.load64, State.store64, State.ea, offset_nat, Option.bind_some, Option.map_some, + gpr_setReg, gpr_arithFlags, mem_setReg, mem_arithFlags, rd_setReg, rd_arithFlags, wr_setReg, wr_arithFlags, + h₁, h₂, h₃, ite_true, ite_false, hp, hp8, hq, hq8, hc, hc8, rp, rp8, rq, rq8, wc, wc8] + rfl, ?_⟩ + refine ⟨rfl, ?_, rfl, rfl⟩ + intro r hr; simp [gpr_setReg, hr] + +theorem cmp16_ok (s : State) {L : Nat} (h8 : s.gpr .r8 = BitVec.ofNat 64 L) (hL : L ≤ 16) : + ∃ s', runBlock isa [.alu .cmp .r8 (.imm 16)] s = some s' ∧ s'.zf = some (decide (L = 16)) ∧ + s'.gpr = s.gpr ∧ s'.mem = s.mem ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by simp only [runBlock_cons, runStep_some, runBlock_nil, exec, execAlu, readSrc, Option.bind_some]; rfl, + ?_⟩ + refine ⟨?_, rfl, rfl, rfl, rfl⟩ + rw [zf_arithFlags, h8, show BitVec.signExtend 64 (16 : BitVec 32) = BitVec.ofNat 64 16 from rfl, + Offset.ofNat_sub_ofNat_beq (by omega) (by decide)] + +/-- The memory after zeroing the block at `c`. -/ +def zero2 (m : Mem) (c : Addr) : Mem := + (m.writeW c (BitVec.setWidth 64 (0 : BitVec 32))).writeW (c + BitVec.ofNat 64 8) (BitVec.setWidth 64 (0 : BitVec 32)) + +theorem zero2_bytes (m : Mem) (c : Addr) : Spec.Aes.bytesAt (zero2 m c) c 16 = Spec.Cmac.zeros 16 := by + rw [zero2, Proof.Cmac.bytesAt_store2, zero_le8, zeros_8_8] + +theorem zero_ok (s : State) {C : Addr} {L : Nat} (hc : s.gpr .r9 + BitVec.ofNat 64 2048 = C) + (hc8 : s.gpr .r9 + BitVec.ofNat 64 2056 = C + BitVec.ofNat 64 8) (h8 : s.gpr .r8 = BitVec.ofNat 64 L) + (hL : L < 2 ^ 64) (wc : InRegions s.wr C 8) (wc8 : InRegions s.wr (C + BitVec.ofNat 64 8) 8) : + ∃ s', runBlock isa zero s = some s' ∧ s'.zf = some (decide (L = 0)) ∧ s'.mem = zero2 s.mem C ∧ + (∀ r, r ≠ .rax → s'.gpr r = s.gpr r) ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + simp (config := {decide := true}) only [zero, cOff, runBlock_cons, runStep_some, runBlock_nil, at_, exec, + readSrc, readSrc32, execAlu, State.store64, State.ea, State.setReg32, offset_nat, Option.bind_some, + Option.map_some, gpr_setReg, mem_setReg, rd_setReg, wr_setReg, ite_true, ite_false, hc, hc8, wc, wc8] + rfl, ?_⟩ + refine ⟨?_, rfl, ?_, rfl, rfl⟩ + · rw [zf_arithFlags] + simp only [h8, BitVec.and_self] + rw [beq_zero hL] + · intro r hr; simp [gpr_setReg, hr] + +theorem pad_ok (s : State) {C : Addr} {L : Nat} (hc : s.gpr .r9 + s.gpr .r8 * BitVec.ofNat 64 1 + + BitVec.ofInt 64 (cOff : Int) = C + BitVec.ofNat 64 L) (wc : InRegions s.wr (C + BitVec.ofNat 64 L) 1) : + ∃ s', runBlock isa [.mov32 .rax (.imm 0x80), .store8 { base := .r9, index := some .r8, disp := cOff } .rax] s = + some s' ∧ s'.mem = s.mem.writeW (C + BitVec.ofNat 64 L) (0x80 : Byte) ∧ + (∀ r, r ≠ .rax → s'.gpr r = s.gpr r) ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + simp (config := {decide := true}) only [runBlock_cons, runStep_some, runBlock_nil, exec, readSrc32, + State.store8, State.ea, State.setReg32, Option.map_some, gpr_setReg, mem_setReg, rd_setReg, wr_setReg, + ite_true, ite_false, hc, wc] + rfl, ?_⟩ + refine ⟨rfl, ?_, rfl, rfl⟩ + intro r hr; simp [gpr_setReg, hr] + +theorem args_ok (s : State) {D : Addr} (hd : s.gpr .rdx = D) + (wd : InRegions s.wr (D + BitVec.ofNat 64 0) 8) (wd8 : InRegions s.wr (D + BitVec.ofNat 64 8) 8) : + ∃ s', runBlock isa [.mov32 .rax (.imm 0), .store (at_ .rdx 0) .rax, .store (at_ .rdx 8) .rax, + .mov .rcx (.reg .rdx), .mov .rdx (.reg .r9), .alu .add .rdx (.imm (BitVec.ofNat 32 cOff)), + .mov32 .r8 (.imm 1)] s = some s' ∧ + s'.mem = (s.mem.writeW (D + BitVec.ofNat 64 0) (BitVec.setWidth 64 (0 : BitVec 32))).writeW + (D + BitVec.ofNat 64 8) (BitVec.setWidth 64 (0 : BitVec 32)) ∧ + s'.gpr .rcx = D ∧ s'.gpr .rdx = s.gpr .r9 + BitVec.ofNat 64 2048 ∧ s'.gpr .r8 = 1 ∧ + (∀ r, r ≠ .rax → r ≠ .rcx → r ≠ .rdx → r ≠ .r8 → s'.gpr r = s.gpr r) ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + simp (config := {decide := true}) only [cOff, runBlock_cons, runStep_some, runBlock_nil, at_, exec, + readSrc, readSrc32, execAlu, State.store64, State.ea, State.setReg32, offset_nat, Option.bind_some, + Option.map_some, gpr_setReg, mem_setReg, rd_setReg, wr_setReg, ite_true, ite_false, hd, wd, wd8] + rfl, ?_⟩ + refine ⟨rfl, ?_, ?_, ?_, ?_, rfl, rfl⟩ + · simp [gpr_setReg] + · simp [gpr_setReg] + · simp [gpr_setReg] + · intro r h₁ h₂ h₃ h₄; simp [gpr_setReg, h₁, h₂, h₃, h₄] + +open VG.WriteBytes in +/-- The padded last block `Mₙ* ‖ 10ʲ`, from the bytes copied onto zeros. -/ +theorem padded_bytes (m : Mem) (C : Addr) (xs : List Byte) (hL : xs.length < 16) + (hz : Spec.Aes.bytesAt m C 16 = Spec.Cmac.zeros 16) : + Spec.Aes.bytesAt ((writeBytes m C xs).writeW (C + BitVec.ofNat 64 xs.length) (0x80 : Byte)) C 16 = + xs ++ [0x80] ++ Spec.Cmac.zeros (16 - xs.length - 1) := by + refine Proof.Cmac.ext16 (by simp [Spec.Aes.bytesAt]) (by simp [Spec.Cmac.zeros]; omega) fun k hk => ?_ + rw [Proof.Cmac.getD_bytesAt _ _ hk, writeW8_apply] + have hz' : m (C + BitVec.ofNat 64 k) = 0 := by + have := congrArg (fun l => l.getD k 0) hz + rw [Proof.Cmac.getD_bytesAt _ _ hk] at this + rw [this]; simp only [Spec.Cmac.zeros, List.getD_eq_getElem?_getD, List.getElem?_replicate, hk, + ite_true, Option.getD_some] + have hsub : (C + BitVec.ofNat 64 k - C).toNat = k := Mem.sub_ofNat_toNat C (by omega) + have heq : (C + BitVec.ofNat 64 k = C + BitVec.ofNat 64 xs.length) ↔ k = xs.length := by + constructor + · intro h + have := congrArg (fun a => (a - C).toNat) h + simp only [Mem.sub_ofNat_toNat C (show k < 2 ^ 64 by omega), + Mem.sub_ofNat_toNat C (show xs.length < 2 ^ 64 by omega)] at this + exact this + · intro h; rw [h] + rcases Nat.lt_trichotomy k xs.length with h | h | h + · have hne : ¬ (C + BitVec.ofNat 64 k = C + BitVec.ofNat 64 xs.length) := by rw [heq]; omega + simp only [hne, ite_false, writeBytes, hsub, h, ite_true] + simp [List.getD_eq_getElem?_getD, List.getElem?_append_left h] + · subst h + simp [List.getD_eq_getElem?_getD] + · have hne : ¬ (C + BitVec.ofNat 64 k = C + BitVec.ofNat 64 xs.length) := by rw [heq]; omega + simp only [hne, ite_false, writeBytes, hsub, show ¬ k < xs.length by omega, hz'] + obtain ⟨j, hj⟩ : ∃ j, k - xs.length = j + 1 := ⟨k - xs.length - 1, by omega⟩ + rw [List.getD_eq_getElem?_getD, List.append_assoc, List.getElem?_append_right (show xs.length ≤ k by omega), + hj, List.singleton_append, List.getElem?_cons_succ, Spec.Cmac.zeros, List.getElem?_replicate] + simp only [show j < 16 - xs.length - 1 by omega, ite_true, Option.getD_some] + +end VG.Proof.CmacAes.X86_64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/X86_64/FinalizeCT.lean b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/FinalizeCT.lean new file mode 100644 index 000000000..536121cb5 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/FinalizeCT.lean @@ -0,0 +1,47 @@ +import VerifiedGarbage.Proof.CmacAes.X86_64.FinalizeCorrect +import VerifiedGarbage.Proof.Framework.X86_64.Taint + +/-! +# AES-CMAC on x86-64: `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 (its branches and the copy loop depend only on +`last_len`), and the call of `vg_aes_ctr32` is constant time by its own proof +(`ctr_rel`), its arguments pinned by the correctness proof (`FMid`). +-/ + +namespace VG.Proof.CmacAes.X86_64 + +open VG VG.X86_64 VG.Impl.CmacAes.X86_64 +open VG.Proof.Aes.X86_64 (Ctr32Impl) + +theorem finalize_rel (v : Ctr32Impl) {s₀ s₀' : State} (h0 : finalizeX86_64.pre s₀) + (h0' : finalizeX86_64.pre s₀') (hq : finalizeX86_64.pub s₀ s₀') : + RelCT isa (fun a b => a = s₀ ∧ b = s₀') (finalize v.callee) fun _ _ => True := by + obtain ⟨q1, q2, q3, q4, q5, q6, q7⟩ := hq + have hp := FPre.of h0 + have hp' : FPre s₀' (s₀.gpr .rdi) (s₀.gpr .rdx) (s₀.gpr .rcx) (s₀.gpr .r9) (s₀.gpr .r8).toNat + (s₀.gpr .rsi).toNat := by + rw [q1, q2, q3, q4, q5, q6]; exact FPre.of h0' + obtain ⟨_, hA⟩ : ∃ h, (taint.check (Taint.ofRegs [.rdi, .rsi, .rdx, .rcx, .r8, .r9, .rsp]) finPre h).isSome = + true := ⟨_, by taint_decide⟩ + have a := (RelCT.taint (A := taint) (P := fun a b => a = s₀ ∧ b = s₀') _ + (fun a b h => by + obtain ⟨rfl, rfl⟩ := h + 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 | rfl | rfl | rfl <;> assumption) hA).wp + (F₁ := FMid s₀ _ _ _ _ _ _) (F₂ := FMid s₀' _ _ _ _ _ _) fun a b h => by + obtain ⟨rfl, rfl⟩ := h; exact ⟨finPre_wp hp, finPre_wp hp'⟩ + have c := ctr_rel v (P := fun s₁ s₂ => + FMid s₀ (s₀.gpr .rdi) (s₀.gpr .rdx) (s₀.gpr .rcx) (s₀.gpr .r9) (s₀.gpr .r8).toNat (s₀.gpr .rsi).toNat s₁ ∧ + FMid s₀' (s₀.gpr .rdi) (s₀.gpr .rdx) (s₀.gpr .rcx) (s₀.gpr .r9) (s₀.gpr .r8).toNat (s₀.gpr .rsi).toNat s₂) + fun s₁ s₂ h => ⟨_, _, _, _, _, h.1.pre, h.2.pre, by + rw [h.1.saved _ (by simp [calleeSaved]), h.2.saved _ (by simp [calleeSaved]), q7]⟩ + exact (a.mono (fun _ _ h => h) fun _ _ h => h.2).seq c + +theorem finalize_ct (v : Ctr32Impl) : + ConstantTime isa finalizeX86_64.pre finalizeX86_64.pub (finalize v.callee) := + fun _ _ _ _ _ _ h₁ h₂ hq e₁ e₂ => (finalize_rel v h₁ h₂ hq _ _ _ _ _ _ ⟨rfl, rfl⟩ e₁ e₂).1 + +end VG.Proof.CmacAes.X86_64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/X86_64/FinalizeCorrect.lean b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/FinalizeCorrect.lean new file mode 100644 index 000000000..081157176 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/FinalizeCorrect.lean @@ -0,0 +1,386 @@ +import VerifiedGarbage.Proof.CmacAes.X86_64.Finalize + +/-! +# AES-CMAC on x86-64: `vg_cmac_aes_finalize` is correct + +Untrusted: everything here is checked by Lean. Before the call, the +counter block (at `S + 2048`) 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 (`macFull_split`). +-/ + +namespace VG.Proof.CmacAes.X86_64 + +open VG VG.X86_64 VG.X86_64.RegUpd VG.Impl.CmacAes.X86_64 +open VG.Proof.Aes.X86_64 (Ctr32Impl) + +/-- The precondition, by name: the key (schedule and subkeys) `W`, the state +`St`, the last bytes `P` (`L` of them), the scratch buffer `S` and the +rounds `R`. -/ +structure FPre (s₀ : State) (W St P S : Addr) (L R : Nat) : Prop where + rdi : s₀.gpr .rdi = W + rdx : s₀.gpr .rdx = St + rcx : s₀.gpr .rcx = P + r8 : (s₀.gpr .r8).toNat = L + r9 : s₀.gpr .r9 = S + rsi : (s₀.gpr .rsi).toNat = R + rd : s₀.rd = [⟨W, 272⟩, ⟨P, L⟩] + wr : s₀.wr = [⟨St, 16⟩, ⟨S, 2176⟩] + key_st : (⟨W, 272⟩ : Region).Disjoint ⟨St, 16⟩ + key_scr : (⟨W, 272⟩ : Region).Disjoint ⟨S, 2176⟩ + last_st : (⟨P, L⟩ : Region).Disjoint ⟨St, 16⟩ + last_scr : (⟨P, L⟩ : Region).Disjoint ⟨S, 2176⟩ + st_scr : (⟨St, 16⟩ : Region).Disjoint ⟨S, 2176⟩ + ret_st : (⟨s₀.gpr .rsp, 8⟩ : Region).Disjoint ⟨St, 16⟩ + ret_scr : (⟨s₀.gpr .rsp, 8⟩ : Region).Disjoint ⟨S, 2176⟩ + stk_key : (below (s₀.gpr .rsp) 8).Disjoint ⟨W, 272⟩ + stk_last : (below (s₀.gpr .rsp) 8).Disjoint ⟨P, L⟩ + stk_st : (below (s₀.gpr .rsp) 8).Disjoint ⟨St, 16⟩ + stk_scr : (below (s₀.gpr .rsp) 8).Disjoint ⟨S, 2176⟩ + key_wrap : W.toNat + 272 ≤ 2 ^ 64 + st_wrap : St.toNat + 16 ≤ 2 ^ 64 + last_wrap : P.toNat + L ≤ 2 ^ 64 + scr_wrap : S.toNat + 2176 ≤ 2 ^ 64 + rounds : R = 10 ∨ R = 12 ∨ R = 14 + len : L ≤ 16 + +theorem FPre.of {s₀ : State} (h : finalizeX86_64.pre s₀) : + FPre s₀ (s₀.gpr .rdi) (s₀.gpr .rdx) (s₀.gpr .rcx) (s₀.gpr .r9) (s₀.gpr .r8).toNat (s₀.gpr .rsi).toNat := + let ⟨a, b, c, d, e, f, g, h, i, j, k, l, m, n, o, p, q, r, s⟩ := h + ⟨rfl, rfl, rfl, rfl, rfl, rfl, a, b, c, d, e, f, g, h, i, j, k, l, m, n, o, p, q, r, s⟩ + +/-- The last block `Mₙ` (§6.2 step 4), from the key and the last bytes in `m`. -/ +abbrev mn (m : Mem) (W P : Addr) (L : Nat) : List Byte := + Spec.Cmac.lastBlock 16 (Spec.Aes.bytesAt m (W + BitVec.ofNat 64 240) 16) + (Spec.Aes.bytesAt m (W + BitVec.ofNat 64 256) 16) (Spec.Aes.bytesAt m P L) + +/-- What the branch on the length leaves: `Mₙ` in the counter block. -/ +structure BPost (s₀ : State) (W St P S : Addr) (L : Nat) (s : State) : Prop where + rdi : s.gpr .rdi = W + rdx : s.gpr .rdx = St + r9 : s.gpr .r9 = S + rsi : s.gpr .rsi = s₀.gpr .rsi + saved : ∀ r ∈ calleeSaved, s.gpr r = s₀.gpr r + rd : s.rd = s₀.rd + wr : s.wr = s₀.wr + frame : Frame [⟨S + BitVec.ofNat 64 2048, 16⟩] s₀.mem s.mem + blk : Spec.Aes.bytesAt s.mem (S + BitVec.ofNat 64 2048) 16 = mn s₀.mem W P L + +section +variable {s₀ : State} {W St P S : Addr} {L R : Nat} (hp : FPre s₀ W St P S L R) +include hp + +theorem FPre.inScr {d n : Nat} (h : d + n ≤ 2176) : InRegions s₀.wr (S + BitVec.ofNat 64 d) n := by + rw [hp.wr]; exact in_rw (r := ⟨S, 2176⟩) (by simp) (Offset.contains_base _ h (by have := hp.scr_wrap; omega)) + +theorem FPre.inKey {d n : Nat} (h : d + n ≤ 272) : InRegions (s₀.rd ++ s₀.wr) (W + BitVec.ofNat 64 d) n := by + rw [hp.rd]; exact in_rw (r := ⟨W, 272⟩) (by simp) (Offset.contains_base _ h (by have := hp.key_wrap; omega)) + +theorem FPre.inLast {d n : Nat} (h : d + n ≤ L) : InRegions (s₀.rd ++ s₀.wr) (P + BitVec.ofNat 64 d) n := by + rw [hp.rd]; exact in_rw (r := ⟨P, L⟩) (by simp) (Offset.contains_base _ h (by have := hp.len; omega)) + +end + +theorem FPre.scrD {S : Addr} {d n : Nat} (h : d + n ≤ 2176) : Region.Sub ⟨S + BitVec.ofNat 64 d, n⟩ ⟨S, 2176⟩ := + Offset.sub_base _ h + +theorem wr_in {s : State} {a : Addr} {n : Nat} (h : InRegions s.wr a n) : InRegions (s.rd ++ s.wr) a n := by + obtain ⟨r, hr, hc⟩ := h; exact ⟨r, List.mem_append_right _ hr, hc⟩ + +theorem full_wp {s₀ : State} {W St P S : Addr} {L R : Nat} (hp : FPre s₀ W St P S L R) (hL : L = 16) + {s : State} (hg : s.gpr = s₀.gpr) (hm : s.mem = s₀.mem) (hrd : s.rd = s₀.rd) (hwr : s.wr = s₀.wr) : + WP isa (.block full) s (BPost s₀ W St P S L) := by + subst hL + have e : full = [.mov .rax (.mem (at_ .rcx 0)), .alu .xor .rax (.mem (at_ .rdi 240)), .store (at_ .r9 2048) .rax, + .mov .rax (.mem (at_ .rcx (0 + 8))), .alu .xor .rax (.mem (at_ .rdi (240 + 8))), + .store (at_ .r9 (2048 + 8)) .rax] := rfl + have sw := hp.scr_wrap + have kw := hp.key_wrap + obtain ⟨s', run, mem, g, rd, wr⟩ := xor2_ok s .rcx .rdi .r9 0 240 2048 + (P := P) (Q := W + BitVec.ofNat 64 240) (C := S + BitVec.ofNat 64 2048) + (by rw [hg, hp.rcx, k0]) (by rw [hg, hp.rcx]) + (by rw [hg, hp.rdi]) (by rw [hg, hp.rdi, Offset.add_add]) + (by rw [hg, hp.r9]) (by rw [hg, hp.r9, Offset.add_add]) + ⟨by decide, by decide, by decide⟩ + (by rw [hrd, hwr]; simpa using hp.inLast (d := 0) (n := 8) (by decide)) + (by rw [hrd, hwr]; exact hp.inLast (d := 8) (n := 8) (by decide)) + (by rw [hrd, hwr]; exact hp.inKey (d := 240) (n := 8) (by decide)) + (by rw [hrd, hwr, Offset.add_add]; exact hp.inKey (d := 248) (n := 8) (by decide)) + (by rw [hwr]; exact hp.inScr (d := 2048) (n := 8) (by decide)) + (by rw [hwr, Offset.add_add]; exact hp.inScr (d := 2056) (n := 8) (by decide)) + rw [e] + refine WP.of_runBlock ⟨s', run, ?_⟩ + have gg (r : Reg) (hr : r ≠ .rax) : s'.gpr r = s₀.gpr r := by rw [g r hr, hg] + refine ⟨by rw [gg _ (by decide), hp.rdi], by rw [gg _ (by decide), hp.rdx], by rw [gg _ (by decide), hp.r9], + gg _ (by decide), fun r hr => gg r (by rintro rfl; simp [calleeSaved] at hr), by rw [rd, hrd], + by rw [wr, hwr], by rw [mem, hm]; exact xor2Mem_frame _ _ _ _, ?_⟩ + rw [mem, hm, xor2Mem_bytes] + · simp only [mn, Spec.Cmac.lastBlock, Proof.Cmac.bytesAt_length, ite_true] + exact xor_comm _ _ + · exact (hp.last_scr.symm.sub_left (FPre.scrD (d := 2048) (n := 8) (by decide))).sub_right + (Offset.sub_base P (d := 8) (n := 8) (k := 16) (by decide)) + · rw [Offset.add_add] + exact (hp.key_scr.symm.sub_left (FPre.scrD (d := 2048) (n := 8) (by decide))).sub_right + (Offset.sub_base W (d := 248) (n := 8) (k := 272) (by decide)) + +open VG.WriteBytes in +theorem partial_wp {s₀ : State} {W St P S : Addr} {L R : Nat} (hp : FPre s₀ W St P S L R) (hL : L < 16) + {s : State} (hg : s.gpr = s₀.gpr) (hm : s.mem = s₀.mem) (hrd : s.rd = s₀.rd) (hwr : s.wr = s₀.wr) : + WP isa partialBlock s (BPost s₀ W St P S L) := by + have sw := hp.scr_wrap + have kw := hp.key_wrap + have h8 : s.gpr .r8 = BitVec.ofNat 64 L := by + rw [hg, ← hp.r8]; apply BitVec.eq_of_toNat_eq; simp + obtain ⟨C, hC⟩ : ∃ C, S + BitVec.ofNat 64 2048 = C := ⟨_, rfl⟩ + have hc : s.gpr .r9 + BitVec.ofNat 64 2048 = C := by rw [hg, hp.r9, hC] + have hc8 : s.gpr .r9 + BitVec.ofNat 64 2056 = C + BitVec.ofNat 64 8 := by rw [hg, hp.r9, ← hC, Offset.add_add] + have dPC : (⟨P, L⟩ : Region).Disjoint ⟨C, 16⟩ := by rw [← hC]; exact hp.last_scr.sub_right (FPre.scrD (by decide)) + have dKC (d n : Nat) (h : d + n ≤ 272) : (⟨W + BitVec.ofNat 64 d, n⟩ : Region).Disjoint ⟨C, 16⟩ := by + rw [← hC]; exact (hp.key_scr.sub_left (Offset.sub_base _ h)).sub_right (FPre.scrD (by decide)) + -- Zero the block. + obtain ⟨s₁, run₁, zf₁, mem₁, g₁, rd₁, wr₁⟩ := zero_ok s hc hc8 h8 (by omega) + (by rw [hwr, ← hC]; exact hp.inScr (d := 2048) (n := 8) (by decide)) + (by rw [hwr, ← hC, Offset.add_add]; exact hp.inScr (d := 2056) (n := 8) (by decide)) + refine WP.seq (WP.of_runBlock ⟨s₁, run₁, ?_⟩) + have zf : s₁.mem = zero2 s₀.mem C := by rw [mem₁, hm] + have fz : Frame [⟨C, 16⟩] s₀.mem (zero2 s₀.mem C) := frame_store2 _ _ _ + have lastZ : Spec.Aes.bytesAt (zero2 s₀.mem C) P L = Spec.Aes.bytesAt s₀.mem P L := + bytesAt_frame fz (fun r hr => by simp only [List.mem_singleton] at hr; subst hr; exact dPC) (by omega) + -- Copy the last bytes. + refine WP.seq (WP.mono (Q := fun (s₂ : State) => s₂.mem = writeBytes (zero2 s₀.mem C) C (Spec.Aes.bytesAt s₀.mem P L) ∧ + (∀ r, r ≠ .rax → r ≠ .r10 → s₂.gpr r = s₀.gpr r) ∧ s₂.rd = s₀.rd ∧ s₂.wr = s₀.wr) ?_ fun s₂ h₂ => ?_) + · by_cases hL0 : L = 0 + · subst hL0 + refine WP.ite true (by show s₁.zf = _; rw [zf₁]; rfl) (fun _ => WP.block_nil ?_) (fun h => by cases h) + refine ⟨by rw [zf]; simp [Spec.Aes.bytesAt, writeBytes_nil], fun r h₁ _ => by rw [g₁ r h₁, hg], + by rw [rd₁, hrd], by rw [wr₁, hwr]⟩ + · refine WP.ite false (by show s₁.zf = _; rw [zf₁]; simp [hL0]) (fun h => by cases h) fun _ => ?_ + refine WP.mono (copy_ok s₁ (by omega) hL (by rw [g₁ _ (by decide), hg, hp.rcx]) + (by rw [g₁ _ (by decide)]; exact hc) (by rw [g₁ _ (by decide)]; exact h8) + (fun i hi => by rw [rd₁, wr₁, hrd, hwr]; exact hp.inLast (d := i) (n := 1) (by omega)) + (fun i hi => by + rw [wr₁, hwr, ← hC, Offset.add_add]; exact hp.inScr (d := 2048 + i) (n := 1) (by omega)) dPC) ?_ + rintro s₂ ⟨m₂, g₂, rd₂, wr₂⟩ + refine ⟨by rw [m₂, zf, lastZ], fun r h₁ h₂ => by rw [g₂ r h₁ h₂, g₁ r h₁, hg], by rw [rd₂, rd₁, hrd], + by rw [wr₂, wr₁, hwr]⟩ + · obtain ⟨m₂, g₂, rd₂, wr₂⟩ := h₂ + have e : padK2 = [.mov32 .rax (.imm 0x80), .store8 { base := .r9, index := some .r8, disp := cOff } .rax] ++ + [.mov .rax (.mem (at_ .r9 2048)), .alu .xor .rax (.mem (at_ .rdi 256)), .store (at_ .r9 2048) .rax, + .mov .rax (.mem (at_ .r9 (2048 + 8))), .alu .xor .rax (.mem (at_ .rdi (256 + 8))), + .store (at_ .r9 (2048 + 8)) .rax] := rfl + rw [e, WP.block_append_iff] + have r9₂ : s₂.gpr .r9 = S := by rw [g₂ _ (by decide) (by decide), hp.r9] + have r8₂ : s₂.gpr .r8 = BitVec.ofNat 64 L := by rw [g₂ _ (by decide) (by decide), ← hg]; exact h8 + obtain ⟨s₃, run₃, m₃, g₃, rd₃, wr₃⟩ := pad_ok s₂ (C := C) (L := L) + (by rw [r9₂, r8₂, BitVec.mul_one, offset_nat, ← hC, BitVec.add_assoc, BitVec.add_comm (BitVec.ofNat 64 L), + ← BitVec.add_assoc]; rfl) + (by rw [wr₂, ← hC, Offset.add_add]; exact hp.inScr (d := 2048 + L) (n := 1) (by omega)) + refine WP.of_runBlock ⟨s₃, run₃, ?_⟩ + have r9₃ : s₃.gpr .r9 = S := by rw [g₃ _ (by decide), r9₂] + have rdi₃ : s₃.gpr .rdi = W := by rw [g₃ _ (by decide), g₂ _ (by decide) (by decide), hp.rdi] + obtain ⟨s₄, run₄, m₄, g₄, rd₄, wr₄⟩ := xor2_ok s₃ .r9 .rdi .r9 2048 256 2048 + (P := C) (Q := W + BitVec.ofNat 64 256) (C := C) + (by rw [r9₃, hC]) (by rw [r9₃, ← hC, Offset.add_add]) (by rw [rdi₃]) (by rw [rdi₃, Offset.add_add]) + (by rw [r9₃, hC]) (by rw [r9₃, ← hC, Offset.add_add]) ⟨by decide, by decide, by decide⟩ + (by rw [rd₃, wr₃, rd₂, wr₂, ← hC]; exact wr_in (hp.inScr (d := 2048) (n := 8) (by decide))) + (by rw [rd₃, wr₃, rd₂, wr₂, ← hC, Offset.add_add]; exact wr_in (hp.inScr (d := 2056) (n := 8) (by decide))) + (by rw [rd₃, wr₃, rd₂, wr₂]; exact hp.inKey (d := 256) (n := 8) (by decide)) + (by rw [rd₃, wr₃, rd₂, wr₂, Offset.add_add]; exact hp.inKey (d := 264) (n := 8) (by decide)) + (by rw [wr₃, wr₂, ← hC]; exact hp.inScr (d := 2048) (n := 8) (by decide)) + (by rw [wr₃, wr₂, ← hC, Offset.add_add]; exact hp.inScr (d := 2056) (n := 8) (by decide)) + refine WP.of_runBlock ⟨s₄, run₄, ?_⟩ + have gg (r : Reg) (h₁ : r ≠ .rax) (h₂ : r ≠ .r10) : s₄.gpr r = s₀.gpr r := by rw [g₄ r h₁, g₃ r h₁, g₂ r h₁ h₂] + have hlen : (Spec.Aes.bytesAt s₀.mem P L).length = L := Proof.Cmac.bytesAt_length _ _ _ + have fW : Frame [⟨C, 16⟩] (writeBytes (zero2 s₀.mem C) C (Spec.Aes.bytesAt s₀.mem P L)) s₃.mem := by + rw [m₃, m₂] + exact (Frame.refl _ _).writeW (List.mem_singleton_self _) _ (Offset.contains_base _ (by omega) (by omega)) + have fB : Frame [⟨C, 16⟩] (zero2 s₀.mem C) (writeBytes (zero2 s₀.mem C) C (Spec.Aes.bytesAt s₀.mem P L)) := + writeBytes_frame _ _ _ (by + rw [hlen]; simpa using Offset.contains_base C (d := 0) (n := L) (k := 16) (by omega) (by decide)) + have f₃ : Frame [⟨C, 16⟩] s₀.mem s₃.mem := (fz.trans fB).trans fW + have k2 : Spec.Aes.bytesAt s₃.mem (W + BitVec.ofNat 64 256) 16 = Spec.Aes.bytesAt s₀.mem (W + BitVec.ofNat 64 256) 16 := + bytesAt_frame' f₃ fun r hr => by simp only [List.mem_singleton] at hr; subst hr; exact dKC 256 16 (by decide) + have pad : Spec.Aes.bytesAt s₃.mem C 16 = + Spec.Aes.bytesAt s₀.mem P L ++ [0x80] ++ Spec.Cmac.zeros (16 - L - 1) := by + have := padded_bytes (zero2 s₀.mem C) C (Spec.Aes.bytesAt s₀.mem P L) (by rw [hlen]; exact hL) (zero2_bytes _ _) + rw [hlen] at this + rw [m₃, m₂]; exact this + refine ⟨by rw [gg _ (by decide) (by decide), hp.rdi], by rw [gg _ (by decide) (by decide), hp.rdx], + by rw [gg _ (by decide) (by decide), hp.r9], gg _ (by decide) (by decide), + fun r hr => gg r (by rintro rfl; simp [calleeSaved] at hr) (by rintro rfl; simp [calleeSaved] at hr), + by rw [rd₄, rd₃, rd₂], by rw [wr₄, wr₃, wr₂], ?_, ?_⟩ + · rw [m₄, hC]; exact f₃.trans (xor2Mem_frame _ _ _ _) + · rw [m₄, hC, xor2Mem_bytes, pad, k2] + · simp only [mn, Spec.Cmac.lastBlock, hlen, show L ≠ 16 by omega, ite_false] + exact xor_comm _ _ + · simpa using Offset.disjoint C (d := 0) (n := 8) (e := 8) (k := 8) (by decide) (by decide) (by decide) + · rw [Offset.add_add] + exact (dKC 264 8 (by decide)).symm.sub_left (Region.sub_prefix (by decide)) + +/-! ## Up to the call -/ + +/-- What the code before the call leaves. -/ +structure FMid (s₀ : State) (W St P S : Addr) (L R : Nat) (s : State) : Prop where + pre : CallPre s W (S + BitVec.ofNat 64 2048) St S R + blk : Spec.Aes.bytesAt s.mem (S + BitVec.ofNat 64 2048) 16 = + Spec.Cmac.xor (mn s₀.mem W P L) (Spec.Aes.bytesAt s₀.mem St 16) + frame : Frame [⟨S + BitVec.ofNat 64 2048, 16⟩, ⟨St, 16⟩] s₀.mem s.mem + saved : ∀ r ∈ calleeSaved, s.gpr r = s₀.gpr r + +theorem finArgs_wp {s₀ : State} {W St P S : Addr} {L R : Nat} (hp : FPre s₀ W St P S L R) {s : State} + (h : BPost s₀ W St P S L s) : WP isa (.block finArgs) s (FMid s₀ W St P S L R) := by + have sw := hp.scr_wrap + have tw := hp.st_wrap + have e : finArgs = [.mov .rax (.mem (at_ .r9 2048)), .alu .xor .rax (.mem (at_ .rdx 0)), .store (at_ .r9 2048) .rax, + .mov .rax (.mem (at_ .r9 (2048 + 8))), .alu .xor .rax (.mem (at_ .rdx (0 + 8))), + .store (at_ .r9 (2048 + 8)) .rax] ++ + [.mov32 .rax (.imm 0), .store (at_ .rdx 0) .rax, .store (at_ .rdx 8) .rax, + .mov .rcx (.reg .rdx), .mov .rdx (.reg .r9), .alu .add .rdx (.imm (BitVec.ofNat 32 cOff)), + .mov32 .r8 (.imm 1)] := rfl + rw [e, WP.block_append_iff] + have hRegs : s.rd ++ s.wr = [⟨W, 272⟩, ⟨P, L⟩, ⟨St, 16⟩, ⟨S, 2176⟩] := by rw [h.rd, h.wr, hp.rd, hp.wr]; rfl + have inC (d : Nat) (hd : d + 8 ≤ 2176) : InRegions s.wr (S + BitVec.ofNat 64 d) 8 := by + rw [h.wr]; exact hp.inScr (by omega) + have inSt (d : Nat) (hd : d + 8 ≤ 16) : InRegions s.wr (St + BitVec.ofNat 64 d) 8 := by + rw [h.wr, hp.wr]; exact in_rw (r := ⟨St, 16⟩) (by simp) (Offset.contains_base _ hd (by omega)) + obtain ⟨s₁, run₁, m₁, g₁, rd₁, wr₁⟩ := xor2_ok s .r9 .rdx .r9 2048 0 2048 + (P := S + BitVec.ofNat 64 2048) (Q := St) (C := S + BitVec.ofNat 64 2048) + (by rw [h.r9]) (by rw [h.r9, Offset.add_add]) (by rw [h.rdx, k0]) (by rw [h.rdx]) + (by rw [h.r9]) (by rw [h.r9, Offset.add_add]) ⟨by decide, by decide, by decide⟩ + (wr_in (inC 2048 (by decide))) (by rw [Offset.add_add]; exact wr_in (inC 2056 (by decide))) + (by simpa using wr_in (inSt 0 (by decide))) (wr_in (inSt 8 (by decide))) + (inC 2048 (by decide)) (by rw [Offset.add_add]; exact inC 2056 (by decide)) + refine WP.of_runBlock ⟨s₁, run₁, ?_⟩ + obtain ⟨s₂, run₂, m₂, rcx₂, rdx₂, r8₂, g₂, rd₂, wr₂⟩ := args_ok s₁ (D := St) + (by rw [g₁ _ (by decide), h.rdx]) (by rw [wr₁]; exact inSt 0 (by decide)) (by rw [wr₁]; exact inSt 8 (by decide)) + refine WP.of_runBlock ⟨s₂, run₂, ?_⟩ + have g (r : Reg) (h₁ : r ≠ .rax) (h₂ : r ≠ .rcx) (h₃ : r ≠ .rdx) (h₄ : r ≠ .r8) : s₂.gpr r = s.gpr r := by + rw [g₂ r h₁ h₂ h₃ h₄, g₁ r h₁] + have rsp₂ : s₂.gpr .rsp = s₀.gpr .rsp := by + rw [g _ (by decide) (by decide) (by decide) (by decide), h.saved _ (by simp [calleeSaved])] + have dCSt : (⟨S + BitVec.ofNat 64 2048, 16⟩ : Region).Disjoint ⟨St, 16⟩ := + hp.st_scr.symm.sub_left (FPre.scrD (by decide)) + have fA : Frame [⟨St, 16⟩] s₁.mem s₂.mem := by + rw [m₂, k0]; exact frame_store2 _ _ _ + have stS : Spec.Aes.bytesAt s.mem St 16 = Spec.Aes.bytesAt s₀.mem St 16 := + bytesAt_frame' h.frame fun r hr => by simp only [List.mem_singleton] at hr; subst hr; exact dCSt.symm + have hR := hp.rounds + refine ⟨?_, ?_, ?_, ?_⟩ + · exact + { rdi := by rw [g _ (by decide) (by decide) (by decide) (by decide), h.rdi] + rsi := by + rw [g _ (by decide) (by decide) (by decide) (by decide), h.rsi] + apply BitVec.eq_of_toNat_eq; simp [hp.rsi]; omega + rdx := by rw [rdx₂, g₁ _ (by decide), h.r9] + rcx := rcx₂ + r8 := r8₂ + r9 := by rw [g _ (by decide) (by decide) (by decide) (by decide), h.r9] + rounds := hR + wc := (hp.key_scr.sub_left (Region.sub_prefix (by decide))).sub_right (FPre.scrD (by decide)) + 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 := dCSt + cs := Offset.disjoint_base _ (by decide) (by have := hp.scr_wrap; omega) + ds := hp.st_scr.sub_right (Region.sub_prefix (by decide)) + stkW := by rw [rsp₂]; exact hp.stk_key.sub_right (Region.sub_prefix (by decide)) + stkC := by rw [rsp₂]; exact hp.stk_scr.sub_right (FPre.scrD (by decide)) + stkD := by rw [rsp₂]; exact hp.stk_st + stkS := by rw [rsp₂]; exact hp.stk_scr.sub_right (Region.sub_prefix (by decide)) + wrap := hp.st_wrap + reads := by + rw [rd₂, wr₂, rd₁, wr₁, hRegs] + refine Covers.of_sub fun r hr => ?_ + 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 + · exact ⟨⟨W, 272⟩, by simp, 0, by simp, by simp⟩ + · exact ⟨⟨S, 2176⟩, by simp, 2048, rfl, by simp⟩ + · exact ⟨⟨St, 16⟩, by simp, 0, by simp, by simp⟩ + · exact ⟨⟨S, 2176⟩, by simp, 0, by simp, by simp⟩ + writes := by + rw [wr₂, wr₁, h.wr, hp.wr] + 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 ⟨⟨S, 2176⟩, by simp, 2048, rfl, by simp⟩ + · exact ⟨⟨St, 16⟩, by simp, 0, by simp, by simp⟩ + · exact ⟨⟨S, 2176⟩, by simp, 0, by simp, by simp⟩ + zero := by rw [m₂, k0, Proof.Cmac.bytesAt_store2, zero_le8, zeros_8_8] } + · rw [bytesAt_frame' fA (fun r hr => by simp only [List.mem_singleton] at hr; subst hr; exact dCSt), m₁, + xor2Mem_bytes, h.blk, stS] + · simpa using Offset.disjoint (S + BitVec.ofNat 64 2048) (d := 0) (n := 8) (e := 8) (k := 8) (by decide) + (by decide) (by decide) + · exact (dCSt.sub_left (Region.sub_prefix (by decide))).sub_right (Offset.sub_base _ (by decide)) + · refine (h.frame.trans (m₁ ▸ xor2Mem_frame _ _ _ _)).mono (by simp) |>.trans (fA.mono (by simp)) + · intro r hr + rw [g r (by rintro rfl; simp [calleeSaved] at hr) (by rintro rfl; simp [calleeSaved] at hr) + (by rintro rfl; simp [calleeSaved] at hr) (by rintro rfl; simp [calleeSaved] at hr), h.saved r hr] + +theorem finPre_wp {s₀ : State} {W St P S : Addr} {L R : Nat} (hp : FPre s₀ W St P S L R) : + WP isa finPre s₀ (FMid s₀ W St P S L R) := by + have h8 : s₀.gpr .r8 = BitVec.ofNat 64 L := by rw [← hp.r8]; apply BitVec.eq_of_toNat_eq; simp + obtain ⟨s₁, run₁, zf₁, g₁, m₁, rd₁, wr₁⟩ := cmp16_ok s₀ h8 hp.len + refine WP.seq (WP.of_runBlock ⟨s₁, run₁, ?_⟩) + refine WP.seq (WP.mono (Q := BPost s₀ W St P S L) ?_ fun _ h => finArgs_wp hp h) + by_cases hL : L = 16 + · exact WP.ite true (by show s₁.zf = _; rw [zf₁]; simp [hL]) (fun _ => full_wp hp hL g₁ m₁ rd₁ wr₁) + (fun h => by cases h) + · exact WP.ite false (by show s₁.zf = _; rw [zf₁]; simp [hL]) (fun h => by cases h) + (fun _ => partial_wp hp (by have := hp.len; omega) g₁ m₁ rd₁ wr₁) + +theorem k1k2 {m : Mem} {W : Addr} {k1 k2 : List Byte} (h1 : k1.length = 16) + (h : Spec.Aes.bytesAt m (W + BitVec.ofNat 64 240) 32 = k1 ++ k2) : + Spec.Aes.bytesAt m (W + BitVec.ofNat 64 240) 16 = k1 ∧ Spec.Aes.bytesAt m (W + BitVec.ofNat 64 256) 16 = k2 := by + rw [bytesAt_32, Offset.add_add] at h + exact List.append_inj h (by rw [Proof.Cmac.bytesAt_length, h1]) + +theorem finalize_wp (v : Ctr32Impl) {s₀ : State} (h0 : finalizeX86_64.pre s₀) : + WP isa (finalize v.callee) s₀ fun s' => gprPreserved s₀ s' ∧ finalizeX86_64.post s₀ s' := by + have hp := FPre.of h0 + generalize s₀.gpr .rdi = W at hp + generalize s₀.gpr .rdx = St at hp + generalize s₀.gpr .rcx = P at hp + generalize s₀.gpr .r9 = S at hp + generalize (s₀.gpr .r8).toNat = L at hp + generalize (s₀.gpr .rsi).toNat = R at hp + have hR := hp.rounds + have hRb : 16 * (R + 1) ≤ 240 := by rcases hR with h | h | h <;> omega + refine WP.seq (WP.mono (finPre_wp hp) fun s₁ h₁ => ?_) + refine WP.mono (ctr_call v h₁.pre) fun s₂ h₂ => ?_ + have rsp₁ : s₁.gpr .rsp = s₀.gpr .rsp := h₁.saved _ (by simp [calleeSaved]) + -- The key and the frame. + have big : Frame [⟨St, 16⟩, ⟨S, 2176⟩, below (s₀.gpr .rsp) 8] s₀.mem s₂.mem := by + refine (h₁.frame.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 ⟨⟨S, 2176⟩, by simp, FPre.scrD (by decide)⟩ + · exact ⟨⟨St, 16⟩, 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 ⟨⟨S, 2176⟩, by simp, FPre.scrD (by decide)⟩ + · exact ⟨⟨St, 16⟩, by simp, fun _ h => h⟩ + · exact ⟨⟨S, 2176⟩, by simp, Region.sub_prefix (by decide)⟩ + · exact ⟨below (s₀.gpr .rsp) 8, by simp, by rw [rsp₁]; exact fun _ h => h⟩ + have sch : Spec.Aes.bytesAt s₁.mem W (16 * (R + 1)) = Spec.Aes.bytesAt s₀.mem W (16 * (R + 1)) := + bytesAt_frame h₁.frame (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))).sub_right (FPre.scrD (by decide)) + · exact hp.key_st.sub_left (Region.sub_prefix (by omega))) (by omega) + refine ⟨⟨fun r hr => by rw [h₂.saved r hr, h₁.saved r hr], ?_⟩, ?_⟩ + · refine big.readW (r := ⟨s₀.gpr .rsp, 8⟩) (Region.contains_self _ _) (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.ret_st + · exact hp.ret_scr + · exact Offset.base_disjoint_below _ (by decide) + · intro hk msg hm hne hst + rw [hp.rdi, hp.rsi] at hk hst ⊢ + rw [hp.rdx] at hst ⊢ + rw [hp.r8] at hne + rw [hp.rcx, hp.r8] + obtain ⟨e1, e2⟩ := k1k2 (Proof.Cmac.subkeys_aes_length _ _) hk + rw [h₂.out, sch, 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), xor_comm] + +end VG.Proof.CmacAes.X86_64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Subkeys.lean b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Subkeys.lean new file mode 100644 index 000000000..8649665d8 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Subkeys.lean @@ -0,0 +1,400 @@ +import VerifiedGarbage.Proof.CmacAes.X86_64.Dbl +import VerifiedGarbage.Proof.CmacAes.X86_64.UpdateCorrect + +/-! +# AES-CMAC on x86-64: `vg_cmac_aes_subkeys` + +Untrusted: everything here is checked by Lean. +-/ + +namespace VG.Proof.CmacAes.X86_64 + +open VG VG.X86_64 VG.X86_64.RegUpd VG.Impl.CmacAes.X86_64 +open VG.Proof.Aes.X86_64 (Ctr32Impl) + +/-! ## Doubling a block -/ + +/-- The words `dbl` stores, from the halves `hi` and `lo` it loads. -/ +def dblHi (hi lo : BitVec 64) : BitVec 64 := (hi + hi) ||| (lo >>> 63) +def dblLo (hi lo : BitVec 64) : BitVec 64 := + (lo + lo) ^^^ (((0 : BitVec 32).setWidth 64 - (hi >>> 63)) &&& BitVec.signExtend 64 (0x87 : BitVec 32)) + +/-- The memory after `dbl src dst`, from `rbx = K`. -/ +def dblMem (m : Mem) (K : Addr) (src dst : Nat) : Mem := + let hi := bswap64 (m.readW (K + BitVec.ofNat 64 src) 64) + let lo := bswap64 (m.readW (K + BitVec.ofNat 64 (src + 8)) 64) + (m.writeW (K + BitVec.ofNat 64 dst) (bswap64 (dblHi hi lo))).writeW (K + BitVec.ofNat 64 (dst + 8)) + (bswap64 (dblLo hi lo)) + +theorem dbl_ok (s : State) {K : Addr} (hb : s.gpr .rbx = K) {src dst : Nat} + (r₀ : InRegions (s.rd ++ s.wr) (K + BitVec.ofNat 64 src) 8) + (r₁ : InRegions (s.rd ++ s.wr) (K + BitVec.ofNat 64 (src + 8)) 8) + (w₀ : InRegions s.wr (K + BitVec.ofNat 64 dst) 8) (w₁ : InRegions s.wr (K + BitVec.ofNat 64 (dst + 8)) 8) : + ∃ s', runBlock isa (dbl src dst) s = some s' ∧ s'.mem = dblMem s.mem K src dst ∧ + (∀ r, r ≠ .rax → r ≠ .rdx → r ≠ .rcx → r ≠ .r8 → s'.gpr r = s.gpr r) ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + simp (config := {decide := true}) only [dbl, runBlock_cons, runStep_some, runBlock_nil, at_, exec, + readSrc, readSrc32, execAlu, execShift, State.load64, State.store64, State.ea, State.setReg32, offset_nat, + Option.bind_some, Option.map_some, gpr_setReg, gpr_arithFlags, gpr_setFlags, mem_setReg, mem_arithFlags, + mem_setFlags, rd_setReg, rd_arithFlags, rd_setFlags, wr_setReg, wr_arithFlags, wr_setFlags, + ite_true, ite_false, hb, r₀, r₁, w₀, w₁] + rfl, ?_⟩ + refine ⟨?_, ?_, rfl, rfl⟩ + · rfl + · intro r h₁ h₂ h₃ h₄ + simp [gpr_setReg, gpr_setFlags, h₁, h₂, h₃, h₄] + +theorem dbl_words' (hi lo : BitVec 64) : dblHi hi lo ++ dblLo hi lo = Proof.Cmac.dbl128 (hi ++ lo) := by + rw [dblHi, dblLo, show (0 : BitVec 32).setWidth 64 = 0 from rfl] + exact dbl_words hi lo + +theorem dblMem_frame (m : Mem) (K : Addr) (src dst : Nat) : + Frame [⟨K + BitVec.ofNat 64 dst, 16⟩] m (dblMem m K src dst) := by + rw [dblMem, show K + BitVec.ofNat 64 (dst + 8) = K + BitVec.ofNat 64 dst + BitVec.ofNat 64 8 from + (Offset.add_add _ _ _).symm] + exact frame_store2 _ _ _ + +theorem dblMem_bytes (m : Mem) (K : Addr) (src dst : Nat) : + Spec.Aes.bytesAt (dblMem m K src dst) (K + BitVec.ofNat 64 dst) 16 = + Spec.Cmac.dbl 16 (Spec.Aes.bytesAt m (K + BitVec.ofNat 64 src) 16) := by + rw [dblMem, show K + BitVec.ofNat 64 (dst + 8) = K + BitVec.ofNat 64 dst + BitVec.ofNat 64 8 from + (Offset.add_add _ _ _).symm, + show K + BitVec.ofNat 64 (src + 8) = K + BitVec.ofNat 64 src + BitVec.ofNat 64 8 from + (Offset.add_add _ _ _).symm, + Proof.Cmac.bytesAt_store2, le8_bswap, dbl_words', Proof.Cmac.dbl_eq (Proof.Cmac.bytesAt_length _ _ _), + ← Spec.Gcm.blockAt, ← Proof.Gcm.X86_64.blockAt_bswap, BitVec.add_zero] + +/-! ## Before the call -/ + +/-- The memory after `subkeysPre`. -/ +def preMem (s : State) : Mem := + (((((s.mem.writeW (s.gpr .rcx + BitVec.ofNat 64 2064) (s.gpr .rbx)).writeW + (s.gpr .rcx + BitVec.ofNat 64 2072) (s.gpr .rbp)).writeW + (s.gpr .rcx + BitVec.ofNat 64 2048) (BitVec.setWidth 64 (0 : BitVec 32))).writeW + (s.gpr .rcx + BitVec.ofNat 64 2056) (BitVec.setWidth 64 (0 : BitVec 32))).writeW + (s.gpr .rdx + BitVec.ofNat 64 0) (BitVec.setWidth 64 (0 : BitVec 32))).writeW + (s.gpr .rdx + BitVec.ofNat 64 8) (BitVec.setWidth 64 (0 : BitVec 32)) + +theorem subkeysPre_ok (s : State) + (w₁ : InRegions s.wr (s.gpr .rcx + BitVec.ofNat 64 2064) 8) + (w₂ : InRegions s.wr (s.gpr .rcx + BitVec.ofNat 64 2072) 8) + (w₃ : InRegions s.wr (s.gpr .rcx + BitVec.ofNat 64 2048) 8) + (w₄ : InRegions s.wr (s.gpr .rcx + BitVec.ofNat 64 2056) 8) + (w₅ : InRegions s.wr (s.gpr .rdx + BitVec.ofNat 64 0) 8) + (w₆ : InRegions s.wr (s.gpr .rdx + BitVec.ofNat 64 8) 8) : + ∃ s', runBlock isa subkeysPre s = some s' ∧ + s'.gpr .rdi = s.gpr .rdi ∧ s'.gpr .rsi = s.gpr .rsi ∧ + s'.gpr .rdx = s.gpr .rcx + BitVec.ofNat 64 2048 ∧ s'.gpr .rcx = s.gpr .rdx ∧ s'.gpr .r8 = 1 ∧ + s'.gpr .r9 = s.gpr .rcx ∧ s'.gpr .rbx = s.gpr .rdx ∧ s'.gpr .rbp = s.gpr .rcx ∧ + (∀ r ∈ calleeSaved, r ≠ .rbx → r ≠ .rbp → s'.gpr r = s.gpr r) ∧ + s'.mem = preMem s ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + simp (config := {decide := true}) only [subkeysPre, ctrArgs, cOff, List.cons_append, List.nil_append, + runBlock_cons, runStep_some, runBlock_nil, at_, exec, readSrc, readSrc32, execAlu, State.store64, + State.ea, State.setReg32, offset_nat, Option.bind_some, Option.map_some, gpr_setReg, gpr_arithFlags, + mem_setReg, rd_setReg, wr_setReg, ite_true, ite_false, w₁, w₂, w₃, w₄, w₅, w₆] + rfl, ?_⟩ + refine ⟨?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_⟩ + all_goals first + | rfl + | (simp [gpr_setReg]; done) + | (intro r hr h₁ h₂ + simp only [calleeSaved, List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl | rfl | rfl | rfl <;> simp_all [gpr_setReg]) + +/-! ## The whole function -/ + +/-- The precondition, by name: the schedule `W`, the subkeys `K`, the scratch +buffer `S` and the rounds `R`. -/ +structure SPre (s₀ : State) (W K S : Addr) (R : Nat) : Prop where + rdi : s₀.gpr .rdi = W + rdx : s₀.gpr .rdx = K + rcx : s₀.gpr .rcx = S + rsi : (s₀.gpr .rsi).toNat = R + rd : s₀.rd = [⟨W, 240⟩] + wr : s₀.wr = [⟨K, 32⟩, ⟨S, 2176⟩] + sch_k : (⟨W, 240⟩ : Region).Disjoint ⟨K, 32⟩ + sch_scr : (⟨W, 240⟩ : Region).Disjoint ⟨S, 2176⟩ + k_scr : (⟨K, 32⟩ : Region).Disjoint ⟨S, 2176⟩ + ret_k : (⟨s₀.gpr .rsp, 8⟩ : Region).Disjoint ⟨K, 32⟩ + ret_scr : (⟨s₀.gpr .rsp, 8⟩ : Region).Disjoint ⟨S, 2176⟩ + stk_sch : (below (s₀.gpr .rsp) 8).Disjoint ⟨W, 240⟩ + stk_k : (below (s₀.gpr .rsp) 8).Disjoint ⟨K, 32⟩ + stk_scr : (below (s₀.gpr .rsp) 8).Disjoint ⟨S, 2176⟩ + k_wrap : K.toNat + 32 ≤ 2 ^ 64 + scr_wrap : S.toNat + 2176 ≤ 2 ^ 64 + rounds : R = 10 ∨ R = 12 ∨ R = 14 + +theorem SPre.of {s₀ : State} (h : subkeysX86_64.pre s₀) : + SPre s₀ (s₀.gpr .rdi) (s₀.gpr .rdx) (s₀.gpr .rcx) (s₀.gpr .rsi).toNat := + let ⟨a, b, c, d, e, f, g, h, i, j, k, l, m⟩ := h + ⟨rfl, rfl, rfl, rfl, a, b, c, d, e, f, g, h, i, j, k, l, m⟩ + +theorem bytesAt_32 (m : Mem) (p : Addr) : + Spec.Aes.bytesAt m p 32 = Spec.Aes.bytesAt m p 16 ++ Spec.Aes.bytesAt m (p + BitVec.ofNat 64 16) 16 := by + simp only [Spec.Aes.bytesAt] + rw [show (32 : Nat) = 16 + 16 from rfl, 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] + +theorem k0 (K : Addr) : K + BitVec.ofNat 64 0 = K := BitVec.add_zero K + +theorem zeros_8_8 : Spec.Cmac.zeros 8 ++ Spec.Cmac.zeros 8 = Spec.Cmac.zeros 16 := by decide + +theorem zero_le8 : Proof.Cmac.le8 (BitVec.setWidth 64 (0 : BitVec 32)) = Spec.Cmac.zeros 8 := by decide + +theorem scr_sub' {S : Addr} {d n : Nat} (h : d + n ≤ 2176) : + Region.Sub ⟨S + BitVec.ofNat 64 d, n⟩ ⟨S, 2176⟩ := Offset.sub_base _ h + + +theorem preMem_frame (s : State) : + Frame [⟨s.gpr .rcx + BitVec.ofNat 64 2048, 32⟩, ⟨s.gpr .rdx, 16⟩] s.mem (preMem s) := by + have c (d : Nat) (h : d + 8 ≤ 32) : + (⟨s.gpr .rcx + BitVec.ofNat 64 2048, 32⟩ : Region).Contains (s.gpr .rcx + BitVec.ofNat 64 (2048 + d)) 8 := by + rw [← Offset.add_add]; exact Offset.contains_base _ h (by omega) + have k (d : Nat) (h : d + 8 ≤ 16) : (⟨s.gpr .rdx, 16⟩ : Region).Contains (s.gpr .rdx + BitVec.ofNat 64 d) 8 := + Offset.contains_base _ h (by omega) + exact ((((((Frame.refl _ _).writeW (by simp) _ (c 16 (by decide))).writeW (by simp) _ (c 24 (by decide))).writeW + (by simp) _ (c 0 (by decide))).writeW (by simp) _ (c 8 (by decide))).writeW (by simp) _ (k 0 (by decide))).writeW + (by simp) _ (k 8 (by decide)) + +theorem frame_store2' {m : Mem} (p : Addr) (w₀ w₁ : BitVec 64) : + Frame [⟨p, 16⟩] m ((m.writeW (p + BitVec.ofNat 64 0) w₀).writeW (p + BitVec.ofNat 64 8) w₁) := by + rw [k0]; exact frame_store2 _ _ _ + +theorem restore2_ok (s : State) {B : Addr} (hb : s.gpr .rbp = B) + (r₁ : InRegions (s.rd ++ s.wr) (B + BitVec.ofNat 64 2064) 8) + (r₂ : InRegions (s.rd ++ s.wr) (B + BitVec.ofNat 64 2072) 8) : + ∃ s', runBlock isa [.mov .rbx (.mem (at_ .rbp 2064)), .mov .rbp (.mem (at_ .rbp 2072))] s = some s' ∧ + s'.gpr .rbx = s.mem.readW (B + BitVec.ofNat 64 2064) 64 ∧ + s'.gpr .rbp = s.mem.readW (B + BitVec.ofNat 64 2072) 64 ∧ + (∀ r, r ≠ .rbx → r ≠ .rbp → s'.gpr r = s.gpr r) ∧ s'.mem = s.mem := by + refine ⟨_, by + simp (config := {decide := true}) only [runBlock_cons, runStep_some, runBlock_nil, at_, exec, readSrc, + State.load64, State.ea, offset_nat, gpr_setReg, mem_setReg, rd_setReg, wr_setReg, ite_true, ite_false, + Option.map_some, hb, r₁, r₂] + rfl, ?_⟩ + refine ⟨?_, ?_, ?_, rfl⟩ + · simp [gpr_setReg] + · simp [gpr_setReg] + · intro r h₁ h₂; simp [gpr_setReg, h₁, h₂] + +theorem preMem_slot (s : State) {d : Nat} (hd : d = 2064 ∨ d = 2072) + (hks : (⟨s.gpr .rdx, 32⟩ : Region).Disjoint ⟨s.gpr .rcx, 2176⟩) : + (preMem s).readW (s.gpr .rcx + BitVec.ofNat 64 d) 64 = if d = 2064 then s.gpr .rbx else s.gpr .rbp := by + have kd (e : Nat) (he : e + 8 ≤ 32) : Mem.Sep (s.gpr .rcx + BitVec.ofNat 64 d) (64 / 8) + (s.gpr .rdx + BitVec.ofNat 64 e) (64 / 8) := + hks.symm.sep (Offset.contains_base _ (by omega) (by omega)) (Offset.contains_base _ he (by omega)) + rw [preMem, Mem.readW_writeW_sep (kd 8 (by decide)) (by decide), Mem.readW_writeW_sep (kd 0 (by decide)) (by decide), + readW_writeW_other _ _ _ (by omega) (by omega) (by decide), + readW_writeW_other _ _ _ (by omega) (by omega) (by decide)] + rcases hd with rfl | rfl + · rw [readW_writeW_other _ _ _ (by decide) (by decide) (by decide), Mem.readW_writeW_self64]; rfl + · rw [Mem.readW_writeW_self64]; rfl + +theorem callPre_of {s₀ : State} {W K S : Addr} {R : Nat} (hp : SPre s₀ W K S R) {s₁ : State} + (rdi₁ : s₁.gpr .rdi = s₀.gpr .rdi) (rsi₁ : s₁.gpr .rsi = s₀.gpr .rsi) + (rdx₁ : s₁.gpr .rdx = s₀.gpr .rcx + BitVec.ofNat 64 2048) (rcx₁ : s₁.gpr .rcx = s₀.gpr .rdx) + (r8₁ : s₁.gpr .r8 = 1) (r9₁ : s₁.gpr .r9 = s₀.gpr .rcx) (rsp₁ : s₁.gpr .rsp = s₀.gpr .rsp) + (mem₁ : s₁.mem = preMem s₀) (rd₁ : s₁.rd = s₀.rd) (wr₁ : s₁.wr = s₀.wr) : + CallPre s₁ W (S + BitVec.ofNat 64 2048) K S R := by + have hR := hp.rounds + have kw := hp.k_wrap + have sw := hp.scr_wrap + have cK : (⟨K, 16⟩ : Region).Disjoint ⟨S + BitVec.ofNat 64 2048, 16⟩ := + (hp.k_scr.sub_left (Region.sub_prefix (by decide))).sub_right (scr_sub' (by decide)) + have zK : Spec.Aes.bytesAt s₁.mem K 16 = Spec.Cmac.zeros 16 := by + rw [mem₁, preMem, hp.rdx, k0, Proof.Cmac.bytesAt_store2, zero_le8, zeros_8_8] + exact + { rdi := by rw [rdi₁, hp.rdi] + rsi := by rw [rsi₁]; apply BitVec.eq_of_toNat_eq; simp [hp.rsi]; omega + rdx := by rw [rdx₁, hp.rcx] + rcx := by rw [rcx₁, hp.rdx] + r8 := r8₁ + r9 := by rw [r9₁, hp.rcx] + rounds := hR + wc := hp.sch_scr.sub_right (scr_sub' (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 := cK.symm + cs := 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)) + stkW := by rw [rsp₁]; exact hp.stk_sch + stkC := by rw [rsp₁]; exact hp.stk_scr.sub_right (scr_sub' (by decide)) + stkD := by rw [rsp₁]; exact hp.stk_k.sub_right (Region.sub_prefix (by decide)) + stkS := by rw [rsp₁]; exact hp.stk_scr.sub_right (Region.sub_prefix (by decide)) + wrap := by omega + reads := by + rw [rd₁, wr₁, hp.rd, hp.wr] + refine Covers.of_sub fun r hr => ?_ + 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 + · exact ⟨⟨W, 240⟩, by simp, 0, by simp, by simp⟩ + · exact ⟨⟨S, 2176⟩, by simp, 2048, rfl, by simp⟩ + · exact ⟨⟨K, 32⟩, by simp, 0, by simp, by simp⟩ + · exact ⟨⟨S, 2176⟩, by simp, 0, by simp, by simp⟩ + writes := by + rw [wr₁, hp.wr] + 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 ⟨⟨S, 2176⟩, by simp, 2048, rfl, by simp⟩ + · exact ⟨⟨K, 32⟩, by simp, 0, by simp, by simp⟩ + · exact ⟨⟨S, 2176⟩, by simp, 0, by simp, by simp⟩ + zero := zK } + +theorem subkeys_wp (v : Ctr32Impl) {s₀ : State} (h0 : subkeysX86_64.pre s₀) : + WP isa (subkeys v.callee) s₀ fun s' => gprPreserved s₀ s' ∧ subkeysX86_64.post s₀ s' := by + have hp := SPre.of h0 + generalize s₀.gpr .rdi = W at hp + generalize s₀.gpr .rdx = K at hp + generalize s₀.gpr .rcx = S at hp + generalize (s₀.gpr .rsi).toNat = R at hp + have hR := hp.rounds + have hRb : 16 * (R + 1) ≤ 240 := by rcases hR with h | h | h <;> omega + have kw := hp.k_wrap + have sw := hp.scr_wrap + have inS (d : Nat) (h : d + 8 ≤ 2176) : InRegions s₀.wr (S + BitVec.ofNat 64 d) 8 := by + rw [hp.wr]; exact in_rw (r := ⟨S, 2176⟩) (by simp) (Offset.contains_base _ h (by omega)) + have inK (d : Nat) (h : d + 8 ≤ 32) : InRegions s₀.wr (K + BitVec.ofNat 64 d) 8 := by + rw [hp.wr]; exact in_rw (r := ⟨K, 32⟩) (by simp) (Offset.contains_base _ h (by omega)) + -- Before the call. + obtain ⟨s₁, run₁, rdi₁, rsi₁, rdx₁, rcx₁, r8₁, r9₁, rbx₁, rbp₁, cs₁, mem₁, rd₁, wr₁⟩ := + subkeysPre_ok s₀ (by rw [hp.rcx]; exact inS _ (by decide)) (by rw [hp.rcx]; exact inS _ (by decide)) + (by rw [hp.rcx]; exact inS _ (by decide)) (by rw [hp.rcx]; exact inS _ (by decide)) + (by rw [hp.rdx]; exact inK _ (by decide)) (by rw [hp.rdx]; exact inK _ (by decide)) + refine WP.seq (WP.of_runBlock ⟨s₁, run₁, ?_⟩) + -- The memory before the call. + have f₁ : Frame [⟨S + BitVec.ofNat 64 2048, 32⟩, ⟨K, 16⟩] s₀.mem s₁.mem := by + rw [mem₁, ← hp.rcx, ← hp.rdx]; exact preMem_frame s₀ + have cK : (⟨K, 16⟩ : Region).Disjoint ⟨S + BitVec.ofNat 64 2048, 16⟩ := + (hp.k_scr.sub_left (Region.sub_prefix (by decide))).sub_right (scr_sub' (by decide)) + have zC : Spec.Aes.bytesAt s₁.mem (S + BitVec.ofNat 64 2048) 16 = Spec.Cmac.zeros 16 := by + rw [mem₁, preMem, hp.rcx, hp.rdx] + rw [bytesAt_frame' (frame_store2' K _ _) (by + intro r hr; simp only [List.mem_singleton] at hr; subst hr; exact cK.symm)] + rw [(Offset.add_add_eq S (a := 2048) (b := 8) (c := 2056) rfl).symm, Proof.Cmac.bytesAt_store2, zero_le8, + zeros_8_8] + have zK : Spec.Aes.bytesAt s₁.mem K 16 = Spec.Cmac.zeros 16 := by + rw [mem₁, preMem, hp.rdx, k0, Proof.Cmac.bytesAt_store2, zero_le8, zeros_8_8] + have schB : ∀ m : Mem, Frame [⟨S + BitVec.ofNat 64 2048, 32⟩, ⟨K, 16⟩] s₀.mem m → + Spec.Aes.bytesAt m W (16 * (R + 1)) = Spec.Aes.bytesAt s₀.mem W (16 * (R + 1)) := fun m hf => + bytesAt_frame hf (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)).sub_right (scr_sub' (by decide)) + · exact (hp.sch_k.sub_left (Region.sub_prefix hRb)).sub_right (Region.sub_prefix (by decide))) (by omega) + -- The call. + have rsp₁ : s₁.gpr .rsp = s₀.gpr .rsp := cs₁ .rsp (by simp [calleeSaved]) (by decide) (by decide) + have pre := callPre_of hp rdi₁ rsi₁ rdx₁ rcx₁ r8₁ r9₁ rsp₁ mem₁ rd₁ wr₁ + refine WP.seq (WP.mono (ctr_call v pre) fun s₂ h₂ => ?_) + -- After the call. + have rbx₂ : s₂.gpr .rbx = K := by rw [h₂.saved .rbx (by simp [calleeSaved]), rbx₁, hp.rdx] + have rbp₂ : s₂.gpr .rbp = S := by rw [h₂.saved .rbp (by simp [calleeSaved]), rbp₁, hp.rcx] + 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 rIn (a : Addr) (h : InRegions s₀.wr a 8) : InRegions (s₀.rd ++ s₀.wr) a 8 := by + obtain ⟨r, hr, hc⟩ := h; exact ⟨r, List.mem_append_right _ hr, hc⟩ + rw [subkeysPost, WP.block_append_iff, WP.block_append_iff] + obtain ⟨s₃, run₃, mem₃, g₃, rd₃, wr₃⟩ := dbl_ok s₂ rbx₂ (src := 0) (dst := 0) + (by rw [rdwr₂]; exact rIn _ (inK 0 (by decide))) (by rw [rdwr₂]; exact rIn _ (inK 8 (by decide))) + (by rw [wr₂]; exact inK 0 (by decide)) (by rw [wr₂]; exact inK 8 (by decide)) + refine WP.of_runBlock ⟨s₃, run₃, ?_⟩ + have rbx₃ : s₃.gpr .rbx = K := by rw [g₃ _ (by decide) (by decide) (by decide) (by decide), rbx₂] + obtain ⟨s₄, run₄, mem₄, g₄, rd₄, wr₄⟩ := dbl_ok s₃ rbx₃ (src := 0) (dst := 16) + (by rw [rd₃, wr₃, rdwr₂]; exact rIn _ (inK 0 (by decide))) + (by rw [rd₃, wr₃, rdwr₂]; exact rIn _ (inK 8 (by decide))) + (by rw [wr₃, wr₂]; exact inK 16 (by decide)) (by rw [wr₃, wr₂]; exact inK 24 (by decide)) + refine WP.of_runBlock ⟨s₄, run₄, ?_⟩ + have rbp₄ : s₄.gpr .rbp = S := by + rw [g₄ _ (by decide) (by decide) (by decide) (by decide), g₃ _ (by decide) (by decide) (by decide) (by decide), + rbp₂] + obtain ⟨s₅, run₅, rbx₅, rbp₅, g₅, mem₅⟩ := restore2_ok s₄ rbp₄ + (by rw [rd₄, wr₄, rd₃, wr₃, rdwr₂]; exact rIn _ (inS 2064 (by decide))) + (by rw [rd₄, wr₄, rd₃, wr₃, rdwr₂]; exact rIn _ (inS 2072 (by decide))) + refine WP.of_runBlock ⟨s₅, run₅, ?_⟩ + -- Memory. + have f₂ := h₂.frame + have f₃ : Frame [⟨K + BitVec.ofNat 64 0, 16⟩] s₂.mem s₃.mem := by rw [mem₃]; exact dblMem_frame _ _ _ _ + have f₄ : Frame [⟨K + BitVec.ofNat 64 16, 16⟩] s₃.mem s₄.mem := by rw [mem₄]; exact dblMem_frame _ _ _ _ + have slotD : ∀ r ∈ [⟨S + BitVec.ofNat 64 2048, 16⟩, ⟨K, 16⟩, ⟨S, 2048⟩, below (s₁.gpr .rsp) 8], + (⟨S + BitVec.ofNat 64 2064, 16⟩ : Region).Disjoint r := by + intro r hr + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl + · exact Offset.disjoint S (by decide) (by omega) (by omega) + · exact (hp.k_scr.symm.sub_left (scr_sub' (by decide))).sub_right (Region.sub_prefix (by decide)) + · exact Offset.disjoint_base _ (by decide) (by omega) + · rw [rsp₁]; exact (hp.stk_scr.symm.sub_left (scr_sub' (by decide))) + have slotK (e : Nat) (he : e ≤ 16) : (⟨S + BitVec.ofNat 64 2064, 16⟩ : Region).Disjoint ⟨K + BitVec.ofNat 64 e, 16⟩ := + (hp.k_scr.symm.sub_left (scr_sub' (by decide))).sub_right (Offset.sub_base _ (by omega)) + have slot (d : Nat) (h₁ : 2064 ≤ d) (h₂' : d + 8 ≤ 2080) : + s₄.mem.readW (S + BitVec.ofNat 64 d) 64 = s₁.mem.readW (S + BitVec.ofNat 64 d) 64 := by + have c : (⟨S + BitVec.ofNat 64 2064, 16⟩ : Region).Contains (S + BitVec.ofNat 64 d) (64 / 8) := by + rw [show S + BitVec.ofNat 64 d = S + BitVec.ofNat 64 2064 + BitVec.ofNat 64 (d - 2064) from + (Offset.add_add_eq S (by omega)).symm] + exact Offset.contains_base _ (by omega) (by omega) + rw [f₄.readW c (fun r hr => by simp only [List.mem_singleton] at hr; subst hr; exact slotK 16 (by decide)) + (by decide), + f₃.readW c (fun r hr => by simp only [List.mem_singleton] at hr; subst hr; exact slotK 0 (by decide)) + (by decide), + f₂.readW c slotD (by decide)] + have pslot := fun d hd => preMem_slot s₀ (d := d) hd (by rw [hp.rdx, hp.rcx]; exact hp.k_scr) + rw [hp.rcx] at pslot + refine ⟨⟨fun r hr => ?_, ?_⟩, ?_⟩ + · by_cases hb : r = .rbx + · subst hb; rw [rbx₅, slot 2064 (by decide) (by decide), mem₁, pslot 2064 (.inl rfl)]; rfl + by_cases hb' : r = .rbp + · subst hb'; rw [rbp₅, slot 2072 (by decide) (by decide), mem₁, pslot 2072 (.inr rfl)]; rfl + rw [g₅ r hb hb', g₄ r (by rintro rfl; simp [calleeSaved] at hr) (by rintro rfl; simp [calleeSaved] at hr) + (by rintro rfl; simp [calleeSaved] at hr) (by rintro rfl; simp [calleeSaved] at hr), + g₃ r (by rintro rfl; simp [calleeSaved] at hr) (by rintro rfl; simp [calleeSaved] at hr) + (by rintro rfl; simp [calleeSaved] at hr) (by rintro rfl; simp [calleeSaved] at hr), + h₂.saved r hr, cs₁ r hr hb hb'] + · have fall : Frame [⟨K, 32⟩, ⟨S, 2176⟩, below (s₀.gpr .rsp) 8] s₀.mem s₅.mem := by + rw [mem₅] + refine ((f₁.sub fun r hr => ?_).trans (f₂.sub fun r hr => ?_)).trans + ((f₃.sub fun r hr => ?_).trans (f₄.sub fun r hr => ?_)) + · simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl + · exact ⟨⟨S, 2176⟩, by simp, scr_sub' (by decide)⟩ + · exact ⟨⟨K, 32⟩, by simp, Region.sub_prefix (by decide)⟩ + · simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl + · exact ⟨⟨S, 2176⟩, by simp, scr_sub' (by decide)⟩ + · exact ⟨⟨K, 32⟩, by simp, Region.sub_prefix (by decide)⟩ + · exact ⟨⟨S, 2176⟩, by simp, Region.sub_prefix (by decide)⟩ + · exact ⟨below (s₀.gpr .rsp) 8, by simp, by rw [rsp₁]; exact fun _ h => h⟩ + · simp only [List.mem_singleton] at hr; subst hr + exact ⟨⟨K, 32⟩, by simp, Offset.sub_base _ (by decide)⟩ + · simp only [List.mem_singleton] at hr; subst hr + exact ⟨⟨K, 32⟩, by simp, Offset.sub_base _ (by decide)⟩ + refine fall.readW (r := ⟨s₀.gpr .rsp, 8⟩) (Region.contains_self _ _) (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.ret_k + · exact hp.ret_scr + · exact Offset.base_disjoint_below _ (by decide) + · show Spec.Aes.bytesAt s₅.mem (s₀.gpr .rdx) 32 = _ + rw [hp.rdx, hp.rdi, hp.rsi, mem₅, bytesAt_32] + have L : Spec.Aes.bytesAt s₂.mem K 16 = ciphAt s₀.mem W R (Spec.Cmac.zeros 16) := by + rw [h₂.out, schB _ f₁, zC] + have b3 : Spec.Aes.bytesAt s₃.mem K 16 = Spec.Cmac.dbl 16 (Spec.Aes.bytesAt s₂.mem K 16) := by + have := dblMem_bytes s₂.mem K 0 0 + rw [k0] at this; rw [mem₃, this] + have b4lo : Spec.Aes.bytesAt s₄.mem K 16 = Spec.Aes.bytesAt s₃.mem K 16 := + bytesAt_frame' f₄ fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact (Offset.disjoint_base K (by decide) (by omega)).symm + have b4hi : Spec.Aes.bytesAt s₄.mem (K + BitVec.ofNat 64 16) 16 = + Spec.Cmac.dbl 16 (Spec.Aes.bytesAt s₃.mem K 16) := by + have := dblMem_bytes s₃.mem K 0 16 + rw [k0] at this; rw [mem₄, this] + rw [b4lo, b4hi, b3, L] + rfl + +end VG.Proof.CmacAes.X86_64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/X86_64/SubkeysCT.lean b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/SubkeysCT.lean new file mode 100644 index 000000000..b3e51a06d --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/SubkeysCT.lean @@ -0,0 +1,86 @@ +import VerifiedGarbage.Proof.CmacAes.X86_64.Subkeys +import VerifiedGarbage.Proof.Framework.X86_64.Taint + +/-! +# AES-CMAC on x86-64: `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 `rbx` and `rbp`); the call of +`vg_aes_ctr32` is constant time by its own proof (`ctr_rel`). +-/ + +namespace VG.Proof.CmacAes.X86_64 + +open VG VG.X86_64 VG.Impl.CmacAes.X86_64 +open VG.Proof.Aes.X86_64 (Ctr32Impl) + +/-- What is known between the code before the call and the call. -/ +structure SMid (s₀ : State) (W K S : Addr) (R : Nat) (s : State) : Prop where + pre : CallPre s W (S + BitVec.ofNat 64 2048) K S R + rbx : s.gpr .rbx = K + rbp : s.gpr .rbp = S + rsp : s.gpr .rsp = s₀.gpr .rsp + +theorem smid_wp {s₀ : State} {W K S : Addr} {R : Nat} (hp : SPre s₀ W K S R) : + WP isa (.block subkeysPre) s₀ (SMid s₀ W K S R) := by + have sw := hp.scr_wrap + have kw := hp.k_wrap + have inS (d : Nat) (h : d + 8 ≤ 2176) : InRegions s₀.wr (s₀.gpr .rcx + BitVec.ofNat 64 d) 8 := by + rw [hp.wr, hp.rcx]; exact in_rw (r := ⟨S, 2176⟩) (by simp) (Offset.contains_base _ h (by omega)) + have inK (d : Nat) (h : d + 8 ≤ 32) : InRegions s₀.wr (s₀.gpr .rdx + BitVec.ofNat 64 d) 8 := by + rw [hp.wr, hp.rdx]; exact in_rw (r := ⟨K, 32⟩) (by simp) (Offset.contains_base _ h (by omega)) + obtain ⟨s₁, run₁, rdi₁, rsi₁, rdx₁, rcx₁, r8₁, r9₁, rbx₁, rbp₁, cs₁, mem₁, rd₁, wr₁⟩ := + subkeysPre_ok s₀ (inS _ (by decide)) (inS _ (by decide)) (inS _ (by decide)) (inS _ (by decide)) + (inK _ (by decide)) (inK _ (by decide)) + have rsp₁ : s₁.gpr .rsp = s₀.gpr .rsp := cs₁ .rsp (by simp [calleeSaved]) (by decide) (by decide) + exact WP.of_runBlock ⟨s₁, run₁, + callPre_of hp rdi₁ rsi₁ rdx₁ rcx₁ r8₁ r9₁ rsp₁ mem₁ rd₁ wr₁, by rw [rbx₁, hp.rdx], by rw [rbp₁, hp.rcx], rsp₁⟩ + +theorem subkeys_rel (v : Ctr32Impl) {s₀ s₀' : State} (h0 : subkeysX86_64.pre s₀) + (h0' : subkeysX86_64.pre s₀') (hq : subkeysX86_64.pub s₀ s₀') : + RelCT isa (fun a b => a = s₀ ∧ b = s₀') (subkeys v.callee) fun _ _ => True := by + obtain ⟨q1, q2, q3, q4, q5⟩ := hq + have hp := SPre.of h0 + have hp' : SPre s₀' (s₀.gpr .rdi) (s₀.gpr .rdx) (s₀.gpr .rcx) (s₀.gpr .rsi).toNat := by + rw [q1, q2, q3, q4]; exact SPre.of h0' + obtain ⟨_, hA⟩ : ∃ h, (taint.check (Taint.ofRegs [.rdi, .rsi, .rdx, .rcx, .rsp]) (.block subkeysPre) h).isSome = + true := ⟨_, by taint_decide⟩ + obtain ⟨_, hB⟩ : ∃ h, (taint.check (Taint.ofRegs [.rbx, .rbp]) (.block subkeysPost) h).isSome = true := + ⟨_, by taint_decide⟩ + have a := (RelCT.taint (A := taint) (P := fun a b => a = s₀ ∧ b = s₀') _ + (fun a b h => by + obtain ⟨rfl, rfl⟩ := h + 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 | rfl <;> assumption) hA).wp + (F₁ := SMid s₀ _ _ _ _) (F₂ := SMid s₀' _ _ _ _) fun a b h => by + obtain ⟨rfl, rfl⟩ := h; exact ⟨smid_wp hp, smid_wp hp'⟩ + have c := (ctr_rel v (P := fun s₁ s₂ => + SMid s₀ (s₀.gpr .rdi) (s₀.gpr .rdx) (s₀.gpr .rcx) (s₀.gpr .rsi).toNat s₁ ∧ + SMid s₀' (s₀.gpr .rdi) (s₀.gpr .rdx) (s₀.gpr .rcx) (s₀.gpr .rsi).toNat s₂) fun s₁ s₂ h => + ⟨_, _, _, _, _, h.1.pre, h.2.pre, by rw [h.1.rsp, h.2.rsp, q5]⟩).wp + (F₁ := fun (s : State) => s.gpr .rbx = s₀.gpr .rdx ∧ s.gpr .rbp = s₀.gpr .rcx) + (F₂ := fun (s : State) => s.gpr .rbx = s₀.gpr .rdx ∧ s.gpr .rbp = s₀.gpr .rcx) fun s₁ s₂ h => + ⟨WP.mono (ctr_call v h.1.pre) fun _ hc => + ⟨by rw [hc.saved .rbx (by simp [calleeSaved]), h.1.rbx], + by rw [hc.saved .rbp (by simp [calleeSaved]), h.1.rbp]⟩, + WP.mono (ctr_call v h.2.pre) fun _ hc => + ⟨by rw [hc.saved .rbx (by simp [calleeSaved]), h.2.rbx], + by rw [hc.saved .rbp (by simp [calleeSaved]), h.2.rbp]⟩⟩ + have b := RelCT.taint (A := taint) + (P := fun s₁ s₂ => (s₁.gpr .rbx = s₀.gpr .rdx ∧ s₁.gpr .rbp = s₀.gpr .rcx) ∧ + (s₂.gpr .rbx = s₀.gpr .rdx ∧ s₂.gpr .rbp = s₀.gpr .rcx)) _ + (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.1, h.2.1] + · rw [h.1.2, h.2.2]) hB + exact (a.mono (fun _ _ h => h) fun _ _ h => h.2).seq + ((c.mono (fun _ _ h => h) fun _ _ h => h.2).seq b) + +theorem subkeys_ct (v : Ctr32Impl) : + ConstantTime isa subkeysX86_64.pre subkeysX86_64.pub (subkeys v.callee) := + fun _ _ _ _ _ _ h₁ h₂ hq e₁ e₂ => (subkeys_rel v h₁ h₂ hq _ _ _ _ _ _ ⟨rfl, rfl⟩ e₁ e₂).1 + +end VG.Proof.CmacAes.X86_64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Update.lean b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Update.lean new file mode 100644 index 000000000..ba252f7aa --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Update.lean @@ -0,0 +1,128 @@ +import VerifiedGarbage.Proof.CmacAes.X86_64.Contract +import VerifiedGarbage.Proof.Cmac.Mem +import VerifiedGarbage.Proof.Framework.X86_64.Exec +import VerifiedGarbage.Proof.Framework.X86_64.RegUpd +import VerifiedGarbage.Proof.Framework.X86_64.Abi + +/-! +# AES-CMAC on x86-64: `vg_cmac_aes_update` + +Untrusted: everything here is checked by Lean. +-/ + +namespace VG.Proof.CmacAes.X86_64 + +open VG VG.X86_64 VG.X86_64.RegUpd VG.Impl.CmacAes.X86_64 + +theorem offset_nat (i : Nat) : BitVec.ofInt 64 (i : Int) = BitVec.ofNat 64 i := rfl + +/-- The memory after saving the registers. -/ +def savedMem (s : State) : Mem := + saved.foldl (fun m (r, d) => m.writeW (s.gpr .r9 + BitVec.ofNat 64 d) (s.gpr r)) s.mem + +theorem prologue_ok (s : State) + (hw : ∀ d, 2064 ≤ d → d + 8 ≤ 2112 → InRegions s.wr (s.gpr .r9 + BitVec.ofNat 64 d) 8) : + ∃ s', runBlock isa (save ++ setup) s = some s' ∧ + s'.gpr .rbx = s.gpr .rdi ∧ s'.gpr .rbp = s.gpr .rsi ∧ s'.gpr .r12 = s.gpr .rdx ∧ + s'.gpr .r13 = s.gpr .rcx ∧ s'.gpr .r14 = s.gpr .r8 ∧ s'.gpr .r15 = s.gpr .r9 ∧ + s'.gpr .rsp = s.gpr .rsp ∧ s'.zf = some (s.gpr .r8 == 0) ∧ + s'.mem = savedMem s ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + simp only [save, setup, saved, List.map, List.cons_append, List.nil_append, runBlock_cons, + runStep_some, runBlock_nil, at_, exec, readSrc, State.store64, State.ea, offset_nat, + hw 2064 (by decide) (by decide), hw 2072 (by decide) (by decide), hw 2080 (by decide) (by decide), + hw 2088 (by decide) (by decide), hw 2096 (by decide) (by decide), hw 2104 (by decide) (by decide), + ite_true, Option.map_some, execAlu, Option.bind_some] + rfl, ?_⟩ + simp (config := {decide := true}) only [gpr_setReg, gpr_arithFlags, zf_arithFlags, mem_setReg, + mem_arithFlags, rd_setReg, rd_arithFlags, wr_setReg, wr_arithFlags, ite_true, ite_false, + BitVec.and_self] + trivial + +end VG.Proof.CmacAes.X86_64 + +namespace VG.Proof.CmacAes.X86_64 + +open VG VG.X86_64 VG.X86_64.RegUpd VG.Impl.CmacAes.X86_64 + +/-- The memory after `chainIn`: the counter block at `c` is the state at `p` +XORed with the block at `q`, and the state is zeroed. -/ +def chainMem (m : Mem) (c p q : Addr) : Mem := + let m₁ := m.writeW c (m.readW p 64 ^^^ m.readW q 64) + let m₂ := m₁.writeW (c + BitVec.ofNat 64 8) (m₁.readW (p + BitVec.ofNat 64 8) 64 ^^^ m₁.readW (q + BitVec.ofNat 64 8) 64) + (m₂.writeW p (0 : BitVec 64)).writeW (p + BitVec.ofNat 64 8) (0 : BitVec 64) + +theorem chainIn_ok (s : State) {C P Q : Addr} (hc : s.gpr .r15 + BitVec.ofNat 64 2048 = C) + (hp : s.gpr .r12 = P) (hq : s.gpr .r13 = Q) + (rp : InRegions (s.rd ++ s.wr) P 8) (rp8 : InRegions (s.rd ++ s.wr) (P + BitVec.ofNat 64 8) 8) + (rq : InRegions (s.rd ++ s.wr) Q 8) (rq8 : InRegions (s.rd ++ s.wr) (Q + BitVec.ofNat 64 8) 8) + (wc : InRegions s.wr C 8) (wc8 : InRegions s.wr (C + BitVec.ofNat 64 8) 8) + (wp : InRegions s.wr P 8) (wp8 : InRegions s.wr (P + BitVec.ofNat 64 8) 8) : + ∃ s', runBlock isa (chainIn ++ updArgs) s = some s' ∧ + s'.gpr .rdi = s.gpr .rbx ∧ s'.gpr .rsi = s.gpr .rbp ∧ s'.gpr .rdx = C ∧ s'.gpr .rcx = P ∧ + s'.gpr .r8 = 1 ∧ s'.gpr .r9 = s.gpr .r15 ∧ + (∀ r ∈ calleeSaved, s'.gpr r = s.gpr r) ∧ + s'.mem = chainMem s.mem C P Q ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + have hc' : s.gpr .r15 + BitVec.ofNat 64 2056 = C + BitVec.ofNat 64 8 := by + rw [← hc, BitVec.add_assoc]; rfl + refine ⟨_, by + simp (config := {decide := true}) only [chainIn, updArgs, ctrArgs, cOff, List.cons_append, + List.nil_append, runBlock_cons, runStep_some, runBlock_nil, at_, exec, readSrc, readSrc32, + State.load64, State.store64, State.ea, offset_nat, execAlu, Option.bind_some, Option.map_some, + gpr_setReg, gpr_arithFlags, mem_setReg, mem_arithFlags, rd_setReg, rd_arithFlags, wr_setReg, + wr_arithFlags, State.setReg32, ite_true, ite_false, hc, hc', hp, hq, BitVec.add_zero, + rp, rp8, rq, rq8, wc, wc8, wp, wp8] + rfl, ?_⟩ + refine ⟨?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, rfl, rfl⟩ + · simp [gpr_setReg] + · simp [gpr_setReg] + · simp [gpr_setReg, ← hc] + · simp [gpr_setReg] + · simp [gpr_setReg] + · simp [gpr_setReg] + · intro r hr + simp only [calleeSaved, List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl | rfl | rfl | rfl <;> simp [gpr_setReg] + · rfl + +theorem frame_store2 {m : Mem} (p : Addr) (w₀ w₁ : BitVec 64) : + Frame [⟨p, 16⟩] m ((m.writeW p w₀).writeW (p + BitVec.ofNat 64 8) w₁) := + ((Frame.refl _ _).writeW (List.mem_singleton_self _) _ + (by simpa using Offset.contains_base p (d := 0) (n := 8) (k := 16) (by decide) (by decide))).writeW + (List.mem_singleton_self _) _ (Offset.contains_base p (d := 8) (n := 8) (k := 16) (by decide) (by decide)) + +theorem chainMem_frame (m : Mem) (C P Q : Addr) : Frame [⟨C, 16⟩, ⟨P, 16⟩] m (chainMem m C P Q) := by + have f₁ : Frame [⟨C, 16⟩, ⟨P, 16⟩] m _ := + (frame_store2 (m := m) C (m.readW P 64 ^^^ m.readW Q 64) + ((m.writeW C (m.readW P 64 ^^^ m.readW Q 64)).readW (P + BitVec.ofNat 64 8) 64 ^^^ + (m.writeW C (m.readW P 64 ^^^ m.readW Q 64)).readW (Q + BitVec.ofNat 64 8) 64)).mono + (fun r hr => by simp only [List.mem_singleton] at hr; simp [hr]) + exact f₁.trans ((frame_store2 P 0 0).mono (fun r hr => by simp only [List.mem_singleton] at hr; simp [hr])) + +theorem chainMem_state (m : Mem) (C P Q : Addr) : + Spec.Aes.bytesAt (chainMem m C P Q) P 16 = Spec.Cmac.zeros 16 := by + rw [chainMem, Proof.Cmac.bytesAt_store2, Proof.Cmac.le8_zero]; rfl + +theorem bytesAt_frame' {rs : List Region} {m m' : Mem} (hf : Frame rs m m') {p : Addr} + (hd : ∀ r ∈ rs, (⟨p, 16⟩ : Region).Disjoint r) : Spec.Aes.bytesAt m' p 16 = Spec.Aes.bytesAt m p 16 := + bytesAt_frame hf hd (by decide) + +theorem readW_frame16 {rs : List Region} {m m' : Mem} (hf : Frame rs m m') {p : Addr} {d : Nat} (hd8 : d + 8 ≤ 16) + (hd : ∀ r ∈ rs, (⟨p, 16⟩ : Region).Disjoint r) : + m'.readW (p + BitVec.ofNat 64 d) 64 = m.readW (p + BitVec.ofNat 64 d) 64 := + hf.readW (r := ⟨p + BitVec.ofNat 64 d, 8⟩) (Region.contains_self _ _) + (fun r hr => (hd r hr).sub_left (Offset.sub_base p hd8)) (by decide) + +theorem chainMem_counter (m : Mem) {C P Q : Addr} (hcp : (⟨C, 16⟩ : Region).Disjoint ⟨P, 16⟩) + (hcq : (⟨C, 16⟩ : Region).Disjoint ⟨Q, 16⟩) : + Spec.Aes.bytesAt (chainMem m C P Q) C 16 = + Spec.Cmac.xor (Spec.Aes.bytesAt m P 16) (Spec.Aes.bytesAt m Q 16) := by + rw [chainMem, bytesAt_frame' (frame_store2 P 0 0) (by simpa using hcp), Proof.Cmac.bytesAt_store2] + have g : Frame [⟨C, 16⟩] m (m.writeW C (m.readW P 64 ^^^ m.readW Q 64)) := + (Frame.refl _ _).writeW (List.mem_singleton_self _) _ + (by simpa using Offset.contains_base C (d := 0) (n := 8) (k := 16) (by decide) (by decide)) + rw [readW_frame16 g (d := 8) (by decide) (by simpa using hcp.symm), + readW_frame16 g (d := 8) (by decide) (by simpa using hcq.symm)] + exact Proof.Cmac.xor_words m P Q + +end VG.Proof.CmacAes.X86_64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/X86_64/UpdateCT.lean b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/UpdateCT.lean new file mode 100644 index 000000000..aee89e45e --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/UpdateCT.lean @@ -0,0 +1,193 @@ +import VerifiedGarbage.Proof.CmacAes.X86_64.UpdateCorrect +import VerifiedGarbage.Proof.Framework.X86_64.Taint + +/-! +# AES-CMAC on x86-64: `vg_cmac_aes_update` is constant time + +Untrusted: everything here is checked by Lean. 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 registers the +correctness proof pins to the public arguments (`LInv`), and each call of +`vg_aes_ctr32` is constant time by its own proof (`ctr_rel`). +-/ + +namespace VG.Proof.CmacAes.X86_64 + +open VG VG.X86_64 VG.Impl.CmacAes.X86_64 +open VG.Proof.Aes.X86_64 (Ctr32Impl) + +section +variable {s₀ s₀' : State} (hq : updateX86_64.pub s₀ s₀') +include hq + +theorem pub_W : W s₀ = W s₀' := hq.1 +theorem pub_rsi : s₀.gpr .rsi = s₀'.gpr .rsi := hq.2.1 +theorem pub_R : R s₀ = R s₀' := by rw [R, R, pub_rsi hq] +theorem pub_St : St s₀ = St s₀' := hq.2.2.1 +theorem pub_Dp : Dp s₀ = Dp s₀' := hq.2.2.2.1 +theorem pub_N : N s₀ = N s₀' := by rw [N, N, hq.2.2.2.2.1] +theorem pub_S : S s₀ = S s₀' := hq.2.2.2.2.2.1 +theorem pub_rsp : s₀.gpr .rsp = s₀'.gpr .rsp := hq.2.2.2.2.2.2 + +/-- 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 ∈ [Reg.rbx, .rbp, .r12, .r13, .r14, .r15, .rsp], s₁.gpr r = s₂.gpr r := by + intro r hr + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl | rfl | rfl | rfl + · rw [h₁.rbx, h₂.rbx, pub_W hq] + · rw [h₁.rbp, h₂.rbp, pub_rsi hq] + · rw [h₁.r12, h₂.r12, pub_St hq] + · rw [h₁.r13, h₂.r13, pub_Dp hq] + · rw [h₁.r14, h₂.r14, pub_N hq] + · rw [h₁.r15, h₂.r15, pub_S hq] + · rw [h₁.rsp, h₂.rsp, pub_rsp 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₀) (S s₀ + BitVec.ofNat 64 2048) (St s₀) (S s₀) (R s₀) + r13 : s.gpr .r13 = Dp s₀ + BitVec.ofNat 64 (16 * k) + r14 : s.gpr .r14 = BitVec.ofNat 64 (N s₀ - k) + rsp : s.gpr .rsp = s₀.gpr .rsp + +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.saved .r13 (by simp [calleeSaved]), h.r13], + by rw [hb.saved .r14 (by simp [calleeSaved]), h.r14], by rw [hb.saved .rsp (by simp [calleeSaved]), h.rsp]⟩ + +/-- What is known after the call. -/ +structure After (s₀ : State) (k : Nat) (s : State) : Prop where + r13 : s.gpr .r13 = Dp s₀ + BitVec.ofNat 64 (16 * k) + r14 : s.gpr .r14 = BitVec.ofNat 64 (N s₀ - k) + +theorem body_ct (v : Ctr32Impl) {s₀ s₀' : State} (hp : UPre s₀) (hp' : UPre s₀') + (hq : updateX86_64.pub s₀ s₀') (k : Nat) : + RelCT isa (fun s₁ s₂ => k < N s₀ ∧ LInv s₀ k s₁ ∧ LInv s₀' k s₂) (body v.callee) fun _ _ => True := by + obtain ⟨_, hA⟩ : ∃ h, (taint.check (Taint.ofRegs [.rbx, .rbp, .r12, .r13, .r14, .r15, .rsp]) + (.block (chainIn ++ updArgs)) h).isSome = true := ⟨_, by taint_decide⟩ + obtain ⟨_, hB⟩ : ∃ h, (taint.check (Taint.ofRegs [.r13, .r14]) (.block advance) h).isSome = true := + ⟨_, by taint_decide⟩ + have a := (RelCT.taint (A := taint) (P := fun s₁ s₂ => k < N s₀ ∧ LInv s₀ k s₁ ∧ LInv s₀' k s₂) _ + (fun _ _ h => Taint.agree_ofRegs (LInv.agree hq h.2.1 h.2.2)) hA).wp + (F₁ := Mid s₀ k) (F₂ := Mid s₀' k) fun _ _ h => + ⟨bodyMid_wp hp h.1 h.2.1, bodyMid_wp hp' (by rw [← pub_N hq]; exact h.1) h.2.2⟩ + have c := (ctr_rel v (P := fun s₁ s₂ => Mid s₀ k s₁ ∧ Mid s₀' k s₂) fun s₁ s₂ h => + ⟨_, _, _, _, _, h.1.pre, by + rw [pub_W hq, pub_S hq, pub_St hq, pub_R hq]; exact h.2.pre, + by rw [h.1.rsp, h.2.rsp, pub_rsp hq]⟩).wp + (F₁ := After s₀ k) (F₂ := After s₀' k) fun s₁ s₂ h => + ⟨WP.mono (ctr_call v h.1.pre) fun _ hc => + ⟨by rw [hc.saved .r13 (by simp [calleeSaved]), h.1.r13], + by rw [hc.saved .r14 (by simp [calleeSaved]), h.1.r14]⟩, + WP.mono (ctr_call v h.2.pre) fun _ hc => + ⟨by rw [hc.saved .r13 (by simp [calleeSaved]), h.2.r13], + by rw [hc.saved .r14 (by simp [calleeSaved]), h.2.r14]⟩⟩ + have b := RelCT.taint (A := taint) (P := fun s₁ s₂ => After s₀ k s₁ ∧ After s₀' k s₂) _ + (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.r13, h.2.r13, pub_Dp hq] + · rw [h.1.r14, h.2.r14, pub_N hq]) hB + exact (a.mono (fun _ _ h => h) fun _ _ h => h.2).seq + ((c.mono (fun _ _ h => h) fun _ _ h => h.2).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 ∧ k < N s₀ ∧ LInv s₀ k s₁ ∧ LInv s₀' k s₂ + +theorem loop_ct (v : Ctr32Impl) {s₀ s₀' : State} (hp : UPre s₀) (hp' : UPre s₀') + (hq : updateX86_64.pub s₀ s₀') (n : Nat) : + RelCT isa (LRel s₀ s₀' n) (.loop (body v.callee) .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 + by_cases hk : k < N s₀ + swap + · exact RelCT.of_false fun _ _ h => hk h.2.1 + have ct := (body_ct v hp hp' hq k).wp + (F₁ := fun (s' : State) => LInv s₀ (k + 1) s' ∧ s'.zf = some (decide (N s₀ - (k + 1) = 0))) + (F₂ := fun (s' : State) => LInv s₀' (k + 1) s' ∧ s'.zf = some (decide (N s₀' - (k + 1) = 0))) + fun _ _ h => ⟨body_ok v hp h.1 h.2.1, body_ok v hp' (by rw [← hN]; exact h.1) h.2.2⟩ + refine ct.mono (fun _ _ h => h.2) fun s₁ s₂ ⟨_, ⟨l₁, z₁⟩, ⟨l₂, z₂⟩⟩ => ?_ + rw [← hN] at z₂ + have e₁ : isa.eval .ne s₁ = some (!decide (N s₀ - (k + 1) = 0)) := by + show s₁.zf.map _ = _; rw [z₁]; rfl + have e₂ : isa.eval .ne s₂ = some (!decide (N s₀ - (k + 1) = 0)) := by + show s₂.zf.map _ = _; rw [z₂]; rfl + 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₁, l₂⟩ + +/-! ## The whole function -/ + +theorem update_rel (v : Ctr32Impl) {s₀ s₀' : State} (h0 : updateX86_64.pre s₀) (h0' : updateX86_64.pre s₀') + (hq : updateX86_64.pub s₀ s₀') : + RelCT isa (fun a b => a = s₀ ∧ b = s₀') (update v.callee) fun _ _ => True := by + have hp := UPre.of h0 + have hp' := UPre.of h0' + have hN := pub_N hq + obtain ⟨_, hpro⟩ : ∃ h, (taint.check (Taint.ofRegs [.rdi, .rsi, .rdx, .rcx, .r8, .r9, .rsp]) + (.block (save ++ setup)) h).isSome = true := ⟨_, by taint_decide⟩ + obtain ⟨_, hepi⟩ : ∃ h, (taint.check (Taint.ofRegs [.r15]) (.block restore) h).isSome = true := + ⟨_, by taint_decide⟩ + obtain ⟨_, hnil⟩ : ∃ h, (taint.check (Taint.ofRegs []) (.block []) h).isSome = true := + ⟨_, by taint_decide⟩ + have pro := (RelCT.taint (A := taint) (P := fun a b => a = s₀ ∧ b = s₀') _ + (fun a b h => by + obtain ⟨rfl, rfl⟩ := h + refine Taint.agree_ofRegs fun r hr => ?_ + obtain ⟨h1, h2, h3, h4, h5, h6, h7⟩ := hq + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl | rfl | rfl | rfl <;> assumption) hpro).wp + (F₁ := fun (s : State) => LInv s₀ 0 s ∧ s.zf = some (decide (N s₀ = 0))) + (F₂ := fun (s : State) => LInv s₀' 0 s ∧ s.zf = some (decide (N s₀' = 0))) + fun a b h => by obtain ⟨rfl, rfl⟩ := h; exact ⟨prologue_wp hp, prologue_wp hp'⟩ + have nil := RelCT.taint (A := taint) + (P := fun a b => ((LInv s₀ 0 a ∧ a.zf = some (decide (N s₀ = 0))) ∧ + (LInv s₀' 0 b ∧ b.zf = some (decide (N s₀' = 0)))) ∧ isa.eval .e a = some true) _ + (fun _ _ _ => Taint.agree_ofRegs fun r hr => by simp at hr) hnil + have mid : RelCT isa (fun a b => (LInv s₀ 0 a ∧ a.zf = some (decide (N s₀ = 0))) ∧ + (LInv s₀' 0 b ∧ b.zf = some (decide (N s₀' = 0)))) + (.ite .e (.block []) (.loop (body v.callee) .ne)) + (fun a b => LInv s₀ (N s₀) a ∧ LInv s₀' (N s₀') b) := by + refine RelCT.ite (fun a b h => ?_) ?_ ?_ + · show a.zf = b.zf; rw [h.1.2, h.2.2, hN] + · 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; change a.zf = _ at this; rw [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 v hp hp' hq (N s₀ - 0)).mono (fun a b h => ⟨0, rfl, ?_, h.1.1.1, h.1.2.1⟩) + fun _ _ h => h + have := h.2; change a.zf = _ at this; rw [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) _ + (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.r15, h.2.r15, pub_S hq]) hepi + exact (pro.mono (fun _ _ h => h) fun _ _ h => h.2).seq (mid.seq epi) + +theorem update_ct (v : Ctr32Impl) : + ConstantTime isa updateX86_64.pre updateX86_64.pub (update v.callee) := + fun _ _ _ _ _ _ h₁ h₂ hq e₁ e₂ => (update_rel v h₁ h₂ hq _ _ _ _ _ _ ⟨rfl, rfl⟩ e₁ e₂).1 + +end VG.Proof.CmacAes.X86_64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/X86_64/UpdateCorrect.lean b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/UpdateCorrect.lean new file mode 100644 index 000000000..895820a95 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/UpdateCorrect.lean @@ -0,0 +1,181 @@ +import VerifiedGarbage.Proof.CmacAes.X86_64.UpdateLoop + +/-! +# AES-CMAC on x86-64: `vg_cmac_aes_update` is correct + +Untrusted: everything here is checked by Lean. +-/ + +namespace VG.Proof.CmacAes.X86_64 + +open VG VG.X86_64 VG.X86_64.RegUpd VG.Impl.CmacAes.X86_64 +open VG.Proof.Aes.X86_64 (Ctr32Impl) + +theorem beq_zero {x : Nat} (hx : x < 2 ^ 64) : (BitVec.ofNat 64 x == 0) = decide (x = 0) := by + rw [Bool.eq_iff_iff, beq_iff_eq, decide_eq_true_iff] + constructor + · intro he + have := congrArg BitVec.toNat he + rw [BitVec.toNat_ofNat, Nat.mod_eq_of_lt hx] at this + simpa using this + · intro he; rw [he]; rfl + +theorem loop_ok (v : Ctr32Impl) {s₀ : State} (hp : UPre s₀) {k : Nat} (hk : k < N s₀) {s : State} + (h : LInv s₀ k s) : WP isa (.loop (body v.callee) .ne) s (LInv s₀ (N s₀)) := by + refine WP.loop (M := isa) (body := body v.callee) (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 v hp hk h) fun s' ⟨h', zf'⟩ => ?_ + by_cases hz : N s₀ - (k + 1) = 0 + · left + refine ⟨by simp [eval, zf', hz], ?_⟩ + rwa [show N s₀ = k + 1 by omega] + · right + refine ⟨by simp [eval, zf', hz], N s₀ - (k + 1), by omega, k + 1, rfl, by omega, h'⟩ + +/-! ## Saving and restoring the registers -/ + +theorem readW_writeW_other (m : Mem) (b : Addr) {d e : Nat} (v : BitVec 64) (h : d + 8 ≤ e ∨ e + 8 ≤ d) + (hd : d + 8 ≤ 2 ^ 64) (he : e + 8 ≤ 2 ^ 64) : + (m.writeW (b + BitVec.ofNat 64 e) v).readW (b + BitVec.ofNat 64 d) 64 = m.readW (b + BitVec.ofNat 64 d) 64 := + Mem.readW_writeW_sep (Offset.sep b h hd he) (by decide) + +theorem savedMem_rbx (s : State) : (savedMem s).readW (s.gpr .r9 + BitVec.ofNat 64 2064) 64 = s.gpr .rbx := by + simp only [savedMem, saved, List.foldl] + rw [readW_writeW_other _ _ _ (by decide) (by decide) (by decide), + readW_writeW_other _ _ _ (by decide) (by decide) (by decide), + readW_writeW_other _ _ _ (by decide) (by decide) (by decide), + readW_writeW_other _ _ _ (by decide) (by decide) (by decide), + readW_writeW_other _ _ _ (by decide) (by decide) (by decide), Mem.readW_writeW_self64] + +theorem savedMem_rbp (s : State) : (savedMem s).readW (s.gpr .r9 + BitVec.ofNat 64 2072) 64 = s.gpr .rbp := by + simp only [savedMem, saved, List.foldl] + rw [readW_writeW_other _ _ _ (by decide) (by decide) (by decide), + readW_writeW_other _ _ _ (by decide) (by decide) (by decide), + readW_writeW_other _ _ _ (by decide) (by decide) (by decide), + readW_writeW_other _ _ _ (by decide) (by decide) (by decide), Mem.readW_writeW_self64] + +theorem savedMem_r12 (s : State) : (savedMem s).readW (s.gpr .r9 + BitVec.ofNat 64 2080) 64 = s.gpr .r12 := by + simp only [savedMem, saved, List.foldl] + rw [readW_writeW_other _ _ _ (by decide) (by decide) (by decide), + readW_writeW_other _ _ _ (by decide) (by decide) (by decide), + readW_writeW_other _ _ _ (by decide) (by decide) (by decide), Mem.readW_writeW_self64] + +theorem savedMem_r13 (s : State) : (savedMem s).readW (s.gpr .r9 + BitVec.ofNat 64 2088) 64 = s.gpr .r13 := by + simp only [savedMem, saved, List.foldl] + rw [readW_writeW_other _ _ _ (by decide) (by decide) (by decide), + readW_writeW_other _ _ _ (by decide) (by decide) (by decide), Mem.readW_writeW_self64] + +theorem savedMem_r14 (s : State) : (savedMem s).readW (s.gpr .r9 + BitVec.ofNat 64 2096) 64 = s.gpr .r14 := by + simp only [savedMem, saved, List.foldl] + rw [readW_writeW_other _ _ _ (by decide) (by decide) (by decide), Mem.readW_writeW_self64] + +theorem savedMem_r15 (s : State) : (savedMem s).readW (s.gpr .r9 + BitVec.ofNat 64 2104) 64 = s.gpr .r15 := by + simp only [savedMem, saved, List.foldl] + rw [Mem.readW_writeW_self64] + +theorem restore_ok (s : State) {B : Addr} (hb : s.gpr .r15 = B) + (hr : ∀ d, 2064 ≤ d → d + 8 ≤ 2112 → InRegions (s.rd ++ s.wr) (B + BitVec.ofNat 64 d) 8) : + ∃ s', runBlock isa restore s = some s' ∧ + s'.gpr .rbx = s.mem.readW (B + BitVec.ofNat 64 2064) 64 ∧ + s'.gpr .rbp = s.mem.readW (B + BitVec.ofNat 64 2072) 64 ∧ + s'.gpr .r12 = s.mem.readW (B + BitVec.ofNat 64 2080) 64 ∧ + s'.gpr .r13 = s.mem.readW (B + BitVec.ofNat 64 2088) 64 ∧ + s'.gpr .r14 = s.mem.readW (B + BitVec.ofNat 64 2096) 64 ∧ + s'.gpr .r15 = s.mem.readW (B + BitVec.ofNat 64 2104) 64 ∧ + s'.gpr .rsp = s.gpr .rsp ∧ s'.mem = s.mem := by + refine ⟨_, by + simp (config := {decide := true}) only [restore, saved, List.map, runBlock_cons, runStep_some, + runBlock_nil, at_, exec, readSrc, State.load64, State.ea, offset_nat, gpr_setReg, mem_setReg, + rd_setReg, wr_setReg, ite_true, ite_false, Option.map_some, hb, + hr 2064 (by decide) (by decide), hr 2072 (by decide) (by decide), hr 2080 (by decide) (by decide), + hr 2088 (by decide) (by decide), hr 2096 (by decide) (by decide), hr 2104 (by decide) (by decide)] + rfl, ?_⟩ + simp (config := {decide := true}) only [gpr_setReg, mem_setReg, ite_true, ite_false] + +/-! ## The whole function -/ + +theorem r8_ofNat (s₀ : State) : s₀.gpr .r8 = BitVec.ofNat 64 (N s₀) := by + apply BitVec.eq_of_toNat_eq; simp [N] + +theorem slots_disj {s₀ : State} (hp : UPre s₀) : + ∀ r ∈ [stR s₀, ⟨S s₀, 2064⟩, stkR s₀], (⟨S s₀ + BitVec.ofNat 64 2064, 48⟩ : Region).Disjoint r := by + intro r hr + 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 decide)) + · exact Offset.disjoint_base _ (by decide) (by have := hp.scr_wrap; omega) + · exact hp.stk_scr.symm.sub_left (UPre.scr_sub (by decide)) + +theorem slot_read {s₀ : State} (hp : UPre s₀) {m : Mem} + (hf : Frame [stR s₀, ⟨S s₀, 2064⟩, stkR s₀] (savedMem s₀) m) {d : Nat} (h₁ : 2064 ≤ d) (h₂ : d + 8 ≤ 2112) : + m.readW (S s₀ + BitVec.ofNat 64 d) 64 = (savedMem s₀).readW (S s₀ + BitVec.ofNat 64 d) 64 := + hf.readW (r := ⟨S s₀ + BitVec.ofNat 64 2064, 48⟩) (slot_contains _ h₁ h₂) (slots_disj hp) (by decide) + +theorem prologue_wp {s₀ : State} (hp : UPre s₀) : + WP isa (.block (save ++ setup)) s₀ fun s₁ => LInv s₀ 0 s₁ ∧ s₁.zf = some (decide (N s₀ = 0)) := by + have hN := (s₀.gpr .r8).isLt + obtain ⟨s₁, run₁, rbx₁, rbp₁, r12₁, r13₁, r14₁, r15₁, rsp₁, zf₁, mem₁, rd₁, wr₁⟩ := + prologue_ok s₀ fun d _ h₂ => by + rw [hp.wr]; exact in_rw (r := scrR s₀) (by simp) (Offset.contains_base _ (by omega) (by omega)) + refine WP.of_runBlock ⟨s₁, run₁, ?_, ?_⟩ + · have stSaved : Spec.Aes.bytesAt (savedMem s₀) (St s₀) 16 = Spec.Aes.bytesAt s₀.mem (St s₀) 16 := + bytesAt_frame' (savedMem_frame s₀) fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact hp.st_scr.sub_right (UPre.scr_sub (by decide)) + exact { rbx := rbx₁, rbp := rbp₁, r12 := r12₁ + r13 := by rw [r13₁]; simp + r14 := by rw [r14₁, r8_ofNat]; rfl + r15 := r15₁, rsp := rsp₁, rd := rd₁, wr := wr₁ + frame := by rw [mem₁]; exact Frame.refl _ _ + state := by rw [mem₁, stSaved]; rfl } + · rw [zf₁, r8_ofNat, beq_zero hN] + +theorem mid_wp (v : Ctr32Impl) {s₀ : State} (hp : UPre s₀) {s₁ : State} (h : LInv s₀ 0 s₁) + (hz : s₁.zf = some (decide (N s₀ = 0))) : + WP isa (.ite .e (.block []) (.loop (body v.callee) .ne)) s₁ (LInv s₀ (N s₀)) := by + have ev : isa.eval .e s₁ = some (decide (N s₀ = 0)) := hz + by_cases hn : N s₀ = 0 + · refine WP.ite true (by rw [ev, hn]; rfl) (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 v hp (by omega) h + +theorem epilogue_wp {s₀ : State} (hp : UPre s₀) {s₂ : State} (h₂ : LInv s₀ (N s₀) s₂) : + WP isa (.block restore) s₂ fun s' => gprPreserved s₀ s' ∧ updateX86_64.post s₀ s' := by + have rdwr : s₂.rd ++ s₂.wr = s₀.rd ++ s₀.wr := by rw [h₂.rd, h₂.wr] + obtain ⟨s₃, run₃, rbx₃, rbp₃, r12₃, r13₃, r14₃, r15₃, rsp₃, mem₃⟩ := + restore_ok s₂ h₂.r15 fun d _ h₂' => by + rw [rdwr, hp.rd, hp.wr] + exact in_rw (r := scrR s₀) (by simp) (Offset.contains_base _ (by omega) (by have := hp.scr_wrap; omega)) + refine WP.of_runBlock ⟨s₃, run₃, ?_⟩ + have rd' (d : Nat) (h₁ : 2064 ≤ d) (h₂' : d + 8 ≤ 2112) := slot_read hp h₂.frame h₁ h₂' + refine ⟨⟨fun r hr => ?_, ?_⟩, ?_⟩ + · simp only [calleeSaved, List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl | rfl | rfl | rfl + · rw [rbx₃, rd' 2064 (by decide) (by decide), savedMem_rbx] + · rw [rbp₃, rd' 2072 (by decide) (by decide), savedMem_rbp] + · rw [rsp₃, h₂.rsp] + · rw [r12₃, rd' 2080 (by decide) (by decide), savedMem_r12] + · rw [r13₃, rd' 2088 (by decide) (by decide), savedMem_r13] + · rw [r14₃, rd' 2096 (by decide) (by decide), savedMem_r14] + · rw [r15₃, rd' 2104 (by decide) (by decide), savedMem_r15] + · rw [mem₃] + refine (UPre.big_of h₂.frame).readW (r := ⟨s₀.gpr .rsp, 8⟩) (Region.contains_self _ _) + (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.ret_st + · exact hp.ret_scr + · exact Offset.base_disjoint_below _ (by decide) + · show Spec.Aes.bytesAt s₃.mem (St s₀) 16 = Spec.Cmac.chain (ciph s₀) _ (blks s₀) + rw [mem₃, h₂.state, List.take_of_length_le (by simp [Spec.Cmac.blocksAt])] + +theorem update_wp (v : Ctr32Impl) {s₀ : State} (h0 : updateX86_64.pre s₀) : + WP isa (update v.callee) s₀ fun s' => gprPreserved s₀ s' ∧ updateX86_64.post s₀ s' := by + have hp := UPre.of h0 + exact WP.seq (WP.mono (prologue_wp hp) fun s₁ ⟨h₁, z₁⟩ => + WP.seq (WP.mono (mid_wp v hp h₁ z₁) fun _ h₂ => epilogue_wp hp h₂)) + +end VG.Proof.CmacAes.X86_64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/X86_64/UpdateLoop.lean b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/UpdateLoop.lean new file mode 100644 index 000000000..1dd6e44ee --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/UpdateLoop.lean @@ -0,0 +1,332 @@ +import VerifiedGarbage.Proof.CmacAes.X86_64.Update + +/-! +# AES-CMAC on x86-64: the loop of `vg_cmac_aes_update` + +Untrusted: everything here is checked by Lean. The invariant after `k` +blocks (`LInv`): the registers hold the arguments (`r13` the next block, +`r14` the blocks left), only the state, the first 2064 bytes of the scratch +buffer and the stack below the return address have changed since the +registers were saved, and the state is the chaining value after the first +`k` blocks. +-/ + +namespace VG.Proof.CmacAes.X86_64 + +open VG VG.X86_64 VG.X86_64.RegUpd VG.Impl.CmacAes.X86_64 +open VG.Proof.Aes.X86_64 (Ctr32Impl) + +section +variable (s₀ : State) + +abbrev W : Addr := s₀.gpr .rdi +abbrev R : Nat := (s₀.gpr .rsi).toNat +abbrev St : Addr := s₀.gpr .rdx +abbrev Dp : Addr := s₀.gpr .rcx +abbrev N : Nat := (s₀.gpr .r8).toNat +abbrev S : Addr := s₀.gpr .r9 + +abbrev schR : Region := ⟨W s₀, 240⟩ +abbrev stR : Region := ⟨St s₀, 16⟩ +abbrev dataR : Region := ⟨Dp s₀, 16 * N s₀⟩ +abbrev scrR : Region := ⟨S s₀, 2176⟩ +abbrev stkR : Region := below (s₀.gpr .rsp) 8 + +/-- The cipher. -/ +abbrev ciph : Spec.Cmac.Cipher := ciphAt s₀.mem (W s₀) (R s₀) + +/-- The message blocks. -/ +abbrev blks : List (List Byte) := Spec.Cmac.blocksAt s₀.mem (Dp s₀) 16 (N s₀) + +end + +/-- The precondition, by name. -/ +structure UPre (s₀ : State) : Prop where + rd : s₀.rd = [schR s₀, dataR 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₀) + ret_st : (⟨s₀.gpr .rsp, 8⟩ : Region).Disjoint (stR s₀) + ret_scr : (⟨s₀.gpr .rsp, 8⟩ : Region).Disjoint (scrR s₀) + stk_sch : (stkR s₀).Disjoint (schR s₀) + stk_data : (stkR s₀).Disjoint (dataR s₀) + stk_st : (stkR s₀).Disjoint (stR s₀) + stk_scr : (stkR s₀).Disjoint (scrR s₀) + st_wrap : (St s₀).toNat + 16 ≤ 2 ^ 64 + data_wrap : (Dp s₀).toNat + 16 * N s₀ ≤ 2 ^ 64 + scr_wrap : (S s₀).toNat + 2176 ≤ 2 ^ 64 + rounds : R s₀ = 10 ∨ R s₀ = 12 ∨ R s₀ = 14 + +theorem UPre.of {s₀ : State} (h : updateX86_64.pre s₀) : UPre s₀ := + let ⟨a, b, c, d, e, f, g, h, i, j, k, l, m, n, o, p, q⟩ := h + ⟨a, b, c, d, e, f, g, h, i, j, k, l, m, n, o, p, q⟩ + +/-- The loop invariant, after `k` blocks. -/ +structure LInv (s₀ : State) (k : Nat) (s : State) : Prop where + rbx : s.gpr .rbx = W s₀ + rbp : s.gpr .rbp = s₀.gpr .rsi + r12 : s.gpr .r12 = St s₀ + r13 : s.gpr .r13 = Dp s₀ + BitVec.ofNat 64 (16 * k) + r14 : s.gpr .r14 = BitVec.ofNat 64 (N s₀ - k) + r15 : s.gpr .r15 = S s₀ + rsp : s.gpr .rsp = s₀.gpr .rsp + rd : s.rd = s₀.rd + wr : s.wr = s₀.wr + frame : Frame [stR s₀, ⟨S s₀, 2064⟩, stkR s₀] (savedMem s₀) s.mem + state : Spec.Aes.bytesAt s.mem (St s₀) 16 = + Spec.Cmac.chain (ciph s₀) (Spec.Aes.bytesAt s₀.mem (St s₀) 16) ((blks s₀).take k) + +/-! ## Regions -/ + +section +variable {s₀ : State} + +theorem UPre.scr_sub {d n : Nat} (h : d + n ≤ 2176) : Region.Sub ⟨S s₀ + BitVec.ofNat 64 d, n⟩ (scrR s₀) := + Offset.sub_base _ h + +theorem UPre.data_sub {k : Nat} (hk : k < N s₀) : + Region.Sub ⟨Dp s₀ + BitVec.ofNat 64 (16 * k), 16⟩ (dataR s₀) := + Offset.sub_base _ (by omega) + +theorem UPre.sch_sub {R' : Nat} (h : R' ≤ 240) : Region.Sub ⟨W s₀, R'⟩ (schR s₀) := + Region.sub_prefix h + +end + +theorem slot_contains (b : Addr) {d : Nat} (h₁ : 2064 ≤ d) (h₂ : d + 8 ≤ 2112) : + (⟨b + BitVec.ofNat 64 2064, 48⟩ : Region).Contains (b + BitVec.ofNat 64 d) 8 := by + rw [show b + BitVec.ofNat 64 d = (b + BitVec.ofNat 64 2064) + BitVec.ofNat 64 (d - 2064) from + (Offset.add_add_eq b (by omega)).symm] + exact Offset.contains_base _ (by omega) (by omega) + +/-- Saving the registers changes only their slots. -/ +theorem savedMem_frame (s : State) : Frame [⟨s.gpr .r9 + BitVec.ofNat 64 2064, 48⟩] s.mem (savedMem s) := by + simp only [savedMem, saved, List.foldl] + exact (((((((Frame.refl _ _).writeW (List.mem_singleton_self _) _ (slot_contains _ (by decide) (by decide))).writeW + (List.mem_singleton_self _) _ (slot_contains _ (by decide) (by decide))).writeW + (List.mem_singleton_self _) _ (slot_contains _ (by decide) (by decide))).writeW (List.mem_singleton_self _) _ + (slot_contains _ (by decide) (by decide))).writeW + (List.mem_singleton_self _) _ (slot_contains _ (by decide) (by decide))).writeW (List.mem_singleton_self _) _ + (slot_contains _ (by decide) (by decide))) + +theorem advance_ok (s : State) : + ∃ s', runBlock isa advance s = some s' ∧ + s'.gpr .r13 = s.gpr .r13 + BitVec.ofNat 64 16 ∧ s'.gpr .r14 = s.gpr .r14 - 1 ∧ + s'.zf = some ((s.gpr .r14 - 1) == 0) ∧ (∀ r, r ≠ .r13 → r ≠ .r14 → s'.gpr r = s.gpr r) ∧ + s'.mem = s.mem ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + simp (config := {decide := true}) only [advance, runBlock_cons, runStep_some, runBlock_nil, exec, + execAlu, readSrc, Option.bind_some, gpr_setReg, gpr_arithFlags, ite_false] + rfl, ?_⟩ + refine ⟨?_, ?_, ?_, ?_, rfl, rfl, rfl⟩ + · simp [gpr_setReg] + · exact gpr_setReg_self _ _ _ + · rw [zf_setReg, zf_arithFlags]; simp + · intro r h₁ h₂; simp [gpr_setReg, h₁, h₂] + +theorem rsi_ofNat (s₀ : State) : s₀.gpr .rsi = BitVec.ofNat 64 (R s₀) := by + apply BitVec.eq_of_toNat_eq; simp [R] + +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 (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] + +theorem in_rw {rs : List Region} {r : Region} (hr : r ∈ rs) {a : Addr} {n : Nat} (hc : r.Contains a n) : + InRegions rs a n := ⟨r, hr, hc⟩ + +/-! ## Memory outside the writable regions -/ + +/-- The regions the function writes. -/ +abbrev Big (s₀ : State) : List Region := [stR s₀, scrR s₀, stkR 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 (W s₀) (16 * (R s₀ + 1)) = Spec.Aes.bytesAt s₀.mem (W s₀) (16 * (R s₀ + 1)) := by + have hR : 16 * (R s₀ + 1) ≤ 240 := by rcases hp.rounds with h | h | h <;> omega + refine 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.stk_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 (Dp s₀ + BitVec.ofNat 64 (16 * k)) 16 = + Spec.Aes.bytesAt s₀.mem (Dp s₀ + BitVec.ofNat 64 (16 * k)) 16 := by + refine 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.stk_data.symm.sub_left (UPre.data_sub hk) + +omit hp in +theorem UPre.big_of {m : Mem} (hf : Frame [stR s₀, ⟨S s₀, 2064⟩, stkR 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, UPre.scr_sub (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 ⟨stkR s₀, by simp, fun _ h => h⟩) + +end + +/-! ## One block -/ + +/-- What the code before the call leaves. -/ +structure BodyA (s₀ : State) (k : Nat) (s s₁ : State) : Prop where + pre : CallPre s₁ (W s₀) (S s₀ + BitVec.ofNat 64 2048) (St s₀) (S s₀) (R s₀) + saved : ∀ r ∈ calleeSaved, s₁.gpr r = s.gpr r + mem : s₁.mem = chainMem s.mem (S s₀ + BitVec.ofNat 64 2048) (St s₀) (Dp s₀ + BitVec.ofNat 64 (16 * k)) + rd : s₁.rd = s.rd + wr : s₁.wr = s.wr + +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₀, 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 h16k : 16 * k + 16 ≤ 16 * N s₀ := by omega + have hdw := hp.data_wrap + have cSt0 : (stR s₀).Contains (St s₀) 8 := by + simpa using Offset.contains_base (St s₀) (d := 0) (n := 8) (k := 16) (by decide) (by decide) + have cSt8 : (stR s₀).Contains (St s₀ + BitVec.ofNat 64 8) 8 := Offset.contains_base _ (by decide) (by decide) + have cQ0 : (dataR s₀).Contains (Dp s₀ + BitVec.ofNat 64 (16 * k)) 8 := + Offset.contains_base _ (by omega) (by omega) + have cQ8 : (dataR s₀).Contains (Dp s₀ + BitVec.ofNat 64 (16 * k) + BitVec.ofNat 64 8) 8 := by + rw [Offset.add_add]; exact Offset.contains_base _ (by omega) (by omega) + have cC0 : (scrR s₀).Contains (S s₀ + BitVec.ofNat 64 2048) 8 := Offset.contains_base _ (by decide) (by decide) + have cC8 : (scrR s₀).Contains (S s₀ + BitVec.ofNat 64 2048 + BitVec.ofNat 64 8) 8 := by + rw [Offset.add_add]; exact Offset.contains_base _ (by decide) (by decide) + obtain ⟨s₁, run₁, rdi₁, rsi₁, rdx₁, rcx₁, r8₁, r9₁, cs₁, mem₁, rd₁, wr₁⟩ := + chainIn_ok s (C := S s₀ + BitVec.ofNat 64 2048) (P := St s₀) (Q := Dp s₀ + BitVec.ofNat 64 (16 * k)) + (by rw [h.r15]) h.r12 h.r13 + (by rw [hRegs]; exact in_rw (by simp) cSt0) (by rw [hRegs]; exact in_rw (by simp) cSt8) + (by rw [hRegs]; exact in_rw (by simp) cQ0) (by rw [hRegs]; exact in_rw (by simp) cQ8) + (by rw [hW]; exact in_rw (by simp) cC0) (by rw [hW]; exact in_rw (by simp) cC8) + (by rw [hW]; exact in_rw (by simp) cSt0) (by rw [hW]; exact in_rw (by simp) cSt8) + refine WP.of_runBlock ⟨s₁, run₁, ?_⟩ + have rsp₁ : s₁.gpr .rsp = s₀.gpr .rsp := by rw [cs₁ .rsp (by simp [calleeSaved]), h.rsp] + have hR : 16 * (R s₀ + 1) ≤ 240 := by rcases hp.rounds with h | h | h <;> omega + have cDis : (⟨S s₀ + BitVec.ofNat 64 2048, 16⟩ : Region).Disjoint ⟨S s₀, 2048⟩ := + Offset.disjoint_base _ (by decide) (by have := hp.scr_wrap; omega) + have pre : CallPre s₁ (W s₀) (S s₀ + BitVec.ofNat 64 2048) (St s₀) (S s₀) (R s₀) := + { rdi := by rw [rdi₁, h.rbx] + rsi := by rw [rsi₁, h.rbp, rsi_ofNat] + rdx := rdx₁ + rcx := rcx₁ + r8 := r8₁ + r9 := by rw [r9₁, h.r15] + rounds := hp.rounds + wc := hp.sch_scr.sub_right (UPre.scr_sub (by decide)) + wd := hp.sch_st + ws := hp.sch_scr.sub_right (Region.sub_prefix (by decide)) + cd := hp.st_scr.symm.sub_left (UPre.scr_sub (by decide)) + cs := cDis + ds := hp.st_scr.sub_right (Region.sub_prefix (by decide)) + stkW := by rw [rsp₁]; exact hp.stk_sch + stkC := by rw [rsp₁]; exact hp.stk_scr.sub_right (UPre.scr_sub (by decide)) + stkD := by rw [rsp₁]; exact hp.stk_st + stkS := by rw [rsp₁]; exact hp.stk_scr.sub_right (Region.sub_prefix (by decide)) + wrap := hp.st_wrap + reads := by + rw [rd₁, wr₁, hRegs] + refine Covers.of_sub fun r hr => ?_ + 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 + · exact ⟨schR s₀, by simp, 0, by simp, by simp⟩ + · 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⟩ + writes := by + rw [wr₁, hW] + 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 chainMem_state _ _ _ _ } + exact ⟨pre, cs₁, mem₁, rd₁, wr₁⟩ + +theorem body_ok (v : Ctr32Impl) {s₀ : State} (hp : UPre s₀) {k : Nat} (hk : k < N s₀) {s : State} + (h : LInv s₀ k s) : + WP isa (body v.callee) s fun s' => LInv s₀ (k + 1) s' ∧ s'.zf = some (decide (N s₀ - (k + 1) = 0)) := by + have h16k : 16 * k + 16 ≤ 16 * N s₀ := by omega + have hdw := hp.data_wrap + refine WP.seq (WP.mono (bodyA_wp hp hk h) fun s₁ ⟨pre, cs₁, mem₁, rd₁, wr₁⟩ => ?_) + have rsp₁ : s₁.gpr .rsp = s₀.gpr .rsp := by rw [cs₁ .rsp (by simp [calleeSaved]), h.rsp] + refine WP.seq (WP.mono (ctr_call v pre) fun s₂ h₂ => ?_) + obtain ⟨s₃, run₃, r13₃, r14₃, zf₃, keep₃, mem₃, rd₃, wr₃⟩ := advance_ok s₂ + refine WP.of_runBlock ⟨s₃, run₃, ?_⟩ + have g (r : Reg) (hr : r ∈ calleeSaved) (h13 : r ≠ .r13) (h14 : r ≠ .r14) : s₃.gpr r = s.gpr r := by + rw [keep₃ r h13 h14, h₂.saved r hr, cs₁ r hr] + have r13₂ : s₂.gpr .r13 = Dp s₀ + BitVec.ofNat 64 (16 * k) := by + rw [h₂.saved .r13 (by simp [calleeSaved]), cs₁ .r13 (by simp [calleeSaved]), h.r13] + have r14₂ : s₂.gpr .r14 = BitVec.ofNat 64 (N s₀ - k) := by + rw [h₂.saved .r14 (by simp [calleeSaved]), cs₁ .r14 (by simp [calleeSaved]), h.r14] + have hN := (s₀.gpr .r8).isLt + have dec : BitVec.ofNat 64 (N s₀ - k) - 1 = BitVec.ofNat 64 (N s₀ - (k + 1)) := by + rw [show (1 : BitVec 64) = BitVec.ofNat 64 1 from rfl, Offset.ofNat_sub_ofNat (by omega)]; rfl + -- Memory. + have bigS := UPre.big_of h.frame + have f₁ : Frame [⟨S s₀ + BitVec.ofNat 64 2048, 16⟩, ⟨St s₀, 16⟩] s.mem s₁.mem := by + rw [mem₁]; exact chainMem_frame _ _ _ _ + have big₁ : Frame (Big s₀) s₀.mem s₁.mem := bigS.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 : (⟨S s₀ + BitVec.ofNat 64 2048, 16⟩ : Region).Disjoint (stR s₀) := + hp.st_scr.symm.sub_left (UPre.scr_sub (by decide)) + have cq : (⟨S s₀ + BitVec.ofNat 64 2048, 16⟩ : Region).Disjoint ⟨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₁, mem₁, chainMem_counter _ cst cq, h.state, + UPre.block_bytes hp bigS hk] at out + refine ⟨⟨?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_⟩, ?_⟩ + · rw [g .rbx (by simp [calleeSaved]) (by decide) (by decide), h.rbx] + · rw [g .rbp (by simp [calleeSaved]) (by decide) (by decide), h.rbp] + · rw [g .r12 (by simp [calleeSaved]) (by decide) (by decide), h.r12] + · rw [r13₃, r13₂, Offset.add_add_eq _ (c := 16 * (k + 1)) (by omega)] + · rw [r14₃, r14₂, dec] + · rw [g .r15 (by simp [calleeSaved]) (by decide) (by decide), h.r15] + · rw [g .rsp (by simp [calleeSaved]) (by decide) (by decide), h.rsp] + · rw [rd₃, h₂.rd, rd₁, h.rd] + · rw [wr₃, h₂.wr, wr₁, h.wr] + · rw [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 ⟨⟨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 ⟨⟨S s₀, 2064⟩, by simp, Offset.sub_base _ (by decide)⟩ + · exact ⟨stR s₀, by simp, fun _ h => h⟩ + · exact ⟨⟨S s₀, 2064⟩, by simp, Region.sub_prefix (by decide)⟩ + · exact ⟨stkR s₀, by simp, by rw [rsp₁]; exact fun _ h => h⟩ + · rw [mem₃, out, take_succ_blks s₀ hk, Proof.Cmac.chain_append, Proof.Cmac.chain_single] + · rw [zf₃, r14₂, dec] + congr 1 + rw [Bool.eq_iff_iff, beq_iff_eq, decide_eq_true_iff] + constructor + · intro he + have := congrArg BitVec.toNat he + rw [BitVec.toNat_ofNat, Nat.mod_eq_of_lt (by omega)] at this + simpa using this + · intro he; rw [he]; rfl + +end VG.Proof.CmacAes.X86_64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Verified.lean b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Verified.lean new file mode 100644 index 000000000..1cafeb32d --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/X86_64/Verified.lean @@ -0,0 +1,108 @@ +import VerifiedGarbage.Proof.CmacAes.X86_64.UpdateCT +import VerifiedGarbage.Proof.CmacAes.X86_64.SubkeysCT +import VerifiedGarbage.Proof.CmacAes.X86_64.FinalizeCT +import VerifiedGarbage.Proof.Framework.Contract +import VerifiedGarbage.Spec.Cmac.Contract + +/-! +# AES-CMAC on x86-64: `Verified` + +Untrusted: everything here is checked by Lean. Correctness and constant time +(for any implementation `v` of `vg_aes_ctr32`), a state satisfying each +precondition, and the shared contracts of `Spec/Cmac/Contract.lean` (with +8 bytes of stack, for the return address of the call of `vg_aes_ctr32`). +-/ + +namespace VG.Proof.CmacAes.X86_64 + +open VG VG.X86_64 VG.Impl.CmacAes.X86_64 +open VG.Proof.Aes.X86_64 (Ctr32Impl) + +theorem update_mx (v : Ctr32Impl) : (update v.callee).allInstrs (fun i => !loadsMxcsr i) = true := by + simp only [update, body, Code.allInstrs, v.mxcsr]; decide +kernel + +theorem subkeys_mx (v : Ctr32Impl) : (subkeys v.callee).allInstrs (fun i => !loadsMxcsr i) = true := by + simp only [subkeys, Code.allInstrs, v.mxcsr]; decide +kernel + +theorem finalize_mx (v : Ctr32Impl) : (finalize v.callee).allInstrs (fun i => !loadsMxcsr i) = true := by + simp only [finalize, Code.allInstrs, v.mxcsr]; decide +kernel + +theorem update_spSafe (v : Ctr32Impl) : (update v.callee).all (fun i => !X86_64.isa.writesSp i) = true := by + simp only [update, body, Code.all, v.spSafe]; decide +kernel + +theorem subkeys_spSafe (v : Ctr32Impl) : (subkeys v.callee).all (fun i => !X86_64.isa.writesSp i) = true := by + simp only [subkeys, Code.all, v.spSafe]; decide +kernel + +theorem finalize_spSafe (v : Ctr32Impl) : (finalize v.callee).all (fun i => !X86_64.isa.writesSp i) = true := by + simp only [finalize, Code.all, v.spSafe]; decide +kernel + +theorem update_correct (v : Ctr32Impl) (s : State) (hs : updateX86_64.pre s) : + ∃ t s', Exec isa (update v.callee) s t s' ∧ abiPreserved s s' ∧ updateX86_64.post s s' := by + obtain ⟨t, s', he, hg, hp⟩ := update_wp v hs + exact ⟨t, s', he, abiPreserved_of_exec (update_mx v) he hg, hp⟩ + +theorem subkeys_correct (v : Ctr32Impl) (s : State) (hs : subkeysX86_64.pre s) : + ∃ t s', Exec isa (subkeys v.callee) s t s' ∧ abiPreserved s s' ∧ subkeysX86_64.post s s' := by + obtain ⟨t, s', he, hg, hp⟩ := subkeys_wp v hs + exact ⟨t, s', he, abiPreserved_of_exec (subkeys_mx v) he hg, hp⟩ + +theorem finalize_correct (v : Ctr32Impl) (s : State) (hs : finalizeX86_64.pre s) : + ∃ t s', Exec isa (finalize v.callee) s t s' ∧ abiPreserved s s' ∧ finalizeX86_64.post s s' := by + obtain ⟨t, s', he, hg, hp⟩ := finalize_wp v hs + exact ⟨t, s', he, abiPreserved_of_exec (finalize_mx v) he hg, hp⟩ + +/-- A state satisfying `vg_cmac_aes_update`'s precondition (with no blocks). -/ +def updSat : State where + gpr r := match r with + | .rdi => 0x1000 | .rsi => 10 | .rdx => 0x2000 | .rcx => 0x3000 | .r9 => 0x4000 | .rsp => 0x8000 | _ => 0 + cf := none + zf := none + sf := none + of := none + mem _ := 0 + rd := [⟨0x1000, 240⟩, ⟨0x3000, 0⟩] + wr := [⟨0x2000, 16⟩, ⟨0x4000, 2176⟩] + +theorem update_verified (v : Ctr32Impl) : + Verified X86_64.target (update v.callee) (Spec.Cmac.aesUpdateContract X86_64.abi 8) := + Verified.of_correct (update_correct v) (update_ct v) (by + sig_implies [Spec.Cmac.aesUpdateContract, Spec.Cmac.aesUpdateSig, updateX86_64, X86_64.abi, + X86_64.argRegs] [updSat] using updSat) + +/-- A state satisfying `vg_cmac_aes_subkeys`'s precondition. -/ +def subSat : State where + gpr r := match r with + | .rdi => 0x1000 | .rsi => 10 | .rdx => 0x2000 | .rcx => 0x4000 | .rsp => 0x8000 | _ => 0 + cf := none + zf := none + sf := none + of := none + mem _ := 0 + rd := [⟨0x1000, 240⟩] + wr := [⟨0x2000, 32⟩, ⟨0x4000, 2176⟩] + +theorem subkeys_verified (v : Ctr32Impl) : + Verified X86_64.target (subkeys v.callee) (Spec.Cmac.aesSubkeysContract X86_64.abi 8) := + Verified.of_correct (subkeys_correct v) (subkeys_ct v) (by + sig_implies [Spec.Cmac.aesSubkeysContract, Spec.Cmac.aesSubkeysSig, subkeysX86_64, X86_64.abi, + X86_64.argRegs] [subSat] using subSat) + +/-- A state satisfying `vg_cmac_aes_finalize`'s precondition (with no last bytes). -/ +def finSat : State where + gpr r := match r with + | .rdi => 0x1000 | .rsi => 10 | .rdx => 0x2000 | .rcx => 0x3000 | .r9 => 0x4000 | .rsp => 0x8000 | _ => 0 + cf := none + zf := none + sf := none + of := none + mem _ := 0 + rd := [⟨0x1000, 272⟩, ⟨0x3000, 0⟩] + wr := [⟨0x2000, 16⟩, ⟨0x4000, 2176⟩] + +theorem finalize_verified (v : Ctr32Impl) : + Verified X86_64.target (finalize v.callee) (Spec.Cmac.aesFinalizeContract X86_64.abi 8) := + Verified.of_correct (finalize_correct v) (finalize_ct v) (by + sig_implies [Spec.Cmac.aesFinalizeContract, Spec.Cmac.aesFinalizeSig, finalizeX86_64, X86_64.abi, + X86_64.argRegs] [finSat] using finSat) + +end VG.Proof.CmacAes.X86_64 diff --git a/lean/VerifiedGarbage/Variants/AesCtr32/X86_64/AesNi.lean b/lean/VerifiedGarbage/Variants/AesCtr32/X86_64/AesNi.lean new file mode 100644 index 000000000..bab089793 --- /dev/null +++ b/lean/VerifiedGarbage/Variants/AesCtr32/X86_64/AesNi.lean @@ -0,0 +1,14 @@ +import VerifiedGarbage.Proof.Aes.X86_64.Variant + +/-! +# `vg_aes_ctr32` on x86-64: with AES-NI + +A variant of `AesCtr32` on x86-64 (see `TCB/Emit.lean`): +`vg_aes_ctr32_aesni`, which needs AES-NI and SSSE3. +-/ + +namespace VG.Variants.AesCtr32.X86_64.AesNi + +def variant : Proof.Aes.X86_64.Ctr32Impl := .aesni + +end VG.Variants.AesCtr32.X86_64.AesNi diff --git a/lean/VerifiedGarbage/Variants/AesCtr32/X86_64/Scalar.lean b/lean/VerifiedGarbage/Variants/AesCtr32/X86_64/Scalar.lean new file mode 100644 index 000000000..84a44b79b --- /dev/null +++ b/lean/VerifiedGarbage/Variants/AesCtr32/X86_64/Scalar.lean @@ -0,0 +1,14 @@ +import VerifiedGarbage.Proof.Aes.X86_64.Variant + +/-! +# `vg_aes_ctr32` on x86-64: the bitsliced implementation + +A variant of `AesCtr32` on x86-64 (see `TCB/Emit.lean`): `vg_aes_ctr32`, in +the baseline ISA. +-/ + +namespace VG.Variants.AesCtr32.X86_64.Scalar + +def variant : Proof.Aes.X86_64.Ctr32Impl := .scalar + +end VG.Variants.AesCtr32.X86_64.Scalar diff --git a/src/asm/x86_64/cmac_aes.rs b/src/asm/x86_64/cmac_aes.rs new file mode 100644 index 000000000..789e8d8a6 --- /dev/null +++ b/src/asm/x86_64/cmac_aes.rs @@ -0,0 +1,453 @@ +// @generated from lean/VerifiedGarbage/Artifacts.lean by lean/Emit.lean. DO NOT EDIT. +//! Verified `cmac_aes` functions for `x86_64`. +#![allow(dead_code)] + +/// The CPU features `vg_cmac_aes_subkeys_aesni` requires (`Artifact.features`). +pub(crate) const VG_CMAC_AES_SUBKEYS_AESNI_FEATURES: &[&str] = &["aes", "ssse3"]; + +/// 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_aesni`. +/// +/// # 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 return address on the stack or the 8 bytes of stack below it, or wrap around the end of the address space (no Rust object does). +/// * The CPU must support the `aes` and `ssse3` target features. +#[unsafe(naked)] +pub(crate) unsafe extern "sysv64" fn vg_cmac_aes_subkeys_aesni(schedule: *const [u8; 240], rounds: usize, subkeys: *mut [u8; 32], scratch: *mut [u64; 272]) { + core::arch::naked_asm!( + "mov QWORD PTR [rcx+2064], rbx", + "mov QWORD PTR [rcx+2072], rbp", + "mov rbx, rdx", + "mov rbp, rcx", + "mov eax, 0", + "mov QWORD PTR [rcx+2048], rax", + "mov QWORD PTR [rcx+2056], rax", + "mov QWORD PTR [rdx], rax", + "mov QWORD PTR [rdx+8], rax", + "mov r9, rbp", + "mov rdx, rbp", + "add rdx, 2048", + "mov rcx, rbx", + "mov r8d, 1", + "call {vg_aes_ctr32_aesni}", + "mov rax, QWORD PTR [rbx]", + "bswap rax", + "mov rdx, QWORD PTR [rbx+8]", + "bswap rdx", + "mov rcx, rax", + "shr rcx, 63", + "mov r8d, 0", + "sub r8, rcx", + "and r8, 135", + "mov rcx, rdx", + "shr rcx, 63", + "add rax, rax", + "or rax, rcx", + "add rdx, rdx", + "xor rdx, r8", + "bswap rax", + "bswap rdx", + "mov QWORD PTR [rbx], rax", + "mov QWORD PTR [rbx+8], rdx", + "mov rax, QWORD PTR [rbx]", + "bswap rax", + "mov rdx, QWORD PTR [rbx+8]", + "bswap rdx", + "mov rcx, rax", + "shr rcx, 63", + "mov r8d, 0", + "sub r8, rcx", + "and r8, 135", + "mov rcx, rdx", + "shr rcx, 63", + "add rax, rax", + "or rax, rcx", + "add rdx, rdx", + "xor rdx, r8", + "bswap rax", + "bswap rdx", + "mov QWORD PTR [rbx+16], rax", + "mov QWORD PTR [rbx+24], rdx", + "mov rbx, QWORD PTR [rbp+2064]", + "mov rbp, QWORD PTR [rbp+2072]", + "ret", + vg_aes_ctr32_aesni = sym super::aes::vg_aes_ctr32_aesni, + ) +} + +/// The CPU features `vg_cmac_aes_update_aesni` requires (`Artifact.features`). +pub(crate) const VG_CMAC_AES_UPDATE_AESNI_FEATURES: &[&str] = &["aes", "ssse3"]; + +/// 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_aesni`. +/// +/// # 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` or `data` (distinct Rust objects never do). +/// * None of `schedule`, `state`, `data` and `scratch` may overlap the return address on the stack or the 8 bytes of stack below it, or wrap around the end of the address space (no Rust object does). +/// * The CPU must support the `aes` and `ssse3` target features. +#[unsafe(naked)] +pub(crate) unsafe extern "sysv64" fn vg_cmac_aes_update_aesni(schedule: *const [u8; 240], rounds: usize, state: *mut [u8; 16], data: *const [u8; 16], n: usize, scratch: *mut [u64; 272]) { + core::arch::naked_asm!( + "mov QWORD PTR [r9+2064], rbx", + "mov QWORD PTR [r9+2072], rbp", + "mov QWORD PTR [r9+2080], r12", + "mov QWORD PTR [r9+2088], r13", + "mov QWORD PTR [r9+2096], r14", + "mov QWORD PTR [r9+2104], r15", + "mov rbx, rdi", + "mov rbp, rsi", + "mov r12, rdx", + "mov r13, rcx", + "mov r14, r8", + "mov r15, r9", + "test r14, r14", + "je 20f", + "22:", + "mov rax, QWORD PTR [r12]", + "xor rax, QWORD PTR [r13]", + "mov QWORD PTR [r15+2048], rax", + "mov rax, QWORD PTR [r12+8]", + "xor rax, QWORD PTR [r13+8]", + "mov QWORD PTR [r15+2056], rax", + "mov eax, 0", + "mov QWORD PTR [r12], rax", + "mov QWORD PTR [r12+8], rax", + "mov rdi, rbx", + "mov rsi, rbp", + "mov r9, r15", + "mov rdx, r15", + "add rdx, 2048", + "mov rcx, r12", + "mov r8d, 1", + "call {vg_aes_ctr32_aesni}", + "add r13, 16", + "sub r14, 1", + "jne 22b", + "jmp 21f", + "20:", + "21:", + "mov rbx, QWORD PTR [r15+2064]", + "mov rbp, QWORD PTR [r15+2072]", + "mov r12, QWORD PTR [r15+2080]", + "mov r13, QWORD PTR [r15+2088]", + "mov r14, QWORD PTR [r15+2096]", + "mov r15, QWORD PTR [r15+2104]", + "ret", + vg_aes_ctr32_aesni = sym super::aes::vg_aes_ctr32_aesni, + ) +} + +/// The CPU features `vg_cmac_aes_finalize_aesni` requires (`Artifact.features`). +pub(crate) const VG_CMAC_AES_FINALIZE_AESNI_FEATURES: &[&str] = &["aes", "ssse3"]; + +/// 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_aesni`. +/// +/// # 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` or `last` (distinct Rust objects never do). +/// * None of `key`, `state`, `last` and `scratch` may overlap the return address on the stack or the 8 bytes of stack below it, or wrap around the end of the address space (no Rust object does). +/// * The CPU must support the `aes` and `ssse3` target features. +#[unsafe(naked)] +pub(crate) unsafe extern "sysv64" fn vg_cmac_aes_finalize_aesni(key: *const [u8; 272], rounds: usize, state: *mut [u8; 16], last: *const u8, last_len: usize, scratch: *mut [u64; 272]) { + core::arch::naked_asm!( + "cmp r8, 16", + "je 20f", + "mov eax, 0", + "mov QWORD PTR [r9+2048], rax", + "mov QWORD PTR [r9+2056], rax", + "test r8, r8", + "je 22f", + "mov r10d, 0", + "24:", + "movzx eax, BYTE PTR [rcx+r10*1]", + "mov BYTE PTR [r9+r10*1+2048], al", + "add r10, 1", + "cmp r10, r8", + "jne 24b", + "jmp 23f", + "22:", + "23:", + "mov eax, 128", + "mov BYTE PTR [r9+r8*1+2048], al", + "mov rax, QWORD PTR [r9+2048]", + "xor rax, QWORD PTR [rdi+256]", + "mov QWORD PTR [r9+2048], rax", + "mov rax, QWORD PTR [r9+2056]", + "xor rax, QWORD PTR [rdi+264]", + "mov QWORD PTR [r9+2056], rax", + "jmp 21f", + "20:", + "mov rax, QWORD PTR [rcx]", + "xor rax, QWORD PTR [rdi+240]", + "mov QWORD PTR [r9+2048], rax", + "mov rax, QWORD PTR [rcx+8]", + "xor rax, QWORD PTR [rdi+248]", + "mov QWORD PTR [r9+2056], rax", + "21:", + "mov rax, QWORD PTR [r9+2048]", + "xor rax, QWORD PTR [rdx]", + "mov QWORD PTR [r9+2048], rax", + "mov rax, QWORD PTR [r9+2056]", + "xor rax, QWORD PTR [rdx+8]", + "mov QWORD PTR [r9+2056], rax", + "mov eax, 0", + "mov QWORD PTR [rdx], rax", + "mov QWORD PTR [rdx+8], rax", + "mov rcx, rdx", + "mov rdx, r9", + "add rdx, 2048", + "mov r8d, 1", + "call {vg_aes_ctr32_aesni}", + "ret", + vg_aes_ctr32_aesni = sym super::aes::vg_aes_ctr32_aesni, + ) +} + +/// 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 return address on the stack or the 8 bytes of stack below it, or wrap around the end of the address space (no Rust object does). +#[unsafe(naked)] +pub(crate) unsafe extern "sysv64" fn vg_cmac_aes_subkeys(schedule: *const [u8; 240], rounds: usize, subkeys: *mut [u8; 32], scratch: *mut [u64; 272]) { + core::arch::naked_asm!( + "mov QWORD PTR [rcx+2064], rbx", + "mov QWORD PTR [rcx+2072], rbp", + "mov rbx, rdx", + "mov rbp, rcx", + "mov eax, 0", + "mov QWORD PTR [rcx+2048], rax", + "mov QWORD PTR [rcx+2056], rax", + "mov QWORD PTR [rdx], rax", + "mov QWORD PTR [rdx+8], rax", + "mov r9, rbp", + "mov rdx, rbp", + "add rdx, 2048", + "mov rcx, rbx", + "mov r8d, 1", + "call {vg_aes_ctr32}", + "mov rax, QWORD PTR [rbx]", + "bswap rax", + "mov rdx, QWORD PTR [rbx+8]", + "bswap rdx", + "mov rcx, rax", + "shr rcx, 63", + "mov r8d, 0", + "sub r8, rcx", + "and r8, 135", + "mov rcx, rdx", + "shr rcx, 63", + "add rax, rax", + "or rax, rcx", + "add rdx, rdx", + "xor rdx, r8", + "bswap rax", + "bswap rdx", + "mov QWORD PTR [rbx], rax", + "mov QWORD PTR [rbx+8], rdx", + "mov rax, QWORD PTR [rbx]", + "bswap rax", + "mov rdx, QWORD PTR [rbx+8]", + "bswap rdx", + "mov rcx, rax", + "shr rcx, 63", + "mov r8d, 0", + "sub r8, rcx", + "and r8, 135", + "mov rcx, rdx", + "shr rcx, 63", + "add rax, rax", + "or rax, rcx", + "add rdx, rdx", + "xor rdx, r8", + "bswap rax", + "bswap rdx", + "mov QWORD PTR [rbx+16], rax", + "mov QWORD PTR [rbx+24], rdx", + "mov rbx, QWORD PTR [rbp+2064]", + "mov rbp, QWORD PTR [rbp+2072]", + "ret", + 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` or `data` (distinct Rust objects never do). +/// * None of `schedule`, `state`, `data` and `scratch` may overlap the return address on the stack or the 8 bytes of stack below it, or wrap around the end of the address space (no Rust object does). +#[unsafe(naked)] +pub(crate) unsafe extern "sysv64" 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!( + "mov QWORD PTR [r9+2064], rbx", + "mov QWORD PTR [r9+2072], rbp", + "mov QWORD PTR [r9+2080], r12", + "mov QWORD PTR [r9+2088], r13", + "mov QWORD PTR [r9+2096], r14", + "mov QWORD PTR [r9+2104], r15", + "mov rbx, rdi", + "mov rbp, rsi", + "mov r12, rdx", + "mov r13, rcx", + "mov r14, r8", + "mov r15, r9", + "test r14, r14", + "je 20f", + "22:", + "mov rax, QWORD PTR [r12]", + "xor rax, QWORD PTR [r13]", + "mov QWORD PTR [r15+2048], rax", + "mov rax, QWORD PTR [r12+8]", + "xor rax, QWORD PTR [r13+8]", + "mov QWORD PTR [r15+2056], rax", + "mov eax, 0", + "mov QWORD PTR [r12], rax", + "mov QWORD PTR [r12+8], rax", + "mov rdi, rbx", + "mov rsi, rbp", + "mov r9, r15", + "mov rdx, r15", + "add rdx, 2048", + "mov rcx, r12", + "mov r8d, 1", + "call {vg_aes_ctr32}", + "add r13, 16", + "sub r14, 1", + "jne 22b", + "jmp 21f", + "20:", + "21:", + "mov rbx, QWORD PTR [r15+2064]", + "mov rbp, QWORD PTR [r15+2072]", + "mov r12, QWORD PTR [r15+2080]", + "mov r13, QWORD PTR [r15+2088]", + "mov r14, QWORD PTR [r15+2096]", + "mov r15, QWORD PTR [r15+2104]", + "ret", + 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` or `last` (distinct Rust objects never do). +/// * None of `key`, `state`, `last` and `scratch` may overlap the return address on the stack or the 8 bytes of stack below it, or wrap around the end of the address space (no Rust object does). +#[unsafe(naked)] +pub(crate) unsafe extern "sysv64" 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!( + "cmp r8, 16", + "je 20f", + "mov eax, 0", + "mov QWORD PTR [r9+2048], rax", + "mov QWORD PTR [r9+2056], rax", + "test r8, r8", + "je 22f", + "mov r10d, 0", + "24:", + "movzx eax, BYTE PTR [rcx+r10*1]", + "mov BYTE PTR [r9+r10*1+2048], al", + "add r10, 1", + "cmp r10, r8", + "jne 24b", + "jmp 23f", + "22:", + "23:", + "mov eax, 128", + "mov BYTE PTR [r9+r8*1+2048], al", + "mov rax, QWORD PTR [r9+2048]", + "xor rax, QWORD PTR [rdi+256]", + "mov QWORD PTR [r9+2048], rax", + "mov rax, QWORD PTR [r9+2056]", + "xor rax, QWORD PTR [rdi+264]", + "mov QWORD PTR [r9+2056], rax", + "jmp 21f", + "20:", + "mov rax, QWORD PTR [rcx]", + "xor rax, QWORD PTR [rdi+240]", + "mov QWORD PTR [r9+2048], rax", + "mov rax, QWORD PTR [rcx+8]", + "xor rax, QWORD PTR [rdi+248]", + "mov QWORD PTR [r9+2056], rax", + "21:", + "mov rax, QWORD PTR [r9+2048]", + "xor rax, QWORD PTR [rdx]", + "mov QWORD PTR [r9+2048], rax", + "mov rax, QWORD PTR [r9+2056]", + "xor rax, QWORD PTR [rdx+8]", + "mov QWORD PTR [r9+2056], rax", + "mov eax, 0", + "mov QWORD PTR [rdx], rax", + "mov QWORD PTR [rdx+8], rax", + "mov rcx, rdx", + "mov rdx, r9", + "add rdx, 2048", + "mov r8d, 1", + "call {vg_aes_ctr32}", + "ret", + vg_aes_ctr32 = sym super::aes::vg_aes_ctr32, + ) +} diff --git a/src/asm/x86_64/mod.rs b/src/asm/x86_64/mod.rs index a5c654b2a..7816e4666 100644 --- a/src/asm/x86_64/mod.rs +++ b/src/asm/x86_64/mod.rs @@ -19,6 +19,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 new file mode 100644 index 000000000..ce29b48ab --- /dev/null +++ b/src/cmac/aes.rs @@ -0,0 +1,305 @@ +//! AES-CMAC (NIST SP 800-38B, RFC 4493) with 128-, 192- and 256-bit AES +//! keys and the full 16-byte MAC. +//! +//! Everything but buffering is the verified assembly for the target +//! architecture: `vg_aes_expand_key` (contract +//! `VG.Spec.Aes.expandKeyContract`) expands the key, `vg_cmac_aes_subkeys` +//! (`VG.Spec.Cmac.aesSubkeysContract`) derives the subkeys `K1` and `K2`, +//! `vg_cmac_aes_update` (`VG.Spec.Cmac.aesUpdateContract`) chains whole +//! blocks, and `vg_cmac_aes_finalize` (`VG.Spec.Cmac.aesFinalizeContract`) +//! masks and pads the last block and encrypts it. This module only keeps the +//! last (possibly whole) block of what it has absorbed back for `finalize`. +//! +//! On x86-64, CPUs with AES-NI and SSSE3 run `vg_aes_expand_key_aesni` and +//! the `_aesni` CMAC functions instead, which have the same contracts: the +//! same verified CMAC code, calling `vg_aes_ctr32_aesni` rather than +//! `vg_aes_ctr32` to encrypt each block. + +#![cfg(target_arch = "x86_64")] + +use super::{InvalidKeyLength, InvalidMac}; +use crate::arch::aes::vg_aes_expand_key; +#[cfg(target_arch = "x86_64")] +use crate::arch::aes::{VG_AES_EXPAND_KEY_AESNI_FEATURES, vg_aes_expand_key_aesni}; +#[cfg(target_arch = "x86_64")] +use crate::arch::cmac_aes::{ + VG_CMAC_AES_FINALIZE_AESNI_FEATURES, VG_CMAC_AES_SUBKEYS_AESNI_FEATURES, + VG_CMAC_AES_UPDATE_AESNI_FEATURES, vg_cmac_aes_finalize_aesni, vg_cmac_aes_subkeys_aesni, + vg_cmac_aes_update_aesni, +}; +use crate::arch::cmac_aes::{vg_cmac_aes_finalize, vg_cmac_aes_subkeys, vg_cmac_aes_update}; +use crate::cpu::{Features, detected}; +use crate::zeroize::zeroize; +use core::mem::MaybeUninit; + +/// A 16-byte block. +type Block = [u8; 16]; + +/// The working space of the CMAC functions, in 64-bit words. +const SCRATCH: usize = 272; + +/// The implementations of the primitives. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum Backend { + /// Constant-time scalar code, for the target's baseline ISA. + Scalar, + /// AES-NI. + #[cfg(target_arch = "x86_64")] + AesNi, +} + +impl Backend { + /// The best implementation a CPU with the features `f` can run. + #[cfg(target_arch = "x86_64")] + fn select(f: Features) -> Backend { + if f.contains(Features::all(&[ + VG_AES_EXPAND_KEY_AESNI_FEATURES, + VG_CMAC_AES_SUBKEYS_AESNI_FEATURES, + VG_CMAC_AES_UPDATE_AESNI_FEATURES, + VG_CMAC_AES_FINALIZE_AESNI_FEATURES, + ])) { + Backend::AesNi + } else { + Backend::Scalar + } + } +} + +/// An incremental AES-CMAC computation. +/// +/// A computation that has absorbed nothing yet can be cloned to MAC several +/// messages with the same key without expanding it again. +#[derive(Clone)] +pub struct AesCmac { + /// The AES key schedule (240 bytes), then the subkeys `K1 ‖ K2`. + key: [u8; 272], + rounds: usize, + backend: Backend, + /// The chaining value: the encryption of the blocks absorbed so far, + /// chained from the zero block. + state: Block, + /// What has been absorbed after them… + buf: Block, + /// …of which this many bytes: at most a block, and more than none once + /// anything has been absorbed. + buf_len: usize, +} + +impl Drop for AesCmac { + /// Wipes the key schedule, the subkeys, the chaining value and the + /// buffered data. + fn drop(&mut self) { + zeroize(&mut self.key); + zeroize(&mut self.state); + zeroize(&mut self.buf); + } +} + +impl AesCmac { + /// The size of the MAC, in bytes. + pub const MAC_SIZE: usize = 16; + + /// Starts an AES-CMAC computation with `key`, which must be 16, 24 or 32 + /// bytes long (AES-128, AES-192 or AES-256). + pub fn new(key: &[u8]) -> Result { + if !matches!(key.len(), 16 | 24 | 32) { + return Err(InvalidKeyLength); + } + let mut c = AesCmac { + key: [0; 272], + rounds: key.len() / 4 + 6, + backend: Backend::select(detected()), + state: [0; 16], + buf: [0; 16], + buf_len: 0, + }; + let (schedule, subkeys) = c.key.split_first_chunk_mut::<240>().unwrap(); + let subkeys: &mut [u8; 32] = subkeys.try_into().unwrap(); + let mut expand_scratch = MaybeUninit::<[u64; 64]>::uninit(); + let mut scratch = MaybeUninit::<[u64; SCRATCH]>::uninit(); + let (k, rounds) = (key.as_ptr(), c.rounds); + let (e, s) = (expand_scratch.as_mut_ptr(), scratch.as_mut_ptr()); + // SAFETY: `key` is valid for reads of `key.len()` bytes, which is 16, + // 24 or 32; `schedule` for reads and writes of 240 bytes, `subkeys` + // of 32, `expand_scratch` of 512 and `scratch` of 2176. `schedule` + // and `subkeys` are disjoint parts of `c.key`, and the scratch + // buffers locals, so no two overlap each other, `key` or the return + // address. Key expansion leaves the schedule for `rounds` (10, 12 or + // 14) rounds in `schedule`, which the subkey derivation reads. The + // CPU has the features of the implementation selected. The scratch + // buffers are uninitialized: they are only working space, and + // neither contract's result depends on what they hold. + unsafe { + match c.backend { + Backend::Scalar => { + vg_aes_expand_key(k, key.len(), schedule, e); + vg_cmac_aes_subkeys(schedule, rounds, subkeys, s); + } + #[cfg(target_arch = "x86_64")] + Backend::AesNi => { + vg_aes_expand_key_aesni(k, key.len(), schedule, e); + vg_cmac_aes_subkeys_aesni(schedule, rounds, subkeys, s); + } + } + } + Ok(c) + } + + /// Chains the whole blocks `blocks` into the state. + fn blocks(&mut self, blocks: &[Block]) { + if blocks.is_empty() { + return; + } + let mut scratch = MaybeUninit::<[u64; SCRATCH]>::uninit(); + let f = match self.backend { + Backend::Scalar => vg_cmac_aes_update, + #[cfg(target_arch = "x86_64")] + Backend::AesNi => vg_cmac_aes_update_aesni, + }; + let schedule = self.key.first_chunk::<240>().unwrap(); + // SAFETY: `schedule` holds the key schedule for `self.rounds` (10, + // 12 or 14) rounds, written by key expansion (every implementation + // writes the same one); it is valid for reads of 240 bytes, + // `self.state` for reads and writes of 16, `blocks` for reads of + // `16 * blocks.len()` and `scratch` (uninitialized working space, as + // in `new`) for reads and writes of 2176. `self.state` is a mutable + // borrow and `scratch` a local, so neither overlaps another argument + // or the return address. The CPU has the features of the + // implementation selected. + unsafe { + f( + schedule, + self.rounds, + &mut self.state, + blocks.as_ptr(), + blocks.len(), + scratch.as_mut_ptr(), + ) + }; + } + + /// Absorbs `data`. + pub fn update(&mut self, mut data: &[u8]) { + if data.is_empty() { + return; + } + // Fill the buffer, and chain it only if more data follows: the last + // block, even a whole one, is `finalize`'s. + if self.buf_len > 0 { + let n = data.len().min(16 - self.buf_len); + self.buf[self.buf_len..self.buf_len + n].copy_from_slice(&data[..n]); + self.buf_len += n; + data = &data[n..]; + if data.is_empty() { + return; + } + let buf = self.buf; + self.blocks(&[buf]); + } + let (blocks, _) = data[..data.len() - 1].as_chunks::<16>(); + self.blocks(blocks); + let rest = &data[16 * blocks.len()..]; + self.buf[..rest.len()].copy_from_slice(rest); + self.buf_len = rest.len(); + } + + /// Returns the MAC of everything absorbed. + pub fn finalize(mut self) -> [u8; 16] { + let mut scratch = MaybeUninit::<[u64; SCRATCH]>::uninit(); + let f = match self.backend { + Backend::Scalar => vg_cmac_aes_finalize, + #[cfg(target_arch = "x86_64")] + Backend::AesNi => vg_cmac_aes_finalize_aesni, + }; + // SAFETY: `self.key` holds the key schedule for `self.rounds` (10, 12 + // or 14) rounds and then its subkeys, written by `new`; it is valid + // for reads of 272 bytes, `self.state` for reads and writes of 16, + // `self.buf` for reads of `self.buf_len` (at most 16) and `scratch` + // (uninitialized working space, as in `new`) for reads and writes of + // 2176. `self.state` is a mutable borrow and `scratch` a local, so + // neither overlaps another argument or the return address. The state + // is the chaining of the blocks before the buffered ones, and the + // buffer holds at least a byte if any were chained, as the contract's + // postcondition requires to give the MAC. The CPU has the features + // of the implementation selected. + unsafe { + f( + &self.key, + self.rounds, + &mut self.state, + self.buf.as_ptr(), + self.buf_len, + scratch.as_mut_ptr(), + ) + }; + self.state + } + + /// Checks that `mac` is the MAC of everything absorbed, in constant + /// time: the time taken does not depend on where, or whether, `mac` + /// differs from it (its length is public). `mac` must be the whole MAC, + /// of [`MAC_SIZE`](Self::MAC_SIZE) bytes; a truncated one is rejected. + pub fn verify(self, mac: &[u8]) -> Result<(), InvalidMac> { + if crate::ct::eq(&self.finalize(), mac) { + Ok(()) + } else { + Err(InvalidMac) + } + } + + /// The MAC of `data` with `key`, which must be 16, 24 or 32 bytes long. + pub fn mac(key: &[u8], data: &[u8]) -> Result<[u8; 16], InvalidKeyLength> { + let mut c = Self::new(key)?; + c.update(data); + Ok(c.finalize()) + } +} + +#[cfg(test)] +mod tests { + use super::{AesCmac, Backend}; + use crate::cmac::InvalidKeyLength; + use crate::cpu::Features; + + /// Keys of other lengths are rejected. + #[test] + fn key_lengths() { + for len in [0, 15, 17, 23, 25, 31, 33, 64] { + assert_eq!(AesCmac::new(&[0; 64][..len]).err(), Some(InvalidKeyLength)); + assert_eq!( + AesCmac::mac(&[0; 64][..len], b"").err(), + Some(InvalidKeyLength) + ); + } + } + + /// Absorbing a message in two pieces, split anywhere, gives the MAC of + /// the whole, also from a clone of a computation that absorbed nothing. + #[test] + fn splits() { + let key = [0x2b; 16]; + let msg: [u8; 50] = core::array::from_fn(|i| i as u8); + let fresh = AesCmac::new(&key).unwrap(); + for len in 0..=msg.len() { + let mac = AesCmac::mac(&key, &msg[..len]).unwrap(); + for split in 0..=len { + let mut c = fresh.clone(); + c.update(&msg[..split]); + c.update(&msg[split..len]); + assert_eq!(c.finalize(), mac, "{split} of {len}"); + } + } + } + + /// Each implementation is selected exactly when the CPU has its + /// features. + #[test] + fn select() { + assert_eq!(Backend::select(Features::of(&[])), Backend::Scalar); + assert_eq!( + Backend::select(Features::of(&["aes", "ssse3"])), + Backend::AesNi + ); + assert_eq!(Backend::select(Features::of(&["aes"])), Backend::Scalar); + } +} diff --git a/src/cmac/mod.rs b/src/cmac/mod.rs new file mode 100644 index 000000000..27ddd55b2 --- /dev/null +++ b/src/cmac/mod.rs @@ -0,0 +1,22 @@ +//! CMAC (NIST SP 800-38B, RFC 4493): a MAC built from a block cipher. +//! +//! Each block cipher with a verified CMAC implementation has a module of its +//! own here: [`aes`]. + +#![cfg(any( + target_arch = "x86_64", + target_arch = "aarch64", + target_arch = "arm", + target_arch = "x86" +))] + +pub mod aes; + +/// The key does not have a length the block cipher accepts. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct InvalidKeyLength; + +/// The MAC did not match: the message or the key is not what was +/// authenticated. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct InvalidMac; diff --git a/src/lib.rs b/src/lib.rs index c4d397489..210dc7a78 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -81,6 +81,7 @@ compile_error!("32-bit ARM needs an AAPCS target (not Apple's armv7s or armv7k)" pub mod aes_gcm; pub mod chacha20; pub mod chacha20poly1305; +pub mod cmac; mod ct; pub mod ed25519; pub mod hashes; diff --git a/tests/cavp/cmac_aes.rs b/tests/cavp/cmac_aes.rs new file mode 100644 index 000000000..1ae0ddf5e --- /dev/null +++ b/tests/cavp/cmac_aes.rs @@ -0,0 +1,96 @@ +//! 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(target_arch = "x86_64")] + +use verified_garbage::cmac::aes::AesCmac; + +use super::{fields, unhex}; + +/// The vectors of a CMAC response file: the fields of each, from its +/// `Count` line to the next one. +fn vectors(text: &str) -> Vec> { + let mut vs: Vec> = Vec::new(); + for (k, v) in fields(text) { + if k == "Count" { + vs.push(Vec::new()); + } + vs.last_mut().unwrap().push((k, v)); + } + vs +} + +/// The value of the field `key` of a vector. +fn field<'a>(v: &[(&str, &'a str)], key: &str) -> &'a str { + v.iter().find(|(k, _)| *k == key).unwrap().1 +} + +/// Whether the MAC of a vector's message (truncated to its `Tlen`) is its +/// `Mac`; its key must be `bytes` long. +fn matches(v: &[(&str, &str)], bytes: usize) -> bool { + let key = unhex(field(v, "Key")); + assert_eq!(key.len(), bytes); + assert_eq!(field(v, "Klen").parse::().unwrap(), bytes); + // `Mlen` is in bytes; the empty message is written as `Msg = 00`. + let len: usize = field(v, "Mlen").parse().unwrap(); + let msg = unhex(field(v, "Msg")); + let tlen: usize = field(v, "Tlen").parse().unwrap(); + let mac = unhex(field(v, "Mac")); + assert_eq!(mac.len(), tlen); + AesCmac::mac(&key, &msg[..len]).unwrap()[..tlen] == mac[..] +} + +#[test] +fn generate() { + let files = [ + ( + include_str!("../../vectors/nist-cavp/cmac-aes/CMACGenAES128.rsp"), + 16, + ), + ( + include_str!("../../vectors/nist-cavp/cmac-aes/CMACGenAES192.rsp"), + 24, + ), + ( + include_str!("../../vectors/nist-cavp/cmac-aes/CMACGenAES256.rsp"), + 32, + ), + ]; + let mut n = 0; + for (text, bytes) in files { + for v in vectors(text) { + assert!(matches(&v, bytes), "{v:?}"); + n += 1; + } + } + assert_eq!(n, 96 + 144 + 96); +} + +#[test] +fn verify() { + let files = [ + ( + include_str!("../../vectors/nist-cavp/cmac-aes/CMACVerAES128.rsp"), + 16, + ), + ( + include_str!("../../vectors/nist-cavp/cmac-aes/CMACVerAES256.rsp"), + 32, + ), + ]; + let (mut pass, mut fail) = (0, 0); + for (text, bytes) in files { + for v in vectors(text) { + let result = field(&v, "Result"); + if result == "P" { + assert!(matches(&v, bytes), "{v:?}"); + pass += 1; + } else { + assert!(result.starts_with("F "), "{v:?}"); + assert!(!matches(&v, bytes), "{v:?}"); + fail += 1; + } + } + } + assert_eq!((pass, fail), (48 + 48, 192 + 192)); +} diff --git a/tests/cavp/main.rs b/tests/cavp/main.rs index 98a9c7226..e3547f522 100644 --- a/tests/cavp/main.rs +++ b/tests/cavp/main.rs @@ -14,6 +14,7 @@ ))] mod aes_gcm; +mod cmac_aes; mod rc2_cbc; mod sha1; mod sha256; diff --git a/tests/wycheproof/cmac_aes.rs b/tests/wycheproof/cmac_aes.rs new file mode 100644 index 000000000..5ec8ee84b --- /dev/null +++ b/tests/wycheproof/cmac_aes.rs @@ -0,0 +1,72 @@ +//! AES-CMAC (`MacTest` vectors, `aes_cmac_test.json`). +//! +//! A valid vector's tag must be the MAC of its message, computed at once, +//! in one `update` and a byte at a time, and `verify` must accept it. An +//! 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(target_arch = "x86_64")] + +use serde::Deserialize; +use verified_garbage::cmac::InvalidKeyLength; +use verified_garbage::cmac::aes::AesCmac; + +use crate::harness::{self, Expectation, Hex}; +use crate::require_vectors; + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase")] +struct Group { + key_size: usize, + tag_size: usize, +} + +#[derive(Deserialize)] +struct Case { + key: Hex, + msg: Hex, + tag: Hex, +} + +#[test] +fn cmac_aes() { + require_vectors!(); + let file = harness::load::("aes_cmac_test.json"); + let (mut valid, mut modified, mut bad_keys) = (0, 0, 0); + for (group, test) in file.tests() { + let Case { key, msg, tag } = &test.case; + let id = test.tc_id; + assert_eq!(key.0.len() * 8, group.params.key_size, "tcId {id}"); + assert_eq!(group.params.tag_size, 128, "tcId {id}"); + if !matches!(key.0.len(), 16 | 24 | 32) { + assert_eq!(test.result, Expectation::Invalid, "tcId {id}"); + let err = AesCmac::new(&key.0).err(); + assert_eq!(err, Some(InvalidKeyLength), "tcId {id}"); + bad_keys += 1; + continue; + } + let full = AesCmac::mac(&key.0, &msg.0).unwrap(); + let mut c = AesCmac::new(&key.0).unwrap(); + c.update(&msg.0); + assert_eq!(c.finalize(), full, "tcId {id}"); + let mut c = AesCmac::new(&key.0).unwrap(); + for byte in &msg.0 { + c.update(core::slice::from_ref(byte)); + } + assert_eq!(c.finalize(), full, "tcId {id}"); + let mut c = AesCmac::new(&key.0).unwrap(); + c.update(&msg.0); + let ok = c.verify(&tag.0).is_ok(); + if test.result == Expectation::Valid { + assert!(ok, "tcId {id}"); + assert_eq!(&full[..], &tag.0[..], "tcId {id}"); + valid += 1; + } else { + assert_eq!(test.result, Expectation::Invalid, "tcId {id}"); + assert!(!ok, "tcId {id}"); + assert_ne!(&full[..], &tag.0[..], "tcId {id}"); + modified += 1; + } + } + assert_eq!((valid, modified, bad_keys), (63, 243, 5)); +} diff --git a/tests/wycheproof/main.rs b/tests/wycheproof/main.rs index 4dad9837e..3cb2a8c01 100644 --- a/tests/wycheproof/main.rs +++ b/tests/wycheproof/main.rs @@ -12,6 +12,7 @@ mod aes_gcm; mod chacha20; mod chacha20poly1305; +mod cmac_aes; mod ed25519; mod harness; mod hmac;