diff --git a/mROA/Implementation/Backend/NetworkGatewayModule.cs b/mROA/Implementation/Backend/NetworkGatewayModule.cs index 1fe4169..b01786c 100644 --- a/mROA/Implementation/Backend/NetworkGatewayModule.cs +++ b/mROA/Implementation/Backend/NetworkGatewayModule.cs @@ -1,4 +1,5 @@ using System; +using System.Collections.Generic; using System.Net; using System.Net.Sockets; using System.Threading; @@ -15,6 +16,7 @@ namespace mROA.Implementation.Backend private readonly TcpListener _tcpListener; private IConnectionHub? _hub; private ISerializationToolkit? _serialization; + private Dictionary _extractorsCTS = new(); public NetworkGatewayModule(IPEndPoint endpoint, Type interactionModuleType, IInjectableModule[] injectableModules) @@ -82,14 +84,16 @@ namespace mROA.Implementation.Backend streamExtractor.SingleReceive(); var connectionRequest = interaction.GetNextMessageReceiving(false) .GetAwaiter().GetResult()!; + var cts = new CancellationTokenSource(); switch (connectionRequest.MessageType) { case EMessageType.ClientConnect: - Task.Run(async () => await streamExtractor.LoopedReceive()); - streamExtractor.SendFromChannel(interaction.TrustedPostChanel); + Task.Run(async () => await streamExtractor.LoopedReceive(cts.Token)); + _ = streamExtractor.SendFromChannel(interaction.TrustedPostChanel, cts.Token); interaction.PostMessageAsync(new NetworkMessageHeader(_serialization!, new IdAssignment { Id = -interaction.ConnectionId })); + _extractorsCTS[interaction.ConnectionId] = cts; _hub!.RegisterInteraction(interaction); Console.WriteLine("Client registered"); break; @@ -97,13 +101,17 @@ namespace mROA.Implementation.Backend { var recoveryRequest = _serialization!.Deserialize(connectionRequest.Data)!; var recoveryInteraction = _hub.GetInteraction(recoveryRequest.Id); - + + _extractorsCTS[recoveryRequest.Id].Cancel(); + recoveryInteraction.IsConnected = () => streamExtractor.IsConnected; streamExtractor.MessageReceived = message => { recoveryInteraction.ReceiveChanel.Writer.WriteAsync(message); }; - Task.Run(async () => await streamExtractor.LoopedReceive()); + _ = streamExtractor.SendFromChannel(recoveryInteraction.TrustedPostChanel, cts.Token); + + Task.Run(async () => await streamExtractor.LoopedReceive(cts.Token)); recoveryInteraction.Restart(false); diff --git a/mROA/Implementation/ChannelInteractionModule.cs b/mROA/Implementation/ChannelInteractionModule.cs index e706dfb..4657c62 100644 --- a/mROA/Implementation/ChannelInteractionModule.cs +++ b/mROA/Implementation/ChannelInteractionModule.cs @@ -53,8 +53,7 @@ namespace mROA.Implementation public ChannelReader TrustedPostChanel => _outputTrustedChannel.Reader; public ChannelReader UntrustedPostChanel => _outputUntrustedChannel.Reader; public Func IsConnected { get; set; } - - + public void Inject(T dependency) { switch (dependency) diff --git a/mROA/mROA.csproj b/mROA/mROA.csproj index e52150a..c0682c2 100644 --- a/mROA/mROA.csproj +++ b/mROA/mROA.csproj @@ -21,7 +21,7 @@ - TRACE; + ;