use crate::error::{Result, SZipError};
use aes::Aes256;
use ctr::{
cipher::{KeyIvInit, StreamCipher},
Ctr128BE,
};
use hmac::{Hmac, Mac};
use pbkdf2::pbkdf2_hmac;
use sha1::Sha1;
type HmacSha1 = Hmac<Sha1>;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AesStrength {
Aes256,
}
impl AesStrength {
pub fn salt_size(&self) -> usize {
match self {
AesStrength::Aes256 => 16,
}
}
pub fn key_size(&self) -> usize {
match self {
AesStrength::Aes256 => 32,
}
}
pub fn derived_key_size(&self) -> usize {
self.key_size() * 2 + 2 }
pub fn to_winzip_code(&self) -> u16 {
match self {
AesStrength::Aes256 => 0x03,
}
}
}
pub struct AesEncryptor {
strength: AesStrength,
salt: Vec<u8>,
password_verify: [u8; 2],
encryption_key: Vec<u8>,
#[allow(dead_code)] auth_key: Vec<u8>,
hmac: HmacSha1,
byte_offset: u64,
}
impl AesEncryptor {
pub fn new(password: &str, strength: AesStrength) -> Result<Self> {
let salt = generate_salt(strength.salt_size())?;
let derived_key_size = strength.derived_key_size();
let mut derived_keys = vec![0u8; derived_key_size];
pbkdf2_hmac::<Sha1>(password.as_bytes(), &salt, 1000, &mut derived_keys);
let key_size = strength.key_size();
let encryption_key = derived_keys[..key_size].to_vec();
let auth_key = derived_keys[key_size..key_size * 2].to_vec();
let password_verify = [derived_keys[key_size * 2], derived_keys[key_size * 2 + 1]];
let hmac = HmacSha1::new_from_slice(&auth_key)
.map_err(|e| SZipError::InvalidFormat(format!("HMAC init failed: {}", e)))?;
Ok(Self {
strength,
salt,
password_verify,
encryption_key,
auth_key,
hmac,
byte_offset: 0,
})
}
pub fn salt(&self) -> &[u8] {
&self.salt
}
pub fn password_verify(&self) -> &[u8; 2] {
&self.password_verify
}
pub fn strength(&self) -> AesStrength {
self.strength
}
pub fn update_hmac(&mut self, data: &[u8]) {
self.hmac.update(data);
}
pub fn encrypt(&mut self, data: &mut [u8]) -> Result<()> {
let block_number = self.byte_offset / 16;
let mut iv = [0u8; 16];
iv[8..16].copy_from_slice(&block_number.to_be_bytes());
let key = self.encryption_key.as_slice();
let mut cipher = Ctr128BE::<Aes256>::new(key.into(), &iv.into());
let partial = (self.byte_offset % 16) as usize;
if partial != 0 {
let mut discard = vec![0u8; partial];
cipher.apply_keystream(&mut discard);
}
cipher.apply_keystream(data);
self.byte_offset += data.len() as u64;
Ok(())
}
pub fn finalize(self) -> Vec<u8> {
let mac = self.hmac.finalize();
mac.into_bytes()[..10].to_vec()
}
}
pub struct AesDecryptor {
#[allow(dead_code)] strength: AesStrength,
encryption_key: Vec<u8>,
#[allow(dead_code)] auth_key: Vec<u8>,
#[allow(dead_code)] password_verify: [u8; 2],
hmac: HmacSha1,
byte_offset: u64,
}
impl AesDecryptor {
pub fn new(
password: &str,
strength: AesStrength,
salt: &[u8],
password_verify: &[u8; 2],
) -> Result<Self> {
if salt.len() != strength.salt_size() {
return Err(SZipError::InvalidFormat(format!(
"Invalid salt size: expected {}, got {}",
strength.salt_size(),
salt.len()
)));
}
let derived_key_size = strength.derived_key_size();
let mut derived_keys = vec![0u8; derived_key_size];
pbkdf2_hmac::<Sha1>(password.as_bytes(), salt, 1000, &mut derived_keys);
let key_size = strength.key_size();
let encryption_key = derived_keys[..key_size].to_vec();
let auth_key = derived_keys[key_size..key_size * 2].to_vec();
let expected_pw_verify = [derived_keys[key_size * 2], derived_keys[key_size * 2 + 1]];
if &expected_pw_verify != password_verify {
return Err(SZipError::IncorrectPassword);
}
let hmac = HmacSha1::new_from_slice(&auth_key)
.map_err(|e| SZipError::InvalidFormat(format!("HMAC init failed: {}", e)))?;
Ok(Self {
strength,
encryption_key,
auth_key,
password_verify: *password_verify,
hmac,
byte_offset: 0,
})
}
pub fn decrypt(&mut self, data: &mut [u8]) -> Result<()> {
let block_number = self.byte_offset / 16;
let mut iv = [0u8; 16];
iv[8..16].copy_from_slice(&block_number.to_be_bytes());
let key = self.encryption_key.as_slice();
let mut cipher = Ctr128BE::<Aes256>::new(key.into(), &iv.into());
let partial = (self.byte_offset % 16) as usize;
if partial != 0 {
let mut discard = vec![0u8; partial];
cipher.apply_keystream(&mut discard);
}
cipher.apply_keystream(data);
self.byte_offset += data.len() as u64;
Ok(())
}
pub fn update_hmac(&mut self, data: &[u8]) {
self.hmac.update(data);
}
pub fn verify_auth_code(&self, auth_code: &[u8]) -> Result<()> {
let expected = self.hmac.clone().finalize();
let expected_bytes = &expected.into_bytes()[..10];
if expected_bytes != auth_code {
return Err(SZipError::InvalidFormat(
"Authentication failed: file may be corrupted or password is incorrect".to_string(),
));
}
Ok(())
}
}
fn generate_salt(size: usize) -> Result<Vec<u8>> {
#[cfg(feature = "encryption")]
{
let mut salt = vec![0u8; size];
getrandom::getrandom(&mut salt).map_err(|e| {
SZipError::EncryptionError(format!("Failed to generate random salt: {}", e))
})?;
Ok(salt)
}
#[cfg(not(feature = "encryption"))]
{
use std::time::{SystemTime, UNIX_EPOCH};
let mut salt = vec![0u8; size];
let seed = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos() as u64;
for (i, byte) in salt.iter_mut().enumerate() {
*byte = ((seed.wrapping_mul(i as u64 + 1).wrapping_add(i as u64)) % 256) as u8;
}
Ok(salt)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_aes_strength_sizes() {
assert_eq!(AesStrength::Aes256.salt_size(), 16);
assert_eq!(AesStrength::Aes256.key_size(), 32);
assert_eq!(AesStrength::Aes256.to_winzip_code(), 0x03);
}
#[test]
fn test_encrypt_decrypt_roundtrip() {
let password = "test_password_123";
let plaintext = b"Hello, encrypted world!";
let mut encryptor = AesEncryptor::new(password, AesStrength::Aes256).unwrap();
let salt = encryptor.salt().to_vec();
let password_verify = *encryptor.password_verify();
let mut encrypted = plaintext.to_vec();
encryptor.encrypt(&mut encrypted).unwrap();
let auth_code = encryptor.finalize();
assert_ne!(encrypted, plaintext);
let mut decryptor =
AesDecryptor::new(password, AesStrength::Aes256, &salt, &password_verify).unwrap();
decryptor.decrypt(&mut encrypted).unwrap();
decryptor.verify_auth_code(&auth_code).unwrap();
assert_eq!(encrypted, plaintext);
}
#[test]
fn test_wrong_password() {
let password = "correct_password";
let wrong_password = "wrong_password";
let plaintext = b"Secret data";
let mut encryptor = AesEncryptor::new(password, AesStrength::Aes256).unwrap();
let salt = encryptor.salt().to_vec();
let password_verify = *encryptor.password_verify();
let mut encrypted = plaintext.to_vec();
encryptor.encrypt(&mut encrypted).unwrap();
let result =
AesDecryptor::new(wrong_password, AesStrength::Aes256, &salt, &password_verify);
assert!(result.is_err(), "Expected password verification to fail");
}
#[test]
fn test_multi_chunk_encrypt_decrypt_roundtrip() {
let password = "multi_chunk_password";
let chunk1 = vec![0xAAu8; 3 * 1024 * 1024]; let chunk2 = vec![0xBBu8; 2 * 1024 * 1024];
let mut encryptor = AesEncryptor::new(password, AesStrength::Aes256).unwrap();
let salt = encryptor.salt().to_vec();
let password_verify = *encryptor.password_verify();
let mut enc1 = chunk1.clone();
encryptor.encrypt(&mut enc1).unwrap();
let mut enc2 = chunk2.clone();
encryptor.encrypt(&mut enc2).unwrap();
let auth_code = encryptor.finalize();
assert_ne!(enc1, chunk1, "chunk1 should be encrypted");
assert_ne!(enc2, chunk2, "chunk2 should be encrypted");
let mut decryptor =
AesDecryptor::new(password, AesStrength::Aes256, &salt, &password_verify).unwrap();
decryptor.decrypt(&mut enc1).unwrap();
decryptor.decrypt(&mut enc2).unwrap();
decryptor.verify_auth_code(&auth_code).unwrap();
assert_eq!(enc1, chunk1, "chunk1 decryption mismatch");
assert_eq!(enc2, chunk2, "chunk2 decryption mismatch");
}
#[test]
fn test_ctr_keystreams_differ_across_chunks() {
let password = "keystream_test_password";
let plaintext_block = vec![0xCCu8; 1024];
let mut encryptor = AesEncryptor::new(password, AesStrength::Aes256).unwrap();
let mut c1 = plaintext_block.clone();
encryptor.encrypt(&mut c1).unwrap();
let mut c2 = plaintext_block.clone();
encryptor.encrypt(&mut c2).unwrap();
assert_ne!(c1, c2, "CTR keystream must not repeat across calls");
assert_ne!(c1, plaintext_block);
assert_ne!(c2, plaintext_block);
}
#[test]
fn test_single_chunk_still_works_after_offset_fix() {
let password = "single_chunk_regression";
let plaintext = b"The quick brown fox jumps over the lazy dog.".to_vec();
let mut encryptor = AesEncryptor::new(password, AesStrength::Aes256).unwrap();
let salt = encryptor.salt().to_vec();
let password_verify = *encryptor.password_verify();
let mut ciphertext = plaintext.clone();
encryptor.encrypt(&mut ciphertext).unwrap();
let auth_code = encryptor.finalize();
assert_ne!(ciphertext, plaintext);
let mut decryptor =
AesDecryptor::new(password, AesStrength::Aes256, &salt, &password_verify).unwrap();
decryptor.decrypt(&mut ciphertext).unwrap();
decryptor.verify_auth_code(&auth_code).unwrap();
assert_eq!(ciphertext, plaintext);
}
}