use std::{
error::Error,
fmt::{self, Display},
marker::PhantomData,
};
use crate::serializer::MutableSerializer;
use bevy::log;
use serde::{Deserialize, Serialize};
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
pub enum CustomCryptClientPacket {
String(String),
}
impl Default for CustomCryptClientPacket {
fn default() -> Self {
CustomCryptClientPacket::String(String::new())
}
}
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
pub enum CustomCryptServerPacket {
String(String),
}
impl Default for CustomCryptServerPacket {
fn default() -> Self {
CustomCryptServerPacket::String(String::new())
}
}
#[derive(Debug)]
pub struct CustomSerializationError;
impl Display for CustomSerializationError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "SerializationFailed")
}
}
impl Error for CustomSerializationError {}
pub trait CryptEngine<ReceivingPacket, SendingPacket>: Default {
fn encrypt(&mut self, packet: SendingPacket) -> Result<Vec<u8>, CustomSerializationError>;
fn decrypt(&mut self, packet: &[u8]) -> Result<ReceivingPacket, CustomSerializationError>;
}
#[derive(Clone, Debug, Default)]
pub struct ExampleKeyPair(u64, u64);
#[derive(Clone, Debug, Default)]
pub struct CustomCryptEngine {
key_pair: ExampleKeyPair,
}
impl CustomCryptEngine {
fn xor_encrypt(&mut self, data: Vec<u8>) -> Vec<u8> {
let mut key = self.key_pair.0;
let encrypted: Vec<u8> = data
.into_iter()
.map(|byte| {
let result = byte ^ (key as u8);
key = key.wrapping_add(1);
result
})
.collect();
self.key_pair.0 = key;
encrypted
}
fn xor_decrypt(&mut self, data: Vec<u8>) -> Vec<u8> {
let mut key = self.key_pair.1;
let decrypted: Vec<u8> = data
.into_iter()
.map(|byte| {
let result = byte ^ (key as u8);
key = key.wrapping_add(1);
result
})
.collect();
self.key_pair.1 = key;
decrypted
}
}
impl CryptEngine<CustomCryptClientPacket, CustomCryptServerPacket> for CustomCryptEngine {
fn encrypt(
&mut self,
packet: CustomCryptServerPacket,
) -> Result<Vec<u8>, CustomSerializationError> {
let packet_data = bincode::serialize(&packet).unwrap();
let encrypted_data = self.xor_encrypt(packet_data);
Ok(encrypted_data)
}
fn decrypt(
&mut self,
packet: &[u8],
) -> Result<CustomCryptClientPacket, CustomSerializationError> {
let decrypted_data = self.xor_decrypt(packet.to_vec());
let packet = bincode::deserialize(&decrypted_data).unwrap();
Ok(packet)
}
}
impl CryptEngine<CustomCryptServerPacket, CustomCryptClientPacket> for CustomCryptEngine {
fn encrypt(
&mut self,
packet: CustomCryptClientPacket,
) -> Result<Vec<u8>, CustomSerializationError> {
let packet_data = bincode::serialize(&packet).unwrap();
let encrypted_data = self.xor_encrypt(packet_data);
Ok(encrypted_data)
}
fn decrypt(
&mut self,
packet: &[u8],
) -> Result<CustomCryptServerPacket, CustomSerializationError> {
let decrypted_data = self.xor_decrypt(packet.to_vec());
let packet = bincode::deserialize(&decrypted_data).unwrap();
Ok(packet)
}
}
#[derive(Clone, Default)]
pub struct CustomCryptSerializer<C, ReceivingPacket, SendingPacket>
where
C: Send + Sync + 'static + CryptEngine<ReceivingPacket, SendingPacket>,
{
crypt_engine: C,
_client: PhantomData<ReceivingPacket>,
_server: PhantomData<SendingPacket>,
}
impl<
C: Send + Sync + 'static + CryptEngine<ReceivingPacket, SendingPacket>,
SendingPacket,
ReceivingPacket,
> CustomCryptSerializer<C, ReceivingPacket, SendingPacket>
where
C: Send + Sync + 'static + CryptEngine<ReceivingPacket, SendingPacket>,
{
pub fn new(crypt_engine: C) -> Self {
Self {
crypt_engine,
_client: PhantomData,
_server: PhantomData,
}
}
}
impl<ReceivingPacket, SendingPacket, C> MutableSerializer<ReceivingPacket, SendingPacket>
for CustomCryptSerializer<C, ReceivingPacket, SendingPacket>
where
C: Send + Sync + 'static + CryptEngine<ReceivingPacket, SendingPacket>,
ReceivingPacket: Send + Sync + 'static,
SendingPacket: Send + Sync + 'static,
{
type Error = CustomSerializationError;
fn serialize(&mut self, packet: SendingPacket) -> Result<Vec<u8>, Self::Error> {
Ok(self.crypt_engine.encrypt(packet).unwrap())
}
fn deserialize(&mut self, buffer: &[u8]) -> Result<ReceivingPacket, Self::Error> {
match self.crypt_engine.decrypt(buffer) {
Ok(encrypted) => Ok(encrypted),
Err(e) => {
log::error!("{}", e);
Err(e)
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_xor_encrypt_decrypt() {
let mut engine = CustomCryptEngine::default();
let data = vec![1, 2, 3, 4, 5];
let encrypted = engine.xor_encrypt(data.clone());
assert_ne!(
encrypted, data,
"Encrypted data should not be equal to original data"
);
let decrypted = engine.xor_decrypt(encrypted);
assert_eq!(
decrypted, data,
"Decrypted data should be equal to original data"
);
}
#[test]
fn test_crypt_engine_encrypt_decrypt() {
let mut engine = CustomCryptEngine::default();
let client_packet = CustomCryptClientPacket::String("Hello, Server!".to_string());
let server_packet = CustomCryptServerPacket::String("Hello, Client!".to_string());
let encrypted = engine.encrypt(client_packet.clone()).unwrap();
let decrypted: CustomCryptClientPacket = engine.decrypt(&encrypted).unwrap();
assert_eq!(
decrypted, client_packet,
"Decrypted client packet should be equal to original packet"
);
let encrypted = engine.encrypt(server_packet.clone()).unwrap();
let decrypted: CustomCryptServerPacket = engine.decrypt(&encrypted).unwrap();
assert_eq!(
decrypted, server_packet,
"Decrypted server packet should be equal to original packet"
);
}
}