支持sm4-cbc加密

This commit is contained in:
lubeilin
2023-09-26 21:53:07 +08:00
parent 9c098c55c9
commit ca4e8d14f0
5 changed files with 187 additions and 4 deletions
+2 -2
View File
@@ -1,6 +1,6 @@
[package]
name = "vnt"
version = "1.2.4"
version = "1.2.5"
edition = "2021"
# See more keys and their definitions at https://doc.rust-lang.org/cargo/reference/manifest.html
@@ -29,7 +29,7 @@ stun-format = { version = "1.0.1", features = ["fmt", "rfc3489"] }
rsa = { version = "0.7.2", features = [] }
spki = { version = "0.6.0", features = ["fingerprint", "alloc"] }
openssl-sys = { git = "https://github.com/lbl8603/rust-openssl" ,optional = true}
libsm = {git="https://github.com/lbl8603/libsm"}
[target.'cfg(any(target_os = "linux",target_os = "macos"))'.dependencies]
tun = { path = "./rust-tun" }
+12
View File
@@ -6,6 +6,7 @@ use crate::cipher::aes_gcm_cipher::AesGcmCipher;
use crate::cipher::openssl_aes_ecb::AesEcbCipher;
#[cfg(feature = "ring-cipher")]
use crate::cipher::ring_aes_gcm_cipher::AesGcmCipher;
use crate::cipher::sm4_cbc::Sm4CbcCipher;
use crate::cipher::{aes_cbc, Finger};
use crate::protocol::NetPacket;
use aes_cbc::AesCbcCipher;
@@ -18,6 +19,7 @@ pub enum CipherModel {
AesGcm,
AesCbc,
AesEcb,
Sm4Cbc,
}
impl FromStr for CipherModel {
@@ -28,6 +30,7 @@ impl FromStr for CipherModel {
"aes_gcm" => Ok(CipherModel::AesGcm),
"aes_cbc" => Ok(CipherModel::AesCbc),
"aes_ecb" => Ok(CipherModel::AesEcb),
"sm4_cbc" => Ok(CipherModel::Sm4Cbc),
_ => Err(format!("not match '{}', enum:aes_gcm/aes_cbc/aes_ecb", s)),
}
}
@@ -38,6 +41,7 @@ pub enum Cipher {
AesGcm((AesGcmCipher, Vec<u8>)),
AesCbc(AesCbcCipher),
AesEcb(AesEcbCipher),
Sm4Cbc(Sm4CbcCipher),
None,
}
@@ -80,6 +84,10 @@ impl Cipher {
Cipher::AesEcb(aes)
}
}
CipherModel::Sm4Cbc => {
let aes = Sm4CbcCipher::new_128(key[..16].try_into().unwrap(), finger);
Cipher::Sm4Cbc(aes)
}
}
} else {
Cipher::None
@@ -107,6 +115,7 @@ impl Cipher {
Cipher::AesGcm((aes_gcm, _)) => aes_gcm.decrypt_ipv4(net_packet),
Cipher::AesCbc(aes_cbc) => aes_cbc.decrypt_ipv4(net_packet),
Cipher::AesEcb(aes_ecb) => aes_ecb.decrypt_ipv4(net_packet),
Cipher::Sm4Cbc(sm4_cbc) => sm4_cbc.decrypt_ipv4(net_packet),
Cipher::None => {
if net_packet.is_encrypt() {
return Err(io::Error::new(io::ErrorKind::Other, "not key"));
@@ -123,6 +132,7 @@ impl Cipher {
Cipher::AesGcm((aes_gcm, _)) => aes_gcm.encrypt_ipv4(net_packet),
Cipher::AesCbc(aes_cbc) => aes_cbc.encrypt_ipv4(net_packet),
Cipher::AesEcb(aes_ecb) => aes_ecb.encrypt_ipv4(net_packet),
Cipher::Sm4Cbc(sm4_cbc) => sm4_cbc.encrypt_ipv4(net_packet),
Cipher::None => Ok(()),
}
}
@@ -131,6 +141,7 @@ impl Cipher {
Cipher::AesGcm((aes_gcm, _)) => aes_gcm.finger.as_ref(),
Cipher::AesCbc(aes_cbc) => aes_cbc.finger.as_ref(),
Cipher::AesEcb(aes_ecb) => aes_ecb.finger.as_ref(),
Cipher::Sm4Cbc(sm4_cbc) => sm4_cbc.finger.as_ref(),
Cipher::None => None,
};
if let Some(finger) = finger {
@@ -144,6 +155,7 @@ impl Cipher {
Cipher::AesGcm((_, key)) => Some(key),
Cipher::AesCbc(aes_cbc) => Some(aes_cbc.key()),
Cipher::AesEcb(aes_ecb) => Some(aes_ecb.key()),
Cipher::Sm4Cbc(sm4_cbc) => Some(sm4_cbc.key()),
Cipher::None => None,
}
}
+1 -1
View File
@@ -11,7 +11,7 @@ mod rsa_cipher;
#[cfg(any(feature = "openssl-vendored", feature = "openssl"))]
mod openssl_aes_ecb;
mod sm4_cbc;
pub use cipher::Cipher;
pub use cipher::CipherModel;
pub use finger::Finger;
+171
View File
@@ -0,0 +1,171 @@
use crate::cipher::Finger;
use crate::protocol::{NetPacket, HEAD_LEN};
use libsm::sm4::cipher_mode::CipherMode;
use libsm::sm4::Sm4CipherMode;
use rand::RngCore;
use std::io;
pub struct Sm4CbcCipher {
key: [u8; 16],
pub(crate) cipher: Sm4CipherMode,
pub(crate) finger: Option<Finger>,
}
impl Clone for Sm4CbcCipher {
fn clone(&self) -> Self {
let cipher = Sm4CipherMode::new(&self.key, CipherMode::Cbc).unwrap();
Self {
key: self.key,
cipher,
finger: self.finger.clone(),
}
}
}
impl Sm4CbcCipher {
pub fn key(&self) -> &[u8] {
&self.key
}
}
impl Sm4CbcCipher {
pub fn new_128(key: [u8; 16], finger: Option<Finger>) -> Self {
let cipher = Sm4CipherMode::new(&key, CipherMode::Cbc).unwrap();
Self {
key,
cipher,
finger,
}
}
pub fn decrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
if !net_packet.is_encrypt() {
//未加密的数据直接丢弃
return Err(io::Error::new(io::ErrorKind::Other, "not encrypt"));
}
if let Some(finger) = &self.finger {
let mut nonce_raw = [0; 12];
nonce_raw[0..4].copy_from_slice(&net_packet.source().octets());
nonce_raw[4..8].copy_from_slice(&net_packet.destination().octets());
nonce_raw[8] = net_packet.protocol().into();
nonce_raw[9] = net_packet.transport_protocol();
nonce_raw[10] = net_packet.is_gateway() as u8;
nonce_raw[11] = net_packet.source_ttl();
let len = net_packet.payload().len();
if len < 12 {
return Err(io::Error::new(io::ErrorKind::Other, "payload len <12"));
}
let secret_body = &net_packet.payload()[..len - 12];
let finger = finger.calculate_finger(&nonce_raw, secret_body);
if &finger != &net_packet.payload()[len - 12..] {
return Err(io::Error::new(io::ErrorKind::Other, "finger err"));
}
net_packet.set_data_len(net_packet.data_len() - finger.len())?;
}
let payload = net_packet.payload();
let len = payload.len();
if len < 16 || len > 1024 * 4 {
log::error!("数据异常,长度{}小于16或大于4096", len);
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
let mut out = [0u8; 1024 * 4];
let data = &payload[..len - 16];
let iv = &payload[len - 16..];
match self.cipher.decrypt(data, iv, &mut out) {
Ok(len) => {
let src_net_packet = NetPacket::new(&out[..len])?;
if src_net_packet.source() != net_packet.source() {
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
if src_net_packet.destination() != net_packet.destination() {
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
if src_net_packet.protocol() != net_packet.protocol() {
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
if src_net_packet.transport_protocol() != net_packet.transport_protocol() {
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
if src_net_packet.is_gateway() != net_packet.is_gateway() {
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
if src_net_packet.source_ttl() != net_packet.source_ttl() {
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
net_packet.set_data_len(len)?;
net_packet.set_payload(src_net_packet.payload())?;
net_packet.set_encrypt_flag(false);
Ok(())
}
Err(e) => Err(io::Error::new(
io::ErrorKind::Other,
format!("sm4_cbc解密失败:{}", e),
)),
}
}
/// net_packet 必须预留足够长度
/// data_len是有效载荷的长度
pub fn encrypt_ipv4<B: AsRef<[u8]> + AsMut<[u8]>>(
&self,
net_packet: &mut NetPacket<B>,
) -> io::Result<()> {
let mut out = [0u8; 1024 * 4];
let mut iv = [0u8; 16];
rand::thread_rng().fill_bytes(&mut iv);
if net_packet.buffer().len() > 1024 * 4 - 32 {
log::error!(
"数据异常,长度{}大于1024 * 4 - 32",
net_packet.buffer().len()
);
return Err(io::Error::new(io::ErrorKind::Other, "data err"));
}
match self.cipher.encrypt(net_packet.buffer(), &iv, &mut out) {
Ok(len) => {
out[len..len + 16].copy_from_slice(&iv);
net_packet.set_data_len(HEAD_LEN + len + 16)?;
net_packet.set_payload(&out[..len + 16])?;
net_packet.set_encrypt_flag(true);
if let Some(finger) = &self.finger {
let mut nonce_raw = [0; 12];
nonce_raw[0..4].copy_from_slice(&net_packet.source().octets());
nonce_raw[4..8].copy_from_slice(&net_packet.destination().octets());
nonce_raw[8] = net_packet.protocol().into();
nonce_raw[9] = net_packet.transport_protocol();
nonce_raw[10] = net_packet.is_gateway() as u8;
nonce_raw[11] = net_packet.source_ttl();
let finger = finger.calculate_finger(&nonce_raw, net_packet.payload());
let src_data_len = net_packet.data_len();
//设置实际长度
net_packet.set_data_len(src_data_len + finger.len())?;
net_packet.buffer_mut()[src_data_len..].copy_from_slice(&finger);
}
Ok(())
}
Err(e) => Err(io::Error::new(
io::ErrorKind::Other,
format!("sm4_cbc加密失败:{}", e),
)),
}
}
}
#[test]
fn test_sm4_ecb() {
let d = Sm4CbcCipher::new_128([0; 16], Some(Finger::new("123")));
let mut p = NetPacket::new_encrypt([1; 1024]).unwrap();
let src = p.buffer().to_vec();
d.encrypt_ipv4(&mut p).unwrap();
d.decrypt_ipv4(&mut p).unwrap();
assert_eq!(p.buffer(), &src);
let d = Sm4CbcCipher::new_128([0; 16], None);
let mut p = NetPacket::new_encrypt([1; 102]).unwrap();
let src = p.buffer().to_vec();
d.encrypt_ipv4(&mut p).unwrap();
d.decrypt_ipv4(&mut p).unwrap();
assert_eq!(p.buffer(), &src)
}
+1 -1
View File
@@ -1,6 +1,6 @@
use std::{fmt, io};
pub const ENCRYPTION_RESERVED: usize = 32 + 12;
pub const ENCRYPTION_RESERVED: usize = 16 + 32 + 12;
pub const AES_GCM_ENCRYPTION_RESERVED: usize = 32;
pub const RSA_ENCRYPTION_RESERVED: usize = 32;