/*copyright*/
/*********************************************************************
 * ASCON
 *********************************************************************/
#include <stdint.h>
//#include "api.h"
#define CRYPTO_VERSION "1.3.0"
#define CRYPTO_KEYBYTES 16
#define CRYPTO_NSECBYTES 0
#define CRYPTO_NPUBBYTES 16

//#define CRYPTO_ABYTES 16
#define CRYPTO_ABYTES 8   // do not have extra space in packet

#define CRYPTO_NOOVERLAP 1
#define ASCON_AEAD_RATE 16
#define ASCON_VARIANT 1
//#include "ascon.h"
typedef struct {
  uint64_t x[5];
} ascon_state_t;
//#include "crypto_aead.h"

//#include "permutations.h"
//#include "constants.h"
#define ASCON_80PQ_VARIANT 0
#define ASCON_AEAD_VARIANT 1
#define ASCON_HASH_VARIANT 2
#define ASCON_XOF_VARIANT 3
#define ASCON_CXOF_VARIANT 4
#define ASCON_MAC_VARIANT 5
#define ASCON_PRF_VARIANT 6
#define ASCON_PRFS_VARIANT 7

#define ASCON_TAG_SIZE 16
#define ASCON_HASH_SIZE 32

#define ASCON_128_RATE 8
#define ASCON_128A_RATE 16
#define ASCON_HASH_RATE 8
#define ASCON_PRF_IN_RATE 32
#define ASCON_PRFA_IN_RATE 40
#define ASCON_PRF_OUT_RATE 16

#define ASCON_PA_ROUNDS 12
#define ASCON_128_PB_ROUNDS 6
#define ASCON_128A_PB_ROUNDS 8
#define ASCON_HASH_PB_ROUNDS 12
#define ASCON_HASHA_PB_ROUNDS 8
#define ASCON_PRF_PB_ROUNDS 12
#define ASCON_PRFA_PB_ROUNDS 8

#define ASCON_128_IV                         \
  (((uint64_t)(ASCON_AEAD_VARIANT) << 0) |   \
   ((uint64_t)(ASCON_PA_ROUNDS) << 16) |     \
   ((uint64_t)(ASCON_128_PB_ROUNDS) << 20) | \
   ((uint64_t)(ASCON_TAG_SIZE * 8) << 24) |  \
   ((uint64_t)(ASCON_128_RATE) << 40))

#define ASCON_128A_IV                         \
  (((uint64_t)(ASCON_AEAD_VARIANT) << 0) |    \
   ((uint64_t)(ASCON_PA_ROUNDS) << 16) |      \
   ((uint64_t)(ASCON_128A_PB_ROUNDS) << 20) | \
   ((uint64_t)(ASCON_TAG_SIZE * 8) << 24) |   \
   ((uint64_t)(ASCON_128A_RATE) << 40))

#define ASCON_80PQ_IV                                                          \
  (((uint64_t)(ASCON_80PQ_VARIANT) << 0) | ((uint64_t)(ASCON_128_RATE) << 8) | \
   ((uint64_t)(ASCON_PA_ROUNDS) << 16) |                                       \
   ((uint64_t)(ASCON_128_PB_ROUNDS) << 20) |                                   \
   ((uint64_t)(ASCON_TAG_SIZE * 8) << 24))

#define ASCON_HASH_IV                         \
  (((uint64_t)(ASCON_HASH_VARIANT) << 0) |    \
   ((uint64_t)(ASCON_PA_ROUNDS) << 16) |      \
   ((uint64_t)(ASCON_HASH_PB_ROUNDS) << 20) | \
   ((uint64_t)(ASCON_HASH_SIZE * 8) << 24) |  \
   ((uint64_t)(ASCON_HASH_RATE) << 40))

#define ASCON_HASHA_IV                         \
  (((uint64_t)(ASCON_HASH_VARIANT) << 0) |     \
   ((uint64_t)(ASCON_PA_ROUNDS) << 16) |       \
   ((uint64_t)(ASCON_HASHA_PB_ROUNDS) << 20) | \
   ((uint64_t)(ASCON_HASH_SIZE * 8) << 24) |   \
   ((uint64_t)(ASCON_HASH_RATE) << 40))

#define ASCON_XOF_IV                          \
  (((uint64_t)(ASCON_XOF_VARIANT) << 0) |     \
   ((uint64_t)(ASCON_PA_ROUNDS) << 16) |      \
   ((uint64_t)(ASCON_HASH_PB_ROUNDS) << 20) | \
   ((uint64_t)(ASCON_HASH_RATE) << 40))

#define ASCON_XOFA_IV                          \
  (((uint64_t)(ASCON_XOF_VARIANT) << 0) |      \
   ((uint64_t)(ASCON_PA_ROUNDS) << 16) |       \
   ((uint64_t)(ASCON_HASHA_PB_ROUNDS) << 20) | \
   ((uint64_t)(ASCON_HASH_RATE) << 40))

#define ASCON_CXOF_IV                         \
  (((uint64_t)(ASCON_CXOF_VARIANT) << 0) |    \
   ((uint64_t)(ASCON_PA_ROUNDS) << 16) |      \
   ((uint64_t)(ASCON_HASH_PB_ROUNDS) << 20) | \
   ((uint64_t)(ASCON_HASH_RATE) << 40))

#define ASCON_CXOFA_IV                         \
  (((uint64_t)(ASCON_CXOF_VARIANT) << 0) |     \
   ((uint64_t)(ASCON_PA_ROUNDS) << 16) |       \
   ((uint64_t)(ASCON_HASHA_PB_ROUNDS) << 20) | \
   ((uint64_t)(ASCON_HASH_RATE) << 40))

#define ASCON_MAC_IV                         \
  (((uint64_t)(ASCON_MAC_VARIANT) << 0) |    \
   ((uint64_t)(ASCON_PA_ROUNDS) << 16) |     \
   ((uint64_t)(ASCON_PRF_PB_ROUNDS) << 20) | \
   ((uint64_t)(ASCON_TAG_SIZE * 8) << 24) |  \
   ((uint64_t)(ASCON_PRF_IN_RATE) << 40) |   \
   ((uint64_t)(ASCON_PRF_OUT_RATE) << 48))

#define ASCON_MACA_IV                         \
  (((uint64_t)(ASCON_MAC_VARIANT) << 0) |     \
   ((uint64_t)(ASCON_PA_ROUNDS) << 16) |      \
   ((uint64_t)(ASCON_PRFA_PB_ROUNDS) << 20) | \
   ((uint64_t)(ASCON_TAG_SIZE * 8) << 24) |   \
   ((uint64_t)(ASCON_PRFA_IN_RATE) << 40) |   \
   ((uint64_t)(ASCON_PRF_OUT_RATE) << 48))

#define ASCON_PRF_IV                         \
  (((uint64_t)(ASCON_PRF_VARIANT) << 0) |    \
   ((uint64_t)(ASCON_PA_ROUNDS) << 16) |     \
   ((uint64_t)(ASCON_PRF_PB_ROUNDS) << 20) | \
   ((uint64_t)(ASCON_PRF_IN_RATE) << 40) |   \
   ((uint64_t)(ASCON_PRF_OUT_RATE) << 48))

#define ASCON_PRFA_IV                         \
  (((uint64_t)(ASCON_PRF_VARIANT) << 0) |     \
   ((uint64_t)(ASCON_PA_ROUNDS) << 16) |      \
   ((uint64_t)(ASCON_PRFA_PB_ROUNDS) << 20) | \
   ((uint64_t)(ASCON_PRFA_IN_RATE) << 40) |   \
   ((uint64_t)(ASCON_PRF_OUT_RATE) << 48))

#define ASCON_PRFS_IV                      \
  (((uint64_t)(ASCON_PRFS_VARIANT) << 0) | \
   ((uint64_t)(ASCON_PA_ROUNDS) << 16) |   \
   ((uint64_t)(ASCON_TAG_SIZE * 8) << 24))

//#include "printstate.h"
//#include "word.h"
/* get byte from 64-bit Ascon word */
#define GETBYTE(x, i) ((uint8_t)((uint64_t)(x) >> (8 * (i))))

/* set byte in 64-bit Ascon word */
#define SETBYTE(b, i) ((uint64_t)(b) << (8 * (i)))

/* set padding byte in 64-bit Ascon word */
#define PAD(i) SETBYTE(0x01, i)

/* define domain separation bit in 64-bit Ascon word */
#define DSEP() SETBYTE(0x80, 7)

/* load bytes into 64-bit Ascon word */
static inline uint64_t LOADBYTES(const uint8_t* bytes, int n) {
  int i;
  uint64_t x = 0;
  for (i = 0; i < n; ++i) x |= SETBYTE(bytes[i], i);
  return x;
}

/* store bytes from 64-bit Ascon word */
static inline void STOREBYTES(uint8_t* bytes, uint64_t x, int n) {
  int i;
  for (i = 0; i < n; ++i) bytes[i] = GETBYTE(x, i);
}

/* clear bytes in 64-bit Ascon word */
static inline uint64_t CLEARBYTES(uint64_t x, int n) {
  int i;
  for (i = 0; i < n; ++i) x &= ~SETBYTE(0xff, i);
  return x;
}

#ifdef ASCON_PRINT_STATE

//#include "ascon.h"
//#include "word.h"

void print(const char* text);
void printbytes(const char* text, const uint8_t* b, uint64_t len);
void printword(const char* text, const uint64_t x);
void printstate(const char* text, const ascon_state_t* s);

#else

#define print(text) \
  do {              \
  } while (0)

#define printbytes(text, b, l) \
  do {                         \
  } while (0)

#define printword(text, w) \
  do {                     \
  } while (0)

#define printstate(text, s) \
  do {                      \
  } while (0)

#endif

//#include "round.h"
static inline uint64_t ROR(uint64_t x, int n) {
  return x >> n | x << (-n & 63);
}

static inline void ROUND(ascon_state_t* s, uint8_t C) {
  ascon_state_t t;
  /* addition of round constant */
  s->x[2] ^= C;
  /* printstate(" round constant", s); */
  /* substitution layer */
  s->x[0] ^= s->x[4];
  s->x[4] ^= s->x[3];
  s->x[2] ^= s->x[1];
  /* start of keccak s-box */
  t.x[0] = s->x[0] ^ (~s->x[1] & s->x[2]);
  t.x[1] = s->x[1] ^ (~s->x[2] & s->x[3]);
  t.x[2] = s->x[2] ^ (~s->x[3] & s->x[4]);
  t.x[3] = s->x[3] ^ (~s->x[4] & s->x[0]);
  t.x[4] = s->x[4] ^ (~s->x[0] & s->x[1]);
  /* end of keccak s-box */
  t.x[1] ^= t.x[0];
  t.x[0] ^= t.x[4];
  t.x[3] ^= t.x[2];
  t.x[2] = ~t.x[2];
  /* printstate(" substitution layer", &t); */
  /* linear diffusion layer */
  s->x[0] = t.x[0] ^ ROR(t.x[0], 19) ^ ROR(t.x[0], 28);
  s->x[1] = t.x[1] ^ ROR(t.x[1], 61) ^ ROR(t.x[1], 39);
  s->x[2] = t.x[2] ^ ROR(t.x[2], 1) ^ ROR(t.x[2], 6);
  s->x[3] = t.x[3] ^ ROR(t.x[3], 10) ^ ROR(t.x[3], 17);
  s->x[4] = t.x[4] ^ ROR(t.x[4], 7) ^ ROR(t.x[4], 41);
  printstate(" round output", s);
}



static inline void P12(ascon_state_t* s) {
  ROUND(s, 0xf0);
  ROUND(s, 0xe1);
  ROUND(s, 0xd2);
  ROUND(s, 0xc3);
  ROUND(s, 0xb4);
  ROUND(s, 0xa5);
  ROUND(s, 0x96);
  ROUND(s, 0x87);
  ROUND(s, 0x78);
  ROUND(s, 0x69);
  ROUND(s, 0x5a);
  ROUND(s, 0x4b);
}

static inline void P8(ascon_state_t* s) {
  ROUND(s, 0xb4);
  ROUND(s, 0xa5);
  ROUND(s, 0x96);
  ROUND(s, 0x87);
  ROUND(s, 0x78);
  ROUND(s, 0x69);
  ROUND(s, 0x5a);
  ROUND(s, 0x4b);
}

static inline void P6(ascon_state_t* s) {
  ROUND(s, 0x96);
  ROUND(s, 0x87);
  ROUND(s, 0x78);
  ROUND(s, 0x69);
  ROUND(s, 0x5a);
  ROUND(s, 0x4b);
}
//#include "printstate.h"
//#include "word.h"

int crypto_aead_encrypt(unsigned char* c, unsigned long long* clen,
                        const unsigned char* m, unsigned long long mlen,
                        const unsigned char* ad, unsigned long long adlen,
                        const unsigned char* nsec, const unsigned char* npub,
                        const unsigned char* k) {
  (void)nsec;

  /* set ciphertext size */
  *clen = mlen + CRYPTO_ABYTES;

  /* print input bytes */
  print("encrypt\n");
  printbytes("k", k, CRYPTO_KEYBYTES);
  printbytes("n", npub, CRYPTO_NPUBBYTES);
  printbytes("a", ad, adlen);
  printbytes("m", m, mlen);

  /* load key and nonce */
  const uint64_t K0 = LOADBYTES(k, 8);
  const uint64_t K1 = LOADBYTES(k + 8, 8);
  const uint64_t N0 = LOADBYTES(npub, 8);
  const uint64_t N1 = LOADBYTES(npub + 8, 8);

  /* initialize */
  ascon_state_t s;
  s.x[0] = ASCON_128A_IV;
  s.x[1] = K0;
  s.x[2] = K1;
  s.x[3] = N0;
  s.x[4] = N1;
  printstate("init 1st key xor", &s);
  P12(&s);
  s.x[3] ^= K0;
  s.x[4] ^= K1;
  printstate("init 2nd key xor", &s);

  if (adlen) {
    /* full associated data blocks */
    while (adlen >= ASCON_128A_RATE) {
      s.x[0] ^= LOADBYTES(ad, 8);
      s.x[1] ^= LOADBYTES(ad + 8, 8);
      printstate("absorb adata", &s);
      P8(&s);
      ad += ASCON_128A_RATE;
      adlen -= ASCON_128A_RATE;
    }
    /* final associated data block */
    if (adlen >= 8) {
      s.x[0] ^= LOADBYTES(ad, 8);
      s.x[1] ^= LOADBYTES(ad + 8, adlen - 8);
      s.x[1] ^= PAD(adlen - 8);
    } else {
      s.x[0] ^= LOADBYTES(ad, adlen);
      s.x[0] ^= PAD(adlen);
    }
    printstate("pad adata", &s);
    P8(&s);
  }
  /* domain separation */
  s.x[4] ^= DSEP();
  printstate("domain separation", &s);

  /* full plaintext blocks */
  while (mlen >= ASCON_128A_RATE) {
    s.x[0] ^= LOADBYTES(m, 8);
    s.x[1] ^= LOADBYTES(m + 8, 8);
    STOREBYTES(c, s.x[0], 8);
    STOREBYTES(c + 8, s.x[1], 8);
    printstate("absorb plaintext", &s);
    P8(&s);
    m += ASCON_128A_RATE;
    c += ASCON_128A_RATE;
    mlen -= ASCON_128A_RATE;
  }
  /* final plaintext block */
  if (mlen >= 8) {
    s.x[0] ^= LOADBYTES(m, 8);
    s.x[1] ^= LOADBYTES(m + 8, mlen - 8);
    STOREBYTES(c, s.x[0], 8);
    STOREBYTES(c + 8, s.x[1], mlen - 8);
    s.x[1] ^= PAD(mlen - 8);
  } else {
    s.x[0] ^= LOADBYTES(m, mlen);
    STOREBYTES(c, s.x[0], mlen);
    s.x[0] ^= PAD(mlen);
  }
  m += mlen;
  c += mlen;
  printstate("pad plaintext", &s);

  /* finalize */
  s.x[2] ^= K0;
  s.x[3] ^= K1;
  printstate("final 1st key xor", &s);
  P12(&s);
  s.x[3] ^= K0;
  s.x[4] ^= K1;
  printstate("final 2nd key xor", &s);

  /* get tag */
  STOREBYTES(c, s.x[3], 8);
  STOREBYTES(c + 8, s.x[4], 8);

  /* print output bytes */
  printbytes("c", c - *clen + CRYPTO_ABYTES, *clen - CRYPTO_ABYTES);
  printbytes("t", c, CRYPTO_ABYTES);
  print("\n");

  return 0;
}

int crypto_aead_decrypt(unsigned char* m, unsigned long long* mlen,
                        unsigned char* nsec, const unsigned char* c,
                        unsigned long long clen, const unsigned char* ad,
                        unsigned long long adlen, const unsigned char* npub,
                        const unsigned char* k) {
  (void)nsec;

  if (clen < CRYPTO_ABYTES) return -1;

  /* set plaintext size */
  *mlen = clen - CRYPTO_ABYTES;

  /* print input bytes */
  print("decrypt\n");
  printbytes("k", k, CRYPTO_KEYBYTES);
  printbytes("n", npub, CRYPTO_NPUBBYTES);
  printbytes("a", ad, adlen);
  printbytes("c", c, *mlen);
  printbytes("t", c + *mlen, CRYPTO_ABYTES);

  /* load key and nonce */
  const uint64_t K0 = LOADBYTES(k, 8);
  const uint64_t K1 = LOADBYTES(k + 8, 8);
  const uint64_t N0 = LOADBYTES(npub, 8);
  const uint64_t N1 = LOADBYTES(npub + 8, 8);

  /* initialize */
  ascon_state_t s;
  s.x[0] = ASCON_128A_IV;
  s.x[1] = K0;
  s.x[2] = K1;
  s.x[3] = N0;
  s.x[4] = N1;
  printstate("init 1st key xor", &s);
  P12(&s);
  s.x[3] ^= K0;
  s.x[4] ^= K1;
  printstate("init 2nd key xor", &s);

  if (adlen) {
    /* full associated data blocks */
    while (adlen >= ASCON_128A_RATE) {
      s.x[0] ^= LOADBYTES(ad, 8);
      s.x[1] ^= LOADBYTES(ad + 8, 8);
      printstate("absorb adata", &s);
      P8(&s);
      ad += ASCON_128A_RATE;
      adlen -= ASCON_128A_RATE;
    }
    /* final associated data block */
    if (adlen >= 8) {
      s.x[0] ^= LOADBYTES(ad, 8);
      s.x[1] ^= LOADBYTES(ad + 8, adlen - 8);
      s.x[1] ^= PAD(adlen - 8);
    } else {
      s.x[0] ^= LOADBYTES(ad, adlen);
      s.x[0] ^= PAD(adlen);
    }
    printstate("pad adata", &s);
    P8(&s);
  }
  /* domain separation */
  s.x[4] ^= DSEP();
  printstate("domain separation", &s);

  /* full ciphertext blocks */
  clen -= CRYPTO_ABYTES;
  while (clen >= ASCON_128A_RATE) {
    uint64_t c0 = LOADBYTES(c, 8);
    uint64_t c1 = LOADBYTES(c + 8, 8);
    STOREBYTES(m, s.x[0] ^ c0, 8);
    STOREBYTES(m + 8, s.x[1] ^ c1, 8);
    s.x[0] = c0;
    s.x[1] = c1;
    printstate("insert ciphertext", &s);
    P8(&s);
    m += ASCON_128A_RATE;
    c += ASCON_128A_RATE;
    clen -= ASCON_128A_RATE;
  }
  /* final ciphertext block */
  if (clen >= 8) {
    uint64_t c0 = LOADBYTES(c, 8);
    uint64_t c1 = LOADBYTES(c + 8, clen - 8);
    STOREBYTES(m, s.x[0] ^ c0, 8);
    STOREBYTES(m + 8, s.x[1] ^ c1, clen - 8);
    s.x[0] = c0;
    s.x[1] = CLEARBYTES(s.x[1], clen - 8);
    s.x[1] |= c1;
    s.x[1] ^= PAD(clen - 8);
  } else {
    uint64_t c0 = LOADBYTES(c, clen);
    STOREBYTES(m, s.x[0] ^ c0, clen);
    s.x[0] = CLEARBYTES(s.x[0], clen);
    s.x[0] |= c0;
    s.x[0] ^= PAD(clen);
  }
  m += clen;
  c += clen;
  printstate("pad ciphertext", &s);

  /* finalize */
  s.x[2] ^= K0;
  s.x[3] ^= K1;
  printstate("final 1st key xor", &s);
  P12(&s);
  s.x[3] ^= K0;
  s.x[4] ^= K1;
  printstate("final 2nd key xor", &s);

  /* get tag */
  uint8_t t[16];
  STOREBYTES(t, s.x[3], 8);
  STOREBYTES(t + 8, s.x[4], 8);
  
  /* verify should be constant time, check compiler output */
  int i;
  int result = 0;
  for (i = 0; i < CRYPTO_ABYTES; ++i) result |= c[i] ^ t[i];
  result = (((result - 1) >> 8) & 1) - 1;

  /* print output bytes */
  printbytes("m", m - *mlen, *mlen);
  print("\n");

  return result;
}

#define CRYPTO_BYTES 32
int crypto_hash(unsigned char* out, const unsigned char* in,
                unsigned long long len) {
  printbytes("m", in, len);
  /* initialize */
  ascon_state_t s;
  s.x[0] = ASCON_HASH_IV;
  s.x[1] = 0;
  s.x[2] = 0;
  s.x[3] = 0;
  s.x[4] = 0;
  printstate("initial value", &s);
  P12(&s);
  printstate("initialization", &s);

  /* absorb full plaintext blocks */
  while (len >= ASCON_HASH_RATE) {
    s.x[0] ^= LOADBYTES(in, 8);
    printstate("absorb plaintext", &s);
    P12(&s);
    in += ASCON_HASH_RATE;
    len -= ASCON_HASH_RATE;
  }
  /* absorb final plaintext block */
  s.x[0] ^= LOADBYTES(in, len);
  s.x[0] ^= PAD(len);
  printstate("pad plaintext", &s);
  P12(&s);

  /* squeeze full output blocks */
  len = CRYPTO_BYTES;
  while (len > ASCON_HASH_RATE) {
    STOREBYTES(out, s.x[0], 8);
    printstate("squeeze output", &s);
    P12(&s);
    out += ASCON_HASH_RATE;
    len -= ASCON_HASH_RATE;
  }
  /* squeeze final output block */
  STOREBYTES(out, s.x[0], len);
  printstate("squeeze output", &s);
  printbytes("h", out + len - CRYPTO_BYTES, CRYPTO_BYTES);

  return 0;
}

static int asconencryptinplace(uint8_t *in,int n,uint8_t *ki)
{
  // in=|pubnonce(4)|plaintext(n-4)|space for tag(16)|. ki=|privnonce(16)|key(16)|
  // return |pubnonce(4)|ciphertext(n-4)|tag(8)|
  unsigned long long clen=0;
  uint8_t nonce[16];
  memcpy(nonce,ki,16);
  nonce[0] ^= *in++;
  nonce[1] ^= *in++;
  nonce[2] ^= *in++;
  nonce[3] ^= *in++;
  n -= 4;
  if(crypto_aead_encrypt(in,&clen,in,n,NULL,0,(void*)0,nonce,&ki[16])) return -1;
  return clen+4;
}

static int ascondecryptinplace(uint8_t *in,int n,uint8_t *ki)
{
  // in=|pubnonce(4)|ciphertext(n-4-8)|tag(8)| ki=|privnonce(16)|key(16)|
  // return |pubnonce(4)|plaintext(n-4)|
  unsigned long long mlen=0;
  uint8_t nonce[16];
  memcpy(nonce,ki,16);
  nonce[0] ^= *in++;
  nonce[1] ^= *in++;
  nonce[2] ^= *in++;
  nonce[3] ^= *in++;
  n -= 4;
  if(crypto_aead_decrypt(in,&mlen,(void*)0,in,(uint64_t)n,NULL,0,nonce,&ki[16])) return -1;
  return mlen+4;
}
static int asconhash(uint8_t *out,uint8_t *in,int n)
{
  if(crypto_hash(out,in,(unsigned long long)n)) return -1;
  return 0;
}
