diff --git a/vnt/Cargo.toml b/vnt/Cargo.toml index 9a02adc..07fbc14 100644 --- a/vnt/Cargo.toml +++ b/vnt/Cargo.toml @@ -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" } diff --git a/vnt/src/cipher/cipher.rs b/vnt/src/cipher/cipher.rs index 9998300..780934c 100644 --- a/vnt/src/cipher/cipher.rs +++ b/vnt/src/cipher/cipher.rs @@ -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)), 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, } } diff --git a/vnt/src/cipher/mod.rs b/vnt/src/cipher/mod.rs index 1f4e66d..99fe605 100644 --- a/vnt/src/cipher/mod.rs +++ b/vnt/src/cipher/mod.rs @@ -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; diff --git a/vnt/src/cipher/sm4_cbc.rs b/vnt/src/cipher/sm4_cbc.rs new file mode 100644 index 0000000..9b89680 --- /dev/null +++ b/vnt/src/cipher/sm4_cbc.rs @@ -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, +} + +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) -> Self { + let cipher = Sm4CipherMode::new(&key, CipherMode::Cbc).unwrap(); + Self { + key, + cipher, + finger, + } + } + + pub fn decrypt_ipv4 + AsMut<[u8]>>( + &self, + net_packet: &mut NetPacket, + ) -> 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 + AsMut<[u8]>>( + &self, + net_packet: &mut NetPacket, + ) -> 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) +} diff --git a/vnt/src/protocol/body.rs b/vnt/src/protocol/body.rs index d35cea4..384175d 100644 --- a/vnt/src/protocol/body.rs +++ b/vnt/src/protocol/body.rs @@ -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;