From 3d3772a341fd8166eeb1941239b1038b9286370b Mon Sep 17 00:00:00 2001 From: Mikhail Mitrofanov Date: Sat, 12 Apr 2025 12:44:38 +0300 Subject: [PATCH] Primary implementation of untrusted message channel --- mROA/Abstract/IInteractionModule.cs | 11 +++- .../Backend/NetworkGatewayModule.cs | 21 +++++-- .../Frontend/NetworkFrontendBridge.cs | 21 +++++-- .../NextGenerationInteractionModule.cs | 62 ++++++++++++++++--- mROA/mROA.csproj | 1 + 5 files changed, 96 insertions(+), 20 deletions(-) diff --git a/mROA/Abstract/IInteractionModule.cs b/mROA/Abstract/IInteractionModule.cs index d797ae0..fae3768 100644 --- a/mROA/Abstract/IInteractionModule.cs +++ b/mROA/Abstract/IInteractionModule.cs @@ -1,5 +1,6 @@ using System; using System.IO; +using System.Threading.Channels; using System.Threading.Tasks; using mROA.Implementation; @@ -8,13 +9,17 @@ namespace mROA.Abstract public interface INextGenerationInteractionModule : IInjectableModule, IDisposable { int ConnectionId { get; set; } - public Stream? BaseStream { get; set; } + Stream? BaseStream { get; set; } + ChannelReader UntrustedReceiveChanel { get; set; } + ChannelWriter<(int clientId, NetworkMessageHeader messageHeader)> UntrustedPostChanel { get; set; } + Task GetNextMessageReceiving(bool infinite = true); Task PostMessageAsync(NetworkMessageHeader messageHeader); + Task PostMessageUntrustedAsync(NetworkMessageHeader messageHeader); void HandleMessage(NetworkMessageHeader messageHeader); - NetworkMessageHeader[] UnhandledMessages { get; } + // NetworkMessageHeader[] UnhandledMessages { get; } NetworkMessageHeader? FirstByFilter(Predicate predicate); - event Action OnDisconected; + event Action OnDisconnected; Task Restart(bool sendRecovery); } } \ No newline at end of file diff --git a/mROA/Implementation/Backend/NetworkGatewayModule.cs b/mROA/Implementation/Backend/NetworkGatewayModule.cs index a3236a8..f756920 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.Channels; using System.Threading.Tasks; using mROA.Abstract; @@ -67,13 +68,18 @@ namespace mROA.Implementation.Backend var client = _tcpListener.AcceptTcpClient(); Console.WriteLine($"Client connected from {client.Client.RemoteEndPoint}"); var interaction = Activator.CreateInstance(_interactionModuleType!) as INextGenerationInteractionModule; - + foreach (var injectableModule in _injectableModules!) interaction!.Inject(injectableModule); interaction!.Inject(_serialization); interaction.BaseStream = client.GetStream(); - + interaction.UntrustedReceiveChanel = Channel.CreateUnbounded(new UnboundedChannelOptions + { + SingleWriter = false, + SingleReader = false, + AllowSynchronousContinuations = true + }).Reader; var connectionRequest = interaction.GetNextMessageReceiving(false) .GetAwaiter().GetResult()!; @@ -87,12 +93,19 @@ 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(); - + recoveryInteraction.Restart(false); Console.WriteLine("Connection recovery for client {0} finished", recoveryRequest.Id); break; diff --git a/mROA/Implementation/Frontend/NetworkFrontendBridge.cs b/mROA/Implementation/Frontend/NetworkFrontendBridge.cs index ca32354..27e8219 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.Channels; using System.Threading.Tasks; using mROA.Abstract; using Exception = System.Exception; @@ -11,7 +12,7 @@ namespace mROA.Implementation.Frontend { private readonly IPEndPoint _serverEndPoint; private TcpClient _tcpClient = new(); - private NextGenerationInteractionModule? _interactionModule; + private INextGenerationInteractionModule? _interactionModule; private ISerializationToolkit? _serialization; public NetworkFrontendBridge(IPEndPoint serverEndPoint) @@ -42,12 +43,14 @@ namespace mROA.Implementation.Frontend _tcpClient.Connect(_serverEndPoint); _interactionModule.BaseStream = _tcpClient.GetStream(); - - _interactionModule.OnDisconected += id => + _interactionModule.UntrustedReceiveChanel = Channel.CreateUnbounded(new UnboundedChannelOptions { - Reconnect(); - }; - + SingleWriter = false, + SingleReader = false, + AllowSynchronousContinuations = true + }).Reader; + _interactionModule.OnDisconnected += id => { Reconnect(); }; + _interactionModule.PostMessageAsync(new NetworkMessageHeader(_serialization, new ClientConnect())).Wait(); var idMessage = _interactionModule.GetNextMessageReceiving(false).GetAwaiter().GetResult(); if (idMessage.MessageType != EMessageType.IdAssigning) @@ -67,6 +70,12 @@ namespace mROA.Implementation.Frontend _tcpClient = new TcpClient(); _tcpClient.Connect(_serverEndPoint); _interactionModule.BaseStream = _tcpClient.GetStream(); + _interactionModule.UntrustedReceiveChanel = Channel.CreateUnbounded(new UnboundedChannelOptions + { + SingleWriter = false, + SingleReader = false, + AllowSynchronousContinuations = true + }).Reader; await _interactionModule.Restart(true); } diff --git a/mROA/Implementation/NextGenerationInteractionModule.cs b/mROA/Implementation/NextGenerationInteractionModule.cs index 57ba7dd..92d63a3 100644 --- a/mROA/Implementation/NextGenerationInteractionModule.cs +++ b/mROA/Implementation/NextGenerationInteractionModule.cs @@ -2,6 +2,7 @@ using System.Collections.Generic; using System.IO; using System.Linq; +using System.Threading.Channels; using System.Threading.Tasks; using mROA.Abstract; @@ -20,6 +21,8 @@ namespace mROA.Implementation private bool _isInReconnectionState; private bool _isActive = true; private TaskCompletionSource _reconnection; + private ValueTask? _trustedReceive; + private TaskCompletionSource _untrustedReceive; public NextGenerationInteractionModule() { @@ -30,9 +33,13 @@ namespace mROA.Implementation public Stream? BaseStream { - get => _baseStream; set => _baseStream = value; + get => _baseStream; + set => _baseStream = value; } + public ChannelReader UntrustedReceiveChanel { get; set; } + public ChannelWriter<(int clientId, NetworkMessageHeader messageHeader)> UntrustedPostChanel { get; set; } + public void Inject(T dependency) { @@ -53,7 +60,6 @@ namespace mROA.Implementation if (_currentReceiving != null) return _currentReceiving; _currentReceiving = Task.Run(async () => await GetNextMessage()); return _currentReceiving; - } #pragma warning disable CS8602 // Dereference of a possibly null reference. private async ValueTask PostMessageInternal(NetworkMessageHeader messageHeader) @@ -68,7 +74,7 @@ namespace mROA.Implementation if (!_baseStream.CanWrite) return false; - + await BaseStream.WriteAsync(header); await BaseStream.WriteAsync(rawMessage); return true; @@ -101,25 +107,31 @@ namespace mROA.Implementation { return; } + _isConnected = false; withError = true; await MakeRecovery("OUT"); } } + public async Task PostMessageUntrustedAsync(NetworkMessageHeader messageHeader) + { + await UntrustedPostChanel.WriteAsync((ConnectionId, messageHeader)); + } + public void HandleMessage(NetworkMessageHeader messageHeader) { _messageBuffer.Remove(messageHeader); } - public NetworkMessageHeader[] UnhandledMessages => _messageBuffer.ToArray(); + // public NetworkMessageHeader[] UnhandledMessages => _messageBuffer.ToArray(); public NetworkMessageHeader? FirstByFilter(Predicate predicate) { return _messageBuffer.FirstOrDefault(m => predicate(m)); } - public event Action? OnDisconected; + public event Action? OnDisconnected; private async Task GetNextMessage() { @@ -140,7 +152,41 @@ namespace mROA.Implementation try { - var message = await Receive(); + NetworkMessageHeader message; + + var wasNull = _untrustedReceive is null; + + _trustedReceive ??= Receive(); + _untrustedReceive = new TaskCompletionSource(); + + + if (!wasNull) + { + if (_trustedReceive.Value.IsCompleted) + { + _trustedReceive = Receive(); + } + + if (_untrustedReceive.Task.IsCompleted) + { + _untrustedReceive = new TaskCompletionSource(); + _ = UntrustedReceiveChanel.ReadAsync().AsTask() + .ContinueWith(task => _untrustedReceive.SetResult(task.Result)); + } + } + else + { + _ = UntrustedReceiveChanel.ReadAsync().AsTask() + .ContinueWith(task => _untrustedReceive.SetResult(task.Result)); + } + + + await Task.WhenAny(_trustedReceive.Value.AsTask() , _untrustedReceive.Task); + + message = _trustedReceive.Value.IsCompleted + ? _trustedReceive.Value.Result + : _untrustedReceive.Task.Result; + _currentReceiving = Task.Run(async () => await GetNextMessage()); return message; @@ -151,6 +197,7 @@ namespace mROA.Implementation { return NetworkMessageHeader.Null; } + withError = true; await MakeRecovery("IN"); } @@ -238,7 +285,7 @@ namespace mROA.Implementation Console.WriteLine("Call OnDisconnected from {0}", source); _isInReconnectionState = true; - OnDisconected?.Invoke(ConnectionId); + OnDisconnected?.Invoke(ConnectionId); } Console.WriteLine("Waiting for reconnect from {0}", source); @@ -263,6 +310,7 @@ namespace mROA.Implementation { _currentReceiving?.Dispose(); } + _baseStream?.Dispose(); } } diff --git a/mROA/mROA.csproj b/mROA/mROA.csproj index 5c9aea9..3c172de 100644 --- a/mROA/mROA.csproj +++ b/mROA/mROA.csproj @@ -26,6 +26,7 @@ +