【干货】如何手搓AES加密算法

· · 科技·工程

【干货】如何手搓 AES 加密算法

:::info[AI 声明]{open} 本文章及代码由 DeepSeek V4 优化。 :::

0 前置芝士

大多 AES 算法是在一个特定的有限域(一般为 \mathrm{GF}(2^8) )上进行,加密过程是在一个 4\times 4 的矩阵上进行的,这个矩阵叫做“体(state)”,每一个体都是单独的一个加密块。

AES 算法分为 AES-128AES-192AES-256,分别对应 128 位、192 位、256 位密钥,其中,不同密钥长度的加密轮次不一样,分别为 10 轮、12 轮、14 轮。

2 AES 加密的步骤

AES 加密的步骤为:

  1. 轮密钥加
  2. 字节替换
  3. 行变换
  4. 列混淆
  5. 轮密钥加
  6. 返回第 2 步,循环次数按算法定。

    2.1 Add Round Key(轮密钥加)

    在这一步骤里,上一轮加密后的密文中的每一个字节将会与这一轮的轮密钥(round key)的对应字节进行异或操作。其中,轮秘钥由主密钥经过 Rijndael 密钥生成方案生成。

轮秘钥是一个与密文大小相同的 4\times 4 矩阵。

2.2 Sub Bytes(字节替换)

在这一步骤里,经过“轮密钥加”后的密文将会经过一个固定的查找表(下称 S-box),把密文中的每一个字节进行对应替换。

S-box 是一个非线性的替换函数,已知具有良好的非线性特性,使用 S-box 进行替换后相当于一个错排的效果。

其设计逻辑为: 它是一个8位非线性替换表,其设计基于有限域 GF(2^8) 上的数学运算,主要分为两步:

  1. 求乘法逆元

对每个输入字节 x(值域 0 \sim 255),在有限域 GF(2^8) 中计算其乘法逆元 x^{-1},满足 x \cdot x^{-1} = 1 \pmod{m(x)}。此处采用的不可约多项式为

m(x) = x^8 + x^4 + x^3 + x + 1

0x11B
特殊地,x=0 的逆元定义为 0 本身。求逆运算提供了极强的非线性,是抵抗差分和线性密码分析的核心。

  1. 仿射变换

对逆元结果 b = (b_7, b_6, \dots, b_0)(按位向量表示)施加一个可逆的仿射变换,公式为

y = M \cdot b \oplus c

其中 M8\times8 的固定二进制矩阵,c 是常数向量 c = (1,1,0,0,0,1,1,0)^\mathsf{T}(即十六进制 0x63)。矩阵 M 的具体形式为:

M = \begin{bmatrix} 1 & 0 & 0 & 0 & 1 & 1 & 1 & 1 \\ 1 & 1 & 0 & 0 & 0 & 1 & 1 & 1 \\ 1 & 1 & 1 & 0 & 0 & 0 & 1 & 1 \\ 1 & 1 & 1 & 1 & 0 & 0 & 0 & 1 \\ 1 & 1 & 1 & 1 & 1 & 0 & 0 & 0 \\ 0 & 1 & 1 & 1 & 1 & 1 & 0 & 0 \\ 0 & 0 & 1 & 1 & 1 & 1 & 1 & 0 \\ 0 & 0 & 0 & 1 & 1 & 1 & 1 & 1 \end{bmatrix}

该变换确保 S(0) \ne 0S(1) \ne 1,消除了不动点,并进一步增加了代数复杂性,使 S-box 不存在简单的代数关系。

2.3 Shift Rows(行变换)

在这一步骤里,经过“字节替换”后的密文将会以固定偏移量进行循环移位,其中第一行左移 0 位,第二行左移 1 位,第三行左移 2 位,第四行左移 3 位。

这一步是为了扩散字节位置,为了和下面的“Mix Columns”产生雪崩效应。

2.4 Mix Columns(列混淆)

在这一步骤中,将会进行如下运算(ans 数组为答案体,a 数组为一个密文体,b 数组位指定的列混淆矩阵,在 GF(2^8) 中进行运算)

ans_{i, j}=(a_{0, j}\times b_{i, 0})\oplus (a_{1, j}\times b_{i, 1})\oplus (a_{2, j}\times b_{i, 2})\oplus (a_{3, j}\times b_{i, 3})

最后答案为 ans 数组。

在最后一轮中,Mix Columns 步骤不执行!

3 代码实现

本代码以 AES-256 为例。

3.1 代码中的缩写含义

  1. vivector<int>
  2. vvivector<vector<int> >
  3. migmulti_in_GF256
  4. enMC4x4enMixColumns4x4
  5. deMC4x4deMixColumns4x4

    3.2 前置步骤的实现

    3.2.1 生成主密钥

    这一部分很简单,只需要生成随机数,在 random 库中就可以直接使用 mt_19937_64 随机数引擎,再加上 uniform_int_distribution就可以生成随机数(我这里用的种子(seed)是现在时间到 1970 年 1 月 1 日的纳秒数):

    void gen_key() {
    chrono::system_clock::time_point now = chrono::system_clock::now();
    unsigned long long seed=chrono::duration_cast<chrono::nanoseconds>(now.time_since_epoch()).count();  // 现在时间到 1970 年 1 月 1 日的纳秒数
    mt19937_64 mt(seed);
    uniform_int_distribution<> dist(0, 255);
    for(int i = 0;i<32;i++) { // 这里的32是AES-256的字数,如果是AES-128/192,那么这里为16/24。
        key[i]=static_cast<unsigned char>(dist(mt));
    }
    key_expansion();// 生成轮秘钥
    }

    3.2.2 生成轮秘钥

    这一部分分为两个步骤:密钥扩展轮秘钥选择,定义一个“字”的大小为 4 字节。

首先是密钥扩展

定义函数 rotl(x, y)4 字节变量 x 循环左移 y 位,例如:rotl(b, 1) 就相当于将 4 字节字 [b0, b1, b2, b3] 左移一个字节,变成 [b1, b2, b3, b0]

函数 SubWord(x) 为把变量 x 经过 S-box 中字节替换。

数组 Rcon 代表轮系数数组:

如果为 AES-128 那么:

1. 把主密钥的 4 个字直接放进 W_0\!\sim\!W_3
2. 剩下的 40 个字,我们这样生成:

  • tmp = W[i-1]
  • 如果 i % 4 == 0
    tmp = SubWord(RotWord(tmp)) ^ Rcon[i/4]
  • 否则 tmp 不变
  • 最后 W[i] = W[i-4] ^ tmp

如果为 AES-192,则步骤类似,只是初始密钥为 6 个字(放入 W_0\!\sim\!W_5),轮密钥总字数变为 52,每 6 个字进行一次操作:

如果为 AES-256,则步骤类似,但初始密钥为 8 个字(放入 W_0\!\sim\!W_7),轮密钥总字数变为 60,每 8 个字进行一次主操作,并且中间要额外多做一次 SubWord

然后是轮秘钥选择

这一步骤就很简单了,我们只需要在生成的 W 数组中顺序截取相应的长度就是每一轮的密钥。

即设 Nb 为密钥的字数,那么第 0 轮密钥为 W_0\!\sim\!W_{Nb-1},第 1 轮密钥为 W_{Nb}\!\sim\!W_{2*Nb-1},以此类推,第 r 轮密钥为 W_{r\times Nb}\!\sim\!W_{(r+1)\times Nb-1}

轮秘钥生成代码:

void rot_word(unsigned char w[4]) {
    unsigned char tmp = w[0];
    w[0] = w[1]; w[1] = w[2]; w[2] = w[3]; w[3] = tmp;
}
void sub_word(unsigned char w[4]) {
    for(int i = 0; i < 4; i++) {
        int l = w[i] >> 4, r = w[i] & 15;
        w[i] = S_Box[l][r];
    }
}
unsigned char rcon_val(int j) {
    unsigned char val = 1;
    for(int k = 1; k < j; k++) {
        unsigned char carry = val & 0x80;
        val <<= 1;
        if(carry) val ^= 0x1b;
    }
    return val;
}
void key_expansion() {
    unsigned char w[60][4];
    for(int i = 0; i < 8; i++) {
        for(int j = 0; j < 4; j++) {
            w[i][j] = key[i * 4 + j];
        }
    }
    for(int i = 8; i < 60; i++) {
        unsigned char tmp[4];
        for(int j = 0; j < 4; j++) tmp[j] = w[i - 1][j];
        if(i % 8 == 0) {
            rot_word(tmp);
            sub_word(tmp);
            unsigned char rc = rcon_val(i / 8);
            tmp[0] ^= rc;
        } else if(i % 8 == 4) {
            sub_word(tmp);
        }
        for(int j = 0; j < 4; j++) {
            w[i][j] = w[i - 8][j] ^ tmp[j];
        }
    }
    for(int r = 0; r < 15; r++) {
        for(int j = 0; j < 4; j++) {
            for(int k = 0; k < 4; k++) {
                roundKeys[r][j * 4 + k] = w[r * 4 + j][k];
            }
        }
    }
}

3.2.3 在有限域内乘法

在有限域 GF(2^8) 内做运算与普通运算不同,它的运算规则为异或代替加法,GF(2^8) 乘法代替普通乘法。

其不可约多项式为

x^8+x^4+x^3+x+1

0x11B,也记为 m(x)

他的乘法由两部分构成

  1. 函数 xtime,这个函数负责实现字节 x\times 2 模不可约多项式。xtime 函数负责计算 GF(2^8) 域中乘以 2(即乘以多项式 x 的运算。 将 x 左移 1 位,若原最高位为 1(即左移后超出 8 位),则再异或常数 0x1B(即 m(x)-x^8)。
    1. 无进位乘法使用 xtime 累加的方法,类似二进制乘法,但无进位,加法使用异或。 所以代码实现:
      unsigned char multi_in_GF256(unsigned char a, unsigned char b) {
      unsigned char res = 0;
      while(b) {
      if(b & 1) res ^= a;
      unsigned char carry = a & 0x80;
      a <<= 1;
      if(carry) a ^= 0x1b; // xtime函数集成到此函数里
      b >>= 1;
      }
      return res;
      }

      3.2.4 其他数据格式转换函数

      这些函数都在加解密中使用,不需要使用算法,只需要模拟。

  2. 16 个字符的字符串转换为体:
    vvi char_to_piece(const char* str) {
    vvi ans(4, vector<int>(4, 0));
    for(int i = 0; i < 4; i++) {
        for(int j = 0; j < 4; j++) {
            ans[i][j] = static_cast<unsigned char>(str[i * 4 + j]);
        }
    }
    return ans;
    }
  3. 把体转换为字符串
    string piece_to_str(vvi piece) {
    string ans;
    for(int i = 0; i < 4; i++) {
        for(int j = 0; j < 4; j++) {
            ans.push_back(static_cast<unsigned char>(piece[i][j]));
        }
    }
    return ans;
    }
  4. 把体转换为一维 vector
    vi piece_to_vector(vvi piece) {
    vector<int> ans;
    for(int i = 0; i < 4; i++) {
        for(int j = 0; j < 4; j++) {
            ans.push_back(piece[i][j]);
        }
    }
    return ans;
    }

    3.2.5 Base64 编码

    这里只需要了解其原理,就很好写出来,原理不再赘述。

    static string base64_encode(const vi& data) {
    const char table[] = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
    string ans;
    int i = 0, n = data.size();
    while(i < n) {
        unsigned char a = data[i++];
        bool has_b = (i < n);
        unsigned char b = has_b ? data[i++] : 0;
        bool has_c = (i < n);
        unsigned char c = has_c ? data[i++] : 0;
        ans.push_back(table[a >> 2]);
        ans.push_back(table[((a & 3) << 4) | (b >> 4)]);
        ans.push_back(has_b ? table[((b & 15) << 2) | (c >> 6)] : '=');
        ans.push_back(has_c ? table[c & 63] : '=');
    }
    return ans;
    }
    static vi base64_decode(const string& str) {
    auto idx = [](unsigned char c) -> int {
        if(c >= 'A' && c <= 'Z') return c - 'A';
        if(c >= 'a' && c <= 'z') return c - 'a' + 26;
        if(c >= '0' && c <= '9') return c - '0' + 52;
        if(c == '+') return 62;
        if(c == '/') return 63;
        return -1;
    };
    vi ans;
    int i = 0, n = str.length();
    while(i < n && str[i] != '=') {
        int a = idx(str[i++]); if(a == -1) continue;
        int b = idx(str[i++]); if(b == -1) continue;
        ans.push_back((a << 2) | (b >> 4));
        if(i < n && str[i] != '=') {
            int c = idx(str[i++]); if(c == -1) continue;
            ans.push_back(((b & 15) << 4) | (c >> 2));
            if(i < n && str[i] != '=') {
                int d = idx(str[i++]); if(d == -1) continue;
                ans.push_back(((c & 3) << 6) | d);
            }
        }
    }
    return ans;
    }

    3.3 加密步骤的实现

    3.3.1 Add Round Key 的实现

    这里我们只需要把轮密钥与明文体逐字节做异或,所以代码就很简单:

    vvi add_round_key(const vvi state, int round) {
    vvi ans = state;
    for(int i = 0; i < 4; i++) {
        for(int j = 0; j < 4; j++) {
            ans[i][j] ^= static_cast<unsigned char>(roundKeys[round][j * 4 + i]);
        }
    }
    return ans;
    }

    3.3.2 Sub Bytes 的实现

    上图为 Sub Bytes 中的 S-box,只需要甩给AI有耐心地把这些字母抄下来,然后再逐个查表即可(代码中 S-box 略):

    vvi enSubBytes(const vvi str) {
    vvi ans(4, vi(4, 0));
    for(int i = 0;i<4;i++) {
        for(int j = 0;j<4;j++) {
            int l=static_cast<unsigned char>(str[i][j])>>4, r=static_cast<unsigned char>(str[i][j])&15;
            ans[i][j]=(int)S_Box[l][r];
        }
    }
    return ans;
    }

    3.3.3 Shift Rows 的实现

    这里我们只需要打一个模拟,很简单,所以直接放出来,注意,移动顺序不要搞混,不然就会有覆盖的情况:

    vvi enShiftRows(const vvi str) {
    vvi ans = str;
    // 第1行: 左移1格
    int tmp1 = ans[1][0];
    ans[1][0] = ans[1][1];
    ans[1][1] = ans[1][2];
    ans[1][2] = ans[1][3];
    ans[1][3] = tmp1;
    // 第2行: 左移2格
    tmp1 = ans[2][0];
    int tmp2 = ans[2][1];
    ans[2][0] = ans[2][2];
    ans[2][1] = ans[2][3];
    ans[2][2] = tmp1;
    ans[2][3] = tmp2;
    // 第3行: 左移3格
    tmp1 = ans[3][0], tmp2 = ans[3][1];
    int tmp3 = ans[3][2];
    ans[3][0] = ans[3][3];
    ans[3][1] = tmp1;
    ans[3][2] = tmp2;
    ans[3][3] = tmp3;
    return ans;
    }

    3.3.4 Mix Columns 的实现

    这里只需要套公式,再加上 GF(2^8) 乘法即可,公式为:

    ans_{i, j}=(a_{0, j}\times b_{i, 0})\oplus (a_{1, j}\times b_{i, 1})\oplus (a_{2, j}\times b_{i, 2})\oplus (a_{3, j}\times b_{i, 3}) b=\begin{pmatrix} 2 & 3 & 1 & 1 \\ 1 & 2 & 3 & 1 \\ 1 & 1 & 2 & 3 \\ 3 & 1 & 1 & 2 \end{pmatrix}
    int enMixColumns4x4[4][4] = {
    {0x02, 0x03, 0x01, 0x01},
    {0x01, 0x02, 0x03, 0x01},
    {0x01, 0x01, 0x02, 0x03},
    {0x03, 0x01, 0x01, 0x02}
    }; // 这个是如上的矩阵b
    vvi enMixColumns(const vvi str) {
    vvi ans(4, vi(4, 0));
    for(int c = 0; c < 4; c++) {
        ans[0][c] = mig(str[0][c], enMC4x4[0][0]) ^ mig(str[1][c], enMC4x4[0][1])
                  ^ mig(str[2][c], enMC4x4[0][2]) ^ mig(str[3][c], enMC4x4[0][3]);
        ans[1][c] = mig(str[0][c], enMC4x4[1][0]) ^ mig(str[1][c], enMC4x4[1][1])
                  ^ mig(str[2][c], enMC4x4[1][2]) ^ mig(str[3][c], enMC4x4[1][3]);
        ans[2][c] = mig(str[0][c], enMC4x4[2][0]) ^ mig(str[1][c], enMC4x4[2][1])
                  ^ mig(str[2][c], enMC4x4[2][2]) ^ mig(str[3][c], enMC4x4[2][3]);
        ans[3][c] = mig(str[0][c], enMC4x4[3][0]) ^ mig(str[1][c], enMC4x4[3][1])
                  ^ mig(str[2][c], enMC4x4[3][2]) ^ mig(str[3][c], enMC4x4[3][3]);
    }
    return ans;
    }

    3.3.5 总体加密实现

    只需要按照步骤一个一个模拟即可,注意数据的转换以及密钥的输出:

    pair<string, string> encode(const char* msg, const string& key_b64 = "") {
    if(key_b64 != "") {
        vi k = base64_decode(key_b64);
        unsigned char k_arr[32];
        for(int i = 0; i < 32; i++) k_arr[i] = static_cast<unsigned char>(k[i]);
        set_key(k_arr);
    } else {
        gen_key();
    }
    string str(msg);
    vi raw;
    int len = 16 - str.length() % 16;
    for(int i = 1; i <= len; i++) {
        str.push_back(static_cast<char>(len));
    }
    for(int i = 0; i < (int)str.length(); i += 16) {
        vvi state = char_to_piece(str.substr(i, 16).c_str());
        state = add_round_key(state, 0);
        for(int r = 1; r <= 13; r++) {
            state = enSubBytes(state);
            state = enShiftRows(state);
            state = enMixColumns(state);
            state = add_round_key(state, r);
        }
        state = enSubBytes(state);
        state = enShiftRows(state);
        state = add_round_key(state, 14);
        vi v = piece_to_vector(state);
        for(int j = 0; j < 16; j++) raw.push_back(v[j]);
    }
    return {base64_encode(raw), get_key_b64()};
    }

    3.4 解密步骤的实现

    3.4.1 Add Round Key 解密

    我们都知道异或运算存在自相反性,所以这个和加密的代码一样。

    3.4.2 Sub Bytes 解密

    这里想不到什么好办法,除非打表,所以暴力枚举体的每一个字节,与 S_box 作对比,一模一样就存进答案数组里。

    vvi deSubBytes(const vvi str) {
    vvi ans(4, vi(4, 0));
    for(int i = 0; i < 4; i++) {
        for(int j = 0; j < 4; j++) {
            for(int k = 0; k < 16; k++) {
                for(int l = 0; l < 16; l++) {
                    if(str[i][j] == S_Box[k][l]) {
                        ans[i][j] = (k << 4) + l;
                    }
                }
            }
        }
    }
    return ans;
    }

    3.4.3 Shift Rows 解密

    这里只需要把原来向左移变为向右移即可:

    vvi deShiftRows(const vvi str) {
    vvi ans=str;
    int tmp1=ans[1][3];
    ans[1][3]=ans[1][2];
    ans[1][2]=ans[1][1];
    ans[1][1]=ans[1][0];
    ans[1][0]=tmp1;
    tmp1=ans[2][3];
    int tmp2=ans[2][2];
    ans[2][3]=ans[2][1];
    ans[2][2]=ans[2][0];
    ans[2][1]=tmp1;
    ans[2][0]=tmp2;
    tmp1=ans[3][3], tmp2=ans[3][2];
    int tmp3=ans[3][1];
    ans[3][3]=ans[3][0];
    ans[3][2]=tmp1;
    ans[3][1]=tmp2;
    ans[3][0]=tmp3;
    return ans;
    }

    3.4.4 Mix Columns 解密

    这里的公式和加密的一样,只不过 b 数组改变,即:

    ans_{i, j}=(a_{0, j}\times b_{i, 0})\oplus (a_{1, j}\times b_{i, 1})\oplus (a_{2, j}\times b_{i, 2})\oplus (a_{3, j}\times b_{i, 3}) b=\begin{pmatrix} 14 & 11 & 13 & 9 \\ 9 & 14 & 11 & 13 \\ 13 & 9 & 14 & 11 \\ 11 & 13 & 9 & 14 \end{pmatrix}
    int deMixColumns4x4[4][4] = {
    {0x0e, 0x0b, 0x0d, 0x09},
    {0x09, 0x0e, 0x0b, 0x0d},
    {0x0d, 0x09, 0x0e, 0x0b},
    {0x0b, 0x0d, 0x09, 0x0e}
    };
    vvi deMixColumns(const vvi str) {
    vvi ans(4, vi(4, 0));
    for(int c = 0; c < 4; c++) {
        ans[0][c] = mig(str[0][c], deMC4x4[0][0]) ^ mig(str[1][c], deMC4x4[0][1])
                  ^ mig(str[2][c], deMC4x4[0][2]) ^ mig(str[3][c], deMC4x4[0][3]);
        ans[1][c] = mig(str[0][c], deMC4x4[1][0]) ^ mig(str[1][c], deMC4x4[1][1])
                  ^ mig(str[2][c], deMC4x4[1][2]) ^ mig(str[3][c], deMC4x4[1][3]);
        ans[2][c] = mig(str[0][c], deMC4x4[2][0]) ^ mig(str[1][c], deMC4x4[2][1])
                  ^ mig(str[2][c], deMC4x4[2][2]) ^ mig(str[3][c], deMC4x4[2][3]);
        ans[3][c] = mig(str[0][c], deMC4x4[3][0]) ^ mig(str[1][c], deMC4x4[3][1])
                  ^ mig(str[2][c], deMC4x4[3][2]) ^ mig(str[3][c], deMC4x4[3][3]);
    }
    return ans;
    }

    3.4.5 总体解密

    解密的步骤就是加密的步骤完全反过来,注意数据类型的转换。

    string decode(const string& b64, const string& key_b64="") {
    if(key_b64 != "") {
        vi k = base64_decode(key_b64);
        unsigned char k_arr[32];
        for(int i = 0; i < 32; i++) k_arr[i] = static_cast<unsigned char>(k[i]);
        set_key(k_arr);
    }
    vi msg = base64_decode(b64);
    string ans;
    for(int i = 0; i < (int)msg.size(); i += 16) {
        vvi state(4, vi(4, 0));
        for(int j = 0; j < 4; j++) {
            for(int k = 0; k < 4; k++) {
                state[j][k] = msg[j * 4 + k + i];
            }
        }
        state = add_round_key(state, 14);
        state = deShiftRows(state);
        state = deSubBytes(state);
        for(int r = 13; r >= 1; r--) {
            state = add_round_key(state, r);
            state = deMixColumns(state);
            state = deShiftRows(state);
            state = deSubBytes(state);
        }
        state = add_round_key(state, 0);
        string str = piece_to_str(state);
        ans += str;
    }
    int pad_len = static_cast<unsigned char>(ans[ans.length() - 1]);
    return ans.substr(0, ans.length() - pad_len);
    }

    3.5 完整代码

    完整代码中需要添加其他的一些函数,如输入密钥等,这里不在赘述: 完整代码

    4 后记&彩蛋

    我编写这个代码纯闲的想来看一看加密的原理,虽然我以后也不想学密码学,我估计以后还会讲 SHA-256 的哈希运算。不说了,来看看彩蛋:

无奖竞猜加密密钥,并解密出来(提示:密钥是关于我、文章以及某编程网站连起来的密钥,密钥需经过 Base-64 编码再输入进程序里): IoJ/398Z3e4921cIpgBU8gH3TYT3TpoSRIq7KZboYdXZSCaQAQJ0e1bfMkRJcrPd