Primary implementation of untrusted message channel
This commit is contained in:
@@ -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;
|
||||||
|
|
||||||
@@ -67,13 +68,18 @@ namespace mROA.Implementation.Backend
|
|||||||
var client = _tcpListener.AcceptTcpClient();
|
var client = _tcpListener.AcceptTcpClient();
|
||||||
Console.WriteLine($"Client connected from {client.Client.RemoteEndPoint}");
|
Console.WriteLine($"Client connected from {client.Client.RemoteEndPoint}");
|
||||||
var interaction = Activator.CreateInstance(_interactionModuleType!) as INextGenerationInteractionModule;
|
var interaction = Activator.CreateInstance(_interactionModuleType!) as INextGenerationInteractionModule;
|
||||||
|
|
||||||
foreach (var injectableModule in _injectableModules!)
|
foreach (var injectableModule in _injectableModules!)
|
||||||
interaction!.Inject(injectableModule);
|
interaction!.Inject(injectableModule);
|
||||||
|
|
||||||
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,12 +93,19 @@ 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);
|
||||||
Console.WriteLine("Connection recovery for client {0} finished", recoveryRequest.Id);
|
Console.WriteLine("Connection recovery for client {0} finished", recoveryRequest.Id);
|
||||||
break;
|
break;
|
||||||
|
|||||||
@@ -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,12 +43,14 @@ 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();
|
||||||
if (idMessage.MessageType != EMessageType.IdAssigning)
|
if (idMessage.MessageType != EMessageType.IdAssigning)
|
||||||
@@ -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)
|
||||||
@@ -68,7 +74,7 @@ namespace mROA.Implementation
|
|||||||
|
|
||||||
if (!_baseStream.CanWrite)
|
if (!_baseStream.CanWrite)
|
||||||
return false;
|
return false;
|
||||||
|
|
||||||
await BaseStream.WriteAsync(header);
|
await BaseStream.WriteAsync(header);
|
||||||
await BaseStream.WriteAsync(rawMessage);
|
await BaseStream.WriteAsync(rawMessage);
|
||||||
return true;
|
return true;
|
||||||
@@ -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();
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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>
|
||||||
|
|||||||
Reference in New Issue
Block a user