use crate::crypto_systems::mix::{Mix, MixKey};
use crate::error::MixError;
use crate::prelude::ALPHABET_LEN;
use crate::Traits::{BruteForce, Decrypt};
use itertools::Itertools;
use rayon::prelude::*;
use std::collections::{HashMap, HashSet};
use std::sync::Mutex;
use strsim::levenshtein;
pub struct PermutationIterator {
vec_len: usize,
current: [Vec<usize>; 3],
max_value: usize,
exhausted: bool,
known_keys: [Option<Vec<usize>>; 3],
}
impl PermutationIterator {
fn new(vec_len: usize, known_keys: [Option<Vec<usize>>; 3]) -> Self {
let mut current = [vec![0; vec_len], vec![0; vec_len], vec![0; vec_len]];
for (i, key) in known_keys.iter().enumerate() {
if let Some(ref key_val) = key {
current[i] = key_val.clone();
}
}
Self {
vec_len,
current,
max_value: 27,
exhausted: false,
known_keys,
}
}
pub fn all_permutations(&self) -> impl ParallelIterator<Item = [Vec<usize>; 3]> {
let iterators: Vec<_> = self
.known_keys
.iter()
.map(|key| {
if let Some(ref key_val) = key {
vec![key_val.clone()].into_iter()
} else {
Self::generate_permutations(0, self.max_value, self.vec_len).into_iter()
}
})
.collect();
iterators
.into_iter()
.multi_cartesian_product()
.par_bridge()
.map(|perm| [perm[0].clone(), perm[1].clone(), perm[2].clone()])
}
fn generate_permutations(min: usize, max: usize, length: usize) -> Vec<Vec<usize>> {
if min > max || length == 0 {
return vec![];
}
if length == 1 {
return (min..=max).map(|x| vec![x]).collect();
}
let mut permutations = vec![];
for perm in Self::generate_permutations(min, max, length - 1) {
for i in min..=max {
let mut extended_perm = perm.clone();
extended_perm.push(i);
permutations.push(extended_perm);
}
}
permutations
}
fn increment(&mut self, idx: usize) -> bool {
if idx >= self.current.len() {
self.exhausted = true;
return false;
}
if self.known_keys[idx].is_some() {
return self.increment(idx + 1);
}
let vec = &mut self.current[idx];
for i in 0..vec.len() {
if vec[i] < self.max_value {
vec[i] += 1;
return true;
} else {
vec[i] = 0; }
}
self.increment(idx + 1)
}
}
impl Iterator for PermutationIterator {
type Item = [Vec<usize>; 3];
fn next(&mut self) -> Option<Self::Item> {
if self.exhausted {
None
} else {
let result = self.current.clone();
self.increment(0);
Some(result)
}
}
}
type BruteForceResult = HashMap<MixKey, String>;
impl BruteForce<BruteForceResult, MixError, PermutationIterator, [Option<Vec<usize>>; 3]> for Mix {
#[cfg(not(feature = "parallel"))]
fn brute_force(
&mut self,
cipher_text: String,
clear_text: Option<String>,
known_keys: [Option<Vec<usize>>; 3],
) -> Result<BruteForceResult, MixError> {
let mut iterator = self.gen_permutations(known_keys)?;
let mut result = HashMap::new();
let mut filtered_result = HashMap::new();
let similarity_threshold = 0.90;
while let Some(permutation) = iterator.next() {
let a = permutation[0].clone();
let b = permutation[1].clone();
let c = permutation[2].clone();
let key = MixKey::new(a, b, c).map_err(|_| MixError::InvalidKey)?;
let mix = Mix::new(key.clone());
let decryption_result = mix.decrypt(cipher_text.clone())?;
if let Some(ref plain_text) = clear_text {
let similarity = calculate_similarity(&decryption_result, plain_text);
if similarity >= similarity_threshold {
filtered_result.insert(key.clone(), decryption_result.clone());
if similarity == 1.0 {
println!("Found exact match: {:?}", key);
break;
}
}
} else {
result.insert(key.clone(), decryption_result);
}
}
let final_result = if !filtered_result.is_empty() {
filtered_result
} else if clear_text.is_some() {
println!(
"No matches found with similarity >= {}.",
similarity_threshold
);
result
} else {
HashMap::new()
};
Ok(final_result)
}
#[cfg(feature = "parallel")]
fn brute_force(
&mut self,
cipher_text: String,
clear_text: Option<String>,
known_keys: [Option<Vec<usize>>; 3],
) -> Result<BruteForceResult, MixError> {
let iterator = self.gen_permutations(known_keys)?;
let all_permutations = iterator.all_permutations();
let results = Mutex::new(HashMap::new());
all_permutations.for_each(|permutation| {
let a = permutation[0].clone();
let b = permutation[1].clone();
let c = permutation[2].clone();
let key = match MixKey::new(a, b, c) {
Ok(k) => k,
Err(_) => return, };
let mix = Mix::new(key.clone());
let decryption_result = match mix.decrypt(cipher_text.clone()) {
Ok(res) => res,
Err(_) => return, };
if let Some(ref plain_text) = clear_text {
let similarity = calculate_similarity(&decryption_result, plain_text);
if similarity >= 0.90 {
let mut res = results.lock().unwrap();
res.insert(key.clone(), decryption_result.clone());
}
}
});
Ok(results.into_inner().unwrap())
}
fn gen_permutations(
&mut self,
known_keys: [Option<Vec<usize>>; 3],
) -> Result<PermutationIterator, MixError> {
Ok(PermutationIterator::new(self.key.len, known_keys))
}
}
fn calculate_similarity(text1: &str, text2: &str) -> f64 {
let distance = levenshtein(text1, text2);
let max_len = text1.len().max(text2.len());
if max_len == 0 {
return 1.0; }
1.0 - (distance as f64 / max_len as f64)
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
#[test]
fn brute_force_returns_correct_decryption_for_known_key() {
let k1 = vec![13, 5, 1, 0];
let k2 = vec![12, 4, 16, 8];
let k3 = vec![1, 24, 2, 21];
let clear_text = "orosmoln".to_string();
let key = MixKey::new(k1.clone(), k2.clone(), k3.clone()).expect("Key is invalid");
let mut mix = Mix::new(key.clone());
let cipher_text = String::from("HJUMTKLC");
let known_keys: [Option<Vec<usize>>; 3] = [Some(k1), None, Some(k3)];
let result = mix
.brute_force(cipher_text, Some(clear_text.clone()), known_keys)
.unwrap();
let result_key = result.keys().find(|&k| k == &key).unwrap();
assert_eq!(result_key, &key);
assert_eq!(result.values().next().unwrap(), &clear_text);
}
}