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.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<NetworkMessageHeader> UntrustedReceiveChanel { get; set; }
ChannelWriter<(int clientId, NetworkMessageHeader messageHeader)> UntrustedPostChanel { get; set; }
Task<NetworkMessageHeader> 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<NetworkMessageHeader> predicate);
event Action<int> OnDisconected;
event Action<int> OnDisconnected;
Task Restart(bool sendRecovery);
}
}
@@ -1,6 +1,7 @@
using System;
using System.Net;
using System.Net.Sockets;
using System.Threading.Channels;
using System.Threading.Tasks;
using mROA.Abstract;
@@ -73,7 +74,12 @@ namespace mROA.Implementation.Backend
interaction!.Inject(_serialization);
interaction.BaseStream = client.GetStream();
interaction.UntrustedReceiveChanel = Channel.CreateUnbounded<NetworkMessageHeader>(new UnboundedChannelOptions
{
SingleWriter = false,
SingleReader = false,
AllowSynchronousContinuations = true
}).Reader;
var connectionRequest = interaction.GetNextMessageReceiving(false)
.GetAwaiter().GetResult()!;
@@ -87,10 +93,17 @@ namespace mROA.Implementation.Backend
break;
case EMessageType.ClientRecovery:
{
interaction.BaseStream = null;
var recoveryRequest = _serialization!.Deserialize<ClientRecovery>(connectionRequest.Data)!;
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.Restart(false);
@@ -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,11 +43,13 @@ namespace mROA.Implementation.Frontend
_tcpClient.Connect(_serverEndPoint);
_interactionModule.BaseStream = _tcpClient.GetStream();
_interactionModule.OnDisconected += id =>
_interactionModule.UntrustedReceiveChanel = Channel.CreateUnbounded<NetworkMessageHeader>(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();
@@ -67,6 +70,12 @@ namespace mROA.Implementation.Frontend
_tcpClient = new TcpClient();
_tcpClient.Connect(_serverEndPoint);
_interactionModule.BaseStream = _tcpClient.GetStream();
_interactionModule.UntrustedReceiveChanel = Channel.CreateUnbounded<NetworkMessageHeader>(new UnboundedChannelOptions
{
SingleWriter = false,
SingleReader = false,
AllowSynchronousContinuations = true
}).Reader;
await _interactionModule.Restart(true);
}
@@ -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<Stream> _reconnection;
private ValueTask<NetworkMessageHeader>? _trustedReceive;
private TaskCompletionSource<NetworkMessageHeader> _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<NetworkMessageHeader> UntrustedReceiveChanel { get; set; }
public ChannelWriter<(int clientId, NetworkMessageHeader messageHeader)> UntrustedPostChanel { get; set; }
public void Inject<T>(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<bool> PostMessageInternal(NetworkMessageHeader messageHeader)
@@ -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<NetworkMessageHeader> predicate)
{
return _messageBuffer.FirstOrDefault(m => predicate(m));
}
public event Action<int>? OnDisconected;
public event Action<int>? OnDisconnected;
private async Task<NetworkMessageHeader> 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<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());
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();
}
}
+1
View File
@@ -26,6 +26,7 @@
<ItemGroup>
<PackageReference Include="System.Text.Json" Version="9.0.2"/>
<PackageReference Include="System.Threading.Channels" Version="9.0.4" />
</ItemGroup>
</Project>