diff --git a/Cargo.lock b/Cargo.lock index 90aa616..3e74005 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4147,9 +4147,9 @@ checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] name = "tokio" -version = "1.53.0" +version = "1.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d988bcd52dbe076d3d46903332f58c912b87a2c49b1428419a5845154762ffee" +checksum = "202caea871b69668250d242070849eb495be178ed697a3e98aebce5bc81a0bed" dependencies = [ "bytes", "libc", diff --git a/prudpv0/src/server.rs b/prudpv0/src/server.rs index 96b6416..3c790cb 100644 --- a/prudpv0/src/server.rs +++ b/prudpv0/src/server.rs @@ -38,6 +38,7 @@ pub struct InternalConnection { server_packet_counter: u16, client_packet_counter: u16, unacknowledged_packets: HashMap>>, + packet_buffer: Vec, packet_queue: HashMap>)>, } pub struct Connection { @@ -96,7 +97,7 @@ impl Server { .expect("packet malformed in creation"), );*/ let mut inner = conn.inner.lock().await; - let pieces = data.chunks(700); + let pieces = data.chunks(962); let max_piece = pieces.len() - 1; let mut frag_num = 1; for (i, piece) in pieces.enumerate() { @@ -140,9 +141,18 @@ impl Server { .send_to(&data, conn.addr.regular_socket_addr) .await .ok(); - - break; + sleep(Duration::from_millis(500)).await; } + println!("connection exceeded max fail count, disconnecting"); + let Some(conn) = conn.upgrade() else { + return; + }; + let Some(this) = this.upgrade() else { + return; + }; + let mut conns = this.connections.write().await; + conns.remove(&(conn.addr, conn.session_id)); + drop(conns); }); frag_num += 1; } @@ -282,6 +292,7 @@ impl Server { server_packet_counter: 1, unacknowledged_packets: HashMap::new(), packet_queue: HashMap::new(), + packet_buffer: vec![], }), }); @@ -334,6 +345,13 @@ impl Server { warn!("data packet on inactive connection from: {:?}", addr); return; }; + + if header.type_flags.get_flags() & ACK != 0 { + let mut inner = res.inner.lock().await; + inner.unacknowledged_packets.remove(&header.sequence_id); + return; + } + info!("frag: {}", frag_id); let mut conn = res.inner.lock().await; let ack = new_data_packet( @@ -368,9 +386,16 @@ impl Server { }; conn.crypto_instance.decrypt_incoming(payload); - - res.target.send(payload.to_owned()).await; + conn.packet_buffer.extend_from_slice(payload); conn.client_packet_counter += 1; + if *packet.fragment_id().unwrap() != 0 { + info!("handeling fragmented packet"); + continue; + } + + res.target + .send(std::mem::take(&mut conn.packet_buffer)) + .await; } info!("finished handeling packets, dropping inner connection"); drop(conn); @@ -472,8 +497,8 @@ impl Server { inner.last_action = Instant::now(); drop(inner); }; - if header.type_flags.get_flags() & ACK != 0 { - info!("got ack(acks are ignored for now)"); + if header.type_flags.get_flags() & ACK != 0 && header.type_flags.get_types() != DATA { + info!("got ack(acks are ignored for now(unless they are data acks))"); return; } println!("{:?}", header); diff --git a/prudpv1/src/prudp/socket.rs b/prudpv1/src/prudp/socket.rs index 57f4ce4..0059882 100644 --- a/prudpv1/src/prudp/socket.rs +++ b/prudpv1/src/prudp/socket.rs @@ -50,6 +50,7 @@ struct InternalConnection { socket: Arc, packet_queue: HashMap, last_packet_time: Instant, + partial_packet: Vec, unacknowleged_packets: Vec<(Instant, PRUDPV1Packet)>, } @@ -431,6 +432,7 @@ impl InternalSocket { packet_queue: Default::default(), last_packet_time: Instant::now(), unacknowleged_packets: Vec::new(), + partial_packet: Vec::new(), supported_function_version, }; @@ -573,11 +575,24 @@ impl InternalSocket { while let Some(mut packet) = conn.packet_queue.remove(&counter) { conn.crypto_handler_instance .decrypt_incoming(packet.header.substream_id, &mut packet.payload[..]); - - conn.data_sender.send(packet.payload).await.ok(); - + conn.partial_packet + .extend_from_slice(&mut packet.payload[..]); conn.reliable_client_counter = conn.reliable_client_counter.overflowing_add(1).0; counter = conn.reliable_client_counter; + if packet.options.iter().any(|v| { + if let FragmentId(f) = v { + *f != 0 + } else { + false + } + }) { + println!("handeling fragmented packet"); + continue; + } + + let packet = std::mem::take(&mut conn.partial_packet); + + conn.data_sender.send(packet).await.ok(); } } @@ -690,19 +705,7 @@ impl AnyInternalSocket for InternalSocket { let conn = &**conn; let mut conn = conn.lock().await; - if conn.supported_function_version == 1 { - let mut collected_ids: Vec = Vec::new(); - let mut cursor = Cursor::new(&packet.payload); - - while let Ok(v) = read_u16(&mut cursor) { - collected_ids.push(v); - } - - conn.unacknowleged_packets.retain_mut(|(_, up)| { - !(collected_ids.iter().any(|id| up.header.sequence_id == *id) - || up.header.sequence_id <= packet.header.sequence_id) - }); - } else { + if packet.header.substream_id == 1 { let mut collected_ids: Vec = Vec::new(); let mut cursor = Cursor::new(&packet.payload); @@ -729,10 +732,22 @@ impl AnyInternalSocket for InternalSocket { collected_ids.push(additional_sequence_id); } - conn.unacknowleged_packets.retain_mut(|(_, up)| { + conn.unacknowleged_packets.retain(|(_, up)| { !(collected_ids.iter().any(|id| up.header.sequence_id == *id) || up.header.sequence_id <= sequence_id) }); + } else { + let mut collected_ids: Vec = Vec::new(); + let mut cursor = Cursor::new(&packet.payload); + + while let Ok(v) = read_u16(&mut cursor) { + collected_ids.push(v); + } + + conn.unacknowleged_packets.retain(|(_, up)| { + !(collected_ids.iter().any(|id| up.header.sequence_id == *id) + || up.header.sequence_id <= packet.header.sequence_id) + }); } } else { error!("non connection acknowledgement packet on nonexistent connection...")