NetSharp

NetSharp.git
git clone git://git.lenczewski.org/NetSharp.git
Log | Files | Refs | README | LICENSE

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:
MNetSharp/NetSharp/Sockets/Datagram/DatagramSocketServer.cs | 389++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++-------
MNetSharp/NetSharp/Sockets/SocketServer.cs | 6+++---
ANetSharp/NetSharp/Utils/WaitHandleExtensions.cs | 45+++++++++++++++++++++++++++++++++++++++++++++
MNetSharp/NetSharpExamples/Program.cs | 28++++++++++++++++++++--------
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