Primary implementation of untrusted message channel

This commit is contained in:
2025-04-12 12:44:38 +03:00
parent 3f92cc0f64
commit 3d3772a341
5 changed files with 96 additions and 20 deletions
+8 -3
View File
@@ -1,5 +1,6 @@
using System; using System;
using System.IO; using System.IO;
using System.Threading.Channels;
using System.Threading.Tasks; using System.Threading.Tasks;
using mROA.Implementation; using mROA.Implementation;
@@ -8,13 +9,17 @@ namespace mROA.Abstract
public interface INextGenerationInteractionModule : IInjectableModule, IDisposable public interface INextGenerationInteractionModule : IInjectableModule, IDisposable
{ {
int ConnectionId { get; set; } int ConnectionId { get; set; }
public Stream? BaseStream { get; set; } Stream? BaseStream { get; set; }
ChannelReader<NetworkMessageHeader> UntrustedReceiveChanel { get; set; }
ChannelWriter<(int clientId, NetworkMessageHeader messageHeader)> UntrustedPostChanel { get; set; }
Task<NetworkMessageHeader> GetNextMessageReceiving(bool infinite = true); Task<NetworkMessageHeader> GetNextMessageReceiving(bool infinite = true);
Task PostMessageAsync(NetworkMessageHeader messageHeader); Task PostMessageAsync(NetworkMessageHeader messageHeader);
Task PostMessageUntrustedAsync(NetworkMessageHeader messageHeader);
void HandleMessage(NetworkMessageHeader messageHeader); void HandleMessage(NetworkMessageHeader messageHeader);
NetworkMessageHeader[] UnhandledMessages { get; } // NetworkMessageHeader[] UnhandledMessages { get; }
NetworkMessageHeader? FirstByFilter(Predicate<NetworkMessageHeader> predicate); NetworkMessageHeader? FirstByFilter(Predicate<NetworkMessageHeader> predicate);
event Action<int> OnDisconected; event Action<int> OnDisconnected;
Task Restart(bool sendRecovery); Task Restart(bool sendRecovery);
} }
} }
@@ -1,6 +1,7 @@
using System; using System;
using System.Net; using System.Net;
using System.Net.Sockets; using System.Net.Sockets;
using System.Threading.Channels;
using System.Threading.Tasks; using System.Threading.Tasks;
using mROA.Abstract; using mROA.Abstract;
@@ -73,7 +74,12 @@ namespace mROA.Implementation.Backend
interaction!.Inject(_serialization); interaction!.Inject(_serialization);
interaction.BaseStream = client.GetStream(); interaction.BaseStream = client.GetStream();
interaction.UntrustedReceiveChanel = Channel.CreateUnbounded<NetworkMessageHeader>(new UnboundedChannelOptions
{
SingleWriter = false,
SingleReader = false,
AllowSynchronousContinuations = true
}).Reader;
var connectionRequest = interaction.GetNextMessageReceiving(false) var connectionRequest = interaction.GetNextMessageReceiving(false)
.GetAwaiter().GetResult()!; .GetAwaiter().GetResult()!;
@@ -87,10 +93,17 @@ namespace mROA.Implementation.Backend
break; break;
case EMessageType.ClientRecovery: case EMessageType.ClientRecovery:
{ {
interaction.BaseStream = null; interaction.BaseStream = null;
var recoveryRequest = _serialization!.Deserialize<ClientRecovery>(connectionRequest.Data)!; var recoveryRequest = _serialization!.Deserialize<ClientRecovery>(connectionRequest.Data)!;
var recoveryInteraction = _hub.GetInteraction(recoveryRequest.Id); var recoveryInteraction = _hub.GetInteraction(recoveryRequest.Id);
recoveryInteraction.UntrustedReceiveChanel =
Channel.CreateUnbounded<NetworkMessageHeader>(new UnboundedChannelOptions
{
SingleWriter = false,
SingleReader = false,
AllowSynchronousContinuations = true,
}).Reader;
recoveryInteraction.BaseStream = client.GetStream(); recoveryInteraction.BaseStream = client.GetStream();
recoveryInteraction.Restart(false); recoveryInteraction.Restart(false);
@@ -1,6 +1,7 @@
using System; using System;
using System.Net; using System.Net;
using System.Net.Sockets; using System.Net.Sockets;
using System.Threading.Channels;
using System.Threading.Tasks; using System.Threading.Tasks;
using mROA.Abstract; using mROA.Abstract;
using Exception = System.Exception; using Exception = System.Exception;
@@ -11,7 +12,7 @@ namespace mROA.Implementation.Frontend
{ {
private readonly IPEndPoint _serverEndPoint; private readonly IPEndPoint _serverEndPoint;
private TcpClient _tcpClient = new(); private TcpClient _tcpClient = new();
private NextGenerationInteractionModule? _interactionModule; private INextGenerationInteractionModule? _interactionModule;
private ISerializationToolkit? _serialization; private ISerializationToolkit? _serialization;
public NetworkFrontendBridge(IPEndPoint serverEndPoint) public NetworkFrontendBridge(IPEndPoint serverEndPoint)
@@ -42,11 +43,13 @@ namespace mROA.Implementation.Frontend
_tcpClient.Connect(_serverEndPoint); _tcpClient.Connect(_serverEndPoint);
_interactionModule.BaseStream = _tcpClient.GetStream(); _interactionModule.BaseStream = _tcpClient.GetStream();
_interactionModule.UntrustedReceiveChanel = Channel.CreateUnbounded<NetworkMessageHeader>(new UnboundedChannelOptions
_interactionModule.OnDisconected += id =>
{ {
Reconnect(); SingleWriter = false,
}; SingleReader = false,
AllowSynchronousContinuations = true
}).Reader;
_interactionModule.OnDisconnected += id => { Reconnect(); };
_interactionModule.PostMessageAsync(new NetworkMessageHeader(_serialization, new ClientConnect())).Wait(); _interactionModule.PostMessageAsync(new NetworkMessageHeader(_serialization, new ClientConnect())).Wait();
var idMessage = _interactionModule.GetNextMessageReceiving(false).GetAwaiter().GetResult(); var idMessage = _interactionModule.GetNextMessageReceiving(false).GetAwaiter().GetResult();
@@ -67,6 +70,12 @@ namespace mROA.Implementation.Frontend
_tcpClient = new TcpClient(); _tcpClient = new TcpClient();
_tcpClient.Connect(_serverEndPoint); _tcpClient.Connect(_serverEndPoint);
_interactionModule.BaseStream = _tcpClient.GetStream(); _interactionModule.BaseStream = _tcpClient.GetStream();
_interactionModule.UntrustedReceiveChanel = Channel.CreateUnbounded<NetworkMessageHeader>(new UnboundedChannelOptions
{
SingleWriter = false,
SingleReader = false,
AllowSynchronousContinuations = true
}).Reader;
await _interactionModule.Restart(true); await _interactionModule.Restart(true);
} }
@@ -2,6 +2,7 @@
using System.Collections.Generic; using System.Collections.Generic;
using System.IO; using System.IO;
using System.Linq; using System.Linq;
using System.Threading.Channels;
using System.Threading.Tasks; using System.Threading.Tasks;
using mROA.Abstract; using mROA.Abstract;
@@ -20,6 +21,8 @@ namespace mROA.Implementation
private bool _isInReconnectionState; private bool _isInReconnectionState;
private bool _isActive = true; private bool _isActive = true;
private TaskCompletionSource<Stream> _reconnection; private TaskCompletionSource<Stream> _reconnection;
private ValueTask<NetworkMessageHeader>? _trustedReceive;
private TaskCompletionSource<NetworkMessageHeader> _untrustedReceive;
public NextGenerationInteractionModule() public NextGenerationInteractionModule()
{ {
@@ -30,9 +33,13 @@ namespace mROA.Implementation
public Stream? BaseStream public Stream? BaseStream
{ {
get => _baseStream; set => _baseStream = value; get => _baseStream;
set => _baseStream = value;
} }
public ChannelReader<NetworkMessageHeader> UntrustedReceiveChanel { get; set; }
public ChannelWriter<(int clientId, NetworkMessageHeader messageHeader)> UntrustedPostChanel { get; set; }
public void Inject<T>(T dependency) public void Inject<T>(T dependency)
{ {
@@ -53,7 +60,6 @@ namespace mROA.Implementation
if (_currentReceiving != null) return _currentReceiving; if (_currentReceiving != null) return _currentReceiving;
_currentReceiving = Task.Run(async () => await GetNextMessage()); _currentReceiving = Task.Run(async () => await GetNextMessage());
return _currentReceiving; return _currentReceiving;
} }
#pragma warning disable CS8602 // Dereference of a possibly null reference. #pragma warning disable CS8602 // Dereference of a possibly null reference.
private async ValueTask<bool> PostMessageInternal(NetworkMessageHeader messageHeader) private async ValueTask<bool> PostMessageInternal(NetworkMessageHeader messageHeader)
@@ -101,25 +107,31 @@ namespace mROA.Implementation
{ {
return; return;
} }
_isConnected = false; _isConnected = false;
withError = true; withError = true;
await MakeRecovery("OUT"); await MakeRecovery("OUT");
} }
} }
public async Task PostMessageUntrustedAsync(NetworkMessageHeader messageHeader)
{
await UntrustedPostChanel.WriteAsync((ConnectionId, messageHeader));
}
public void HandleMessage(NetworkMessageHeader messageHeader) public void HandleMessage(NetworkMessageHeader messageHeader)
{ {
_messageBuffer.Remove(messageHeader); _messageBuffer.Remove(messageHeader);
} }
public NetworkMessageHeader[] UnhandledMessages => _messageBuffer.ToArray(); // public NetworkMessageHeader[] UnhandledMessages => _messageBuffer.ToArray();
public NetworkMessageHeader? FirstByFilter(Predicate<NetworkMessageHeader> predicate) public NetworkMessageHeader? FirstByFilter(Predicate<NetworkMessageHeader> predicate)
{ {
return _messageBuffer.FirstOrDefault(m => predicate(m)); return _messageBuffer.FirstOrDefault(m => predicate(m));
} }
public event Action<int>? OnDisconected; public event Action<int>? OnDisconnected;
private async Task<NetworkMessageHeader> GetNextMessage() private async Task<NetworkMessageHeader> GetNextMessage()
{ {
@@ -140,7 +152,41 @@ namespace mROA.Implementation
try try
{ {
var message = await Receive(); NetworkMessageHeader message;
var wasNull = _untrustedReceive is null;
_trustedReceive ??= Receive();
_untrustedReceive = new TaskCompletionSource<NetworkMessageHeader>();
if (!wasNull)
{
if (_trustedReceive.Value.IsCompleted)
{
_trustedReceive = Receive();
}
if (_untrustedReceive.Task.IsCompleted)
{
_untrustedReceive = new TaskCompletionSource<NetworkMessageHeader>();
_ = 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()); _currentReceiving = Task.Run(async () => await GetNextMessage());
return message; return message;
@@ -151,6 +197,7 @@ namespace mROA.Implementation
{ {
return NetworkMessageHeader.Null; return NetworkMessageHeader.Null;
} }
withError = true; withError = true;
await MakeRecovery("IN"); await MakeRecovery("IN");
} }
@@ -238,7 +285,7 @@ namespace mROA.Implementation
Console.WriteLine("Call OnDisconnected from {0}", source); Console.WriteLine("Call OnDisconnected from {0}", source);
_isInReconnectionState = true; _isInReconnectionState = true;
OnDisconected?.Invoke(ConnectionId); OnDisconnected?.Invoke(ConnectionId);
} }
Console.WriteLine("Waiting for reconnect from {0}", source); Console.WriteLine("Waiting for reconnect from {0}", source);
@@ -263,6 +310,7 @@ namespace mROA.Implementation
{ {
_currentReceiving?.Dispose(); _currentReceiving?.Dispose();
} }
_baseStream?.Dispose(); _baseStream?.Dispose();
} }
} }
+1
View File
@@ -26,6 +26,7 @@
<ItemGroup> <ItemGroup>
<PackageReference Include="System.Text.Json" Version="9.0.2"/> <PackageReference Include="System.Text.Json" Version="9.0.2"/>
<PackageReference Include="System.Threading.Channels" Version="9.0.4" />
</ItemGroup> </ItemGroup>
</Project> </Project>