diff --git a/mROA.Test/NextGenTest.cs b/mROA.Test/NextGenTest.cs index 64c6d86..6a36ef5 100644 --- a/mROA.Test/NextGenTest.cs +++ b/mROA.Test/NextGenTest.cs @@ -33,7 +33,7 @@ namespace mROA.Test Task.Run(() => { _listener.Start(); - _interactionModuleB.BaseStream = _listener.AcceptTcpClient().GetStream(); + // _interactionModuleB.BaseStream = _listener.AcceptTcpClient().GetStream(); foreach (var guid in guids) { @@ -43,7 +43,7 @@ namespace mROA.Test var client = new TcpClient(); client.Connect(IPAddress.Loopback, 4567); - _interactionModuleA.BaseStream = client.GetStream(); + // _interactionModuleA.BaseStream = client.GetStream(); var tasks = guids.Select(ReadStream); diff --git a/mROA/Abstract/IInteractionModule.cs b/mROA/Abstract/IChannelInteractionModule.cs similarity index 69% rename from mROA/Abstract/IInteractionModule.cs rename to mROA/Abstract/IChannelInteractionModule.cs index 4a884e8..e9030d1 100644 --- a/mROA/Abstract/IInteractionModule.cs +++ b/mROA/Abstract/IChannelInteractionModule.cs @@ -9,15 +9,13 @@ namespace mROA.Abstract public interface IChannelInteractionModule : IInjectableModule, IDisposable { int ConnectionId { get; set; } - ChannelWriter ReceiveChanel { get; } + Channel ReceiveChanel { get; } ChannelReader TrustedPostChanel { get; } ChannelReader UntrustedPostChanel { get; } Func IsConnected { get; set; } - Task GetNextMessageReceiving(bool infinite = true); + ValueTask GetNextMessageReceiving(bool infinite = true); Task PostMessageAsync(NetworkMessageHeader messageHeader); Task PostMessageUntrustedAsync(NetworkMessageHeader messageHeader); - void HandleMessage(NetworkMessageHeader messageHeader); - NetworkMessageHeader? FirstByFilter(Predicate predicate); event Action OnDisconnected; Task Restart(bool sendRecovery); } diff --git a/mROA/Abstract/IRepresentationModule.cs b/mROA/Abstract/IRepresentationModule.cs index 2d9fc06..bca9a76 100644 --- a/mROA/Abstract/IRepresentationModule.cs +++ b/mROA/Abstract/IRepresentationModule.cs @@ -10,7 +10,8 @@ namespace mROA.Abstract { int Id { get; } - Task<(object parced, EMessageType originalType)> GetSingle(Predicate rule, CancellationToken token, + Task<(object? Deserialized, EMessageType MessageType)> GetSingle(Predicate rule, + CancellationToken token, params Func[] converter); IAsyncEnumerable<(object parced, EMessageType originalType)> GetStream(Predicate rule, CancellationToken token, diff --git a/mROA/Implementation/Backend/NetworkGatewayModule.cs b/mROA/Implementation/Backend/NetworkGatewayModule.cs index 52ac70c..c5d15b8 100644 --- a/mROA/Implementation/Backend/NetworkGatewayModule.cs +++ b/mROA/Implementation/Backend/NetworkGatewayModule.cs @@ -78,7 +78,7 @@ namespace mROA.Implementation.Backend var streamExtractor = new StreamExtractor(client.GetStream(), _serialization); interaction.IsConnected = () => streamExtractor.IsConnected; - streamExtractor.MessageReceived += message => interaction.ReceiveChanel.WriteAsync(message); + streamExtractor.MessageReceived = message => interaction.ReceiveChanel.Writer.WriteAsync(message); streamExtractor.SingleReceive(); var connectionRequest = interaction.GetNextMessageReceiving(false) .GetAwaiter().GetResult()!; @@ -98,10 +98,8 @@ namespace mROA.Implementation.Backend var recoveryRequest = _serialization!.Deserialize(connectionRequest.Data)!; var recoveryInteraction = _hub.GetInteraction(recoveryRequest.Id); - streamExtractor = new StreamExtractor(client.GetStream(), _serialization); - streamExtractor.MessageReceived += message => recoveryInteraction.ReceiveChanel.WriteAsync(message); + streamExtractor.MessageReceived = message => recoveryInteraction.ReceiveChanel.Writer.WriteAsync(message); - _ = streamExtractor.LoopedReceive(); recoveryInteraction.Restart(false); Console.WriteLine("Connection recovery for client {0} finished", recoveryRequest.Id); break; diff --git a/mROA/Implementation/ChannelInteractionModule.cs b/mROA/Implementation/ChannelInteractionModule.cs index c3f637b..b52d385 100644 --- a/mROA/Implementation/ChannelInteractionModule.cs +++ b/mROA/Implementation/ChannelInteractionModule.cs @@ -2,7 +2,6 @@ using System.Collections.Generic; using System.IO; using System.Linq; -using System.Threading; using System.Threading.Channels; using System.Threading.Tasks; using mROA.Abstract; @@ -14,7 +13,6 @@ namespace mROA.Implementation private readonly ChannelReader _receiveReader; private readonly ChannelWriter _trustedWriter; private readonly ChannelWriter _untrustedWriter; - private readonly Channel _inputChannel; private readonly Channel _outputTrustedChannel; private readonly Channel _outputUntrustedChannel; private readonly List _messageBuffer = new(128); @@ -26,13 +24,13 @@ namespace mROA.Implementation public ChannelInteractionModule() { - _inputChannel = Channel.CreateUnbounded(new UnboundedChannelOptions + ReceiveChanel = Channel.CreateUnbounded(new UnboundedChannelOptions { SingleReader = false, SingleWriter = false, AllowSynchronousContinuations = true }); - _receiveReader = _inputChannel.Reader; + _receiveReader = ReceiveChanel.Reader; _outputTrustedChannel = Channel.CreateBounded(new BoundedChannelOptions(1) { SingleReader = true, @@ -52,7 +50,8 @@ namespace mROA.Implementation public int ConnectionId { get; set; } - public ChannelWriter ReceiveChanel => _inputChannel.Writer; + public Channel ReceiveChanel { get; } + public ChannelReader TrustedPostChanel => _outputTrustedChannel.Reader; public ChannelReader UntrustedPostChanel => _outputUntrustedChannel.Reader; public Func IsConnected { get; set; } @@ -71,12 +70,12 @@ namespace mROA.Implementation } } - public Task GetNextMessageReceiving(bool infinite = true) + public ValueTask GetNextMessageReceiving(bool infinite = true) { - if (!infinite) return _receiveReader.ReadAsync().AsTask(); - if (_currentReceiving != null) return _currentReceiving; - _currentReceiving = Task.Run(async () => await GetNextMessage()); - return _currentReceiving; + return _receiveReader.ReadAsync(); + // 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) @@ -192,6 +191,7 @@ namespace mROA.Implementation private async Task MakeRecovery(string source) { //TODO переделать реконнект + // Console.WriteLine("Staring recovery from {0}", source); // // lock (_reconnection) diff --git a/mROA/Implementation/Frontend/NetworkFrontendBridge.cs b/mROA/Implementation/Frontend/NetworkFrontendBridge.cs index 4b98fab..f66281e 100644 --- a/mROA/Implementation/Frontend/NetworkFrontendBridge.cs +++ b/mROA/Implementation/Frontend/NetworkFrontendBridge.cs @@ -78,7 +78,7 @@ namespace mROA.Implementation.Frontend _ = _currentExtractor.SendFromChannel(_interactionModule.TrustedPostChanel, _rawExtractorCancellation.Token); - _currentExtractor.MessageReceived += message => _interactionModule.ReceiveChanel.WriteAsync(message); + _currentExtractor.MessageReceived = message => _interactionModule.ReceiveChanel.Writer.WriteAsync(message); } private async Task Reconnect() diff --git a/mROA/Implementation/Frontend/RequestExtractor.cs b/mROA/Implementation/Frontend/RequestExtractor.cs index 5b348d9..f5616fc 100644 --- a/mROA/Implementation/Frontend/RequestExtractor.cs +++ b/mROA/Implementation/Frontend/RequestExtractor.cs @@ -47,6 +47,9 @@ namespace mROA.Implementation.Frontend public async Task StartExtraction() { ThrowIfNotInjected(); + + TransmissionConfig.OwnershipRepository = new StaticOwnershipRepository(_representationModule.Id); + var multiClientOwnershipRepository = TransmissionConfig.OwnershipRepository as MultiClientOwnershipRepository; multiClientOwnershipRepository?.RegisterOwnership(_representationModule!.Id); @@ -54,11 +57,11 @@ namespace mROA.Implementation.Frontend try { #if TRACE - var sw = new Stopwatch(); + var sw = new Stopwatch(); #endif var streamTokenSource = new CancellationTokenSource(); - + var query = _representationModule!.GetStream(m => m.MessageType is EMessageType.CallRequest or EMessageType.CancelRequest or EMessageType.EventRequest or EMessageType.ClientDisconnect, streamTokenSource.Token, @@ -71,17 +74,18 @@ namespace mROA.Implementation.Frontend await foreach (var command in query) { #if TRACE - Console.WriteLine("Waiting for request..."); - if (sw.IsRunning) - { - sw.Stop(); - Console.WriteLine($"Request handling took {Math.Round(sw.Elapsed.TotalMilliseconds * 1000.0)} microseconds."); - } + Console.WriteLine("Waiting for request..."); + if (sw.IsRunning) + { + sw.Stop(); + Console.WriteLine( + $"Request handling took {Math.Round(sw.Elapsed.TotalMilliseconds * 1000.0)} microseconds."); + } #endif #if TRACE - Console.WriteLine("Request received"); - sw.Restart(); + Console.WriteLine("Request received"); + sw.Restart(); #endif switch (command.originalType) { @@ -95,7 +99,6 @@ namespace mROA.Implementation.Frontend break; case EMessageType.CancelRequest: HandleCancelRequest((command.parced as CancelRequest)!); - break; default: continue; diff --git a/mROA/Implementation/RemoteObjectBase.cs b/mROA/Implementation/RemoteObjectBase.cs index 941cb52..c12a07b 100644 --- a/mROA/Implementation/RemoteObjectBase.cs +++ b/mROA/Implementation/RemoteObjectBase.cs @@ -82,14 +82,14 @@ namespace mROA.Implementation var response = await responseRequestTask; - if (response.parced is FinalCommandExecution successResponse) + if (response.Deserialized is FinalCommandExecution successResponse) { localTokenSource.Cancel(); return successResponse.Result!; } localTokenSource.Cancel(); - throw (response.parced as ExceptionCommandExecution)!.GetException(); + throw (response.Deserialized as ExceptionCommandExecution)!.GetException(); } protected async Task CallAsync(int methodId, object?[]? parameters = null, @@ -103,7 +103,7 @@ namespace mROA.Implementation var localTokenSource = new CancellationTokenSource(); - var responceRequestTask = _representationModule.GetSingle( + var responseRequestTask = _representationModule.GetSingle( m => m.MessageType is EMessageType.FinishedCommandExecution or EMessageType.ExceptionCommandExecution, localTokenSource.Token, m => m.MessageType is EMessageType.FinishedCommandExecution ? typeof(FinalCommandExecution) : null, @@ -124,16 +124,16 @@ namespace mROA.Implementation }).ContinueWith(_ => localTokenSource.Cancel()); }); - var responseRequest = await responceRequestTask; + var responseRequest = await responseRequestTask; #if TRACE Console.WriteLine($"Handling message"); #endif - switch (responseRequest.originalType) + switch (responseRequest.MessageType) { case EMessageType.FinishedCommandExecution: return; case EMessageType.ExceptionCommandExecution: - throw (responseRequest.parced as ExceptionCommandExecution)!.GetException(); + throw (responseRequest.Deserialized as ExceptionCommandExecution)!.GetException(); } } diff --git a/mROA/Implementation/RepresentationModule.cs b/mROA/Implementation/RepresentationModule.cs index 2e80308..acf8878 100644 --- a/mROA/Implementation/RepresentationModule.cs +++ b/mROA/Implementation/RepresentationModule.cs @@ -1,5 +1,7 @@ using System; using System.Collections.Generic; +using System.Linq; +using System.Runtime.CompilerServices; using System.Threading; using System.Threading.Tasks; using mROA.Abstract; @@ -24,72 +26,48 @@ namespace mROA.Implementation } } - - + public int Id => (_interaction ?? throw new NullReferenceException("Interaction is not initialized")) .ConnectionId; - public Task<(object parced, EMessageType originalType)> GetSingle(Predicate rule, CancellationToken token, params Func[] converter) + public async Task<(object? Deserialized, EMessageType MessageType)> GetSingle( + Predicate rule, + CancellationToken token = default, params Func[] converter) { - throw new NotImplementedException(); - } - - public IAsyncEnumerable<(object parced, EMessageType originalType)> GetStream(Predicate rule, CancellationToken token, params Func[] converter) - { - throw new NotImplementedException(); - } - - public async Task GetMessageAsync(Guid? requestId, EMessageType? messageType, - CancellationToken token = default) - { - if (_serialization == null) - throw new NullReferenceException("Serialization toolkit is not initialized"); - - var rawMessage = await GetRawMessage(requestId, messageType, token); - return _serialization.Deserialize(rawMessage)!; - } - - public T GetMessage(Guid? requestId = null, EMessageType? messageType = null) - { - if (_serialization == null) - throw new NullReferenceException("Serialization toolkit is not initialized"); - - var rawMessage = GetRawMessage(requestId, messageType).GetAwaiter().GetResult(); - return _serialization.Deserialize(rawMessage)!; - } - - public async Task GetRawMessage(Guid? requestId = null, EMessageType? messageType = null, - CancellationToken token = default) - { - if (_interaction == null) - throw new NullReferenceException("Interaction toolkit is not initialized"); - - var fromBuffer = - _interaction.FirstByFilter(message => - (requestId is null || message.Id == requestId) && - (messageType is null || message.MessageType == messageType)); - - if (fromBuffer == null) + var writer = _interaction.ReceiveChanel.Writer; + await foreach (var message in _interaction.ReceiveChanel.Reader.ReadAllAsync(token)) { - while (token.IsCancellationRequested == false) + if (!rule(message)) { - var message = await _interaction.GetNextMessageReceiving(); - if ((requestId is not null && message.Id != requestId) || - (messageType is not null && message.MessageType != messageType)) - continue; - - _interaction.HandleMessage(message); - return message.Data; + await writer.WriteAsync(message, token); + continue; } + + var type = converter.Select(i => i(message)).First(i => i != null)!; + var deserialized = _serialization.Deserialize(message.Data, type); + return (deserialized, message.MessageType); } - if (fromBuffer == null) + return (null, EMessageType.Unknown); + } + + public async IAsyncEnumerable<(object parced, EMessageType originalType)> GetStream( + Predicate rule, [EnumeratorCancellation] CancellationToken token = default, + params Func[] converter) + { + var writer = _interaction?.ReceiveChanel.Writer; + await foreach (var message in _interaction.ReceiveChanel.Reader.ReadAllAsync(token)) { - return Array.Empty(); - } + if (!rule(message)) + { + await writer.WriteAsync(message, token); + continue; + } - _interaction.HandleMessage(fromBuffer); - return fromBuffer.Data; + var type = converter.Select(i => i(message)).First(i => i != null)!; + var deserialized = _serialization.Deserialize(message.Data, type); + yield return (deserialized, message.MessageType)!; + } } public async Task PostCallMessageAsync(Guid id, EMessageType eMessageType, T payload) where T : notnull diff --git a/mROA/Implementation/StreamExtractor.cs b/mROA/Implementation/StreamExtractor.cs index 53d345d..5e748c0 100644 --- a/mROA/Implementation/StreamExtractor.cs +++ b/mROA/Implementation/StreamExtractor.cs @@ -21,7 +21,7 @@ namespace mROA.Implementation _serializationToolkit = serializationToolkit; } - public event Action MessageReceived; + public Action MessageReceived = _ => { }; private ushort ReadMessageLength() {