diff --git a/RUDPSharp.Tests/ChannelTests.cs b/RUDPSharp.Tests/ChannelTests.cs index 81caa42..90575ed 100644 --- a/RUDPSharp.Tests/ChannelTests.cs +++ b/RUDPSharp.Tests/ChannelTests.cs @@ -236,5 +236,263 @@ public void TestReliableInOrderChannelHandlesSequenceWrapping () Assert.AreEqual (expectedSequence, packet.Sequence, $"Packet sequence should be {expectedSequence} but was {packet.Sequence}"); } } + + [Test] + public void TestFragmentationBoundaryNonFragmented() + { + // Test packet just below fragmentation threshold (1016 bytes) + // Should NOT be fragmented (threshold is 1017 bytes) + var channel = new UnreliableChannel(); + byte[] payload = new byte[1016]; + for (int i = 0; i < payload.Length; i++) + { + payload[i] = (byte)(i % 256); + } + + var packet = new Packet(PacketType.Data, Channel.None, 1, payload); + var pendingPacket = channel.QueueOutgoingPacket(EndPoint, packet); + + Assert.IsNotNull(pendingPacket); + + // Should get exactly one packet (not fragmented) + var outgoingPackets = channel.GetPendingOutgoingPackets().ToArray(); + Assert.AreEqual(1, outgoingPackets.Length, "Should have exactly 1 packet (not fragmented)"); + } + + [Test] + public void TestFragmentationBoundaryExact() + { + // Test packet at exact fragmentation threshold (1017 bytes) + // Should NOT be fragmented (threshold is >1017, not >=1017) + var channel = new UnreliableChannel(); + byte[] payload = new byte[1017]; + for (int i = 0; i < payload.Length; i++) + { + payload[i] = (byte)(i % 256); + } + + var packet = new Packet(PacketType.Data, Channel.None, 1, payload); + var pendingPacket = channel.QueueOutgoingPacket(EndPoint, packet); + + Assert.IsNotNull(pendingPacket); + + // Should get exactly one packet (not fragmented) + var outgoingPackets = channel.GetPendingOutgoingPackets().ToArray(); + Assert.AreEqual(1, outgoingPackets.Length, "Should have exactly 1 packet at boundary"); + } + + [Test] + public void TestFragmentationBoundaryJustAbove() + { + // Test packet just above fragmentation threshold (1018 bytes) + // Should be fragmented into 2 packets + var channel = new UnreliableChannel(); + byte[] payload = new byte[1018]; + for (int i = 0; i < payload.Length; i++) + { + payload[i] = (byte)(i % 256); + } + + var packet = new Packet(PacketType.Data, Channel.None, 1, payload); + var pendingPacket = channel.QueueOutgoingPacket(EndPoint, packet); + + Assert.IsNotNull(pendingPacket); + + // Should get exactly two fragments + var outgoingPackets = channel.GetPendingOutgoingPackets().ToArray(); + Assert.AreEqual(2, outgoingPackets.Length, "Should have 2 fragments for 1018 bytes"); + } + + [Test] + public void TestFragmentationReassemblyBoundary() + { + // Test that fragments at boundary are correctly reassembled + var channel = new UnreliableChannel(); + byte[] payload = new byte[1018]; + for (int i = 0; i < payload.Length; i++) + { + payload[i] = (byte)(i % 256); + } + + var outPacket = new Packet(PacketType.Data, Channel.None, 1, payload); + channel.QueueOutgoingPacket(EndPoint, outPacket); + + var outgoingPackets = channel.GetPendingOutgoingPackets().ToArray(); + Assert.AreEqual(2, outgoingPackets.Length); + + // Simulate receiving the fragments + var receiveChannel = new UnreliableChannel(); + foreach (var fragment in outgoingPackets) + { + var receivedPacket = new Packet(fragment.Data, fragment.Data.Length); + receiveChannel.QueueIncomingPacket(EndPoint, receivedPacket); + } + + // Should get one reassembled packet + var incomingPackets = receiveChannel.GetPendingIncomingPackets().ToArray(); + Assert.AreEqual(1, incomingPackets.Length, "Should have 1 reassembled packet"); + Assert.AreEqual(1018, incomingPackets[0].Data.Length, "Reassembled packet should be 1018 bytes"); + + // Verify data integrity + for (int i = 0; i < payload.Length; i++) + { + Assert.AreEqual((byte)(i % 256), incomingPackets[0].Data[i], $"Data mismatch at index {i}"); + } + } + + [Test] + public void TestFragmentationMultipleFragments() + { + // Test packet that requires multiple fragments (2048 bytes = 3 fragments) + var channel = new UnreliableChannel(); + byte[] payload = new byte[2048]; + var rnd = new Random(12345); + rnd.NextBytes(payload); + + var packet = new Packet(PacketType.Data, Channel.None, 1, payload); + channel.QueueOutgoingPacket(EndPoint, packet); + + var outgoingPackets = channel.GetPendingOutgoingPackets().ToArray(); + // 2048 bytes / 1017 bytes per fragment = 3 fragments (1017 + 1017 + 14) + Assert.AreEqual(3, outgoingPackets.Length, "Should have 3 fragments for 2048 bytes"); + + // Simulate receiving the fragments in order + var receiveChannel = new UnreliableChannel(); + foreach (var fragment in outgoingPackets) + { + var receivedPacket = new Packet(fragment.Data, fragment.Data.Length); + receiveChannel.QueueIncomingPacket(EndPoint, receivedPacket); + } + + var incomingPackets = receiveChannel.GetPendingIncomingPackets().ToArray(); + Assert.AreEqual(1, incomingPackets.Length, "Should have 1 reassembled packet"); + Assert.AreEqual(2048, incomingPackets[0].Data.Length, "Reassembled packet should be 2048 bytes"); + Assert.AreEqual(payload, incomingPackets[0].Data, "Reassembled data should match original"); + } + + [Test] + public void TestFragmentationOutOfOrderReassembly() + { + // Test that fragments received out of order are correctly reassembled + var channel = new UnreliableChannel(); + byte[] payload = new byte[3000]; + var rnd = new Random(54321); + rnd.NextBytes(payload); + + var packet = new Packet(PacketType.Data, Channel.None, 1, payload); + channel.QueueOutgoingPacket(EndPoint, packet); + + var outgoingPackets = channel.GetPendingOutgoingPackets().ToArray(); + Assert.AreEqual(3, outgoingPackets.Length, "Should have 3 fragments for 3000 bytes"); + + // Simulate receiving the fragments OUT OF ORDER (2, 0, 1) + var receiveChannel = new UnreliableChannel(); + var receivedPacket2 = new Packet(outgoingPackets[2].Data, outgoingPackets[2].Data.Length); + receiveChannel.QueueIncomingPacket(EndPoint, receivedPacket2); + + var receivedPacket0 = new Packet(outgoingPackets[0].Data, outgoingPackets[0].Data.Length); + receiveChannel.QueueIncomingPacket(EndPoint, receivedPacket0); + + var receivedPacket1 = new Packet(outgoingPackets[1].Data, outgoingPackets[1].Data.Length); + receiveChannel.QueueIncomingPacket(EndPoint, receivedPacket1); + + var incomingPackets = receiveChannel.GetPendingIncomingPackets().ToArray(); + Assert.AreEqual(1, incomingPackets.Length, "Should have 1 reassembled packet"); + Assert.AreEqual(3000, incomingPackets[0].Data.Length, "Reassembled packet should be 3000 bytes"); + Assert.AreEqual(payload, incomingPackets[0].Data, "Reassembled data should match original despite out-of-order receipt"); + } + + [Test] + public void TestFragmentationLargePacket() + { + // Test large packet (4096 bytes = 5 fragments) + var channel = new UnreliableChannel(); + byte[] payload = new byte[4096]; + for (int i = 0; i < payload.Length; i++) + { + payload[i] = (byte)(i % 256); + } + + var packet = new Packet(PacketType.Data, Channel.None, 1, payload); + channel.QueueOutgoingPacket(EndPoint, packet); + + var outgoingPackets = channel.GetPendingOutgoingPackets().ToArray(); + // 4096 / 1017 = ~5 fragments + Assert.AreEqual(5, outgoingPackets.Length, "Should have 5 fragments for 4096 bytes"); + + // Verify reassembly + var receiveChannel = new UnreliableChannel(); + foreach (var fragment in outgoingPackets) + { + var receivedPacket = new Packet(fragment.Data, fragment.Data.Length); + receiveChannel.QueueIncomingPacket(EndPoint, receivedPacket); + } + + var incomingPackets = receiveChannel.GetPendingIncomingPackets().ToArray(); + Assert.AreEqual(1, incomingPackets.Length); + Assert.AreEqual(4096, incomingPackets[0].Data.Length); + + for (int i = 0; i < payload.Length; i++) + { + Assert.AreEqual((byte)(i % 256), incomingPackets[0].Data[i], $"Data mismatch at index {i}"); + } + } + + [Test] + public void TestFragmentationMaxFragmentCount() + { + // Test approaching maximum fragment count (255 fragments) + // 255 * 1017 = 259335 bytes + var channel = new UnreliableChannel(); + int payloadSize = 255 * 1017; // Exactly 255 fragments + byte[] payload = new byte[payloadSize]; + + // Fill with pattern for verification + for (int i = 0; i < payload.Length; i++) + { + payload[i] = (byte)(i % 256); + } + + var packet = new Packet(PacketType.Data, Channel.None, 1, payload); + var pendingPacket = channel.QueueOutgoingPacket(EndPoint, packet); + + Assert.IsNotNull(pendingPacket); + var outgoingPackets = channel.GetPendingOutgoingPackets().ToArray(); + Assert.AreEqual(255, outgoingPackets.Length, "Should have exactly 255 fragments"); + + // Verify reassembly + var receiveChannel = new UnreliableChannel(); + foreach (var fragment in outgoingPackets) + { + var receivedPacket = new Packet(fragment.Data, fragment.Data.Length); + receiveChannel.QueueIncomingPacket(EndPoint, receivedPacket); + } + + var incomingPackets = receiveChannel.GetPendingIncomingPackets().ToArray(); + Assert.AreEqual(1, incomingPackets.Length); + Assert.AreEqual(payloadSize, incomingPackets[0].Data.Length); + } + + [Test] + public void TestFragmentationExceedsMaxFragmentCount() + { + // Test that exceeding maximum fragment count throws exception + var channel = new UnreliableChannel(); + int payloadSize = 256 * 1017; // Would require 256 fragments (exceeds 255 max) + byte[] payload = new byte[payloadSize]; + + bool exceptionThrown = false; + try + { + var packet = new Packet(PacketType.Data, Channel.None, 1, payload); + channel.QueueOutgoingPacket(EndPoint, packet); + } + catch (InvalidOperationException) + { + exceptionThrown = true; + } + + Assert.IsTrue(exceptionThrown, "Should throw InvalidOperationException when exceeding max fragment count"); + } } } \ No newline at end of file diff --git a/RUDPSharp.Tests/RUDPTests.cs b/RUDPSharp.Tests/RUDPTests.cs index c622d07..24ec968 100644 --- a/RUDPSharp.Tests/RUDPTests.cs +++ b/RUDPSharp.Tests/RUDPTests.cs @@ -103,6 +103,7 @@ public void TestClientCanConnectAndDisconnect () rUDPClient.Disconnect (); serverWait.WaitOne (1000); + Thread.Sleep (100); // Give Poll() loop time to remove remote after disconnect event Assert.AreEqual (0, rUDPServer.Remotes.Count); wait.WaitOne (1000); Assert.AreEqual (0, rUDPClient.Remotes.Count); @@ -261,10 +262,10 @@ public void TestLargePacketIsDelivered () rnd.NextBytes (message); dataReceived = null; remote = null; - serverWait.Reset (); + wait.Reset (); // Reset wait since client will receive the data Assert.IsTrue (rUDPServer.SendToAll (Channel.None, message)); - serverWait.WaitOne (5000); + wait.WaitOne (5000); // Wait for client to receive data Assert.AreEqual (message, dataReceived, $"({(string.Join (",", dataReceived ?? Array.Empty ()))}) != ({(string.Join (",", message))})"); Assert.AreEqual (rUDPServer.EndPoint, remote); diff --git a/RUDPSharp/FragmentAssembler.cs b/RUDPSharp/FragmentAssembler.cs new file mode 100644 index 0000000..1b0437b --- /dev/null +++ b/RUDPSharp/FragmentAssembler.cs @@ -0,0 +1,111 @@ +using System; +using System.Collections.Generic; +using System.Linq; + +namespace RUDPSharp +{ + internal class FragmentAssembler + { + private class FragmentCollection + { + public byte[][] Fragments { get; set; } + public byte TotalFragments { get; set; } + public int ReceivedCount { get; set; } + public long LastUpdate { get; set; } + public PacketType PacketType { get; set; } + public Channel Channel { get; set; } + public ushort Sequence { get; set; } + } + + private Dictionary fragmentBuffers = new Dictionary(); + private static readonly long FRAGMENT_TIMEOUT_TICKS = TimeSpan.FromSeconds(5).Ticks; + + public void AddFragment(Packet packet) + { + if (!packet.Fragmented) + { + throw new ArgumentException("Packet is not fragmented", nameof(packet)); + } + + ushort fragmentId = packet.FragmentId; + + if (!fragmentBuffers.ContainsKey(fragmentId)) + { + fragmentBuffers[fragmentId] = new FragmentCollection + { + Fragments = new byte[packet.TotalFragments][], + TotalFragments = packet.TotalFragments, + ReceivedCount = 0, + LastUpdate = DateTime.Now.Ticks, + PacketType = packet.PacketType, + Channel = packet.Channel, + Sequence = packet.Sequence + }; + } + + var collection = fragmentBuffers[fragmentId]; + + if (collection.Fragments[packet.FragmentIndex] == null) + { + collection.Fragments[packet.FragmentIndex] = packet.Payload.ToArray(); + collection.ReceivedCount++; + collection.LastUpdate = DateTime.Now.Ticks; + } + } + + public bool TryGetCompleteMessage(out byte[] data, out PacketType packetType, out Channel channel, out ushort sequence) + { + data = null; + packetType = PacketType.Data; + channel = Channel.None; + sequence = 0; + + // Find a complete fragment collection + foreach (var kvp in fragmentBuffers) + { + var collection = kvp.Value; + if (collection.ReceivedCount == collection.TotalFragments) + { + // All fragments received, reassemble + int totalLength = collection.Fragments.Sum(f => f.Length); + data = new byte[totalLength]; + int offset = 0; + + for (int i = 0; i < collection.TotalFragments; i++) + { + Buffer.BlockCopy(collection.Fragments[i], 0, data, offset, collection.Fragments[i].Length); + offset += collection.Fragments[i].Length; + } + + packetType = collection.PacketType; + channel = collection.Channel; + sequence = collection.Sequence; + + fragmentBuffers.Remove(kvp.Key); + return true; + } + } + + return false; + } + + public void CleanupOldFragments() + { + long now = DateTime.Now.Ticks; + var toRemove = new List(); + + foreach (var kvp in fragmentBuffers) + { + if (now - kvp.Value.LastUpdate > FRAGMENT_TIMEOUT_TICKS) + { + toRemove.Add(kvp.Key); + } + } + + foreach (var key in toRemove) + { + fragmentBuffers.Remove(key); + } + } + } +} diff --git a/RUDPSharp/InOrderChannel.cs b/RUDPSharp/InOrderChannel.cs index 558f86a..e9c120b 100644 --- a/RUDPSharp/InOrderChannel.cs +++ b/RUDPSharp/InOrderChannel.cs @@ -55,6 +55,11 @@ public override PendingPacket QueueOutgoingPacket (EndPoint endPoint, Packet pac } public override PendingPacket QueueIncomingPacket (EndPoint endPoint, Packet packet) { + // Handle fragmentation first before sequence checking + if (packet.Fragmented) { + return base.QueueIncomingPacket (endPoint, packet); + } + if (QueueOrDiscardPendingPackages (endPoint, PendingPacket.FromPacket (endPoint, packet))) { return base.QueueIncomingPacket (endPoint, packet); } diff --git a/RUDPSharp/Packet.cs b/RUDPSharp/Packet.cs index e130d79..5454098 100644 --- a/RUDPSharp/Packet.cs +++ b/RUDPSharp/Packet.cs @@ -5,15 +5,25 @@ namespace RUDPSharp /* * Packet Structure * - * byte header PacketType|Channel|FragmetedBit - * byte sequence - * byte sequence + * byte header PacketType|Channel|FragmentedBit + * byte sequence (low) + * byte sequence (high) + * -- If fragmented: + * byte fragmentId (low) + * byte fragmentId (high) + * byte fragmentIndex + * byte totalFragments + * -- End if * byte+ payload */ public ref struct Packet { const int HEADER_OFFSET = 0; const int SEQUENCE_OFFSET = 1; const int PAYLOAD_OFFSET = 3; + const int FRAGMENT_ID_OFFSET = 3; + const int FRAGMENT_INDEX_OFFSET = 5; + const int FRAGMENT_TOTAL_OFFSET = 6; + const int FRAGMENT_PAYLOAD_OFFSET = 7; byte [] rawData; Span span; @@ -26,19 +36,33 @@ public Packet (byte[] data, int length) PacketType = header.type; Channel = header.channel; Fragmented = header.fragmented; + + if (Fragmented) { + FragmentId = BitConverter.ToUInt16(rawData, FRAGMENT_ID_OFFSET); + FragmentIndex = rawData[FRAGMENT_INDEX_OFFSET]; + TotalFragments = rawData[FRAGMENT_TOTAL_OFFSET]; + } else { + FragmentId = 0; + FragmentIndex = 0; + TotalFragments = 1; + } } public Packet (PacketType packetType, Channel channel, ReadOnlySpan payload, bool fragmented = false) { - byte header = EncodeHeader (packetType, channel); - rawData = new byte[payload.Length + PAYLOAD_OFFSET]; - rawData[HEADER_OFFSET] = header; + byte header = EncodeHeader (packetType, channel, fragmented); + int offset = fragmented ? FRAGMENT_PAYLOAD_OFFSET : PAYLOAD_OFFSET; + rawData = new byte[payload.Length + offset]; span = new Span (rawData); - payload.TryCopyTo (span.Slice (PAYLOAD_OFFSET)); + rawData[HEADER_OFFSET] = header; PacketType = packetType; Channel = channel; Fragmented = fragmented; + FragmentId = 0; + FragmentIndex = 0; + TotalFragments = 1; Sequence = 0; + payload.TryCopyTo (span.Slice (offset)); } public Packet (PacketType packetType, Channel channel, ushort sequence, ReadOnlySpan payload, bool fragmented = false) @@ -46,6 +70,27 @@ public Packet (PacketType packetType, Channel channel, ushort sequence, ReadOnly { Sequence = sequence; } + + public void SetFragmentInfo(ushort fragmentId, byte fragmentIndex, byte totalFragments) + { + if (!Fragmented) { + throw new InvalidOperationException("Cannot set fragment info on non-fragmented packet"); + } + FragmentId = fragmentId; + FragmentIndex = fragmentIndex; + TotalFragments = totalFragments; + // Use BitConverter to properly handle ushort + byte[] fragmentIdBytes = BitConverter.GetBytes(fragmentId); + if (BitConverter.IsLittleEndian) { + rawData[FRAGMENT_ID_OFFSET] = fragmentIdBytes[0]; + rawData[FRAGMENT_ID_OFFSET + 1] = fragmentIdBytes[1]; + } else { + rawData[FRAGMENT_ID_OFFSET] = fragmentIdBytes[1]; + rawData[FRAGMENT_ID_OFFSET + 1] = fragmentIdBytes[0]; + } + rawData[FRAGMENT_INDEX_OFFSET] = fragmentIndex; + rawData[FRAGMENT_TOTAL_OFFSET] = totalFragments; + } public byte Header { get { return span[HEADER_OFFSET]; }} @@ -54,13 +99,24 @@ public Packet (PacketType packetType, Channel channel, ushort sequence, ReadOnly public Channel Channel { get; private set; } public bool Fragmented { get; private set; } + + public ushort FragmentId { get; private set; } + + public byte FragmentIndex { get; private set; } + + public byte TotalFragments { get; private set; } public ushort Sequence { get { return DecodeSequence (); } set { EncodeSequence (value);} } - public ReadOnlySpan Payload { get { return span.Slice (PAYLOAD_OFFSET); }} + public ReadOnlySpan Payload { + get { + int offset = Fragmented ? FRAGMENT_PAYLOAD_OFFSET : PAYLOAD_OFFSET; + return span.Slice (offset); + } + } public byte[] Data { get { return rawData; }} @@ -75,7 +131,7 @@ static byte EncodeHeader(PacketType type, Channel channel, bool fragmented = fal byte t = (byte)type; byte t1 = (byte)channel; byte header = (byte)(t1 << 5); - byte f = (byte)(fragmented ? 1 : 0 << 7); + byte f = (byte)(fragmented ? 1 << 7 : 0); header = (byte)(f | header | t); return header; } diff --git a/RUDPSharp/ReliableChannel.cs b/RUDPSharp/ReliableChannel.cs index 13aecda..3117d73 100644 --- a/RUDPSharp/ReliableChannel.cs +++ b/RUDPSharp/ReliableChannel.cs @@ -19,6 +19,13 @@ public override PendingPacket QueueOutgoingPacket(EndPoint endPoint, Packet pack } public override PendingPacket QueueIncomingPacket (EndPoint endPoint, Packet packet) { + // Don't process ACKs or handle acknowledgment for fragments + // Fragments will be handled in the base class and acknowledgment will happen + // after reassembly + if (packet.Fragmented) { + return base.QueueIncomingPacket (endPoint, packet); + } + if (!acknowledgement.HandleIncommingPacket(packet)) { QueueOutgoingPacket(endPoint, new Packet(PacketType.Ack, packet.Channel, packet.Sequence, Array.Empty ())); diff --git a/RUDPSharp/UDPSocket.cs b/RUDPSharp/UDPSocket.cs index 5d8e044..9cf2d74 100644 --- a/RUDPSharp/UDPSocket.cs +++ b/RUDPSharp/UDPSocket.cs @@ -11,7 +11,7 @@ namespace RUDPSharp public class UDPSocket : IDisposable { Socket socketIP4; Socket socketIP6; - const int BufferSize = 1024; + const int BufferSize = 8192; const int SioUdpConnreset = -1744830452; //SIO_UDP_CONNRESET = IOC_IN | IOC_VENDOR | 12 const int SocketTTL = 255; string name; diff --git a/RUDPSharp/UnreliableChannel.cs b/RUDPSharp/UnreliableChannel.cs index d7eb1cc..8c99b8d 100644 --- a/RUDPSharp/UnreliableChannel.cs +++ b/RUDPSharp/UnreliableChannel.cs @@ -10,11 +10,13 @@ public class UnreliableChannel { ConcurrentQueue outgoing; ConcurrentQueue incoming; + FragmentAssembler fragmentAssembler = new FragmentAssembler(); protected ConcurrentQueue Outgoing => outgoing; protected ConcurrentQueue Incoming => incoming; protected int MaxBufferSize = 1024; + private ushort nextFragmentId = 0; public UnreliableChannel (int maxBufferSize = 1024) { @@ -39,19 +41,70 @@ public virtual bool TryGetNextIncomingPacket (out PendingPacket packet) } public virtual PendingPacket QueueOutgoingPacket (EndPoint endPoint, Packet packet) { - var pendingPacket = PendingPacket.FromPacket (endPoint, packet, incoming: false); - outgoing.Enqueue (pendingPacket); - return pendingPacket; + // Check if packet needs fragmentation (leaving room for headers) + // Fragment overhead = header(1) + sequence(2) + fragmentId(2) + fragmentIndex(1) + totalFragments(1) = 7 bytes + const int NON_FRAGMENT_OVERHEAD = 3; // header(1) + sequence(2) + const int FRAGMENT_METADATA_SIZE = 4; // fragmentId(2) + fragmentIndex(1) + totalFragments(1) + int maxPayloadSize = MaxBufferSize - NON_FRAGMENT_OVERHEAD - FRAGMENT_METADATA_SIZE; + + if (packet.Payload.Length > maxPayloadSize && !packet.Fragmented) + { + // Need to fragment this packet + ushort fragmentId = nextFragmentId++; + byte[] payload = packet.Payload.ToArray(); + int totalFragments = (payload.Length + maxPayloadSize - 1) / maxPayloadSize; // ceiling division + + if (totalFragments > 255) + { + throw new InvalidOperationException($"Payload too large: {payload.Length} bytes would require {totalFragments} fragments (max 255)"); + } + + PendingPacket lastPendingPacket = null; + for (int i = 0; i < totalFragments; i++) + { + int offset = i * maxPayloadSize; + int length = Math.Min(maxPayloadSize, payload.Length - offset); + ReadOnlySpan fragmentPayload = new ReadOnlySpan(payload, offset, length); + + var fragmentPacket = new Packet(packet.PacketType, packet.Channel, packet.Sequence, fragmentPayload, fragmented: true); + fragmentPacket.SetFragmentInfo(fragmentId, (byte)i, (byte)totalFragments); + + lastPendingPacket = PendingPacket.FromPacket(endPoint, fragmentPacket, incoming: false); + outgoing.Enqueue(lastPendingPacket); + } + + return lastPendingPacket; + } + else + { + var pendingPacket = PendingPacket.FromPacket (endPoint, packet, incoming: false); + outgoing.Enqueue (pendingPacket); + return pendingPacket; + } } public virtual PendingPacket QueueIncomingPacket (EndPoint endPoint, Packet packet) { if (packet.Fragmented) { - //TODO put the fragment in a bucket for processing when we get the rest of it. + // Add fragment to assembler + fragmentAssembler.AddFragment(packet); + + // Check if we have a complete message + if (fragmentAssembler.TryGetCompleteMessage(out byte[] data, out PacketType packetType, out Channel channel, out ushort sequence)) + { + // Create a non-fragmented packet with the reassembled data + var reassembledPacket = new Packet(packetType, channel, sequence, data, fragmented: false); + var pendingPacket = PendingPacket.FromPacket(endPoint, reassembledPacket); + incoming.Enqueue(pendingPacket); + return pendingPacket; + } + + // Fragment stored but message not complete yet + return null; } - var pendingPacket = PendingPacket.FromPacket (endPoint, packet); - incoming.Enqueue (pendingPacket); - return pendingPacket; + var pending = PendingPacket.FromPacket (endPoint, packet); + incoming.Enqueue (pending); + return pending; } public virtual IEnumerable GetPendingOutgoingPackets () @@ -63,6 +116,9 @@ public virtual IEnumerable GetPendingOutgoingPackets () public virtual IEnumerable GetPendingIncomingPackets () { + // Cleanup old fragments periodically + fragmentAssembler.CleanupOldFragments(); + while (TryGetNextIncomingPacket (out PendingPacket pendingPacket)) { yield return pendingPacket; }