~/bend-docscommunity

libs/AES256GCMCore.bend source

libs/AES256GCMCore.bend on the hub · documented module

import Base
import ./AES256SBox.bend as SBox

# Byte arithmetic. Inputs and outputs are represented by U32 values in [0,255].
# Turn a boolean into an all-zero or all-one word to keep field arithmetic branchless.
def xor_if(a: U32, b: U32, yes: Bool) -> U32:
    mask = U32.sub(0, Bool.to_u32(yes))
    U32.xor(a, U32.and(mask, b))

def gf8_xtime.high(+a: U32, high: Bool) -> U32:
    xor_if(U32.and(255, U32.shln(a, 1n)), 27, high)

def gf8_xtime(+a: U32) -> U32:
    gf8_xtime.high(a, U32.is_ne(U32.and(128, a), 0))

def gf8_mul.go(fuel: Nat, +a: U32, +b: U32, acc: U32) -> U32:
    match fuel:
        case 0n: acc
        case 1n+p:
            bit = U32.is_ne(U32.and(1, b), 0)
            next = xor_if(acc, a, bit)
            gf8_mul.go(p, gf8_xtime(a), U32.shrn(b, 1n), next)

def gf8_mul(a: U32, b: U32) -> U32:
    gf8_mul.go(8n, a, b, 0)

def gf8_pow.go(fuel: Nat, +base: U32, +exp: U32, +acc: U32) -> U32:
    match fuel:
        case 0n: acc
        case 1n+p:
            next = Bool.pick(U32, U32.is_ne(U32.and(1, exp), 0), gf8_mul(acc, base), acc)
            gf8_pow.go(p, gf8_mul(base, base), U32.shrn(exp, 1n), next)

def gf8_pow(a: U32, exp: U32) -> U32:
    gf8_pow.go(8n, a, exp, 1)

def rotl8(+a: U32, +shift: Nat) -> U32:
    U32.and(255, U32.or(U32.shln(a, shift), U32.shrn(a, (8n - shift : Nat))))

def aes_sbox_algebraic(+x: U32) -> U32:
    +inv = gf8_pow(x, 254)
    U32.xor(U32.xor(U32.xor(U32.xor(inv, rotl8(inv, 1n)),
        rotl8(inv, 2n)), rotl8(inv, 3n)),
        U32.xor(rotl8(inv, 4n), 99))

def aes_sbox(+x: U32) -> U32:
    SBox.lookup(x)

def aes_subword(+w: U32) -> U32:
    a b c d = aes_sbox(U32.shrn(w, 24n)) aes_sbox(U32.and(255, U32.shrn(w, 16n)))
        aes_sbox(U32.and(255, U32.shrn(w, 8n))) aes_sbox(U32.and(255, w))
    U32.or(U32.or(U32.shln(a, 24n), U32.shln(b, 16n)),
        U32.or(U32.shln(c, 8n), d))

def aes_rotword(+w: U32) -> U32:
    U32.or(U32.shln(w, 8n), U32.shrn(w, 24n))

def aes_rcon_step(r: U32) -> U32:
    gf8_xtime(r)

def aes_word(a: U32, b: U32, c: U32, d: U32) -> U32:
    U32.or(U32.or(U32.shln(a, 24n), U32.shln(b, 16n)),
        U32.or(U32.shln(c, 8n), d))

def aes_words(bytes: List<&2, U32>) -> List<&2, U32>:
    match bytes:
        case a <> b <> c <> d <> tail:
            aes_word(a, b, c, d) <> aes_words(tail)
        case _: Nil{}

def aes_word_from_maybe(value: Maybe<&2, U32>) -> U32:
    match value:
        case None{}: 0
        case Some{w}: w

def aes_word_at(words: List<&2, U32>, i: Nat) -> U32:
    aes_word_from_maybe(List.get(&2, U32, words, i))

def aes_key_word_kind.step(previous: U32, rcon: U32, mod4: Bool) -> U32:
    match mod4:
        case True{}: aes_subword(previous)
        case False{}: previous

def aes_key_word_kind.rcon(previous: U32, rcon: U32, mod4: Bool) -> U32:
    U32.xor(aes_subword(aes_rotword(previous)), U32.shln(rcon, 24n))

def aes_key_word_kind(previous: U32, rcon: U32, mod0: Bool, mod4: Bool) -> U32:
    match mod0:
        case True{}: aes_key_word_kind.rcon(previous, rcon, mod4)
        case False{}: aes_key_word_kind.step(previous, rcon, mod4)

def aes_rcon_kind(rcon: U32, advance: Bool) -> U32:
    match advance:
        case True{}: aes_rcon_step(rcon)
        case False{}: rcon

def aes_expand.go(fuel: Nat, +i: U32, +rcon: U32,
                  +words: List<&2, U32>) -> List<&2, U32>:
    match fuel:
        case 0n: words
        case 1n+p:
            previous = aes_word_at(words, U32.to_nat(U32.sub(i, 1)))
            back8 = aes_word_at(words, U32.to_nat(U32.sub(i, 8)))
            part = aes_key_word_kind(previous, rcon, U32.is_eq(U32.mod(i, 8), 0), U32.is_eq(U32.mod(i, 8), 4))
            next = U32.xor(back8, part)
            next_rcon = aes_rcon_kind(rcon, U32.is_eq(U32.mod(i, 8), 0))
            aes_expand.go(p, U32.add(i, 1), next_rcon,
                List.append(&2, U32, words, [next]))

def aes256_expand_impl(+key: List<&2, U32>) -> List<&2, U32>:
    aes_expand.go(52n, 8, 1, aes_words(key))

# An abstract key remains neutral. The explicit fuel parameter also lets
# computation certificates compose without unfolding all 52 expansion steps.
def aes_expand_key(+fuel: Nat, +key: List<&2, U32>) -> List<&2, U32>:
    match key:
        case Nil{}: aes_expand.go(fuel, 8, 1, aes_words(key))
        case head <> tail: aes_expand.go(fuel, 8, 1, aes_words(key))

def aes256_expand(+key: List<&2, U32>) -> List<&2, U32>:
    aes_expand_key(52n, key)

def aes_add_key.go(+state: List<&2, U32>, +words: List<&2, U32>,
                   +round: Nat, +index: Nat, fuel: Nat,
                   acc: List<&2, U32>) -> List<&2, U32>:
    match state fuel:
        case Nil{} _: List.reverse(&2, U32, acc)
        case byte <> tail 0n: List.reverse(&2, U32, acc)
        case +byte <> tail 1n+p:
            word_index = ((round * 4n + (index / 4n : Nat)) : Nat)
            word = aes_word_at(words, word_index)
            shift = ((3n - (index % 4n : Nat)) * 8n : Nat)
            keybyte = U32.and(255, U32.shrn(word, shift))
            aes_add_key.go(tail, words, round, (index + 1n : Nat), p,
                U32.xor(byte, keybyte) <> acc)

def aes_add_key(state: List<&2, U32>, words: List<&2, U32>, round: Nat) -> List<&2, U32>:
    aes_add_key.go(state, words, round, 0n, 16n, Nil{})

def aes_sub_state.go(+state: List<&2, U32>, +acc: List<&2, U32>) -> List<&2, U32>:
    match state:
        case Nil{}: List.reverse(&2, U32, acc)
        case byte <> tail:
            aes_sub_state.go(tail, aes_sbox(byte) <> acc)

def aes_sub_state(state: List<&2, U32>) -> List<&2, U32>:
    aes_sub_state.go(state, Nil{})

def aes_byte_from_maybe(value: Maybe<&2, U32>) -> U32:
    match value:
        case None{}: 0
        case Some{byte}: byte

def aes_byte_at(state: List<&2, U32>, index: Nat) -> U32:
    aes_byte_from_maybe(List.get(&2, U32, state, index))

def aes_shift_rows(+state: List<&2, U32>) -> List<&2, U32>:
    [aes_byte_at(state, 0n), aes_byte_at(state, 5n), aes_byte_at(state, 10n), aes_byte_at(state, 15n),
     aes_byte_at(state, 4n), aes_byte_at(state, 9n), aes_byte_at(state, 14n), aes_byte_at(state, 3n),
     aes_byte_at(state, 8n), aes_byte_at(state, 13n), aes_byte_at(state, 2n), aes_byte_at(state, 7n),
     aes_byte_at(state, 12n), aes_byte_at(state, 1n), aes_byte_at(state, 6n), aes_byte_at(state, 11n)]

type AesColumn is Data:
    AesColumn{a: U32, b: U32, c: U32, d: U32}

def aes_mix_column(col: AesColumn) -> AesColumn:
    match col:
        case AesColumn{+a, +b, +c, +d}:
            AesColumn{
                U32.xor(U32.xor(U32.xor(gf8_mul(a, 2), gf8_mul(b, 3)), c), d),
                U32.xor(U32.xor(U32.xor(a, gf8_mul(b, 2)), gf8_mul(c, 3)), d),
                U32.xor(U32.xor(U32.xor(a, b), gf8_mul(c, 2)), gf8_mul(d, 3)),
                U32.xor(U32.xor(U32.xor(gf8_mul(a, 3), b), c), gf8_mul(d, 2))}

def aes_mix_column.first(col: AesColumn) -> U32:
    match col:
        case AesColumn{+a, +b, +c, +d}:
            U32.xor(U32.xor(U32.xor(gf8_mul(a, 2), gf8_mul(b, 3)), c), d)

def aes_mix_column.second(col: AesColumn) -> U32:
    match col:
        case AesColumn{+a, +b, +c, +d}:
            U32.xor(U32.xor(U32.xor(a, gf8_mul(b, 2)), gf8_mul(c, 3)), d)

def aes_mix_column.third(col: AesColumn) -> U32:
    match col:
        case AesColumn{+a, +b, +c, +d}:
            U32.xor(U32.xor(U32.xor(a, b), gf8_mul(c, 2)), gf8_mul(d, 3))

def aes_mix_column.fourth(col: AesColumn) -> U32:
    match col:
        case AesColumn{+a, +b, +c, +d}:
            U32.xor(U32.xor(U32.xor(gf8_mul(a, 3), b), c), gf8_mul(d, 2))

def aes_mix_row(col: AesColumn, row: Nat) -> U32:
    match row:
        case 0n: aes_mix_column.first(col)
        case 1n: aes_mix_column.second(col)
        case 2n: aes_mix_column.third(col)
        case _: aes_mix_column.fourth(col)

def aes_mix_byte(+state: List<&2, U32>, row: Nat, +col: Nat) -> U32:
    column: AesColumn = AesColumn{
        aes_byte_at(state, (col * 4n : Nat)),
        aes_byte_at(state, (col * 4n + 1n : Nat)),
        aes_byte_at(state, (col * 4n + 2n : Nat)),
        aes_byte_at(state, (col * 4n + 3n : Nat))}
    aes_mix_row(column, row)

def aes_mix_state.finish(col0: AesColumn, col1: AesColumn, col2: AesColumn, col3: AesColumn) -> List<&2, U32>:
    match col0 col1 col2 col3:
        case AesColumn{a0, b0, c0, d0} AesColumn{a1, b1, c1, d1} AesColumn{a2, b2, c2, d2} AesColumn{a3, b3, c3, d3}:
            [a0,b0,c0,d0,a1,b1,c1,d1,a2,b2,c2,d2,a3,b3,c3,d3]

def aes_mix_state(+state: List<&2, U32>) -> List<&2, U32>:
    col0 col1 col2 col3 =
        aes_mix_column(AesColumn{aes_byte_at(state, 0n), aes_byte_at(state, 1n), aes_byte_at(state, 2n), aes_byte_at(state, 3n)})
        aes_mix_column(AesColumn{aes_byte_at(state, 4n), aes_byte_at(state, 5n), aes_byte_at(state, 6n), aes_byte_at(state, 7n)})
        aes_mix_column(AesColumn{aes_byte_at(state, 8n), aes_byte_at(state, 9n), aes_byte_at(state, 10n), aes_byte_at(state, 11n)})
        aes_mix_column(AesColumn{aes_byte_at(state, 12n), aes_byte_at(state, 13n), aes_byte_at(state, 14n), aes_byte_at(state, 15n)})
    aes_mix_state.finish(col0, col1, col2, col3)

def aes_middle_round(state: List<&2, U32>, words: List<&2, U32>, round: Nat) -> List<&2, U32>:
    aes_add_key(aes_mix_state(aes_shift_rows(aes_sub_state(state))), words, round)

def aes_final_round(state: List<&2, U32>, words: List<&2, U32>) -> List<&2, U32>:
    aes_add_key(aes_shift_rows(aes_sub_state(state)), words, 14n)

def aes_rounds.go(fuel: Nat, +state: List<&2, U32>, +words: List<&2, U32>, +round: Nat) -> List<&2, U32>:
    match fuel:
        case 0n: aes_final_round(state, words)
        case 1n+p:
            aes_rounds.go(p, aes_middle_round(state, words, round), words, (round + 1n : Nat))

def aes256_encrypt_expanded(+words: List<&2, U32>, block: List<&2, U32>) -> List<&2, U32>:
    match words:
        case Nil{}: aes_rounds.go(13n, aes_add_key(block, words, 0n), words, 1n)
        case head <> tail: aes_rounds.go(13n, aes_add_key(block, words, 0n), words, 1n)

def aes256_encrypt_block(key: List<&2, U32>, block: List<&2, U32>) -> List<&2, U32>:
    aes256_encrypt_expanded(aes256_expand(key), block)

# GCM counter mode and GHASH use 16-byte blocks in network byte order.
def gcm_hex_mask(value: U32, mask: U32) -> U32:
    U32.and(mask, value)

def gcm_mask_byte_valid(+value: U32) ->
    {U32.is_lt(gcm_hex_mask(value, 255), 256) == True{} : Bool}:
    match value:
        case U32{WCon{b0, WCon{b1, WCon{b2, WCon{b3, WCon{b4, WCon{b5, WCon{b6, WCon{b7, WCon{b8, WCon{b9, WCon{b10, WCon{b11, WCon{b12, WCon{b13, WCon{b14, WCon{b15, WCon{b16, WCon{b17, WCon{b18, WCon{b19, WCon{b20, WCon{b21, WCon{b22, WCon{b23, WCon{b24, WCon{b25, WCon{b26, WCon{b27, WCon{b28, WCon{b29, WCon{b30, WCon{b31, WNil{}}}}}}}}}}}}}}}}}}}}}}}}}}}}}}}}}}: {==}

def gcm_bytes_valid(bytes: List<&2, U32>) -> Bool:
    match bytes:
        case Nil{}: True{}
        case byte <> tail:
            U32.is_lt(byte, 256) && gcm_bytes_valid(tail)

def gcm_mask_bytes(+bytes: List<&2, U32>) -> List<&2, U32>:
    match bytes:
        case Nil{}: Nil{}
        case byte <> tail:
            gcm_hex_mask(byte, 255) <> gcm_mask_bytes(tail)

def gcm_mask_bytes_valid(+bytes: List<&2, U32>) ->
    {gcm_bytes_valid(gcm_mask_bytes(bytes)) == True{} : Bool}:
    match bytes:
        case Nil{}: {==}
        case byte <> tail:
            Equal.trans(Bool,
                gcm_bytes_valid(gcm_mask_bytes(byte <> tail)),
                Bool.and(True{}, gcm_bytes_valid(gcm_mask_bytes(tail))),
                True{},
                Equal.cong(Bool, Bool,
                    head => Bool.and(head,
                        gcm_bytes_valid(gcm_mask_bytes(tail))),
                    U32.is_lt(gcm_hex_mask(byte, 255), 256), True{},
                    gcm_mask_byte_valid(byte)),
                Equal.cong(Bool, Bool, tail_ok => Bool.and(True{}, tail_ok),
                    gcm_bytes_valid(gcm_mask_bytes(tail)), True{},
                    gcm_mask_bytes_valid(tail)))

def gcm_mask_bytes_length(+bytes: List<&2, U32>) ->
    {List.length(&2, U32, gcm_mask_bytes(bytes)) ==
     List.length(&2, U32, bytes) : Nat}:
    match bytes:
        case Nil{}: {==}
        case byte <> tail:
            Equal.cong(Nat, Nat, n => 1n+n,
                List.length(&2, U32, gcm_mask_bytes(tail)),
                List.length(&2, U32, tail),
                gcm_mask_bytes_length(tail))

type GcmTag is Data:
    GcmTag{
        bytes: List<&2, U32>,
        valid: {gcm_bytes_valid(bytes) == True{} : Bool},
        size: {List.length(&2, U32, bytes) == 16n : Nat}
    }

type GcmCiphertext is Data:
    GcmCiphertext{
        bytes: List<&2, U32>,
        valid: {gcm_bytes_valid(bytes) == True{} : Bool}
    }

def gcm_ciphertext_bytes(ciphertext: GcmCiphertext) -> List<&2, U32>:
    match ciphertext:
        case GcmCiphertext{bytes, valid}: bytes


def ghash_xor_if(a: U32, b: U32, yes: Bool) -> U32:
    match yes:
        case True{}: U32.xor(a, b)
        case False{}: a

def ghash_xor_list.go(a: List<&2, U32>, b: List<&2, U32>, +yes: Bool,
                      +acc: List<&2, U32>) -> List<&2, U32>:
    match a b:
        case Nil{} _: List.reverse(&2, U32, acc)
        case _ Nil{}: List.reverse(&2, U32, acc)
        case x <> xt y <> yt:
            ghash_xor_list.go(xt, yt, yes, ghash_xor_if(x, y, yes) <> acc)

def ghash_xor_list(a: List<&2, U32>, b: List<&2, U32>, yes: Bool) -> List<&2, U32>:
    ghash_xor_list.go(a, b, yes, Nil{})

# Keep the 128-bit GHASH accumulator in four machine words during its hot loop.
type GhashField is Data:
    GhashField{w0: U32, w1: U32, w2: U32, w3: U32}

def ghash_field(+bytes: List<&2, U32>) -> GhashField:
    +words = aes_words(bytes)
    GhashField{aes_word_at(words, 0n), aes_word_at(words, 1n),
        aes_word_at(words, 2n), aes_word_at(words, 3n)}

def ghash_field_bytes(value: GhashField) -> List<&2, U32>:
    match value:
        case GhashField{+w0, +w1, +w2, +w3}:
            [U32.and(255, U32.shrn(w0, 24n)), U32.and(255, U32.shrn(w0, 16n)),
             U32.and(255, U32.shrn(w0, 8n)), U32.and(255, w0),
             U32.and(255, U32.shrn(w1, 24n)), U32.and(255, U32.shrn(w1, 16n)),
             U32.and(255, U32.shrn(w1, 8n)), U32.and(255, w1),
             U32.and(255, U32.shrn(w2, 24n)), U32.and(255, U32.shrn(w2, 16n)),
             U32.and(255, U32.shrn(w2, 8n)), U32.and(255, w2),
             U32.and(255, U32.shrn(w3, 24n)), U32.and(255, U32.shrn(w3, 16n)),
             U32.and(255, U32.shrn(w3, 8n)), U32.and(255, w3)]

def ghash_field_word(value: GhashField, index: Nat) -> U32:
    match value index:
        case GhashField{w0, _, _, _} 0n: w0
        case GhashField{_, w1, _, _} 1n: w1
        case GhashField{_, _, w2, _} 2n: w2
        case GhashField{_, _, _, w3} _: w3

def ghash_field_shift_top(w0: U32, reduce: Bool) -> U32:
    reduction_mask = U32.sub(0, Bool.to_u32(reduce))
    U32.xor(U32.shrn(w0, 1n), U32.and(3774873600, reduction_mask))

def ghash_field_shift_right(value: GhashField) -> GhashField:
    match value:
        case GhashField{+w0, +w1, +w2, +w3}:
            GhashField{
                ghash_field_shift_top(w0, U32.is_ne(U32.and(1, w3), 0)),
                U32.or(U32.shrn(w1, 1n), U32.shln(U32.and(1, w0), 31n)),
                U32.or(U32.shrn(w2, 1n), U32.shln(U32.and(1, w1), 31n)),
                U32.or(U32.shrn(w3, 1n), U32.shln(U32.and(1, w2), 31n))}

def ghash_field_xor_if(a: U32, b: U32, yes: Bool) -> U32:
    mask = U32.sub(0, Bool.to_u32(yes))
    U32.xor(a, U32.and(mask, b))

def ghash_field_xor_if.all(a: GhashField, b: GhashField, +yes: Bool) -> GhashField:
    match a b:
        case GhashField{a0, a1, a2, a3} GhashField{b0, b1, b2, b3}:
            GhashField{ghash_field_xor_if(a0, b0, yes), ghash_field_xor_if(a1, b1, yes),
                ghash_field_xor_if(a2, b2, yes), ghash_field_xor_if(a3, b3, yes)}

def ghash_mul.go(fuel: Nat, +x: GhashField, +v: GhashField,
                 +z: GhashField, +bit_index: Nat) -> GhashField:
    match fuel:
        case 0n: z
        case 1n+p:
            input_word = ghash_field_word(x, (bit_index / 32n : Nat))
            shift = (31n - (bit_index % 32n : Nat) : Nat)
            set = U32.is_ne(U32.and(1, U32.shrn(input_word, shift)), 0)
            next_z = ghash_field_xor_if.all(z, v, set)
            ghash_mul.go(p, x, ghash_field_shift_right(v), next_z, (bit_index + 1n : Nat))

def ghash_mul(x: List<&2, U32>, h: List<&2, U32>) -> List<&2, U32>:
    ghash_field_bytes(ghash_mul.go(128n, ghash_field(x), ghash_field(h),
        GhashField{0,0,0,0}, 0n))

def ghash_update(y: List<&2, U32>, h: List<&2, U32>, block: List<&2, U32>) -> List<&2, U32>:
    ghash_mul(ghash_xor_list(y, block, True{}), h)

def ghash_zero_pad.go(fuel: Nat, +bytes: List<&2, U32>) -> List<&2, U32>:
    match fuel:
        case 0n: bytes
        case 1n+p: ghash_zero_pad.go(p, List.append(&2, U32, bytes, [0]))

def ghash_finish(space: Nat, block_rev: List<&2, U32>, h: List<&2, U32>, y: List<&2, U32>) -> List<&2, U32>:
    match space:
        case 16n: y
        case 0n: ghash_update(y, h, List.reverse(&2, U32, block_rev))
        case 1n+p:
            ghash_update(y, h, ghash_zero_pad.go(space, List.reverse(&2, U32, block_rev)))

def ghash_bytes.go(fuel: Nat, bytes: List<&2, U32>, space: Nat,
                   +block_rev: List<&2, U32>, +h: List<&2, U32>, +y: List<&2, U32>) -> List<&2, U32>:
    match fuel bytes space:
        case 0n Nil{} _: ghash_finish(space, block_rev, h, y)
        case 0n _ _: ghash_finish(space, block_rev, h, y)
        case 1n+p Nil{} _: ghash_finish(space, block_rev, h, y)
        case 1n+p byte <> tail 0n:
            y2 = ghash_update(y, h, List.reverse(&2, U32, block_rev))
            ghash_bytes.go(p, tail, 15n, [byte], h, y2)
        case 1n+p byte <> tail 1n+s:
            ghash_bytes.go(p, tail, s, byte <> block_rev, h, y)

def ghash_bytes(+bytes: List<&2, U32>, h: List<&2, U32>, y: List<&2, U32>) -> List<&2, U32>:
    ghash_bytes.go(List.length(&2, U32, bytes), bytes, 16n, Nil{}, h, y)

def ghash_nat_bytes.go(fuel: Nat, qr: Nat & Nat, +acc: List<&2, U32>) -> List<&2, U32>:
    match fuel:
        case 0n: acc
        case 1n+p:
            (q, r) = qr
            ghash_nat_bytes.go(p, Nat.divmod(q, 256n), U32.from_nat(r) <> acc)

def ghash_nat_bytes(n: Nat) -> List<&2, U32>:
    ghash_nat_bytes.go(8n, Nat.divmod(n, 256n), Nil{})

def ghash_auth_with_hash_impl(+h: List<&2, U32>, +aad: List<&2, U32>, +ciphertext: List<&2, U32>) -> List<&2, U32>:
    y0: List<&2, U32> = [0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0]
    y1 = ghash_bytes(aad, h, y0)
    y2 = ghash_bytes(ciphertext, h, y1)
    aad_bits = ghash_nat_bytes((List.length(&2, U32, aad) * 8n : Nat))
    ct_bits = ghash_nat_bytes((List.length(&2, U32, ciphertext) * 8n : Nat))
    ghash_update(y2, h, List.append(&2, U32, aad_bits, ct_bits))

# Do not expand a symbolic GHASH multiplication while transporting its hash.
def ghash_auth_with_hash(+h: List<&2, U32>, +aad: List<&2, U32>, +ciphertext: List<&2, U32>) -> List<&2, U32>:
    match aad:
        case Nil{}: ghash_auth_with_hash_impl(h, aad, ciphertext)
        case head <> tail: ghash_auth_with_hash_impl(h, aad, ciphertext)

def ghash_auth_expanded(words: List<&2, U32>, aad: List<&2, U32>, ciphertext: List<&2, U32>) -> List<&2, U32>:
    ghash_auth_with_hash(aes256_encrypt_expanded(words,
        [0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0]), aad, ciphertext)

def gcm_j0(nonce: List<&2, U32>) -> List<&2, U32>:
    List.append(&2, U32, nonce, [0,0,0,1])

def gcm_inc32(+counter: List<&2, U32>) -> List<&2, U32>:
    hi = U32.shln(aes_byte_at(counter, 12n), 24n)
    mid1 = U32.shln(aes_byte_at(counter, 13n), 16n)
    mid2 = U32.shln(aes_byte_at(counter, 14n), 8n)
    lo = aes_byte_at(counter, 15n)
    +word = U32.add(U32.or(U32.or(hi, mid1), U32.or(mid2, lo)), 1)
    List.append(&2, U32, List.take(&2, U32, counter, 12n),
        [U32.and(255, U32.shrn(word, 24n)), U32.and(255, U32.shrn(word, 16n)),
         U32.and(255, U32.shrn(word, 8n)), U32.and(255, word)])

def gcm_ctr_stream.go(fuel: Nat, +words: List<&2, U32>, +counter: List<&2, U32>,
                      +acc: List<&2, U32>) -> List<&2, U32>:
    match fuel:
        case 0n: List.reverse(&2, U32, acc)
        case 1n+p:
            stream = aes256_encrypt_expanded(words, counter)
            next_acc = List.append(&2, U32, List.reverse(&2, U32, stream), acc)
            gcm_ctr_stream.go(p, words, gcm_inc32(counter), next_acc)

def gcm_ctr_advance(fuel: Nat, +counter: List<&2, U32>) -> List<&2, U32>:
    match fuel:
        case 0n: counter
        case 1n+p: gcm_ctr_advance(p, gcm_inc32(counter))

type GcmCtrStep is Data:
    GcmCtrStep{byte: U32, counter: List<&2, U32>, rest: List<&2, U32>}

def gcm_ctr_step(+words: List<&2, U32>, +counter: List<&2, U32>,
    block: List<&2, U32>) -> GcmCtrStep:
    match block:
        case Nil{}:
            +next_block = aes256_encrypt_expanded(words, counter)
            GcmCtrStep{aes_byte_at(next_block, 0n), gcm_inc32(counter),
                List.tail(&2, U32, next_block)}
        case byte <> rest: GcmCtrStep{byte, counter, rest}

def gcm_ctr_step_byte(step: GcmCtrStep) -> U32:
    match step:
        case GcmCtrStep{byte, _, _}: byte

def gcm_ctr_step_counter(step: GcmCtrStep) -> List<&2, U32>:
    match step:
        case GcmCtrStep{_, counter, _}: counter

def gcm_ctr_step_rest(step: GcmCtrStep) -> List<&2, U32>:
    match step:
        case GcmCtrStep{_, _, rest}: rest

# Generate each AES counter block once and consume all of its bytes. Fuel is
# still a byte count, so partial final blocks have exactly the requested length.
def gcm_ctr_stream_exact.go(fuel: Nat, +words: List<&2, U32>,
    counter: List<&2, U32>, block: List<&2, U32>) -> List<&2, U32>:
    match fuel:
        case 0n: Nil{}
        case 1n+p:
            +step = gcm_ctr_step(words, counter, block)
            gcm_ctr_step_byte(step) <>
                gcm_ctr_stream_exact.go(p, words, gcm_ctr_step_counter(step),
                    gcm_ctr_step_rest(step))

def gcm_ctr_stream_exact(fuel: Nat, +words: List<&2, U32>,
    +start: List<&2, U32>) -> List<&2, U32>:
    gcm_ctr_stream_exact.go(fuel, words, start, Nil{})

def gcm_zip_xor(+bytes: List<&2, U32>, +stream: List<&2, U32>) -> List<&2, U32>:
    match bytes stream:
        case Nil{} _: Nil{}
        case _ Nil{}: Nil{}
        case byte <> tail stream_byte <> stream_tail:
            gcm_hex_mask(U32.xor(byte, stream_byte), 255) <>
                gcm_zip_xor(tail, stream_tail)

def gcm_zip_xor_nil_left(+stream: List<&2, U32>) ->
    {gcm_zip_xor(Nil{}, stream) == Nil{} : List<&2, U32>}:
    match stream:
        case Nil{}: {==}
        case _ <> _: {==}

def gcm_zip_xor_nil_right(+bytes: List<&2, U32>) ->
    {gcm_zip_xor(bytes, Nil{}) == Nil{} : List<&2, U32>}:
    match bytes:
        case Nil{}: {==}
        case _ <> _: {==}

def gcm_zip_xor_valid(+bytes: List<&2, U32>, +stream: List<&2, U32>) ->
    {gcm_bytes_valid(gcm_zip_xor(bytes, stream)) == True{} : Bool}:
    match bytes stream:
        case Nil{} _:
            Equal.trans(Bool,
                gcm_bytes_valid(gcm_zip_xor(Nil{}, stream)),
                gcm_bytes_valid(Nil{}), True{},
                Equal.cong(List<&2, U32>, Bool, gcm_bytes_valid,
                    gcm_zip_xor(Nil{}, stream), Nil{},
                    gcm_zip_xor_nil_left(stream)),
                {==})
        case _ Nil{}:
            Equal.trans(Bool,
                gcm_bytes_valid(gcm_zip_xor(bytes, Nil{})),
                gcm_bytes_valid(Nil{}), True{},
                Equal.cong(List<&2, U32>, Bool, gcm_bytes_valid,
                    gcm_zip_xor(bytes, Nil{}), Nil{},
                    gcm_zip_xor_nil_right(bytes)),
                {==})
        case byte <> tail stream_byte <> stream_tail:
            recursive = gcm_zip_xor_valid(tail, stream_tail)
            range = gcm_mask_byte_valid(U32.xor(byte, stream_byte))
            head_changed = Equal.cong(Bool, Bool,
                head => Bool.and(head,
                    gcm_bytes_valid(gcm_zip_xor(tail, stream_tail))),
                U32.is_lt(gcm_hex_mask(U32.xor(byte, stream_byte), 255), 256),
                True{}, range)
            tail_changed = Equal.cong(Bool, Bool,
                valid => Bool.and(True{}, valid),
                gcm_bytes_valid(gcm_zip_xor(tail, stream_tail)), True{},
                recursive)
            Equal.trans(Bool,
                Bool.and(
                    U32.is_lt(gcm_hex_mask(U32.xor(byte, stream_byte), 255), 256),
                    gcm_bytes_valid(gcm_zip_xor(tail, stream_tail))),
                Bool.and(True{}, gcm_bytes_valid(gcm_zip_xor(tail, stream_tail))),
                True{}, head_changed,
                Equal.trans(Bool,
                    Bool.and(True{}, gcm_bytes_valid(gcm_zip_xor(tail, stream_tail))),
                    Bool.and(True{}, True{}), True{}, tail_changed, {==}))

# A GCM tag is exactly sixteen bytes. Expose that shape without unfolding
# abstract AES rounds or GHASH values when checking its certificates.
def gcm_tag_parts(+auth: List<&2, U32>, +mask: List<&2, U32>) -> List<&2, U32>:
    [U32.xor(aes_byte_at(mask, 0n), aes_byte_at(auth, 0n)),
     U32.xor(aes_byte_at(mask, 1n), aes_byte_at(auth, 1n)),
     U32.xor(aes_byte_at(mask, 2n), aes_byte_at(auth, 2n)),
     U32.xor(aes_byte_at(mask, 3n), aes_byte_at(auth, 3n)),
     U32.xor(aes_byte_at(mask, 4n), aes_byte_at(auth, 4n)),
     U32.xor(aes_byte_at(mask, 5n), aes_byte_at(auth, 5n)),
     U32.xor(aes_byte_at(mask, 6n), aes_byte_at(auth, 6n)),
     U32.xor(aes_byte_at(mask, 7n), aes_byte_at(auth, 7n)),
     U32.xor(aes_byte_at(mask, 8n), aes_byte_at(auth, 8n)),
     U32.xor(aes_byte_at(mask, 9n), aes_byte_at(auth, 9n)),
     U32.xor(aes_byte_at(mask, 10n), aes_byte_at(auth, 10n)),
     U32.xor(aes_byte_at(mask, 11n), aes_byte_at(auth, 11n)),
     U32.xor(aes_byte_at(mask, 12n), aes_byte_at(auth, 12n)),
     U32.xor(aes_byte_at(mask, 13n), aes_byte_at(auth, 13n)),
     U32.xor(aes_byte_at(mask, 14n), aes_byte_at(auth, 14n)),
     U32.xor(aes_byte_at(mask, 15n), aes_byte_at(auth, 15n))]

def gcm_tag_output_impl(+words: List<&2, U32>, nonce: List<&2, U32>,
                   aad: List<&2, U32>, ciphertext: List<&2, U32>) -> GcmTag:
    auth mask = ghash_auth_expanded(words, aad, ciphertext)
        aes256_encrypt_expanded(words, gcm_j0(nonce))
    +raw_tag = gcm_tag_parts(auth, mask)
    GcmTag{gcm_mask_bytes(raw_tag), gcm_mask_bytes_valid(raw_tag), {==}}

# Keep an abstract key schedule abstract during proof normalization. Splitting
# on the nonce leaves the key schedule abstract in both identical branches.
def gcm_tag_output(+words: List<&2, U32>, nonce: List<&2, U32>,
                   aad: List<&2, U32>, ciphertext: List<&2, U32>) -> GcmTag:
    match nonce:
        case Nil{}: gcm_tag_output_impl(words, nonce, aad, ciphertext)
        case head <> tail: gcm_tag_output_impl(words, nonce, aad, ciphertext)

def gcm_tag_bytes(tag_output: GcmTag) -> List<&2, U32>:
    match tag_output:
        case GcmTag{bytes, valid, size}: bytes

def gcm_tag_valid(tag_output: GcmTag) ->
    {gcm_bytes_valid(gcm_tag_bytes(tag_output)) == True{} : Bool}:
    match tag_output:
        case GcmTag{bytes, valid, size}: valid

def gcm_tag_size(tag_output: GcmTag) ->
    {List.length(&2, U32, gcm_tag_bytes(tag_output)) == 16n : Nat}:
    match tag_output:
        case GcmTag{bytes, valid, size}: size

type GcmOutput is Data:
    GcmOutput{
        ciphertext: List<&2, U32>,
        ciphertext_valid: {gcm_bytes_valid(ciphertext) == True{} : Bool},
        tag_bytes: List<&2, U32>,
        tag_valid: {gcm_bytes_valid(tag_bytes) == True{} : Bool},
        tag_size: {List.length(&2, U32, tag_bytes) == 16n : Nat},
        tag_words: List<&2, U32>,
        tag_nonce: List<&2, U32>,
        tag_aad: List<&2, U32>,
        tag_matches: {tag_bytes == gcm_tag_bytes(
            gcm_tag_output(tag_words, tag_nonce, tag_aad, ciphertext)) :
            List<&2, U32>}
    }

def gcm_output_tag_bytes(output: GcmOutput) -> List<&2, U32>:
    match output:
        case GcmOutput{ciphertext, ciphertext_valid, tag_bytes,
            tag_valid, tag_size, tag_words, tag_nonce, tag_aad, tag_matches}:
            tag_bytes

def gcm_output_tag_valid(output: GcmOutput) ->
    {gcm_bytes_valid(gcm_output_tag_bytes(output)) == True{} : Bool}:
    match output:
        case GcmOutput{ciphertext, ciphertext_valid, tag_bytes,
            tag_valid, tag_size, tag_words, tag_nonce, tag_aad, tag_matches}:
            tag_valid

def gcm_output_tag_size(output: GcmOutput) ->
    {List.length(&2, U32, gcm_output_tag_bytes(output)) == 16n : Nat}:
    match output:
        case GcmOutput{ciphertext, ciphertext_valid, tag_bytes,
            tag_valid, tag_size, tag_words, tag_nonce, tag_aad, tag_matches}:
            tag_size

def gcm_output_from_tag(+words: List<&2, U32>, +nonce: List<&2, U32>,
    +aad: List<&2, U32>, +ciphertext: List<&2, U32>,
    ciphertext_valid: {gcm_bytes_valid(ciphertext) == True{} : Bool}) ->
    GcmOutput:
    +tag_output = gcm_tag_output(words, nonce, aad, ciphertext)
    GcmOutput{ciphertext, ciphertext_valid, gcm_tag_bytes(tag_output),
        gcm_tag_valid(tag_output), gcm_tag_size(tag_output), words, nonce,
        aad, {==}}

def gcm_output_from_ciphertext(+words: List<&2, U32>,
    nonce: List<&2, U32>, aad: List<&2, U32>,
    cipher_output: GcmCiphertext) -> GcmOutput:
    match cipher_output:
        case GcmCiphertext{+ciphertext, valid}:
            gcm_output_from_tag(words, nonce, aad, ciphertext, valid)

def gcm_ctr_encrypt(+words: List<&2, U32>, nonce: List<&2, U32>,
    +plaintext: List<&2, U32>) -> GcmCiphertext:
    +start = gcm_inc32(gcm_j0(nonce))
    +stream = gcm_ctr_stream_exact(List.length(&2, U32, plaintext),
        words, start)
    GcmCiphertext{gcm_zip_xor(plaintext, stream),
        gcm_zip_xor_valid(plaintext, stream)}

def gcm_encrypt_raw(+key: List<&2, U32>, +nonce: List<&2, U32>,
                    +aad: List<&2, U32>, +plaintext: List<&2, U32>) -> GcmOutput:
    +words = aes256_expand(key)
    gcm_output_from_ciphertext(words, nonce, aad,
        gcm_ctr_encrypt(words, nonce, plaintext))

def gcm_decrypt_expanded(+words: List<&2, U32>, nonce: List<&2, U32>, +ciphertext: List<&2, U32>) -> List<&2, U32>:
    +start = gcm_inc32(gcm_j0(nonce))
    +stream = gcm_ctr_stream_exact(List.length(&2, U32, ciphertext),
        words, start)
    gcm_zip_xor(ciphertext, stream)

def gcm_decrypt_raw(+key: List<&2, U32>, nonce: List<&2, U32>, +ciphertext: List<&2, U32>) -> List<&2, U32>:
    gcm_decrypt_expanded(aes256_expand(key), nonce, ciphertext)

# Authenticate ciphertext directly; expand the key once for both AES calls.
def gcm_tag_expanded(+words: List<&2, U32>, nonce: List<&2, U32>, aad: List<&2, U32>, ciphertext: List<&2, U32>) -> List<&2, U32>:
    gcm_tag_bytes(gcm_tag_output(words, nonce, aad, ciphertext))

def gcm_tag(+key: List<&2, U32>, nonce: List<&2, U32>, aad: List<&2, U32>, ciphertext: List<&2, U32>) -> List<&2, U32>:
    gcm_tag_expanded(aes256_expand(key), nonce, aad, ciphertext)

def ghash_auth(key: List<&2, U32>, aad: List<&2, U32>, ciphertext: List<&2, U32>) -> List<&2, U32>:
    ghash_auth_expanded(aes256_expand(key), aad, ciphertext)