From 619992da806e25aa924607c5e1e732f2f9e47679 Mon Sep 17 00:00:00 2001 From: Alexandru Macocian Date: Fri, 12 Mar 2021 14:05:55 +0100 Subject: [PATCH] Switched server resources to Slim.ServiceManager. Changed HttpRoutingHandler to use the general Slim.ServiceManager. Changed WebsocketRoutingHandler to a similar approach as HttpRoutingHandler. Implemented IRunOnStartup interface for handlers. --- MTSC.UnitTests/E2ETests.cs | 43 ++- MTSC.UnitTests/EchoWebsocketModule.cs | 15 +- MTSC.UnitTests/EchoWebsocketModule2.cs | 24 ++ MTSC.UnitTests/ResourcesTests.cs | 35 --- .../RoutingModules/HelloWorldMessage.cs | 10 + .../HelloWorldMessageConverter.cs | 27 ++ .../RoutingModules/HelloWorldModule.cs | 24 ++ .../RoutingModules/AdhocConverter.cs | 26 ++ .../RoutingModules/ISetWebsocketContext.cs | 12 + .../IWebsocketMessageConverter.cs | 8 + .../WebsocketMessageConvertAttribute.cs | 21 ++ .../RoutingModules/WebsocketRouteBase.cs | 238 +++++++++++---- MTSC/MTSC.csproj | 6 +- MTSC/ServerSide/ClientData.cs | 2 +- .../ServerSide/Handlers/HttpRoutingHandler.cs | 70 ++--- MTSC/ServerSide/Handlers/IRunOnStartup.cs | 10 + .../Handlers/WebsocketRoutingHandler.cs | 158 ++++++---- MTSC/ServerSide/Server.cs | 275 ++++++++++++------ 18 files changed, 702 insertions(+), 302 deletions(-) create mode 100644 MTSC.UnitTests/EchoWebsocketModule2.cs delete mode 100644 MTSC.UnitTests/ResourcesTests.cs create mode 100644 MTSC.UnitTests/RoutingModules/HelloWorldMessage.cs create mode 100644 MTSC.UnitTests/RoutingModules/HelloWorldMessageConverter.cs create mode 100644 MTSC.UnitTests/RoutingModules/HelloWorldModule.cs create mode 100644 MTSC/Common/WebSockets/RoutingModules/AdhocConverter.cs create mode 100644 MTSC/Common/WebSockets/RoutingModules/ISetWebsocketContext.cs create mode 100644 MTSC/Common/WebSockets/RoutingModules/IWebsocketMessageConverter.cs create mode 100644 MTSC/Common/WebSockets/RoutingModules/WebsocketMessageConvertAttribute.cs create mode 100644 MTSC/ServerSide/Handlers/IRunOnStartup.cs diff --git a/MTSC.UnitTests/E2ETests.cs b/MTSC.UnitTests/E2ETests.cs index b666aa3..124dc4e 100644 --- a/MTSC.UnitTests/E2ETests.cs +++ b/MTSC.UnitTests/E2ETests.cs @@ -41,15 +41,9 @@ namespace MTSC.UnitTests .WithReadTimeout(TimeSpan.FromMilliseconds(1000)) .WithClientCertificate(false) .AddHandler(new WebsocketRoutingHandler() - .AddRoute("echo", new EchoWebsocketModule() - .WithReceiveTemplateProvider((message) => Encoding.UTF8.GetString(message.Data)) - .WithSendTemplateProvider((s) => - { - WebsocketMessage websocketMessage = new WebsocketMessage(); - websocketMessage.Data = Encoding.UTF8.GetBytes(s); - websocketMessage.Opcode = WebsocketMessage.Opcodes.Text; - return websocketMessage; - }))) + .AddRoute("echo") + .AddRoute("echo2") + .AddRoute("hello-world")) .AddHandler(new HttpRoutingHandler() .AddRoute(HttpMessage.HttpMethods.Get, "") .AddRoute(HttpMessage.HttpMethods.Get, "query") @@ -291,17 +285,38 @@ namespace MTSC.UnitTests } [TestMethod] - public void EchoWebsocket() + [DataRow("echo")] + [DataRow("echo2")] + public async Task EchoWebsocket(string endpoint) { - byte[] bytes = new byte[100]; + var bytes = new byte[100]; ClientWebSocket client = new ClientWebSocket(); - client.ConnectAsync(new Uri("ws://localhost:800/echo"), CancellationToken.None).Wait(); - client.SendAsync(Encoding.ASCII.GetBytes("Hello world!"), WebSocketMessageType.Text, true, CancellationToken.None).Wait(); - client.ReceiveAsync(bytes, CancellationToken.None).Wait(); + await client.ConnectAsync(new Uri($"ws://localhost:800/{endpoint}"), CancellationToken.None); + await client.SendAsync(Encoding.ASCII.GetBytes("Hello world!"), WebSocketMessageType.Text, true, CancellationToken.None); + await client.ReceiveAsync(bytes, CancellationToken.None); var resultString = Encoding.ASCII.GetString(bytes, 0, 12); Assert.AreEqual(resultString, "Hello world!"); } + [TestMethod] + public async Task HelloWorldWebsocket() + { + var bytes = new byte[100]; + var client = new ClientWebSocket(); + await client.ConnectAsync(new Uri($"ws://localhost:800/hello-world"), CancellationToken.None); + await client.SendAsync(Encoding.ASCII.GetBytes("Hello world!"), WebSocketMessageType.Text, true, CancellationToken.None); + await client.ReceiveAsync(bytes, CancellationToken.None); + var resultString = Encoding.ASCII.GetString(bytes, 0, 12); + Assert.AreEqual(resultString, "Hello world!"); + + client = new ClientWebSocket(); + await client.ConnectAsync(new Uri($"ws://localhost:800/hello-world"), CancellationToken.None); + await client.SendAsync(Encoding.ASCII.GetBytes("Something else"), WebSocketMessageType.Text, true, CancellationToken.None); + await client.ReceiveAsync(bytes, CancellationToken.None); + resultString = Encoding.ASCII.GetString(bytes, 0, 16); + Assert.AreEqual(resultString, "Not hello world!"); + } + [ClassCleanup] public static void CleanupServer() { diff --git a/MTSC.UnitTests/EchoWebsocketModule.cs b/MTSC.UnitTests/EchoWebsocketModule.cs index e4ffea3..ca9a848 100644 --- a/MTSC.UnitTests/EchoWebsocketModule.cs +++ b/MTSC.UnitTests/EchoWebsocketModule.cs @@ -1,29 +1,24 @@ using MTSC.Common.WebSockets.RoutingModules; -using MTSC.ServerSide; -using MTSC.ServerSide.Handlers; namespace MTSC.UnitTests { public class EchoWebsocketModule : WebsocketRouteBase { - public override void ConnectionClosed(Server server, WebsocketRoutingHandler handler, ClientData client) + public override void ConnectionClosed() { - } - public override void ConnectionInitialized(Server server, WebsocketRoutingHandler handler, ClientData client) + public override void ConnectionInitialized() { - } - public override void HandleReceivedMessage(Server server, WebsocketRoutingHandler handler, ClientData client, string message) + public override void HandleReceivedMessage(string message) { - this.SendMessage(message, client, handler); + this.SendMessage(message); } - public override void Tick(Server server, WebsocketRoutingHandler handler) + public override void Tick() { - } } } diff --git a/MTSC.UnitTests/EchoWebsocketModule2.cs b/MTSC.UnitTests/EchoWebsocketModule2.cs new file mode 100644 index 0000000..49d099d --- /dev/null +++ b/MTSC.UnitTests/EchoWebsocketModule2.cs @@ -0,0 +1,24 @@ +using MTSC.Common.WebSockets.RoutingModules; + +namespace MTSC.UnitTests +{ + public class EchoWebsocketModule2 : WebsocketRouteBase + { + public override void ConnectionClosed() + { + } + + public override void ConnectionInitialized() + { + } + + public override void HandleReceivedMessage(byte[] message) + { + this.SendMessage(message); + } + + public override void Tick() + { + } + } +} diff --git a/MTSC.UnitTests/ResourcesTests.cs b/MTSC.UnitTests/ResourcesTests.cs deleted file mode 100644 index 52d9580..0000000 --- a/MTSC.UnitTests/ResourcesTests.cs +++ /dev/null @@ -1,35 +0,0 @@ -using Microsoft.VisualStudio.TestTools.UnitTesting; -using MTSC.Exceptions; -using MTSC.Logging; -using MTSC.ServerSide.Handlers; -using MTSC.ServerSide.UsageMonitors; - -namespace MTSC.UnitTests -{ - [TestClass] - public class ResourcesTests - { - public static ServerSide.Server server; - [ClassInitialize] - public static void InitializeServer(TestContext testContext) - { - server = new ServerSide.Server(); - } - [TestMethod] - public void AddAndGetResource() - { - StringResource resource = new StringResource { Value = "hello" }; - server.WithResource(resource); - server.AddHandler(new HttpHandler()) - .AddExceptionHandler(new ExceptionConsoleLogger()) - .AddLogger(new ConsoleLogger()) - .AddServerUsageMonitor(new TickrateEnforcer()); - var gotResource = server.GetResource(); - Assert.AreEqual(gotResource, resource); - Assert.IsNotNull(server.GetExceptionHandler()); - Assert.IsNotNull(server.GetHandler()); - Assert.IsNotNull(server.GetLogger()); - Assert.IsNotNull(server.GetServerUsageMonitor()); - } - } -} diff --git a/MTSC.UnitTests/RoutingModules/HelloWorldMessage.cs b/MTSC.UnitTests/RoutingModules/HelloWorldMessage.cs new file mode 100644 index 0000000..87117c8 --- /dev/null +++ b/MTSC.UnitTests/RoutingModules/HelloWorldMessage.cs @@ -0,0 +1,10 @@ +using MTSC.Common.WebSockets.RoutingModules; + +namespace MTSC.UnitTests.RoutingModules +{ + [WebsocketMessageConvert(typeof(HelloWorldMessageConverter))] + public class HelloWorldMessage + { + public bool HelloWorld { get; set; } + } +} diff --git a/MTSC.UnitTests/RoutingModules/HelloWorldMessageConverter.cs b/MTSC.UnitTests/RoutingModules/HelloWorldMessageConverter.cs new file mode 100644 index 0000000..c8c9501 --- /dev/null +++ b/MTSC.UnitTests/RoutingModules/HelloWorldMessageConverter.cs @@ -0,0 +1,27 @@ +using MTSC.Common.WebSockets; +using MTSC.Common.WebSockets.RoutingModules; +using System.Text; + +namespace MTSC.UnitTests.RoutingModules +{ + public class HelloWorldMessageConverter : IWebsocketMessageConverter + { + public HelloWorldMessage ConvertFromWebsocketMessage(WebsocketMessage websocketMessage) + { + var str = Encoding.UTF8.GetString(websocketMessage.Data); + return new HelloWorldMessage + { + HelloWorld = str == "Hello world!" + }; + } + + public WebsocketMessage ConvertToWebsocketMessage(HelloWorldMessage message) + { + return new WebsocketMessage + { + Data = message.HelloWorld ? Encoding.UTF8.GetBytes("Hello world!") : Encoding.UTF8.GetBytes("Not hello world!"), + Opcode = WebsocketMessage.Opcodes.Text + }; + } + } +} diff --git a/MTSC.UnitTests/RoutingModules/HelloWorldModule.cs b/MTSC.UnitTests/RoutingModules/HelloWorldModule.cs new file mode 100644 index 0000000..a8457f8 --- /dev/null +++ b/MTSC.UnitTests/RoutingModules/HelloWorldModule.cs @@ -0,0 +1,24 @@ +using MTSC.Common.WebSockets.RoutingModules; + +namespace MTSC.UnitTests.RoutingModules +{ + public class HelloWorldModule : WebsocketRouteBase + { + public override void ConnectionClosed() + { + } + + public override void ConnectionInitialized() + { + } + + public override void HandleReceivedMessage(HelloWorldMessage message) + { + this.SendMessage(message); + } + + public override void Tick() + { + } + } +} diff --git a/MTSC/Common/WebSockets/RoutingModules/AdhocConverter.cs b/MTSC/Common/WebSockets/RoutingModules/AdhocConverter.cs new file mode 100644 index 0000000..2ca0074 --- /dev/null +++ b/MTSC/Common/WebSockets/RoutingModules/AdhocConverter.cs @@ -0,0 +1,26 @@ +using System; + +namespace MTSC.Common.WebSockets.RoutingModules +{ + internal class AdhocConverter : IWebsocketMessageConverter + { + private readonly Func convertFrom; + private readonly Func convertTo; + + public AdhocConverter(Func convertFrom, Func convertTo) + { + this.convertFrom = convertFrom; + this.convertTo = convertTo; + } + + public T ConvertFromWebsocketMessage(WebsocketMessage websocketMessage) + { + return this.convertFrom(websocketMessage); + } + + public WebsocketMessage ConvertToWebsocketMessage(T message) + { + return this.convertTo(message); + } + } +} diff --git a/MTSC/Common/WebSockets/RoutingModules/ISetWebsocketContext.cs b/MTSC/Common/WebSockets/RoutingModules/ISetWebsocketContext.cs new file mode 100644 index 0000000..da1d536 --- /dev/null +++ b/MTSC/Common/WebSockets/RoutingModules/ISetWebsocketContext.cs @@ -0,0 +1,12 @@ +using MTSC.ServerSide; +using MTSC.ServerSide.Handlers; + +namespace MTSC.Common.WebSockets.RoutingModules +{ + internal interface ISetWebsocketContext + { + void SetServer(Server server); + void SetHandler(WebsocketRoutingHandler websocketRoutingHandler); + void SetClient(ClientData clientData); + } +} diff --git a/MTSC/Common/WebSockets/RoutingModules/IWebsocketMessageConverter.cs b/MTSC/Common/WebSockets/RoutingModules/IWebsocketMessageConverter.cs new file mode 100644 index 0000000..7ad947b --- /dev/null +++ b/MTSC/Common/WebSockets/RoutingModules/IWebsocketMessageConverter.cs @@ -0,0 +1,8 @@ +namespace MTSC.Common.WebSockets.RoutingModules +{ + public interface IWebsocketMessageConverter + { + T ConvertFromWebsocketMessage(WebsocketMessage websocketMessage); + WebsocketMessage ConvertToWebsocketMessage(T message); + } +} diff --git a/MTSC/Common/WebSockets/RoutingModules/WebsocketMessageConvertAttribute.cs b/MTSC/Common/WebSockets/RoutingModules/WebsocketMessageConvertAttribute.cs new file mode 100644 index 0000000..5ceae76 --- /dev/null +++ b/MTSC/Common/WebSockets/RoutingModules/WebsocketMessageConvertAttribute.cs @@ -0,0 +1,21 @@ +using System; +using System.Linq; + +namespace MTSC.Common.WebSockets.RoutingModules +{ + [AttributeUsage(AttributeTargets.Class)] + public sealed class WebsocketMessageConvertAttribute : Attribute + { + public Type ConverterType { get; } + + public WebsocketMessageConvertAttribute(Type type) + { + if (!type.GetInterfaces().Any(x => x.IsGenericType && x.GetGenericTypeDefinition() == typeof(IWebsocketMessageConverter<>))) + { + throw new InvalidOperationException($"{type.FullName} is not a {typeof(IWebsocketMessageConverter<>).FullName}"); + } + + this.ConverterType = type; + } + } +} diff --git a/MTSC/Common/WebSockets/RoutingModules/WebsocketRouteBase.cs b/MTSC/Common/WebSockets/RoutingModules/WebsocketRouteBase.cs index ac3da40..230c9f7 100644 --- a/MTSC/Common/WebSockets/RoutingModules/WebsocketRouteBase.cs +++ b/MTSC/Common/WebSockets/RoutingModules/WebsocketRouteBase.cs @@ -1,99 +1,233 @@ using MTSC.ServerSide; using MTSC.ServerSide.Handlers; using System; +using System.Linq; +using System.Text; namespace MTSC.Common.WebSockets.RoutingModules { - public abstract class WebsocketRouteBase + public abstract class WebsocketRouteBase : ISetWebsocketContext { - public void CallConnectionInitialized(Server server, WebsocketRoutingHandler handler, ClientData client) + protected Server Server { get; private set; } + protected WebsocketRoutingHandler WebsocketRoutingHandler { get; private set; } + protected ClientData ClientData { get; private set; } + + public void CallConnectionInitialized() { - ConnectionInitialized(server, handler, client); + ConnectionInitialized(); } - public void CallHandleReceivedMessage(Server server, WebsocketRoutingHandler handler, ClientData client, WebsocketMessage receivedMessage) + public void CallHandleReceivedMessage(WebsocketMessage receivedMessage) { - HandleReceivedMessage(server, handler, client, receivedMessage); + this.HandleReceivedMessage(receivedMessage); } - public void CallConnectionClosed(Server server, WebsocketRoutingHandler handler, ClientData client) + public void CallConnectionClosed() { - ConnectionClosed(server, handler, client); + this.ConnectionClosed(); } - - public void SendMessage(WebsocketMessage message, ClientData client, WebsocketRoutingHandler handler) + public void SendMessage(WebsocketMessage message) { - handler.QueueMessage(client, message); + this.WebsocketRoutingHandler.QueueMessage(this.ClientData, message); } - public abstract void ConnectionInitialized(Server server, WebsocketRoutingHandler handler, ClientData client); - public abstract void HandleReceivedMessage(Server server, WebsocketRoutingHandler handler, ClientData client, WebsocketMessage receivedMessage); - public abstract void ConnectionClosed(Server server, WebsocketRoutingHandler handler, ClientData client); - public abstract void Tick(Server server, WebsocketRoutingHandler handler); + public abstract void ConnectionInitialized(); + public abstract void HandleReceivedMessage(WebsocketMessage receivedMessage); + public abstract void ConnectionClosed(); + public abstract void Tick(); + + void ISetWebsocketContext.SetServer(Server server) + { + this.Server = server; + } + void ISetWebsocketContext.SetHandler(WebsocketRoutingHandler websocketRoutingHandler) + { + this.WebsocketRoutingHandler = websocketRoutingHandler; + } + void ISetWebsocketContext.SetClient(ClientData clientData) + { + this.ClientData = clientData; + } + + internal static IWebsocketMessageConverter GetStringAdhocConverter() + { + return new AdhocConverter( + convertFrom: message => (T)(Encoding.UTF8.GetString(message.Data) as object), + convertTo: message => + { + var str = (string)(message as object); + return new WebsocketMessage + { + Data = Encoding.UTF8.GetBytes(str), + Opcode = WebsocketMessage.Opcodes.Text + }; + }); + } + internal static IWebsocketMessageConverter GetByteArrayAdhocConverter() + { + return new AdhocConverter( + convertFrom: message => (T)(message.Data as object), + convertTo: message => + { + return new WebsocketMessage + { + Data = (byte[])(message as object), + Opcode = WebsocketMessage.Opcodes.Binary + }; + }); + } } public abstract class WebsocketRouteBase : WebsocketRouteBase { - private Func receiveTemplate; + private readonly static object cachedLock = new object(); + private static IWebsocketMessageConverter CachedConverter { get; set; } - public WebsocketRouteBase(Func receiveTemplate) + public sealed override void HandleReceivedMessage(WebsocketMessage receivedMessage) { - this.receiveTemplate = receiveTemplate; - } + lock (cachedLock) + { + if (CachedConverter is null) + { + CachedConverter = ImplementConverter(); + } + } - public WebsocketRouteBase() + this.HandleReceivedMessage(CachedConverter.ConvertFromWebsocketMessage(receivedMessage)); + } + public abstract void HandleReceivedMessage(TReceive message); + + private static bool MatchesRequiredType(WebsocketMessageConvertAttribute attribute) { + if (attribute.ConverterType.GetInterfaces().Any(x => x.IsGenericType && x.GetGenericTypeDefinition() == typeof(IWebsocketMessageConverter))) + { + return false; + } + return true; } - - public WebsocketRouteBase WithReceiveTemplateProvider(Func templateProvider) + private static IWebsocketMessageConverter ImplementConverter() { - this.receiveTemplate = templateProvider; - return this; - } + var converterType = typeof(TReceive) + .GetCustomAttributes(true) + .OfType() + .Where(MatchesRequiredType) + .Select(attribute => attribute.ConverterType) + .FirstOrDefault(); + if (converterType is null) + { + if (typeof(TReceive) == typeof(string)) + { + return GetStringAdhocConverter(); + } + else if (typeof(TReceive) == typeof(byte[])) + { + return GetByteArrayAdhocConverter(); + } - public override void HandleReceivedMessage(Server server, WebsocketRoutingHandler handler, ClientData client, WebsocketMessage receivedMessage) - { - HandleReceivedMessage(server, handler, client, receiveTemplate.Invoke(receivedMessage)); - } + throw new InvalidOperationException($"No converter found for type {typeof(TReceive).FullName}"); + } - public abstract void HandleReceivedMessage(Server server, WebsocketRoutingHandler handler, ClientData client, TReceive message); + var converter = Activator.CreateInstance(converterType) as IWebsocketMessageConverter; + return converter; + } } public abstract class WebsocketRouteBase : WebsocketRouteBase { - private Func receiveTemplate; - private Func sendTemplate; + private readonly static object recLock = new object(), sendLock = new object(); + private static IWebsocketMessageConverter CachedReceiveConverter { get; set; } + private static IWebsocketMessageConverter CachedSendConverter { get; set; } - public WebsocketRouteBase(Func receiveTemplate, Func sendTemplate) + public void SendMessage(TSend message) { - this.receiveTemplate = receiveTemplate; - this.sendTemplate = sendTemplate; - } + lock (sendLock) + { + if (CachedSendConverter is null) + { + CachedSendConverter = ImplementSendConverter(); + } + } - public WebsocketRouteBase() + base.SendMessage(CachedSendConverter.ConvertToWebsocketMessage(message)); + } + public sealed override void HandleReceivedMessage(WebsocketMessage receivedMessage) { - + lock (recLock) + { + if (CachedReceiveConverter is null) + { + CachedReceiveConverter = ImplementReceiveConverter(); + } + } + + this.HandleReceivedMessage(CachedReceiveConverter.ConvertFromWebsocketMessage(receivedMessage)); } + public abstract void HandleReceivedMessage(TReceive message); - public WebsocketRouteBase WithReceiveTemplateProvider(Func templateProvider) + private static bool MatchesRequiredReceiveType(WebsocketMessageConvertAttribute attribute) { - this.receiveTemplate = templateProvider; - return this; - } + if (attribute.ConverterType.GetInterfaces().Any(x => x.IsGenericType && x.GetGenericTypeDefinition() == typeof(IWebsocketMessageConverter))) + { + return false; + } - public WebsocketRouteBase WithSendTemplateProvider(Func templateProvider) + return true; + } + private static bool MatchesRequiredSendType(WebsocketMessageConvertAttribute attribute) { - this.sendTemplate = templateProvider; - return this; - } + if (attribute.ConverterType.GetInterfaces().Any(x => x.IsGenericType && x.GetGenericTypeDefinition() == typeof(IWebsocketMessageConverter))) + { + return false; + } - public void SendMessage(TSend message, ClientData client, WebsocketRoutingHandler handler) + return true; + } + private static IWebsocketMessageConverter ImplementSendConverter() { - base.SendMessage(sendTemplate.Invoke(message), client, handler); - } + var converterType = typeof(TSend) + .GetCustomAttributes(true) + .OfType() + .Where(MatchesRequiredSendType) + .Select(attribute => attribute.ConverterType) + .FirstOrDefault(); + if (converterType is null) + { + if (typeof(TSend) == typeof(string)) + { + return GetStringAdhocConverter(); + } + else if (typeof(TSend) == typeof(byte[])) + { + return GetByteArrayAdhocConverter(); + } - public override void HandleReceivedMessage(Server server, WebsocketRoutingHandler handler, ClientData client, WebsocketMessage receivedMessage) + throw new InvalidOperationException($"No converter found for type {typeof(TSend).FullName}"); + } + + var converter = Activator.CreateInstance(converterType) as IWebsocketMessageConverter; + return converter; + } + private static IWebsocketMessageConverter ImplementReceiveConverter() { - HandleReceivedMessage(server, handler, client, receiveTemplate.Invoke(receivedMessage)); - } + var converterType = typeof(TReceive) + .GetCustomAttributes(true) + .OfType() + .Where(MatchesRequiredReceiveType) + .Select(attribute => attribute.ConverterType) + .FirstOrDefault(); + if (converterType is null) + { + if (typeof(TReceive) == typeof(string)) + { + return GetStringAdhocConverter(); + } + else if (typeof(TReceive) == typeof(byte[])) + { + return GetByteArrayAdhocConverter(); + } - public abstract void HandleReceivedMessage(Server server, WebsocketRoutingHandler handler, ClientData client, TReceive message); + throw new InvalidOperationException($"No converter found for type {typeof(TReceive).FullName}"); + } + + var converter = Activator.CreateInstance(converterType) as IWebsocketMessageConverter; + return converter; + } } } diff --git a/MTSC/MTSC.csproj b/MTSC/MTSC.csproj index e623936..39e818f 100644 --- a/MTSC/MTSC.csproj +++ b/MTSC/MTSC.csproj @@ -5,13 +5,13 @@ netcoreapp2.1;net48;netstandard2.0;netcoreapp3.1;net5.0 - 3.1 + 3.2 latest Alexandru-Victor Macocian MTSC Modular TCP Server and Client - 0.3.1 - 0.3.1 + 0.3.2 + 0.3.2 true AnyCPU;x64 https://github.com/AlexMacocian/MTSC diff --git a/MTSC/ServerSide/ClientData.cs b/MTSC/ServerSide/ClientData.cs index eb26197..07a4d31 100644 --- a/MTSC/ServerSide/ClientData.cs +++ b/MTSC/ServerSide/ClientData.cs @@ -11,7 +11,7 @@ namespace MTSC.ServerSide /// public class ClientData : IDisposable, IActiveClient, IQueueHolder { - private ProducerConsumerQueue messageQueue = new ProducerConsumerQueue(); + private readonly ProducerConsumerQueue messageQueue = new ProducerConsumerQueue(); public TcpClient TcpClient { get; } /// diff --git a/MTSC/ServerSide/Handlers/HttpRoutingHandler.cs b/MTSC/ServerSide/Handlers/HttpRoutingHandler.cs index 9206c79..57d1ca5 100644 --- a/MTSC/ServerSide/Handlers/HttpRoutingHandler.cs +++ b/MTSC/ServerSide/Handlers/HttpRoutingHandler.cs @@ -1,21 +1,16 @@ using MTSC.Common.Http; using MTSC.Common.Http.RoutingModules; using MTSC.Common.Http.Telemetry; -using MTSC.Exceptions; -using Slim; using System; using System.Collections.Concurrent; using System.Collections.Generic; -using System.Text; using static MTSC.Common.Http.HttpMessage; namespace MTSC.ServerSide.Handlers { - public sealed class HttpRoutingHandler : IHandler + public sealed class HttpRoutingHandler : IHandler, IRunOnStartup { private static readonly Func alwaysEnabled = (server, request, client) => RouteEnablerResponse.Accept; - - private readonly ServiceManager serviceManager = new ServiceManager(); private readonly ConcurrentQueue> messageOutQueue = new ConcurrentQueue>(); private readonly List httpLoggers = new List(); private readonly Dictionary)>>(); - private bool initialized = false; - public TimeSpan FragmentsExpirationTime { get; set; } = TimeSpan.FromSeconds(15); public double MaximumRequestSize { get; set; } = double.MaxValue; @@ -32,16 +25,14 @@ namespace MTSC.ServerSide.Handlers { foreach (HttpMethods method in (HttpMethods[])Enum.GetValues(typeof(HttpMethods))) { - moduleDictionary[method] = new Dictionary)>(); } - - this.serviceManager.RegisterServiceManager(); } public HttpRoutingHandler AddHttpLogger(IHttpLogger logger) { - httpLoggers.Add(logger); + this.httpLoggers.Add(logger); return this; } public HttpRoutingHandler AddRoute( @@ -120,7 +111,7 @@ namespace MTSC.ServerSide.Handlers * Else, parse the messages into a partial request. If this causes an exception, let it throw out of the handler, * cause the handler needs at least valid headers to work. */ - PartialHttpRequest request = null; + PartialHttpRequest request; if (client.Resources.TryGetResource(out var fragmentedMessage)) { var bytesToBeAdded = message.MessageBytes.TrimTrailingNullBytes(); @@ -128,7 +119,7 @@ namespace MTSC.ServerSide.Handlers fragmentedMessage.PartialRequest.Body.Length + message.MessageLength > this.MaximumRequestSize) { - QueueResponse(client, new HttpResponse { StatusCode = StatusCodes.BadRequest, BodyString = $"Request exceeded [{MaximumRequestSize}] bytes!" }); + this.QueueResponse(client, new HttpResponse { StatusCode = StatusCodes.BadRequest, BodyString = $"Request exceeded [{MaximumRequestSize}] bytes!" }); client.ResetAffinityIfMe(this); client.Resources.RemoveResource(); client.Resources.RemoveResourceIfExists(); @@ -140,7 +131,7 @@ namespace MTSC.ServerSide.Handlers { if (message.MessageLength > this.MaximumRequestSize) { - QueueResponse(client, new HttpResponse { StatusCode = StatusCodes.BadRequest, BodyString = $"Request exceeded [{MaximumRequestSize}] bytes!" }); + this.QueueResponse(client, new HttpResponse { StatusCode = StatusCodes.BadRequest, BodyString = $"Request exceeded [{MaximumRequestSize}] bytes!" }); return false; } @@ -155,7 +146,7 @@ namespace MTSC.ServerSide.Handlers server.LogDebug("Returning 100-Continue"); var contResponse = new HttpResponse { StatusCode = HttpMessage.StatusCodes.Continue }; contResponse.Headers[HttpMessage.GeneralHeaders.Connection] = "keep-alive"; - QueueResponse(client, contResponse); + this.QueueResponse(client, contResponse); } } } @@ -174,12 +165,12 @@ namespace MTSC.ServerSide.Handlers client.Resources.RemoveResource(); client.Resources.RemoveResource(); client.ResetAffinityIfMe(this); - HandleCompleteRequest(client, server, request.ToRequest(), mapping.MappedModule, mapping.RouteEnabler); + this.HandleCompleteRequest(client, server, request.ToRequest(), mapping.MappedModule, mapping.RouteEnabler); return true; } else { - HandleIncompleteRequest(client, server, client.Resources.GetResource()); + this.HandleIncompleteRequest(server, client.Resources.GetResource()); return true; } } @@ -193,7 +184,7 @@ namespace MTSC.ServerSide.Handlers { var httpRequest = request.ToRequest(); client.ResetAffinityIfMe(this); - return HandleCompleteRequest(client, server, httpRequest, module, routeEnabler); + return this.HandleCompleteRequest(client, server, httpRequest, module, routeEnabler); } else { @@ -216,20 +207,9 @@ namespace MTSC.ServerSide.Handlers void IHandler.Tick(Server server) { - if (this.initialized is false) + while (this.messageOutQueue.Count > 0) { - this.initialized = true; - this.serviceManager.RegisterSingleton(typeof(Server), typeof(Server), (sp) => server); - this.serviceManager.RegisterSingleton(typeof(HttpRoutingHandler), typeof(HttpRoutingHandler), sp => this); - foreach(var resource in server.Resources.Values) - { - this.serviceManager.RegisterSingleton(resource.GetType(), resource.GetType(), (sp) => resource); - } - } - - while (messageOutQueue.Count > 0) - { - if (messageOutQueue.TryDequeue(out Tuple tuple)) + if (this.messageOutQueue.TryDequeue(out Tuple tuple)) { server.QueueMessage(tuple.Item1, tuple.Item2.GetPackedResponse(true)); } @@ -238,16 +218,27 @@ namespace MTSC.ServerSide.Handlers { if (client.Resources.TryGetResource(out var fragmentedMessage)) { - if ((DateTime.Now - fragmentedMessage.LastReceived) > FragmentsExpirationTime) + if ((DateTime.Now - fragmentedMessage.LastReceived) > this.FragmentsExpirationTime) { client.Resources.RemoveResource(); - QueueResponse(client, new HttpResponse { StatusCode = StatusCodes.BadRequest, BodyString = $"Request timed out in [{FragmentsExpirationTime.TotalMilliseconds}] ms!" }); + this.QueueResponse(client, new HttpResponse { StatusCode = StatusCodes.BadRequest, BodyString = $"Request timed out in [{this.FragmentsExpirationTime.TotalMilliseconds}] ms!" }); } } } } - private void HandleIncompleteRequest(ClientData client, Server server, FragmentedMessage fragmentedMessage) + void IRunOnStartup.OnStartup(Server server) + { + foreach (var routes in this.moduleDictionary.Values) + { + foreach ((var routeType, _) in routes.Values) + { + server.ServiceManager.RegisterTransient(routeType, routeType); + } + } + } + + private void HandleIncompleteRequest(Server server, FragmentedMessage fragmentedMessage) { fragmentedMessage.LastReceived = DateTime.Now; server.LogDebug("Incomplete request received!"); @@ -270,7 +261,7 @@ namespace MTSC.ServerSide.Handlers module.CallHandleRequest(request).ContinueWith((task) => { foreach (var httpLogger in this.httpLoggers) httpLogger.LogResponse(server, this, client, task.Result); - QueueResponse(client, task.Result); + this.QueueResponse(client, task.Result); }); } catch (Exception e) @@ -279,7 +270,7 @@ namespace MTSC.ServerSide.Handlers server.LogDebug("Stacktrace: " + e.StackTrace); var response = new HttpResponse() { StatusCode = StatusCodes.InternalServerError }; foreach (var httpLogger in this.httpLoggers) httpLogger.LogResponse(server, this, client, response); - QueueResponse(client, response); + this.QueueResponse(client, response); } return true; } @@ -291,7 +282,7 @@ namespace MTSC.ServerSide.Handlers { foreach (var httpLogger in this.httpLoggers) httpLogger.LogResponse(server, this, client, (routeEnablerResponse as RouteEnablerResponse.RouteEnablerResponseError).Response); - QueueResponse(client, (routeEnablerResponse as RouteEnablerResponse.RouteEnablerResponseError).Response); + this.QueueResponse(client, (routeEnablerResponse as RouteEnablerResponse.RouteEnablerResponseError).Response); return true; } else @@ -325,7 +316,7 @@ namespace MTSC.ServerSide.Handlers throw new InvalidOperationException($"Cannot create new route of type {routeType.FullName}. Not of type {typeof(HttpRouteBase).FullName}"); } - var module = this.serviceManager.GetService(routeType) as HttpRouteBase; + var module = server.ServiceManager.GetService(routeType) as HttpRouteBase; (module as ISetHttpContext).SetClientData(client); (module as ISetHttpContext).SetServer(server); (module as ISetHttpContext).SetHttpRoutingHandler(this); @@ -339,7 +330,6 @@ namespace MTSC.ServerSide.Handlers throw new InvalidOperationException($"{routeType.FullName} must be of type {typeof(HttpRouteBase).FullName}"); } - this.serviceManager.RegisterSingleton(routeType, routeType); this.moduleDictionary[method][uri] = (routeType, routeEnabler); } } diff --git a/MTSC/ServerSide/Handlers/IRunOnStartup.cs b/MTSC/ServerSide/Handlers/IRunOnStartup.cs new file mode 100644 index 0000000..800b35e --- /dev/null +++ b/MTSC/ServerSide/Handlers/IRunOnStartup.cs @@ -0,0 +1,10 @@ +namespace MTSC.ServerSide.Handlers +{ + /// + /// Implement this interface in handlers that need to run a procedure on server startup. + /// + public interface IRunOnStartup + { + void OnStartup(Server server); + } +} diff --git a/MTSC/ServerSide/Handlers/WebsocketRoutingHandler.cs b/MTSC/ServerSide/Handlers/WebsocketRoutingHandler.cs index 6c9f88f..6b4cc57 100644 --- a/MTSC/ServerSide/Handlers/WebsocketRoutingHandler.cs +++ b/MTSC/ServerSide/Handlers/WebsocketRoutingHandler.cs @@ -4,21 +4,22 @@ using MTSC.Common.WebSockets.RoutingModules; using System; using System.Collections.Concurrent; using System.Collections.Generic; +using System.Linq; using System.Security.Cryptography; using System.Text; namespace MTSC.ServerSide.Handlers { - public class WebsocketRoutingHandler : IHandler + public class WebsocketRoutingHandler : IHandler, IRunOnStartup { - private static Func alwaysEnabled = (server, message, client) => RouteEnablerResponse.Accept; - private static readonly string WebsocketHeaderAcceptKey = "Sec-WebSocket-Accept"; - private static readonly string WebsocketHeaderKey = "Sec-WebSocket-Key"; - private static readonly string WebsocketProtocolKey = "Sec-WebSocket-Protocol"; - private static readonly string WebsocketProtocolVersionKey = "Sec-WebSocket-Version"; - private static readonly string GlobalUniqueIdentifier = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"; - private static SHA1 sha1Provider = SHA1.Create(); - private static RNGCryptoServiceProvider rng = new RNGCryptoServiceProvider(); + private const string WebsocketHeaderAcceptKey = "Sec-WebSocket-Accept"; + private const string WebsocketHeaderKey = "Sec-WebSocket-Key"; + private const string WebsocketProtocolKey = "Sec-WebSocket-Protocol"; + private const string WebsocketProtocolVersionKey = "Sec-WebSocket-Version"; + private const string GlobalUniqueIdentifier = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"; + private static readonly Func alwaysEnabled = (server, message, client) => RouteEnablerResponse.Accept; + private static readonly SHA1 sha1Provider = SHA1.Create(); + private static readonly RNGCryptoServiceProvider rng = new RNGCryptoServiceProvider(); public enum SocketState { Initial, @@ -27,55 +28,71 @@ namespace MTSC.ServerSide.Handlers Closed } #region Fields - Dictionary)> moduleDictionary = - new Dictionary)>(); - ConcurrentQueue> messageQueue = new ConcurrentQueue>(); + private readonly Dictionary)> moduleDictionary = + new Dictionary)>(); + private readonly ConcurrentQueue> messageQueue = new ConcurrentQueue>(); #endregion #region Public Methods - public WebsocketRoutingHandler AddRoute(string uri, WebsocketRouteBase module) + public WebsocketRoutingHandler AddRoute(string uri) + where T : WebsocketRouteBase { - this.moduleDictionary.Add(uri, (module, alwaysEnabled)); + this.RegisterRoute(uri, typeof(T), alwaysEnabled); + return this; + } + public WebsocketRoutingHandler AddRoute( + string uri, + Func routeEnabler) + where T : WebsocketRouteBase + { + this.RegisterRoute(uri, typeof(T), routeEnabler); + return this; + } + public WebsocketRoutingHandler AddRoute(string uri, Type routeType) + { + this.RegisterRoute(uri, routeType, alwaysEnabled); return this; } - public WebsocketRoutingHandler AddRoute( string uri, - WebsocketRouteBase module, + Type routeType, Func routeEnabler) { - this.moduleDictionary.Add(uri, (module, routeEnabler)); + this.RegisterRoute(uri, routeType, routeEnabler); return this; } - public WebsocketRoutingHandler RemoveRoute(string uri) { - moduleDictionary.Remove(uri); + this.moduleDictionary.Remove(uri); return this; } public void QueueMessage(ClientData client, byte[] message, WebsocketMessage.Opcodes opcode = WebsocketMessage.Opcodes.Text) { - WebsocketMessage sendMessage = new WebsocketMessage(); - sendMessage.Data = message; - sendMessage.FIN = true; - sendMessage.Masked = false; + WebsocketMessage sendMessage = new WebsocketMessage + { + Data = message, + FIN = true, + Masked = false + }; rng.GetBytes(sendMessage.Mask); sendMessage.Opcode = opcode; - messageQueue.Enqueue(new Tuple(client, sendMessage)); + this.messageQueue.Enqueue(new Tuple(client, sendMessage)); } public void QueueMessage(ClientData client, WebsocketMessage message) { - messageQueue.Enqueue(new Tuple(client, message)); + this.messageQueue.Enqueue(new Tuple(client, message)); } public void CloseConnection(ClientData client) { - WebsocketMessage websocketMessage = new WebsocketMessage(); - websocketMessage.FIN = true; - websocketMessage.Opcode = WebsocketMessage.Opcodes.Close; - websocketMessage.Masked = false; - QueueMessage(client, websocketMessage); + WebsocketMessage websocketMessage = new WebsocketMessage + { + FIN = true, + Opcode = WebsocketMessage.Opcodes.Close, + Masked = false + }; + this.QueueMessage(client, websocketMessage); } #endregion #region Handler Implementation @@ -83,7 +100,7 @@ namespace MTSC.ServerSide.Handlers { if (client.Resources.TryGetResource(out WebsocketRouteBase route)) { - route.CallConnectionClosed(server, this, client); + route.CallConnectionClosed(); } } @@ -112,11 +129,12 @@ namespace MTSC.ServerSide.Handlers request.Headers[HttpMessage.GeneralHeaders.Connection].ToLower() == "upgrade" && request.Headers.ContainsHeader(WebsocketProtocolVersionKey) && request.Headers[WebsocketProtocolVersionKey] == "13") { - if (!moduleDictionary.ContainsKey(request.RequestURI)) + if (!this.moduleDictionary.ContainsKey(request.RequestURI)) { - QueueMessage(client, new HttpResponse { StatusCode = HttpMessage.StatusCodes.NotFound, BodyString = "URI not found" }.GetPackedResponse(true)); + this.QueueMessage(client, new HttpResponse { StatusCode = HttpMessage.StatusCodes.NotFound, BodyString = "URI not found" }.GetPackedResponse(true)); } - (var module, var routeEnabler) = moduleDictionary[request.RequestURI]; + + (var moduleType, var routeEnabler) = moduleDictionary[request.RequestURI]; var routeEnablerResponse = routeEnabler.Invoke(server, request.ToRequest(), client); if(routeEnablerResponse is RouteEnablerResponse.RouteEnablerResponseIgnore) { @@ -124,31 +142,45 @@ namespace MTSC.ServerSide.Handlers } else if(routeEnablerResponse is RouteEnablerResponse.RouteEnablerResponseError) { - QueueMessage(client, (routeEnablerResponse as RouteEnablerResponse.RouteEnablerResponseError).Response.GetPackedResponse(true)); + this.QueueMessage(client, (routeEnablerResponse as RouteEnablerResponse.RouteEnablerResponseError).Response.GetPackedResponse(true)); return true; } + /* * The RouteEnabler accepted the request. * Prepare the handshake string. */ - string base64Key = request.Headers[WebsocketHeaderKey]; + var base64Key = request.Headers[WebsocketHeaderKey]; base64Key = base64Key.Trim(); - string handshakeKey = base64Key + GlobalUniqueIdentifier; - string returnBase64Key = Convert.ToBase64String(sha1Provider.ComputeHash(Encoding.UTF8.GetBytes(handshakeKey))); + var handshakeKey = base64Key + GlobalUniqueIdentifier; + var returnBase64Key = Convert.ToBase64String(sha1Provider.ComputeHash(Encoding.UTF8.GetBytes(handshakeKey))); /* * Prepare the response. */ - HttpResponse response = new HttpResponse(); - response.StatusCode = HttpMessage.StatusCodes.SwitchingProtocols; + var response = new HttpResponse + { + StatusCode = HttpMessage.StatusCodes.SwitchingProtocols + }; response.Headers[HttpMessage.GeneralHeaders.Upgrade] = "websocket"; response.Headers[HttpMessage.GeneralHeaders.Connection] = "Upgrade"; response.Headers[WebsocketHeaderAcceptKey] = returnBase64Key; server.QueueMessage(client, response.GetPackedResponse(false)); client.Resources.SetResource(SocketState.Established); server.LogDebug("Websocket initialized " + client.TcpClient.Client.RemoteEndPoint.ToString()); + /* + * Create and assign route module to client. + */ + if (server.ServiceManager.GetService(moduleType) is not WebsocketRouteBase module) + { + throw new InvalidOperationException($"Unexpected error during websocket module initialization. {moduleType.FullName} is not of type {typeof(WebsocketRouteBase).FullName}"); + } + + (module as ISetWebsocketContext).SetClient(client); + (module as ISetWebsocketContext).SetHandler(this); + (module as ISetWebsocketContext).SetServer(server); client.Resources.SetResource(module); - module.CallConnectionInitialized(server, this, client); + module.CallConnectionInitialized(); return true; } else @@ -158,7 +190,7 @@ namespace MTSC.ServerSide.Handlers } else if (socketState == SocketState.Established) { - WebsocketMessage receivedMessage = null; + WebsocketMessage receivedMessage; try { receivedMessage = new WebsocketMessage(message.MessageBytes); @@ -172,16 +204,18 @@ namespace MTSC.ServerSide.Handlers if (receivedMessage.Opcode == WebsocketMessage.Opcodes.Close) { client.ToBeRemoved = true; - WebsocketMessage closeFrame = new WebsocketMessage(); - closeFrame.Opcode = WebsocketMessage.Opcodes.Close; - QueueMessage(client, closeFrame); + WebsocketMessage closeFrame = new WebsocketMessage + { + Opcode = WebsocketMessage.Opcodes.Close + }; + this.QueueMessage(client, closeFrame); return true; } else { try { - client.Resources.GetResource().CallHandleReceivedMessage(server, this, client, receivedMessage); + client.Resources.GetResource().CallHandleReceivedMessage(receivedMessage); return true; } catch(Exception e) @@ -209,26 +243,48 @@ namespace MTSC.ServerSide.Handlers void IHandler.Tick(Server server) { - foreach((var module, var _) in moduleDictionary.Values) + foreach(var client in server.Clients) { - module.Tick(server, this); + if (client.Resources.TryGetResource(out var route)) + { + route.Tick(); + } } - while (messageQueue.Count > 0) + + while (this.messageQueue.Count > 0) { - if (messageQueue.TryDequeue(out Tuple tuple)) + if (this.messageQueue.TryDequeue(out Tuple tuple)) { server.QueueMessage(tuple.Item1, tuple.Item2.GetMessageBytes()); if (tuple.Item2.Opcode == WebsocketMessage.Opcodes.Close) { if (tuple.Item1.Resources.TryGetResource(out var route)) { - route.CallConnectionClosed(server, this, tuple.Item1); + route.CallConnectionClosed(); } tuple.Item1.ToBeRemoved = true; } } } } + + void IRunOnStartup.OnStartup(Server server) + { + foreach((var routeType, _) in this.moduleDictionary.Values) + { + server.ServiceManager.RegisterTransient(routeType, routeType); + } + } #endregion + + private void RegisterRoute(string uri, Type moduleType, Func routeEnabler) + { + if (!typeof(WebsocketRouteBase).IsAssignableFrom(moduleType)) + { + throw new InvalidOperationException($"{moduleType.FullName} must be of type {typeof(WebsocketRouteBase).FullName}"); + } + + this.moduleDictionary.Add(uri, (moduleType, routeEnabler)); + } } } diff --git a/MTSC/ServerSide/Server.cs b/MTSC/ServerSide/Server.cs index b14c989..71ca24b 100644 --- a/MTSC/ServerSide/Server.cs +++ b/MTSC/ServerSide/Server.cs @@ -5,6 +5,7 @@ using MTSC.ServerSide.Handlers; using MTSC.ServerSide.Resources; using MTSC.ServerSide.Schedulers; using MTSC.ServerSide.UsageMonitors; +using Slim; using System; using System.Collections.Generic; using System.Linq; @@ -24,22 +25,23 @@ namespace MTSC.ServerSide public sealed class Server { #region Fields - bool running; - X509Certificate2 certificate; - TcpListener listener; - ProducerConsumerQueue addQueue = new ProducerConsumerQueue(); - List clients = new List(); - List toRemove = new List(); - List handlers = new List(); - List loggers = new List(); - List exceptionHandlers = new List(); - List serverUsageMonitors = new List(); - ProducerConsumerQueue<(ClientData, byte[])> messageOutQueue = new ProducerConsumerQueue<(ClientData, byte[])>(); + private bool running; + private X509Certificate2 certificate; + private TcpListener listener; + private readonly ProducerConsumerQueue addQueue = new ProducerConsumerQueue(); + private readonly List clients = new List(); + private readonly List toRemove = new List(); + private readonly List handlers = new List(); + private readonly List loggers = new List(); + private readonly List exceptionHandlers = new List(); + private readonly List serverUsageMonitors = new List(); + private readonly ProducerConsumerQueue<(ClientData, byte[])> messageOutQueue = new ProducerConsumerQueue<(ClientData, byte[])>(); + private readonly IServiceManager serviceManager = new ServiceManager(); #endregion #region Private Properties - private IConsumerQueue _ConsumerClientQueue { get => addQueue; } - private IProducerQueue _ProducerClientQueue { get => addQueue; } - private IConsumerQueue<(ClientData, byte[])> _ConsumerMessageOutQueue { get => messageOutQueue; } + private IConsumerQueue ConsumerClientQueue { get => addQueue; } + private IProducerQueue ProducerClientQueue { get => addQueue; } + private IConsumerQueue<(ClientData, byte[])> ConsumerMessageOutQueue { get => messageOutQueue; } #endregion #region Public Properties /// @@ -93,11 +95,11 @@ namespace MTSC.ServerSide /// /// List of clients currently connected to the server. /// - public IReadOnlyCollection Clients { get => clients.AsReadOnly(); } + public IReadOnlyCollection Clients { get => this.clients.AsReadOnly(); } /// - /// Dictionary of resources + /// for configuring and retrieving services. /// - public Dictionary Resources { get; } = new Dictionary(); + public IServiceManager ServiceManager { get => this.serviceManager; } #endregion #region Constructors /// @@ -105,7 +107,6 @@ namespace MTSC.ServerSide /// public Server() { - } /// /// Creates an instance of server. @@ -128,6 +129,90 @@ namespace MTSC.ServerSide #endregion #region Public Methods /// + /// Adds a service with transient lifetime. + /// + /// This server object. + public Server AddTransientService() + where TService : TInterface + where TInterface : class + { + this.serviceManager.RegisterTransient(); + return this; + } + /// + /// Adds a service with transient lifetime. Registers the service for all the interfaces it implements. + /// + /// This server object. + public Server AddTransientService() + where TService : class + { + this.serviceManager.RegisterTransient(); + return this; + } + /// + /// Adds a service with transient lifetime. + /// + /// This server object. + public Server AddTransientService(Func serviceFactory) + where TService : TInterface + where TInterface : class + { + this.serviceManager.RegisterTransient(serviceFactory); + return this; + } + /// + /// Adds a service with transient lifetime. Registers the service for all the interfaces it implements. + /// + /// This server object. + public Server AddTransientService(Func serviceFactory) + where TService : class + { + this.serviceManager.RegisterTransient(serviceFactory); + return this; + } + /// + /// Adds a service with singleton lifetime. + /// + /// This server object. + public Server AddSingletonService() + where TService : TInterface + where TInterface : class + { + this.serviceManager.RegisterSingleton(); + return this; + } + /// + /// Adds a service with singleton lifetime. Registers the service for all the interfaces it implements. + /// + /// This server object. + public Server AddSingletonService() + where TService : class + { + this.serviceManager.RegisterSingleton(); + return this; + } + /// + /// Adds a service with singleton lifetime. + /// + /// This server object. + public Server AddSingletonService(Func serviceFactory) + where TService : TInterface + where TInterface : class + { + this.serviceManager.RegisterSingleton(serviceFactory); + return this; + } + /// + /// Adds a service with singleton lifetime. Registers the service for all the interfaces it implements. + /// + /// This server object. + public Server AddSingletonService(Func serviceFactory) + where TService : class + { + this.serviceManager.RegisterSingleton(serviceFactory); + return this; + } + /// /// Sets the property. /// /// @@ -158,16 +243,6 @@ namespace MTSC.ServerSide return this; } /// - /// Adds a resource to the server. - /// - /// resource to be added to the server. - /// . - public Server WithResource(IResource resource) - { - Resources[resource.GetType()] = resource; - return this; - } - /// /// Requests that the client provides a certificate. /// /// @@ -254,7 +329,7 @@ namespace MTSC.ServerSide /// This server object. public Server AddHandler(IHandler handler) { - handlers.Add(handler); + this.handlers.Add(handler); return this; } /// @@ -264,7 +339,7 @@ namespace MTSC.ServerSide /// This server object. public Server AddLogger(ILogger logger) { - loggers.Add(logger); + this.loggers.Add(logger); return this; } /// @@ -274,7 +349,7 @@ namespace MTSC.ServerSide /// This server object. public Server AddExceptionHandler(IExceptionHandler handler) { - exceptionHandlers.Add(handler); + this.exceptionHandlers.Add(handler); return this; } /// @@ -284,17 +359,18 @@ namespace MTSC.ServerSide /// This server object. public Server AddServerUsageMonitor(IServerUsageMonitor serverUsageMonitor) { - serverUsageMonitors.Add(serverUsageMonitor); + this.serverUsageMonitors.Add(serverUsageMonitor); return this; } /// - /// Get the resource of provided type + /// Gets a required service. /// - /// + /// Type of the service used during registration. /// - public T GetResource() + public T GetService() + where T : class { - return (T)Resources[typeof(T)]; + return this.serviceManager.GetService(); } /// /// Get handler of provided type @@ -303,7 +379,7 @@ namespace MTSC.ServerSide /// public T GetHandler() where T : class { - foreach(var handler in handlers) + foreach(var handler in this.handlers) { if(handler is T) { @@ -319,7 +395,7 @@ namespace MTSC.ServerSide /// public T GetExceptionHandler() where T : class { - foreach (var exceptionHandler in exceptionHandlers) + foreach (var exceptionHandler in this.exceptionHandlers) { if (exceptionHandler is T) { @@ -335,7 +411,7 @@ namespace MTSC.ServerSide /// public T GetLogger() where T : class { - foreach (var logger in loggers) + foreach (var logger in this.loggers) { if (logger is T) { @@ -351,7 +427,7 @@ namespace MTSC.ServerSide /// public T GetServerUsageMonitor() where T : class { - foreach (var serverMonitor in serverUsageMonitors) + foreach (var serverMonitor in this.serverUsageMonitors) { if(serverMonitor is T) { @@ -375,7 +451,7 @@ namespace MTSC.ServerSide /// Message to be logged public void Log(string log) { - foreach (ILogger logger in loggers) + foreach (ILogger logger in this.loggers) { if (logger.Log(log)) { @@ -389,7 +465,7 @@ namespace MTSC.ServerSide /// Debug message to be logged public void LogDebug(string debugMessage) { - foreach (ILogger logger in loggers) + foreach (ILogger logger in this.loggers) { if (logger.LogDebug(debugMessage)) { @@ -402,11 +478,18 @@ namespace MTSC.ServerSide /// public void Run() { - listener?.Stop(); - listener = new TcpListener(IPAddress.Any, Port); - listener.Start(); - running = true; - Log("Server started on: " + listener.LocalEndpoint.ToString()); + this.listener?.Stop(); + this.listener = new TcpListener(IPAddress.Any, Port); + this.listener.Start(); + this.running = true; + this.Log("Server started on: " + this.listener.LocalEndpoint.ToString()); + foreach(var toBeRunOnStartup in this.handlers.OfType()) + { + toBeRunOnStartup.OnStartup(this); + } + + this.serviceManager.RegisterServiceManager(); + this.serviceManager.RegisterSingleton(sp => this); DateTime startLoopTime; while (running) { @@ -417,11 +500,11 @@ namespace MTSC.ServerSide */ try { - CheckAndRemoveInactiveClients(); + this.CheckAndRemoveInactiveClients(); } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in exceptionHandlers) + foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) { if (exceptionHandler.HandleException(e)) { @@ -432,23 +515,23 @@ namespace MTSC.ServerSide /* * Check and gather messages from clients and place them in their queues. */ - CheckAndGatherMessages(); + this.CheckAndGatherMessages(); /* * Check if the server has any pending connections. * If it has a new connection, process it. */ try { - while (listener.Pending()) + while (this.listener.Pending()) { - TcpClient tcpClient = listener.AcceptTcpClient(); - ClientData clientStruct = new ClientData(tcpClient); - Task.Run(() => AcceptClient(clientStruct)); + var tcpClient = this.listener.AcceptTcpClient(); + var clientStruct = new ClientData(tcpClient); + Task.Run(() => this.AcceptClient(clientStruct)); } } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in exceptionHandlers) + foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) { if (exceptionHandler.HandleException(e)) { @@ -459,11 +542,11 @@ namespace MTSC.ServerSide /* * Add all accepted clients to the list */ - while(this._ConsumerClientQueue.TryDequeue(out var client)) + while(this.ConsumerClientQueue.TryDequeue(out var client)) { - Log("Accepted new connection: " + client.TcpClient.Client.RemoteEndPoint.ToString()); - clients.Add(client); - foreach (IHandler handler in handlers) + this.Log("Accepted new connection: " + client.TcpClient.Client.RemoteEndPoint.ToString()); + this.clients.Add(client); + foreach (IHandler handler in this.handlers) { if (handler.HandleClient(this, client)) { @@ -476,30 +559,30 @@ namespace MTSC.ServerSide * Call the scheduler to handle all received messages and distribute them to the handlers */ - Scheduler.ScheduleHandling( - clients + this.Scheduler.ScheduleHandling( + this.clients .Where(client => client.ToBeRemoved == false) .Select(client => (client, (client as IQueueHolder).ConsumerQueue)) .ToList(), - HandleClientMessages); + this.HandleClientMessages); /* * Iterate through all the handlers, running periodic operations. */ - foreach(IHandler handler in handlers) + foreach(IHandler handler in this.handlers) { - TickHandler(handler); + this.TickHandler(handler); } /* * Check if there are messages queued to be sent. */ - SendQueuedMessages(); + this.SendQueuedMessages(); /* * Call the usage monitors and let them scale or determine current resource usage. */ - foreach (IServerUsageMonitor usageMonitor in serverUsageMonitors) + foreach (IServerUsageMonitor usageMonitor in this.serverUsageMonitors) { try { @@ -507,7 +590,7 @@ namespace MTSC.ServerSide } catch(Exception e) { - foreach (IExceptionHandler exceptionHandler in exceptionHandlers) + foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) { if (exceptionHandler.HandleException(e)) { @@ -517,8 +600,8 @@ namespace MTSC.ServerSide } } } - listener.Stop(); - foreach (var client in Clients) + this.listener.Stop(); + foreach (var client in this.Clients) { try { @@ -527,13 +610,13 @@ namespace MTSC.ServerSide } catch (Exception e) { - foreach (var handler in exceptionHandlers) + foreach (var handler in this.exceptionHandlers) { handler.HandleException(e); } } } - listener = null; + this.listener = null; } /// /// Runs the server async. @@ -556,7 +639,7 @@ namespace MTSC.ServerSide #region Private Methods private void SendQueuedMessages() { - while (this._ConsumerMessageOutQueue.TryDequeue(out var tuple)) + while (this.ConsumerMessageOutQueue.TryDequeue(out var tuple)) { (var client, var bytes) = tuple; if (client.TcpClient.Available > 0) @@ -569,19 +652,19 @@ namespace MTSC.ServerSide try { Message sendMessage = CommunicationPrimitives.BuildMessage(bytes); - for (int i = handlers.Count - 1; i >= 0; i--) + for (int i = this.handlers.Count - 1; i >= 0; i--) { IHandler handler = handlers[i]; handler.HandleSendMessage(this, client, ref sendMessage); } CommunicationPrimitives.SendMessage(client.TcpClient, sendMessage, client.SslStream); (client as IActiveClient).UpdateLastActivity(); - LogDebug("Sent message to " + client.TcpClient.Client.RemoteEndPoint.ToString() + + this.LogDebug("Sent message to " + client.TcpClient.Client.RemoteEndPoint.ToString() + "\nMessage length: " + sendMessage.MessageLength); } catch(Exception e) { - foreach (IExceptionHandler exceptionHandler in exceptionHandlers) + foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) { if (exceptionHandler.HandleException(e)) { @@ -593,27 +676,27 @@ namespace MTSC.ServerSide } private void CheckAndRemoveInactiveClients() { - foreach (ClientData client in Clients) + foreach (ClientData client in this.Clients) { if (!client.TcpClient.Connected || client.ToBeRemoved) { - toRemove.Add(client); + this.toRemove.Add(client); } } - foreach (ClientData client in toRemove) + foreach (ClientData client in this.toRemove) { try { - foreach (IHandler handler in handlers) + foreach (IHandler handler in this.handlers) { handler.ClientRemoved(this, client); } - LogDebug("Client removed: " + client.TcpClient?.Client?.RemoteEndPoint?.ToString()); + this.LogDebug("Client removed: " + client.TcpClient?.Client?.RemoteEndPoint?.ToString()); client.Dispose(); } catch(Exception e) { - foreach (IExceptionHandler exceptionHandler in exceptionHandlers) + foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) { if (exceptionHandler.HandleException(e)) { @@ -621,9 +704,9 @@ namespace MTSC.ServerSide } } } - clients.Remove(client); + this.clients.Remove(client); } - toRemove.Clear(); + this.toRemove.Clear(); } private void TickHandler(IHandler handler) { @@ -633,7 +716,7 @@ namespace MTSC.ServerSide } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in exceptionHandlers) + foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) { if (exceptionHandler.HandleException(e)) { @@ -644,7 +727,7 @@ namespace MTSC.ServerSide } private void CheckAndGatherMessages() { - foreach(var client in Clients) + foreach(var client in this.Clients) { if (client.TcpClient.Available > 0 && !(client as IActiveClient).ReadingData) { @@ -683,11 +766,11 @@ namespace MTSC.ServerSide { if (client.Affinity is null) { - HandleClientMessage(client, message); + this.HandleClientMessage(client, message); } else { - AffinityHandleClientMessage(client, message); + this.AffinityHandleClientMessage(client, message); } } } @@ -700,7 +783,7 @@ namespace MTSC.ServerSide } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in exceptionHandlers) + foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) { if (exceptionHandler.HandleException(e)) { @@ -714,7 +797,7 @@ namespace MTSC.ServerSide } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in exceptionHandlers) + foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) { if (exceptionHandler.HandleException(e)) { @@ -725,7 +808,7 @@ namespace MTSC.ServerSide } private void HandleClientMessage(ClientData client, Message message) { - foreach (IHandler handler in handlers) + foreach (IHandler handler in this.handlers) { try { @@ -736,7 +819,7 @@ namespace MTSC.ServerSide } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in exceptionHandlers) + foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) { if (exceptionHandler.HandleException(e)) { @@ -745,7 +828,7 @@ namespace MTSC.ServerSide } } } - foreach (IHandler handler in handlers) + foreach (IHandler handler in this.handlers) { try { @@ -756,7 +839,7 @@ namespace MTSC.ServerSide } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in exceptionHandlers) + foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) { if (exceptionHandler.HandleException(e)) { @@ -784,7 +867,7 @@ namespace MTSC.ServerSide /* * Client authenticated in the alloted time */ - this._ProducerClientQueue.Enqueue(client); + this.ProducerClientQueue.Enqueue(client); } else { @@ -793,12 +876,12 @@ namespace MTSC.ServerSide } else { - this._ProducerClientQueue.Enqueue(client); + this.ProducerClientQueue.Enqueue(client); } } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in exceptionHandlers) + foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) { if (exceptionHandler.HandleException(e)) {