finish simplistic rust client implementation
This commit is contained in:
@@ -0,0 +1,68 @@
|
||||
use std::fmt;
|
||||
use std::marker::PhantomData;
|
||||
use serde::ser::{Serialize, Serializer, SerializeTuple};
|
||||
use serde::de::{Deserialize, Deserializer, Visitor, SeqAccess, Error};
|
||||
|
||||
pub trait BigArray<'de>: Sized {
|
||||
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
|
||||
where S: Serializer;
|
||||
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
||||
where D: Deserializer<'de>;
|
||||
}
|
||||
|
||||
macro_rules! big_array {
|
||||
($($len:expr,)+) => {
|
||||
$(
|
||||
impl<'de, T> BigArray<'de> for [T; $len]
|
||||
where T: Default + Copy + Serialize + Deserialize<'de>
|
||||
{
|
||||
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
|
||||
where S: Serializer
|
||||
{
|
||||
let mut seq = serializer.serialize_tuple(self.len())?;
|
||||
for elem in &self[..] {
|
||||
seq.serialize_element(elem)?;
|
||||
}
|
||||
seq.end()
|
||||
}
|
||||
|
||||
fn deserialize<D>(deserializer: D) -> Result<[T; $len], D::Error>
|
||||
where D: Deserializer<'de>
|
||||
{
|
||||
struct ArrayVisitor<T> {
|
||||
element: PhantomData<T>,
|
||||
}
|
||||
|
||||
impl<'de, T> Visitor<'de> for ArrayVisitor<T>
|
||||
where T: Default + Copy + Deserialize<'de>
|
||||
{
|
||||
type Value = [T; $len];
|
||||
|
||||
fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
|
||||
formatter.write_str(concat!("an array of length ", $len))
|
||||
}
|
||||
|
||||
fn visit_seq<A>(self, mut seq: A) -> Result<[T; $len], A::Error>
|
||||
where A: SeqAccess<'de>
|
||||
{
|
||||
let mut arr = [T::default(); $len];
|
||||
for i in 0..$len {
|
||||
arr[i] = seq.next_element()?
|
||||
.ok_or_else(|| Error::invalid_length(i, &self))?;
|
||||
}
|
||||
Ok(arr)
|
||||
}
|
||||
}
|
||||
|
||||
let visitor = ArrayVisitor { element: PhantomData };
|
||||
deserializer.deserialize_tuple($len, visitor)
|
||||
}
|
||||
}
|
||||
)+
|
||||
}
|
||||
}
|
||||
|
||||
big_array! {
|
||||
40, 48, 50, 56, 64, 72, 96, 100, 128, 160, 192, 200, 224, 256, 384, 488, 512,
|
||||
768, 1024, 2048, 4096, 8192, 16384, 32768, 65536,
|
||||
}
|
||||
+134
-18
@@ -1,6 +1,11 @@
|
||||
use bincode::{self, Error, options, Options};
|
||||
use std::net::{UdpSocket, SocketAddr};
|
||||
use std::{net::{UdpSocket, SocketAddr}, io::Write};
|
||||
use serde::{Serialize, Deserialize};
|
||||
use sha3::{Digest, Sha3_512};
|
||||
use sha2::Sha512;
|
||||
|
||||
mod big_array;
|
||||
use big_array::BigArray;
|
||||
|
||||
const MAX_FRAME_PAYLOAD:u16=508;
|
||||
const MAX_FRAME_PAYLOAD_U:usize=MAX_FRAME_PAYLOAD as usize;
|
||||
@@ -15,10 +20,12 @@ struct PacketInformation {
|
||||
response_filename_checksum: u64, //8 bytes
|
||||
} //14 bytes
|
||||
|
||||
struct packet {
|
||||
#[derive(Serialize, Deserialize)]
|
||||
struct Packet {
|
||||
packet_number: u32, //4 bytes
|
||||
payload_hash: u128, //16 bytes
|
||||
payload: [u8; MAX_PAYLOAD_U], //508 bytes
|
||||
payload_hash: [u8; 16], //16 bytes
|
||||
#[serde(with = "BigArray")]
|
||||
payload: [u8; MAX_PAYLOAD_U], //488 bytes
|
||||
} //512 bytes
|
||||
|
||||
fn main() {
|
||||
@@ -28,8 +35,8 @@ fn main() {
|
||||
|
||||
let filename = "data.bin";
|
||||
|
||||
let local_addr: SocketAddr = "0.0.0.0:26000".parse().expect("Failed to parse address");
|
||||
let socket = UdpSocket::bind(local_addr).expect("Failed to bind socket");
|
||||
let local_addr: SocketAddr = "0.0.0.0:0".parse().expect("Failed to parse address");
|
||||
let socket = UdpSocket::bind(local_addr).expect("Failed to bind socket to port 26000");
|
||||
|
||||
socket.set_read_timeout(Some(std::time::Duration::new(timeout, 0))).expect("set_read_timeout call failed");
|
||||
//socket.set_nonblocking(true).expect("set_nonblocking call failed");
|
||||
@@ -74,28 +81,137 @@ fn main() {
|
||||
}
|
||||
|
||||
//create vector to store the packets
|
||||
let mut packets: Vec<u8> = Vec::new();
|
||||
let mut packets: Vec<Vec<u8>> = Vec::new();
|
||||
for _ in 0..packet_info.packet_numbers {
|
||||
packets.push(0);
|
||||
packets.push(Vec::new());
|
||||
}
|
||||
|
||||
let received_packets = 0;
|
||||
let mut received_packets = 0;
|
||||
|
||||
let mut server_hash = [0u8; 64];
|
||||
let mut server_hash_received = false;
|
||||
|
||||
//receive the packets
|
||||
while received_packets < packet_info.packet_numbers {
|
||||
let mut buffer = [0u8; MAX_PAYLOAD_U];
|
||||
let (received_bytes, remote_addr) = socket.recv_from(&mut buffer).expect("Failed to receive data");
|
||||
while received_packets < packet_info.packet_numbers-1 || !server_hash_received {
|
||||
let mut buffer = [0u8; MAX_FRAME_PAYLOAD_U];
|
||||
let res = socket.recv_from(&mut buffer);
|
||||
if let Ok((received_bytes, remote_addr)) = res {
|
||||
let filled_buffer = &buffer;//[..received_bytes];
|
||||
|
||||
let filled_buffer = &buffer[..received_bytes];
|
||||
//print the filled buffer
|
||||
//println!("Received data: {:?} {}", filled_buffer, filled_buffer.len());
|
||||
|
||||
//print the filled buffer
|
||||
println!("Received data: {:?} {}", filled_buffer, filled_buffer.len());
|
||||
if remote_addr != server_addr {
|
||||
panic!("Received data from unknown address");
|
||||
}
|
||||
|
||||
if remote_addr != server_addr {
|
||||
panic!("Received data from unknown address");
|
||||
if received_bytes == 65 {
|
||||
//println!("Received hash packet, ignoring for now");
|
||||
server_hash = filled_buffer[0..64].try_into().expect("Failed to convert hash");
|
||||
//println!("Received hash: {:?}", server_hash);
|
||||
server_hash_received = true;
|
||||
continue;
|
||||
}
|
||||
|
||||
if received_bytes != MAX_FRAME_PAYLOAD_U {
|
||||
println!("Received packet with invalid size {} not {} | ignoring", received_bytes, MAX_FRAME_PAYLOAD_U);
|
||||
continue;
|
||||
}
|
||||
|
||||
let packet_result: Result<Packet, Error> = options.deserialize(filled_buffer);
|
||||
let packet: Packet;
|
||||
match packet_result {
|
||||
Ok(p) => {
|
||||
//println!("Packet {}", p.packet_number);
|
||||
//check checksum with sum
|
||||
if p.packet_number != packet_info.packet_numbers - 1 {
|
||||
let mut hasher = Sha3_512::new();
|
||||
hasher.update(&p.payload);
|
||||
let hash = hasher.finalize();
|
||||
|
||||
if p.payload_hash != hash[..16] {
|
||||
continue;
|
||||
}
|
||||
|
||||
packet = p;
|
||||
} else {
|
||||
//cut packet down to size
|
||||
packets[p.packet_number as usize] = p.payload[..packet_info.last_packet_size as usize].to_vec();
|
||||
continue;
|
||||
}
|
||||
}
|
||||
Err(err) => {
|
||||
panic!("Failed to deserialize data: {}", err);
|
||||
}
|
||||
}
|
||||
|
||||
if packets[packet.packet_number as usize].len() != 0 {
|
||||
//println!("Packet already received, ignoring");
|
||||
continue;
|
||||
}
|
||||
|
||||
packets[packet.packet_number as usize] = packet.payload.to_vec();
|
||||
received_packets += 1;
|
||||
} else {
|
||||
//println!("Timeout, requesting again {}/{}\r", received_packets, packet_info.packet_numbers);
|
||||
//collect packets that were not received and send them in n messages
|
||||
//where n is the minimum amount of messages needed to request all packets
|
||||
let mut missing_packets: Vec<u32> = Vec::new();
|
||||
for i in 0..packet_info.packet_numbers {
|
||||
if packets[i as usize].len() == 0 {
|
||||
missing_packets.push(i);
|
||||
}
|
||||
}
|
||||
|
||||
//split lost_packets into groups of size 508-filename.len() bytes
|
||||
let mut missing_packet_groups: Vec<String> = Vec::new();
|
||||
let mut current_group: String = filename.to_string();
|
||||
for i in 0..missing_packets.len() {
|
||||
if current_group.len() + missing_packets[i].to_string().len() + 1 > MAX_PAYLOAD_U {
|
||||
missing_packet_groups.push(current_group);
|
||||
current_group = filename.to_string();
|
||||
}
|
||||
current_group.push('/');
|
||||
current_group.push_str(&missing_packets[i].to_string());
|
||||
}
|
||||
|
||||
if current_group.len() > filename.len() {
|
||||
missing_packet_groups.push(current_group);
|
||||
}
|
||||
|
||||
for i in 0..missing_packet_groups.len() {
|
||||
let message = &missing_packet_groups[i];
|
||||
//println!("Requesting packets: {}", message);
|
||||
socket.send_to(message.as_bytes(), server_addr).expect("Failed to send data");
|
||||
}
|
||||
|
||||
if !server_hash_received {
|
||||
let message = filename.to_string()+":";
|
||||
socket.send_to(message.as_bytes(), server_addr).expect("Failed to send data");
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
print!("Packet {}/{}\r", received_packets, packet_info.packet_numbers);
|
||||
}
|
||||
|
||||
//check hash via sha512
|
||||
|
||||
let mut hasher = Sha512::new();
|
||||
for i in 0..packets.len() {
|
||||
hasher.update(&packets[i]);
|
||||
}
|
||||
let client_hash = hasher.finalize();
|
||||
|
||||
if client_hash[..64] != server_hash {
|
||||
panic!("Hashes do not match, correct hash: {:?}, client hash: {:?}", server_hash, client_hash);
|
||||
}
|
||||
|
||||
|
||||
println!("Received all packets, writing to file");
|
||||
//write packets to file
|
||||
let mut file = std::fs::File::create("received/".to_string()+filename).expect("Failed to create file");
|
||||
for i in 0..packets.len() {
|
||||
file.write_all(&packets[i]).expect("Failed to write to file");
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user