src/crypto/aes/aes.bend source
src/crypto/aes/aes.bend on the hub · documented module
import Baseimport ./types.bend as Timport ./sbox.bend as B# AES (FIPS 197) block encryption for 128-, 192- and 256-bit keys. The# state is four columns of four bytes (U32 values below 256); the S-box and# the field doublings are the constant-time circuits of sbox.bend. The key# schedule is expanded once into the list of round keys.def sub_word(w: T.Quad) -> T.Quad: match w: case T.W{a, b, c, d}: T.W{B.sbox(a), B.sbox(b), B.sbox(c), B.sbox(d)}def sub_bytes(s: T.State) -> T.State: match s: case T.S{c0, c1, c2, c3}: T.S{sub_word(c0), sub_word(c1), sub_word(c2), sub_word(c3)}# Row r moves r columns to the left.def shift_rows(s: T.State) -> T.State: match s: case T.S{T.W{a0, a1, a2, a3}, T.W{b0, b1, b2, b3}, T.W{c0, c1, c2, c3}, T.W{d0, d1, d2, d3}}: T.S{T.W{a0, b1, c2, d3}, T.W{b0, c1, d2, a3}, T.W{c0, d1, a2, b3}, T.W{d0, a1, b2, c3}}def mix_column(w: T.Quad) -> T.Quad: match w: case T.W{+a, +b, +c, +d}: T.W{U32.xor(U32.xor(U32.xor(B.xtime(a), B.mul3(b)), c), d), U32.xor(U32.xor(U32.xor(a, B.xtime(b)), B.mul3(c)), d), U32.xor(U32.xor(U32.xor(a, b), B.xtime(c)), B.mul3(d)), U32.xor(U32.xor(U32.xor(B.mul3(a), b), c), B.xtime(d))}def mix_columns(s: T.State) -> T.State: match s: case T.S{c0, c1, c2, c3}: T.S{mix_column(c0), mix_column(c1), mix_column(c2), mix_column(c3)}def xor_word(a: T.Quad, b: T.Quad) -> T.Quad: match a b: case T.W{a0, a1, a2, a3} T.W{b0, b1, b2, b3}: T.W{U32.xor(a0, b0), U32.xor(a1, b1), U32.xor(a2, b2), U32.xor(a3, b3)}def add_round_key(s: T.State, k: T.State) -> T.State: match s k: case T.S{c0, c1, c2, c3} T.S{k0, k1, k2, k3}: T.S{xor_word(c0, k0), xor_word(c1, k1), xor_word(c2, k2), xor_word(c3, k3)}def round(s: T.State, k: T.State) -> T.State: add_round_key(mix_columns(shift_rows(sub_bytes(s))), k)def last_round(s: T.State, k: T.State) -> T.State: add_round_key(shift_rows(sub_bytes(s)), k)# n full rounds, then the final round.def rounds(n: Nat, s: T.State, ks: List<&2, T.State>) -> T.State: match n ks: case 0n k <> rest: last_round(s, k) case 1n+p k <> rest: rounds(p, round(s, k), rest) case _ _: sdef cipher(+nr: Nat, s: T.State, ks: List<&2, T.State>) -> T.State: match ks: case k <> rest: rounds(Nat.sub(nr, 1n), add_round_key(s, k), rest) case Nil{}: s# ---- key schedule ----def rot_word(w: T.Quad) -> T.Quad: match w: case T.W{a, b, c, d}: T.W{b, c, d, a}# x^n in GF(2^8), by doubling.def xpow(n: Nat) -> U32: match n: case 0n: 1 case 1n+p: B.xtime(xpow(p))def rcon(j: Nat) -> T.Quad: T.W{xpow(Nat.sub(j, 1n)), 0, 0, 0}def key_words(key: List<&2, U32>) -> List<&2, T.Quad>: match key: case a <> b <> c <> d <> rest: T.W{a, b, c, d} <> key_words(rest) case _: Nil{}def get_word(ws: List<&2, T.Quad>, n: Nat) -> T.Quad: match ws n: case Nil{} _: T.W{0, 0, 0, 0} case w <> rest 0n: w case w <> rest 1n+p: get_word(rest, p)def schedule_core(k: Nat, big: Bool, +q: Nat, temp: T.Quad) -> T.Quad: match k: case 0n: xor_word(sub_word(rot_word(temp)), rcon(q)) case 1n: temp case 2n: temp case 3n: temp case 4n: match big: case True{}: sub_word(temp) case False{}: temp case 5n+p: temp# history is the words so far, the newest first.def next_word(+nk: Nat, +i: Nat, +history: List<&2, T.Quad>) -> T.Quad: xor_word(get_word(history, Nat.sub(nk, 1n)), schedule_core(Nat.mod(i, nk), Nat.is_lt(6n, nk), Nat.div(i, nk), get_word(history, 0n)))def grow(n: Nat, +nk: Nat, +i: Nat, +history: List<&2, T.Quad>) -> List<&2, T.Quad>: match n: case 0n: history case 1n+p: grow(p, nk, 1n+i, next_word(nk, i, history) <> history)# Four consecutive words make a round key.def round_keys(ws: List<&2, T.Quad>) -> List<&2, T.State>: match ws: case w0 <> w1 <> w2 <> w3 <> rest: T.S{w0, w1, w2, w3} <> round_keys(rest) case _: Nil{}def expand(+nk: Nat, +nr: Nat, key: List<&2, U32>) -> List<&2, T.State>: round_keys(List.reverse(&2, T.Quad, grow(Nat.sub(Nat.mul(4n, 1n+nr), nk), nk, nk, List.reverse(&2, T.Quad, key_words(key)))))def encrypt(+nk: Nat, +nr: Nat, key: List<&2, U32>, block: T.State) -> T.State: cipher(nr, block, expand(nk, nr, key))# ---- byte API ----# An expanded key: the number of rounds and the round keys.type Schedule is Data: Schedule{rounds: Nat, keys: List<&2, T.State>}# FIPS 197 Figure 4: a key of 16, 24 or 32 bytes; Nk = its length in words,# Nr = Nk + 6.def valid_length(+n: Nat) -> Bool: Bool.or(Nat.is_eq(n, 16n), Bool.or(Nat.is_eq(n, 24n), Nat.is_eq(n, 32n)))def schedule_if(ok: Bool, +nk: Nat, +key: List<&2, U32>) -> Maybe<&2, Schedule>: match ok: case True{}: Some{Schedule{Nat.add(nk, 6n), expand(nk, Nat.add(nk, 6n), key)}} case False{}: None{}# A 16-, 24- or 32-byte key (AES-128, AES-192, AES-256); None otherwise.def expand_key(+key: List<&2, U32>) -> Maybe<&2, Schedule>: +n = List.length(&2, U32, key) schedule_if(valid_length(n), Nat.div(n, 4n), key)def encrypt_state(sched: Schedule, s: T.State) -> T.State: match sched: case Schedule{nr, ks}: cipher(nr, s, ks)def state_of(bytes: List<&2, U32>) -> Maybe<&2, T.State>: match bytes: case a0 <> a1 <> a2 <> a3 <> b0 <> b1 <> b2 <> b3 <> c0 <> c1 <> c2 <> c3 <> d0 <> d1 <> d2 <> d3 <> Nil{}: Some{T.S{T.W{a0, a1, a2, a3}, T.W{b0, b1, b2, b3}, T.W{c0, c1, c2, c3}, T.W{d0, d1, d2, d3}}} case _: None{}def bytes_of(s: T.State) -> List<&2, U32>: match s: case T.S{T.W{a0, a1, a2, a3}, T.W{b0, b1, b2, b3}, T.W{c0, c1, c2, c3}, T.W{d0, d1, d2, d3}}: [a0, a1, a2, a3, b0, b1, b2, b3, c0, c1, c2, c3, d0, d1, d2, d3]def encrypt_bytes(sched: Schedule, m: Maybe<&2, T.State>) -> Maybe<&2, List<&2, U32>>: match m: case None{}: None{} case Some{s}: Some{bytes_of(encrypt_state(sched, s))}# Encrypts one 16-byte block; None for another length.def encrypt_block(sched: Schedule, block: List<&2, U32>) -> Maybe<&2, List<&2, U32>>: encrypt_bytes(sched, state_of(block))def hex_digit_if(x: U32, small: Bool) -> Char: match small: case True{}: Chr{U32.add(48, x)} case False{}: Chr{U32.add(87, x)}def hex_digit(+x: U32) -> Char: hex_digit_if(x, U32.is_lt(x, 10))# Lowercase hexadecimal of a byte list.def hex(bytes: List<&2, U32>) -> String: match bytes: case Nil{}: "" case +b <> rest: SCon{hex_digit(U32.and(15, U32.shrn(b, 4n))), SCon{hex_digit(U32.and(15, b)), hex(rest)}}