commit b59b0e98dc9f339088acb5fc966143be8a2d9ce1
parent 18ba6e50c408176e86b5dea2e0278ecc86d6dc11
Author: Mikolaj Lenczewski <mikolaj.lenczewski308@gmail.com>
Date: Fri, 10 Apr 2020 11:00:54 +0100
Redoing SocketAsyncEventArgs architecture to use events, not async/await
Diffstat:
4 files changed, 427 insertions(+), 41 deletions(-)
diff --git a/NetSharp/NetSharp/Sockets/Datagram/DatagramSocketServer.cs b/NetSharp/NetSharp/Sockets/Datagram/DatagramSocketServer.cs
@@ -1,21 +1,28 @@
-using NetSharp.Packets;
-using NetSharp.Utils;
+using NetSharp.Utils;
using System;
using System.Collections.Concurrent;
+using System.Collections.Generic;
using System.Net;
using System.Net.Sockets;
using System.Text;
using System.Threading;
using System.Threading.Channels;
using System.Threading.Tasks;
+using Microsoft.Extensions.ObjectPool;
+using NetSharp.Deprecated;
+using NetworkPacket = NetSharp.Packets.NetworkPacket;
namespace NetSharp.Sockets.Datagram
{
public class DatagramSocketServer : SocketServer
{
+ private static readonly EndPoint AnyRemoteEndPoint = new IPEndPoint(IPAddress.Any, 0);
+
private readonly ConcurrentDictionary<EndPoint, RemoteDatagramClientToken> connectedClientTokens;
+ private readonly ObjectPool<SocketAsyncEventArgs> ArgsPool;
+
private readonly struct RemoteDatagramClientToken
{
private readonly Channel<NetworkPacket> PacketChannel;
@@ -36,6 +43,34 @@ namespace NetSharp.Sockets.Datagram
: base(in connectionAddressFamily, SocketType.Dgram, in connectionProtocolType)
{
connectedClientTokens = new ConcurrentDictionary<EndPoint, RemoteDatagramClientToken>();
+
+ SocketAsyncEventArgs CreateArgs()
+ {
+ SocketAsyncEventArgs args = new SocketAsyncEventArgs();
+
+ args.Completed += HandleIoCompleted;
+
+ return args;
+ }
+
+ static void ResetArgs(SocketAsyncEventArgs args)
+ {
+
+ }
+
+ void DestroyArgs(SocketAsyncEventArgs args)
+ {
+ args.Completed -= HandleIoCompleted;
+
+ args.Dispose();
+ }
+
+ static bool ReBufferArgsPredicate(in SocketAsyncEventArgs args)
+ {
+ return true;
+ }
+
+ ArgsPool = new ObjectPool<SocketAsyncEventArgs>(CreateArgs, ResetArgs, DestroyArgs, ReBufferArgsPredicate);
}
protected override SocketAsyncEventArgs GenerateConnectionArgs(EndPoint remoteEndPoint)
@@ -115,15 +150,13 @@ namespace NetSharp.Sockets.Datagram
}
}
- private async Task HandleClientRequest(object clientRequestObj)
+ private async Task HandleClientRequest(ClientRequest clientRequest)
{
- SocketAsyncEventArgs clientArgs = SocketArgsPool.Get();
+ SocketAsyncEventArgs clientArgs = TransmissionArgsPool.Get();
byte[] responseBuffer = BufferPool.Rent(NetworkPacket.TotalSize);
Memory<byte> responseBufferMemory = new Memory<byte>(responseBuffer);
- ClientRequest clientRequest = (ClientRequest) clientRequestObj;
-
NetworkPacket request = clientRequest.RequestPacket;
EndPoint remoteEndPoint = clientRequest.ClientEndPoint;
CancellationToken cancellationToken = clientRequest.CancellationToken;
@@ -148,11 +181,283 @@ namespace NetSharp.Sockets.Datagram
BufferPool.Return(responseBuffer, true);
- SocketArgsPool.Return(clientArgs);
+ TransmissionArgsPool.Return(clientArgs);
+ }
+
+ private readonly struct ClientPacket
+ {
+ public readonly SocketAsyncEventArgs RentedArgs;
+
+ public readonly byte[] RentedBuffer;
+
+ public readonly Memory<byte> RentedBufferMemory;
+
+ public readonly EndPoint RemoteEndPoint;
+
+ public ClientPacket(in SocketAsyncEventArgs rentedArgs, in byte[] rentedBuffer, in Memory<byte> rentedBufferMemory, in EndPoint remoteEndPoint)
+ {
+ RentedArgs = rentedArgs;
+
+ RentedBuffer = rentedBuffer;
+
+ RentedBufferMemory = rentedBufferMemory;
+
+ RemoteEndPoint = remoteEndPoint;
+ }
+ }
+
+ private readonly struct SocketOperationToken
+ {
+ public readonly byte[] RentedBuffer;
+
+ public readonly Memory<byte> RentedBufferMemory;
+
+ public SocketOperationToken(in byte[] rentedBuffer, in Memory<byte> rentedBufferMemory)
+ {
+ RentedBuffer = rentedBuffer;
+
+ RentedBufferMemory = rentedBufferMemory;
+ }
+ }
+
+ private void HandleIoCompleted(object sender, SocketAsyncEventArgs args)
+ {
+ switch (args.LastOperation)
+ {
+ case SocketAsyncOperation.SendTo:
+ CompleteSendTo(args);
+ break;
+
+ case SocketAsyncOperation.ReceiveFrom:
+ CompleteReceiveFrom(args);
+ break;
+
+ default:
+ throw new NotSupportedException($"{nameof(HandleIoCompleted)} doesn't support {args.LastOperation}");
+ }
+ }
+
+ private void DoSendTo(SocketAsyncEventArgs sendArgs)
+ {
+ bool completesAsync = connection.SendToAsync(sendArgs);
+
+ if (!completesAsync)
+ {
+ CompleteSendTo(sendArgs);
+ }
}
- public override async Task RunAsync(CancellationToken cancellationToken = default)
+ private void CompleteSendTo(SocketAsyncEventArgs sendArgs)
{
+ SocketOperationToken sendToken = (SocketOperationToken) sendArgs.UserToken;
+
+ TransmissionResult sendResult = new TransmissionResult(sendArgs);
+
+#if DEBUG
+ lock (typeof(Console))
+ {
+ Console.WriteLine($"[Server] Sent {sendResult.Count} bytes to {sendResult.RemoteEndPoint}");
+ Console.WriteLine($"[Server] >>>> {Encoding.UTF8.GetString(sendResult.Buffer.Span)}");
+ }
+#endif
+
+ BufferPool.Return(sendToken.RentedBuffer, true);
+
+ ArgsPool.Return(sendArgs);
+ }
+
+ private void DoReceiveFrom(EndPoint remoteEndPoint)
+ {
+ SocketAsyncEventArgs args = ArgsPool.Rent();
+
+ byte[] receiveBuffer = BufferPool.Rent(NetworkPacket.TotalSize);
+ Memory<byte> receiveBufferMemory = new Memory<byte>(receiveBuffer);
+
+ args.SetBuffer(receiveBufferMemory);
+ args.RemoteEndPoint = remoteEndPoint;
+ args.UserToken = new SocketOperationToken(in receiveBuffer, in receiveBufferMemory);
+
+ bool completesAsync = connection.ReceiveFromAsync(args);
+
+ if (!completesAsync)
+ {
+ CompleteReceiveFrom(args);
+ }
+ }
+
+ private void CompleteReceiveFrom(SocketAsyncEventArgs receiveArgs)
+ {
+ DoReceiveFrom(AnyRemoteEndPoint); // start a new receive from operation immediately, to not drop any packets
+
+ SocketOperationToken receiveToken = (SocketOperationToken)receiveArgs.UserToken;
+
+ TransmissionResult receiveResult = new TransmissionResult(receiveArgs);
+
+#if DEBUG
+ lock (typeof(Console))
+ {
+ Console.WriteLine($"[Server] Received {receiveResult.Count} bytes from {receiveResult.RemoteEndPoint}");
+ Console.WriteLine($"[Server] <<<< {Encoding.UTF8.GetString(receiveResult.Buffer.Span)}");
+ }
+#endif
+
+ NetworkPacket request = NetworkPacket.Deserialise(receiveArgs.MemoryBuffer);
+
+ NetworkPacket response = request;
+
+ SocketAsyncEventArgs sendArgs = ArgsPool.Rent();
+
+ byte[] sendBuffer = BufferPool.Rent(NetworkPacket.TotalSize);
+ Memory<byte> sendBufferMemory = new Memory<byte>(sendBuffer);
+
+ NetworkPacket.Serialise(response, sendBufferMemory);
+
+ sendArgs.SetBuffer(sendBufferMemory);
+ sendArgs.RemoteEndPoint = receiveResult.RemoteEndPoint;
+ sendArgs.UserToken = new SocketOperationToken(in sendBuffer, in sendBufferMemory);
+
+ DoSendTo(sendArgs);
+
+ BufferPool.Return(receiveToken.RentedBuffer, true);
+
+ ArgsPool.Return(receiveArgs);
+ }
+
+ public override Task RunAsync(CancellationToken cancellationToken = default)
+ {
+ for (int i = 0; i < 10; i++)
+ {
+ DoReceiveFrom(AnyRemoteEndPoint);
+ }
+
+ return cancellationToken.WaitHandle.WaitOneAsync();
+ /*
+ UnboundedChannelOptions requestChannelOptions = new UnboundedChannelOptions()
+ {
+ AllowSynchronousContinuations = true,
+ SingleReader = true,
+ SingleWriter = true,
+ };
+ Channel<ClientPacket> requestChannel = Channel.CreateUnbounded<ClientPacket>(requestChannelOptions);
+
+ UnboundedChannelOptions responseChannelOptions = new UnboundedChannelOptions()
+ {
+ AllowSynchronousContinuations = true,
+ SingleReader = true,
+ SingleWriter = true,
+ };
+ Channel<ClientPacket> responseChannel = Channel.CreateUnbounded<ClientPacket>(responseChannelOptions);
+
+ async Task ReceivePacketTask()
+ {
+ EndPoint anyRemoteEndPoint = new IPEndPoint(IPAddress.Any, 0);
+
+ while (!cancellationToken.IsCancellationRequested)
+ {
+ SocketAsyncEventArgs transmissionArgs = TransmissionArgsPool.Get();
+
+ byte[] requestBuffer = BufferPool.Rent(NetworkPacket.TotalSize);
+ Memory<byte> requestBufferMemory = new Memory<byte>(requestBuffer);
+
+ TransmissionResult receiveResult =
+ await SocketAsyncOperations
+ .ReceiveFromAsync(transmissionArgs, connection, anyRemoteEndPoint, SocketFlags.None, requestBufferMemory, cancellationToken)
+ .ConfigureAwait(false);
+
+#if DEBUG
+ lock (typeof(Console))
+ {
+ Console.WriteLine($"[Server] Received {receiveResult.Count} bytes from {receiveResult.RemoteEndPoint}");
+ Console.WriteLine($"[Server] <<<< {Encoding.UTF8.GetString(receiveResult.Buffer.Span)}");
+ }
+#endif
+ ClientPacket request = new ClientPacket(in transmissionArgs, in requestBuffer, in requestBufferMemory, in receiveResult.RemoteEndPoint);
+
+ await requestChannel.Writer.WriteAsync(request, cancellationToken);
+ }
+ }
+
+ async Task HandleRequestTask()
+ {
+ while (!cancellationToken.IsCancellationRequested)
+ {
+ ClientPacket request = await requestChannel.Reader.ReadAsync(cancellationToken);
+
+ NetworkPacket requestPacket = NetworkPacket.Deserialise(request.RentedBufferMemory);
+
+ // TODO implement actual request handling, besides just an echo
+ NetworkPacket responsePacket = requestPacket;
+
+ byte[] responseBuffer = BufferPool.Rent(NetworkPacket.TotalSize);
+ Memory<byte> responseBufferMemory = new Memory<byte>(responseBuffer);
+
+ NetworkPacket.Serialise(responsePacket, responseBufferMemory);
+
+ // after the response has been serialised to the response buffer, the request buffer is done with and can be freed
+ BufferPool.Return(request.RentedBuffer, true);
+
+ ClientPacket response = new ClientPacket(in request.RentedArgs, in responseBuffer, in responseBufferMemory, in request.RemoteEndPoint);
+
+ await responseChannel.Writer.WriteAsync(response, cancellationToken);
+ }
+ }
+
+ async Task SendResponseTask()
+ {
+ while (!cancellationToken.IsCancellationRequested)
+ {
+ ClientPacket response = await responseChannel.Reader.ReadAsync(cancellationToken);
+
+ EndPoint remoteEndPoint = response.RemoteEndPoint;
+ Memory<byte> responseBufferMemory = response.RentedBufferMemory;
+
+ TransmissionResult sendResult =
+ await SocketAsyncOperations
+ .SendToAsync(response.RentedArgs, connection, remoteEndPoint, SocketFlags.None, responseBufferMemory, cancellationToken)
+ .ConfigureAwait(false);
+
+#if DEBUG
+ lock (typeof(Console))
+ {
+ Console.WriteLine($"[Server] Sent {sendResult.Count} bytes to {sendResult.RemoteEndPoint}");
+ Console.WriteLine($"[Server] >>>> {Encoding.UTF8.GetString(sendResult.Buffer.Span)}");
+ }
+#endif
+
+ BufferPool.Return(response.RentedBuffer, true);
+
+ TransmissionArgsPool.Return(response.RentedArgs);
+ }
+ }
+
+ // TODO this still functions as the simple example below, in fact it performs worse, with a worse bandwidth :(
+ List<Task> completeTaskList = new List<Task>();
+
+ Task[] readRequestTasks = new Task[2];
+ for (int i = 0; i < readRequestTasks.Length; i++)
+ {
+ readRequestTasks[i] = ReceivePacketTask();
+ completeTaskList.Add(readRequestTasks[i]);
+ }
+
+ Task[] handleRequestTasks = new Task[2];
+ for (int i = 0; i < handleRequestTasks.Length; i++)
+ {
+ handleRequestTasks[i] = HandleRequestTask();
+ completeTaskList.Add(handleRequestTasks[i]);
+ }
+
+ Task[] writeResponseTasks = new Task[2];
+ for (int i = 0; i < writeResponseTasks.Length; i++)
+ {
+ writeResponseTasks[i] = SendResponseTask();
+ completeTaskList.Add(writeResponseTasks[i]);
+ }
+
+ await Task.WhenAll(completeTaskList);
+ */
+
+ /*
EndPoint remoteEndPoint = new IPEndPoint(IPAddress.Any, 0);
using SocketAsyncEventArgs remoteArgs = GenerateConnectionArgs(remoteEndPoint);
@@ -182,6 +487,7 @@ namespace NetSharp.Sockets.Datagram
Task _ = HandleClientRequest(request);
}
+ */
/*
EndPoint remoteEndPoint = new IPEndPoint(IPAddress.Any, 0);
@@ -227,33 +533,56 @@ namespace NetSharp.Sockets.Datagram
await connectedClientTokens[clientEndPoint].PacketWriter.WriteAsync(request, cancellationToken);
}
*/
+ }
+ }
- /*
- while (true)
- {
- TransmissionResult receiveResult =
- await SocketAsyncEventArgs.ReceiveAsync(remoteEndPoint, SocketFlags.None, requestBufferMemory);
+ internal class ObjectPool<T> where T : class
+ {
+ internal delegate T CreateObjectDelegate();
-#if DEBUG
- lock (typeof(Console))
- {
- Console.WriteLine($"[Server] Received {receiveResult.Count} bytes from {receiveResult.RemoteEndPoint}");
- Console.WriteLine($"[Server] <<<< {Encoding.UTF8.GetString(receiveResult.Buffer.Span)}");
- }
-#endif
+ internal delegate bool KeepObjectPredicate(in T instance);
- TransmissionResult sendResult =
- await SocketAsyncOperations.SendAsync(receiveResult.RemoteEndPoint, SocketFlags.None, requestBuffer);
+ internal delegate void ResetObjectDelegate(T instance);
-#if DEBUG
- lock (typeof(Console))
- {
- Console.WriteLine($"[Server] Sent {sendResult.Count} bytes to {sendResult.RemoteEndPoint}");
- Console.WriteLine($"[Server] >>>> {Encoding.UTF8.GetString(sendResult.Buffer.Span)}");
- }
-#endif
+ internal delegate void DestroyObjectDelegate(T instance);
+
+ private readonly CreateObjectDelegate createObjectDelegate;
+ private readonly KeepObjectPredicate rebufferObjectPredicate;
+ private readonly ResetObjectDelegate resetObjectDelegate;
+ private readonly DestroyObjectDelegate destroyObjectDelegate;
+
+ private readonly ConcurrentBag<T> objectBuffer;
+
+ public ObjectPool(in CreateObjectDelegate createDelegate, in ResetObjectDelegate resetDelegate, in DestroyObjectDelegate destroyDelegate, in KeepObjectPredicate keepObjectPredicate)
+ {
+ createObjectDelegate = createDelegate;
+
+ resetObjectDelegate = resetDelegate;
+
+ destroyObjectDelegate = destroyDelegate;
+
+ rebufferObjectPredicate = keepObjectPredicate;
+
+ objectBuffer = new ConcurrentBag<T>();
+ }
+
+ public T Rent()
+ {
+ return objectBuffer.TryTake(out T result) ? result : createObjectDelegate();
+ }
+
+ public void Return(T instance)
+ {
+ if (rebufferObjectPredicate(instance))
+ {
+ resetObjectDelegate(instance);
+
+ objectBuffer.Add(instance);
+ }
+ else
+ {
+ destroyObjectDelegate(instance);
}
- */
}
}
}
\ No newline at end of file
diff --git a/NetSharp/NetSharp/Sockets/SocketServer.cs b/NetSharp/NetSharp/Sockets/SocketServer.cs
@@ -16,14 +16,14 @@ namespace NetSharp.Sockets
protected readonly ArrayPool<byte> BufferPool;
- protected readonly ObjectPool<SocketAsyncEventArgs> SocketArgsPool;
+ protected readonly ObjectPool<SocketAsyncEventArgs> TransmissionArgsPool;
protected SocketServer(in AddressFamily connectionAddressFamily, in SocketType connectionSocketType, in ProtocolType connectionProtocolType)
: base(in connectionAddressFamily, in connectionSocketType, in connectionProtocolType)
{
- BufferPool = ArrayPool<byte>.Create(NetworkPacket.TotalSize, 10);
+ BufferPool = ArrayPool<byte>.Create(NetworkPacket.TotalSize, 100);
- SocketArgsPool = new DefaultObjectPool<SocketAsyncEventArgs>(new PooledSocketAsyncEventArgsPolicy());
+ TransmissionArgsPool = new DefaultObjectPool<SocketAsyncEventArgs>(new PooledSocketAsyncEventArgsPolicy());
ConnectedClientHandlerTasks = new ConcurrentDictionary<EndPoint, Task>();
}
diff --git a/NetSharp/NetSharp/Utils/WaitHandleExtensions.cs b/NetSharp/NetSharp/Utils/WaitHandleExtensions.cs
@@ -0,0 +1,44 @@
+using System;
+using System.Threading;
+using System.Threading.Tasks;
+
+namespace NetSharp.Utils
+{
+ public static class WaitHandleExtensions
+ {
+ public static Task<bool> WaitOneAsync(this WaitHandle instance, TimeSpan timeout, CancellationToken cancellationToken = default)
+ {
+ TaskCompletionSource<bool> tcs = new TaskCompletionSource<bool>();
+
+ RegisteredWaitHandle? registeredHandle = default;
+ CancellationTokenRegistration tokenRegistration = default;
+
+ try
+ {
+ registeredHandle = ThreadPool.RegisterWaitForSingleObject(
+ instance,
+ (state, timedOut) => ((TaskCompletionSource<bool>)state).TrySetResult(!timedOut),
+ tcs,
+ timeout,
+ true);
+
+ tokenRegistration = cancellationToken.Register(
+ state => ((TaskCompletionSource<bool>)state).TrySetCanceled(),
+ tcs);
+
+ return tcs.Task;
+ }
+ finally
+ {
+ registeredHandle?.Unregister(null);
+ tokenRegistration.Dispose();
+ }
+ }
+
+ public static Task<bool> WaitOneAsync(this WaitHandle instance, int timeoutMs, CancellationToken cancellationToken = default)
+ => instance.WaitOneAsync(TimeSpan.FromMilliseconds(timeoutMs), cancellationToken);
+
+ public static Task<bool> WaitOneAsync(this WaitHandle instance, CancellationToken cancellationToken = default)
+ => instance.WaitOneAsync(Timeout.InfiniteTimeSpan, cancellationToken);
+ }
+}
+\ No newline at end of file
diff --git a/NetSharp/NetSharpExamples/Program.cs b/NetSharp/NetSharpExamples/Program.cs
@@ -23,15 +23,19 @@ namespace NetSharpExamples
{
private const int NetworkTimeout = 1_000_000;
private const int ServerPort = 12374;
- private static readonly IPAddress ServerAddress = IPAddress.Parse("192.168.0.15");
+ private static readonly IPAddress ServerAddress = IPAddress.Parse("192.168.0.10");
private static readonly EndPoint ServerEndPoint = new IPEndPoint(ServerAddress, ServerPort);
private static async Task Main()
{
Console.WriteLine("Hello World!");
- await Task.Factory.StartNew(TestSocketServer);
- await Task.Factory.StartNew(TestSocketClient).Result;
+ Console.WriteLine("Starting socket server test...");
+ Task serverTest = TestSocketServer();
+
+ Console.WriteLine("Starting socket client test...");
+ Task clientTest = TestSocketClient();
+ await clientTest;
Console.ReadLine();
}
@@ -40,7 +44,7 @@ namespace NetSharpExamples
private static async Task TestSocketClient()
{
- const int clientCount = 16;
+ const int clientCount = 20;
const long packetsToSend = 100_000;
Task[] clientTasks = new Task[clientCount];
@@ -220,15 +224,23 @@ namespace NetSharpExamples
private static async Task TestSocketServer()
{
+ try
+ {
#if TCP
- using SocketServer server = new StreamSocketServer(AddressFamily.InterNetwork, ProtocolType.Tcp);
+ using SocketServer server = new StreamSocketServer(AddressFamily.InterNetwork, ProtocolType.Tcp);
#else
- using SocketServer server = new DatagramSocketServer(AddressFamily.InterNetwork, ProtocolType.Udp);
+ using SocketServer server = new DatagramSocketServer(AddressFamily.InterNetwork, ProtocolType.Udp);
#endif
- server.Bind(in ServerEndPoint);
+ server.Bind(in ServerEndPoint);
- await server.RunAsync();
+ await server.RunAsync();
+ }
+ catch (Exception ex)
+ {
+ Console.WriteLine(ex);
+ throw;
+ }
}
#endregion Socket Tests