// SM4 分组密码算法
// 依据:GB/T 32907-2016《信息安全技术 SM4 分组密码算法》
// 纯 MoonBit 实现,零依赖。提供 128 位分组加解密、
// 密钥扩展、ECB/CBC 工作模式与 PKCS7 填充。
///|
/// SM4 相关错误类型。
pub(all) enum Sm4Error {
/// 密钥长度不是 16 字节,实际为 `Int` 字节
InvalidKeyLength(Int)
/// IV 长度不是 16 字节,实际为 `Int` 字节
InvalidIvLength(Int)
/// 密文长度不是 16 的整数倍
InvalidCiphertextLength(Int)
/// 密文 PKCS7 填充不合法
InvalidPadding
} derive(Eq, Debug)
///|
pub extend Sm4Error with Eq::{not_equal, equal}
///|
pub extend Sm4Error with @moonbitlang/core/debug.Debug::{to_repr}
///|
/// 将错误转换为可读字符串。
pub fn Sm4Error::to_string(self : Sm4Error) -> String {
match self {
InvalidKeyLength(n) =>
"sm4: 密钥长度必须为 16 字节,实际 \{n} 字节"
InvalidIvLength(n) =>
"sm4: IV 长度必须为 16 字节,实际 \{n} 字节"
InvalidCiphertextLength(n) =>
"sm4: 密文长度必须为 16 的整数倍,实际 \{n} 字节"
InvalidPadding => "sm4: PKCS7 填充不合法"
}
}
///|
/// 加密轮合成变换 T:按字节查合成表(S 盒替换 + 线性变换 L 已合并)。
fn t_round(x : UInt) -> UInt {
sm4_te0[((x >> 24) & 0xffU).reinterpret_as_int()] ^
sm4_te3[((x >> 16) & 0xffU).reinterpret_as_int()] ^
sm4_te2[((x >> 8) & 0xffU).reinterpret_as_int()] ^
sm4_te1[(x & 0xffU).reinterpret_as_int()]
}
///|
/// 密钥扩展合成变换 T':按字节查合成表(S 盒替换 + 线性变换 L' 已合并)。
fn t_key(x : UInt) -> UInt {
sm4_td0[((x >> 24) & 0xffU).reinterpret_as_int()] ^
sm4_td3[((x >> 16) & 0xffU).reinterpret_as_int()] ^
sm4_td2[((x >> 8) & 0xffU).reinterpret_as_int()] ^
sm4_td1[(x & 0xffU).reinterpret_as_int()]
}
///|
/// SM4 密码器:保存 32 个轮密钥,可反复加解密数据块。
pub struct Sm4 {
/// 轮密钥(加密序;解密时倒序使用)
rk : FixedArray[UInt]
}
///|
/// 从 16 字节密钥创建 SM4 密码器。密钥长度错误时返回 [InvalidKeyLength]。
pub fn Sm4::new(key : Bytes) -> Result[Sm4, Sm4Error] {
if key.length() != 16 {
return Err(InvalidKeyLength(key.length()))
}
let mk : FixedArray[UInt] = FixedArray::make(4, 0U)
for i in 0..<4 {
mk[i] = ((key[i * 4].to_int() << 24) |
(key[i * 4 + 1].to_int() << 16) |
(key[i * 4 + 2].to_int() << 8) |
key[i * 4 + 3].to_int()).reinterpret_as_uint()
}
// 密钥扩展:K[i+4] = K[i] ^ T'(K[i+1] ^ K[i+2] ^ K[i+3] ^ CK[i])
let rk : FixedArray[UInt] = FixedArray::make(32, 0U)
let mut k0 = mk[0] ^ sm4_fk[0]
let mut k1 = mk[1] ^ sm4_fk[1]
let mut k2 = mk[2] ^ sm4_fk[2]
let mut k3 = mk[3] ^ sm4_fk[3]
for i in 0..<32 {
let nk = k0 ^ t_key(k1 ^ k2 ^ k3 ^ sm4_ck[i])
rk[i] = nk
k0 = k1
k1 = k2
k2 = k3
k3 = nk
}
Ok({ rk, })
}
///|
/// 32 轮迭代核心。`keys` 为本轮使用的轮密钥序列。
fn Sm4::crypt_block(self : Sm4, block : Bytes, reverse : Bool) -> Bytes {
let mut x0 = ((block[0].to_int() << 24) |
(block[1].to_int() << 16) |
(block[2].to_int() << 8) |
block[3].to_int()).reinterpret_as_uint()
let mut x1 = ((block[4].to_int() << 24) |
(block[5].to_int() << 16) |
(block[6].to_int() << 8) |
block[7].to_int()).reinterpret_as_uint()
let mut x2 = ((block[8].to_int() << 24) |
(block[9].to_int() << 16) |
(block[10].to_int() << 8) |
block[11].to_int()).reinterpret_as_uint()
let mut x3 = ((block[12].to_int() << 24) |
(block[13].to_int() << 16) |
(block[14].to_int() << 8) |
block[15].to_int()).reinterpret_as_uint()
for i in 0..<32 {
let rk = if reverse { self.rk[31 - i] } else { self.rk[i] }
let xn = x0 ^ t_round(x1 ^ x2 ^ x3 ^ rk)
x0 = x1
x1 = x2
x2 = x3
x3 = xn
}
// 输出为反序:X35 X34 X33 X32
let out : FixedArray[Byte] = FixedArray::make(16, b'\x00')
let words = [x3, x2, x1, x0]
for i in 0..<4 {
let w = words[i]
out[i * 4] = ((w >> 24) & 0xffU).reinterpret_as_int().to_byte()
out[i * 4 + 1] = ((w >> 16) & 0xffU).reinterpret_as_int().to_byte()
out[i * 4 + 2] = ((w >> 8) & 0xffU).reinterpret_as_int().to_byte()
out[i * 4 + 3] = (w & 0xffU).reinterpret_as_int().to_byte()
}
Bytes::from_array(out)
}
///|
/// 加密一个 16 字节块。
pub fn Sm4::encrypt_block(self : Sm4, block : Bytes) -> Bytes {
self.crypt_block(block, false)
}
///|
/// 解密一个 16 字节块。
pub fn Sm4::decrypt_block(self : Sm4, block : Bytes) -> Bytes {
self.crypt_block(block, true)
}
///|
/// PKCS7 填充:追加 `16 - len % 16` 个值为该数的字节(恒追加 1..16 字节)。
fn pkcs7_pad(data : Bytes) -> Bytes {
let pad = 16 - data.length() % 16
let out : FixedArray[Byte] = FixedArray::make(data.length() + pad, b'\x00')
for i in 0.. Result[Bytes, Sm4Error] {
let n = data.length()
let pad = data[n - 1].to_int()
if pad < 1 || pad > 16 || pad > n {
return Err(InvalidPadding)
}
for i in 0.. Result[Sm4, Sm4Error] {
Sm4::new(key)
}
///|
/// ECB 模式加密(含 PKCS7 填充)。
pub fn encrypt_ecb(key : Bytes, plaintext : Bytes) -> Result[Bytes, Sm4Error] {
let c = match sm4_of(key) {
Ok(c) => c
Err(e) => return Err(e)
}
let padded = pkcs7_pad(plaintext)
let out : FixedArray[Byte] = FixedArray::make(padded.length(), b'\x00')
let mut off = 0
while off < padded.length() {
let ct = c.encrypt_block(padded[off:off + 16].to_owned())
for i in 0..<16 {
out[off + i] = ct[i]
}
off = off + 16
}
Ok(Bytes::from_array(out))
}
///|
/// ECB 模式解密(去 PKCS7 填充)。
pub fn decrypt_ecb(key : Bytes, ciphertext : Bytes) -> Result[Bytes, Sm4Error] {
let c = match sm4_of(key) {
Ok(c) => c
Err(e) => return Err(e)
}
if ciphertext.length() == 0 || ciphertext.length() % 16 != 0 {
return Err(InvalidCiphertextLength(ciphertext.length()))
}
let out : FixedArray[Byte] = FixedArray::make(ciphertext.length(), b'\x00')
let mut off = 0
while off < ciphertext.length() {
let pt = c.decrypt_block(ciphertext[off:off + 16].to_owned())
for i in 0..<16 {
out[off + i] = pt[i]
}
off = off + 16
}
pkcs7_unpad(Bytes::from_array(out))
}
///|
/// CBC 模式加密(含 PKCS7 填充)。`iv` 必须为 16 字节。
pub fn encrypt_cbc(
key : Bytes,
iv : Bytes,
plaintext : Bytes,
) -> Result[Bytes, Sm4Error] {
if iv.length() != 16 {
return Err(InvalidIvLength(iv.length()))
}
let c = match sm4_of(key) {
Ok(c) => c
Err(e) => return Err(e)
}
let padded = pkcs7_pad(plaintext)
let out : FixedArray[Byte] = FixedArray::make(padded.length(), b'\x00')
let prev : FixedArray[Byte] = FixedArray::make(16, b'\x00')
for i in 0..<16 {
prev[i] = iv[i]
}
let mut off = 0
while off < padded.length() {
let xored : FixedArray[Byte] = FixedArray::make(16, b'\x00')
for i in 0..<16 {
xored[i] = (padded[off + i].to_int() ^ prev[i].to_int()).to_byte()
}
let ct = c.encrypt_block(Bytes::from_array(xored))
for i in 0..<16 {
out[off + i] = ct[i]
prev[i] = ct[i]
}
off = off + 16
}
Ok(Bytes::from_array(out))
}
///|
/// CBC 模式解密(去 PKCS7 填充)。`iv` 必须为 16 字节。
pub fn decrypt_cbc(
key : Bytes,
iv : Bytes,
ciphertext : Bytes,
) -> Result[Bytes, Sm4Error] {
if iv.length() != 16 {
return Err(InvalidIvLength(iv.length()))
}
let c = match sm4_of(key) {
Ok(c) => c
Err(e) => return Err(e)
}
if ciphertext.length() == 0 || ciphertext.length() % 16 != 0 {
return Err(InvalidCiphertextLength(ciphertext.length()))
}
let out : FixedArray[Byte] = FixedArray::make(ciphertext.length(), b'\x00')
let prev : FixedArray[Byte] = FixedArray::make(16, b'\x00')
for i in 0..<16 {
prev[i] = iv[i]
}
let mut off = 0
while off < ciphertext.length() {
let pt = c.decrypt_block(ciphertext[off:off + 16].to_owned())
for i in 0..<16 {
out[off + i] = (pt[i].to_int() ^ prev[i].to_int()).to_byte()
prev[i] = ciphertext[off + i]
}
off = off + 16
}
pkcs7_unpad(Bytes::from_array(out))
}