NetSharp

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

commit 95065dd32dd38dbce75b79d55dc3572ac47964de
parent d6c8b600ef67aa55e97011a2dbab601df9826607
Author: Mikolaj Lenczewski <mikolaj.lenczewski308@gmail.com>
Date:   Fri,  1 Jan 2021 14:56:07 +0000

Reimplemented RawStreamConnection
- Switched to Monitor for synchronisation
- Now successfully disposes when has active remote connections
- Passes the DisposesCleanly test
- Finished conversion to iteration over recursion

Diffstat:
MNetSharp/NetSharp.Examples/Examples/Raw Stream Examples/RawStreamClient.cs | 3+++
MNetSharp/NetSharp.Examples/Examples/Raw Stream Examples/RawStreamServer.cs | 3++-
MNetSharp/NetSharp.Tests/RawDatagramConnectionTests.cs | 2+-
MNetSharp/NetSharp.Tests/RawStreamConnectionTests.cs | 2+-
MNetSharp/NetSharp/NetSharp.xml | 53++++++++++++++++++++---------------------------------
MNetSharp/NetSharp/Raw/RawConnectionBase.cs | 17+++++++++++++----
MNetSharp/NetSharp/Raw/Stream/IRawStreamWriter.cs | 2+-
MNetSharp/NetSharp/Raw/Stream/RawStreamConnection.cs | 252++++++++++++++++++++++++++++++++++++++++++++++++++++++++-----------------------
8 files changed, 219 insertions(+), 115 deletions(-)

diff --git a/NetSharp/NetSharp.Examples/Examples/Raw Stream Examples/RawStreamClient.cs b/NetSharp/NetSharp.Examples/Examples/Raw Stream Examples/RawStreamClient.cs @@ -1,5 +1,6 @@ using System; using System.Net.Sockets; +using System.Threading; using System.Threading.Tasks; using NetSharp.Raw.Stream; @@ -31,6 +32,8 @@ namespace NetSharp.Examples.Examples.Raw_Stream_Examples sentBytes = client.SendAsync(0, packet).GetAwaiter().GetResult(); Console.WriteLine($"[Client] Sent {sentBytes} bytes to {client.RemoteEndPoint}"); + + Thread.Sleep(1000); } while (sentBytes > 0); client.DisconnectAsync().GetAwaiter().GetResult(); diff --git a/NetSharp/NetSharp.Examples/Examples/Raw Stream Examples/RawStreamServer.cs b/NetSharp/NetSharp.Examples/Examples/Raw Stream Examples/RawStreamServer.cs @@ -37,8 +37,9 @@ namespace NetSharp.Examples.Examples.Raw_Stream_Examples Console.WriteLine("[Server] Press enter to stop the server..."); Console.ReadLine(); - server.DeregisterHandler(0, PacketHandler); + Console.WriteLine("[Server] Attempting to stop the server..."); server.Close(); + Console.WriteLine("[Server] Successfully shut down the server!"); return Task.CompletedTask; } diff --git a/NetSharp/NetSharp.Tests/RawDatagramConnectionTests.cs b/NetSharp/NetSharp.Tests/RawDatagramConnectionTests.cs @@ -21,7 +21,7 @@ namespace NetSharp.Tests { using RawDatagramConnection conn = ConnectionFactory(); - conn.Bind(ServerLocalEndPoint); + conn.Bind(ClientLocalEndPoint); conn.Start(); conn.Close(); diff --git a/NetSharp/NetSharp.Tests/RawStreamConnectionTests.cs b/NetSharp/NetSharp.Tests/RawStreamConnectionTests.cs @@ -22,7 +22,7 @@ namespace NetSharp.Tests static void Instantiate() { using RawStreamConnection conn = ConnectionFactory(); - conn.Bind(ServerLocalEndPoint); + conn.Bind(ClientLocalEndPoint); conn.Start(); diff --git a/NetSharp/NetSharp/NetSharp.xml b/NetSharp/NetSharp/NetSharp.xml @@ -52,7 +52,7 @@ <member name="M:NetSharp.Raw.Datagram.RawDatagramConnection.ResetSocketArgsHook(System.Net.Sockets.SocketAsyncEventArgs@)"> <inheritdoc /> </member> - <member name="M:NetSharp.Raw.Datagram.RawDatagramConnection.HandlerThreadWork"> + <member name="M:NetSharp.Raw.Datagram.RawDatagramConnection.HandlerTaskWork"> <inheritdoc /> </member> <member name="T:NetSharp.Raw.RawConnectionBase"> @@ -147,7 +147,7 @@ Whether the <see cref="M:NetSharp.Raw.RawConnectionBase.Dispose" /> method was called. </param> </member> - <member name="M:NetSharp.Raw.RawConnectionBase.HandlerThreadWork"> + <member name="M:NetSharp.Raw.RawConnectionBase.HandlerTaskWork"> <summary> Handler work delegate, started when a call to <see cref="M:NetSharp.Raw.RawConnectionBase.Start(System.Int32)" /> is made. </summary> @@ -389,7 +389,7 @@ <member name="M:NetSharp.Raw.Stream.RawStreamConnection.Dispose(System.Boolean)"> <inheritdoc /> </member> - <member name="M:NetSharp.Raw.Stream.RawStreamConnection.HandlerThreadWork"> + <member name="M:NetSharp.Raw.Stream.RawStreamConnection.HandlerTaskWork"> <inheritdoc /> </member> <member name="M:NetSharp.Raw.Stream.RawStreamConnection.ResetSocketArgsHook(System.Net.Sockets.SocketAsyncEventArgs@)"> @@ -398,7 +398,7 @@ <member name="M:NetSharp.Raw.Stream.RawStreamConnection.StartHook(System.Int32)"> <inheritdoc /> </member> - <member name="M:NetSharp.Raw.Stream.RawStreamConnection.ConfigureSendRequestAsync(System.Net.Sockets.SocketAsyncEventArgs,System.Byte[]@,NetSharp.Raw.RawPacketHeader@,System.ReadOnlyMemory{System.Byte}@,NetSharp.Raw.Stream.RawStreamConnection.StateToken,System.Threading.Tasks.TaskCompletionSource{System.Int32})"> + <member name="M:NetSharp.Raw.Stream.RawStreamConnection.ConfigureSendRequestAsync(System.Net.Sockets.SocketAsyncEventArgs,System.Byte[]@,NetSharp.Raw.RawPacketHeader@,System.ReadOnlyMemory{System.Byte}@,NetSharp.Raw.Stream.RawStreamConnection.WriterStateToken,System.Threading.Tasks.TaskCompletionSource{System.Int32})"> <summary> Prepares the given socket args for sending a request to the network. </summary> @@ -408,17 +408,12 @@ Cleans up and returns the given socket args. </summary> </member> - <member name="M:NetSharp.Raw.Stream.RawStreamConnection.CloseClientConnection(System.Net.Sockets.SocketAsyncEventArgs)"> - <summary> - Closes the remote network connection associated with the given socket args. - </summary> - </member> - <member name="M:NetSharp.Raw.Stream.RawStreamConnection.ConfigureReceiveDataAsync(System.Net.Sockets.SocketAsyncEventArgs,NetSharp.Raw.Stream.RawStreamConnection.StateToken,NetSharp.Raw.RawPacketHeader@)"> + <member name="M:NetSharp.Raw.Stream.RawStreamConnection.ConfigureReceiveDataAsync(System.Net.Sockets.SocketAsyncEventArgs,NetSharp.Raw.Stream.RawStreamConnection.ReaderStateToken,NetSharp.Raw.RawPacketHeader@)"> <summary> Prepares the given socket args for receiving a packet's data from the network. </summary> </member> - <member name="M:NetSharp.Raw.Stream.RawStreamConnection.ConfigureReceiveHeaderAsync(System.Net.Sockets.SocketAsyncEventArgs,NetSharp.Raw.Stream.RawStreamConnection.StateToken)"> + <member name="M:NetSharp.Raw.Stream.RawStreamConnection.ConfigureReceiveHeaderAsync(System.Net.Sockets.SocketAsyncEventArgs,NetSharp.Raw.Stream.RawStreamConnection.ReaderStateToken)"> <summary> Prepares the given socket args for receiving a packet's header from the network. </summary> @@ -448,12 +443,12 @@ Handles the completion of a <see cref="M:System.Net.Sockets.Socket.ReceiveAsync(System.Net.Sockets.SocketAsyncEventArgs)" /> call. </summary> </member> - <member name="M:NetSharp.Raw.Stream.RawStreamConnection.HandleReceivedData(System.Net.Sockets.SocketAsyncEventArgs,NetSharp.Raw.Stream.RawStreamConnection.StateToken)"> + <member name="M:NetSharp.Raw.Stream.RawStreamConnection.HandleReceivedData(System.Net.Sockets.SocketAsyncEventArgs,NetSharp.Raw.Stream.RawStreamConnection.ReaderStateToken)"> <summary> Handles the completion of a <see cref="M:System.Net.Sockets.Socket.ReceiveAsync(System.Net.Sockets.SocketAsyncEventArgs)" /> call, when receiving a packet's data from the network. </summary> </member> - <member name="M:NetSharp.Raw.Stream.RawStreamConnection.HandleReceivedHeader(System.Net.Sockets.SocketAsyncEventArgs,NetSharp.Raw.Stream.RawStreamConnection.StateToken)"> + <member name="M:NetSharp.Raw.Stream.RawStreamConnection.HandleReceivedHeader(System.Net.Sockets.SocketAsyncEventArgs,NetSharp.Raw.Stream.RawStreamConnection.ReaderStateToken)"> <summary> Handles the completion of a <see cref="M:System.Net.Sockets.Socket.ReceiveAsync(System.Net.Sockets.SocketAsyncEventArgs)" /> call, when receiving a packet's header from the network. @@ -464,13 +459,13 @@ Handles the completion of a <see cref="M:System.Net.Sockets.Socket.SendAsync(System.Net.Sockets.SocketAsyncEventArgs)" /> call. </summary> </member> - <member name="M:NetSharp.Raw.Stream.RawStreamConnection.HandleSentRequest(System.Net.Sockets.SocketAsyncEventArgs,NetSharp.Raw.Stream.RawStreamConnection.StateToken)"> + <member name="M:NetSharp.Raw.Stream.RawStreamConnection.HandleSentRequest(System.Net.Sockets.SocketAsyncEventArgs,NetSharp.Raw.Stream.RawStreamConnection.WriterStateToken)"> <summary> Handles the completion of a <see cref="M:System.Net.Sockets.Socket.SendAsync(System.Net.Sockets.SocketAsyncEventArgs)" /> call, when sending a request packet to the network. In this case, the <see cref="P:System.Net.Sockets.SocketAsyncEventArgs.ConnectSocket" /> will be used to perform the transmission. </summary> </member> - <member name="M:NetSharp.Raw.Stream.RawStreamConnection.HandleSentResponse(System.Net.Sockets.SocketAsyncEventArgs,NetSharp.Raw.Stream.RawStreamConnection.StateToken)"> + <member name="M:NetSharp.Raw.Stream.RawStreamConnection.HandleSentResponse(System.Net.Sockets.SocketAsyncEventArgs,NetSharp.Raw.Stream.RawStreamConnection.ReaderStateToken)"> <summary> Handles the completion of a <see cref="M:System.Net.Sockets.Socket.SendAsync(System.Net.Sockets.SocketAsyncEventArgs)" /> call, when sending a response packet to the network. In this case, the <see cref="P:System.Net.Sockets.SocketAsyncEventArgs.AcceptSocket" /> will be used to perform the transmission. @@ -486,30 +481,22 @@ Starts or continues an asynchronous network write operation using the given socket. </summary> </member> - <member name="T:NetSharp.Raw.Stream.RawStreamConnection.StateToken"> + <member name="P:NetSharp.Raw.Stream.RawStreamConnection.OperationStateToken.OperationCompletionSource"> <summary> - State token for the stream network connection. + The <see cref="T:System.Threading.Tasks.TaskCompletionSource`1" /> for asynchronous network operations. </summary> </member> - <member name="P:NetSharp.Raw.Stream.RawStreamConnection.StateToken.BytesToTransfer"> - <summary> - The number of bytes that we need to transfer over the network. - </summary> + <member name="M:NetSharp.Raw.Stream.RawStreamConnection.OperationStateToken.Dispose"> + <inheritdoc /> </member> - <member name="P:NetSharp.Raw.Stream.RawStreamConnection.StateToken.OperationCompletionSource"> - <summary> - The <see cref="T:System.Threading.Tasks.TaskCompletionSource`1" /> for asynchronous network operations. - </summary> + <member name="M:NetSharp.Raw.Stream.RawStreamConnection.ReaderStateToken.Dispose"> + <inheritdoc /> </member> - <member name="P:NetSharp.Raw.Stream.RawStreamConnection.StateToken.RequestCompletionSource"> - <summary> - The <see cref="T:System.Threading.Tasks.TaskCompletionSource`1" /> for asynchronous packet writes. - </summary> + <member name="M:NetSharp.Raw.Stream.RawStreamConnection.RemoteConnectionWrapper.SendAsync(System.UInt16,System.ReadOnlyMemory{System.Byte},System.Net.Sockets.SocketFlags)"> + <inheritdoc /> </member> - <member name="P:NetSharp.Raw.Stream.RawStreamConnection.StateToken.RequestHeader"> - <summary> - The deserialised request packet header. - </summary> + <member name="M:NetSharp.Raw.Stream.RawStreamConnection.WriterStateToken.Dispose"> + <inheritdoc /> </member> <member name="T:NetSharp.Utils.SlimObjectPool`1"> <summary> diff --git a/NetSharp/NetSharp/Raw/RawConnectionBase.cs b/NetSharp/NetSharp/Raw/RawConnectionBase.cs @@ -18,7 +18,8 @@ namespace NetSharp.Raw private readonly Socket connection; private readonly SlimObjectPool<SocketAsyncEventArgs> socketArgsPool; - private int activeOperations; // TODO: consider removing since we have handler threads + private int activeOperations; + private readonly object activeOperationsLock = new object(); private bool isBound; private bool isDisposed; @@ -159,9 +160,13 @@ namespace NetSharp.Raw { connection.Dispose(); - // TODO: find better way of waiting for completion of all tasks - while (activeOperations > 0) + lock (activeOperationsLock) { + while (activeOperations > 0) + { + // will keep reacquiring the lock and blocking until we reach the activeOperations == 0 case + _ = Monitor.Wait(activeOperationsLock); + } } socketArgsPool.Dispose(); @@ -239,7 +244,11 @@ namespace NetSharp.Raw { socketArgsPool.Return(socketArgs); - _ = Interlocked.Decrement(ref activeOperations); + lock (activeOperationsLock) + { + _ = Interlocked.Decrement(ref activeOperations); + Monitor.Pulse(activeOperationsLock); // signals that we might have reached the activeOperations == 0 state + } } /// <summary> diff --git a/NetSharp/NetSharp/Raw/Stream/IRawStreamWriter.cs b/NetSharp/NetSharp/Raw/Stream/IRawStreamWriter.cs @@ -7,7 +7,7 @@ namespace NetSharp.Raw.Stream /// <summary> /// Describes the interface for a stream network connection that can write to the network. /// </summary> - public interface IRawStreamWriter : IRawStreamPacketHandler + public interface IRawStreamWriter { /// <summary> /// Writes the given packet header and data to the network asynchronously, using the given socket flags for the transmission. diff --git a/NetSharp/NetSharp/Raw/Stream/RawStreamConnection.cs b/NetSharp/NetSharp/Raw/Stream/RawStreamConnection.cs @@ -1,9 +1,11 @@ using System; using System.Collections.Concurrent; +using System.Collections.Generic; using System.Diagnostics; using System.Net; using System.Net.Sockets; using System.Runtime.CompilerServices; +using System.Threading; using System.Threading.Tasks; using NetSharp.Utils; @@ -35,13 +37,18 @@ namespace NetSharp.Raw.Stream /// Represents a network connection using a stream-based protocol to interact over the network, that is capable of /// sending raw bytes. /// </summary> - public sealed class RawStreamConnection : RawConnectionBase, IRawStreamWriter + public sealed class RawStreamConnection : RawConnectionBase, IRawStreamWriter, IRawStreamPacketHandler { + private readonly object activeRemoteConnectionsLock = new object(); private readonly SlimObjectPool<OperationStateToken> operationStatePool; private readonly SlimObjectPool<ReaderStateToken> readerStatePool; private readonly ConcurrentDictionary<int, RawStreamPacketHandler> registeredHandlers; + private readonly List<RemoteConnectionWrapper> remoteConnections; + private readonly object remoteConnectionsLock = new object(); private readonly SlimObjectPool<WriterStateToken> writerStatePool; + private int activeRemoteConnections; + /// <summary> /// Initialises a new instance of the <see cref="RawStreamConnection" /> class. /// </summary> @@ -54,6 +61,9 @@ namespace NetSharp.Raw.Stream public RawStreamConnection(ProtocolType connectionProtocolType, EndPoint defaultRemoteEndPoint) : base(SocketType.Stream, connectionProtocolType, defaultRemoteEndPoint) { + activeRemoteConnections = 0; + remoteConnections = new List<RemoteConnectionWrapper>(); + registeredHandlers = new ConcurrentDictionary<int, RawStreamPacketHandler>(); static OperationStateToken CreateOperationToken() @@ -203,21 +213,7 @@ namespace NetSharp.Raw.Stream /// <inheritdoc /> public ValueTask<int> SendAsync(ushort type, ReadOnlyMemory<byte> buffer, SocketFlags flags = SocketFlags.None) { - TaskCompletionSource<int> tcs = new TaskCompletionSource<int>(); - SocketAsyncEventArgs socketArgs = RentSocketArgs(); - - RawPacketHeader header = new RawPacketHeader(type, buffer.Length); - byte[] ownedBuffer = RentBuffer(RawPacket.TotalSize(in header)); - - WriterStateToken writerState = writerStatePool.Rent(); - - ConfigureSendRequestAsync(socketArgs, ref ownedBuffer, in header, in buffer, writerState, tcs); - - socketArgs.SocketFlags = flags; - - StartOrContinueSending(Connection, socketArgs); - - return new ValueTask<int>(tcs.Task); + return DoSendAsync(Connection, type, buffer, flags); } /// <inheritdoc /> @@ -254,14 +250,27 @@ namespace NetSharp.Raw.Stream return; } - base.Dispose(disposing); - if (disposing) { + foreach (RemoteConnectionWrapper remoteConnection in remoteConnections) + { + remoteConnection.LocalShutdown(); + } + + lock (activeRemoteConnectionsLock) + { + while (activeRemoteConnections > 0) + { + _ = Monitor.Wait(activeRemoteConnectionsLock); + } + } + operationStatePool.Dispose(); readerStatePool.Dispose(); writerStatePool.Dispose(); } + + base.Dispose(disposing); } /// <inheritdoc /> @@ -353,21 +362,6 @@ namespace NetSharp.Raw.Stream } /// <summary> - /// Closes the remote network connection associated with the given socket args. - /// </summary> - private void CloseClientConnection(SocketAsyncEventArgs args) - { - Socket connection = args.AcceptSocket; - - connection.Disconnect(false); - connection.Shutdown(SocketShutdown.Both); - connection.Close(); - connection.Dispose(); - - CleanupArgs(args); - } - - /// <summary> /// Prepares the given socket args for receiving a packet's data from the network. /// </summary> private void ConfigureReceiveDataAsync(SocketAsyncEventArgs args, ReaderStateToken readerState, in RawPacketHeader header) @@ -400,6 +394,25 @@ namespace NetSharp.Raw.Stream args.UserToken = readerState; } + private ValueTask<int> DoSendAsync(Socket connection, ushort type, ReadOnlyMemory<byte> buffer, SocketFlags flags) + { + TaskCompletionSource<int> tcs = new TaskCompletionSource<int>(); + SocketAsyncEventArgs socketArgs = RentSocketArgs(); + + RawPacketHeader header = new RawPacketHeader(type, buffer.Length); + byte[] ownedBuffer = RentBuffer(RawPacket.TotalSize(in header)); + + WriterStateToken writerState = writerStatePool.Rent(); + + ConfigureSendRequestAsync(socketArgs, ref ownedBuffer, in header, in buffer, writerState, tcs); + + socketArgs.SocketFlags = flags; + + StartOrContinueSending(connection, socketArgs); + + return new ValueTask<int>(tcs.Task); + } + /// <summary> /// Handles a completed <see cref="Socket.AcceptAsync" /> call. /// </summary> @@ -408,6 +421,8 @@ namespace NetSharp.Raw.Stream switch (args.SocketError) { case SocketError.Success: + _ = Interlocked.Increment(ref activeRemoteConnections); + // the buffer is set to allow a simpler ConfigureReceiveHeader() implementation. Since returning an // empty buffer is ignored in the array pool, this allows us to just return the last assigned buffer // in the ConfigureXXX() method to the pool (this means that usually we will usually be returning @@ -416,9 +431,70 @@ namespace NetSharp.Raw.Stream ReaderStateToken readerState = readerStatePool.Rent(); - // TODO convert into iteration instead of recursion - ConfigureReceiveHeaderAsync(args, readerState); - StartOrContinueReceiving(args); + Socket remoteConnection = args.AcceptSocket; + RemoteConnectionWrapper wrapper = new RemoteConnectionWrapper(this, remoteConnection, readerState); + + lock (remoteConnectionsLock) + { + remoteConnections.Add(wrapper); + } + + WaitHandle[] eventHandles = + { + readerState.ConnectionClosed.WaitHandle, + readerState.RequestReceived.WaitHandle, + }; + + while (true) + { + ConfigureReceiveHeaderAsync(args, readerState); + StartOrContinueReceiving(args); + + int completedEvent = WaitHandle.WaitAny(eventHandles); + + if (completedEvent == 0) + { + // the connection has been closed + break; + } + + // since the request received event has been set, we need to reset it + readerState.RequestReceived.Reset(); + + RawPacketHeader header = readerState.RequestHeader!.Value; + + byte[] dataBuffer = args.Buffer; + ReadOnlyMemory<byte> dataBufferMemory = new ReadOnlyMemory<byte>(dataBuffer, 0, header.DataLength); + + if (registeredHandlers.TryGetValue(header.Type, out RawStreamPacketHandler handler)) + { + handler.Invoke(remoteConnection.RemoteEndPoint, in header, in dataBufferMemory, wrapper); + } + } + + lock (activeRemoteConnectionsLock) + { + _ = Interlocked.Decrement(ref activeRemoteConnections); + Monitor.Pulse(activeRemoteConnectionsLock); // signals that we may have reached the activeRemoteConnections == 0 state + } + + lock (remoteConnectionsLock) + { + // TODO: ensure that this is atomic via locking or something else. are properties inherently atomic??? + if (!readerState.LocalShutdownSignaled) + { + // since we were shutdown remotely, we need to remove the remote connection from the list of tracked connections + // TODO: come up with a better way of tracking active remote connections (have the wrapper deregister itself?) + _ = remoteConnections.Remove(wrapper); + + remoteConnection.Disconnect(false); + remoteConnection.Shutdown(SocketShutdown.Both); + remoteConnection.Close(); + remoteConnection.Dispose(); + } + } + + CleanupArgs(args); break; default: @@ -545,8 +621,7 @@ namespace NetSharp.Raw.Stream break; default: - // TODO break out of iteration in HandleAccepted - CloseClientConnection(args); + readerState.ConnectionClosed.Set(); break; } } @@ -561,21 +636,9 @@ namespace NetSharp.Raw.Stream int totalReceived = previouslyReceived + received; int expected = readerState.BytesToTransfer; - RawPacketHeader header = readerState.RequestHeader!.Value; - - byte[] dataBuffer = args.Buffer; - ReadOnlyMemory<byte> dataBufferMemory = new ReadOnlyMemory<byte>(dataBuffer, 0, header.DataLength); - if (totalReceived == expected) { - // TODO switch out recursion for iteration in HandleAccepted - if (registeredHandlers.TryGetValue(header.Type, out RawStreamPacketHandler handler)) - { - handler.Invoke(args.AcceptSocket.RemoteEndPoint, in header, in dataBufferMemory, this); - } - - ConfigureReceiveHeaderAsync(args, readerState); - StartOrContinueReceiving(args); + readerState.RequestReceived.Set(); } else if (totalReceived > 0 && totalReceived < expected) { @@ -584,8 +647,7 @@ namespace NetSharp.Raw.Stream } else if (received == 0) { - // TODO break out of iteration in HandleAccepted - CloseClientConnection(args); + readerState.ConnectionClosed.Set(); } } @@ -617,8 +679,7 @@ namespace NetSharp.Raw.Stream } else if (received == 0) { - // TODO break out of iteration in HandleAccepted - CloseClientConnection(args); + readerState.ConnectionClosed.Set(); } } @@ -637,8 +698,7 @@ namespace NetSharp.Raw.Stream break; default: - // TODO break out of iteration in HandleAccepted - CloseClientConnection(args); + readerState.ConnectionClosed.Set(); break; } @@ -711,9 +771,7 @@ namespace NetSharp.Raw.Stream if (totalSent == expected) { - // TODO switch out recursion for iteration in HandleAccepted - ConfigureReceiveHeaderAsync(args, readerState); - StartOrContinueReceiving(args); + // TODO do we need to do any notification that a send has completed? } else if (totalSent > 0 && totalSent < expected) { @@ -722,8 +780,7 @@ namespace NetSharp.Raw.Stream } else if (sent == 0) { - // TODO break out of iteration in HandleAccepted - CloseClientConnection(args); + readerState.ConnectionClosed.Set(); } } @@ -776,39 +833,86 @@ namespace NetSharp.Raw.Stream private sealed class ReaderStateToken : IDisposable { - /// <summary> - /// The number of bytes that we need to transfer over the network. - /// </summary> internal int BytesToTransfer { get; set; } - /// <summary> - /// The deserialised request packet header. - /// </summary> + internal ManualResetEventSlim ConnectionClosed { get; } = new ManualResetEventSlim(false); + + internal bool LocalShutdownSignaled { get; set; } + internal RawPacketHeader? RequestHeader { get; set; } + internal ManualResetEventSlim RequestReceived { get; } = new ManualResetEventSlim(false); + /// <inheritdoc /> public void Dispose() { Reset(); + + ConnectionClosed.Dispose(); + RequestReceived.Dispose(); } internal void Reset() { BytesToTransfer = 0; RequestHeader = null; + + ConnectionClosed.Reset(); + RequestReceived.Reset(); + } + } + + private sealed class RemoteConnectionWrapper : IRawStreamWriter + { + private readonly WeakReference<Socket> connectionRef; + private readonly WeakReference<RawStreamConnection> parentRef; + private readonly WeakReference<ReaderStateToken> tokenRef; + + internal RemoteConnectionWrapper(RawStreamConnection parent, Socket connection, ReaderStateToken token) + { + parentRef = new WeakReference<RawStreamConnection>(parent); + connectionRef = new WeakReference<Socket>(connection); + tokenRef = new WeakReference<ReaderStateToken>(token); + } + + public void LocalShutdown() + { + bool connectionDisposed = !connectionRef.TryGetTarget(out Socket connection); + bool tokenDisposed = !tokenRef.TryGetTarget(out ReaderStateToken token); + + if (connectionDisposed || tokenDisposed) + { + return; + } + + token.LocalShutdownSignaled = true; + + connection.Disconnect(false); + connection.Shutdown(SocketShutdown.Both); + connection.Close(); + connection.Dispose(); + } + + /// <inheritdoc /> + public ValueTask<int> SendAsync(ushort type, ReadOnlyMemory<byte> buffer, SocketFlags flags = SocketFlags.None) + { + bool parentDisposed = !parentRef.TryGetTarget(out RawStreamConnection parent); + bool connectionDisposed = !connectionRef.TryGetTarget(out Socket connection); + + if (parentDisposed || connectionDisposed) + { + // TODO: just use a result of -1? some other error code? + throw new ObjectDisposedException(parentDisposed ? nameof(parent) : nameof(connection)); + } + + return parent.DoSendAsync(connection, type, buffer, flags); } } private sealed class WriterStateToken : IDisposable { - /// <summary> - /// The number of bytes that we need to transfer over the network. - /// </summary> internal int BytesToTransfer { get; set; } - /// <summary> - /// The <see cref="TaskCompletionSource{TResult}" /> for asynchronous packet writes. - /// </summary> internal TaskCompletionSource<int>? RequestCompletionSource { get; set; } /// <inheritdoc />