/*
 *   AES-128 and AES-CMAC (RFC 4493).
 *
 *   Portions Copyright RISC OS Developments 2019+, credited to the RISC OS One Project.
 *   Placed in the public domain.
 */

#include <string.h>

#include "aes.h"

static const unsigned char sbox[256] =
{
    0x63,0x7c,0x77,0x7b,0xf2,0x6b,0x6f,0xc5,0x30,0x01,0x67,0x2b,0xfe,0xd7,0xab,0x76,
    0xca,0x82,0xc9,0x7d,0xfa,0x59,0x47,0xf0,0xad,0xd4,0xa2,0xaf,0x9c,0xa4,0x72,0xc0,
    0xb7,0xfd,0x93,0x26,0x36,0x3f,0xf7,0xcc,0x34,0xa5,0xe5,0xf1,0x71,0xd8,0x31,0x15,
    0x04,0xc7,0x23,0xc3,0x18,0x96,0x05,0x9a,0x07,0x12,0x80,0xe2,0xeb,0x27,0xb2,0x75,
    0x09,0x83,0x2c,0x1a,0x1b,0x6e,0x5a,0xa0,0x52,0x3b,0xd6,0xb3,0x29,0xe3,0x2f,0x84,
    0x53,0xd1,0x00,0xed,0x20,0xfc,0xb1,0x5b,0x6a,0xcb,0xbe,0x39,0x4a,0x4c,0x58,0xcf,
    0xd0,0xef,0xaa,0xfb,0x43,0x4d,0x33,0x85,0x45,0xf9,0x02,0x7f,0x50,0x3c,0x9f,0xa8,
    0x51,0xa3,0x40,0x8f,0x92,0x9d,0x38,0xf5,0xbc,0xb6,0xda,0x21,0x10,0xff,0xf3,0xd2,
    0xcd,0x0c,0x13,0xec,0x5f,0x97,0x44,0x17,0xc4,0xa7,0x7e,0x3d,0x64,0x5d,0x19,0x73,
    0x60,0x81,0x4f,0xdc,0x22,0x2a,0x90,0x88,0x46,0xee,0xb8,0x14,0xde,0x5e,0x0b,0xdb,
    0xe0,0x32,0x3a,0x0a,0x49,0x06,0x24,0x5c,0xc2,0xd3,0xac,0x62,0x91,0x95,0xe4,0x79,
    0xe7,0xc8,0x37,0x6d,0x8d,0xd5,0x4e,0xa9,0x6c,0x56,0xf4,0xea,0x65,0x7a,0xae,0x08,
    0xba,0x78,0x25,0x2e,0x1c,0xa6,0xb4,0xc6,0xe8,0xdd,0x74,0x1f,0x4b,0xbd,0x8b,0x8a,
    0x70,0x3e,0xb5,0x66,0x48,0x03,0xf6,0x0e,0x61,0x35,0x57,0xb9,0x86,0xc1,0x1d,0x9e,
    0xe1,0xf8,0x98,0x11,0x69,0xd9,0x8e,0x94,0x9b,0x1e,0x87,0xe9,0xce,0x55,0x28,0xdf,
    0x8c,0xa1,0x89,0x0d,0xbf,0xe6,0x42,0x68,0x41,0x99,0x2d,0x0f,0xb0,0x54,0xbb,0x16
};

/* Round constants, one per key schedule round */
static const unsigned char rcon[AES_ROUNDS] =
{
    0x01,0x02,0x04,0x08,0x10,0x20,0x40,0x80,0x1b,0x36
};

/* Multiply by x in GF(2^8), reducing by the AES polynomial */
static unsigned char xtime(unsigned char a)
{
    return (unsigned char) ((a << 1) ^ ((a & 0x80) ? 0x1b : 0x00));
}

void AESSetKey(aes_key_t *k, const unsigned char key[AES_KEYLEN])
{
    unsigned char *w;
    unsigned char t[4];
    int i, j;

    w = k->round_key;
    memcpy(w, key, AES_KEYLEN);

    for(i = 1; i <= AES_ROUNDS; i++)
    {
        unsigned char *prev = w + (i - 1) * AES_BLOCKLEN;
        unsigned char *cur  = w + i * AES_BLOCKLEN;

        /* RotWord then SubWord of the previous last column, plus Rcon */
        t[0] = (unsigned char) (sbox[prev[13]] ^ rcon[i - 1]);
        t[1] = sbox[prev[14]];
        t[2] = sbox[prev[15]];
        t[3] = sbox[prev[12]];

        for(j = 0; j < 4; j++)
            cur[j] = (unsigned char) (prev[j] ^ t[j]);
        for(j = 4; j < AES_BLOCKLEN; j++)
            cur[j] = (unsigned char) (prev[j] ^ cur[j - 4]);
    }
    memset(t, 0, sizeof(t));
}

void AESEncryptBlock(const aes_key_t *k,
                     const unsigned char in[AES_BLOCKLEN],
                     unsigned char out[AES_BLOCKLEN])
{
    unsigned char s[AES_BLOCKLEN];
    unsigned char a, b, c, d, e;
    const unsigned char *rk;
    int round, i;

    rk = k->round_key;
    for(i = 0; i < AES_BLOCKLEN; i++)
        s[i] = (unsigned char) (in[i] ^ rk[i]);

    for(round = 1; round <= AES_ROUNDS; round++)
    {
        /* SubBytes */
        for(i = 0; i < AES_BLOCKLEN; i++)
            s[i] = sbox[s[i]];

        /* ShiftRows.  The state is held in column order, so row r is the
           bytes at r, r+4, r+8, r+12, and is rotated left by r. */
        a = s[1];  s[1]  = s[5];  s[5]  = s[9];  s[9]  = s[13]; s[13] = a;
        a = s[2];  b = s[6];  s[2]  = s[10]; s[6]  = s[14]; s[10] = a; s[14] = b;
        a = s[15]; s[15] = s[11]; s[11] = s[7];  s[7]  = s[3];  s[3]  = a;

        /* MixColumns, except in the last round */
        if(round != AES_ROUNDS)
        {
            for(i = 0; i < AES_BLOCKLEN; i += 4)
            {
                a = s[i]; b = s[i+1]; c = s[i+2]; d = s[i+3];
                e = (unsigned char) (a ^ b ^ c ^ d);
                s[i]   = (unsigned char) (a ^ e ^ xtime((unsigned char)(a ^ b)));
                s[i+1] = (unsigned char) (b ^ e ^ xtime((unsigned char)(b ^ c)));
                s[i+2] = (unsigned char) (c ^ e ^ xtime((unsigned char)(c ^ d)));
                s[i+3] = (unsigned char) (d ^ e ^ xtime((unsigned char)(d ^ a)));
            }
        }

        rk = k->round_key + round * AES_BLOCKLEN;
        for(i = 0; i < AES_BLOCKLEN; i++)
            s[i] = (unsigned char) (s[i] ^ rk[i]);
    }
    memcpy(out, s, AES_BLOCKLEN);
    memset(s, 0, sizeof(s));
}

/* One left shift of a 128 bit value, with the CMAC reduction applied when a
   bit falls off the top */
static void cmac_shift(const unsigned char in[AES_BLOCKLEN],
                       unsigned char out[AES_BLOCKLEN])
{
    int i, carry;

    carry = in[0] >> 7;
    for(i = 0; i < AES_BLOCKLEN - 1; i++)
        out[i] = (unsigned char) ((in[i] << 1) | (in[i + 1] >> 7));
    out[AES_BLOCKLEN - 1] = (unsigned char) (in[AES_BLOCKLEN - 1] << 1);
    if(carry)
        out[AES_BLOCKLEN - 1] ^= 0x87;
}

void aes_cmac(const unsigned char key[AES_KEYLEN],
              const void *data, size_t len,
              unsigned char mac[AES_BLOCKLEN])
{
    aes_key_t k;
    unsigned char zero[AES_BLOCKLEN];
    unsigned char l[AES_BLOCKLEN], k1[AES_BLOCKLEN], k2[AES_BLOCKLEN];
    unsigned char x[AES_BLOCKLEN], block[AES_BLOCKLEN];
    const unsigned char *p;
    size_t whole;
    int i, last_is_whole;

    AESSetKey(&k, key);

    /* Subkeys, from the encryption of an all zero block */
    memset(zero, 0, sizeof(zero));
    AESEncryptBlock(&k, zero, l);
    cmac_shift(l, k1);
    cmac_shift(k1, k2);

    /* How many whole blocks come before the last one.  A message that is a
       non-zero multiple of the block size ends on a whole block; anything
       else, including an empty message, ends on a padded one. */
    last_is_whole = (len > 0) && ((len % AES_BLOCKLEN) == 0);
    if(last_is_whole)
        whole = (len / AES_BLOCKLEN) - 1;
    else
        whole = len / AES_BLOCKLEN;

    memset(x, 0, sizeof(x));
    p = (const unsigned char *) data;
    while(whole > 0)
    {
        for(i = 0; i < AES_BLOCKLEN; i++)
            block[i] = (unsigned char) (x[i] ^ p[i]);
        AESEncryptBlock(&k, block, x);
        p += AES_BLOCKLEN;
        len -= AES_BLOCKLEN;
        whole--;
    }

    /* The last block, XORed with K1 whole or K2 padded */
    memset(block, 0, sizeof(block));
    if(last_is_whole)
    {
        for(i = 0; i < AES_BLOCKLEN; i++)
            block[i] = (unsigned char) (p[i] ^ k1[i]);
    }
    else
    {
        for(i = 0; i < (int) len; i++)
            block[i] = p[i];
        block[len] = 0x80;
        for(i = 0; i < AES_BLOCKLEN; i++)
            block[i] = (unsigned char) (block[i] ^ k2[i]);
    }
    for(i = 0; i < AES_BLOCKLEN; i++)
        block[i] = (unsigned char) (block[i] ^ x[i]);
    AESEncryptBlock(&k, block, mac);

    memset(&k, 0, sizeof(k));
    memset(l, 0, sizeof(l));  memset(k1, 0, sizeof(k1)); memset(k2, 0, sizeof(k2));
    memset(x, 0, sizeof(x));  memset(block, 0, sizeof(block));
}


/*
    AES-CCM, which is what SMB 3.0 and 3.0.2 encrypt with.

    Counter mode for the data and CBC-MAC for the tag, both built out of
    the forward cipher alone - so nothing here needs the inverse cipher,
    and the block encryption above is the whole of the arithmetic.

    The nonce is shorter than a block, and what is left over counts: with
    an 11 byte nonce, as SMB uses, four bytes remain to hold a length or a
    counter, which is where the q below comes from.

    Lengths are held in size_t and encoded big endian into those q bytes;
    the caller's data must therefore fit in them, which for q = 4 is no
    real limit for a protocol whose messages are bounded anyway.
*/

/* The counter block for keystream position i.  Position 0 makes the block
   that masks the tag; 1 upwards mask the data. */
static void ccm_ctr_block(unsigned char a[AES_BLOCKLEN],
                          const unsigned char *nonce, size_t nonce_len,
                          unsigned int counter)
{
    size_t q, i;

    q = (size_t) AES_BLOCKLEN - 1 - nonce_len;
    memset(a, 0, AES_BLOCKLEN);
    a[0] = (unsigned char) (q - 1);
    memcpy(a + 1, nonce, nonce_len);
    for(i = 0; i < q; i++)
        a[AES_BLOCKLEN - 1 - i] = (unsigned char) ((counter >> (8 * i)) & 0xFF);
}

/* Feed bytes through the CBC-MAC, zero padded up to a whole block */
static void ccm_absorb(const aes_key_t *k, unsigned char x[AES_BLOCKLEN],
                       const unsigned char *p, size_t len)
{
    unsigned char block[AES_BLOCKLEN];
    size_t take;
    int i;

    while(len > 0)
    {
        take = (len < (size_t) AES_BLOCKLEN) ? len : (size_t) AES_BLOCKLEN;
        memset(block, 0, sizeof(block));
        memcpy(block, p, take);
        for(i = 0; i < AES_BLOCKLEN; i++)
            block[i] = (unsigned char) (block[i] ^ x[i]);
        AESEncryptBlock(k, block, x);
        p += take;
        len -= take;
    }
    memset(block, 0, sizeof(block));
}

/*
    The CBC-MAC over the whole of what is being protected: a first block
    describing the message, then the additional data, then the payload.

    The additional data is prefixed with its own length and the two are
    padded together as one run, not separately - padding after the length
    would give a different tag, and the server would reject every message.
*/
static void ccm_mac(const aes_key_t *k,
                    const unsigned char *nonce, size_t nonce_len,
                    const unsigned char *aad, size_t aad_len,
                    const unsigned char *data, size_t len,
                    unsigned char x[AES_BLOCKLEN])
{
    unsigned char block[AES_BLOCKLEN];
    size_t q, i, take;

    q = (size_t) AES_BLOCKLEN - 1 - nonce_len;

    memset(block, 0, sizeof(block));
    block[0] = (unsigned char) ((aad_len ? 0x40 : 0x00)      /* data follows */
                                | (((AES_BLOCKLEN - 2) / 2) << 3)  /* 16 byte tag */
                                | (q - 1));
    memcpy(block + 1, nonce, nonce_len);
    for(i = 0; i < q; i++)
        block[AES_BLOCKLEN - 1 - i] = (unsigned char) ((len >> (8 * i)) & 0xFF);

    memset(x, 0, AES_BLOCKLEN);
    ccm_absorb(k, x, block, AES_BLOCKLEN);

    if(aad_len > 0)
    {
        memset(block, 0, sizeof(block));
        block[0] = (unsigned char) ((aad_len >> 8) & 0xFF);
        block[1] = (unsigned char) (aad_len & 0xFF);
        take = (aad_len < (size_t) AES_BLOCKLEN - 2)
                   ? aad_len : (size_t) AES_BLOCKLEN - 2;
        memcpy(block + 2, aad, take);
        ccm_absorb(k, x, block, AES_BLOCKLEN);
        if(aad_len > take)
            ccm_absorb(k, x, aad + take, aad_len - take);
    }

    ccm_absorb(k, x, data, len);
    memset(block, 0, sizeof(block));
}

/* Counter mode over the payload, which is its own inverse */
static void ccm_crypt(const aes_key_t *k,
                      const unsigned char *nonce, size_t nonce_len,
                      unsigned char *data, size_t len)
{
    unsigned char a[AES_BLOCKLEN], s[AES_BLOCKLEN];
    unsigned int counter;
    size_t take, i;

    counter = 1;
    while(len > 0)
    {
        ccm_ctr_block(a, nonce, nonce_len, counter);
        AESEncryptBlock(k, a, s);
        take = (len < (size_t) AES_BLOCKLEN) ? len : (size_t) AES_BLOCKLEN;
        for(i = 0; i < take; i++)
            data[i] = (unsigned char) (data[i] ^ s[i]);
        data += take;
        len -= take;
        counter++;
    }
    memset(a, 0, sizeof(a));
    memset(s, 0, sizeof(s));
}

void aes_ccm_encrypt(const unsigned char key[AES_KEYLEN],
                     const unsigned char *nonce, size_t nonce_len,
                     const unsigned char *aad, size_t aad_len,
                     unsigned char *data, size_t len,
                     unsigned char tag[AES_BLOCKLEN])
{
    aes_key_t k;
    unsigned char x[AES_BLOCKLEN], a0[AES_BLOCKLEN], s0[AES_BLOCKLEN];
    int i;

    AESSetKey(&k, key);
    /* The tag is taken over the plain text, so it is computed first */
    ccm_mac(&k, nonce, nonce_len, aad, aad_len, data, len, x);
    ccm_ctr_block(a0, nonce, nonce_len, 0);
    AESEncryptBlock(&k, a0, s0);
    for(i = 0; i < AES_BLOCKLEN; i++)
        tag[i] = (unsigned char) (x[i] ^ s0[i]);
    ccm_crypt(&k, nonce, nonce_len, data, len);

    memset(&k, 0, sizeof(k));
    memset(x, 0, sizeof(x));
    memset(a0, 0, sizeof(a0));
    memset(s0, 0, sizeof(s0));
}

/*
    Undo that, and say whether the tag was right.

    Returns 0 and leaves the buffer decrypted but untrustworthy if it was
    not: the caller must throw the message away rather than read it.  The
    comparison is over the whole tag whatever it finds, so that how much
    of it matched cannot be timed.
*/
int aes_ccm_decrypt(const unsigned char key[AES_KEYLEN],
                    const unsigned char *nonce, size_t nonce_len,
                    const unsigned char *aad, size_t aad_len,
                    unsigned char *data, size_t len,
                    const unsigned char tag[AES_BLOCKLEN])
{
    aes_key_t k;
    unsigned char x[AES_BLOCKLEN], a0[AES_BLOCKLEN], s0[AES_BLOCKLEN];
    int i, diff;

    AESSetKey(&k, key);
    ccm_crypt(&k, nonce, nonce_len, data, len);
    ccm_mac(&k, nonce, nonce_len, aad, aad_len, data, len, x);
    ccm_ctr_block(a0, nonce, nonce_len, 0);
    AESEncryptBlock(&k, a0, s0);

    diff = 0;
    for(i = 0; i < AES_BLOCKLEN; i++)
        diff |= (int) (unsigned char) (tag[i] ^ (x[i] ^ s0[i]));

    memset(&k, 0, sizeof(k));
    memset(x, 0, sizeof(x));
    memset(a0, 0, sizeof(a0));
    memset(s0, 0, sizeof(s0));
    return diff == 0;
}
