diff --git a/mROA/Abstract/IInteractionModule.cs b/mROA/Abstract/IInteractionModule.cs index cffade7..4a884e8 100644 --- a/mROA/Abstract/IInteractionModule.cs +++ b/mROA/Abstract/IInteractionModule.cs @@ -12,7 +12,7 @@ namespace mROA.Abstract ChannelWriter ReceiveChanel { get; } ChannelReader TrustedPostChanel { get; } ChannelReader UntrustedPostChanel { get; } - Action IsConnected { get; set; } + Func IsConnected { get; set; } Task GetNextMessageReceiving(bool infinite = true); Task PostMessageAsync(NetworkMessageHeader messageHeader); Task PostMessageUntrustedAsync(NetworkMessageHeader messageHeader); diff --git a/mROA/Implementation/Backend/NetworkGatewayModule.cs b/mROA/Implementation/Backend/NetworkGatewayModule.cs index 4c3f22b..52ac70c 100644 --- a/mROA/Implementation/Backend/NetworkGatewayModule.cs +++ b/mROA/Implementation/Backend/NetworkGatewayModule.cs @@ -1,6 +1,7 @@ using System; using System.Net; using System.Net.Sockets; +using System.Threading; using System.Threading.Channels; using System.Threading.Tasks; using mROA.Abstract; @@ -73,21 +74,20 @@ namespace mROA.Implementation.Backend interaction!.Inject(injectableModule); interaction!.Inject(_serialization); - interaction.BaseStream = client.GetStream(); - var channel = Channel.CreateUnbounded(new UnboundedChannelOptions - { - SingleWriter = false, - SingleReader = false, - AllowSynchronousContinuations = true - }); - interaction.UntrustedReceiveChanel = channel.Reader; - interaction.UntrustedReceiveChanelWriter = channel.Writer; + + + var streamExtractor = new StreamExtractor(client.GetStream(), _serialization); + interaction.IsConnected = () => streamExtractor.IsConnected; + streamExtractor.MessageReceived += message => interaction.ReceiveChanel.WriteAsync(message); + streamExtractor.SingleReceive(); var connectionRequest = interaction.GetNextMessageReceiving(false) .GetAwaiter().GetResult()!; - + switch (connectionRequest.MessageType) { case EMessageType.ClientConnect: + Task.Run(async () => await streamExtractor.LoopedReceive()); + streamExtractor.SendFromChannel(interaction.TrustedPostChanel); interaction.PostMessageAsync(new NetworkMessageHeader(_serialization!, new IdAssignment { Id = -interaction.ConnectionId })); _hub!.RegisterInteraction(interaction); @@ -95,19 +95,13 @@ namespace mROA.Implementation.Backend break; case EMessageType.ClientRecovery: { - interaction.BaseStream = null; var recoveryRequest = _serialization!.Deserialize(connectionRequest.Data)!; var recoveryInteraction = _hub.GetInteraction(recoveryRequest.Id); - recoveryInteraction.UntrustedReceiveChanel = - Channel.CreateUnbounded(new UnboundedChannelOptions - { - SingleWriter = false, - SingleReader = false, - AllowSynchronousContinuations = true, - - }).Reader; - recoveryInteraction.BaseStream = client.GetStream(); + + streamExtractor = new StreamExtractor(client.GetStream(), _serialization); + streamExtractor.MessageReceived += message => recoveryInteraction.ReceiveChanel.WriteAsync(message); + _ = streamExtractor.LoopedReceive(); recoveryInteraction.Restart(false); Console.WriteLine("Connection recovery for client {0} finished", recoveryRequest.Id); break; @@ -118,7 +112,7 @@ namespace mROA.Implementation.Backend } } } - + private void ThrowIfNotInjected() { if (_hub is null) diff --git a/mROA/Implementation/ChannelInteractionModule.cs b/mROA/Implementation/ChannelInteractionModule.cs index e5d51ba..c3f637b 100644 --- a/mROA/Implementation/ChannelInteractionModule.cs +++ b/mROA/Implementation/ChannelInteractionModule.cs @@ -12,22 +12,20 @@ namespace mROA.Implementation public class ChannelInteractionModule : IChannelInteractionModule { private readonly ChannelReader _receiveReader; + private readonly ChannelWriter _trustedWriter; + private readonly ChannelWriter _untrustedWriter; private readonly Channel _inputChannel; private readonly Channel _outputTrustedChannel; private readonly Channel _outputUntrustedChannel; - private const int BufferSize = ushort.MaxValue; - private readonly Memory _buffer = new byte[BufferSize]; private readonly List _messageBuffer = new(128); private Task? _currentReceiving; private ISerializationToolkit? _serialization; private bool _isConnected = true; - private bool _isInReconnectionState; private bool _isActive = true; private TaskCompletionSource _reconnection; public ChannelInteractionModule() { - _reconnection = new TaskCompletionSource(); _inputChannel = Channel.CreateUnbounded(new UnboundedChannelOptions { SingleReader = false, @@ -35,18 +33,21 @@ namespace mROA.Implementation AllowSynchronousContinuations = true }); _receiveReader = _inputChannel.Reader; - _outputTrustedChannel = Channel.CreateUnbounded(new UnboundedChannelOptions + _outputTrustedChannel = Channel.CreateBounded(new BoundedChannelOptions(1) { SingleReader = true, SingleWriter = true, - AllowSynchronousContinuations = true + AllowSynchronousContinuations = true, + }); + _trustedWriter = _outputTrustedChannel.Writer; _outputUntrustedChannel = Channel.CreateUnbounded(new UnboundedChannelOptions { SingleReader = true, SingleWriter = true, AllowSynchronousContinuations = true }); + _untrustedWriter = _outputUntrustedChannel.Writer; } public int ConnectionId { get; set; } @@ -54,7 +55,7 @@ namespace mROA.Implementation public ChannelWriter ReceiveChanel => _inputChannel.Writer; public ChannelReader TrustedPostChanel => _outputTrustedChannel.Reader; public ChannelReader UntrustedPostChanel => _outputUntrustedChannel.Reader; - public Action IsConnected { get; set; } + public Func IsConnected { get; set; } public void Inject(T dependency) @@ -72,7 +73,7 @@ namespace mROA.Implementation public Task GetNextMessageReceiving(bool infinite = true) { - if (!infinite) return Receive().AsTask(); + if (!infinite) return _receiveReader.ReadAsync().AsTask(); if (_currentReceiving != null) return _currentReceiving; _currentReceiving = Task.Run(async () => await GetNextMessage()); return _currentReceiving; @@ -80,19 +81,12 @@ namespace mROA.Implementation #pragma warning disable CS8602 // Dereference of a possibly null reference. private async ValueTask PostMessageInternal(NetworkMessageHeader messageHeader) { -#if TRACE - Console.WriteLine( - $"{DateTime.Now.TimeOfDay} Posting message: {messageHeader.Id} - {messageHeader.MessageType} to {ConnectionId}"); -#endif - - var rawMessage = _serialization.Serialize(messageHeader); - var header = BitConverter.GetBytes((ushort)rawMessage.Length).AsMemory(0, sizeof(ushort)); - - if (!_baseStream.CanWrite) + if (!IsConnected()) + { return false; - - await BaseStream.WriteAsync(header); - await BaseStream.WriteAsync(rawMessage); + } + + await _trustedWriter.WriteAsync(messageHeader); return true; } #pragma warning restore CS8602 // Dereference of a possibly null reference. @@ -100,9 +94,6 @@ namespace mROA.Implementation public async Task PostMessageAsync(NetworkMessageHeader messageHeader) { - if (BaseStream == null) - throw new NullReferenceException("BaseStream is null"); - if (_serialization == null) throw new NullReferenceException("Serialization toolkit is not initialized"); @@ -132,7 +123,7 @@ namespace mROA.Implementation public async Task PostMessageUntrustedAsync(NetworkMessageHeader messageHeader) { - await UntrustedPostChanel.WriteAsync((ConnectionId, messageHeader)); + await _untrustedWriter.WriteAsync(messageHeader); } public void HandleMessage(NetworkMessageHeader messageHeader) @@ -151,9 +142,6 @@ namespace mROA.Implementation private async Task GetNextMessage() { - if (BaseStream == null) - throw new NullReferenceException("BaseStream is null"); - if (_serialization == null) throw new NullReferenceException("Serialization toolkit is null"); @@ -184,65 +172,17 @@ namespace mROA.Implementation } } - private ushort ReadMessageLength() - { - var firstBit = BaseStream.ReadByte(); - if (firstBit == -1) - { - _isConnected = false; - throw new EndOfStreamException(); - } - - _isConnected = true; - var secondBit = (byte)BaseStream.ReadByte(); - - var len = BitConverter.ToUInt16(new[] { (byte)firstBit, secondBit }); - - return len; - } - - private async ValueTask Receive() - { - var len = ReadMessageLength(); - var localSpan = _buffer[..len]; - - await BaseStream.ReadExactlyAsync(localSpan); - - var message = _serialization.Deserialize(localSpan.Span); -#if TRACE - Console.WriteLine($"{DateTime.Now.TimeOfDay} Received Message {message.Id} - {message.MessageType}"); - TransmissionConfig.TotalTransmittedBytes += len; - Console.WriteLine($"Total received bytes are {TransmissionConfig.TotalTransmittedBytes}"); -#endif - _messageBuffer.Add(message); - return message; - } - public async Task Restart(bool sendRecovery) { if (sendRecovery) { await PostMessageAsync( new NetworkMessageHeader(_serialization!, new ClientRecovery(Math.Abs(ConnectionId)))); - var iTest = _baseStream.ReadByte(); - var bTest = (byte)iTest; - _baseStream.WriteByte(bTest); - } - else - { - const byte confirmByte = 128; - _baseStream.WriteByte(confirmByte); - var iPong = _baseStream.ReadByte(); - var bPong = (byte)iPong; - if (confirmByte != bPong) - { - Console.WriteLine("Incorrect byte"); - } } Console.WriteLine("Setting result for reconnection"); - var setting = _reconnection.TrySetResult(BaseStream); - _isInReconnectionState = false; + var setting = _reconnection.TrySetResult(null); + // _isInReconnectionState = false; _isConnected = true; Console.WriteLine($"Set result for reconnection {setting}"); @@ -251,29 +191,30 @@ namespace mROA.Implementation private async Task MakeRecovery(string source) { - Console.WriteLine("Staring recovery from {0}", source); - - lock (_reconnection) - { - Console.WriteLine("Got lock from {0}", source); - - Console.WriteLine("Call OnDisconnected from {0}", source); - _isInReconnectionState = true; - OnDisconnected?.Invoke(ConnectionId); - } - - Console.WriteLine("Waiting for reconnect from {0}", source); - if (!_reconnection.Task.IsCompleted && !_isConnected) - { - Console.WriteLine("Current connection state {0} from {1}", _isConnected, source); - await _reconnection.Task; - } - - Console.WriteLine("Reconnect finished from {0}", source); - lock (_reconnection) - { - _isInReconnectionState = false; - } + //TODO переделать реконнект + // Console.WriteLine("Staring recovery from {0}", source); + // + // lock (_reconnection) + // { + // Console.WriteLine("Got lock from {0}", source); + // + // Console.WriteLine("Call OnDisconnected from {0}", source); + // _isInReconnectionState = true; + // OnDisconnected?.Invoke(ConnectionId); + // } + // + // Console.WriteLine("Waiting for reconnect from {0}", source); + // if (!_reconnection.Task.IsCompleted && !_isConnected) + // { + // Console.WriteLine("Current connection state {0} from {1}", _isConnected, source); + // await _reconnection.Task; + // } + // + // Console.WriteLine("Reconnect finished from {0}", source); + // lock (_reconnection) + // { + // _isInReconnectionState = false; + // } } public void Dispose() diff --git a/mROA/Implementation/Frontend/NetworkFrontendBridge.cs b/mROA/Implementation/Frontend/NetworkFrontendBridge.cs index 2a5ec69..5c48b40 100644 --- a/mROA/Implementation/Frontend/NetworkFrontendBridge.cs +++ b/mROA/Implementation/Frontend/NetworkFrontendBridge.cs @@ -1,6 +1,7 @@ using System; using System.Net; using System.Net.Sockets; +using System.Threading; using System.Threading.Channels; using System.Threading.Tasks; using mROA.Abstract; @@ -14,10 +15,13 @@ namespace mROA.Implementation.Frontend private TcpClient _tcpClient = new(); private IChannelInteractionModule? _interactionModule; private ISerializationToolkit? _serialization; + private StreamExtractor _currentExtractor; + private CancellationTokenSource _rawExtractorCancellation; public NetworkFrontendBridge(IPEndPoint serverEndPoint) { _serverEndPoint = serverEndPoint; + _rawExtractorCancellation = new CancellationTokenSource(); } public void Inject(T dependency) @@ -42,19 +46,15 @@ namespace mROA.Implementation.Frontend _tcpClient.Connect(_serverEndPoint); - _interactionModule.BaseStream = _tcpClient.GetStream(); - var channel = Channel.CreateUnbounded(new UnboundedChannelOptions - { - SingleWriter = false, - SingleReader = false, - AllowSynchronousContinuations = true - }); - _interactionModule.UntrustedReceiveChanel = channel.Reader; - _interactionModule.UntrustedReceiveChanelWriter = channel.Writer; + PrepareExtractor(); + _interactionModule.IsConnected = () => _currentExtractor.IsConnected; _interactionModule.OnDisconnected += id => { Reconnect(); }; _interactionModule.PostMessageAsync(new NetworkMessageHeader(_serialization, new ClientConnect())).Wait(); + + _currentExtractor.SingleReceive(); var idMessage = _interactionModule.GetNextMessageReceiving(false).GetAwaiter().GetResult(); + if (idMessage.MessageType != EMessageType.IdAssigning) { throw new Exception( @@ -62,30 +62,40 @@ namespace mROA.Implementation.Frontend } + var stopToken = _rawExtractorCancellation.Token; + Task.Run(async () => await _currentExtractor.LoopedReceive(stopToken)); + var assignment = _serialization.Deserialize(idMessage.Data)!; _interactionModule.ConnectionId = -assignment.Id; TransmissionConfig.OwnershipRepository = new StaticOwnershipRepository(assignment.Id); } + private void PrepareExtractor() + { + _currentExtractor = new StreamExtractor(_tcpClient.GetStream(), _serialization); + + _ = _currentExtractor.SendFromChannel(_interactionModule.TrustedPostChanel, + _rawExtractorCancellation.Token); + _currentExtractor.MessageReceived += message => _interactionModule.ReceiveChanel.WriteAsync(message); + } + private async Task Reconnect() { _tcpClient = new TcpClient(); _tcpClient.Connect(_serverEndPoint); - _interactionModule.BaseStream = _tcpClient.GetStream(); - var channel = Channel.CreateUnbounded(new UnboundedChannelOptions - { - SingleWriter = false, - SingleReader = false, - AllowSynchronousContinuations = true - }); - _interactionModule.UntrustedReceiveChanel = channel.Reader; - _interactionModule.UntrustedReceiveChanelWriter = channel.Writer; + + _rawExtractorCancellation.Cancel(); + _rawExtractorCancellation = new CancellationTokenSource(); + + PrepareExtractor(); + + _ = _currentExtractor.LoopedReceive(_rawExtractorCancellation.Token); + await _interactionModule.Restart(true); } public void Obstacle() { - _interactionModule!.BaseStream!.Dispose(); _tcpClient.Dispose(); } diff --git a/mROA/Implementation/NetworkMessageHeader.cs b/mROA/Implementation/NetworkMessageHeader.cs index a75cdb0..8e6c643 100644 --- a/mROA/Implementation/NetworkMessageHeader.cs +++ b/mROA/Implementation/NetworkMessageHeader.cs @@ -8,6 +8,24 @@ namespace mROA.Implementation { public class NetworkMessageHeader { + protected bool Equals(NetworkMessageHeader other) + { + return Id.Equals(other.Id) && MessageType == other.MessageType; + } + + public override bool Equals(object? obj) + { + if (obj is null) return false; + if (ReferenceEquals(this, obj)) return true; + if (obj.GetType() != GetType()) return false; + return Equals((NetworkMessageHeader)obj); + } + + public override int GetHashCode() + { + return HashCode.Combine(Id, (int)MessageType); + } + public static readonly NetworkMessageHeader Null = new(); public NetworkMessageHeader() { diff --git a/mROA/Implementation/StreamExtractor.cs b/mROA/Implementation/StreamExtractor.cs index f9f71c5..53d345d 100644 --- a/mROA/Implementation/StreamExtractor.cs +++ b/mROA/Implementation/StreamExtractor.cs @@ -1,6 +1,7 @@ using System; using System.IO; using System.Threading; +using System.Threading.Channels; using System.Threading.Tasks; using mROA.Abstract; @@ -13,6 +14,7 @@ namespace mROA.Implementation private const int BufferSize = ushort.MaxValue; private readonly Memory _buffer = new byte[BufferSize]; private bool _manualConnectionState = true; + public StreamExtractor(Stream ioStream, ISerializationToolkit serializationToolkit) { _ioStream = ioStream; @@ -38,12 +40,12 @@ namespace mROA.Implementation return len; } - public async Task SingleReceive() + public async Task SingleReceive(CancellationToken Token = default) { var len = ReadMessageLength(); var localSpan = _buffer[..len]; - await _ioStream.ReadExactlyAsync(localSpan); + await _ioStream.ReadExactlyAsync(localSpan, cancellationToken: Token); var message = _serializationToolkit.Deserialize(localSpan.Span); #if TRACE @@ -54,23 +56,40 @@ namespace mROA.Implementation MessageReceived(message); } - public async Task InfiniteReceive(CancellationToken token) + public async Task LoopedReceive(CancellationToken token = default) { - while (token.IsCancellationRequested == false) + while (token.IsCancellationRequested == false && IsConnected) { - await SingleReceive(); + await SingleReceive(token); } } - public async Task Send(NetworkMessageHeader message) + public async Task Send(NetworkMessageHeader message, CancellationToken token = default) { var rawMessage = _serializationToolkit.Serialize(message); var header = BitConverter.GetBytes((ushort)rawMessage.Length).AsMemory(0, sizeof(ushort)); - await _ioStream.WriteAsync(header); - await _ioStream.WriteAsync(rawMessage); +#if TRACE + Console.WriteLine($"{DateTime.Now.TimeOfDay} Posting Message {message.Id} - {message.MessageType}"); + TransmissionConfig.TotalTransmittedBytes += rawMessage.Length; + Console.WriteLine($"Total received bytes are {TransmissionConfig.TotalTransmittedBytes}"); +#endif + + await _ioStream.WriteAsync(header, token); + await _ioStream.WriteAsync(rawMessage, token); } + public async Task SendFromChannel(ChannelReader channel, + CancellationToken token = default) + { + while (token.IsCancellationRequested == false && IsConnected) + { + var message = await channel.ReadAsync(token); + await Send(message, token); + } + } + + public bool IsConnected => _ioStream is { CanRead: true, CanWrite: true } && _manualConnectionState; } } \ No newline at end of file