/*
 * @Description: 
 * @Author: yangsw
 * @Date: 2022-05-06 10:06:55
 * @LastEditTime: 2022-05-06 10:06:55
 * @LastEditors: yangsw
 * @Reference: 
 */
#include "util_aes.h"
#include <stdlib.h>


#define WPOLY   0x011b
#define BPOLY     0x1b
#define DPOLY   0x008d

#define f1(x)   (x)
#define f2(x)   ((x << 1) ^ (((x >> 7) & 1) * WPOLY))
#define f4(x)   ((x << 2) ^ (((x >> 6) & 1) * WPOLY) ^ (((x >> 6) & 2) * WPOLY))
#define f8(x)   ((x << 3) ^ (((x >> 5) & 1) * WPOLY) ^ (((x >> 5) & 2) * WPOLY) \
              ^ (((x >> 5) & 4) * WPOLY))
#define d2(x)   (((x) >> 1) ^ ((x) & 1 ? DPOLY : 0))

#define f3(x)   (f2(x) ^ x)
#define f9(x)   (f8(x) ^ x)
#define fb(x)   (f8(x) ^ f2(x) ^ x)
#define fd(x)   (f8(x) ^ f4(x) ^ x)
#define fe(x)   (f8(x) ^ f4(x) ^ f2(x))

static Aes_Context_Struct aes_context_struct;


static uint_8t hibit(const uint_8t x)
{   uint_8t r = (uint_8t)((x >> 1) | (x >> 2));

  r |= (r >> 2);
  r |= (r >> 4);
  return (r + 1) >> 1;
}

static uint_8t gf_inv(const uint_8t x)
{   uint_8t p1 = x, p2 = BPOLY, n1 = hibit(x), n2 = 0x80, v1 = 1, v2 = 0;

  if(x < 2)
    return x;

  for( ; ; )
  {
    if(n1)
      while(n2 >= n1)
      {
        n2 /= n1;
        p2 ^= (p1 * n2) & 0xff;
        v2 ^= (v1 * n2);
        n2 = hibit(p2);
      }
    else
      return v1;

    if(n2)
      while(n1 >= n2)
      {
        n1 /= n2;
        p1 ^= p2 * n1;
        v1 ^= v2 * n1;
        n1 = hibit(p1);
      }
    else
      return v2;
  }
}

uint_8t fwd_affine(const uint_8t x)
{
  return 0x63 ^ x ^ (x << 1) ^ (x << 2) ^ (x << 3) ^ (x << 4)
          ^ (x >> 7) ^ (x >> 6) ^ (x >> 5) ^ (x >> 4);
}

uint_8t inv_affine(const uint_8t x)
{
  return 0x05 ^ (x << 1) ^ (x << 3) ^ (x << 6)
        ^ (x >> 7) ^ (x >> 5) ^ (x >> 2);
}

#define s_box(x)   fwd_affine(gf_inv(x))
#define is_box(x)  gf_inv(inv_affine(x))
#define gfm2_sb(x) f2(s_box(x))
#define gfm3_sb(x) f3(s_box(x))
#define gfm_9(x)   f9(x)
#define gfm_b(x)   fb(x)
#define gfm_d(x)   fd(x)
#define gfm_e(x)   fe(x)


#define block_copy_nn(d, s, l)    copy_block_nn(d, s, l)
#define block_copy(d, s)          copy_block(d, s)

static void copy_block( void *d, const void *s )
{
  ((uint_8t*)d)[ 0] = ((uint_8t*)s)[ 0];
  ((uint_8t*)d)[ 1] = ((uint_8t*)s)[ 1];
  ((uint_8t*)d)[ 2] = ((uint_8t*)s)[ 2];
  ((uint_8t*)d)[ 3] = ((uint_8t*)s)[ 3];
  ((uint_8t*)d)[ 4] = ((uint_8t*)s)[ 4];
  ((uint_8t*)d)[ 5] = ((uint_8t*)s)[ 5];
  ((uint_8t*)d)[ 6] = ((uint_8t*)s)[ 6];
  ((uint_8t*)d)[ 7] = ((uint_8t*)s)[ 7];
  ((uint_8t*)d)[ 8] = ((uint_8t*)s)[ 8];
  ((uint_8t*)d)[ 9] = ((uint_8t*)s)[ 9];
  ((uint_8t*)d)[10] = ((uint_8t*)s)[10];
  ((uint_8t*)d)[11] = ((uint_8t*)s)[11];
  ((uint_8t*)d)[12] = ((uint_8t*)s)[12];
  ((uint_8t*)d)[13] = ((uint_8t*)s)[13];
  ((uint_8t*)d)[14] = ((uint_8t*)s)[14];
  ((uint_8t*)d)[15] = ((uint_8t*)s)[15];
}

static void copy_block_nn( uint_8t * d, const uint_8t *s, uint_8t nn )
{
  while( nn-- )
    *d++ = *s++;
}

// static void xor_block( void *d, const void *s )
// {
//   ((uint_8t*)d)[ 0] ^= ((uint_8t*)s)[ 0];
//   ((uint_8t*)d)[ 1] ^= ((uint_8t*)s)[ 1];
//   ((uint_8t*)d)[ 2] ^= ((uint_8t*)s)[ 2];
//   ((uint_8t*)d)[ 3] ^= ((uint_8t*)s)[ 3];
//   ((uint_8t*)d)[ 4] ^= ((uint_8t*)s)[ 4];
//   ((uint_8t*)d)[ 5] ^= ((uint_8t*)s)[ 5];
//   ((uint_8t*)d)[ 6] ^= ((uint_8t*)s)[ 6];
//   ((uint_8t*)d)[ 7] ^= ((uint_8t*)s)[ 7];
//   ((uint_8t*)d)[ 8] ^= ((uint_8t*)s)[ 8];
//   ((uint_8t*)d)[ 9] ^= ((uint_8t*)s)[ 9];
//   ((uint_8t*)d)[10] ^= ((uint_8t*)s)[10];
//   ((uint_8t*)d)[11] ^= ((uint_8t*)s)[11];
//   ((uint_8t*)d)[12] ^= ((uint_8t*)s)[12];
//   ((uint_8t*)d)[13] ^= ((uint_8t*)s)[13];
//   ((uint_8t*)d)[14] ^= ((uint_8t*)s)[14];
//   ((uint_8t*)d)[15] ^= ((uint_8t*)s)[15];
// }

static void copy_and_key( void *d, const void *s, const void *k )
{
  ((uint_8t*)d)[ 0] = ((uint_8t*)s)[ 0] ^ ((uint_8t*)k)[ 0];
  ((uint_8t*)d)[ 1] = ((uint_8t*)s)[ 1] ^ ((uint_8t*)k)[ 1];
  ((uint_8t*)d)[ 2] = ((uint_8t*)s)[ 2] ^ ((uint_8t*)k)[ 2];
  ((uint_8t*)d)[ 3] = ((uint_8t*)s)[ 3] ^ ((uint_8t*)k)[ 3];
  ((uint_8t*)d)[ 4] = ((uint_8t*)s)[ 4] ^ ((uint_8t*)k)[ 4];
  ((uint_8t*)d)[ 5] = ((uint_8t*)s)[ 5] ^ ((uint_8t*)k)[ 5];
  ((uint_8t*)d)[ 6] = ((uint_8t*)s)[ 6] ^ ((uint_8t*)k)[ 6];
  ((uint_8t*)d)[ 7] = ((uint_8t*)s)[ 7] ^ ((uint_8t*)k)[ 7];
  ((uint_8t*)d)[ 8] = ((uint_8t*)s)[ 8] ^ ((uint_8t*)k)[ 8];
  ((uint_8t*)d)[ 9] = ((uint_8t*)s)[ 9] ^ ((uint_8t*)k)[ 9];
  ((uint_8t*)d)[10] = ((uint_8t*)s)[10] ^ ((uint_8t*)k)[10];
  ((uint_8t*)d)[11] = ((uint_8t*)s)[11] ^ ((uint_8t*)k)[11];
  ((uint_8t*)d)[12] = ((uint_8t*)s)[12] ^ ((uint_8t*)k)[12];
  ((uint_8t*)d)[13] = ((uint_8t*)s)[13] ^ ((uint_8t*)k)[13];
  ((uint_8t*)d)[14] = ((uint_8t*)s)[14] ^ ((uint_8t*)k)[14];
  ((uint_8t*)d)[15] = ((uint_8t*)s)[15] ^ ((uint_8t*)k)[15];
}

// static void add_round_key( uint_8t d[N_BLOCK], const uint_8t k[N_BLOCK] )
// {
//   xor_block(d, k);
// }

static void shift_sub_rows( uint_8t st[N_BLOCK] )
{   uint_8t tt;

  st[ 0] = s_box(st[ 0]); st[ 4] = s_box(st[ 4]);
  st[ 8] = s_box(st[ 8]); st[12] = s_box(st[12]);

  tt = st[1]; st[ 1] = s_box(st[ 5]); st[ 5] = s_box(st[ 9]);
  st[ 9] = s_box(st[13]); st[13] = s_box( tt );

  tt = st[2]; st[ 2] = s_box(st[10]); st[10] = s_box( tt );
  tt = st[6]; st[ 6] = s_box(st[14]); st[14] = s_box( tt );

  tt = st[15]; st[15] = s_box(st[11]); st[11] = s_box(st[ 7]);
  st[ 7] = s_box(st[ 3]); st[ 3] = s_box( tt );
}

static void inv_shift_sub_rows( uint_8t st[N_BLOCK] )
{   uint_8t tt;

  st[ 0] = is_box(st[ 0]); st[ 4] = is_box(st[ 4]);
  st[ 8] = is_box(st[ 8]); st[12] = is_box(st[12]);

  tt = st[13]; st[13] = is_box(st[9]); st[ 9] = is_box(st[5]);
  st[ 5] = is_box(st[1]); st[ 1] = is_box( tt );

  tt = st[2]; st[ 2] = is_box(st[10]); st[10] = is_box( tt );
  tt = st[6]; st[ 6] = is_box(st[14]); st[14] = is_box( tt );

  tt = st[3]; st[ 3] = is_box(st[ 7]); st[ 7] = is_box(st[11]);
  st[11] = is_box(st[15]); st[15] = is_box( tt );
}

  static void mix_sub_columns( uint_8t dt[N_BLOCK], uint_8t st[N_BLOCK] )
  {
  dt[ 0] = gfm2_sb(st[0]) ^ gfm3_sb(st[5]) ^ s_box(st[10]) ^ s_box(st[15]);
  dt[ 1] = s_box(st[0]) ^ gfm2_sb(st[5]) ^ gfm3_sb(st[10]) ^ s_box(st[15]);
  dt[ 2] = s_box(st[0]) ^ s_box(st[5]) ^ gfm2_sb(st[10]) ^ gfm3_sb(st[15]);
  dt[ 3] = gfm3_sb(st[0]) ^ s_box(st[5]) ^ s_box(st[10]) ^ gfm2_sb(st[15]);

  dt[ 4] = gfm2_sb(st[4]) ^ gfm3_sb(st[9]) ^ s_box(st[14]) ^ s_box(st[3]);
  dt[ 5] = s_box(st[4]) ^ gfm2_sb(st[9]) ^ gfm3_sb(st[14]) ^ s_box(st[3]);
  dt[ 6] = s_box(st[4]) ^ s_box(st[9]) ^ gfm2_sb(st[14]) ^ gfm3_sb(st[3]);
  dt[ 7] = gfm3_sb(st[4]) ^ s_box(st[9]) ^ s_box(st[14]) ^ gfm2_sb(st[3]);

  dt[ 8] = gfm2_sb(st[8]) ^ gfm3_sb(st[13]) ^ s_box(st[2]) ^ s_box(st[7]);
  dt[ 9] = s_box(st[8]) ^ gfm2_sb(st[13]) ^ gfm3_sb(st[2]) ^ s_box(st[7]);
  dt[10] = s_box(st[8]) ^ s_box(st[13]) ^ gfm2_sb(st[2]) ^ gfm3_sb(st[7]);
  dt[11] = gfm3_sb(st[8]) ^ s_box(st[13]) ^ s_box(st[2]) ^ gfm2_sb(st[7]);

  dt[12] = gfm2_sb(st[12]) ^ gfm3_sb(st[1]) ^ s_box(st[6]) ^ s_box(st[11]);
  dt[13] = s_box(st[12]) ^ gfm2_sb(st[1]) ^ gfm3_sb(st[6]) ^ s_box(st[11]);
  dt[14] = s_box(st[12]) ^ s_box(st[1]) ^ gfm2_sb(st[6]) ^ gfm3_sb(st[11]);
  dt[15] = gfm3_sb(st[12]) ^ s_box(st[1]) ^ s_box(st[6]) ^ gfm2_sb(st[11]);
  }

  static void inv_mix_sub_columns( uint_8t dt[N_BLOCK], uint_8t st[N_BLOCK] )
  {
  dt[ 0] = is_box(gfm_e(st[ 0]) ^ gfm_b(st[ 1]) ^ gfm_d(st[ 2]) ^ gfm_9(st[ 3]));
  dt[ 5] = is_box(gfm_9(st[ 0]) ^ gfm_e(st[ 1]) ^ gfm_b(st[ 2]) ^ gfm_d(st[ 3]));
  dt[10] = is_box(gfm_d(st[ 0]) ^ gfm_9(st[ 1]) ^ gfm_e(st[ 2]) ^ gfm_b(st[ 3]));
  dt[15] = is_box(gfm_b(st[ 0]) ^ gfm_d(st[ 1]) ^ gfm_9(st[ 2]) ^ gfm_e(st[ 3]));

  dt[ 4] = is_box(gfm_e(st[ 4]) ^ gfm_b(st[ 5]) ^ gfm_d(st[ 6]) ^ gfm_9(st[ 7]));
  dt[ 9] = is_box(gfm_9(st[ 4]) ^ gfm_e(st[ 5]) ^ gfm_b(st[ 6]) ^ gfm_d(st[ 7]));
  dt[14] = is_box(gfm_d(st[ 4]) ^ gfm_9(st[ 5]) ^ gfm_e(st[ 6]) ^ gfm_b(st[ 7]));
  dt[ 3] = is_box(gfm_b(st[ 4]) ^ gfm_d(st[ 5]) ^ gfm_9(st[ 6]) ^ gfm_e(st[ 7]));

  dt[ 8] = is_box(gfm_e(st[ 8]) ^ gfm_b(st[ 9]) ^ gfm_d(st[10]) ^ gfm_9(st[11]));
  dt[13] = is_box(gfm_9(st[ 8]) ^ gfm_e(st[ 9]) ^ gfm_b(st[10]) ^ gfm_d(st[11]));
  dt[ 2] = is_box(gfm_d(st[ 8]) ^ gfm_9(st[ 9]) ^ gfm_e(st[10]) ^ gfm_b(st[11]));
  dt[ 7] = is_box(gfm_b(st[ 8]) ^ gfm_d(st[ 9]) ^ gfm_9(st[10]) ^ gfm_e(st[11]));

  dt[12] = is_box(gfm_e(st[12]) ^ gfm_b(st[13]) ^ gfm_d(st[14]) ^ gfm_9(st[15]));
  dt[ 1] = is_box(gfm_9(st[12]) ^ gfm_e(st[13]) ^ gfm_b(st[14]) ^ gfm_d(st[15]));
  dt[ 6] = is_box(gfm_d(st[12]) ^ gfm_9(st[13]) ^ gfm_e(st[14]) ^ gfm_b(st[15]));
  dt[11] = is_box(gfm_b(st[12]) ^ gfm_d(st[13]) ^ gfm_9(st[14]) ^ gfm_e(st[15]));
  }

/*  加密key初始设置 */

int Util_Aes_Set_Key( const unsigned char key[16], Aes_Padding_Mode padding_mode )
{
  Aes_Context_Struct  *ctx=&aes_context_struct; 
  ctx->padding_mode=padding_mode;
  uint_8t cc, rc, hi;
  uint_8t keylen=16;  
  block_copy_nn(ctx->ksch, key, keylen);
  hi = (keylen + 28) << 2;
  ctx->rnd = (hi >> 4) - 1;
  for( cc = keylen, rc = 1; cc < hi; cc += 4 )
  {   uint_8t tt, t0, t1, t2, t3;

    t0 = ctx->ksch[cc - 4];
    t1 = ctx->ksch[cc - 3];
    t2 = ctx->ksch[cc - 2];
    t3 = ctx->ksch[cc - 1];
    if( cc % keylen == 0 )
    {
      tt = t0;
      t0 = s_box(t1) ^ rc;
      t1 = s_box(t2);
      t2 = s_box(t3);
      t3 = s_box(tt);
      rc = f2(rc);
    }
    else if( keylen > 24 && cc % keylen == 16 )
    {
      t0 = s_box(t0);
      t1 = s_box(t1);
      t2 = s_box(t2);
      t3 = s_box(t3);
    }
    tt = cc - keylen;
    ctx->ksch[cc + 0] = ctx->ksch[tt + 0] ^ t0;
    ctx->ksch[cc + 1] = ctx->ksch[tt + 1] ^ t1;
    ctx->ksch[cc + 2] = ctx->ksch[tt + 2] ^ t2;
    ctx->ksch[cc + 3] = ctx->ksch[tt + 3] ^ t3;
  }
  return 0;
}


/*  只加密16字节的block */

int aes_encrypt( const unsigned char in[N_BLOCK], unsigned char  out[N_BLOCK], const Aes_Context_Struct ctx[1] )
{
  if( ctx->rnd )
  {
    uint_8t s1[N_BLOCK], r;
    copy_and_key( s1, in, ctx->ksch );

    for( r = 1 ; r < ctx->rnd ; ++r )
    {   uint_8t s2[N_BLOCK];
      mix_sub_columns( s2, s1 );
      copy_and_key( s1, s2, ctx->ksch + r * N_BLOCK);
    }
    shift_sub_rows( s1 );
    copy_and_key( out, s1, ctx->ksch + r * N_BLOCK );
  }
  else
    return -1;
  return 0;
}

/* 计算加密后内容的长度 */
int compute_out_length(int in_len) {
  int padding = in_len & (N_BLOCK - 1);
  return padding ? N_BLOCK + in_len - padding : N_BLOCK + in_len;
}

/* pkcs#7填充 */
void pkcs7_padding(const unsigned char *in, unsigned char *out, int in_len, int out_len) {
  uint_8t padding = out_len - in_len;
  uint_8t i = 0;
  block_copy_nn(out, in, in_len);
  out += in_len;
  while(i < padding) {
    *(out + i) = padding;
    i++;
  }
}

/* ecb模式加密 */
int Util_Aes_Ecb_Encrypt(const unsigned char *in, int src_len,unsigned char *out) {
  Aes_Context_Struct *ctx=&aes_context_struct;
  int i = 0;
  int out_len;
  if(src_len%N_BLOCK==0){
      out_len=src_len+N_BLOCK;
  }else{
      out_len=src_len+N_BLOCK-src_len%N_BLOCK;
  }
  uint_8t tmp[N_BLOCK];
  while(i < out_len) {
    block_copy(tmp, in);
    if(ctx->padding_mode==PKCS7_PADDING){
        if(src_len-i<N_BLOCK){
            uint_8t padding_data;
            padding_data=N_BLOCK+i-src_len;
            for(int j=src_len-i;j<N_BLOCK;j++){
                tmp[j]=padding_data;
            }
        }
    }   
    if (aes_encrypt(tmp, tmp, ctx) != 0)
    {
      return -1;
    }
    block_copy(out, tmp);
    in += N_BLOCK;
    out += N_BLOCK;
    i += N_BLOCK;
  }
  return out_len;
}


/*  单个16字节块的解密 */

int aes_decrypt( const unsigned char in[N_BLOCK], unsigned char out[N_BLOCK], const Aes_Context_Struct ctx[1] )
{
  if( ctx->rnd )
  {
    uint_8t s1[N_BLOCK], r;
    copy_and_key( s1, in, ctx->ksch + ctx->rnd * N_BLOCK );
    inv_shift_sub_rows( s1 );

    for( r = ctx->rnd ; --r ; )
    {   uint_8t s2[N_BLOCK];
      copy_and_key( s2, s1, ctx->ksch + r * N_BLOCK );
      inv_mix_sub_columns( s1, s2 );
    }
    copy_and_key( out, s1, ctx->ksch );
  }
  else
    return -1;
  return 0;
}

int check_padding_data(const unsigned char *in,char padding_data,int len){
    for(int i=0;i<len;i++){
        if(in[i]!=padding_data){
            return -1;
        }
    }
    return padding_data;
}

/* ecb模式解密,解密后输出含有padding  */
int Util_Aes_Ecb_Decrypt(const unsigned char *in, unsigned char *out, int in_len) {
  Aes_Context_Struct *ctx=&aes_context_struct;  
  int i = 0;
  while(i < in_len) {
    if (aes_decrypt(in, out, ctx) != 0) {
      return -1;
    }
    in += N_BLOCK;
    out += N_BLOCK;
    i += N_BLOCK;
  }
  out=out-in_len;
  int out_len;
  out_len=in_len;
  for(int j=out_len-1;j>=0;j--){
      char padding_data=N_BLOCK-j%N_BLOCK;
      if(check_padding_data(&out[j],padding_data,padding_data)<0){
          continue;
      }else{
          out_len=out_len-padding_data;
          return out_len;
      }
  }
  return out_len;
}


