use crate::decrypt::des::*;
use crate::encrypt::des::*;
use crate::error::DESError;
use crate::utils::BaseString;
use crate::Traits::{BruteForce, Decrypt, Encrypt};
use gmp_mpfr_sys::mpfr::print_rnd_mode;
#[cfg(feature = "python-integration")]
use pyo3::{pyclass, pymethods};
#[derive(Default)]
#[cfg_attr(feature = "python-integration", pyclass)]
pub struct DES {
main_key: u64, sub_keys: Vec<Vec<u8>>, }
#[cfg(not(feature = "python-integration"))]
impl DES {
pub fn new(main_key: u64) -> Result<Self, DESError> {
let mut des = DES {
main_key,
sub_keys: Vec::new(),
};
des.generate_sub_keys()?;
Ok(des)
}
}
impl DES {
const PC2_TABLE: [usize; 48] = [
13, 16, 10, 23, 0, 4, 2, 27, 14, 5, 20, 9, 22, 18, 11, 3, 25, 7, 15, 6, 26, 19, 12, 1, 40,
51, 30, 36, 46, 54, 29, 39, 50, 44, 32, 47, 43, 48, 38, 55, 33, 52, 45, 41, 49, 35, 28, 31,
];
fn apply_pc2(combined_key: u64) -> u64 {
let mut subkey: u64 = 0;
for &position in DES::PC2_TABLE.iter() {
let bit = (combined_key >> (55 - position)) & 1;
subkey <<= 1;
subkey |= bit;
}
subkey }
fn prepare_key(key: u64) -> Vec<u8> {
let key_bytes = key.to_be_bytes();
let mut prepared_key = Vec::with_capacity(7); for i in 0..8 {
let byte = if i < 7 {
key_bytes[i] & 0b01111111
} else {
key_bytes[i] >> 1
};
prepared_key.push(byte);
}
prepared_key
}
fn prepared_key_to_u64(prepared_key: &[u8]) -> u64 {
assert_eq!(prepared_key.len(), 8);
let mut key_u64: u64 = 0;
for &byte in prepared_key.iter() {
key_u64 <<= 8; key_u64 |= u64::from(byte); }
key_u64 <<= 8;
key_u64
}
fn circular_left_shift(bits: u32, shift: u32) -> u32 {
let bits = bits & 0x0FFFFFFF; (bits << shift | bits >> (28 - shift)) & 0x0FFFFFFF
}
fn combine_u32_to_u64(high: u32, low: u32) -> u64 {
let high_64 = u64::from(high);
(high_64 << 32) | u64::from(low)
}
fn generate_sub_keys(&mut self) -> Result<(), DESError> {
let prepared_key = DES::prepare_key(self.main_key);
let key_56_bit = DES::prepared_key_to_u64(&prepared_key);
let c0 = (key_56_bit >> 28) as u32; let d0 = (key_56_bit & 0x0FFFFFFF) as u32;
for &shift_amount in [1, 1, 2, 2, 2, 2, 2, 2, 1, 2, 2, 2, 2, 2, 2, 1].iter() {
let c_shifted = DES::circular_left_shift(c0, shift_amount);
let d_shifted = DES::circular_left_shift(d0, shift_amount);
let cd_shifted = DES::combine_u32_to_u64(c_shifted, d_shifted);
let sub_key = DES::apply_pc2(cd_shifted);
self.sub_keys.push(sub_key.to_be_bytes().to_vec());
}
Ok(())
}
fn get_blocks(input: BaseString) -> Vec<Vec<u8>> {
let bytes = input.data.as_bytes();
let mut blocks = Vec::new();
let padding_size = (8 - (bytes.len() % 8)) % 8;
let padding = vec![0; padding_size];
let padded_bytes = [bytes, &padding].concat();
for block in padded_bytes.chunks(8) {
blocks.push(block.to_vec());
}
blocks
}
fn initial_permutation(input_block: Vec<u8>) -> Vec<u8> {
let ip_table: [usize; 64] = [
58, 50, 42, 34, 26, 18, 10, 2, 60, 52, 44, 36, 28, 20, 12, 4, 62, 54, 46, 38, 30, 22,
14, 6, 64, 56, 48, 40, 32, 24, 16, 8, 57, 49, 41, 33, 25, 17, 9, 1, 59, 51, 43, 35, 27,
19, 11, 3, 61, 53, 45, 37, 29, 21, 13, 5, 63, 55, 47, 39, 31, 23, 15, 7,
];
let mut output_block = vec![0u8; 8];
for (i, &position) in ip_table.iter().enumerate() {
let bit_position = position - 1; let byte_index = bit_position / 8;
let bit_index = bit_position % 8;
let bit = (input_block[byte_index] >> (7 - bit_index)) & 1;
output_block[i / 8] |= bit << (7 - (i % 8));
}
output_block
}
#[cfg(feature = "debug")]
pub fn print_keys(&self) {
println!("main key: {:?}", self.main_key);
for key in &self.sub_keys {
println!("subkey: {:?}", key);
}
}
}
#[cfg(feature = "python-integration")]
mod python_integration {
use super::*;
use pyo3::prelude::*;
use pyo3::{pyclass, pymethods, PyResult};
use std::collections::HashMap;
#[pymethods]
impl DES {
#[new]
pub fn new(main_key: u64) -> Result<Self, DESError> {
let mut des = DES {
main_key,
sub_keys: Vec::new(),
};
des.generate_sub_keys()?;
Ok(des)
}
pub fn encrypt(&self, input: String) -> PyResult<String> {
match Encrypt::encrypt(self, input) {
Ok(s) => Ok(s),
Err(e) => Err(pyo3::exceptions::PyException::new_err(format!("{:?}", e))),
}
}
pub fn decrypt(&self, input: String) -> PyResult<String> {
match Decrypt::decrypt(self, input) {
Ok(s) => Ok(s),
Err(e) => Err(pyo3::exceptions::PyException::new_err(format!("{:?}", e))),
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn get_blocks_returns_empty_vector_for_empty_input() {
let input = BaseString::new(String::from(""));
assert_eq!(DES::get_blocks(input), Vec::<Vec<u8>>::new());
}
#[test]
fn get_blocks_returns_single_block_for_input_less_than_8_bytes() {
let input = BaseString::new(String::from("1234567"));
assert_eq!(
DES::get_blocks(input),
vec![vec![49, 50, 51, 52, 53, 54, 55, 0]]
);
}
#[test]
fn get_blocks_returns_single_block_for_input_of_8_bytes() {
let input = BaseString::new(String::from("12345678"));
assert_eq!(
DES::get_blocks(input),
vec![vec![49, 50, 51, 52, 53, 54, 55, 56]]
);
}
#[test]
fn get_blocks_returns_two_blocks_for_input_of_9_bytes() {
let input = BaseString::new(String::from("123456789"));
assert_eq!(
DES::get_blocks(input),
vec![
vec![49, 50, 51, 52, 53, 54, 55, 56],
vec![57, 0, 0, 0, 0, 0, 0, 0]
]
);
}
#[test]
fn get_blocks_returns_correct_blocks_for_input_of_16_bytes() {
let input = BaseString::new(String::from("1234567812345678"));
assert_eq!(
DES::get_blocks(input),
vec![
vec![49, 50, 51, 52, 53, 54, 55, 56],
vec![49, 50, 51, 52, 53, 54, 55, 56]
]
);
}
#[test]
fn get_blocks_returns_correct_blocks_for_input_of_17_bytes() {
let input = BaseString::new(String::from("12345678123456789"));
assert_eq!(
DES::get_blocks(input),
vec![
vec![49, 50, 51, 52, 53, 54, 55, 56],
vec![49, 50, 51, 52, 53, 54, 55, 56],
vec![57, 0, 0, 0, 0, 0, 0, 0]
]
);
}
#[test]
fn get_blocks_returns_correct_blocks_for_input_of_24_bytes() {
let input = BaseString::new(String::from("123456781234567812345678"));
assert_eq!(
DES::get_blocks(input),
vec![
vec![49, 50, 51, 52, 53, 54, 55, 56],
vec![49, 50, 51, 52, 53, 54, 55, 56],
vec![49, 50, 51, 52, 53, 54, 55, 56]
]
);
}
#[test]
fn get_blocks_returns_correct_blocks_for_input_of_25_bytes() {
let input = BaseString::new(String::from("1234567812345678123456789"));
assert_eq!(
DES::get_blocks(input),
vec![
vec![49, 50, 51, 52, 53, 54, 55, 56],
vec![49, 50, 51, 52, 53, 54, 55, 56],
vec![49, 50, 51, 52, 53, 54, 55, 56],
vec![57, 0, 0, 0, 0, 0, 0, 0]
]
);
}
#[test]
fn get_blocks_returns_correct_blocks_for_input_of_32_bytes() {
let input = BaseString::new(String::from("12345678123456781234567812345678"));
assert_eq!(
DES::get_blocks(input),
vec![
vec![49, 50, 51, 52, 53, 54, 55, 56],
vec![49, 50, 51, 52, 53, 54, 55, 56],
vec![49, 50, 51, 52, 53, 54, 55, 56],
vec![49, 50, 51, 52, 53, 54, 55, 56]
]
);
}
#[test]
fn get_blocks_handles_ascii_characters_correctly() {
let input = BaseString::new(String::from("abc"));
assert_eq!(
DES::get_blocks(input),
vec![vec![97, 98, 99, 0, 0, 0, 0, 0]]
);
}
#[test]
fn get_blocks_handles_unicode_characters_correctly() {
let input = BaseString::new(String::from("あいう"));
assert_eq!(
DES::get_blocks(input),
vec![
vec![227, 129, 130, 227, 129, 132, 227, 129],
vec![134, 0, 0, 0, 0, 0, 0, 0]
]
);
}
#[test]
fn get_blocks_handles_mixed_ascii_and_unicode_characters_correctly() {
let input = BaseString::new(String::from("aあb"));
assert_eq!(
DES::get_blocks(input),
vec![vec![97, 227, 129, 130, 98, 0, 0, 0]]
);
}
#[test]
fn get_blocks_handles_long_unicode_characters_correctly() {
let input = BaseString::new(String::from(
"åäöåäöåäöåäöååäöåäöåäöåäöåääöåäöåäöåäöåäöåäöå",
));
assert_eq!(
DES::get_blocks(input),
vec![
vec![195, 165, 195, 164, 195, 182, 195, 165],
vec![195, 164, 195, 182, 195, 165, 195, 164],
vec![195, 182, 195, 165, 195, 164, 195, 182],
vec![195, 165, 195, 165, 195, 164, 195, 182],
vec![195, 165, 195, 164, 195, 182, 195, 165],
vec![195, 164, 195, 182, 195, 165, 195, 164],
vec![195, 182, 195, 165, 195, 164, 195, 164],
vec![195, 182, 195, 165, 195, 164, 195, 182],
vec![195, 165, 195, 164, 195, 182, 195, 165],
vec![195, 164, 195, 182, 195, 165, 195, 164],
vec![195, 182, 195, 165, 195, 164, 195, 182],
vec![195, 165, 0, 0, 0, 0, 0, 0]
]
);
}
#[test]
fn initial_permutation_for_single_block_less_than_8_bytes() {
let input = BaseString::new(String::from("1234567"));
let blocks = DES::get_blocks(input);
let block = &blocks[0];
let permuted_block = DES::initial_permutation(block.clone());
let expected_permuted_block: Vec<u8> = vec![0, 127, 120, 85, 0, 127, 0, 102];
assert_eq!(permuted_block, expected_permuted_block);
}
#[test]
fn test_key_gen() {
let key = 0b1111000011110000111100001111000011110000111100001111000011110000u64;
let des = DES::new(key);
assert!(des.is_ok())
}
#[test]
fn test_key_gen_simple() {
let key = 0b0000000;
let des = DES::new(key);
assert!(des.is_ok())
}
}