NetSharp

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

commit 008ea478e6956169ced2b06463a6f116ee5a62df
parent 97cf9b9f20d1a0ee9781222d9a07bb645400fb39
Author: Mikolaj Lenczewski <mikolaj.lenczewski308@gmail.com>
Date:   Fri, 10 Apr 2020 22:15:47 +0100

Improved on StreamSocketClient, implementing synchronous and asynchronous methods. Memory leak issue exists which needs to be debugged.

Diffstat:
MNetSharp/NetSharp/Packets/NetworkPacket.cs | 4++--
MNetSharp/NetSharp/Sockets/Datagram/DatagramSocketClient.cs | 101+++++++++++++++++++++++++++++++++++++++++--------------------------------------
MNetSharp/NetSharp/Sockets/Datagram/DatagramSocketServer.cs | 118+++++++++++++++++++++++++++++++++++++++++--------------------------------------
MNetSharp/NetSharp/Sockets/SocketClient.cs | 67++++++++++++++++++++++++++++++++++++++++++++++---------------------
MNetSharp/NetSharp/Sockets/SocketConnection.cs | 1+
MNetSharp/NetSharp/Sockets/SocketServer.cs | 8++------
MNetSharp/NetSharp/Sockets/Stream/StreamSocketClient.cs | 259++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++-----
MNetSharp/NetSharp/Sockets/Stream/StreamSocketServer.cs | 268++++++++++++++++++++++++++++++++++++++++++++++++++++++++-----------------------
MNetSharp/NetSharpExamples/Program.cs | 14+++++++-------
9 files changed, 604 insertions(+), 236 deletions(-)

diff --git a/NetSharp/NetSharp/Packets/NetworkPacket.cs b/NetSharp/NetSharp/Packets/NetworkPacket.cs @@ -4,13 +4,13 @@ namespace NetSharp.Packets { public readonly struct NetworkPacket { - public const int TotalSize = 4096; + public const int TotalSize = HeaderSize + DataSize + FooterSize; public const int HeaderSize = NetworkPacketHeader.TotalSize; public const int FooterSize = NetworkPacketFooter.TotalSize; - public const int DataSize = TotalSize - HeaderSize - FooterSize; + public const int DataSize = 8192; public readonly ReadOnlyMemory<byte> Data; diff --git a/NetSharp/NetSharp/Sockets/Datagram/DatagramSocketClient.cs b/NetSharp/NetSharp/Sockets/Datagram/DatagramSocketClient.cs @@ -7,23 +7,11 @@ using NetSharp.Utils; namespace NetSharp.Sockets.Datagram { + //TODO fix memory leak issue + //TODO address the need to handle series of network packets, not just single packets //TODO document class public sealed class DatagramSocketClient : SocketClient { - private readonly struct SocketOperationToken - { - public readonly TaskCompletionSource<TransmissionResult> CompletionSource; - - public readonly CancellationToken CancellationToken; - - public SocketOperationToken(in TaskCompletionSource<TransmissionResult> completionSource, in CancellationToken cancellationToken) - { - CompletionSource = completionSource; - - CancellationToken = cancellationToken; - } - } - public DatagramSocketClient(in AddressFamily connectionAddressFamily, in ProtocolType connectionProtocolType) : base(in connectionAddressFamily, SocketType.Dgram, in connectionProtocolType) { @@ -58,34 +46,30 @@ namespace NetSharp.Sockets.Datagram { switch (args.LastOperation) { - case SocketAsyncOperation.SendTo: - SocketOperationToken sendToken = (SocketOperationToken) args.UserToken; - - if (sendToken.CancellationToken.IsCancellationRequested) + case SocketAsyncOperation.Connect: + AsyncOperationToken connectToken = (AsyncOperationToken)args.UserToken; + + if (connectToken.CancellationToken.IsCancellationRequested) { - sendToken.CompletionSource.SetCanceled(); + connectToken.CompletionSource.SetCanceled(); } else if (args.SocketError == SocketError.Success) { - TransmissionResult result = new TransmissionResult(in args); - - sendToken.CompletionSource.SetResult(result); - - TransmissionArgsPool.Return(args); + connectToken.CompletionSource.SetResult(true); } else { - sendToken.CompletionSource.SetException(new SocketException((int)args.SocketError)); - - TransmissionArgsPool.Return(args); + connectToken.CompletionSource.SetException(new SocketException((int)args.SocketError)); } + TransmissionArgsPool.Return(args); + break; case SocketAsyncOperation.ReceiveFrom: - SocketOperationToken receiveToken = (SocketOperationToken)args.UserToken; + AsyncTransmissionToken receiveToken = (AsyncTransmissionToken)args.UserToken; - if (sendToken.CancellationToken.IsCancellationRequested) + if (receiveToken.CancellationToken.IsCancellationRequested) { receiveToken.CompletionSource.SetCanceled(); } @@ -94,15 +78,35 @@ namespace NetSharp.Sockets.Datagram TransmissionResult result = new TransmissionResult(in args); receiveToken.CompletionSource.SetResult(result); - - TransmissionArgsPool.Return(args); } else { receiveToken.CompletionSource.SetException(new SocketException((int)args.SocketError)); + } + + TransmissionArgsPool.Return(args); + + break; - TransmissionArgsPool.Return(args); + case SocketAsyncOperation.SendTo: + AsyncTransmissionToken sendToken = (AsyncTransmissionToken) args.UserToken; + + if (sendToken.CancellationToken.IsCancellationRequested) + { + sendToken.CompletionSource.SetCanceled(); } + else if (args.SocketError == SocketError.Success) + { + TransmissionResult result = new TransmissionResult(in args); + + sendToken.CompletionSource.SetResult(result); + } + else + { + sendToken.CompletionSource.SetException(new SocketException((int)args.SocketError)); + } + + TransmissionArgsPool.Return(args); break; @@ -111,60 +115,59 @@ namespace NetSharp.Sockets.Datagram } } - public TransmissionResult SendTo(EndPoint remoteEndPoint, byte[] sendBuffer, SocketFlags flags = SocketFlags.None) + public TransmissionResult ReceiveFrom(ref EndPoint remoteEndPoint, byte[] receiveBuffer, SocketFlags flags = SocketFlags.None) { - int sentBytes = connection.SendTo(sendBuffer, flags, remoteEndPoint); + int receivedBytes = connection.ReceiveFrom(receiveBuffer, flags, ref remoteEndPoint); - return new TransmissionResult(in sendBuffer, in sentBytes, in remoteEndPoint); + return new TransmissionResult(in receiveBuffer, in receivedBytes, in remoteEndPoint); } - public ValueTask<TransmissionResult> SendToAsync(EndPoint remoteEndPoint, Memory<byte> sendBuffer, + public ValueTask<TransmissionResult> ReceiveFromAsync(EndPoint remoteEndPoint, Memory<byte> receiveBuffer, SocketFlags flags = SocketFlags.None, CancellationToken cancellationToken = default) { TaskCompletionSource<TransmissionResult> tcs = new TaskCompletionSource<TransmissionResult>(); SocketAsyncEventArgs args = TransmissionArgsPool.Rent(); - args.SetBuffer(sendBuffer); + args.SetBuffer(receiveBuffer); args.RemoteEndPoint = remoteEndPoint; args.SocketFlags = flags; - args.UserToken = new SocketOperationToken(in tcs, in cancellationToken); + args.UserToken = new AsyncTransmissionToken(in tcs, in cancellationToken); - if (connection.SendToAsync(args)) return new ValueTask<TransmissionResult>(tcs.Task); + if (connection.ReceiveFromAsync(args)) return new ValueTask<TransmissionResult>(tcs.Task); TransmissionResult result = new TransmissionResult(in args); - + TransmissionArgsPool.Return(args); return new ValueTask<TransmissionResult>(result); - } - public TransmissionResult ReceiveFrom(ref EndPoint remoteEndPoint, byte[] receiveBuffer, SocketFlags flags = SocketFlags.None) + public TransmissionResult SendTo(EndPoint remoteEndPoint, byte[] sendBuffer, SocketFlags flags = SocketFlags.None) { - int readBytes = connection.ReceiveFrom(receiveBuffer, flags, ref remoteEndPoint); + int sentBytes = connection.SendTo(sendBuffer, flags, remoteEndPoint); - return new TransmissionResult(in receiveBuffer, in readBytes, in remoteEndPoint); + return new TransmissionResult(in sendBuffer, in sentBytes, in remoteEndPoint); } - public ValueTask<TransmissionResult> ReceiveFromAsync(EndPoint remoteEndPoint, Memory<byte> receiveBuffer, + public ValueTask<TransmissionResult> SendToAsync(EndPoint remoteEndPoint, Memory<byte> sendBuffer, SocketFlags flags = SocketFlags.None, CancellationToken cancellationToken = default) { TaskCompletionSource<TransmissionResult> tcs = new TaskCompletionSource<TransmissionResult>(); SocketAsyncEventArgs args = TransmissionArgsPool.Rent(); - args.SetBuffer(receiveBuffer); + args.SetBuffer(sendBuffer); args.RemoteEndPoint = remoteEndPoint; args.SocketFlags = flags; - args.UserToken = new SocketOperationToken(in tcs, in cancellationToken); + args.UserToken = new AsyncTransmissionToken(in tcs, in cancellationToken); - if (connection.ReceiveFromAsync(args)) return new ValueTask<TransmissionResult>(tcs.Task); + if (connection.SendToAsync(args)) return new ValueTask<TransmissionResult>(tcs.Task); TransmissionResult result = new TransmissionResult(in args); - + TransmissionArgsPool.Return(args); return new ValueTask<TransmissionResult>(result); diff --git a/NetSharp/NetSharp/Sockets/Datagram/DatagramSocketServer.cs b/NetSharp/NetSharp/Sockets/Datagram/DatagramSocketServer.cs @@ -29,9 +29,10 @@ namespace NetSharp.Sockets.Datagram } } - //TODO address the need for a fixed packet size (NetworkPacket.TotalSize; lines 109 and 145) - //TODO address the need for a fixed number of initial ReceiveFrom method calls + //TODO fix memory leak issue //TODO address the need to handle series of network packets, not just single packets + //TODO allow for the server to do more than just echo packets + //TODO document class public sealed class DatagramSocketServer : SocketServer { private static readonly EndPoint AnyRemoteEndPoint = new IPEndPoint(IPAddress.Any, 0); @@ -39,10 +40,10 @@ namespace NetSharp.Sockets.Datagram public readonly DatagramSocketServerOptions ServerOptions; public DatagramSocketServer(in AddressFamily connectionAddressFamily, in ProtocolType connectionProtocolType, - in DatagramSocketServerOptions serverOptions = default) : base(in connectionAddressFamily, SocketType.Dgram, + in DatagramSocketServerOptions? serverOptions = null) : base(in connectionAddressFamily, SocketType.Dgram, in connectionProtocolType) { - ServerOptions = serverOptions.Equals(default) ? DatagramSocketServerOptions.Defaults : serverOptions; + ServerOptions = serverOptions ?? DatagramSocketServerOptions.Defaults; } private readonly struct SocketOperationToken @@ -84,64 +85,42 @@ namespace NetSharp.Sockets.Datagram { switch (args.LastOperation) { - case SocketAsyncOperation.SendTo: - CompleteSendTo(args); - break; - case SocketAsyncOperation.ReceiveFrom: - ReceiveFrom(AnyRemoteEndPoint); // start a new receive from operation immediately, to not drop any packets + SocketAsyncEventArgs newReceiveArgs = TransmissionArgsPool.Rent(); + newReceiveArgs.RemoteEndPoint = AnyRemoteEndPoint; + + ReceiveFrom(newReceiveArgs); // start a new receive from operation immediately, to not drop any packets CompleteReceiveFrom(args); break; + case SocketAsyncOperation.SendTo: + CompleteSendTo(args); + break; + default: throw new NotSupportedException($"{nameof(HandleIoCompleted)} doesn't support {args.LastOperation}"); } } - private void SendTo(SocketAsyncEventArgs sendArgs) - { - if (!connection.SendToAsync(sendArgs)) - { - CompleteSendTo(sendArgs); - } - } - - private void CompleteSendTo(SocketAsyncEventArgs sendArgs) - { - SocketOperationToken sendToken = (SocketOperationToken) sendArgs.UserToken; - - TransmissionResult sendResult = new TransmissionResult(in 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); - - TransmissionArgsPool.Return(sendArgs); - } - - private void ReceiveFrom(EndPoint remoteEndPoint) + private void ReceiveFrom(SocketAsyncEventArgs receiveArgs) { - SocketAsyncEventArgs args = TransmissionArgsPool.Rent(); - - byte[] receiveBuffer = BufferPool.Rent(NetworkPacket.TotalSize); + byte[] receiveBuffer = BufferPool.Rent(ServerOptions.PacketSize); Memory<byte> receiveBufferMemory = new Memory<byte>(receiveBuffer); - args.SetBuffer(receiveBufferMemory); - args.RemoteEndPoint = remoteEndPoint; - args.UserToken = new SocketOperationToken(in receiveBuffer); + receiveArgs.SetBuffer(receiveBufferMemory); + receiveArgs.UserToken = new SocketOperationToken(in receiveBuffer); - if (!connection.ReceiveFromAsync(args)) + bool operationPending = connection.ReceiveFromAsync(receiveArgs); + + if (!operationPending) { - ReceiveFrom(AnyRemoteEndPoint); // start a new receive from operation immediately, to not drop any packets + SocketAsyncEventArgs newReceiveArgs = TransmissionArgsPool.Rent(); + newReceiveArgs.RemoteEndPoint = AnyRemoteEndPoint; + + ReceiveFrom(newReceiveArgs); // start a new receive from operation immediately, to not drop any packets - CompleteReceiveFrom(args); + CompleteReceiveFrom(receiveArgs); } } @@ -164,29 +143,56 @@ namespace NetSharp.Sockets.Datagram // TODO implement actual request processing, not just an echo server NetworkPacket response = request; - SocketAsyncEventArgs sendArgs = TransmissionArgsPool.Rent(); - - byte[] sendBuffer = BufferPool.Rent(NetworkPacket.TotalSize); + byte[] sendBuffer = BufferPool.Rent(ServerOptions.PacketSize); Memory<byte> sendBufferMemory = new Memory<byte>(sendBuffer); NetworkPacket.Serialise(response, sendBufferMemory); - sendArgs.SetBuffer(sendBufferMemory); - sendArgs.RemoteEndPoint = receiveResult.RemoteEndPoint; - sendArgs.UserToken = new SocketOperationToken(in sendBuffer); + receiveArgs.SetBuffer(sendBufferMemory); + receiveArgs.UserToken = new SocketOperationToken(in sendBuffer); - SendTo(sendArgs); + SendTo(receiveArgs); BufferPool.Return(receiveToken.RentedBuffer, true); + } - TransmissionArgsPool.Return(receiveArgs); + private void SendTo(SocketAsyncEventArgs sendArgs) + { + bool operationPending = connection.SendToAsync(sendArgs); + + if (!operationPending) + { + CompleteSendTo(sendArgs); + } + } + + private void CompleteSendTo(SocketAsyncEventArgs sendArgs) + { + SocketOperationToken sendToken = (SocketOperationToken) sendArgs.UserToken; + + TransmissionResult sendResult = new TransmissionResult(in 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); + + TransmissionArgsPool.Return(sendArgs); } public override Task RunAsync(CancellationToken cancellationToken = default) { - for (int i = 0; i < 10; i++) + for (int i = 0; i < ServerOptions.ConcurrentReceiveFromCalls; i++) { - ReceiveFrom(AnyRemoteEndPoint); + SocketAsyncEventArgs newReceiveArgs = TransmissionArgsPool.Rent(); + newReceiveArgs.RemoteEndPoint = AnyRemoteEndPoint; + + ReceiveFrom(newReceiveArgs); } return cancellationToken.WaitHandle.WaitOneAsync(); diff --git a/NetSharp/NetSharp/Sockets/SocketClient.cs b/NetSharp/NetSharp/Sockets/SocketClient.cs @@ -1,43 +1,68 @@ -using NetSharp.Packets; - -using System; -using System.Buffers; -using System.Net; +using System.Net; using System.Net.Sockets; +using System.Threading; using System.Threading.Tasks; using NetSharp.Utils; namespace NetSharp.Sockets { + //TODO document class public abstract class SocketClient : SocketConnection { - protected readonly SocketAsyncEventArgs Args; + //TODO document + protected readonly struct AsyncTransmissionToken + { + public readonly TaskCompletionSource<TransmissionResult> CompletionSource; + + public readonly CancellationToken CancellationToken; + + public AsyncTransmissionToken(in TaskCompletionSource<TransmissionResult> completionSource, in CancellationToken cancellationToken) + { + CompletionSource = completionSource; + + CancellationToken = cancellationToken; + } + } + + //TODO document + protected readonly struct AsyncOperationToken + { + public readonly TaskCompletionSource<bool> CompletionSource; + + public readonly CancellationToken CancellationToken; + + public AsyncOperationToken(in TaskCompletionSource<bool> completionSource, in CancellationToken cancellationToken) + { + CompletionSource = completionSource; + + CancellationToken = cancellationToken; + } + } protected SocketClient(in AddressFamily connectionAddressFamily, in SocketType connectionSocketType, in ProtocolType connectionProtocolType) : base(in connectionAddressFamily, in connectionSocketType, in connectionProtocolType) { - Args = new SocketAsyncEventArgs(); - Args.Completed += SocketAsyncOperations.HandleIoCompleted; } - public int SendBytes(Memory<byte> outgoingDataBuffer, SocketFlags flags = SocketFlags.None) + public void Connect(in EndPoint remoteEndPoint) { - byte[] temporaryBuffer = BufferPool.Rent(NetworkPacket.TotalSize); - outgoingDataBuffer.CopyTo(temporaryBuffer); - int sentBytes = connection.Send(temporaryBuffer); - BufferPool.Return(temporaryBuffer); - - return sentBytes; + connection.Connect(remoteEndPoint); } - public int ReceiveBytes(Memory<byte> incomingDataBuffer, SocketFlags flags = SocketFlags.None) + public ValueTask ConnectAsync(in EndPoint remoteEndPoint, CancellationToken cancellationToken = default) { - byte[] temporaryBuffer = BufferPool.Rent(NetworkPacket.TotalSize); - int receivedBytes = connection.Receive(temporaryBuffer); - temporaryBuffer.CopyTo(incomingDataBuffer); - BufferPool.Return(temporaryBuffer); + TaskCompletionSource<bool> tcs = new TaskCompletionSource<bool>(); + + SocketAsyncEventArgs args = TransmissionArgsPool.Rent(); + + args.RemoteEndPoint = remoteEndPoint; + args.UserToken = new AsyncOperationToken(in tcs, in cancellationToken); + + if (connection.ConnectAsync(args)) return new ValueTask(tcs.Task); + + TransmissionArgsPool.Return(args); - return receivedBytes; + return new ValueTask(); } } } \ No newline at end of file diff --git a/NetSharp/NetSharp/Sockets/SocketConnection.cs b/NetSharp/NetSharp/Sockets/SocketConnection.cs @@ -7,6 +7,7 @@ using NetSharp.Utils; namespace NetSharp.Sockets { + //TODO document class public abstract class SocketConnection : IDisposable { protected readonly Socket connection; diff --git a/NetSharp/NetSharp/Sockets/SocketServer.cs b/NetSharp/NetSharp/Sockets/SocketServer.cs @@ -1,19 +1,15 @@ -using System.Collections.Concurrent; -using System.Net; -using System.Net.Sockets; +using System.Net.Sockets; using System.Threading; using System.Threading.Tasks; namespace NetSharp.Sockets { + //TODO document class public abstract class SocketServer : SocketConnection { - protected readonly ConcurrentDictionary<EndPoint, Task> ConnectedClientHandlerTasks; - protected SocketServer(in AddressFamily connectionAddressFamily, in SocketType connectionSocketType, in ProtocolType connectionProtocolType) : base(in connectionAddressFamily, in connectionSocketType, in connectionProtocolType) { - ConnectedClientHandlerTasks = new ConcurrentDictionary<EndPoint, Task>(); } public abstract Task RunAsync(CancellationToken cancellationToken = default); diff --git a/NetSharp/NetSharp/Sockets/Stream/StreamSocketClient.cs b/NetSharp/NetSharp/Sockets/Stream/StreamSocketClient.cs @@ -1,48 +1,275 @@ -using System.Net; +using System; +using System.Net; using System.Net.Sockets; +using System.Threading; +using System.Threading.Tasks; +using NetSharp.Utils; namespace NetSharp.Sockets.Stream { - public class StreamSocketClient : SocketClient + //TODO fix memory leak issue + //TODO address the need to handle series of network packets, not just single packets + //TODO document class + public sealed class StreamSocketClient : SocketClient { public StreamSocketClient(in AddressFamily connectionAddressFamily, in ProtocolType connectionProtocolType) : base(in connectionAddressFamily, SocketType.Stream, in connectionProtocolType) { } - public void Connect(in EndPoint remoteEndPoint) + protected override SocketAsyncEventArgs CreateTransmissionArgs() { - connection.Connect(remoteEndPoint); - } + SocketAsyncEventArgs connectionArgs = new SocketAsyncEventArgs(); - public void Disconnect() - { - connection.Disconnect(true); - } + connectionArgs.Completed += HandleIoCompleted; - protected override SocketAsyncEventArgs CreateTransmissionArgs() - { - throw new System.NotImplementedException(); + return connectionArgs; } protected override void ResetTransmissionArgs(SocketAsyncEventArgs args) { - throw new System.NotImplementedException(); } protected override bool CanTransmissionArgsBeReused(in SocketAsyncEventArgs args) { - throw new System.NotImplementedException(); + return true; } protected override void DestroyTransmissionArgs(SocketAsyncEventArgs remoteConnectionArgs) { - throw new System.NotImplementedException(); + remoteConnectionArgs.Completed -= HandleIoCompleted; + + remoteConnectionArgs.Dispose(); } protected override void HandleIoCompleted(object sender, SocketAsyncEventArgs args) { - throw new System.NotImplementedException(); + switch (args.LastOperation) + { + case SocketAsyncOperation.Connect: + AsyncOperationToken connectToken = (AsyncOperationToken) args.UserToken; + + if (connectToken.CancellationToken.IsCancellationRequested) + { + connectToken.CompletionSource.SetCanceled(); + } + else if (args.SocketError == SocketError.Success) + { + connectToken.CompletionSource.SetResult(true); + } + else + { + connectToken.CompletionSource.SetException(new SocketException((int)args.SocketError)); + } + + TransmissionArgsPool.Return(args); + + break; + + case SocketAsyncOperation.Disconnect: + AsyncOperationToken disconnectToken = (AsyncOperationToken) args.UserToken; + + if (disconnectToken.CancellationToken.IsCancellationRequested) + { + disconnectToken.CompletionSource.SetCanceled(); + } + else if (args.SocketError == SocketError.Success) + { + disconnectToken.CompletionSource.SetResult(true); + } + else + { + disconnectToken.CompletionSource.SetException(new SocketException((int)args.SocketError)); + } + + TransmissionArgsPool.Return(args); + + break; + + case SocketAsyncOperation.Receive: + AsyncTransmissionToken receiveToken = (AsyncTransmissionToken) args.UserToken; + + if (receiveToken.CancellationToken.IsCancellationRequested) + { + receiveToken.CompletionSource.SetCanceled(); + + TransmissionArgsPool.Return(args); + } + else if (args.SocketError == SocketError.Success) + { + Memory<byte> transmissionBuffer = args.MemoryBuffer; + int expectedBytes = transmissionBuffer.Length; + + if (args.BytesTransferred == expectedBytes) + { + // buffer was fully received + + TransmissionResult result = new TransmissionResult(in args); + + receiveToken.CompletionSource.SetResult(result); + + TransmissionArgsPool.Return(args); + } + else if (expectedBytes > args.BytesTransferred && args.BytesTransferred > 0) + { + // receive the remaining parts of the buffer + + int receivedBytes = args.BytesTransferred; + + args.SetBuffer(receivedBytes, expectedBytes - receivedBytes); + + connection.ReceiveAsync(args); + } + else + { + // no bytes were received, remote socket is dead + + receiveToken.CompletionSource.SetException(new SocketException((int)SocketError.HostDown)); + + TransmissionArgsPool.Return(args); + } + } + else + { + receiveToken.CompletionSource.SetException(new SocketException((int)args.SocketError)); + + TransmissionArgsPool.Return(args); + } + + break; + + case SocketAsyncOperation.Send: + AsyncTransmissionToken sendToken = (AsyncTransmissionToken) args.UserToken; + + if (sendToken.CancellationToken.IsCancellationRequested) + { + sendToken.CompletionSource.SetCanceled(); + + TransmissionArgsPool.Return(args); + } + else if (args.SocketError == SocketError.Success) + { + Memory<byte> transmissionBuffer = args.MemoryBuffer; + int remainingBytes = transmissionBuffer.Length; + + if (args.BytesTransferred == remainingBytes) + { + // buffer was fully sent + + TransmissionResult result = new TransmissionResult(in args); + + sendToken.CompletionSource.SetResult(result); + + TransmissionArgsPool.Return(args); + } + else if (remainingBytes > args.BytesTransferred && args.BytesTransferred > 0) + { + // send the remaining parts of the buffer + + int sentBytes = args.BytesTransferred; + + args.SetBuffer(sentBytes, remainingBytes - sentBytes); + + connection.SendAsync(args); + } + else + { + // no bytes were sent, remote socket is dead + + sendToken.CompletionSource.SetException(new SocketException((int)SocketError.HostDown)); + + TransmissionArgsPool.Return(args); + } + } + else + { + sendToken.CompletionSource.SetException(new SocketException((int)args.SocketError)); + + TransmissionArgsPool.Return(args); + } + + break; + + default: + throw new NotSupportedException($"{nameof(HandleIoCompleted)} doesn't support {args.LastOperation}"); + } + } + + public void Disconnect(bool allowSocketReuse) + { + connection.Disconnect(allowSocketReuse); + } + + public ValueTask DisconnectAsync(bool allowSocketReuse, CancellationToken cancellationToken = default) + { + TaskCompletionSource<bool> tcs = new TaskCompletionSource<bool>(); + + SocketAsyncEventArgs args = TransmissionArgsPool.Rent(); + + args.DisconnectReuseSocket = allowSocketReuse; + args.UserToken = new AsyncOperationToken(in tcs, in cancellationToken); + + if (connection.DisconnectAsync(args)) return new ValueTask(tcs.Task); + + TransmissionArgsPool.Return(args); + + return new ValueTask(); + } + + public TransmissionResult Receive(byte[] buffer, SocketFlags flags = SocketFlags.None) + { + int receivedBytes = connection.Receive(buffer, flags); + + return new TransmissionResult(in buffer, in receivedBytes, connection.RemoteEndPoint); + } + + public ValueTask<TransmissionResult> ReceiveAsync(Memory<byte> receiveBuffer, SocketFlags flags = SocketFlags.None, + CancellationToken cancellationToken = default) + { + TaskCompletionSource<TransmissionResult> tcs = new TaskCompletionSource<TransmissionResult>(); + + SocketAsyncEventArgs args = TransmissionArgsPool.Rent(); + + args.SetBuffer(receiveBuffer); + + args.SocketFlags = flags; + args.UserToken = new AsyncTransmissionToken(in tcs, in cancellationToken); + + if (connection.ReceiveAsync(args)) return new ValueTask<TransmissionResult>(tcs.Task); + + TransmissionResult result = new TransmissionResult(in args); + + TransmissionArgsPool.Return(args); + + return new ValueTask<TransmissionResult>(result); + } + + public TransmissionResult Send(byte[] buffer, SocketFlags flags = SocketFlags.None) + { + int sentBytes = connection.Send(buffer, flags); + + return new TransmissionResult(in buffer, in sentBytes, connection.RemoteEndPoint); + } + + public ValueTask<TransmissionResult> SendAsync(Memory<byte> sendBuffer, SocketFlags flags = SocketFlags.None, + CancellationToken cancellationToken = default) + { + TaskCompletionSource<TransmissionResult> tcs = new TaskCompletionSource<TransmissionResult>(); + + SocketAsyncEventArgs args = TransmissionArgsPool.Rent(); + + args.SetBuffer(sendBuffer); + + args.SocketFlags = flags; + args.UserToken = new AsyncTransmissionToken(in tcs, in cancellationToken); + + if (connection.SendToAsync(args)) return new ValueTask<TransmissionResult>(tcs.Task); + + TransmissionResult result = new TransmissionResult(in args); + + TransmissionArgsPool.Return(args); + + return new ValueTask<TransmissionResult>(result); } } } \ No newline at end of file diff --git a/NetSharp/NetSharp/Sockets/Stream/StreamSocketServer.cs b/NetSharp/NetSharp/Sockets/Stream/StreamSocketServer.cs @@ -11,6 +11,7 @@ using NetSharp.Utils; namespace NetSharp.Sockets.Stream { + //TODO document public readonly struct StreamSocketServerOptions { public static readonly StreamSocketServerOptions Defaults = @@ -18,45 +19,46 @@ namespace NetSharp.Sockets.Stream public readonly int PacketSize; - public readonly int ConcurrentReceiveCalls; + public readonly int ConcurrentAcceptCalls; - public StreamSocketServerOptions(int packetSize, int concurrentReceiveCalls) + public StreamSocketServerOptions(int packetSize, int concurrentAcceptCalls) { PacketSize = packetSize; - ConcurrentReceiveCalls = concurrentReceiveCalls; + ConcurrentAcceptCalls = concurrentAcceptCalls; } } - public class StreamSocketServer : SocketServer + //TODO fix memory leak issue + //TODO address the need to handle series of network packets, not just single packets + //TODO allow for the server to do more than just echo packets + //TODO document class + public sealed class StreamSocketServer : SocketServer { - private readonly ConcurrentDictionary<EndPoint, RemoteStreamClientToken> connectedClientTokens; - - private readonly struct RemoteStreamClientToken + private class RemoteStreamClientToken : IDisposable { - private readonly Channel<NetworkPacket> PacketChannel; + public readonly Socket ClientSocket; - public readonly ChannelReader<NetworkPacket> PacketReader; + public byte[]? RentedBuffer; - public readonly ChannelWriter<NetworkPacket> PacketWriter; + public RemoteStreamClientToken(in Socket clientSocket) + { + ClientSocket = clientSocket; + } - public RemoteStreamClientToken(in Channel<NetworkPacket> packetChannel) + public void Dispose() { - PacketChannel = packetChannel; - PacketReader = packetChannel.Reader; - PacketWriter = packetChannel.Writer; + ClientSocket.Dispose(); } } public readonly StreamSocketServerOptions ServerOptions; public StreamSocketServer(in AddressFamily connectionAddressFamily, in ProtocolType connectionProtocolType, - in StreamSocketServerOptions serverOptions = default) : base(in connectionAddressFamily, SocketType.Stream, + in StreamSocketServerOptions? serverOptions = null) : base(in connectionAddressFamily, SocketType.Stream, in connectionProtocolType) { - connectedClientTokens = new ConcurrentDictionary<EndPoint, RemoteStreamClientToken>(); - - ServerOptions = serverOptions.Equals(default) ? StreamSocketServerOptions.Defaults : serverOptions; + ServerOptions = serverOptions ?? StreamSocketServerOptions.Defaults; } protected override SocketAsyncEventArgs CreateTransmissionArgs() @@ -75,7 +77,7 @@ namespace NetSharp.Sockets.Stream protected override bool CanTransmissionArgsBeReused(in SocketAsyncEventArgs args) { - return false; + return true; } protected override void DestroyTransmissionArgs(SocketAsyncEventArgs remoteConnectionArgs) @@ -90,98 +92,206 @@ namespace NetSharp.Sockets.Stream protected override void HandleIoCompleted(object sender, SocketAsyncEventArgs args) { - throw new NotImplementedException(); + switch (args.LastOperation) + { + case SocketAsyncOperation.Accept: + SocketAsyncEventArgs newAcceptArgs = TransmissionArgsPool.Rent(); + + Accept(newAcceptArgs); // start a new accept operation to not miss any clients + + CompleteAccept(args); + break; + + case SocketAsyncOperation.Receive: + CompleteReceive(args); + break; + + case SocketAsyncOperation.Send: + CompleteSend(args); + break; + + default: + throw new NotSupportedException($"{nameof(HandleIoCompleted)} doesn't support {args.LastOperation}"); + } } - protected async Task HandleClient(SocketAsyncEventArgs clientArgs, CancellationToken cancellationToken = default) + private void Accept(SocketAsyncEventArgs acceptArgs) { - EndPoint clientEndPoint = clientArgs.AcceptSocket.RemoteEndPoint; - RemoteStreamClientToken clientToken = connectedClientTokens[clientEndPoint]; + bool operationPending = connection.AcceptAsync(acceptArgs); - Socket clientSocket = clientArgs.AcceptSocket; + if (!operationPending) + { + SocketAsyncEventArgs newAcceptArgs = TransmissionArgsPool.Rent(); + + Accept(newAcceptArgs); // start a new accept operation to not miss any clients + + CompleteAccept(acceptArgs); + } + } + + private void CompleteAccept(SocketAsyncEventArgs connectedClientArgs) + { + Socket clientSocket = connectedClientArgs.AcceptSocket; + + RemoteStreamClientToken clientToken = new RemoteStreamClientToken(in clientSocket); + + connectedClientArgs.UserToken = clientToken; + + Receive(connectedClientArgs); + } - byte[] requestBuffer = new byte[NetworkPacket.TotalSize]; + private void Receive(SocketAsyncEventArgs clientArgs) + { + RemoteStreamClientToken clientToken = (RemoteStreamClientToken) clientArgs.UserToken; + + byte[] requestBuffer = BufferPool.Rent(ServerOptions.PacketSize); Memory<byte> requestBufferMemory = new Memory<byte>(requestBuffer); - byte[] responseBuffer = new byte[NetworkPacket.TotalSize]; - Memory<byte> responseBufferMemory = new Memory<byte>(responseBuffer); + clientToken.RentedBuffer = requestBuffer; + clientArgs.SetBuffer(clientToken.RentedBuffer, 0, ServerOptions.PacketSize); - try + bool operationPending = clientToken.ClientSocket.ReceiveAsync(clientArgs); + + if (!operationPending) + { + CompleteReceive(clientArgs); + } + } + + private void CompleteReceive(SocketAsyncEventArgs clientArgs) + { + RemoteStreamClientToken receiveToken = (RemoteStreamClientToken) clientArgs.UserToken; + + if (clientArgs.SocketError == SocketError.Success) { - while (!cancellationToken.IsCancellationRequested) + if (clientArgs.BytesTransferred == ServerOptions.PacketSize) { - TransmissionResult receiveResult = - await SocketAsyncOperations - .ReceiveAsync(clientArgs, clientSocket, clientEndPoint, SocketFlags.None, - requestBufferMemory, cancellationToken) - .ConfigureAwait(false); - - if (receiveResult.Count == 0) - { - break; - } -#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(requestBufferMemory); - - // TODO implement actual request handling, besides just an echo + // buffer was fully received + + NetworkPacket request = NetworkPacket.Deserialise(receiveToken.RentedBuffer); + + // TODO implement actual request processing, not just an echo server NetworkPacket response = request; + byte[] responseBuffer = BufferPool.Rent(ServerOptions.PacketSize); + Memory<byte> responseBufferMemory = new Memory<byte>(responseBuffer); + NetworkPacket.Serialise(response, responseBufferMemory); - TransmissionResult sendResult = - await SocketAsyncOperations - .SendAsync(clientArgs, clientSocket, clientEndPoint, 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(receiveToken.RentedBuffer, true); // at this point the request buffer can be returned + + receiveToken.RentedBuffer = responseBuffer; + clientArgs.SetBuffer(responseBuffer, 0, ServerOptions.PacketSize); + + Send(clientArgs); + + } + else if (ServerOptions.PacketSize > clientArgs.BytesTransferred && clientArgs.BytesTransferred > 0) + { + // receive the remaining parts of the buffer + + int receivedBytes = clientArgs.BytesTransferred; + + clientArgs.SetBuffer(receivedBytes, ServerOptions.PacketSize - receivedBytes); + + Receive(clientArgs); + } + else + { + // no bytes were received, remote socket is dead + + CloseClientSocket(clientArgs); } } - catch (OperationCanceledException) + else { - Console.WriteLine($"Client task for {clientArgs.RemoteEndPoint} cancelled!"); + CloseClientSocket(clientArgs); } - finally + } + + private void Send(SocketAsyncEventArgs clientArgs) + { + RemoteStreamClientToken clientToken = (RemoteStreamClientToken) clientArgs.UserToken; + + bool operationPending = clientToken.ClientSocket.SendAsync(clientArgs); + + if (!operationPending) { - TransmissionArgsPool.Return(clientArgs); + CompleteSend(clientArgs); } } - public override async Task RunAsync(CancellationToken cancellationToken = default) + private void CompleteSend(SocketAsyncEventArgs clientArgs) { - connection.Listen(100); + RemoteStreamClientToken sendToken = (RemoteStreamClientToken) clientArgs.UserToken; - EndPoint remoteEndPoint = new IPEndPoint(IPAddress.Any, 0); + if (clientArgs.SocketError == SocketError.Success) + { + if (clientArgs.BytesTransferred == ServerOptions.PacketSize) + { + // buffer was fully sent + + BufferPool.Return(sendToken.RentedBuffer, true); + + sendToken.RentedBuffer = null; + + Receive(clientArgs); + } + else if (ServerOptions.PacketSize > clientArgs.BytesTransferred && clientArgs.BytesTransferred > 0) + { + // send the remaining parts of the buffer + + int sentBytes = clientArgs.BytesTransferred; - while (!cancellationToken.IsCancellationRequested) + clientArgs.SetBuffer(sentBytes, ServerOptions.PacketSize - sentBytes); + + Send(clientArgs); + } + else + { + // no bytes were sent, remote socket is dead + + CloseClientSocket(clientArgs); + } + } + else + { + CloseClientSocket(clientArgs); + } + } + + private void CloseClientSocket(SocketAsyncEventArgs clientArgs) + { + RemoteStreamClientToken clientToken = (RemoteStreamClientToken) clientArgs.UserToken; + + try + { + clientToken.ClientSocket.Shutdown(SocketShutdown.Both); + } + catch (SocketException ex) { - SocketAsyncEventArgs clientArgs = TransmissionArgsPool.Rent(); + Console.WriteLine(ex); + } - await SocketAsyncOperations.AcceptAsync(clientArgs, connection, cancellationToken); + clientToken.ClientSocket.Close(); - EndPoint clientEndPoint = clientArgs.AcceptSocket.RemoteEndPoint; + TransmissionArgsPool.Return(clientArgs); - BoundedChannelOptions clientChannelOptions = new BoundedChannelOptions(60) - { FullMode = BoundedChannelFullMode.DropOldest, SingleReader = true, SingleWriter = true }; - Channel<NetworkPacket> clientChannel = Channel.CreateBounded<NetworkPacket>(clientChannelOptions); + clientToken.Dispose(); + } - connectedClientTokens[clientEndPoint] = new RemoteStreamClientToken(in clientChannel); + public override async Task RunAsync(CancellationToken cancellationToken = default) + { + connection.Listen(100); - ConnectedClientHandlerTasks[clientEndPoint] = HandleClient(clientArgs, cancellationToken); + for (int i = 0; i < ServerOptions.ConcurrentAcceptCalls; i++) + { + SocketAsyncEventArgs acceptArgs = TransmissionArgsPool.Rent(); + + Accept(acceptArgs); } + + await cancellationToken.WaitHandle.WaitOneAsync(); } } } \ No newline at end of file diff --git a/NetSharp/NetSharpExamples/Program.cs b/NetSharp/NetSharpExamples/Program.cs @@ -1,5 +1,5 @@ #define TCP -#undef TCP +//#undef TCP using NetSharp.Sockets.Datagram; using NetSharp.Sockets.Stream; @@ -44,8 +44,8 @@ namespace NetSharpExamples private static async Task TestSocketClient() { - const int clientCount = 10; - const long packetsToSend = 1_000_000; + const int clientCount = 100; + const long packetsToSend = 100_000; Task[] clientTasks = new Task[clientCount]; double[] clientBandwidths = new double[clientCount]; @@ -104,7 +104,7 @@ namespace NetSharpExamples bandwidthStopwatch.Start(); #if TCP - int sendResult = client.SendBytes(requestBufferMemory); + TransmissionResult sendResult = client.Send(requestBuffer); #else TransmissionResult sendResult = client.SendTo(remoteEndPoint, requestBuffer); #endif @@ -124,7 +124,7 @@ namespace NetSharpExamples bandwidthStopwatch.Start(); #if TCP - int receiveResult = client.ReceiveBytes(responseBufferMemory); + TransmissionResult receiveResult = client.Receive(responseBuffer); #else TransmissionResult receiveResult = client.ReceiveFrom(ref remoteEndPoint, responseBuffer); #endif @@ -135,7 +135,7 @@ namespace NetSharpExamples #if DEBUG lock (typeof(Console)) { - Console.WriteLine($"[Client {id}, Packet {i}] Received {receiveResult.Count} bytes from {serverEndPoint}"); + Console.WriteLine($"[Client {id}, Packet {i}] Received {receiveResult.Count} bytes from {remoteEndPoint}"); Console.WriteLine($"[Client {id}, Packet {i}] <<<< {Encoding.UTF8.GetString(responseBufferMemory.Span)}"); } #endif @@ -171,7 +171,7 @@ namespace NetSharpExamples } #if TCP - client.Disconnect(); + client.Disconnect(true); client.Shutdown(SocketShutdown.Both); #endif