Refactor Network gateway module

This commit is contained in:
2025-07-18 14:08:29 +03:00
parent 85cf1bc48a
commit 0c648f0af0
@@ -16,7 +16,7 @@ namespace mROA.Implementation.Backend
private readonly HubRequestExtractor _hre; private readonly HubRequestExtractor _hre;
private readonly DistributionOptions _distribution; private readonly DistributionOptions _distribution;
private readonly IContextualSerializationToolKit _serialization; private readonly IContextualSerializationToolKit _serialization;
private readonly Dictionary<int, CancellationTokenSource> _extractorsCTS = new(); private readonly Dictionary<int, CancellationTokenSource> _extractorsTokenSources = new();
private readonly ICallIndexProvider _callIndexProvider; private readonly ICallIndexProvider _callIndexProvider;
private readonly IIdentityGenerator _identityGenerator; private readonly IIdentityGenerator _identityGenerator;
public NetworkGatewayModule(IOptions<GatewayOptions> options, IIdentityGenerator identityGenerator, IContextualSerializationToolKit serialization, ICallIndexProvider callIndexProvider, IConnectionHub hub, IOptions<DistributionOptions> distribution, HubRequestExtractor hre) public NetworkGatewayModule(IOptions<GatewayOptions> options, IIdentityGenerator identityGenerator, IContextualSerializationToolKit serialization, ICallIndexProvider callIndexProvider, IConnectionHub hub, IOptions<DistributionOptions> distribution, HubRequestExtractor hre)
@@ -49,11 +49,19 @@ namespace mROA.Implementation.Backend
while (true) while (true)
{ {
var client = await _tcpListener.AcceptTcpClientAsync(); var client = await _tcpListener.AcceptTcpClientAsync();
_ = HandleConnection(client).ConfigureAwait(false);
}
}
private async Task HandleConnection(TcpClient client)
{
Console.WriteLine($"Client connected from {client.Client.RemoteEndPoint}"); Console.WriteLine($"Client connected from {client.Client.RemoteEndPoint}");
var interaction = new ChannelInteractionModule(_serialization, _identityGenerator); var interaction = new ChannelInteractionModule(_serialization, _identityGenerator);
var context = new EndPointContext(null, null); var context = new EndPointContext(null, null)
context.CallIndexProvider = _callIndexProvider; {
CallIndexProvider = _callIndexProvider
};
var streamExtractor = var streamExtractor =
new ChannelInteractionModule.StreamExtractor(client.GetStream(), _serialization, context); new ChannelInteractionModule.StreamExtractor(client.GetStream(), _serialization, context);
interaction.IsConnected = () => streamExtractor.IsConnected; interaction.IsConnected = () => streamExtractor.IsConnected;
@@ -68,6 +76,22 @@ namespace mROA.Implementation.Backend
switch (connectionRequest.MessageType) switch (connectionRequest.MessageType)
{ {
case EMessageType.ClientConnect: case EMessageType.ClientConnect:
HandleNewClient(context, interaction, streamExtractor, cts);
break;
case EMessageType.ClientRecovery:
{
RecoverDisconnectedClient(connectionRequest, streamExtractor, cts);
break;
}
default:
client.Close();
break;
}
}
private void HandleNewClient(EndPointContext context, ChannelInteractionModule interaction,
ChannelInteractionModule.StreamExtractor streamExtractor, CancellationTokenSource cts)
{
context.HostId = 0; context.HostId = 0;
context.OwnerId = -interaction.ConnectionId; context.OwnerId = -interaction.ConnectionId;
interaction.Context = context; interaction.Context = context;
@@ -75,22 +99,24 @@ namespace mROA.Implementation.Backend
_ = streamExtractor.SendFromChannel(interaction.TrustedPostChanel, cts.Token); _ = streamExtractor.SendFromChannel(interaction.TrustedPostChanel, cts.Token);
interaction.PostMessageAsync(new NetworkMessageHeader(_serialization, interaction.PostMessageAsync(new NetworkMessageHeader(_serialization,
new IdAssignment { Id = interaction.ConnectionId }, null)); new IdAssignment { Id = interaction.ConnectionId }, null));
_extractorsCTS[interaction.ConnectionId] = cts; _extractorsTokenSources[interaction.ConnectionId] = cts;
_hub.RegisterInteraction(interaction); _hub.RegisterInteraction(interaction);
_hre.HubOnOnConnected(new RepresentationModule(interaction, _serialization)); _hre.HubOnOnConnected(new RepresentationModule(interaction, _serialization));
if (_distribution.DistributionType != EDistributionType.Channeled) if (_distribution.DistributionType != EDistributionType.Channeled)
{ {
} }
Console.WriteLine("Client registered"); }
break;
case EMessageType.ClientRecovery: private void RecoverDisconnectedClient(NetworkMessageHeader connectionRequest, ChannelInteractionModule.StreamExtractor streamExtractor,
CancellationTokenSource cts)
{ {
var recoveryRequest = _serialization.Deserialize<ClientRecovery>(connectionRequest.Data, null); var recoveryRequest = _serialization.Deserialize<ClientRecovery>(connectionRequest.Data, null);
var recoveryInteraction = _hub.GetInteraction(recoveryRequest.Id); var recoveryInteraction = _hub.GetInteraction(recoveryRequest.Id);
_extractorsCTS[-recoveryRequest.Id].Cancel(); _extractorsTokenSources[-recoveryRequest.Id].Cancel();
recoveryInteraction.IsConnected = () => streamExtractor.IsConnected; recoveryInteraction.IsConnected = () => streamExtractor.IsConnected;
streamExtractor.MessageReceived = message => streamExtractor.MessageReceived = message =>
@@ -101,15 +127,7 @@ namespace mROA.Implementation.Backend
Task.Run(async () => await streamExtractor.LoopedReceive(cts.Token).ConfigureAwait(false)); Task.Run(async () => await streamExtractor.LoopedReceive(cts.Token).ConfigureAwait(false));
recoveryInteraction.Restart(false); recoveryInteraction.Restart(false);
break;
}
default:
client.Close();
break;
}
}
} }
} }