/**
 *@file    sr_aes.c
 *@brief   AES加解密接口
 *@details 提供AES加解密接口
 *@copyright Copyright (c) 2024 lierda. All rights reserved.
 *@author  Lierda-RDC
 *@date    2024-05-07
 *@example ZhaoyangSDK\example\component\crypto\crypto_example.c
 */

#include "sr_aes.h"
#include "sr_log.h"

static int _AesProcess(SR_AES_TYPE_E type, uint8_t mode, uint8_t *key, uint16_t keyBit, uint8_t iv[16], uint8_t *in, uint16_t inLen, uint8_t *out)
{
    int i, ret = 0;
    int olen = 0, len = 0;
    uint8_t *aesIV = iv;
    uint8_t ivLen = 16;

    mbedtls_cipher_context_t ctx; // 创建上下文
    const mbedtls_cipher_info_t *cipherInfo;

    /* 初始化上下文 */
    mbedtls_cipher_init(&ctx);
    cipherInfo = mbedtls_cipher_info_from_type(type);

    // 设置上下文
    ret = mbedtls_cipher_setup(&ctx, cipherInfo);
    if (ret != 0)
        goto _END;

    // 设置密钥,加密or解密
    ret = mbedtls_cipher_setkey(&ctx, key, keyBit, mode);
    if (ret != 0)
        goto _END;

    // AES ECB做些特殊处理
    if ((type >= MBEDTLS_CIPHER_AES_128_ECB) && (type <= MBEDTLS_CIPHER_AES_256_ECB))
    {
        uint8_t padding = 0, padBuf[16] = {0};
        uint16_t padInLen = 0;      // 填充后的长度
        uint16_t poll = inLen / 16; // 未填充前有多少个16字节分组

        if (mode == MBEDTLS_DECRYPT && (inLen % 16)) // AES ECB 解密，inLen必须为16整倍数
        {
            ret = -1;
            goto _END;
        }

        // 先处理前poll个分组
        for (i = 0; i < poll; i++)
        {
            ret = mbedtls_cipher_crypt(&ctx, 0, 0, in + olen, 16, out + olen, &len);
            if (ret != 0)
                goto _END;
            olen += 16;
        }

        if (mode == MBEDTLS_ENCRYPT) // 加密模式
        {
            // 处理剩余数据，不够16字节先进行PKCS7填充;若待加密数据正好为16字节整倍数，仍需填充16个字节
            // 计算剩余长度距离16整倍数的差值，此值即为填充值
            padding = 16 - (inLen % 16);
            // padInLen即为经过填充后待加密数据的长度，值必然为16整倍数，同时此长度也为加密后的密文长度
            padInLen = inLen + padding;

            if (padding != 16)
                memcpy(padBuf, in + olen, 16 - padding);
            memset(padBuf + 16 - padding, padding, padding);

            ret = mbedtls_cipher_crypt(&ctx, 0, 0, padBuf, 16, out + olen, &len);
            olen += 16;
        }
        else // 解密模式
        {
            // 解密后去除填充值
            // 读取最后一个输出值，此值即为填充值的长度
            padding = out[olen - 1];
            if (padding < 1 || padding > 16)
            {
                ret = -1;
                goto _END;
            }
            olen = inLen - padding;
        }
    }
    else
    {
        ret = mbedtls_cipher_crypt(&ctx, aesIV, ivLen, in, inLen, out, &olen);
    }

_END:
    mbedtls_cipher_free(&ctx);

    return (ret != 0) ? SR_FAIL : olen;
}

/**
 * @brief AES加密
 *
 * @param[in] type   AES算法类型，取值SR_AES_TYPE类型
 * @param[in] key    密钥
 * @param[in] keyBit 密钥长度(bit),取值必须为128, 192 or 256
 * @param[in] iv     初始化向量，CBC/CTR等算法需要。固定16字节
 * @param[in] in     输入数据
 * @param[in] inLen  输入数据的长度（字节）
 * @param[out] out   输出数据缓存区,一定要预留足够空间
 *
 * @return 成功(>0)：返回密文长度           失败：返回负值错误码
 */
int SR_AesEncrypt(SR_AES_TYPE_E type, uint8_t *key, uint16_t keyBit, uint8_t iv[16], uint8_t *in, uint16_t inLen, uint8_t *out)
{
    return _AesProcess(type, MBEDTLS_ENCRYPT, key, keyBit, iv, in, inLen, out);
}

/**
 * @brief AES解密
 *
 * @param[in] type   AES算法类型，取值SR_AES_TYPE类型
 * @param[in] key    密钥
 * @param[in] keyBit 密钥长度(bit),取值必须为128, 192 or 256
 * @param[in] iv     初始化向量，CBC/CTR等算法需要。固定16字节
 * @param[in] in     输入数据
 * @param[in] inLen  输入数据的长度（字节）
 * @param[out] out   输出数据缓存区，一定要预留足够空间
 *
 * @return 成功(>0)：返回明文长度           失败：返回负值错误码
 */
int SR_AesDecrypt(SR_AES_TYPE_E type, uint8_t *key, uint16_t keyBit, uint8_t iv[16], uint8_t *in, uint16_t inLen, uint8_t *out)
{
    return _AesProcess(type, MBEDTLS_DECRYPT, key, keyBit, iv, in, inLen, out);
}
