From 159a2c04a017f4bc8e5a4d32fe74ac72007e6252 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 13:24:32 +0000 Subject: [PATCH 1/4] Implement AES-CMAC on x86-64, with an AES-NI variant The three AES-CMAC primitives specified in VG.Spec.Cmac (vg_cmac_aes_subkeys, vg_cmac_aes_update, vg_cmac_aes_finalize), proven against their contracts on x86-64: correct and constant time. Each encrypts a block by calling the verified vg_aes_ctr32 on a zero block with the block as the counter, so the CMAC code is generic over the implementations of ctr32 (a new AesCtr32 interface under Generic/ and Variants/): the emitter also emits the _aesni functions, calling vg_aes_ctr32_aesni. The Rust API is verified_garbage::cmac::aes::AesCmac (new, update, finalize, verify, mac), choosing the implementation by CPU feature. It is tested against Wycheproof's aes_cmac_test.json and every vector of the vendored NIST CAVP CMAC generation and verification files, and benchmarked against OpenSSL. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01WkLN6tAYk76HACiEWLbAMD --- README.md | 2 +- bench/benches/primitives/cmac_aes.rs | 67 +++ bench/benches/primitives/main.rs | 2 + .../Generic/AesCtr32/X86_64/CmacAes.lean | 62 +++ .../Impl/Aes/X86_64/Callee.lean | 23 + lean/VerifiedGarbage/Impl/CmacAes/X86_64.lean | 168 +++++++ .../Proof/Aes/X86_64/Variant.lean | 88 ++++ lean/VerifiedGarbage/Proof/Cmac/Dbl.lean | 102 ++++ lean/VerifiedGarbage/Proof/Cmac/Mem.lean | 85 ++++ lean/VerifiedGarbage/Proof/Cmac/Spec.lean | 114 +++++ .../Proof/CmacAes/X86_64/Call.lean | 141 ++++++ .../Proof/CmacAes/X86_64/Contract.lean | 97 ++++ .../Proof/CmacAes/X86_64/Dbl.lean | 102 ++++ .../Proof/CmacAes/X86_64/Finalize.lean | 257 ++++++++++ .../Proof/CmacAes/X86_64/FinalizeCT.lean | 47 ++ .../Proof/CmacAes/X86_64/FinalizeCorrect.lean | 386 +++++++++++++++ .../Proof/CmacAes/X86_64/Subkeys.lean | 400 ++++++++++++++++ .../Proof/CmacAes/X86_64/SubkeysCT.lean | 86 ++++ .../Proof/CmacAes/X86_64/Update.lean | 128 +++++ .../Proof/CmacAes/X86_64/UpdateCT.lean | 193 ++++++++ .../Proof/CmacAes/X86_64/UpdateCorrect.lean | 181 +++++++ .../Proof/CmacAes/X86_64/UpdateLoop.lean | 332 +++++++++++++ .../Proof/CmacAes/X86_64/Verified.lean | 108 +++++ .../Variants/AesCtr32/X86_64/AesNi.lean | 14 + .../Variants/AesCtr32/X86_64/Scalar.lean | 14 + src/asm/x86_64/cmac_aes.rs | 453 ++++++++++++++++++ src/asm/x86_64/mod.rs | 3 + src/cmac/aes.rs | 305 ++++++++++++ src/cmac/mod.rs | 22 + src/lib.rs | 1 + tests/cavp/cmac_aes.rs | 96 ++++ tests/cavp/main.rs | 1 + tests/wycheproof/cmac_aes.rs | 75 +++ tests/wycheproof/main.rs | 1 + 34 files changed, 4155 insertions(+), 1 deletion(-) create mode 100644 bench/benches/primitives/cmac_aes.rs create mode 100644 lean/VerifiedGarbage/Generic/AesCtr32/X86_64/CmacAes.lean create mode 100644 lean/VerifiedGarbage/Impl/Aes/X86_64/Callee.lean create mode 100644 lean/VerifiedGarbage/Impl/CmacAes/X86_64.lean create mode 100644 lean/VerifiedGarbage/Proof/Aes/X86_64/Variant.lean create mode 100644 lean/VerifiedGarbage/Proof/Cmac/Dbl.lean create mode 100644 lean/VerifiedGarbage/Proof/Cmac/Mem.lean create mode 100644 lean/VerifiedGarbage/Proof/Cmac/Spec.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/X86_64/Call.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/X86_64/Contract.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/X86_64/Dbl.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/X86_64/Finalize.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/X86_64/FinalizeCT.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/X86_64/FinalizeCorrect.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/X86_64/Subkeys.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/X86_64/SubkeysCT.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/X86_64/Update.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/X86_64/UpdateCT.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/X86_64/UpdateCorrect.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/X86_64/UpdateLoop.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/X86_64/Verified.lean create mode 100644 lean/VerifiedGarbage/Variants/AesCtr32/X86_64/AesNi.lean create mode 100644 lean/VerifiedGarbage/Variants/AesCtr32/X86_64/Scalar.lean create mode 100644 src/asm/x86_64/cmac_aes.rs create mode 100644 src/cmac/aes.rs create mode 100644 src/cmac/mod.rs create mode 100644 tests/cavp/cmac_aes.rs create mode 100644 tests/wycheproof/cmac_aes.rs 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..7aa5967c6 --- /dev/null +++ b/tests/wycheproof/cmac_aes.rs @@ -0,0 +1,75 @@ +//! 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}"); + assert_eq!( + AesCmac::new(&key.0).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; From 72b0dac736ffa489762e0172fe01bfb387e55d2c Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 15:09:11 +0000 Subject: [PATCH 2/4] Implement AES-CMAC on AArch64, with an AES-extension variant The three AES-CMAC primitives specified in VG.Spec.Cmac (vg_cmac_aes_subkeys, vg_cmac_aes_update, vg_cmac_aes_finalize), proven against their contracts on AArch64: correct and constant time. As on x86-64, each encrypts a block by calling the verified vg_aes_ctr32 on a zero block with the block as the counter, and is generic over the implementations of ctr32 (the AesCtr32 interface on AArch64): the emitter also emits the _aes functions, which call vg_aes_ctr32_aes. The calls keep the return address in x30, which each function saves in the scratch buffer, so no stack is used. The lemmas about blocks in memory that do not depend on the target move to Proof/Cmac/Frame.lean and Proof/Cmac/Block.lean. The Rust API (verified_garbage::cmac::aes::AesCmac), its tests and its benchmark now cover AArch64 too, choosing the AES extension's implementation when the CPU has it. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01WkLN6tAYk76HACiEWLbAMD --- README.md | 2 +- bench/benches/primitives/cmac_aes.rs | 4 +- .../Generic/AesCtr32/AArch64/CmacAes.lean | 59 +++ .../Impl/Aes/AArch64/Callee.lean | 23 + .../VerifiedGarbage/Impl/CmacAes/AArch64.lean | 172 +++++++ .../Proof/Aes/AArch64/Variant.lean | 64 +++ lean/VerifiedGarbage/Proof/Cmac/Block.lean | 114 ++++ lean/VerifiedGarbage/Proof/Cmac/Frame.lean | 74 +++ .../Proof/CmacAes/AArch64/Call.lean | 107 ++++ .../Proof/CmacAes/AArch64/Contract.lean | 85 +++ .../Proof/CmacAes/AArch64/Dbl.lean | 98 ++++ .../Proof/CmacAes/AArch64/Finalize.lean | 198 +++++++ .../Proof/CmacAes/AArch64/FinalizeCT.lean | 61 +++ .../CmacAes/AArch64/FinalizeCorrect.lean | 450 ++++++++++++++++ .../Proof/CmacAes/AArch64/Subkeys.lean | 372 ++++++++++++++ .../Proof/CmacAes/AArch64/SubkeysCT.lean | 86 ++++ .../Proof/CmacAes/AArch64/Update.lean | 68 +++ .../Proof/CmacAes/AArch64/UpdateCT.lean | 196 +++++++ .../Proof/CmacAes/AArch64/UpdateCorrect.lean | 184 +++++++ .../Proof/CmacAes/AArch64/UpdateLoop.lean | 306 +++++++++++ .../Proof/CmacAes/AArch64/Verified.lean | 88 ++++ .../Variants/AesCtr32/AArch64/Aese.lean | 14 + .../Variants/AesCtr32/AArch64/Scalar.lean | 14 + src/asm/aarch64/cmac_aes.rs | 485 ++++++++++++++++++ src/asm/aarch64/mod.rs | 3 + src/cmac/aes.rs | 51 +- tests/cavp/cmac_aes.rs | 2 +- tests/wycheproof/cmac_aes.rs | 2 +- 28 files changed, 3375 insertions(+), 7 deletions(-) create mode 100644 lean/VerifiedGarbage/Generic/AesCtr32/AArch64/CmacAes.lean create mode 100644 lean/VerifiedGarbage/Impl/Aes/AArch64/Callee.lean create mode 100644 lean/VerifiedGarbage/Impl/CmacAes/AArch64.lean create mode 100644 lean/VerifiedGarbage/Proof/Aes/AArch64/Variant.lean create mode 100644 lean/VerifiedGarbage/Proof/Cmac/Block.lean create mode 100644 lean/VerifiedGarbage/Proof/Cmac/Frame.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/AArch64/Call.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/AArch64/Contract.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/AArch64/Dbl.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/AArch64/Finalize.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/AArch64/FinalizeCT.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/AArch64/FinalizeCorrect.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/AArch64/Subkeys.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/AArch64/SubkeysCT.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/AArch64/Update.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/AArch64/UpdateCT.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/AArch64/UpdateCorrect.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/AArch64/UpdateLoop.lean create mode 100644 lean/VerifiedGarbage/Proof/CmacAes/AArch64/Verified.lean create mode 100644 lean/VerifiedGarbage/Variants/AesCtr32/AArch64/Aese.lean create mode 100644 lean/VerifiedGarbage/Variants/AesCtr32/AArch64/Scalar.lean create mode 100644 src/asm/aarch64/cmac_aes.rs diff --git a/README.md b/README.md index 272a85b9b..93935d73b 100644 --- a/README.md +++ b/README.md @@ -201,7 +201,7 @@ yours to keep: ✅ AES-NI -❌ +✅ AES, PMULL ❌ diff --git a/bench/benches/primitives/cmac_aes.rs b/bench/benches/primitives/cmac_aes.rs index 3de7113fe..51a8dbf1c 100644 --- a/bench/benches/primitives/cmac_aes.rs +++ b/bench/benches/primitives/cmac_aes.rs @@ -8,7 +8,7 @@ pub const USES: &[&str] = &["cmac_aes", "aes"]; /// The MAC of a message with a 16-byte key (setup included), computed and /// verified. -#[cfg(target_arch = "x86_64")] +#[cfg(any(target_arch = "x86_64", target_arch = "aarch64"))] pub fn bench(c: &mut Criterion) { use std::hint::black_box; @@ -63,5 +63,5 @@ pub fn bench(c: &mut Criterion) { g.finish(); } -#[cfg(not(target_arch = "x86_64"))] +#[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))] pub fn bench(_: &mut Criterion) {} diff --git a/lean/VerifiedGarbage/Generic/AesCtr32/AArch64/CmacAes.lean b/lean/VerifiedGarbage/Generic/AesCtr32/AArch64/CmacAes.lean new file mode 100644 index 000000000..0815f3b03 --- /dev/null +++ b/lean/VerifiedGarbage/Generic/AesCtr32/AArch64/CmacAes.lean @@ -0,0 +1,59 @@ +import VerifiedGarbage.TCB.AArch64.Target +import VerifiedGarbage.Proof.CmacAes.AArch64.Verified + +/-! +# AES-CMAC (NIST SP 800-38B) on AArch64 + +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/AArch64/`), named with its suffix (e.g. +`vg_cmac_aes_update_aes`), 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 functions use no stack: their calls (`bl`) keep the return address in +`x30`, which they save in the scratch buffer. +-/ + +namespace VG.Generic.AesCtr32.AArch64.CmacAes + +open VG.Proof.CmacAes.AArch64 + +/-- Which implementation of `vg_aes_ctr32` an instance calls. -/ +def ctrNote (v : Proof.Aes.AArch64.Ctr32Impl) : String := + "This implementation encrypts each block with `" ++ v.callee.name ++ "`." + +def artifacts (v : Proof.Aes.AArch64.Ctr32Impl) : List Artifact := [ + { Spec.Cmac.aesSubkeysApi with + name := Spec.Cmac.aesSubkeysApi.name ++ v.suffix + target := AArch64.target + doc := Spec.Cmac.aesSubkeysApi.doc (notes := [ctrNote v]) + code := Impl.CmacAes.AArch64.subkeys v.callee + contract := Spec.Cmac.aesSubkeysContract AArch64.abi + verified := subkeys_verified v + spSafe := Code.all_of_forall (fun _ => rfl) _ + features := v.features }, + { Spec.Cmac.aesUpdateApi with + name := Spec.Cmac.aesUpdateApi.name ++ v.suffix + target := AArch64.target + doc := Spec.Cmac.aesUpdateApi.doc (notes := [ctrNote v]) + code := Impl.CmacAes.AArch64.update v.callee + contract := Spec.Cmac.aesUpdateContract AArch64.abi + verified := update_verified v + spSafe := Code.all_of_forall (fun _ => rfl) _ + features := v.features }, + { Spec.Cmac.aesFinalizeApi with + name := Spec.Cmac.aesFinalizeApi.name ++ v.suffix + target := AArch64.target + doc := Spec.Cmac.aesFinalizeApi.doc (notes := [ctrNote v]) + code := Impl.CmacAes.AArch64.finalize v.callee + contract := Spec.Cmac.aesFinalizeContract AArch64.abi + verified := finalize_verified v + spSafe := Code.all_of_forall (fun _ => rfl) _ + features := v.features }] + +end VG.Generic.AesCtr32.AArch64.CmacAes diff --git a/lean/VerifiedGarbage/Impl/Aes/AArch64/Callee.lean b/lean/VerifiedGarbage/Impl/Aes/AArch64/Callee.lean new file mode 100644 index 000000000..cbf32e085 --- /dev/null +++ b/lean/VerifiedGarbage/Impl/Aes/AArch64/Callee.lean @@ -0,0 +1,23 @@ +import VerifiedGarbage.Impl.Aes.AArch64.Ctr32 +import VerifiedGarbage.Impl.Aes.AArch64.Aese + +/-! +# The implementations of `vg_aes_ctr32` on AArch64 + +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/AArch64/`). +-/ + +namespace VG.Impl.Aes.AArch64 + +open VG.AArch64 + +/-- 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.aese : Ctr32 := ⟨"vg_aes_ctr32_aes", Aese.ctr32⟩ + +end VG.Impl.Aes.AArch64 diff --git a/lean/VerifiedGarbage/Impl/CmacAes/AArch64.lean b/lean/VerifiedGarbage/Impl/CmacAes/AArch64.lean new file mode 100644 index 000000000..d17591308 --- /dev/null +++ b/lean/VerifiedGarbage/Impl/CmacAes/AArch64.lean @@ -0,0 +1,172 @@ +import VerifiedGarbage.Impl.Aes.AArch64.Callee + +/-! +# AES-CMAC: AArch64 implementation + +`vg_cmac_aes_subkeys(schedule = x0, rounds = x1, subkeys = x2, scratch = x3)`, +`vg_cmac_aes_update(schedule = x0, rounds = x1, state = x2, data = x3, n = x4, scratch = x5)` +and `vg_cmac_aes_finalize(key = x0, rounds = x1, state = x2, last = x3, last_len = x4, scratch = x5)` +(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_aes` calls `vg_aes_ctr32_aes`). + +The scratch buffer (2176 bytes): `[0, 2048)` is the working space of +`vg_aes_ctr32`, `[2048, 2064)` the counter block, and `[2064, 2120)` our +caller's callee-saved registers and our return address `x30`. A call (`bl`) +stores nothing in memory, so no stack is used. + +* `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 `x9:x10`, shifted left by one bit, and + XORed with `0x87` masked by the bit shifted out. `x19` holds `subkeys` + and `x20` the scratch buffer across the call. +* `update` keeps its arguments in `x19` (schedule), `x20` (rounds), `x21` + (state), `x22` (data), `x23` (blocks left) and `x24` (scratch) across the + calls; each block, the counter block is `C ⊕ Mᵢ` and the state, zeroed, + receives `CIPH_K(C ⊕ Mᵢ)`. +* `finalize` forms `Mₙ` in the counter block: `Mₙ* ⊕ K1` for a complete + block, else `Mₙ*` copied a byte at a time onto zeros, `0x80` after it, and + XORed with `K2`. It XORs in the chaining value and calls `vg_aes_ctr32` + last, keeping only the scratch buffer (in `x19`) across the call. + +The model has no flags or register-offset addressing: the branches are +`cbz`/`cbnz`, on `n`, on `last_len - 16` and on `last_len`, and the bytes are +copied through advancing pointers. Only the pointers, `rounds`, `n` and +`last_len` can affect timing: the doubling is masked. +-/ + +namespace VG.Impl.CmacAes.AArch64 + +open VG.AArch64 +open VG.Impl.Aes.AArch64 (Ctr32) + +/-- `mov d, n`. -/ +def mov (d n : Reg) : Instr := .addImm .x d n 0 + +/-- 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 := + [.addImm .x .x2 scr cOff, mov .x3 out, .movz .x .x4 1 0, mov .x5 scr] + +/-! ## `vg_cmac_aes_subkeys` -/ + +/-- Saves `x19`, `x20` and `x30`, keeps `subkeys` in `x19` and the scratch +buffer in `x20`, zeroes the counter block and the first block of `subkeys`, +and sets up the arguments of `vg_aes_ctr32`. -/ +def subkeysPre : List Instr := + [.str .x .x19 .x3 2064, .str .x .x20 .x3 2072, .str .x .x30 .x3 2080, mov .x19 .x2, mov .x20 .x3, + .movz .x .x9 0 0, .str .x .x9 .x3 cOff, .str .x .x9 .x3 (cOff + 8), .str .x .x9 .x2 0, + .str .x .x9 .x2 8] ++ + ctrArgs .x20 .x19 + +/-- The block at `x19 + src`, doubled (`VG.Spec.Cmac.dbl 16`), to `x19 + dst`. -/ +def dbl (src dst : Nat) : List Instr := + [.ldr .x .x9 .x19 src, .ldr .x .x10 .x19 (src + 8), .rev .x9 .x9, .rev .x10 .x10, + .lsr .x .x11 .x9 63, .movz .x .x12 0 0, .sub .x .x11 .x12 .x11, .movz .x .x12 0x87 0, + .logic .and .x .x11 .x11 .x12, + .lsr .x .x12 .x10 63, .lsl .x .x9 .x9 1, .logic .orr .x .x9 .x9 .x12, + .lsl .x .x10 .x10 1, .logic .eor .x .x10 .x10 .x11, + .rev .x9 .x9, .rev .x10 .x10, .str .x .x9 .x19 dst, .str .x .x10 .x19 (dst + 8)] + +/-- `K1` over `L`, `K2` after it, and the saved registers restored. -/ +def subkeysPost : List Instr := + dbl 0 0 ++ dbl 0 16 ++ [.ldr .x .x30 .x20 2080, .ldr .x .x19 .x20 2064, .ldr .x .x20 .x20 2072] + +def subkeys (c : Ctr32) : Prog isa := + .seq (.block subkeysPre) (.seq (.call c.name c.code) (.block subkeysPost)) + +/-! ## `vg_cmac_aes_update` -/ + +/-- The registers saved in the scratch buffer, and where (`x24`, the base of +the restore, last). -/ +def saved : List (Reg × Nat) := + [(.x19, 2064), (.x20, 2072), (.x21, 2080), (.x22, 2088), (.x23, 2096), (.x30, 2104), (.x24, 2112)] + +def save : List Instr := saved.map fun (r, d) => .str .x r .x5 d + +/-- Restores the registers, with `x24` (restored last) the scratch buffer. -/ +def restore : List Instr := saved.map fun (r, d) => .ldr .x r .x24 d + +/-- The arguments to their registers. -/ +def setup : List Instr := + [mov .x19 .x0, mov .x20 .x1, mov .x21 .x2, mov .x22 .x3, mov .x23 .x4, mov .x24 .x5] + +/-- The counter block `C ⊕ Mᵢ` (the state at `x21`, the block at `x22`), and +the state zeroed. -/ +def chainIn : List Instr := + [.ldr .x .x9 .x21 0, .ldr .x .x10 .x22 0, .logic .eor .x .x9 .x9 .x10, .str .x .x9 .x24 cOff, + .ldr .x .x9 .x21 8, .ldr .x .x10 .x22 8, .logic .eor .x .x9 .x9 .x10, .str .x .x9 .x24 (cOff + 8), + .movz .x .x9 0 0, .str .x .x9 .x21 0, .str .x .x9 .x21 8] + +/-- The arguments of `vg_aes_ctr32` for the block. -/ +def updArgs : List Instr := [mov .x0 .x19, mov .x1 .x20] ++ ctrArgs .x24 .x21 + +/-- On to the next block. -/ +def advance : List Instr := [.addImm .x .x22 .x22 16, .subImm .x .x23 .x23 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 (.zero .x .x23) (.block []) (.loop (body c) (.nonzero .x .x23))) (.block restore)) + +/-! ## `vg_cmac_aes_finalize` -/ + +/-- `Mₙ = Mₙ* ⊕ K1` (`K1` at `x0 + 240`), for a complete last block. -/ +def full : List Instr := + [.ldr .x .x9 .x3 0, .ldr .x .x10 .x0 240, .logic .eor .x .x9 .x9 .x10, .str .x .x9 .x5 cOff, + .ldr .x .x9 .x3 8, .ldr .x .x10 .x0 248, .logic .eor .x .x9 .x9 .x10, .str .x .x9 .x5 (cOff + 8)] + +/-- The counter block zeroed, with `x6` pointing at it, `x7` at the last +bytes and `x8` counting them. -/ +def zero : List Instr := + [.movz .x .x9 0 0, .str .x .x9 .x5 cOff, .str .x .x9 .x5 (cOff + 8), .addImm .x .x6 .x5 cOff, + mov .x7 .x3, mov .x8 .x4] + +/-- The `x8` (nonzero) bytes at `x7` copied to `x6`, advancing both. -/ +def copy : Prog isa := + .loop (.block [.ldrb .x9 .x7 0, .strb .x9 .x6 0, .addImm .x .x7 .x7 1, .addImm .x .x6 .x6 1, + .subImm .x .x8 .x8 1]) (.nonzero .x .x8) + +/-- `0x80` after the bytes (at `x6`), and the block XORed with `K2` (at +`x0 + 256`). -/ +def padK2 : List Instr := + [.movz .x .x9 0x80 0, .strb .x9 .x6 0, + .ldr .x .x9 .x5 cOff, .ldr .x .x10 .x0 256, .logic .eor .x .x9 .x9 .x10, .str .x .x9 .x5 cOff, + .ldr .x .x9 .x5 (cOff + 8), .ldr .x .x10 .x0 264, .logic .eor .x .x9 .x9 .x10, + .str .x .x9 .x5 (cOff + 8)] + +/-- `Mₙ = K2 ⊕ (Mₙ* ‖ 10ʲ)`, for a partial last block (`last_len < 16`). -/ +def partialBlock : Prog isa := + .seq (.block zero) (.seq (.ite (.zero .x .x4) (.block []) copy) (.block padK2)) + +/-- The counter block `C ⊕ Mₙ` (the state at `x2`), the state zeroed, `x19` +and `x30` saved and the scratch buffer kept in `x19`, and the arguments of +`vg_aes_ctr32` but the schedule (`x0`), the rounds (`x1`) and the working +space (`x5`), which are ours. -/ +def finArgs : List Instr := + [.ldr .x .x9 .x5 cOff, .ldr .x .x10 .x2 0, .logic .eor .x .x9 .x9 .x10, .str .x .x9 .x5 cOff, + .ldr .x .x9 .x5 (cOff + 8), .ldr .x .x10 .x2 8, .logic .eor .x .x9 .x9 .x10, + .str .x .x9 .x5 (cOff + 8), + .movz .x .x9 0 0, .str .x .x9 .x2 0, .str .x .x9 .x2 8, + .str .x .x19 .x5 2064, .str .x .x30 .x5 2072, mov .x19 .x5, + mov .x3 .x2, .addImm .x .x2 .x5 cOff, .movz .x .x4 1 0] + +/-- Everything before the call. -/ +def finPre : Prog isa := + .seq (.block [.subImm .x .x9 .x4 16]) + (.seq (.ite (.zero .x .x9) (.block full) partialBlock) (.block finArgs)) + +def finalize (c : Ctr32) : Prog isa := + .seq finPre (.seq (.call c.name c.code) (.block [.ldr .x .x30 .x19 2072, .ldr .x .x19 .x19 2064])) + +end VG.Impl.CmacAes.AArch64 diff --git a/lean/VerifiedGarbage/Proof/Aes/AArch64/Variant.lean b/lean/VerifiedGarbage/Proof/Aes/AArch64/Variant.lean new file mode 100644 index 000000000..d8e86fd3d --- /dev/null +++ b/lean/VerifiedGarbage/Proof/Aes/AArch64/Variant.lean @@ -0,0 +1,64 @@ +import VerifiedGarbage.Proof.Aes.AArch64.Ctr32 +import VerifiedGarbage.Proof.Aes.AArch64.Aese.Ctr32 +import VerifiedGarbage.Impl.Aes.AArch64.Callee +import VerifiedGarbage.Proof.Framework.AArch64.Call + +/-! +# Implementations of `vg_aes_ctr32` on AArch64 + +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 AArch64 (`Variants/AesCtr32/AArch64/`), and each +caller (in `Generic/AesCtr32/AArch64/`) is emitted once for each of them (see +`TCB/Emit.lean`). Every implementation is proven against the same contract, +`Proof.Aes.ctr32AArch64`, and has no frames. +-/ + +namespace VG.Proof.Aes.AArch64 + +open VG.AArch64 + +/-- An implementation of `vg_aes_ctr32` on AArch64. -/ +structure Ctr32Impl where + /-- Its symbol and code. -/ + callee : Impl.Aes.AArch64.Ctr32 + /-- It has no frames (`WP.call`). -/ + noFrames : callee.code.noFrames = true + ok : ∀ s, Proof.Aes.ctr32AArch64.pre s → + ∃ t s', Exec isa callee.code s t s' ∧ abiPreserved s s' ∧ Proof.Aes.ctr32AArch64.post s s' + ct : ConstantTime isa Proof.Aes.ctr32AArch64.pre Proof.Aes.ctr32AArch64.pub callee.code + /-- It never writes v8–v15. -/ + keepsV : callee.code.allInstrs keepsV = true + /-- What the names of its callers' instances end with (e.g. `_aes`; + nothing for the baseline implementation). -/ + suffix : String + /-- The CPU features its code requires, which its callers require too. -/ + features : List String + +namespace Ctr32Impl + +/-- The bitsliced implementation, `vg_aes_ctr32`, in the baseline ISA. -/ +def scalar : Ctr32Impl where + callee := .scalar + noFrames := by decide +kernel + ok := ctr32_correct + ct := ctr32_ct + keepsV := by decide +kernel + suffix := "" + features := [] + +/-- The implementation with the Cryptographic Extension, `vg_aes_ctr32_aes`. -/ +def aese : Ctr32Impl where + callee := .aese + noFrames := by decide +kernel + ok := Aese.ctr32_correct + ct := Aese.ctr32_ct + keepsV := by decide +kernel + suffix := "_aes" + features := ["aes"] + +end Ctr32Impl + +end VG.Proof.Aes.AArch64 diff --git a/lean/VerifiedGarbage/Proof/Cmac/Block.lean b/lean/VerifiedGarbage/Proof/Cmac/Block.lean new file mode 100644 index 000000000..f95ecced8 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/Cmac/Block.lean @@ -0,0 +1,114 @@ +import VerifiedGarbage.Proof.Cmac.Frame +import VerifiedGarbage.Proof.Framework.WriteBytes + +/-! +# CMAC: forming the last block in memory + +Untrusted: everything here is checked by Lean. + +What `finalize`'s stores leave, on any target: the XOR of two blocks stored a +word at a time (`xor2Mem`), a zeroed block (`zero2`), and a partial last +block copied onto zeros and padded with `0x80` (`padded_bytes`). +-/ + +namespace VG.Proof.Cmac + +open VG + +/-! ## 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, 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 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) + +/-! ## Zeroing, copying and padding -/ + +theorem zeros_8_8 : Spec.Cmac.zeros 8 ++ Spec.Cmac.zeros 8 = Spec.Cmac.zeros 16 := by decide + +/-- The memory after zeroing the block at `c`. -/ +def zero2 (m : Mem) (c : Addr) : Mem := + (m.writeW c (0 : BitVec 64)).writeW (c + BitVec.ofNat 64 8) (0 : BitVec 64) + +theorem zero2_bytes (m : Mem) (c : Addr) : Spec.Aes.bytesAt (zero2 m c) c 16 = Spec.Cmac.zeros 16 := by + rw [zero2, bytesAt_store2, le8_zero, zeros_8_8] + +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] + +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] + +/-- The two subkeys, from the 32 bytes after the key schedule. -/ +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 [bytesAt_length, h1]) + +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 ext16 (by simp [Spec.Aes.bytesAt]) (by simp [Spec.Cmac.zeros]; omega) fun k hk => ?_ + rw [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 [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.Cmac diff --git a/lean/VerifiedGarbage/Proof/Cmac/Frame.lean b/lean/VerifiedGarbage/Proof/Cmac/Frame.lean new file mode 100644 index 000000000..ba4bcd466 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/Cmac/Frame.lean @@ -0,0 +1,74 @@ +import VerifiedGarbage.Proof.Cmac.Mem + +/-! +# CMAC: blocks in memory under frames + +Untrusted: everything here is checked by Lean. + +What the implementations' stores of 64-bit words leave in memory, on any +target: the bytes outside a frame are unchanged (`bytesAt_frame`), and the +memory after forming a counter block `C = P ⊕ Q` and zeroing `P` +(`chainMem`), as each block of `update` does. +-/ + +namespace VG.Proof.Cmac + +open VG + +/-- 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) + +theorem bytesAt_frame16 {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 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 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) + +/-- The memory after forming a counter block: the block at `c` is the block +at `p` XORed with the block at `q`, and the block at `p` is zeroed. -/ +def 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 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, bytesAt_store2, le8_zero]; rfl + +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_frame16 (frame_store2 P 0 0) (by simpa using hcp), 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 xor_words m P Q + +end VG.Proof.Cmac diff --git a/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Call.lean b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Call.lean new file mode 100644 index 000000000..32fe75841 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Call.lean @@ -0,0 +1,107 @@ +import VerifiedGarbage.Proof.Aes.AArch64.Variant +import VerifiedGarbage.Proof.Cmac.Frame +import VerifiedGarbage.Proof.Framework.AArch64.RelCT + +/-! +# AES-CMAC on AArch64: 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` and `S` change in memory. +-/ + +namespace VG.Proof.CmacAes.AArch64 + +open VG VG.AArch64 +open VG.Proof.Aes.AArch64 (Ctr32Impl) + +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 + x0 : s.gpr .x0 = W + x1 : s.gpr .x1 = BitVec.ofNat 64 R + x2 : s.gpr .x2 = C + x3 : s.gpr .x3 = D + x4 : s.gpr .x4 = 1 + x5 : s.gpr .x5 = 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⟩ + 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 + sp : s'.sp = s.sp + saved : ∀ r ∈ preserved, r ≠ .x30 → s'.gpr r = s.gpr r + frame : Frame [⟨C, 16⟩, ⟨D, 16⟩, ⟨S, 2048⟩] 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) + +theorem callEntry_x0 (s : State) : s.callEntry.gpr .x0 = s.gpr .x0 := s.callEntry_gpr (by decide) +theorem callEntry_x1 (s : State) : s.callEntry.gpr .x1 = s.gpr .x1 := s.callEntry_gpr (by decide) +theorem callEntry_x2 (s : State) : s.callEntry.gpr .x2 = s.gpr .x2 := s.callEntry_gpr (by decide) +theorem callEntry_x3 (s : State) : s.callEntry.gpr .x3 = s.gpr .x3 := s.callEntry_gpr (by decide) +theorem callEntry_x4 (s : State) : s.callEntry.gpr .x4 = s.gpr .x4 := s.callEntry_gpr (by decide) +theorem callEntry_x5 (s : State) : s.callEntry.gpr .x5 = s.gpr .x5 := s.callEntry_gpr (by decide) + +/-- `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.ctr32AArch64.pre + (s.callEntry.withRegions [⟨W, 240⟩] [⟨C, 16⟩, ⟨D, 16⟩, ⟨S, 2048⟩]) := by + have hR := toNat_rounds h.rounds + simp only [Proof.Aes.ctr32AArch64, State.withRegions_gpr, State.withRegions_rd, + State.withRegions_wr, callEntry_x0, callEntry_x1, callEntry_x2, callEntry_x3, callEntry_x4, + callEntry_x5, h.x0, h.x1, h.x2, h.x3, h.x4, h.x5, 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, 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.ctr32AArch64) v.ok (rd := [⟨W, 240⟩]) + (wr := [⟨C, 16⟩, ⟨D, 16⟩, ⟨S, 2048⟩]) h.ctr_pre h.reads h.writes ?_ v.noFrames + intro s' hrd hwr hsp hf hsaved _ hpost + refine ⟨hrd, hwr, hsp, hsaved, hf, ?_⟩ + obtain ⟨hdata, -⟩ := hpost + simp only [State.withRegions_gpr, State.withRegions_mem, State.callEntry_mem, callEntry_x0, + callEntry_x1, callEntry_x2, callEntry_x3, callEntry_x4, h.x0, h.x1, h.x2, h.x3, h.x4, hR, + one_toNat] at hdata + 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.mem D = 0 := by rw [Spec.Gcm.blockAt, h.zero, ofBytes_zeros] + rw [one, one, bD, Proof.Cmac.ctr32_one, List.cons.injEq] at hdata + rw [Proof.Cmac.bytesAt_blockAt, hdata.1, Spec.Gcm.blockAt, + 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) {W C D S : Addr} {R : Nat} {P : State → State → Prop} + (h : ∀ s₁ s₂, P s₁ s₂ → CallPre s₁ W C D S R ∧ CallPre s₂ W C D S R ∧ s₁.sp = s₂.sp) : + RelCT isa P (.call v.callee.name v.callee.code) fun _ _ => True := by + refine RelCT.call v.ok v.ct [⟨W, 240⟩] [⟨C, 16⟩, ⟨D, 16⟩, ⟨S, 2048⟩] fun s₁ s₂ hp => ?_ + obtain ⟨h₁, h₂, hsp⟩ := h s₁ s₂ hp + refine ⟨h₁.ctr_pre, h₂.ctr_pre, ?_, h₁.reads, h₁.writes, h₂.reads, h₂.writes⟩ + simp only [Proof.Aes.ctr32AArch64, State.withRegions_gpr, State.withRegions_sp, State.callEntry_sp, + callEntry_x0, callEntry_x1, callEntry_x2, callEntry_x3, callEntry_x4, callEntry_x5, + h₁.x0, h₁.x1, h₁.x2, h₁.x3, h₁.x4, h₁.x5, h₂.x0, h₂.x1, h₂.x2, h₂.x3, h₂.x4, h₂.x5, hsp] + exact ⟨trivial, trivial, trivial, trivial, trivial, trivial, trivial⟩ + +end VG.Proof.CmacAes.AArch64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Contract.lean b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Contract.lean new file mode 100644 index 000000000..6f86c78ee --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Contract.lean @@ -0,0 +1,85 @@ +import VerifiedGarbage.Proof.CmacAes.AArch64.Call +import VerifiedGarbage.Impl.CmacAes.AArch64 + +/-! +# AES-CMAC on AArch64: 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`). A call (`bl`) stores nothing in memory, so no stack is +used. +-/ + +namespace VG.Proof.CmacAes.AArch64 + +open VG VG.AArch64 + +/-- `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 = x0, rounds = x1, state = x2, data = x3, n = x4, scratch = x5)`. -/ +def updateAArch64 : Contract isa where + pre s := + let sched : Region := ⟨s.gpr .x0, 240⟩ + let state : Region := ⟨s.gpr .x2, 16⟩ + let data : Region := ⟨s.gpr .x3, 16 * (s.gpr .x4).toNat⟩ + let scr : Region := ⟨s.gpr .x5, 2176⟩ + s.rd = [sched, data] ∧ s.wr = [state, scr] ∧ + sched.Disjoint state ∧ sched.Disjoint scr ∧ data.Disjoint state ∧ data.Disjoint scr ∧ + state.Disjoint scr ∧ + (s.gpr .x2).toNat + 16 ≤ 2 ^ 64 ∧ (s.gpr .x3).toNat + 16 * (s.gpr .x4).toNat ≤ 2 ^ 64 ∧ + (s.gpr .x5).toNat + 2176 ≤ 2 ^ 64 ∧ + ((s.gpr .x1).toNat = 10 ∨ (s.gpr .x1).toNat = 12 ∨ (s.gpr .x1).toNat = 14) + post s s' := + Spec.Aes.bytesAt s'.mem (s.gpr .x2) 16 = + Spec.Cmac.chain (ciphAt s.mem (s.gpr .x0) (s.gpr .x1).toNat) (Spec.Aes.bytesAt s.mem (s.gpr .x2) 16) + (Spec.Cmac.blocksAt s.mem (s.gpr .x3) 16 (s.gpr .x4).toNat) + pub s₁ s₂ := + s₁.gpr .x0 = s₂.gpr .x0 ∧ s₁.gpr .x1 = s₂.gpr .x1 ∧ s₁.gpr .x2 = s₂.gpr .x2 ∧ + s₁.gpr .x3 = s₂.gpr .x3 ∧ s₁.gpr .x4 = s₂.gpr .x4 ∧ s₁.gpr .x5 = s₂.gpr .x5 ∧ s₁.sp = s₂.sp + +/-- `vg_cmac_aes_subkeys(schedule = x0, rounds = x1, subkeys = x2, scratch = x3)`. -/ +def subkeysAArch64 : Contract isa where + pre s := + let sched : Region := ⟨s.gpr .x0, 240⟩ + let subk : Region := ⟨s.gpr .x2, 32⟩ + let scr : Region := ⟨s.gpr .x3, 2176⟩ + s.rd = [sched] ∧ s.wr = [subk, scr] ∧ + sched.Disjoint subk ∧ sched.Disjoint scr ∧ subk.Disjoint scr ∧ + (s.gpr .x2).toNat + 32 ≤ 2 ^ 64 ∧ (s.gpr .x3).toNat + 2176 ≤ 2 ^ 64 ∧ + ((s.gpr .x1).toNat = 10 ∨ (s.gpr .x1).toNat = 12 ∨ (s.gpr .x1).toNat = 14) + post s s' := + let ks := Spec.Cmac.subkeys (ciphAt s.mem (s.gpr .x0) (s.gpr .x1).toNat) 16 + Spec.Aes.bytesAt s'.mem (s.gpr .x2) 32 = ks.1 ++ ks.2 + pub s₁ s₂ := + s₁.gpr .x0 = s₂.gpr .x0 ∧ s₁.gpr .x1 = s₂.gpr .x1 ∧ s₁.gpr .x2 = s₂.gpr .x2 ∧ + s₁.gpr .x3 = s₂.gpr .x3 ∧ s₁.sp = s₂.sp + +/-- `vg_cmac_aes_finalize(key = x0, rounds = x1, state = x2, last = x3, last_len = x4, scratch = x5)`. -/ +def finalizeAArch64 : Contract isa where + pre s := + let key : Region := ⟨s.gpr .x0, 272⟩ + let state : Region := ⟨s.gpr .x2, 16⟩ + let last : Region := ⟨s.gpr .x3, (s.gpr .x4).toNat⟩ + let scr : Region := ⟨s.gpr .x5, 2176⟩ + s.rd = [key, last] ∧ s.wr = [state, scr] ∧ + key.Disjoint state ∧ key.Disjoint scr ∧ last.Disjoint state ∧ last.Disjoint scr ∧ + state.Disjoint scr ∧ + (s.gpr .x0).toNat + 272 ≤ 2 ^ 64 ∧ (s.gpr .x2).toNat + 16 ≤ 2 ^ 64 ∧ + (s.gpr .x3).toNat + (s.gpr .x4).toNat ≤ 2 ^ 64 ∧ (s.gpr .x5).toNat + 2176 ≤ 2 ^ 64 ∧ + ((s.gpr .x1).toNat = 10 ∨ (s.gpr .x1).toNat = 12 ∨ (s.gpr .x1).toNat = 14) ∧ + (s.gpr .x4).toNat ≤ 16 + post s s' := + let ciph := ciphAt s.mem (s.gpr .x0) (s.gpr .x1).toNat + let ks := Spec.Cmac.subkeys ciph 16 + Spec.Aes.bytesAt s.mem (s.gpr .x0 + 240) 32 = ks.1 ++ ks.2 → + ∀ msg : List Byte, msg.length % 16 = 0 → (msg = [] ∨ 0 < (s.gpr .x4).toNat) → + Spec.Aes.bytesAt s.mem (s.gpr .x2) 16 = Spec.Cmac.chain ciph (Spec.Cmac.zeros 16) (Spec.Cmac.blocks 16 msg) → + Spec.Aes.bytesAt s'.mem (s.gpr .x2) 16 = + Spec.Cmac.macFull ciph 16 (msg ++ Spec.Aes.bytesAt s.mem (s.gpr .x3) (s.gpr .x4).toNat) + pub s₁ s₂ := + s₁.gpr .x0 = s₂.gpr .x0 ∧ s₁.gpr .x1 = s₂.gpr .x1 ∧ s₁.gpr .x2 = s₂.gpr .x2 ∧ + s₁.gpr .x3 = s₂.gpr .x3 ∧ s₁.gpr .x4 = s₂.gpr .x4 ∧ s₁.gpr .x5 = s₂.gpr .x5 ∧ s₁.sp = s₂.sp + +end VG.Proof.CmacAes.AArch64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Dbl.lean b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Dbl.lean new file mode 100644 index 000000000..ab7104d98 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Dbl.lean @@ -0,0 +1,98 @@ +import VerifiedGarbage.Proof.Cmac.Dbl +import VerifiedGarbage.Proof.Cmac.Mem +import VerifiedGarbage.Proof.Gcm.AArch64.Ghash + +/-! +# AES-CMAC on AArch64: doubling a block in two 64-bit words + +Untrusted: everything here is checked by Lean. `subkeys` loads a block as +two byte-reversed words (`rev`), the high and low halves of the block as a +big-endian integer (`Proof.Gcm.AArch64.blockAt_rev`), doubles the integer a +word at a time (`dbl_words`), and stores the halves byte-reversed again +(`le8_rev`). +-/ + +namespace VG.Proof.CmacAes.AArch64 + +open VG VG.AArch64 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_rev64 (x : BitVec 64) {p : Nat} (hp : p < 64) : + (rev64 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_rev (h l : BitVec 64) : + le8 (rev64 h) ++ le8 (rev64 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_rev64 _ (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_rev64 _ (by omega)] + simp only [hj, decide_true, Bool.true_and, show 8 * (15 - k) + j < 64 by omega, ite_true] + congr 1; omega + +theorem mask_eq (hi : BitVec 64) : + ((0 : BitVec 64) - (hi >>> 63)) &&& 0x87 = 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 + +/-- The words `subkeys` stores, from the halves `hi` and `lo` it loads. -/ +def dblHi (hi lo : BitVec 64) : BitVec 64 := (hi <<< 1) ||| (lo >>> 63) +def dblLo (hi lo : BitVec 64) : BitVec 64 := (lo <<< 1) ^^^ (((0 : BitVec 64) - (hi >>> 63)) &&& 0x87) + +/-- `subkeys`' doubling of `hi ++ lo`. -/ +theorem dbl_words (hi lo : BitVec 64) : dblHi hi lo ++ dblLo hi lo = dbl128 (hi ++ lo) := by + rw [dblHi, dblLo, mask_eq, 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.AArch64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Finalize.lean b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Finalize.lean new file mode 100644 index 000000000..570d3d384 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Finalize.lean @@ -0,0 +1,198 @@ +import VerifiedGarbage.Proof.CmacAes.AArch64.Subkeys +import VerifiedGarbage.Proof.Cmac.Block + +/-! +# AES-CMAC on AArch64: `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.AArch64 + +open VG VG.AArch64 VG.AArch64.RegUpd VG.Impl.CmacAes.AArch64 + +/-- The block XOR `full`, `padK2` and `finArgs` do: the words at `pb + pd` and +`qb + qd` XORed into `cb + cd`. -/ +def xor2 (pb qb cb : Reg) (pd qd cd : Nat) : List Instr := + [.ldr .x .x9 pb pd, .ldr .x .x10 qb qd, .logic .eor .x .x9 .x9 .x10, .str .x .x9 cb cd, + .ldr .x .x9 pb (pd + 8), .ldr .x .x10 qb (qd + 8), .logic .eor .x .x9 .x9 .x10, .str .x .x9 cb (cd + 8)] + +theorem xor2_ok (s : State) (pb qb cb : Reg) (pd qd cd : Nat) {P Q C : Addr} + (hpd : pd % 8 = 0 ∧ pd + 8 < 32768) (hqd : qd % 8 = 0 ∧ qd + 8 < 32768) + (hcd : cd % 8 = 0 ∧ cd + 8 < 32768) + (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) + (hr : pb ≠ .x9 ∧ pb ≠ .x10 ∧ qb ≠ .x9 ∧ qb ≠ .x10 ∧ cb ≠ .x9 ∧ cb ≠ .x10) + (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 (xor2 pb qb cb pd qd cd) s = some s' ∧ + s'.mem = Proof.Cmac.xor2Mem s.mem C P Q ∧ (∀ r, r ≠ .x9 → r ≠ .x10 → s'.gpr r = s.gpr r) ∧ + s'.sp = s.sp ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + obtain ⟨h₁, h₂, h₃, h₄, h₅, h₆⟩ := hr + refine ⟨_, by + simp (config := {decide := true}) only [xor2, runBlock_cons, runStep_some, runBlock_nil, exec, addr, + State.load, State.store, Size.bytes, Size.bits, State.read, gpr_write, mem_write, rd_write, + wr_write, ite_true, ite_false, Option.bind_some, Option.map_some, BitVec.setWidth_eq, + h₁, h₂, h₃, h₄, h₅, h₆, hpd.1, hqd.1, hcd.1, Nat.add_mod_right, + show pd < 32768 by omega, show qd < 32768 by omega, show cd < 32768 by omega, hpd.2, hqd.2, hcd.2, + hp, hp8, hq, hq8, hc, hc8, rp, rp8, rq, rq8, wc, wc8, and_self] + rfl, ?_⟩ + refine ⟨?_, fun r h₁ h₂ => by simp [gpr_write, h₁, h₂], rfl, rfl, rfl⟩ + simp only [Proof.Cmac.xor2Mem, Mem.writeW, Mem.readW, BitVec.setWidth_eq] + +theorem full_eq : full = xor2 .x3 .x0 .x5 0 240 2048 := rfl +theorem padK2_eq : + padK2 = ([.movz .x .x9 0x80 0, .strb .x9 .x6 0] : List Instr) ++ xor2 .x5 .x0 .x5 2048 256 2048 := rfl +theorem finArgs_eq : finArgs = xor2 .x5 .x2 .x5 2048 0 2048 ++ + ([.movz .x .x9 0 0, .str .x .x9 .x2 0, .str .x .x9 .x2 8, .str .x .x19 .x5 2064, .str .x .x30 .x5 2072, + mov .x19 .x5, mov .x3 .x2, .addImm .x .x2 .x5 cOff, .movz .x .x4 1 0] : List Instr) := rfl + +theorem sub16_ok (s : State) {L : Nat} (h4 : s.gpr .x4 = BitVec.ofNat 64 L) (hL : L ≤ 16) : + ∃ s', runBlock isa [.subImm .x .x9 .x4 16] s = some s' ∧ + isa.eval (.zero .x .x9) s' = some (decide (L = 16)) ∧ + (∀ r, r ≠ .x9 → s'.gpr r = s.gpr r) ∧ s'.sp = s.sp ∧ s'.mem = s.mem ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by rw [runBlock_cons, exec_subImm_x (by decide), runStep_some, runBlock_nil], ?_⟩ + refine ⟨?_, fun r h => by simp [gpr_write, h], rfl, rfl, rfl, rfl⟩ + show some (_ == 0) = _ + simp only [State.read, gpr_write_self, BitVec.setWidth_eq, h4] + rw [Offset.ofNat_sub_ofNat_beq (by omega) (by decide)] + +theorem zero_ok (s : State) {C : Addr} (hc : s.gpr .x5 + BitVec.ofNat 64 2048 = C) + (hc8 : s.gpr .x5 + BitVec.ofNat 64 2056 = C + BitVec.ofNat 64 8) + (wc : InRegions s.wr C 8) (wc8 : InRegions s.wr (C + BitVec.ofNat 64 8) 8) : + ∃ s', runBlock isa zero s = some s' ∧ s'.mem = Proof.Cmac.zero2 s.mem C ∧ s'.gpr .x6 = C ∧ + s'.gpr .x7 = s.gpr .x3 ∧ s'.gpr .x8 = s.gpr .x4 ∧ + (∀ r, r ≠ .x6 → r ≠ .x7 → r ≠ .x8 → r ≠ .x9 → s'.gpr r = s.gpr r) ∧ + s'.sp = s.sp ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + simp (config := {decide := true}) only [zero, cOff, mov, runBlock_cons, runStep_some, runBlock_nil, + exec, addr, State.store, Size.bytes, Size.bits, State.read, gpr_write, mem_write, wr_write, + ite_true, ite_false, Option.bind_some, BitVec.setWidth_eq, hc, hc8, wc, wc8] + rfl, ?_⟩ + refine ⟨?_, by simp [gpr_write, ← hc], by simp [gpr_write], by simp [gpr_write], + fun r h₁ h₂ h₃ h₄ => by simp [gpr_write, h₁, h₂, h₃, h₄], rfl, rfl, rfl⟩ + simp only [mem_write, Proof.Cmac.zero2, Mem.writeW, BitVec.setWidth_eq, mz0] + +/-- The copy loop's body. -/ +abbrev copyBody : List Instr := + [.ldrb .x9 .x7 0, .strb .x9 .x6 0, .addImm .x .x7 .x7 1, .addImm .x .x6 .x6 1, .subImm .x .x8 .x8 1] + +theorem byte_rt (b : BitVec (8 * 1)) : + BitVec.setWidth 8 (BitVec.setWidth 32 (BitVec.setWidth 64 (BitVec.setWidth 32 b))) = b := by + apply BitVec.eq_of_toNat_eq + have := b.isLt + simp only [BitVec.toNat_setWidth] + omega + +theorem read_one (m : Mem) (a : Addr) : m.read a 1 = m a := by + have := Mem.extractLsb'_read m a (n := 1) (j := 0) (by decide) + rw [show 8 * 0 = 0 from rfl, BitVec.extractLsb'_eq_self, show BitVec.ofNat 64 0 = 0#64 from rfl, + BitVec.add_zero] at this + exact this + +theorem copyStep_ok (s : State) {A B : Addr} (ha : s.gpr .x7 + BitVec.ofNat 64 0 = A) + (hb : s.gpr .x6 + BitVec.ofNat 64 0 = B) + (r : InRegions (s.rd ++ s.wr) A 1) (w : InRegions s.wr B 1) : + ∃ s', runBlock isa copyBody s = some s' ∧ s'.mem = s.mem.writeW B (s.mem A) ∧ + s'.gpr .x7 = s.gpr .x7 + 1 ∧ s'.gpr .x6 = s.gpr .x6 + 1 ∧ s'.gpr .x8 = s.gpr .x8 - 1 ∧ + (∀ r, r ≠ .x6 → r ≠ .x7 → r ≠ .x8 → r ≠ .x9 → s'.gpr r = s.gpr r) ∧ + s'.sp = s.sp ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + simp (config := {decide := true}) only [copyBody, runBlock_cons, runStep_some, runBlock_nil, + exec, addr, State.load, State.store, Size.bits, State.read, gpr_write, mem_write, + rd_write, wr_write, ite_true, ite_false, Option.bind_some, Option.map_some, BitVec.setWidth_eq, + ha, hb, r, w] + rfl, ?_⟩ + refine ⟨?_, by simp [gpr_write], by simp [gpr_write], by simp [gpr_write], + fun r h₁ h₂ h₃ h₄ => by simp [gpr_write, h₁, h₂, h₃, h₄], rfl, rfl, rfl⟩ + simp only [mem_write, Mem.writeW, byte_rt, read_one, Nat.reduceDiv, Nat.reduceMul, BitVec.setWidth_eq] + +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 copy_ok (s : State) {P C : Addr} {L : Nat} (hL₀ : 0 < L) (hL : L < 16) + (h7 : s.gpr .x7 = P) (h6 : s.gpr .x6 = C) (h8 : s.gpr .x8 = 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) ∧ + s'.gpr .x6 = C + BitVec.ofNat 64 L ∧ + (∀ r, r ≠ .x6 → r ≠ .x7 → r ≠ .x8 → r ≠ .x9 → s'.gpr r = s.gpr r) ∧ + s'.sp = s.sp ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine WP.loop (M := isa) (body := .block copyBody) (c := .nonzero .x .x8) + (fun (n : Nat) (t : State) => ∃ i, n = L - i ∧ i < L ∧ t.gpr .x7 = P + BitVec.ofNat 64 i ∧ + t.gpr .x6 = C + BitVec.ofNat 64 i ∧ t.gpr .x8 = BitVec.ofNat 64 (L - i) ∧ + t.mem = writeBytes s.mem C (Spec.Aes.bytesAt s.mem P i) ∧ + (∀ r, r ≠ .x6 → r ≠ .x7 → r ≠ .x8 → r ≠ .x9 → t.gpr r = s.gpr r) ∧ + t.sp = s.sp ∧ t.rd = s.rd ∧ t.wr = s.wr) ?_ (L - 0) _ + ⟨0, rfl, hL₀, by rw [h7]; simp, by rw [h6]; simp, by rw [h8, Nat.sub_zero], by simp [Spec.Aes.bytesAt, writeBytes_nil], + fun _ _ _ _ _ => rfl, rfl, rfl, rfl⟩ + rintro n t ⟨i, rfl, hi, x7, x6, x8, mem, g, sp, rd, wr⟩ + obtain ⟨t', run', mem', x7', x6', x8', g', sp', rd', wr'⟩ := copyStep_ok t + (A := P + BitVec.ofNat 64 i) (B := C + BitVec.ofNat 64 i) (by rw [x7, BitVec.add_zero]) + (by rw [x6, BitVec.add_zero]) (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, Proof.Cmac.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 hb : L < 2 ^ 64 := by omega + have x8'' : t'.gpr .x8 = BitVec.ofNat 64 (L - (i + 1)) := by + rw [x8', x8, show (1 : BitVec 64) = BitVec.ofNat 64 1 from rfl, Offset.ofNat_sub_ofNat (by omega)]; rfl + have ev : isa.eval (.nonzero .x .x8) t' = some !decide (L - (i + 1) = 0) := by + show some (t'.read .x .x8 != 0) = _ + rw [State.read, x8'', BitVec.setWidth_eq, ofNat_ne_zero (by omega)] + have gg : ∀ r, r ≠ .x6 → r ≠ .x7 → r ≠ .x8 → r ≠ .x9 → t'.gpr r = s.gpr r := fun r h₁ h₂ h₃ h₄ => by + rw [g' r h₁ h₂ h₃ h₄, g r h₁ h₂ h₃ h₄] + by_cases he : i + 1 = L + · left + refine ⟨by rw [ev]; simp [he], by rw [hmem, he], by rw [x6', x6, BitVec.add_assoc, succ_ofNat, he], gg, + by rw [sp', sp], by rw [rd', rd], by rw [wr', wr]⟩ + · right + refine ⟨by rw [ev]; simp; omega, L - (i + 1), by omega, i + 1, rfl, by omega, + by rw [x7', x7, BitVec.add_assoc, succ_ofNat], by rw [x6', x6, BitVec.add_assoc, succ_ofNat], x8'', hmem, gg, + by rw [sp', sp], by rw [rd', rd], by rw [wr', wr]⟩ + +theorem pad_ok (s : State) {B : Addr} (hb : s.gpr .x6 + BitVec.ofNat 64 0 = B) (w : InRegions s.wr B 1) : + ∃ s', runBlock isa [.movz .x .x9 0x80 0, .strb .x9 .x6 0] s = some s' ∧ + s'.mem = s.mem.writeW B (0x80 : Byte) ∧ (∀ r, r ≠ .x9 → s'.gpr r = s.gpr r) ∧ + s'.sp = s.sp ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + simp (config := {decide := true}) only [runBlock_cons, runStep_some, runBlock_nil, exec, addr, + State.store, Size.bits, State.read, gpr_write, wr_write, ite_true, ite_false, Option.bind_some, + hb, w] + rfl, ?_⟩ + refine ⟨?_, fun r h => by simp [gpr_write, h], rfl, rfl, rfl⟩ + simp only [mem_write, Mem.writeW, Nat.reduceDiv, Nat.reduceMul] + rfl + +theorem args_ok (s : State) {D S : Addr} (hd : s.gpr .x2 = D) (hs : s.gpr .x5 = S) + (wd : InRegions s.wr (D + BitVec.ofNat 64 0) 8) (wd8 : InRegions s.wr (D + BitVec.ofNat 64 8) 8) + (ws₁ : InRegions s.wr (S + BitVec.ofNat 64 2064) 8) (ws₂ : InRegions s.wr (S + BitVec.ofNat 64 2072) 8) : + ∃ s', runBlock isa [.movz .x .x9 0 0, .str .x .x9 .x2 0, .str .x .x9 .x2 8, .str .x .x19 .x5 2064, + .str .x .x30 .x5 2072, mov .x19 .x5, mov .x3 .x2, .addImm .x .x2 .x5 cOff, .movz .x .x4 1 0] s = some s' ∧ + s'.mem = (((s.mem.writeW (D + BitVec.ofNat 64 0) (0 : BitVec 64)).writeW (D + BitVec.ofNat 64 8) + (0 : BitVec 64)).writeW (S + BitVec.ofNat 64 2064) (s.gpr .x19)).writeW (S + BitVec.ofNat 64 2072) + (s.gpr .x30) ∧ + s'.gpr .x3 = D ∧ s'.gpr .x2 = S + BitVec.ofNat 64 2048 ∧ s'.gpr .x4 = 1 ∧ s'.gpr .x19 = S ∧ + (∀ r, r ≠ .x2 → r ≠ .x3 → r ≠ .x4 → r ≠ .x9 → r ≠ .x19 → s'.gpr r = s.gpr r) ∧ + s'.sp = s.sp ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + simp (config := {decide := true}) only [cOff, mov, runBlock_cons, runStep_some, runBlock_nil, exec, + addr, State.store, Size.bytes, Size.bits, State.read, gpr_write, mem_write, wr_write, ite_true, + ite_false, Option.bind_some, BitVec.setWidth_eq, hd, hs, wd, wd8, ws₁, ws₂] + rfl, ?_⟩ + refine ⟨?_, by simp [gpr_write], by simp [gpr_write], rfl, by simp [gpr_write], + fun r h₁ h₂ h₃ h₄ h₅ => by simp [gpr_write, h₁, h₂, h₃, h₄, h₅], rfl, rfl, rfl⟩ + simp only [mem_write, Mem.writeW, BitVec.setWidth_eq, mz0] + +end VG.Proof.CmacAes.AArch64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/AArch64/FinalizeCT.lean b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/FinalizeCT.lean new file mode 100644 index 000000000..eb076c1da --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/FinalizeCT.lean @@ -0,0 +1,61 @@ +import VerifiedGarbage.Proof.CmacAes.AArch64.FinalizeCorrect +import VerifiedGarbage.Proof.CmacAes.AArch64.UpdateCT + +/-! +# AES-CMAC on AArch64: `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`), the call of `vg_aes_ctr32` is constant time by its own proof +(`ctr_rel`), its arguments pinned by the correctness proof (`FMid`), and the +restore after it by the taint analysis again, from `x19` (the scratch buffer). +-/ + +namespace VG.Proof.CmacAes.AArch64 + +open VG VG.AArch64 VG.Impl.CmacAes.AArch64 +open VG.Proof.Aes.AArch64 (Ctr32Impl) + +theorem finalize_rel (v : Ctr32Impl) {s₀ s₀' : State} (h0 : finalizeAArch64.pre s₀) + (h0' : finalizeAArch64.pre s₀') (hq : finalizeAArch64.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 .x0) (s₀.gpr .x2) (s₀.gpr .x3) (s₀.gpr .x5) (s₀.gpr .x4).toNat + (s₀.gpr .x1).toNat := by + rw [q1, q2, q3, q4, q5, q6]; exact FPre.of h0' + obtain ⟨_, hA⟩ : ∃ h, (taint.check (Taint.ofRegs [.x0, .x1, .x2, .x3, .x4, .x5]) finPre h).isSome = + true := ⟨_, by taint_decide⟩ + obtain ⟨_, hB⟩ : ∃ h, (taint.check (Taint.ofRegs [.x19]) + (.block [.ldr .x .x30 .x19 2072, .ldr .x .x19 .x19 2064]) 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 agree_of q7 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 <;> 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 .x0) (s₀.gpr .x2) (s₀.gpr .x3) (s₀.gpr .x5) (s₀.gpr .x4).toNat (s₀.gpr .x1).toNat s₁ ∧ + FMid s₀' (s₀.gpr .x0) (s₀.gpr .x2) (s₀.gpr .x3) (s₀.gpr .x5) (s₀.gpr .x4).toNat (s₀.gpr .x1).toNat s₂) + fun s₁ s₂ h => ⟨h.1.pre, h.2.pre, by rw [h.1.sp, h.2.sp, q7]⟩).wp + (F₁ := fun (s : State) => s.gpr .x19 = s₀.gpr .x5 ∧ s.sp = s₀.sp) + (F₂ := fun (s : State) => s.gpr .x19 = s₀.gpr .x5 ∧ s.sp = s₀'.sp) fun s₁ s₂ h => + ⟨WP.mono (ctr_call v h.1.pre) fun _ hc => + ⟨by rw [hc.saved .x19 (by simp [preserved]) (by decide), h.1.x19], by rw [hc.sp, h.1.sp]⟩, + WP.mono (ctr_call v h.2.pre) fun _ hc => + ⟨by rw [hc.saved .x19 (by simp [preserved]) (by decide), h.2.x19], by rw [hc.sp, h.2.sp]⟩⟩ + have b := RelCT.taint (A := taint) + (P := fun s₁ s₂ => (s₁.gpr .x19 = s₀.gpr .x5 ∧ s₁.sp = s₀.sp) ∧ (s₂.gpr .x19 = s₀.gpr .x5 ∧ s₂.sp = s₀'.sp)) _ + (fun s₁ s₂ h => agree_of (by rw [h.1.2, h.2.2, q7]) fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + subst hr; rw [h.1.1, h.2.1]) hB + exact (a.mono (fun _ _ h => h) fun _ _ h => h.2).seq + ((c.mono (fun _ _ h => h) fun _ _ h => h.2).seq b) + +theorem finalize_ct (v : Ctr32Impl) : + ConstantTime isa finalizeAArch64.pre finalizeAArch64.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.AArch64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/AArch64/FinalizeCorrect.lean b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/FinalizeCorrect.lean new file mode 100644 index 000000000..95abfa2bd --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/FinalizeCorrect.lean @@ -0,0 +1,450 @@ +import VerifiedGarbage.Proof.CmacAes.AArch64.Finalize + +/-! +# AES-CMAC on AArch64: `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`, the state is zeroed, and +`x19` and `x30` are saved in the scratch buffer; the call leaves +`CIPH_K(C ⊕ Mₙ)` there, the MAC (`macFull_split`). +-/ + +namespace VG.Proof.CmacAes.AArch64 + +open VG VG.AArch64 VG.AArch64.RegUpd VG.Impl.CmacAes.AArch64 +open VG.Proof.Aes.AArch64 (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 + x0 : s₀.gpr .x0 = W + x2 : s₀.gpr .x2 = St + x3 : s₀.gpr .x3 = P + x4 : (s₀.gpr .x4).toNat = L + x5 : s₀.gpr .x5 = S + x1 : (s₀.gpr .x1).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⟩ + 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 : finalizeAArch64.pre s₀) : + FPre s₀ (s₀.gpr .x0) (s₀.gpr .x2) (s₀.gpr .x3) (s₀.gpr .x5) (s₀.gpr .x4).toNat (s₀.gpr .x1).toNat := + let ⟨a, b, c, d, e, f, g, h, i, j, k, l, m⟩ := h + ⟨rfl, rfl, rfl, rfl, rfl, rfl, a, b, c, d, e, f, g, h, i, j, k, l, m⟩ + +/-- 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 + x0 : s.gpr .x0 = W + x2 : s.gpr .x2 = St + x5 : s.gpr .x5 = S + x1 : s.gpr .x1 = s₀.gpr .x1 + saved : ∀ r ∈ preserved, s.gpr r = s₀.gpr r + sp : s.sp = s₀.sp + 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 not_x9 {r : Reg} (hr : r ∈ preserved) : r ≠ .x9 := by rintro rfl; simp [preserved] at hr +theorem not_x10 {r : Reg} (hr : r ∈ preserved) : r ≠ .x10 := by rintro rfl; simp [preserved] at hr +theorem not_x6 {r : Reg} (hr : r ∈ preserved) : r ≠ .x6 := by rintro rfl; simp [preserved] at hr +theorem not_x7 {r : Reg} (hr : r ∈ preserved) : r ≠ .x7 := by rintro rfl; simp [preserved] at hr +theorem not_x8 {r : Reg} (hr : r ∈ preserved) : r ≠ .x8 := by rintro rfl; simp [preserved] at hr + +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 : ∀ r, r ≠ .x9 → s.gpr r = s₀.gpr r) (hm : s.mem = s₀.mem) (hsp : s.sp = s₀.sp) + (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 sw := hp.scr_wrap + have kw := hp.key_wrap + obtain ⟨s', run, mem, g, sp, rd, wr⟩ := xor2_ok s .x3 .x0 .x5 0 240 2048 + (P := P) (Q := W + BitVec.ofNat 64 240) (C := S + BitVec.ofNat 64 2048) + (by decide) (by decide) (by decide) + (by rw [hg _ (by decide), hp.x3, k0]) (by rw [hg _ (by decide), hp.x3]) + (by rw [hg _ (by decide), hp.x0]) (by rw [hg _ (by decide), hp.x0, Offset.add_add]) + (by rw [hg _ (by decide), hp.x5]) (by rw [hg _ (by decide), hp.x5, Offset.add_add]) + ⟨by decide, by decide, by decide, 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 [full_eq] + refine WP.of_runBlock ⟨s', run, ?_⟩ + have gg (r : Reg) (h₁ : r ≠ .x9) (h₂ : r ≠ .x10) : s'.gpr r = s₀.gpr r := by rw [g r h₁ h₂, hg r h₁] + refine ⟨by rw [gg _ (by decide) (by decide), hp.x0], by rw [gg _ (by decide) (by decide), hp.x2], + by rw [gg _ (by decide) (by decide), hp.x5], gg _ (by decide) (by decide), + fun r hr => gg r (not_x9 hr) (not_x10 hr), by rw [sp, hsp], by rw [rd, hrd], + by rw [wr, hwr], by rw [mem, hm]; exact Proof.Cmac.xor2Mem_frame _ _ _ _, ?_⟩ + rw [mem, hm, Proof.Cmac.xor2Mem_bytes] + · simp only [mn, Spec.Cmac.lastBlock, Proof.Cmac.bytesAt_length, ite_true] + exact Proof.Cmac.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 : ∀ r, r ≠ .x9 → s.gpr r = s₀.gpr r) (hm : s.mem = s₀.mem) (hsp : s.sp = s₀.sp) + (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 hb : L < 2 ^ 64 := by omega + have h4 : s₀.gpr .x4 = BitVec.ofNat 64 L := by rw [← hp.x4]; apply BitVec.eq_of_toNat_eq; simp + obtain ⟨C, hC⟩ : ∃ C, S + BitVec.ofNat 64 2048 = C := ⟨_, rfl⟩ + have hc : s.gpr .x5 + BitVec.ofNat 64 2048 = C := by rw [hg _ (by decide), hp.x5, hC] + have hc8 : s.gpr .x5 + BitVec.ofNat 64 2056 = C + BitVec.ofNat 64 8 := by + rw [hg _ (by decide), hp.x5, ← 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₁, mem₁, x6₁, x7₁, x8₁, g₁, sp₁, rd₁, wr₁⟩ := zero_ok s hc hc8 + (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 = Proof.Cmac.zero2 s₀.mem C := by rw [mem₁, hm] + have fz : Frame [⟨C, 16⟩] s₀.mem (Proof.Cmac.zero2 s₀.mem C) := Proof.Cmac.frame_store2 _ _ _ + have lastZ : Spec.Aes.bytesAt (Proof.Cmac.zero2 s₀.mem C) P L = Spec.Aes.bytesAt s₀.mem P L := + Proof.Cmac.bytesAt_frame fz (fun r hr => by simp only [List.mem_singleton] at hr; subst hr; exact dPC) (by omega) + have g₁' (r : Reg) (h₁ : r ≠ .x6) (h₂ : r ≠ .x7) (h₃ : r ≠ .x8) (h₄ : r ≠ .x9) : s₁.gpr r = s₀.gpr r := by + rw [g₁ r h₁ h₂ h₃ h₄, hg r h₄] + -- Copy the last bytes. + refine WP.seq (WP.mono (Q := fun (s₂ : State) => + s₂.mem = writeBytes (Proof.Cmac.zero2 s₀.mem C) C (Spec.Aes.bytesAt s₀.mem P L) ∧ + s₂.gpr .x6 = C + BitVec.ofNat 64 L ∧ + (∀ r, r ≠ .x6 → r ≠ .x7 → r ≠ .x8 → r ≠ .x9 → s₂.gpr r = s₀.gpr r) ∧ + s₂.sp = s₀.sp ∧ s₂.rd = s₀.rd ∧ s₂.wr = s₀.wr) ?_ fun s₂ h₂ => ?_) + · have ev : isa.eval (.zero .x .x4) s₁ = some (decide (L = 0)) := by + show some (s₁.read .x .x4 == 0) = _ + rw [State.read, g₁' _ (by decide) (by decide) (by decide) (by decide), h4, BitVec.setWidth_eq] + have := ofNat_ne_zero hb + rw [bne] at this + cases hx : (BitVec.ofNat 64 L == 0) <;> rw [hx] at this <;> cases hd : decide (L = 0) <;> simp_all + by_cases hL0 : L = 0 + · subst hL0 + refine WP.ite true (by rw [ev]; rfl) (fun _ => WP.block_nil ?_) (fun h => by cases h) + refine ⟨by rw [zf]; simp [Spec.Aes.bytesAt, writeBytes_nil], by rw [x6₁]; simp, + fun r h₁ h₂ h₃ h₄ => g₁' r h₁ h₂ h₃ h₄, by rw [sp₁, hsp], by rw [rd₁, hrd], by rw [wr₁, hwr]⟩ + · refine WP.ite false (by rw [ev]; simp [hL0]) (fun h => by cases h) fun _ => ?_ + refine WP.mono (copy_ok s₁ (P := P) (C := C) (by omega) hL + (by rw [x7₁, hg _ (by decide), hp.x3]) x6₁ + (by rw [x8₁, hg _ (by decide), h4]) + (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₂, x6₂, g₂, sp₂, rd₂, wr₂⟩ + refine ⟨by rw [m₂, zf, lastZ], x6₂, fun r h₁ h₂ h₃ h₄ => by rw [g₂ r h₁ h₂ h₃ h₄, g₁' r h₁ h₂ h₃ h₄], + by rw [sp₂, sp₁, hsp], by rw [rd₂, rd₁, hrd], by rw [wr₂, wr₁, hwr]⟩ + · obtain ⟨m₂, x6₂, g₂, sp₂, rd₂, wr₂⟩ := h₂ + rw [padK2_eq, WP.block_append_iff] + obtain ⟨s₃, run₃, m₃, g₃, sp₃, rd₃, wr₃⟩ := pad_ok s₂ (B := C + BitVec.ofNat 64 L) + (by rw [x6₂, BitVec.add_zero]) + (by rw [wr₂, ← hC, Offset.add_add]; exact hp.inScr (d := 2048 + L) (n := 1) (by omega)) + refine WP.of_runBlock ⟨s₃, run₃, ?_⟩ + have x5₃ : s₃.gpr .x5 = S := by + rw [g₃ _ (by decide), g₂ _ (by decide) (by decide) (by decide) (by decide), hp.x5] + have x0₃ : s₃.gpr .x0 = W := by + rw [g₃ _ (by decide), g₂ _ (by decide) (by decide) (by decide) (by decide), hp.x0] + obtain ⟨s₄, run₄, m₄, g₄, sp₄, rd₄, wr₄⟩ := xor2_ok s₃ .x5 .x0 .x5 2048 256 2048 + (P := C) (Q := W + BitVec.ofNat 64 256) (C := C) (by decide) (by decide) (by decide) + (by rw [x5₃, hC]) (by rw [x5₃, ← hC, Offset.add_add]) (by rw [x0₃]) (by rw [x0₃, Offset.add_add]) + (by rw [x5₃, hC]) (by rw [x5₃, ← hC, Offset.add_add]) + ⟨by decide, by decide, by decide, 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 ≠ .x6) (h₂ : r ≠ .x7) (h₃ : r ≠ .x8) (h₄ : r ≠ .x9) (h₅ : r ≠ .x10) : + s₄.gpr r = s₀.gpr r := by rw [g₄ r h₄ h₅, g₃ r h₄, g₂ r h₁ h₂ h₃ h₄] + have hlen : (Spec.Aes.bytesAt s₀.mem P L).length = L := Proof.Cmac.bytesAt_length _ _ _ + have fW : Frame [⟨C, 16⟩] (writeBytes (Proof.Cmac.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⟩] (Proof.Cmac.zero2 s₀.mem C) + (writeBytes (Proof.Cmac.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 := + Proof.Cmac.bytesAt_frame16 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 := Proof.Cmac.padded_bytes (Proof.Cmac.zero2 s₀.mem C) C (Spec.Aes.bytesAt s₀.mem P L) + (by rw [hlen]; exact hL) (Proof.Cmac.zero2_bytes _ _) + rw [hlen] at this + rw [m₃, m₂]; exact this + refine ⟨by rw [gg _ (by decide) (by decide) (by decide) (by decide) (by decide), hp.x0], + by rw [gg _ (by decide) (by decide) (by decide) (by decide) (by decide), hp.x2], + by rw [gg _ (by decide) (by decide) (by decide) (by decide) (by decide), hp.x5], + gg _ (by decide) (by decide) (by decide) (by decide) (by decide), + fun r hr => gg r (not_x6 hr) (not_x7 hr) (not_x8 hr) (not_x9 hr) (not_x10 hr), + by rw [sp₄, sp₃, sp₂], by rw [rd₄, rd₃, rd₂], by rw [wr₄, wr₃, wr₂], ?_, ?_⟩ + · rw [m₄, hC]; exact f₃.trans (Proof.Cmac.xor2Mem_frame _ _ _ _) + · rw [m₄, hC, Proof.Cmac.xor2Mem_bytes, pad, k2] + · simp only [mn, Spec.Cmac.lastBlock, hlen, show L ≠ 16 by omega, ite_false] + exact Proof.Cmac.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 + BitVec.ofNat 64 2064, 16⟩] s₀.mem s.mem + slot19 : s.mem.readW (S + BitVec.ofNat 64 2064) 64 = s₀.gpr .x19 + slot30 : s.mem.readW (S + BitVec.ofNat 64 2072) 64 = s₀.gpr .x30 + x19 : s.gpr .x19 = S + saved : ∀ r ∈ preserved, r ≠ .x19 → s.gpr r = s₀.gpr r + sp : s.sp = s₀.sp + rd : s.rd = s₀.rd + wr : s.wr = s₀.wr + +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 + rw [finArgs_eq, 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₁, sp₁, rd₁, wr₁⟩ := xor2_ok s .x5 .x2 .x5 2048 0 2048 + (P := S + BitVec.ofNat 64 2048) (Q := St) (C := S + BitVec.ofNat 64 2048) (by decide) (by decide) (by decide) + (by rw [h.x5]) (by rw [h.x5, Offset.add_add]) (by rw [h.x2, k0]) (by rw [h.x2]) + (by rw [h.x5]) (by rw [h.x5, Offset.add_add]) ⟨by decide, by decide, by decide, 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₂, x3₂, x2₂, x4₂, x19₂, g₂, sp₂, rd₂, wr₂⟩ := args_ok s₁ (D := St) (S := S) + (by rw [g₁ _ (by decide) (by decide), h.x2]) (by rw [g₁ _ (by decide) (by decide), h.x5]) + (by rw [wr₁]; exact inSt 0 (by decide)) (by rw [wr₁]; exact inSt 8 (by decide)) + (by rw [wr₁]; exact inC 2064 (by decide)) (by rw [wr₁]; exact inC 2072 (by decide)) + refine WP.of_runBlock ⟨s₂, run₂, ?_⟩ + have g (r : Reg) (h₁ : r ≠ .x2) (h₂ : r ≠ .x3) (h₃ : r ≠ .x4) (h₄ : r ≠ .x9) (h₅ : r ≠ .x19) (h₆ : r ≠ .x10) : + s₂.gpr r = s.gpr r := by + rw [g₂ r h₁ h₂ h₃ h₄ h₅, g₁ r h₄ h₆] + have dCSt : (⟨S + BitVec.ofNat 64 2048, 16⟩ : Region).Disjoint ⟨St, 16⟩ := + hp.st_scr.symm.sub_left (FPre.scrD (by decide)) + have dSlSt : (⟨S + BitVec.ofNat 64 2064, 16⟩ : Region).Disjoint ⟨St, 16⟩ := + hp.st_scr.symm.sub_left (FPre.scrD (by decide)) + have dSlC : (⟨S + BitVec.ofNat 64 2064, 16⟩ : Region).Disjoint ⟨S + BitVec.ofNat 64 2048, 16⟩ := + Offset.disjoint S (by decide) (by omega) (by omega) + -- The memory: the state zeroed, then the two slots. + obtain ⟨Z, hZ⟩ : ∃ Z, (s₁.mem.writeW (St + BitVec.ofNat 64 0) (0 : BitVec 64)).writeW (St + BitVec.ofNat 64 8) + (0 : BitVec 64) = Z := ⟨_, rfl⟩ + have e72 : S + BitVec.ofNat 64 2072 = S + BitVec.ofNat 64 2064 + BitVec.ofNat 64 8 := + (Offset.add_add_eq S (a := 2064) (b := 8) (c := 2072) rfl).symm + have fSl : Frame [⟨S + BitVec.ofNat 64 2064, 16⟩] Z s₂.mem := by + rw [m₂, hZ, e72]; exact Proof.Cmac.frame_store2 _ _ _ + have fZ : Frame [⟨St, 16⟩] s₁.mem Z := by rw [← hZ]; exact frame_store2' _ _ _ + have stS : Spec.Aes.bytesAt s.mem St 16 = Spec.Aes.bytesAt s₀.mem St 16 := + Proof.Cmac.bytesAt_frame16 h.frame fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact dCSt.symm + have hR := hp.rounds + refine ⟨?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, by rw [rd₂, rd₁, h.rd], by rw [wr₂, wr₁, h.wr]⟩ + · exact + { x0 := by rw [g _ (by decide) (by decide) (by decide) (by decide) (by decide) (by decide), h.x0] + x1 := by + rw [g _ (by decide) (by decide) (by decide) (by decide) (by decide) (by decide), h.x1] + apply BitVec.eq_of_toNat_eq; simp [hp.x1]; omega + x2 := x2₂ + x3 := x3₂ + x4 := x4₂ + x5 := by rw [g _ (by decide) (by decide) (by decide) (by decide) (by decide) (by decide), h.x5] + 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 omega) + ds := hp.st_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 [Proof.Cmac.bytesAt_frame16 fSl (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact dSlSt.symm), ← hZ, k0, + Proof.Cmac.bytesAt_store2, Proof.Cmac.le8_zero, Proof.Cmac.zeros_8_8] } + · rw [Proof.Cmac.bytesAt_frame16 fSl (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact dSlC.symm), + Proof.Cmac.bytesAt_frame16 fZ (fun r hr => by + simp only [List.mem_singleton] at hr; subst hr; exact dCSt), + m₁, Proof.Cmac.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₁ ▸ Proof.Cmac.xor2Mem_frame _ _ _ _)).mono (by simp)).trans + (fZ.mono (by simp))).trans (fSl.mono (by simp)) + · rw [m₂, readW_writeW_other _ _ _ (by decide) (by decide) (by decide), Mem.readW_writeW_self64, + g₁ _ (by decide) (by decide), h.saved _ (by simp [preserved])] + · rw [m₂, Mem.readW_writeW_self64, g₁ _ (by decide) (by decide), h.saved _ (by simp [preserved])] + · exact x19₂ + · intro r hr h19 + have n2 : r ≠ .x2 := by rintro rfl; simp [preserved] at hr + have n3 : r ≠ .x3 := by rintro rfl; simp [preserved] at hr + have n4 : r ≠ .x4 := by rintro rfl; simp [preserved] at hr + rw [g r n2 n3 n4 (not_x9 hr) h19 (not_x10 hr), h.saved r hr] + · rw [sp₂, sp₁, h.sp] + +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 h4 : s₀.gpr .x4 = BitVec.ofNat 64 L := by rw [← hp.x4]; apply BitVec.eq_of_toNat_eq; simp + obtain ⟨s₁, run₁, ev₁, g₁, sp₁, m₁, rd₁, wr₁⟩ := sub16_ok s₀ h4 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 rw [ev₁]; simp [hL]) (fun _ => full_wp hp hL g₁ m₁ sp₁ rd₁ wr₁) + (fun h => by cases h) + · exact WP.ite false (by rw [ev₁]; simp [hL]) (fun h => by cases h) + (fun _ => partial_wp hp (by have := hp.len; omega) g₁ m₁ sp₁ rd₁ wr₁) + +theorem restoreF_ok (s : State) {B : Addr} (hb : s.gpr .x19 = B) + (r₁ : InRegions (s.rd ++ s.wr) (B + BitVec.ofNat 64 2072) 8) + (r₂ : InRegions (s.rd ++ s.wr) (B + BitVec.ofNat 64 2064) 8) : + ∃ s', runBlock isa [.ldr .x .x30 .x19 2072, .ldr .x .x19 .x19 2064] s = some s' ∧ + s'.gpr .x30 = s.mem.readW (B + BitVec.ofNat 64 2072) 64 ∧ + s'.gpr .x19 = s.mem.readW (B + BitVec.ofNat 64 2064) 64 ∧ + (∀ r, r ≠ .x19 → r ≠ .x30 → s'.gpr r = s.gpr r) ∧ s'.sp = s.sp ∧ s'.mem = s.mem := by + refine ⟨_, by + simp (config := {decide := true}) only [runBlock_cons, runStep_some, runBlock_nil, exec, addr, + State.load, Size.bytes, Size.bits, gpr_write, mem_write, rd_write, wr_write, ite_true, ite_false, + Option.bind_some, Option.map_some, hb, r₁, r₂] + rfl, ?_⟩ + refine ⟨by simp [gpr_write, Mem.readW], by simp [gpr_write, Mem.readW], + fun r h₁ h₂ => by simp [gpr_write, h₁, h₂], rfl, rfl⟩ + +theorem finalize_wp (v : Ctr32Impl) {s₀ : State} (h0 : finalizeAArch64.pre s₀) : + WP isa (finalize v.callee) s₀ fun s' => GprAbi s₀ s' ∧ finalizeAArch64.post s₀ s' := by + have hp := FPre.of h0 + generalize s₀.gpr .x0 = W at hp + generalize s₀.gpr .x2 = St at hp + generalize s₀.gpr .x3 = P at hp + generalize s₀.gpr .x5 = S at hp + generalize (s₀.gpr .x4).toNat = L at hp + generalize (s₀.gpr .x1).toNat = R at hp + have hR := hp.rounds + have hRb : 16 * (R + 1) ≤ 240 := by rcases hR with h | h | h <;> omega + have sw := hp.scr_wrap + refine WP.seq (WP.mono (finPre_wp hp) fun s₁ h₁ => ?_) + refine WP.seq (WP.mono (ctr_call v h₁.pre) fun s₂ h₂ => ?_) + have x19₂ : s₂.gpr .x19 = S := by rw [h₂.saved .x19 (by simp [preserved]) (by decide), h₁.x19] + have rdwr₂ : s₂.rd ++ s₂.wr = s₀.rd ++ s₀.wr := by rw [h₂.rd, h₂.wr, h₁.rd, h₁.wr] + obtain ⟨s₃, run₃, x30₃, x19₃, g₃, sp₃, mem₃⟩ := restoreF_ok s₂ x19₂ + (by rw [rdwr₂]; exact wr_in (hp.inScr (d := 2072) (n := 8) (by decide))) + (by rw [rdwr₂]; exact wr_in (hp.inScr (d := 2064) (n := 8) (by decide))) + refine WP.of_runBlock ⟨s₃, run₃, ?_⟩ + -- The slots, which the call does not write. + have slots (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 + refine h₂.frame.readW (r := ⟨S + BitVec.ofNat 64 d, 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 Offset.disjoint S (by omega) (by omega) (by omega) + · exact (hp.st_scr.symm.sub_left (FPre.scrD (by omega))) + · exact Offset.disjoint_base _ (by omega) (by omega) + have big : Frame [⟨St, 16⟩, ⟨S, 2176⟩] 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 | rfl + · exact ⟨⟨S, 2176⟩, by simp, FPre.scrD (by decide)⟩ + · exact ⟨⟨St, 16⟩, by simp, fun _ h => h⟩ + · exact ⟨⟨S, 2176⟩, by simp, FPre.scrD (by decide)⟩ + · simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with 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)⟩ + have sch : Spec.Aes.bytesAt s₁.mem W (16 * (R + 1)) = Spec.Aes.bytesAt s₀.mem W (16 * (R + 1)) := + Proof.Cmac.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 | 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)) + · exact (hp.key_scr.sub_left (Region.sub_prefix (by omega))).sub_right (FPre.scrD (by decide))) (by omega) + refine ⟨⟨fun r hr => ?_, by rw [sp₃, h₂.sp, h₁.sp]⟩, ?_⟩ + · by_cases h19 : r = .x19 + · subst h19; rw [x19₃, slots 2064 (by decide) (by decide), h₁.slot19] + by_cases h30 : r = .x30 + · subst h30; rw [x30₃, slots 2072 (by decide) (by decide), h₁.slot30] + rw [g₃ r h19 h30, h₂.saved r hr h30, h₁.saved r hr h19] + · intro hk msg hm hne hst + rw [hp.x0, hp.x1] at hk hst ⊢ + rw [hp.x2] at hst ⊢ + rw [hp.x4] at hne + rw [hp.x3, hp.x4] + obtain ⟨e1, e2⟩ := Proof.Cmac.k1k2 (Proof.Cmac.subkeys_aes_length _ _) hk + rw [mem₃, 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), Proof.Cmac.xor_comm] + +end VG.Proof.CmacAes.AArch64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Subkeys.lean b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Subkeys.lean new file mode 100644 index 000000000..4ddbf3e6a --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Subkeys.lean @@ -0,0 +1,372 @@ +import VerifiedGarbage.Proof.CmacAes.AArch64.Dbl +import VerifiedGarbage.Proof.CmacAes.AArch64.UpdateCorrect + +/-! +# AES-CMAC on AArch64: `vg_cmac_aes_subkeys` + +Untrusted: everything here is checked by Lean. +-/ + +namespace VG.Proof.CmacAes.AArch64 + +open VG VG.AArch64 VG.AArch64.RegUpd VG.Impl.CmacAes.AArch64 +open VG.Proof.Aes.AArch64 (Ctr32Impl) + +/-! ## Doubling a block -/ + +/-- The memory after `dbl src dst`, from `x19 = K`. -/ +def dblMem (m : Mem) (K : Addr) (src dst : Nat) : Mem := + let hi := rev64 (m.readW (K + BitVec.ofNat 64 src) 64) + let lo := rev64 (m.readW (K + BitVec.ofNat 64 (src + 8)) 64) + (m.writeW (K + BitVec.ofNat 64 dst) (rev64 (dblHi hi lo))).writeW (K + BitVec.ofNat 64 (dst + 8)) + (rev64 (dblLo hi lo)) + +theorem mz0 : BitVec.setWidth 64 (0 : BitVec 16) <<< (16 * 0) = 0 := by decide +theorem mz87 : BitVec.setWidth 64 (135 : BitVec 16) <<< (16 * 0) = 0x87 := by decide + +theorem dbl_ok (s : State) {K : Addr} (hb : s.gpr .x19 = K) {src dst : Nat} + (hs : src % 8 = 0 ∧ src + 8 < 32768) (hd : dst % 8 = 0 ∧ dst + 8 < 32768) + (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 ≠ .x9 → r ≠ .x10 → r ≠ .x11 → r ≠ .x12 → s'.gpr r = s.gpr r) ∧ + s'.sp = s.sp ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + simp (config := {decide := true}) only [dbl, runBlock_cons, runStep_some, runBlock_nil, exec, addr, + State.load, State.store, Size.bytes, Size.bits, State.read, gpr_write, mem_write, rd_write, wr_write, + ite_true, ite_false, Option.bind_some, Option.map_some, hb, hs.1, hd.1, BitVec.setWidth_eq, + show src < 32768 by omega, show src + 8 < 32768 from hs.2, show dst < 32768 by omega, + show dst + 8 < 32768 from hd.2, Nat.add_mod_right, r₀, r₁, w₀, w₁, and_self] + rfl, ?_⟩ + refine ⟨?_, fun r h₁ h₂ h₃ h₄ => by simp [gpr_write, h₁, h₂, h₃, h₄], rfl, rfl, rfl⟩ + simp only [dblMem, dblHi, dblLo, Mem.writeW, Mem.readW, BitVec.setWidth_eq, mz0, mz87] + +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 Proof.Cmac.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_rev, dbl_words, Proof.Cmac.dbl_eq (Proof.Cmac.bytesAt_length _ _ _), + ← Spec.Gcm.blockAt, ← Proof.Gcm.AArch64.blockAt_rev, BitVec.add_zero] + +/-! ## Before the call -/ + +/-- The memory after `subkeysPre`. -/ +def preMem (s : State) : Mem := + ((((((s.mem.writeW (s.gpr .x3 + BitVec.ofNat 64 2064) (s.gpr .x19)).writeW + (s.gpr .x3 + BitVec.ofNat 64 2072) (s.gpr .x20)).writeW + (s.gpr .x3 + BitVec.ofNat 64 2080) (s.gpr .x30)).writeW + (s.gpr .x3 + BitVec.ofNat 64 2048) (0 : BitVec 64)).writeW + (s.gpr .x3 + BitVec.ofNat 64 2056) (0 : BitVec 64)).writeW + (s.gpr .x2 + BitVec.ofNat 64 0) (0 : BitVec 64)).writeW + (s.gpr .x2 + BitVec.ofNat 64 8) (0 : BitVec 64) + +theorem subkeysPre_ok (s : State) + (w₁ : InRegions s.wr (s.gpr .x3 + BitVec.ofNat 64 2064) 8) + (w₂ : InRegions s.wr (s.gpr .x3 + BitVec.ofNat 64 2072) 8) + (w₃ : InRegions s.wr (s.gpr .x3 + BitVec.ofNat 64 2080) 8) + (w₄ : InRegions s.wr (s.gpr .x3 + BitVec.ofNat 64 2048) 8) + (w₅ : InRegions s.wr (s.gpr .x3 + BitVec.ofNat 64 2056) 8) + (w₆ : InRegions s.wr (s.gpr .x2 + BitVec.ofNat 64 0) 8) + (w₇ : InRegions s.wr (s.gpr .x2 + BitVec.ofNat 64 8) 8) : + ∃ s', runBlock isa subkeysPre s = some s' ∧ + s'.gpr .x0 = s.gpr .x0 ∧ s'.gpr .x1 = s.gpr .x1 ∧ + s'.gpr .x2 = s.gpr .x3 + BitVec.ofNat 64 2048 ∧ s'.gpr .x3 = s.gpr .x2 ∧ s'.gpr .x4 = 1 ∧ + s'.gpr .x5 = s.gpr .x3 ∧ s'.gpr .x19 = s.gpr .x2 ∧ s'.gpr .x20 = s.gpr .x3 ∧ + (∀ r ∈ preserved, r ≠ .x19 → r ≠ .x20 → s'.gpr r = s.gpr r) ∧ s'.sp = s.sp ∧ + s'.mem = preMem s ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + simp (config := {decide := true}) only [subkeysPre, ctrArgs, cOff, mov, List.cons_append, + List.nil_append, runBlock_cons, runStep_some, runBlock_nil, exec, addr, State.store, Size.bytes, + Size.bits, State.read, gpr_write, mem_write, wr_write, ite_true, ite_false, Option.bind_some, + BitVec.setWidth_eq, w₁, w₂, w₃, w₄, w₅, w₆, w₇] + rfl, ?_⟩ + refine ⟨by simp [gpr_write], by simp [gpr_write], by simp [gpr_write], by simp [gpr_write], rfl, + by simp [gpr_write], by simp [gpr_write], by simp [gpr_write], fun r hr h₁ h₂ => ?_, rfl, ?_, rfl, rfl⟩ + · simp only [preserved, List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl <;> simp_all [gpr_write] + · simp only [mem_write, preMem, Mem.writeW, BitVec.setWidth_eq, mz0] + +/-! ## 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 + x0 : s₀.gpr .x0 = W + x2 : s₀.gpr .x2 = K + x3 : s₀.gpr .x3 = S + x1 : (s₀.gpr .x1).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⟩ + 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 : subkeysAArch64.pre s₀) : + SPre s₀ (s₀.gpr .x0) (s₀.gpr .x2) (s₀.gpr .x3) (s₀.gpr .x1).toNat := + let ⟨a, b, c, d, e, f, g, h⟩ := h + ⟨rfl, rfl, rfl, rfl, a, b, c, d, e, f, g, h⟩ + +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 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 .x3 + BitVec.ofNat 64 2048, 40⟩, ⟨s.gpr .x2, 16⟩] s.mem (preMem s) := by + have c (d : Nat) (h : d + 8 ≤ 40) : + (⟨s.gpr .x3 + BitVec.ofNat 64 2048, 40⟩ : Region).Contains (s.gpr .x3 + 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 .x2, 16⟩ : Region).Contains (s.gpr .x2 + 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 32 (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 Proof.Cmac.frame_store2 _ _ _ + +theorem restore3_ok (s : State) {B : Addr} (hb : s.gpr .x20 = 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) + (r₃ : InRegions (s.rd ++ s.wr) (B + BitVec.ofNat 64 2080) 8) : + ∃ s', runBlock isa [.ldr .x .x30 .x20 2080, .ldr .x .x19 .x20 2064, .ldr .x .x20 .x20 2072] s = some s' ∧ + s'.gpr .x19 = s.mem.readW (B + BitVec.ofNat 64 2064) 64 ∧ + s'.gpr .x20 = s.mem.readW (B + BitVec.ofNat 64 2072) 64 ∧ + s'.gpr .x30 = s.mem.readW (B + BitVec.ofNat 64 2080) 64 ∧ + (∀ r, r ≠ .x19 → r ≠ .x20 → r ≠ .x30 → s'.gpr r = s.gpr r) ∧ s'.sp = s.sp ∧ s'.mem = s.mem := by + refine ⟨_, by + simp (config := {decide := true}) only [runBlock_cons, runStep_some, runBlock_nil, exec, addr, + State.load, Size.bytes, Size.bits, gpr_write, mem_write, rd_write, wr_write, ite_true, ite_false, + Option.bind_some, Option.map_some, hb, r₁, r₂, r₃] + rfl, ?_⟩ + refine ⟨by simp [gpr_write, Mem.readW], by simp [gpr_write, Mem.readW], by simp [gpr_write, Mem.readW], + fun r h₁ h₂ h₃ => by simp [gpr_write, h₁, h₂, h₃], rfl, rfl⟩ + +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 := + readW_writeW_other m b v h hd he + +theorem preMem_slot (s : State) {d : Nat} (hd : d = 2064 ∨ d = 2072 ∨ d = 2080) + (hks : (⟨s.gpr .x2, 32⟩ : Region).Disjoint ⟨s.gpr .x3, 2176⟩) : + (preMem s).readW (s.gpr .x3 + BitVec.ofNat 64 d) 64 = + if d = 2064 then s.gpr .x19 else if d = 2072 then s.gpr .x20 else s.gpr .x30 := by + have kd (e : Nat) (he : e + 8 ≤ 32) : Mem.Sep (s.gpr .x3 + BitVec.ofNat 64 d) (64 / 8) + (s.gpr .x2 + 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 | rfl + · rw [readW_writeW_other _ _ _ (by decide) (by decide) (by decide), + readW_writeW_other _ _ _ (by decide) (by decide) (by decide), Mem.readW_writeW_self64]; 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} + (x0₁ : s₁.gpr .x0 = s₀.gpr .x0) (x1₁ : s₁.gpr .x1 = s₀.gpr .x1) + (x2₁ : s₁.gpr .x2 = s₀.gpr .x3 + BitVec.ofNat 64 2048) (x3₁ : s₁.gpr .x3 = s₀.gpr .x2) + (x4₁ : s₁.gpr .x4 = 1) (x5₁ : s₁.gpr .x5 = s₀.gpr .x3) + (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.x2, k0, Proof.Cmac.bytesAt_store2, Proof.Cmac.le8_zero, zeros_8_8] + exact + { x0 := by rw [x0₁, hp.x0] + x1 := by rw [x1₁]; apply BitVec.eq_of_toNat_eq; simp [hp.x1]; omega + x2 := by rw [x2₁, hp.x3] + x3 := by rw [x3₁, hp.x2] + x4 := x4₁ + x5 := by rw [x5₁, hp.x3] + 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)) + 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 : subkeysAArch64.pre s₀) : + WP isa (subkeys v.callee) s₀ fun s' => GprAbi s₀ s' ∧ subkeysAArch64.post s₀ s' := by + have hp := SPre.of h0 + generalize s₀.gpr .x0 = W at hp + generalize s₀.gpr .x2 = K at hp + generalize s₀.gpr .x3 = S at hp + generalize (s₀.gpr .x1).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₁, x0₁, x1₁, x2₁, x3₁, x4₁, x5₁, x19₁, x20₁, cs₁, sp₁, mem₁, rd₁, wr₁⟩ := + subkeysPre_ok s₀ (by rw [hp.x3]; exact inS _ (by decide)) (by rw [hp.x3]; exact inS _ (by decide)) + (by rw [hp.x3]; exact inS _ (by decide)) (by rw [hp.x3]; exact inS _ (by decide)) + (by rw [hp.x3]; exact inS _ (by decide)) + (by rw [hp.x2]; exact inK _ (by decide)) (by rw [hp.x2]; 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, 40⟩, ⟨K, 16⟩] s₀.mem s₁.mem := by + rw [mem₁, ← hp.x3, ← hp.x2]; 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.x3, hp.x2] + rw [Proof.Cmac.bytesAt_frame16 (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, + Proof.Cmac.le8_zero, zeros_8_8] + have zK : Spec.Aes.bytesAt s₁.mem K 16 = Spec.Cmac.zeros 16 := by + rw [mem₁, preMem, hp.x2, k0, Proof.Cmac.bytesAt_store2, Proof.Cmac.le8_zero, zeros_8_8] + have schB : ∀ m : Mem, Frame [⟨S + BitVec.ofNat 64 2048, 40⟩, ⟨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 => + Proof.Cmac.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 pre := callPre_of hp x0₁ x1₁ x2₁ x3₁ x4₁ x5₁ mem₁ rd₁ wr₁ + refine WP.seq (WP.mono (ctr_call v pre) fun s₂ h₂ => ?_) + -- After the call. + have x19₂ : s₂.gpr .x19 = K := by rw [h₂.saved .x19 (by simp [preserved]) (by decide), x19₁, hp.x2] + have x20₂ : s₂.gpr .x20 = S := by rw [h₂.saved .x20 (by simp [preserved]) (by decide), x20₁, hp.x3] + 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₃, sp₃, rd₃, wr₃⟩ := dbl_ok s₂ x19₂ (src := 0) (dst := 0) (by decide) (by decide) + (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 x19₃ : s₃.gpr .x19 = K := by rw [g₃ _ (by decide) (by decide) (by decide) (by decide), x19₂] + obtain ⟨s₄, run₄, mem₄, g₄, sp₄, rd₄, wr₄⟩ := dbl_ok s₃ x19₃ (src := 0) (dst := 16) (by decide) (by decide) + (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 x20₄ : s₄.gpr .x20 = S := by + rw [g₄ _ (by decide) (by decide) (by decide) (by decide), g₃ _ (by decide) (by decide) (by decide) (by decide), + x20₂] + obtain ⟨s₅, run₅, x19₅, x20₅, x30₅, g₅, sp₅, mem₅⟩ := restore3_ok s₄ x20₄ + (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))) + (by rw [rd₄, wr₄, rd₃, wr₃, rdwr₂]; exact rIn _ (inS 2080 (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⟩], + (⟨S + BitVec.ofNat 64 2064, 24⟩ : 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 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) + have slotK (e : Nat) (he : e ≤ 16) : (⟨S + BitVec.ofNat 64 2064, 24⟩ : 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 ≤ 2088) : + 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, 24⟩ : 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.x2, hp.x3]; exact hp.k_scr) + rw [hp.x3] at pslot + have gk (r : Reg) (hr : r ∈ preserved) (h19 : r ≠ .x19) (h20 : r ≠ .x20) (h30 : r ≠ .x30) : + s₅.gpr r = s₀.gpr r := by + have n9 : r ≠ .x9 := by rintro rfl; simp [preserved] at hr + have n10 : r ≠ .x10 := by rintro rfl; simp [preserved] at hr + have n11 : r ≠ .x11 := by rintro rfl; simp [preserved] at hr + have n12 : r ≠ .x12 := by rintro rfl; simp [preserved] at hr + rw [g₅ r h19 h20 h30, g₄ r n9 n10 n11 n12, g₃ r n9 n10 n11 n12, h₂.saved r hr h30, cs₁ r hr h19 h20] + refine ⟨⟨fun r hr => ?_, by rw [sp₅, sp₄, sp₃, h₂.sp, sp₁]⟩, ?_⟩ + · by_cases h19 : r = .x19 + · subst h19; rw [x19₅, slot 2064 (by decide) (by decide), mem₁, pslot 2064 (.inl rfl)]; rfl + by_cases h20 : r = .x20 + · subst h20; rw [x20₅, slot 2072 (by decide) (by decide), mem₁, pslot 2072 (.inr (.inl rfl))]; rfl + by_cases h30 : r = .x30 + · subst h30; rw [x30₅, slot 2080 (by decide) (by decide), mem₁, pslot 2080 (.inr (.inr rfl))]; rfl + exact gk r hr h19 h20 h30 + · show Spec.Aes.bytesAt s₅.mem (s₀.gpr .x2) 32 = _ + rw [hp.x2, hp.x0, hp.x1, 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 := + Proof.Cmac.bytesAt_frame16 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.AArch64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/AArch64/SubkeysCT.lean b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/SubkeysCT.lean new file mode 100644 index 000000000..b014a5368 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/SubkeysCT.lean @@ -0,0 +1,86 @@ +import VerifiedGarbage.Proof.CmacAes.AArch64.Subkeys +import VerifiedGarbage.Proof.CmacAes.AArch64.UpdateCT + +/-! +# AES-CMAC on AArch64: `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 `x19` and `x20`); the call of +`vg_aes_ctr32` is constant time by its own proof (`ctr_rel`). +-/ + +namespace VG.Proof.CmacAes.AArch64 + +open VG VG.AArch64 VG.Impl.CmacAes.AArch64 +open VG.Proof.Aes.AArch64 (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 + x19 : s.gpr .x19 = K + x20 : s.gpr .x20 = S + sp : s.sp = s₀.sp + +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 .x3 + BitVec.ofNat 64 d) 8 := by + rw [hp.wr, hp.x3]; 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 .x2 + BitVec.ofNat 64 d) 8 := by + rw [hp.wr, hp.x2]; exact in_rw (r := ⟨K, 32⟩) (by simp) (Offset.contains_base _ h (by omega)) + obtain ⟨s₁, run₁, x0₁, x1₁, x2₁, x3₁, x4₁, x5₁, x19₁, x20₁, _, sp₁, mem₁, rd₁, wr₁⟩ := + subkeysPre_ok s₀ (inS _ (by decide)) (inS _ (by decide)) (inS _ (by decide)) (inS _ (by decide)) + (inS _ (by decide)) (inK _ (by decide)) (inK _ (by decide)) + exact WP.of_runBlock ⟨s₁, run₁, + callPre_of hp x0₁ x1₁ x2₁ x3₁ x4₁ x5₁ mem₁ rd₁ wr₁, by rw [x19₁, hp.x2], by rw [x20₁, hp.x3], sp₁⟩ + +theorem subkeys_rel (v : Ctr32Impl) {s₀ s₀' : State} (h0 : subkeysAArch64.pre s₀) + (h0' : subkeysAArch64.pre s₀') (hq : subkeysAArch64.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 .x0) (s₀.gpr .x2) (s₀.gpr .x3) (s₀.gpr .x1).toNat := by + rw [q1, q2, q3, q4]; exact SPre.of h0' + obtain ⟨_, hA⟩ : ∃ h, (taint.check (Taint.ofRegs [.x0, .x1, .x2, .x3]) (.block subkeysPre) h).isSome = + true := ⟨_, by taint_decide⟩ + obtain ⟨_, hB⟩ : ∃ h, (taint.check (Taint.ofRegs [.x19, .x20]) (.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 agree_of q5 fun r hr => ?_ + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl <;> assumption) hA).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 .x0) (s₀.gpr .x2) (s₀.gpr .x3) (s₀.gpr .x1).toNat s₁ ∧ + SMid s₀' (s₀.gpr .x0) (s₀.gpr .x2) (s₀.gpr .x3) (s₀.gpr .x1).toNat s₂) fun s₁ s₂ h => + ⟨h.1.pre, h.2.pre, by rw [h.1.sp, h.2.sp, q5]⟩).wp + (F₁ := fun (s : State) => s.gpr .x19 = s₀.gpr .x2 ∧ s.gpr .x20 = s₀.gpr .x3 ∧ s.sp = s₀.sp) + (F₂ := fun (s : State) => s.gpr .x19 = s₀.gpr .x2 ∧ s.gpr .x20 = s₀.gpr .x3 ∧ s.sp = s₀'.sp) + fun s₁ s₂ h => + ⟨WP.mono (ctr_call v h.1.pre) fun _ hc => + ⟨by rw [hc.saved .x19 (by simp [preserved]) (by decide), h.1.x19], + by rw [hc.saved .x20 (by simp [preserved]) (by decide), h.1.x20], by rw [hc.sp, h.1.sp]⟩, + WP.mono (ctr_call v h.2.pre) fun _ hc => + ⟨by rw [hc.saved .x19 (by simp [preserved]) (by decide), h.2.x19], + by rw [hc.saved .x20 (by simp [preserved]) (by decide), h.2.x20], by rw [hc.sp, h.2.sp]⟩⟩ + have b := RelCT.taint (A := taint) + (P := fun s₁ s₂ => (s₁.gpr .x19 = s₀.gpr .x2 ∧ s₁.gpr .x20 = s₀.gpr .x3 ∧ s₁.sp = s₀.sp) ∧ + (s₂.gpr .x19 = s₀.gpr .x2 ∧ s₂.gpr .x20 = s₀.gpr .x3 ∧ s₂.sp = s₀'.sp)) _ + (fun s₁ s₂ h => agree_of (by rw [h.1.2.2, h.2.2.2, q5]) 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.1, h.2.2.1]) 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 subkeysAArch64.pre subkeysAArch64.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.AArch64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Update.lean b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Update.lean new file mode 100644 index 000000000..62fb10292 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Update.lean @@ -0,0 +1,68 @@ +import VerifiedGarbage.Proof.CmacAes.AArch64.Contract +import VerifiedGarbage.Proof.Cmac.Frame +import VerifiedGarbage.Proof.Framework.AArch64.Exec +import VerifiedGarbage.Proof.Framework.AArch64.RegUpd + +/-! +# AES-CMAC on AArch64: `vg_cmac_aes_update`, the blocks before and in the loop + +Untrusted: everything here is checked by Lean. +-/ + +namespace VG.Proof.CmacAes.AArch64 + +open VG VG.AArch64 VG.AArch64.RegUpd VG.Impl.CmacAes.AArch64 + +/-- The memory after saving the registers. -/ +def savedMem (s : State) : Mem := + saved.foldl (fun m (r, d) => m.writeW (s.gpr .x5 + BitVec.ofNat 64 d) (s.gpr r)) s.mem + +theorem prologue_ok (s : State) + (hw : ∀ d, 2064 ≤ d → d + 8 ≤ 2120 → InRegions s.wr (s.gpr .x5 + BitVec.ofNat 64 d) 8) : + ∃ s', runBlock isa (save ++ setup) s = some s' ∧ + s'.gpr .x19 = s.gpr .x0 ∧ s'.gpr .x20 = s.gpr .x1 ∧ s'.gpr .x21 = s.gpr .x2 ∧ + s'.gpr .x22 = s.gpr .x3 ∧ s'.gpr .x23 = s.gpr .x4 ∧ s'.gpr .x24 = s.gpr .x5 ∧ + (∀ r, r ≠ .x19 → r ≠ .x20 → r ≠ .x21 → r ≠ .x22 → r ≠ .x23 → r ≠ .x24 → s'.gpr r = s.gpr r) ∧ + s'.sp = s.sp ∧ s'.mem = savedMem s ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + simp (config := {decide := true}) only [save, setup, saved, mov, List.map, List.cons_append, + List.nil_append, runBlock_cons, runStep_some, runBlock_nil, exec, addr, State.store, Size.bytes, + Size.bits, State.read, gpr_write, ite_true, ite_false, Option.bind_some, + 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), + hw 2112 (by decide) (by decide)] + rfl, ?_⟩ + refine ⟨by simp [gpr_write], by simp [gpr_write], by simp [gpr_write], by simp [gpr_write], + by simp [gpr_write], by simp [gpr_write], fun r h₁ h₂ h₃ h₄ h₅ h₆ => ?_, rfl, ?_, rfl, rfl⟩ + · simp [gpr_write, h₁, h₂, h₃, h₄, h₅, h₆] + · simp only [mem_write, savedMem, saved, List.foldl, Mem.writeW, BitVec.setWidth_eq] + +theorem chainIn_ok (s : State) {C P Q : Addr} (hc : s.gpr .x24 + BitVec.ofNat 64 2048 = C) + (hp : s.gpr .x21 = P) (hq : s.gpr .x22 = 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 .x0 = s.gpr .x19 ∧ s'.gpr .x1 = s.gpr .x20 ∧ s'.gpr .x2 = C ∧ s'.gpr .x3 = P ∧ + s'.gpr .x4 = 1 ∧ s'.gpr .x5 = s.gpr .x24 ∧ + (∀ r ∈ preserved, s'.gpr r = s.gpr r) ∧ s'.sp = s.sp ∧ + s'.mem = Proof.Cmac.chainMem s.mem C P Q ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + have hc' : s.gpr .x24 + 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, mov, List.cons_append, + List.nil_append, runBlock_cons, runStep_some, runBlock_nil, exec, addr, State.load, State.store, + Size.bytes, Size.bits, State.read, gpr_write, mem_write, rd_write, wr_write, ite_true, ite_false, + Option.bind_some, Option.map_some, hc, hc', hp, hq, BitVec.add_zero, BitVec.setWidth_eq, + rp, rp8, rq, rq8, wc, wc8, wp, wp8] + rfl, ?_⟩ + simp (config := {decide := true}) only [gpr_write, mem_write, rd_write, wr_write, sp_write, + ite_true, ite_false, BitVec.setWidth_eq] + refine ⟨trivial, trivial, trivial, trivial, trivial, trivial, fun r hr => ?_, trivial, ?_, trivial⟩ + · simp only [preserved, List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl <;> rfl + · simp only [Proof.Cmac.chainMem, Mem.writeW, Mem.readW, BitVec.setWidth_eq] + rfl + +end VG.Proof.CmacAes.AArch64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/AArch64/UpdateCT.lean b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/UpdateCT.lean new file mode 100644 index 000000000..6e2f733a6 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/UpdateCT.lean @@ -0,0 +1,196 @@ +import VerifiedGarbage.Proof.CmacAes.AArch64.UpdateCorrect +import VerifiedGarbage.Proof.Framework.AArch64.Taint + +/-! +# AES-CMAC on AArch64: `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.AArch64 + +open VG VG.AArch64 VG.Impl.CmacAes.AArch64 +open VG.Proof.Aes.AArch64 (Ctr32Impl) + +/-- States agree on the registers `rs` and the stack pointer. -/ +theorem agree_of {rs : List Reg} {s₁ s₂ : State} (hsp : s₁.sp = s₂.sp) + (h : ∀ r ∈ rs, s₁.gpr r = s₂.gpr r) : VG.AArch64.Taint.Agree (Taint.ofRegs rs) s₁ s₂ := + ⟨hsp, fun r hr => h r (VG.AArch64.Taint.mem_ofRegs.mp hr)⟩ + +section +variable {s₀ s₀' : State} (hq : updateAArch64.pub s₀ s₀') +include hq + +theorem pub_W : W s₀ = W s₀' := hq.1 +theorem pub_x1 : s₀.gpr .x1 = s₀'.gpr .x1 := hq.2.1 +theorem pub_R : R s₀ = R s₀' := by rw [R, R, pub_x1 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_sp : s₀.sp = s₀'.sp := 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₂) : + VG.AArch64.Taint.Agree (Taint.ofRegs [.x19, .x20, .x21, .x22, .x23, .x24]) s₁ s₂ := by + refine agree_of (by rw [h₁.sp, h₂.sp, pub_sp hq]) 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 + · rw [h₁.x19, h₂.x19, pub_W hq] + · rw [h₁.x20, h₂.x20, pub_x1 hq] + · rw [h₁.x21, h₂.x21, pub_St hq] + · rw [h₁.x22, h₂.x22, pub_Dp hq] + · rw [h₁.x23, h₂.x23, pub_N hq] + · rw [h₁.x24, h₂.x24, pub_S hq] + +end + +/-! ## One block -/ + +/-- What is known between the code before the call and the call. -/ +structure Mid (s₀ : State) (k : Nat) (s : State) : Prop where + pre : CallPre s (W s₀) (S s₀ + BitVec.ofNat 64 2048) (St s₀) (S s₀) (R s₀) + x22 : s.gpr .x22 = Dp s₀ + BitVec.ofNat 64 (16 * k) + x23 : s.gpr .x23 = BitVec.ofNat 64 (N s₀ - k) + sp : s.sp = s₀.sp + +theorem bodyMid_wp {s₀ : State} (hp : UPre s₀) {k : Nat} (hk : k < N s₀) {s : State} (h : LInv s₀ k s) : + WP isa (.block (chainIn ++ updArgs)) s (Mid s₀ k) := + WP.mono (bodyA_wp hp hk h) fun _ hb => + ⟨hb.pre, by rw [hb.saved .x22 (by simp [preserved]), h.x22], + by rw [hb.saved .x23 (by simp [preserved]), h.x23], by rw [hb.sp, h.sp]⟩ + +/-- What is known after the call. -/ +structure After (s₀ : State) (k : Nat) (s : State) : Prop where + x22 : s.gpr .x22 = Dp s₀ + BitVec.ofNat 64 (16 * k) + x23 : s.gpr .x23 = BitVec.ofNat 64 (N s₀ - k) + sp : s.sp = s₀.sp + +theorem body_ct (v : Ctr32Impl) {s₀ s₀' : State} (hp : UPre s₀) (hp' : UPre s₀') + (hq : updateAArch64.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 [.x19, .x20, .x21, .x22, .x23, .x24]) + (.block (chainIn ++ updArgs)) h).isSome = true := ⟨_, by taint_decide⟩ + obtain ⟨_, hB⟩ : ∃ h, (taint.check (Taint.ofRegs [.x22, .x23]) (.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 => 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.sp, h.2.sp, pub_sp 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 .x22 (by simp [preserved]) (by decide), h.1.x22], + by rw [hc.saved .x23 (by simp [preserved]) (by decide), h.1.x23], by rw [hc.sp, h.1.sp]⟩, + WP.mono (ctr_call v h.2.pre) fun _ hc => + ⟨by rw [hc.saved .x22 (by simp [preserved]) (by decide), h.2.x22], + by rw [hc.saved .x23 (by simp [preserved]) (by decide), h.2.x23], by rw [hc.sp, h.2.sp]⟩⟩ + have b := RelCT.taint (A := taint) (P := fun s₁ s₂ => After s₀ k s₁ ∧ After s₀' k s₂) _ + (fun s₁ s₂ h => agree_of (by rw [h.1.sp, h.2.sp, pub_sp hq]) 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.x22, h.2.x22, pub_Dp hq] + · rw [h.1.x23, h.2.x23, 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 : updateAArch64.pub s₀ s₀') (n : Nat) : + RelCT isa (LRel s₀ s₀' n) (.loop (body v.callee) (.nonzero .x .x23)) + 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₁ := LInv s₀ (k + 1)) (F₂ := LInv s₀' (k + 1)) + 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₁, l₂⟩ => ?_ + have hb : N s₀ < 2 ^ 64 := (s₀.gpr .x4).isLt + have e₁ := eval_x23 (x := N s₀ - (k + 1)) (by omega) l₁.x23 + have e₂ := eval_x23 (x := N s₀ - (k + 1)) (by omega) (by rw [l₂.x23, ← hN]) + refine ⟨by rw [e₁, e₂], fun hf => ?_, fun ht => ?_⟩ + · rw [e₁] at hf + have h0 : N s₀ = k + 1 := by + have : N s₀ - (k + 1) = 0 := by simpa using hf + omega + exact ⟨h0 ▸ l₁, by rw [← hN, h0]; exact l₂⟩ + · rw [e₁] at ht + have h0 : N s₀ - (k + 1) ≠ 0 := by simpa using ht + exact ⟨N s₀ - (k + 1), by omega, k + 1, rfl, by omega, l₁, l₂⟩ + +/-! ## The whole function -/ + +theorem update_rel (v : Ctr32Impl) {s₀ s₀' : State} (h0 : updateAArch64.pre s₀) + (h0' : updateAArch64.pre s₀') (hq : updateAArch64.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 + have hb : N s₀ < 2 ^ 64 := (s₀.gpr .x4).isLt + obtain ⟨_, hpro⟩ : ∃ h, (taint.check (Taint.ofRegs [.x0, .x1, .x2, .x3, .x4, .x5]) + (.block (save ++ setup)) h).isSome = true := ⟨_, by taint_decide⟩ + obtain ⟨_, hepi⟩ : ∃ h, (taint.check (Taint.ofRegs [.x24]) (.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 + obtain ⟨h1, h2, h3, h4, h5, h6, h7⟩ := hq + refine agree_of h7 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 <;> assumption) hpro).wp + (F₁ := LInv s₀ 0) (F₂ := LInv s₀' 0) + fun a b h => by obtain ⟨rfl, rfl⟩ := h; exact ⟨prologue_wp hp, prologue_wp hp'⟩ + have ev {k : Nat} {s : State} (h : LInv s₀ 0 s) : + isa.eval (.zero .x .x23) s = some (decide (N s₀ = 0)) := + eval_zero_x23 hb (by rw [h.x23]; rfl) + have ev' {s : State} (h : LInv s₀' 0 s) : isa.eval (.zero .x .x23) s = some (decide (N s₀ = 0)) := + eval_zero_x23 hb (by rw [h.x23, ← hN]; rfl) + have nil := RelCT.taint (A := taint) + (P := fun a b => (LInv s₀ 0 a ∧ LInv s₀' 0 b) ∧ isa.eval (.zero .x .x23) a = some true) _ + (fun a b h => agree_of (by rw [h.1.1.sp, h.1.2.sp, pub_sp hq]) fun r hr => by simp at hr) hnil + have mid : RelCT isa (fun a b => LInv s₀ 0 a ∧ LInv s₀' 0 b) + (.ite (.zero .x .x23) (.block []) (.loop (body v.callee) (.nonzero .x .x23))) + (fun a b => LInv s₀ (N s₀) a ∧ LInv s₀' (N s₀') b) := by + refine RelCT.ite (fun a b h => by rw [ev (k := 0) h.1, ev' h.2]) ?_ ?_ + · refine (nil.wp (F₁ := LInv s₀ (N s₀)) (F₂ := LInv s₀' (N s₀')) fun a b h => ?_).mono + (fun _ _ h => h) fun _ _ h => h.2 + have h0 : N s₀ = 0 := by + have := h.2; rw [ev (k := 0) h.1.1] at this; simpa using this + exact ⟨WP.block_nil (h0 ▸ h.1.1), WP.block_nil (by rw [← hN, h0]; exact h.1.2)⟩ + · refine (loop_ct v hp hp' hq (N s₀ - 0)).mono (fun a b h => ⟨0, rfl, ?_, h.1.1, h.1.2⟩) + fun _ _ h => h + have := h.2; rw [ev (k := 0) h.1.1] 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 => agree_of (by rw [h.1.sp, h.2.sp, pub_sp hq]) fun r hr => by + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + subst hr; rw [h.1.x24, h.2.x24, 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 updateAArch64.pre updateAArch64.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.AArch64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/AArch64/UpdateCorrect.lean b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/UpdateCorrect.lean new file mode 100644 index 000000000..e634f788c --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/UpdateCorrect.lean @@ -0,0 +1,184 @@ +import VerifiedGarbage.Proof.CmacAes.AArch64.UpdateLoop + +/-! +# AES-CMAC on AArch64: `vg_cmac_aes_update` is correct + +Untrusted: everything here is checked by Lean. +-/ + +namespace VG.Proof.CmacAes.AArch64 + +open VG VG.AArch64 VG.AArch64.RegUpd VG.Impl.CmacAes.AArch64 +open VG.Proof.Aes.AArch64 (Ctr32Impl) + +theorem ofNat_ne_zero {x : Nat} (hx : x < 2 ^ 64) : (BitVec.ofNat 64 x != 0) = !decide (x = 0) := by + have : (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 + rw [bne, this] + +theorem eval_x23 {s : State} {x : Nat} (hx : x < 2 ^ 64) (h : s.gpr .x23 = BitVec.ofNat 64 x) : + isa.eval (.nonzero .x .x23) s = some !decide (x = 0) := by + show some (s.read .x .x23 != 0) = _ + rw [State.read, h, BitVec.setWidth_eq, ofNat_ne_zero hx] + +theorem eval_zero_x23 {s : State} {x : Nat} (hx : x < 2 ^ 64) (h : s.gpr .x23 = BitVec.ofNat 64 x) : + isa.eval (.zero .x .x23) s = some (decide (x = 0)) := by + show some (s.read .x .x23 == 0) = _ + rw [State.read, h, BitVec.setWidth_eq] + have := ofNat_ne_zero hx + rw [bne] at this + cases hb : (BitVec.ofNat 64 x == 0) <;> rw [hb] at this <;> cases hd : decide (x = 0) <;> simp_all + +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) (.nonzero .x .x23)) s (LInv s₀ (N s₀)) := by + refine WP.loop (M := isa) (body := body v.callee) (c := .nonzero .x .x23) (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' => ?_ + have hN : N s₀ < 2 ^ 64 := (s₀.gpr .x4).isLt + have ev := eval_x23 (x := N s₀ - (k + 1)) (by omega) h'.x23 + by_cases hz : N s₀ - (k + 1) = 0 + · left + refine ⟨by rw [ev]; simp [hz], ?_⟩ + rwa [show N s₀ = k + 1 by omega] + · right + refine ⟨by rw [ev]; simp [hz], N s₀ - (k + 1), by omega, k + 1, rfl, by omega, h'⟩ + +/-! ## 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) + +/-- Each slot holds the register saved there. -/ +theorem savedMem_slot (s : State) {r : Reg} {d : Nat} (h : (r, d) ∈ saved) : + (savedMem s).readW (s.gpr .x5 + BitVec.ofNat 64 d) 64 = s.gpr r := by + simp only [saved, List.mem_cons, List.not_mem_nil, or_false, Prod.mk.injEq] at h + simp only [savedMem, saved, List.foldl] + rcases h with ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ <;> + repeat (first + | rw [Mem.readW_writeW_self64] + | rw [readW_writeW_other _ _ _ (by decide) (by decide) (by decide)]) + +theorem restore_ok (s : State) {B : Addr} (hb : s.gpr .x24 = B) + (hr : ∀ d, 2064 ≤ d → d + 8 ≤ 2120 → InRegions (s.rd ++ s.wr) (B + BitVec.ofNat 64 d) 8) : + ∃ s', runBlock isa restore s = some s' ∧ + (∀ r d, (r, d) ∈ saved → s'.gpr r = s.mem.readW (B + BitVec.ofNat 64 d) 64) ∧ + (∀ r, r ≠ .x19 → r ≠ .x20 → r ≠ .x21 → r ≠ .x22 → r ≠ .x23 → r ≠ .x30 → r ≠ .x24 → + s'.gpr r = s.gpr r) ∧ + s'.sp = s.sp ∧ s'.mem = s.mem := by + refine ⟨_, by + simp (config := {decide := true}) only [restore, saved, List.map, runBlock_cons, runStep_some, + runBlock_nil, exec, addr, State.load, Size.bytes, Size.bits, gpr_write, mem_write, rd_write, + wr_write, ite_true, ite_false, Option.bind_some, 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), + hr 2112 (by decide) (by decide)] + rfl, ?_⟩ + refine ⟨fun r d h => ?_, fun r h₁ h₂ h₃ h₄ h₅ h₆ h₇ => ?_, rfl, rfl⟩ + · simp only [saved, List.mem_cons, List.not_mem_nil, or_false, Prod.mk.injEq] at h + rcases h with ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ | ⟨rfl, rfl⟩ <;> + simp [gpr_write, Mem.readW] + · simp [gpr_write, h₁, h₂, h₃, h₄, h₅, h₆, h₇] + +/-! ## The whole function -/ + +theorem x4_ofNat (s₀ : State) : s₀.gpr .x4 = 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⟩], (⟨S s₀ + BitVec.ofNat 64 2064, 56⟩ : 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 + · exact hp.st_scr.symm.sub_left (UPre.scr_sub (by decide)) + · exact Offset.disjoint_base _ (by decide) (by have := hp.scr_wrap; omega) + +theorem slot_read {s₀ : State} (hp : UPre s₀) {m : Mem} + (hf : Frame [stR s₀, ⟨S s₀, 2064⟩] (savedMem s₀) m) {d : Nat} (h₁ : 2064 ≤ d) (h₂ : d + 8 ≤ 2120) : + 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, 56⟩) (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₁ := by + obtain ⟨s₁, run₁, x19₁, x20₁, x21₁, x22₁, x23₁, x24₁, keep₁, sp₁, 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 := + Proof.Cmac.bytesAt_frame16 (savedMem_frame s₀) fun r hr => by + simp only [List.mem_singleton] at hr; subst hr + exact hp.st_scr.sub_right (UPre.scr_sub (by decide)) + exact { x19 := x19₁, x20 := x20₁, x21 := x21₁ + x22 := by rw [x22₁]; simp + x23 := by rw [x23₁, x4_ofNat]; rfl + x24 := x24₁ + other := fun r _ h19 h20 h21 h22 h23 h24 _ => keep₁ r h19 h20 h21 h22 h23 h24 + sp := sp₁, rd := rd₁, wr := wr₁ + frame := by rw [mem₁]; exact Frame.refl _ _ + state := by rw [mem₁, stSaved]; rfl } + +theorem mid_wp (v : Ctr32Impl) {s₀ : State} (hp : UPre s₀) {s₁ : State} (h : LInv s₀ 0 s₁) : + WP isa (.ite (.zero .x .x23) (.block []) (.loop (body v.callee) (.nonzero .x .x23))) s₁ + (LInv s₀ (N s₀)) := by + have hN := (s₀.gpr .x4).isLt + have ev := eval_zero_x23 (x := N s₀) hN (by rw [h.x23]; rfl) + by_cases hn : N s₀ = 0 + · refine WP.ite true (by rw [ev]; simp [hn]) (fun _ => WP.block_nil ?_) (fun h => by cases h) + rw [hn]; exact h + · refine WP.ite false (by rw [ev]; simp [hn]) (fun h => by cases h) fun _ => ?_ + exact loop_ok 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' => GprAbi s₀ s' ∧ updateAArch64.post s₀ s' := by + have rdwr : s₂.rd ++ s₂.wr = s₀.rd ++ s₀.wr := by rw [h₂.rd, h₂.wr] + obtain ⟨s₃, run₃, slot₃, keep₃, sp₃, mem₃⟩ := + restore_ok s₂ h₂.x24 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 sl {r : Reg} {d : Nat} (h : (r, d) ∈ saved) : s₃.gpr r = s₀.gpr r := by + have hd : 2064 ≤ d ∧ d + 8 ≤ 2120 := by + simp only [saved, List.mem_cons, List.not_mem_nil, or_false, Prod.mk.injEq] at h + omega + rw [slot₃ r d h, slot_read hp h₂.frame hd.1 hd.2, savedMem_slot s₀ h] + refine ⟨⟨fun r hr => ?_, by rw [sp₃, h₂.sp]⟩, ?_⟩ + · simp only [preserved, List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl | rfl + · exact sl (d := 2064) (by simp [saved]) + · exact sl (d := 2072) (by simp [saved]) + · exact sl (d := 2080) (by simp [saved]) + · exact sl (d := 2088) (by simp [saved]) + · exact sl (d := 2096) (by simp [saved]) + · exact sl (d := 2112) (by simp [saved]) + · rw [keep₃ _ (by decide) (by decide) (by decide) (by decide) (by decide) (by decide) (by decide)] + exact h₂.other _ (by simp [preserved]) (by decide) (by decide) (by decide) (by decide) + (by decide) (by decide) (by decide) + · rw [keep₃ _ (by decide) (by decide) (by decide) (by decide) (by decide) (by decide) (by decide)] + exact h₂.other _ (by simp [preserved]) (by decide) (by decide) (by decide) (by decide) + (by decide) (by decide) (by decide) + · rw [keep₃ _ (by decide) (by decide) (by decide) (by decide) (by decide) (by decide) (by decide)] + exact h₂.other _ (by simp [preserved]) (by decide) (by decide) (by decide) (by decide) + (by decide) (by decide) (by decide) + · rw [keep₃ _ (by decide) (by decide) (by decide) (by decide) (by decide) (by decide) (by decide)] + exact h₂.other _ (by simp [preserved]) (by decide) (by decide) (by decide) (by decide) + (by decide) (by decide) (by decide) + · exact sl (d := 2104) (by simp [saved]) + · 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 : updateAArch64.pre s₀) : + WP isa (update v.callee) s₀ fun s' => GprAbi s₀ s' ∧ updateAArch64.post s₀ s' := by + have hp := UPre.of h0 + exact WP.seq (WP.mono (prologue_wp hp) fun s₁ h₁ => + WP.seq (WP.mono (mid_wp v hp h₁) fun _ h₂ => epilogue_wp hp h₂)) + +end VG.Proof.CmacAes.AArch64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/AArch64/UpdateLoop.lean b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/UpdateLoop.lean new file mode 100644 index 000000000..bf18e7230 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/UpdateLoop.lean @@ -0,0 +1,306 @@ +import VerifiedGarbage.Proof.CmacAes.AArch64.Update + +/-! +# AES-CMAC on AArch64: 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 (`x22` the next block, +`x23` the blocks left), only the state and the first 2064 bytes of the +scratch buffer have changed since the registers were saved, and the state is +the chaining value after the first `k` blocks. +-/ + +namespace VG.Proof.CmacAes.AArch64 + +open VG VG.AArch64 VG.AArch64.RegUpd VG.Impl.CmacAes.AArch64 +open VG.Proof.Aes.AArch64 (Ctr32Impl) + +section +variable (s₀ : State) + +abbrev W : Addr := s₀.gpr .x0 +abbrev R : Nat := (s₀.gpr .x1).toNat +abbrev St : Addr := s₀.gpr .x2 +abbrev Dp : Addr := s₀.gpr .x3 +abbrev N : Nat := (s₀.gpr .x4).toNat +abbrev S : Addr := s₀.gpr .x5 + +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⟩ + +/-- 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₀) + 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 : updateAArch64.pre s₀) : UPre s₀ := + let ⟨a, b, c, d, e, f, g, h, i, j, k⟩ := h + ⟨a, b, c, d, e, f, g, h, i, j, k⟩ + +/-- The loop invariant, after `k` blocks. -/ +structure LInv (s₀ : State) (k : Nat) (s : State) : Prop where + x19 : s.gpr .x19 = W s₀ + x20 : s.gpr .x20 = s₀.gpr .x1 + x21 : s.gpr .x21 = St s₀ + x22 : s.gpr .x22 = Dp s₀ + BitVec.ofNat 64 (16 * k) + x23 : s.gpr .x23 = BitVec.ofNat 64 (N s₀ - k) + x24 : s.gpr .x24 = S s₀ + other : ∀ r ∈ preserved, r ≠ .x19 → r ≠ .x20 → r ≠ .x21 → r ≠ .x22 → r ≠ .x23 → r ≠ .x24 → + r ≠ .x30 → s.gpr r = s₀.gpr r + sp : s.sp = s₀.sp + rd : s.rd = s₀.rd + wr : s.wr = s₀.wr + frame : Frame [stR s₀, ⟨S s₀, 2064⟩] (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) + +end + +theorem slot_contains (b : Addr) {d : Nat} (h₁ : 2064 ≤ d) (h₂ : d + 8 ≤ 2120) : + (⟨b + BitVec.ofNat 64 2064, 56⟩ : 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 .x5 + BitVec.ofNat 64 2064, 56⟩] 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))).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 .x22 = s.gpr .x22 + BitVec.ofNat 64 16 ∧ s'.gpr .x23 = s.gpr .x23 - 1 ∧ + (∀ r, r ≠ .x22 → r ≠ .x23 → s'.gpr r = s.gpr r) ∧ s'.sp = s.sp ∧ + s'.mem = s.mem ∧ s'.rd = s.rd ∧ s'.wr = s.wr := by + refine ⟨_, by + rw [advance, runBlock_cons, exec_addImm_x (by decide), runStep_some, runBlock_cons, + exec_subImm_x (by decide), runStep_some, runBlock_nil], ?_⟩ + refine ⟨?_, ?_, fun r h₁ h₂ => ?_, rfl, rfl, rfl, rfl⟩ + · simp [gpr_write, State.read] + · simp [gpr_write, State.read] + · simp [gpr_write, h₁, h₂] + +theorem x1_ofNat (s₀ : State) : s₀.gpr .x1 = 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₀] + +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 Proof.Cmac.bytesAt_frame hf (fun r hr => ?_) (by omega) + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl + · exact hp.sch_st.sub_left (Region.sub_prefix hR) + · exact hp.sch_scr.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 Proof.Cmac.bytesAt_frame hf (fun r hr => ?_) (by decide) + simp only [List.mem_cons, List.not_mem_nil, or_false] at hr + rcases hr with rfl | rfl + · exact hp.data_st.sub_left (UPre.data_sub hk) + · exact hp.data_scr.sub_left (UPre.data_sub hk) + +omit hp in +theorem UPre.big_of {m : Mem} (hf : Frame [stR s₀, ⟨S s₀, 2064⟩] (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 + · exact ⟨stR s₀, by simp, fun _ h => h⟩ + · exact ⟨scrR s₀, by simp, Region.sub_prefix (by decide)⟩) + +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 ∈ preserved, s₁.gpr r = s.gpr r + sp : s₁.sp = s.sp + mem : s₁.mem = Proof.Cmac.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₁, x0₁, x1₁, x2₁, x3₁, x4₁, x5₁, cs₁, sp₁, 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.x24]) h.x21 h.x22 + (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 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₀) := + { x0 := by rw [x0₁, h.x19] + x1 := by rw [x1₁, h.x20, x1_ofNat] + x2 := x2₁ + x3 := x3₁ + x4 := x4₁ + x5 := by rw [x5₁, h.x24] + 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)) + 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 Proof.Cmac.chainMem_state _ _ _ _ } + exact ⟨pre, cs₁, sp₁, 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' := 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₁, sp₁, mem₁, rd₁, wr₁⟩ => ?_) + refine WP.seq (WP.mono (ctr_call v pre) fun s₂ h₂ => ?_) + obtain ⟨s₃, run₃, x22₃, x23₃, keep₃, sp₃, mem₃, rd₃, wr₃⟩ := advance_ok s₂ + refine WP.of_runBlock ⟨s₃, run₃, ?_⟩ + have g (r : Reg) (hr : r ∈ preserved) (h30 : r ≠ .x30) (h22 : r ≠ .x22) (h23 : r ≠ .x23) : + s₃.gpr r = s.gpr r := by + rw [keep₃ r h22 h23, h₂.saved r hr h30, cs₁ r hr] + have x22₂ : s₂.gpr .x22 = Dp s₀ + BitVec.ofNat 64 (16 * k) := by + rw [h₂.saved .x22 (by simp [preserved]) (by decide), cs₁ .x22 (by simp [preserved]), h.x22] + have x23₂ : s₂.gpr .x23 = BitVec.ofNat 64 (N s₀ - k) := by + rw [h₂.saved .x23 (by simp [preserved]) (by decide), cs₁ .x23 (by simp [preserved]), h.x23] + have hN := (s₀.gpr .x4).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 Proof.Cmac.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₁, Proof.Cmac.chainMem_counter _ cst cq, h.state, + UPre.block_bytes hp bigS hk] at out + refine ⟨?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_⟩ + · rw [g .x19 (by simp [preserved]) (by decide) (by decide) (by decide), h.x19] + · rw [g .x20 (by simp [preserved]) (by decide) (by decide) (by decide), h.x20] + · rw [g .x21 (by simp [preserved]) (by decide) (by decide) (by decide), h.x21] + · rw [x22₃, x22₂, Offset.add_add_eq _ (c := 16 * (k + 1)) (by omega)] + · rw [x23₃, x23₂, dec] + · rw [g .x24 (by simp [preserved]) (by decide) (by decide) (by decide), h.x24] + · intro r hr h19 h20 h21 h22 h23 h24 h30 + rw [g r hr h30 h22 h23, h.other r hr h19 h20 h21 h22 h23 h24 h30] + · rw [sp₃, h₂.sp, sp₁, h.sp] + · 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 + · 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)⟩ + · rw [mem₃, out, take_succ_blks s₀ hk, Proof.Cmac.chain_append, Proof.Cmac.chain_single] + +end VG.Proof.CmacAes.AArch64 diff --git a/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Verified.lean b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Verified.lean new file mode 100644 index 000000000..5f82b6342 --- /dev/null +++ b/lean/VerifiedGarbage/Proof/CmacAes/AArch64/Verified.lean @@ -0,0 +1,88 @@ +import VerifiedGarbage.Proof.CmacAes.AArch64.UpdateCT +import VerifiedGarbage.Proof.CmacAes.AArch64.SubkeysCT +import VerifiedGarbage.Proof.CmacAes.AArch64.FinalizeCT +import VerifiedGarbage.Proof.Framework.Contract +import VerifiedGarbage.Spec.Cmac.Contract + +/-! +# AES-CMAC on AArch64: `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 no +stack: the calls keep the return address in `x30`, which each function saves +in the scratch buffer). +-/ + +namespace VG.Proof.CmacAes.AArch64 + +open VG VG.AArch64 VG.Impl.CmacAes.AArch64 +open VG.Proof.Aes.AArch64 (Ctr32Impl) + +theorem update_keepsV (v : Ctr32Impl) : (update v.callee).allInstrs keepsV = true := by + simp only [update, body, Code.allInstrs, v.keepsV]; decide +kernel + +theorem subkeys_keepsV (v : Ctr32Impl) : (subkeys v.callee).allInstrs keepsV = true := by + simp only [subkeys, Code.allInstrs, v.keepsV]; decide +kernel + +theorem finalize_keepsV (v : Ctr32Impl) : (finalize v.callee).allInstrs keepsV = true := by + simp only [finalize, finPre, partialBlock, copy, Code.allInstrs, v.keepsV]; decide +kernel + +theorem update_correct (v : Ctr32Impl) (s : State) (hs : updateAArch64.pre s) : + ∃ t s', Exec isa (update v.callee) s t s' ∧ abiPreserved s s' ∧ updateAArch64.post s s' := + WP.withPreservedV (update_wp v hs) (update_keepsV v) + +theorem subkeys_correct (v : Ctr32Impl) (s : State) (hs : subkeysAArch64.pre s) : + ∃ t s', Exec isa (subkeys v.callee) s t s' ∧ abiPreserved s s' ∧ subkeysAArch64.post s s' := + WP.withPreservedV (subkeys_wp v hs) (subkeys_keepsV v) + +theorem finalize_correct (v : Ctr32Impl) (s : State) (hs : finalizeAArch64.pre s) : + ∃ t s', Exec isa (finalize v.callee) s t s' ∧ abiPreserved s s' ∧ finalizeAArch64.post s s' := + WP.withPreservedV (finalize_wp v hs) (finalize_keepsV v) + +/-- A state satisfying `vg_cmac_aes_update`'s precondition (with no blocks). -/ +def updSat : State where + gpr r := match r with + | .x0 => 0x1000 | .x1 => 10 | .x2 => 0x2000 | .x3 => 0x3000 | .x5 => 0x4000 | _ => 0 + sp := 0x8000 + mem _ := 0 + rd := [⟨0x1000, 240⟩, ⟨0x3000, 0⟩] + wr := [⟨0x2000, 16⟩, ⟨0x4000, 2176⟩] + +theorem update_verified (v : Ctr32Impl) : + Verified AArch64.target (update v.callee) (Spec.Cmac.aesUpdateContract AArch64.abi) := + Verified.of_correct (update_correct v) (update_ct v) (by + sig_implies [Spec.Cmac.aesUpdateContract, Spec.Cmac.aesUpdateSig, updateAArch64, AArch64.abi, + AArch64.argRegs] [updSat] using updSat) + +/-- A state satisfying `vg_cmac_aes_subkeys`'s precondition. -/ +def subSat : State where + gpr r := match r with + | .x0 => 0x1000 | .x1 => 10 | .x2 => 0x2000 | .x3 => 0x4000 | _ => 0 + sp := 0x8000 + mem _ := 0 + rd := [⟨0x1000, 240⟩] + wr := [⟨0x2000, 32⟩, ⟨0x4000, 2176⟩] + +theorem subkeys_verified (v : Ctr32Impl) : + Verified AArch64.target (subkeys v.callee) (Spec.Cmac.aesSubkeysContract AArch64.abi) := + Verified.of_correct (subkeys_correct v) (subkeys_ct v) (by + sig_implies [Spec.Cmac.aesSubkeysContract, Spec.Cmac.aesSubkeysSig, subkeysAArch64, AArch64.abi, + AArch64.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 + | .x0 => 0x1000 | .x1 => 10 | .x2 => 0x2000 | .x3 => 0x3000 | .x5 => 0x4000 | _ => 0 + sp := 0x8000 + mem _ := 0 + rd := [⟨0x1000, 272⟩, ⟨0x3000, 0⟩] + wr := [⟨0x2000, 16⟩, ⟨0x4000, 2176⟩] + +theorem finalize_verified (v : Ctr32Impl) : + Verified AArch64.target (finalize v.callee) (Spec.Cmac.aesFinalizeContract AArch64.abi) := + Verified.of_correct (finalize_correct v) (finalize_ct v) (by + sig_implies [Spec.Cmac.aesFinalizeContract, Spec.Cmac.aesFinalizeSig, finalizeAArch64, AArch64.abi, + AArch64.argRegs] [finSat] using finSat) + +end VG.Proof.CmacAes.AArch64 diff --git a/lean/VerifiedGarbage/Variants/AesCtr32/AArch64/Aese.lean b/lean/VerifiedGarbage/Variants/AesCtr32/AArch64/Aese.lean new file mode 100644 index 000000000..d7f7bd886 --- /dev/null +++ b/lean/VerifiedGarbage/Variants/AesCtr32/AArch64/Aese.lean @@ -0,0 +1,14 @@ +import VerifiedGarbage.Proof.Aes.AArch64.Variant + +/-! +# `vg_aes_ctr32` on AArch64: with the Cryptographic Extension + +A variant of `AesCtr32` on AArch64 (see `TCB/Emit.lean`): `vg_aes_ctr32_aes`, +which needs the AES instructions (`FEAT_AES`). +-/ + +namespace VG.Variants.AesCtr32.AArch64.Aese + +def variant : Proof.Aes.AArch64.Ctr32Impl := .aese + +end VG.Variants.AesCtr32.AArch64.Aese diff --git a/lean/VerifiedGarbage/Variants/AesCtr32/AArch64/Scalar.lean b/lean/VerifiedGarbage/Variants/AesCtr32/AArch64/Scalar.lean new file mode 100644 index 000000000..82ce87b95 --- /dev/null +++ b/lean/VerifiedGarbage/Variants/AesCtr32/AArch64/Scalar.lean @@ -0,0 +1,14 @@ +import VerifiedGarbage.Proof.Aes.AArch64.Variant + +/-! +# `vg_aes_ctr32` on AArch64: the bitsliced implementation + +A variant of `AesCtr32` on AArch64 (see `TCB/Emit.lean`): `vg_aes_ctr32`, in +the baseline ISA. +-/ + +namespace VG.Variants.AesCtr32.AArch64.Scalar + +def variant : Proof.Aes.AArch64.Ctr32Impl := .scalar + +end VG.Variants.AesCtr32.AArch64.Scalar diff --git a/src/asm/aarch64/cmac_aes.rs b/src/asm/aarch64/cmac_aes.rs new file mode 100644 index 000000000..e026c5f8b --- /dev/null +++ b/src/asm/aarch64/cmac_aes.rs @@ -0,0 +1,485 @@ +// @generated from lean/VerifiedGarbage/Artifacts.lean by lean/Emit.lean. DO NOT EDIT. +//! Verified `cmac_aes` functions for `aarch64`. +#![allow(dead_code)] + +/// The CPU features `vg_cmac_aes_subkeys_aes` requires (`Artifact.features`). +pub(crate) const VG_CMAC_AES_SUBKEYS_AES_FEATURES: &[&str] = &["aes"]; + +/// 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_aes`. +/// +/// # 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 wrap around the end of the address space (no Rust object does). +/// * The CPU must support the `aes` target feature. +#[unsafe(naked)] +pub(crate) unsafe extern "C" fn vg_cmac_aes_subkeys_aes(schedule: *const [u8; 240], rounds: usize, subkeys: *mut [u8; 32], scratch: *mut [u64; 272]) { + core::arch::naked_asm!( + ".arch_extension aes", + "str x19, [x3, #2064]", + "str x20, [x3, #2072]", + "str x30, [x3, #2080]", + "add x19, x2, #0", + "add x20, x3, #0", + "movz x9, #0, lsl #0", + "str x9, [x3, #2048]", + "str x9, [x3, #2056]", + "str x9, [x2, #0]", + "str x9, [x2, #8]", + "add x2, x20, #2048", + "add x3, x19, #0", + "movz x4, #1, lsl #0", + "add x5, x20, #0", + "bl {vg_aes_ctr32_aes}", + "ldr x9, [x19, #0]", + "ldr x10, [x19, #8]", + "rev x9, x9", + "rev x10, x10", + "lsr x11, x9, #63", + "movz x12, #0, lsl #0", + "sub x11, x12, x11", + "movz x12, #135, lsl #0", + "and x11, x11, x12", + "lsr x12, x10, #63", + "lsl x9, x9, #1", + "orr x9, x9, x12", + "lsl x10, x10, #1", + "eor x10, x10, x11", + "rev x9, x9", + "rev x10, x10", + "str x9, [x19, #0]", + "str x10, [x19, #8]", + "ldr x9, [x19, #0]", + "ldr x10, [x19, #8]", + "rev x9, x9", + "rev x10, x10", + "lsr x11, x9, #63", + "movz x12, #0, lsl #0", + "sub x11, x12, x11", + "movz x12, #135, lsl #0", + "and x11, x11, x12", + "lsr x12, x10, #63", + "lsl x9, x9, #1", + "orr x9, x9, x12", + "lsl x10, x10, #1", + "eor x10, x10, x11", + "rev x9, x9", + "rev x10, x10", + "str x9, [x19, #16]", + "str x10, [x19, #24]", + "ldr x30, [x20, #2080]", + "ldr x19, [x20, #2064]", + "ldr x20, [x20, #2072]", + "ret", + ".arch_extension noaes", + vg_aes_ctr32_aes = sym super::aes::vg_aes_ctr32_aes, + ) +} + +/// The CPU features `vg_cmac_aes_update_aes` requires (`Artifact.features`). +pub(crate) const VG_CMAC_AES_UPDATE_AES_FEATURES: &[&str] = &["aes"]; + +/// 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_aes`. +/// +/// # 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 wrap around the end of the address space (no Rust object does). +/// * The CPU must support the `aes` target feature. +#[unsafe(naked)] +pub(crate) unsafe extern "C" fn vg_cmac_aes_update_aes(schedule: *const [u8; 240], rounds: usize, state: *mut [u8; 16], data: *const [u8; 16], n: usize, scratch: *mut [u64; 272]) { + core::arch::naked_asm!( + ".arch_extension aes", + "str x19, [x5, #2064]", + "str x20, [x5, #2072]", + "str x21, [x5, #2080]", + "str x22, [x5, #2088]", + "str x23, [x5, #2096]", + "str x30, [x5, #2104]", + "str x24, [x5, #2112]", + "add x19, x0, #0", + "add x20, x1, #0", + "add x21, x2, #0", + "add x22, x3, #0", + "add x23, x4, #0", + "add x24, x5, #0", + "cbz x23, 20f", + "22:", + "ldr x9, [x21, #0]", + "ldr x10, [x22, #0]", + "eor x9, x9, x10", + "str x9, [x24, #2048]", + "ldr x9, [x21, #8]", + "ldr x10, [x22, #8]", + "eor x9, x9, x10", + "str x9, [x24, #2056]", + "movz x9, #0, lsl #0", + "str x9, [x21, #0]", + "str x9, [x21, #8]", + "add x0, x19, #0", + "add x1, x20, #0", + "add x2, x24, #2048", + "add x3, x21, #0", + "movz x4, #1, lsl #0", + "add x5, x24, #0", + "bl {vg_aes_ctr32_aes}", + "add x22, x22, #16", + "sub x23, x23, #1", + "cbnz x23, 22b", + "b 21f", + "20:", + "21:", + "ldr x19, [x24, #2064]", + "ldr x20, [x24, #2072]", + "ldr x21, [x24, #2080]", + "ldr x22, [x24, #2088]", + "ldr x23, [x24, #2096]", + "ldr x30, [x24, #2104]", + "ldr x24, [x24, #2112]", + "ret", + ".arch_extension noaes", + vg_aes_ctr32_aes = sym super::aes::vg_aes_ctr32_aes, + ) +} + +/// The CPU features `vg_cmac_aes_finalize_aes` requires (`Artifact.features`). +pub(crate) const VG_CMAC_AES_FINALIZE_AES_FEATURES: &[&str] = &["aes"]; + +/// 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_aes`. +/// +/// # 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 wrap around the end of the address space (no Rust object does). +/// * The CPU must support the `aes` target feature. +#[unsafe(naked)] +pub(crate) unsafe extern "C" fn vg_cmac_aes_finalize_aes(key: *const [u8; 272], rounds: usize, state: *mut [u8; 16], last: *const u8, last_len: usize, scratch: *mut [u64; 272]) { + core::arch::naked_asm!( + ".arch_extension aes", + "sub x9, x4, #16", + "cbz x9, 20f", + "movz x9, #0, lsl #0", + "str x9, [x5, #2048]", + "str x9, [x5, #2056]", + "add x6, x5, #2048", + "add x7, x3, #0", + "add x8, x4, #0", + "cbz x4, 22f", + "24:", + "ldrb w9, [x7, #0]", + "strb w9, [x6, #0]", + "add x7, x7, #1", + "add x6, x6, #1", + "sub x8, x8, #1", + "cbnz x8, 24b", + "b 23f", + "22:", + "23:", + "movz x9, #128, lsl #0", + "strb w9, [x6, #0]", + "ldr x9, [x5, #2048]", + "ldr x10, [x0, #256]", + "eor x9, x9, x10", + "str x9, [x5, #2048]", + "ldr x9, [x5, #2056]", + "ldr x10, [x0, #264]", + "eor x9, x9, x10", + "str x9, [x5, #2056]", + "b 21f", + "20:", + "ldr x9, [x3, #0]", + "ldr x10, [x0, #240]", + "eor x9, x9, x10", + "str x9, [x5, #2048]", + "ldr x9, [x3, #8]", + "ldr x10, [x0, #248]", + "eor x9, x9, x10", + "str x9, [x5, #2056]", + "21:", + "ldr x9, [x5, #2048]", + "ldr x10, [x2, #0]", + "eor x9, x9, x10", + "str x9, [x5, #2048]", + "ldr x9, [x5, #2056]", + "ldr x10, [x2, #8]", + "eor x9, x9, x10", + "str x9, [x5, #2056]", + "movz x9, #0, lsl #0", + "str x9, [x2, #0]", + "str x9, [x2, #8]", + "str x19, [x5, #2064]", + "str x30, [x5, #2072]", + "add x19, x5, #0", + "add x3, x2, #0", + "add x2, x5, #2048", + "movz x4, #1, lsl #0", + "bl {vg_aes_ctr32_aes}", + "ldr x30, [x19, #2072]", + "ldr x19, [x19, #2064]", + "ret", + ".arch_extension noaes", + vg_aes_ctr32_aes = sym super::aes::vg_aes_ctr32_aes, + ) +} + +/// 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 wrap around the end of the address space (no Rust object does). +#[unsafe(naked)] +pub(crate) unsafe extern "C" fn vg_cmac_aes_subkeys(schedule: *const [u8; 240], rounds: usize, subkeys: *mut [u8; 32], scratch: *mut [u64; 272]) { + core::arch::naked_asm!( + "str x19, [x3, #2064]", + "str x20, [x3, #2072]", + "str x30, [x3, #2080]", + "add x19, x2, #0", + "add x20, x3, #0", + "movz x9, #0, lsl #0", + "str x9, [x3, #2048]", + "str x9, [x3, #2056]", + "str x9, [x2, #0]", + "str x9, [x2, #8]", + "add x2, x20, #2048", + "add x3, x19, #0", + "movz x4, #1, lsl #0", + "add x5, x20, #0", + "bl {vg_aes_ctr32}", + "ldr x9, [x19, #0]", + "ldr x10, [x19, #8]", + "rev x9, x9", + "rev x10, x10", + "lsr x11, x9, #63", + "movz x12, #0, lsl #0", + "sub x11, x12, x11", + "movz x12, #135, lsl #0", + "and x11, x11, x12", + "lsr x12, x10, #63", + "lsl x9, x9, #1", + "orr x9, x9, x12", + "lsl x10, x10, #1", + "eor x10, x10, x11", + "rev x9, x9", + "rev x10, x10", + "str x9, [x19, #0]", + "str x10, [x19, #8]", + "ldr x9, [x19, #0]", + "ldr x10, [x19, #8]", + "rev x9, x9", + "rev x10, x10", + "lsr x11, x9, #63", + "movz x12, #0, lsl #0", + "sub x11, x12, x11", + "movz x12, #135, lsl #0", + "and x11, x11, x12", + "lsr x12, x10, #63", + "lsl x9, x9, #1", + "orr x9, x9, x12", + "lsl x10, x10, #1", + "eor x10, x10, x11", + "rev x9, x9", + "rev x10, x10", + "str x9, [x19, #16]", + "str x10, [x19, #24]", + "ldr x30, [x20, #2080]", + "ldr x19, [x20, #2064]", + "ldr x20, [x20, #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 wrap around the end of the address space (no Rust object does). +#[unsafe(naked)] +pub(crate) unsafe extern "C" fn vg_cmac_aes_update(schedule: *const [u8; 240], rounds: usize, state: *mut [u8; 16], data: *const [u8; 16], n: usize, scratch: *mut [u64; 272]) { + core::arch::naked_asm!( + "str x19, [x5, #2064]", + "str x20, [x5, #2072]", + "str x21, [x5, #2080]", + "str x22, [x5, #2088]", + "str x23, [x5, #2096]", + "str x30, [x5, #2104]", + "str x24, [x5, #2112]", + "add x19, x0, #0", + "add x20, x1, #0", + "add x21, x2, #0", + "add x22, x3, #0", + "add x23, x4, #0", + "add x24, x5, #0", + "cbz x23, 20f", + "22:", + "ldr x9, [x21, #0]", + "ldr x10, [x22, #0]", + "eor x9, x9, x10", + "str x9, [x24, #2048]", + "ldr x9, [x21, #8]", + "ldr x10, [x22, #8]", + "eor x9, x9, x10", + "str x9, [x24, #2056]", + "movz x9, #0, lsl #0", + "str x9, [x21, #0]", + "str x9, [x21, #8]", + "add x0, x19, #0", + "add x1, x20, #0", + "add x2, x24, #2048", + "add x3, x21, #0", + "movz x4, #1, lsl #0", + "add x5, x24, #0", + "bl {vg_aes_ctr32}", + "add x22, x22, #16", + "sub x23, x23, #1", + "cbnz x23, 22b", + "b 21f", + "20:", + "21:", + "ldr x19, [x24, #2064]", + "ldr x20, [x24, #2072]", + "ldr x21, [x24, #2080]", + "ldr x22, [x24, #2088]", + "ldr x23, [x24, #2096]", + "ldr x30, [x24, #2104]", + "ldr x24, [x24, #2112]", + "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 wrap around the end of the address space (no Rust object does). +#[unsafe(naked)] +pub(crate) unsafe extern "C" fn vg_cmac_aes_finalize(key: *const [u8; 272], rounds: usize, state: *mut [u8; 16], last: *const u8, last_len: usize, scratch: *mut [u64; 272]) { + core::arch::naked_asm!( + "sub x9, x4, #16", + "cbz x9, 20f", + "movz x9, #0, lsl #0", + "str x9, [x5, #2048]", + "str x9, [x5, #2056]", + "add x6, x5, #2048", + "add x7, x3, #0", + "add x8, x4, #0", + "cbz x4, 22f", + "24:", + "ldrb w9, [x7, #0]", + "strb w9, [x6, #0]", + "add x7, x7, #1", + "add x6, x6, #1", + "sub x8, x8, #1", + "cbnz x8, 24b", + "b 23f", + "22:", + "23:", + "movz x9, #128, lsl #0", + "strb w9, [x6, #0]", + "ldr x9, [x5, #2048]", + "ldr x10, [x0, #256]", + "eor x9, x9, x10", + "str x9, [x5, #2048]", + "ldr x9, [x5, #2056]", + "ldr x10, [x0, #264]", + "eor x9, x9, x10", + "str x9, [x5, #2056]", + "b 21f", + "20:", + "ldr x9, [x3, #0]", + "ldr x10, [x0, #240]", + "eor x9, x9, x10", + "str x9, [x5, #2048]", + "ldr x9, [x3, #8]", + "ldr x10, [x0, #248]", + "eor x9, x9, x10", + "str x9, [x5, #2056]", + "21:", + "ldr x9, [x5, #2048]", + "ldr x10, [x2, #0]", + "eor x9, x9, x10", + "str x9, [x5, #2048]", + "ldr x9, [x5, #2056]", + "ldr x10, [x2, #8]", + "eor x9, x9, x10", + "str x9, [x5, #2056]", + "movz x9, #0, lsl #0", + "str x9, [x2, #0]", + "str x9, [x2, #8]", + "str x19, [x5, #2064]", + "str x30, [x5, #2072]", + "add x19, x5, #0", + "add x3, x2, #0", + "add x2, x5, #2048", + "movz x4, #1, lsl #0", + "bl {vg_aes_ctr32}", + "ldr x30, [x19, #2072]", + "ldr x19, [x19, #2064]", + "ret", + vg_aes_ctr32 = sym super::aes::vg_aes_ctr32, + ) +} diff --git a/src/asm/aarch64/mod.rs b/src/asm/aarch64/mod.rs index e410818ac..e8f82190d 100644 --- a/src/asm/aarch64/mod.rs +++ b/src/asm/aarch64/mod.rs @@ -16,6 +16,9 @@ pub(crate) mod chacha20; #[rustfmt::skip] pub(crate) mod chacha20poly1305; +#[rustfmt::skip] +pub(crate) mod cmac_aes; + #[rustfmt::skip] pub(crate) mod ct; diff --git a/src/cmac/aes.rs b/src/cmac/aes.rs index ce29b48ab..2c9cd0f2d 100644 --- a/src/cmac/aes.rs +++ b/src/cmac/aes.rs @@ -13,14 +13,24 @@ //! 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. +//! `vg_aes_ctr32` to encrypt each block. On AArch64, CPUs with the AES +//! extension run `vg_aes_expand_key_aes` and the `_aes` CMAC functions, +//! calling `vg_aes_ctr32_aes`. -#![cfg(target_arch = "x86_64")] +#![cfg(any(target_arch = "x86_64", target_arch = "aarch64"))] use super::{InvalidKeyLength, InvalidMac}; use crate::arch::aes::vg_aes_expand_key; +#[cfg(target_arch = "aarch64")] +use crate::arch::aes::{VG_AES_EXPAND_KEY_AES_FEATURES, vg_aes_expand_key_aes}; #[cfg(target_arch = "x86_64")] use crate::arch::aes::{VG_AES_EXPAND_KEY_AESNI_FEATURES, vg_aes_expand_key_aesni}; +#[cfg(target_arch = "aarch64")] +use crate::arch::cmac_aes::{ + VG_CMAC_AES_FINALIZE_AES_FEATURES, VG_CMAC_AES_SUBKEYS_AES_FEATURES, + VG_CMAC_AES_UPDATE_AES_FEATURES, vg_cmac_aes_finalize_aes, vg_cmac_aes_subkeys_aes, + vg_cmac_aes_update_aes, +}; #[cfg(target_arch = "x86_64")] use crate::arch::cmac_aes::{ VG_CMAC_AES_FINALIZE_AESNI_FEATURES, VG_CMAC_AES_SUBKEYS_AESNI_FEATURES, @@ -46,6 +56,9 @@ enum Backend { /// AES-NI. #[cfg(target_arch = "x86_64")] AesNi, + /// The AES extension. + #[cfg(target_arch = "aarch64")] + ArmCrypto, } impl Backend { @@ -63,6 +76,21 @@ impl Backend { Backend::Scalar } } + + /// The best implementation a CPU with the features `f` can run. + #[cfg(target_arch = "aarch64")] + fn select(f: Features) -> Backend { + if f.contains(Features::all(&[ + VG_AES_EXPAND_KEY_AES_FEATURES, + VG_CMAC_AES_SUBKEYS_AES_FEATURES, + VG_CMAC_AES_UPDATE_AES_FEATURES, + VG_CMAC_AES_FINALIZE_AES_FEATURES, + ])) { + Backend::ArmCrypto + } else { + Backend::Scalar + } + } } /// An incremental AES-CMAC computation. @@ -140,6 +168,11 @@ impl AesCmac { vg_aes_expand_key_aesni(k, key.len(), schedule, e); vg_cmac_aes_subkeys_aesni(schedule, rounds, subkeys, s); } + #[cfg(target_arch = "aarch64")] + Backend::ArmCrypto => { + vg_aes_expand_key_aes(k, key.len(), schedule, e); + vg_cmac_aes_subkeys_aes(schedule, rounds, subkeys, s); + } } } Ok(c) @@ -155,6 +188,8 @@ impl AesCmac { Backend::Scalar => vg_cmac_aes_update, #[cfg(target_arch = "x86_64")] Backend::AesNi => vg_cmac_aes_update_aesni, + #[cfg(target_arch = "aarch64")] + Backend::ArmCrypto => vg_cmac_aes_update_aes, }; let schedule = self.key.first_chunk::<240>().unwrap(); // SAFETY: `schedule` holds the key schedule for `self.rounds` (10, @@ -210,6 +245,8 @@ impl AesCmac { Backend::Scalar => vg_cmac_aes_finalize, #[cfg(target_arch = "x86_64")] Backend::AesNi => vg_cmac_aes_finalize_aesni, + #[cfg(target_arch = "aarch64")] + Backend::ArmCrypto => vg_cmac_aes_finalize_aes, }; // 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 @@ -293,6 +330,7 @@ mod tests { /// Each implementation is selected exactly when the CPU has its /// features. + #[cfg(target_arch = "x86_64")] #[test] fn select() { assert_eq!(Backend::select(Features::of(&[])), Backend::Scalar); @@ -302,4 +340,13 @@ mod tests { ); assert_eq!(Backend::select(Features::of(&["aes"])), Backend::Scalar); } + + /// Each implementation is selected exactly when the CPU has its + /// features. + #[cfg(target_arch = "aarch64")] + #[test] + fn select() { + assert_eq!(Backend::select(Features::of(&[])), Backend::Scalar); + assert_eq!(Backend::select(Features::of(&["aes"])), Backend::ArmCrypto); + } } diff --git a/tests/cavp/cmac_aes.rs b/tests/cavp/cmac_aes.rs index 1ae0ddf5e..490db371b 100644 --- a/tests/cavp/cmac_aes.rs +++ b/tests/cavp/cmac_aes.rs @@ -1,7 +1,7 @@ //! AES-CMAC: every vector of the CMAC generation and verification files, //! for 128-, 192- and 256-bit keys (the MAC truncated to `Tlen` bytes). -#![cfg(target_arch = "x86_64")] +#![cfg(any(target_arch = "x86_64", target_arch = "aarch64"))] use verified_garbage::cmac::aes::AesCmac; diff --git a/tests/wycheproof/cmac_aes.rs b/tests/wycheproof/cmac_aes.rs index 7aa5967c6..a88559f24 100644 --- a/tests/wycheproof/cmac_aes.rs +++ b/tests/wycheproof/cmac_aes.rs @@ -5,7 +5,7 @@ //! invalid one is either a modified tag, which `verify` must reject, or a //! key of a length AES does not take, which `new` must reject. -#![cfg(target_arch = "x86_64")] +#![cfg(any(target_arch = "x86_64", target_arch = "aarch64"))] use serde::Deserialize; use verified_garbage::cmac::InvalidKeyLength; From 5fae20403bb9b42d74ac5ee85367ebfc6900b3b7 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 15:10:24 +0000 Subject: [PATCH 3/4] Test AES-CMAC in each CPU-feature configuration, and cover its key-length check The rust-cpu-features jobs run only the tests they name; without cmac, the scalar implementation (chosen under SDE's Pentium 4) never ran under coverage. The rejected key's assertion keeps its message on one line, which coverage counts as run. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01WkLN6tAYk76HACiEWLbAMD --- .github/workflows/ci.yml | 2 +- tests/wycheproof/cmac_aes.rs | 7 ++----- 2 files changed, 3 insertions(+), 6 deletions(-) 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/tests/wycheproof/cmac_aes.rs b/tests/wycheproof/cmac_aes.rs index 7aa5967c6..5ec8ee84b 100644 --- a/tests/wycheproof/cmac_aes.rs +++ b/tests/wycheproof/cmac_aes.rs @@ -40,11 +40,8 @@ fn cmac_aes() { 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}"); - assert_eq!( - AesCmac::new(&key.0).err(), - Some(InvalidKeyLength), - "tcId {id}" - ); + let err = AesCmac::new(&key.0).err(); + assert_eq!(err, Some(InvalidKeyLength), "tcId {id}"); bad_keys += 1; continue; } From 0d7fb4fcc44602b3794f0dcd115611d5c776764d Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 1 Oct 2026 15:10:38 +0000 Subject: [PATCH 4/4] Test AES-CMAC with the AES extension on the ARM64 runner Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01WkLN6tAYk76HACiEWLbAMD --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 77d025c99..8210e2fa3 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -416,7 +416,7 @@ jobs: tests: x25519 cpu - chip: native os: ubuntu-24.04-arm - tests: aes_gcm sha1 sha256 sha384 sha512 sha3 shake mlkem mldsa scrypt chacha20 ed25519 message_boundaries_and_unaligned_inputs signing_with_aliased_read_only_inputs cpu + tests: aes_gcm cmac sha1 sha256 sha384 sha512 sha3 shake mlkem mldsa scrypt chacha20 ed25519 message_boundaries_and_unaligned_inputs signing_with_aliased_read_only_inputs cpu name: Rust (${{ matrix.os && format('native, {0}', matrix.os) || matrix.chip == 'native' && 'native' || format('sde -{0}', matrix.chip) }}) runs-on: ${{ matrix.os || 'ubuntu-latest' }} env: