diff --git a/MTSC.UnitTests/E2ETests.cs b/MTSC.UnitTests/E2ETests.cs index 124dc4e..51edd77 100644 --- a/MTSC.UnitTests/E2ETests.cs +++ b/MTSC.UnitTests/E2ETests.cs @@ -43,8 +43,12 @@ namespace MTSC.UnitTests .AddHandler(new WebsocketRoutingHandler() .AddRoute("echo") .AddRoute("echo2") - .AddRoute("hello-world")) + .AddRoute("hello-world") + .WithHeartbeatEnabled(true) + .WithHeartbeatFrequency(TimeSpan.FromMilliseconds(100))) .AddHandler(new HttpRoutingHandler() + .WithReturn500OnException(true) + .AddRoute(HttpMessage.HttpMethods.Get, "throw") .AddRoute(HttpMessage.HttpMethods.Get, "") .AddRoute(HttpMessage.HttpMethods.Get, "query") .AddRoute(HttpMessage.HttpMethods.Post, "echo") @@ -60,6 +64,16 @@ namespace MTSC.UnitTests .WithSslAuthenticationTimeout(TimeSpan.FromMilliseconds(100)); Server.RunAsync(); } + + [TestMethod] + public async Task ServerReturns500OnError() + { + HttpClient httpClient = new HttpClient(); + httpClient.BaseAddress = new Uri("http://localhost:800"); + var response = await httpClient.GetAsync("throw"); + Assert.AreEqual(response.StatusCode, HttpStatusCode.InternalServerError); + } + [TestMethod] public async Task ServerParsesRequestAndResponse() { diff --git a/MTSC.UnitTests/RoutingModules/ExceptionThrowingModule.cs b/MTSC.UnitTests/RoutingModules/ExceptionThrowingModule.cs new file mode 100644 index 0000000..ac5ed25 --- /dev/null +++ b/MTSC.UnitTests/RoutingModules/ExceptionThrowingModule.cs @@ -0,0 +1,14 @@ +using MTSC.Common.Http; +using MTSC.Common.Http.RoutingModules; +using System.Threading.Tasks; + +namespace MTSC.UnitTests.RoutingModules +{ + public class ExceptionThrowingModule : HttpRouteBase + { + public override Task HandleRequest(HttpRequest request) + { + throw new System.NotImplementedException(); + } + } +} diff --git a/MTSC/Common/Http/RoutingModules/HttpRouteBase.cs b/MTSC/Common/Http/RoutingModules/HttpRouteBase.cs index 5de02db..2d3b55e 100644 --- a/MTSC/Common/Http/RoutingModules/HttpRouteBase.cs +++ b/MTSC/Common/Http/RoutingModules/HttpRouteBase.cs @@ -8,13 +8,32 @@ namespace MTSC.Common.Http.RoutingModules { public abstract class HttpRouteBase : ISetHttpContext { + private static HttpResponse InternalServerError500 { get; } = + new HttpResponse + { + StatusCode = HttpMessage.StatusCodes.InternalServerError, + BodyString = "An exception ocurred while processing the request" + }; + public ClientData ClientData { get; private set; } public HttpRoutingHandler HttpRoutingHandler { get; private set; } public Server Server { get; private set; } public async Task CallHandleRequest(HttpRequest request) { - return await this.HandleRequest(request); + try + { + return await this.HandleRequest(request); + } + catch + { + if (this.HttpRoutingHandler.Return500OnException is true) + { + return InternalServerError500; + } + + throw; + } } public abstract Task HandleRequest(HttpRequest request); diff --git a/MTSC/Common/TimeoutSuppressedStream.cs b/MTSC/Common/TimeoutSuppressedStream.cs index e034da7..07e226b 100644 --- a/MTSC/Common/TimeoutSuppressedStream.cs +++ b/MTSC/Common/TimeoutSuppressedStream.cs @@ -8,23 +8,22 @@ namespace MTSC.Common { public class TimeoutSuppressedStream : Stream { - NetworkStream innerStream; + private readonly NetworkStream innerStream; public TimeoutSuppressedStream(TcpClient tcpClient) { - innerStream = tcpClient.GetStream(); + this.innerStream = tcpClient.GetStream(); } public override int Read(byte[] buffer, int offset, int count) { try { - return innerStream.Read(buffer, offset, count); + return this.innerStream.Read(buffer, offset, count); } catch (IOException lException) { - SocketException lInnerException = lException.InnerException as SocketException; - if (lInnerException != null && lInnerException.SocketErrorCode == SocketError.TimedOut) + if (lException.InnerException is SocketException lInnerException && lInnerException.SocketErrorCode == SocketError.TimedOut) { // Normally, a simple TimeOut on the read will cause SslStream to flip its lid // However, if we suppress the IOException and just return 0 bytes read, this is ok. @@ -37,38 +36,38 @@ namespace MTSC.Common } - public override bool CanRead => innerStream.CanRead; - public override bool CanSeek => innerStream.CanSeek; - public override bool CanTimeout => innerStream.CanTimeout; - public override bool CanWrite => innerStream.CanWrite; - public virtual bool DataAvailable => innerStream.DataAvailable; - public override long Length => innerStream.Length; - public override IAsyncResult BeginRead(byte[] buffer, int offset, int size, AsyncCallback callback, object state) => innerStream.BeginRead(buffer, offset, size, callback, state); - public override IAsyncResult BeginWrite(byte[] buffer, int offset, int size, AsyncCallback callback, object state) => innerStream.BeginWrite(buffer, offset, size, callback, state); - public override int EndRead(IAsyncResult asyncResult) => innerStream.EndRead(asyncResult); - public override void EndWrite(IAsyncResult asyncResult) => innerStream.EndWrite(asyncResult); - public override void Flush() => innerStream.Flush(); - public override Task FlushAsync(CancellationToken cancellationToken) => innerStream.FlushAsync(cancellationToken); - public override long Seek(long offset, SeekOrigin origin) => innerStream.Seek(offset, origin); - public override void SetLength(long value) => innerStream.SetLength(value); - public override void Write(byte[] buffer, int offset, int count) => innerStream.Write(buffer, offset, count); + public override bool CanRead => this.innerStream.CanRead; + public override bool CanSeek => this.innerStream.CanSeek; + public override bool CanTimeout => this.innerStream.CanTimeout; + public override bool CanWrite => this.innerStream.CanWrite; + public virtual bool DataAvailable => this.innerStream.DataAvailable; + public override long Length => this.innerStream.Length; + public override IAsyncResult BeginRead(byte[] buffer, int offset, int size, AsyncCallback callback, object state) => this.innerStream.BeginRead(buffer, offset, size, callback, state); + public override IAsyncResult BeginWrite(byte[] buffer, int offset, int size, AsyncCallback callback, object state) => this.innerStream.BeginWrite(buffer, offset, size, callback, state); + public override int EndRead(IAsyncResult asyncResult) => this.innerStream.EndRead(asyncResult); + public override void EndWrite(IAsyncResult asyncResult) => this.innerStream.EndWrite(asyncResult); + public override void Flush() => this.innerStream.Flush(); + public override Task FlushAsync(CancellationToken cancellationToken) => this.innerStream.FlushAsync(cancellationToken); + public override long Seek(long offset, SeekOrigin origin) => this.innerStream.Seek(offset, origin); + public override void SetLength(long value) => this.innerStream.SetLength(value); + public override void Write(byte[] buffer, int offset, int count) => this.innerStream.Write(buffer, offset, count); public override long Position { - get { return innerStream.Position; } - set { innerStream.Position = value; } + get { return this.innerStream.Position; } + set { this.innerStream.Position = value; } } public override int ReadTimeout { - get { return innerStream.ReadTimeout; } - set { innerStream.ReadTimeout = value; } + get { return this.innerStream.ReadTimeout; } + set { this.innerStream.ReadTimeout = value; } } public override int WriteTimeout { - get { return innerStream.WriteTimeout; } - set { innerStream.WriteTimeout = value; } + get { return this.innerStream.WriteTimeout; } + set { this.innerStream.WriteTimeout = value; } } } } diff --git a/MTSC/MTSC.csproj b/MTSC/MTSC.csproj index 13bc7e4..d664945 100644 --- a/MTSC/MTSC.csproj +++ b/MTSC/MTSC.csproj @@ -5,13 +5,13 @@ netcoreapp2.1;net48;netstandard2.0;netcoreapp3.1;net5.0 - 3.2.3 + 3.3 latest Alexandru-Victor Macocian MTSC Modular TCP Server and Client - 0.3.2.3 - 0.3.2.3 + 3.3.0.0 + 3.3.0.0 true AnyCPU;x64 https://github.com/AlexMacocian/MTSC diff --git a/MTSC/ServerSide/Handlers/HttpHandler.cs b/MTSC/ServerSide/Handlers/HttpHandler.cs index c87b098..2f81905 100644 --- a/MTSC/ServerSide/Handlers/HttpHandler.cs +++ b/MTSC/ServerSide/Handlers/HttpHandler.cs @@ -14,14 +14,13 @@ namespace MTSC.ServerSide.Handlers /// public sealed class HttpHandler : IHandler { - private static readonly string urlEncodedHeader = "application/x-www-form-urlencoded"; - private static readonly string multipartHeader = "multipart/form-data"; #region Fields - private List httpLoggers = new List(); - private List httpModules = new List(); - private ConcurrentQueue> messageOutQueue = new ConcurrentQueue>(); + private readonly List httpLoggers = new(); + private readonly List httpModules = new(); + private readonly ConcurrentQueue> messageOutQueue = new(); #endregion #region Public Properties + public bool Return500OnException { get; set; } = true; public TimeSpan FragmentsExpirationTime { get; set; } = TimeSpan.FromSeconds(15); public double MaximumRequestSize { get; set; } = 15000; #endregion @@ -42,6 +41,11 @@ namespace MTSC.ServerSide.Handlers this.MaximumRequestSize = size; return this; } + public HttpHandler WithReturn500OnException(bool return500OnException) + { + this.Return500OnException = return500OnException; + return this; + } /// /// The amount of time fragments are kept in the buffer before being discarded. /// @@ -180,7 +184,7 @@ namespace MTSC.ServerSide.Handlers client.Resources.RemoveResource(); } - HttpResponse response = new HttpResponse(); + HttpResponse response = new(); if (request.Headers.ContainsHeader(HttpMessage.GeneralHeaders.Connection) && request.Headers[HttpMessage.GeneralHeaders.Connection].ToLower() == "close") { @@ -194,13 +198,27 @@ namespace MTSC.ServerSide.Handlers foreach (var httpLogger in this.httpLoggers) httpLogger.LogRequest(server, this, client, request); foreach (IHttpModule module in httpModules) { - if (module.HandleRequest(server, this, client, request, ref response)) + try { - break; + if (module.HandleRequest(server, this, client, request, ref response)) + { + break; + } + } + catch + { + if (this.Return500OnException) + { + response.StatusCode = StatusCodes.InternalServerError; + response.BodyString = "An exception ocurred while processing the request"; + break; + } + + throw; } } foreach (var httpLogger in this.httpLoggers) httpLogger.LogResponse(server, this, client, response); - QueueResponse(client, response); + this.QueueResponse(client, response); return true; } /// diff --git a/MTSC/ServerSide/Handlers/HttpRoutingHandler.cs b/MTSC/ServerSide/Handlers/HttpRoutingHandler.cs index 57d1ca5..632f55d 100644 --- a/MTSC/ServerSide/Handlers/HttpRoutingHandler.cs +++ b/MTSC/ServerSide/Handlers/HttpRoutingHandler.cs @@ -12,14 +12,13 @@ namespace MTSC.ServerSide.Handlers { private static readonly Func alwaysEnabled = (server, request, client) => RouteEnablerResponse.Accept; private readonly ConcurrentQueue> messageOutQueue = new ConcurrentQueue>(); - private readonly List httpLoggers = new List(); + private readonly List httpLoggers = new(); private readonly Dictionary)>> moduleDictionary = - new Dictionary)>>(); + Func)>> moduleDictionary = new(); public TimeSpan FragmentsExpirationTime { get; set; } = TimeSpan.FromSeconds(15); public double MaximumRequestSize { get; set; } = double.MaxValue; + public bool Return500OnException { get; set; } = true; public HttpRoutingHandler() { @@ -30,6 +29,11 @@ namespace MTSC.ServerSide.Handlers } } + public HttpRoutingHandler WithReturn500OnException(bool return500OnException) + { + this.Return500OnException = return500OnException; + return this; + } public HttpRoutingHandler AddHttpLogger(IHttpLogger logger) { this.httpLoggers.Add(logger); diff --git a/MTSC/ServerSide/Handlers/WebsocketRoutingHandler.cs b/MTSC/ServerSide/Handlers/WebsocketRoutingHandler.cs index ca253b4..d470f10 100644 --- a/MTSC/ServerSide/Handlers/WebsocketRoutingHandler.cs +++ b/MTSC/ServerSide/Handlers/WebsocketRoutingHandler.cs @@ -28,11 +28,28 @@ namespace MTSC.ServerSide.Handlers Closed } #region Fields + private byte[] emptyData = new byte[0]; + private DateTime ellapsedTime = DateTime.Now; + private DateTime previousHeartbeatProc = DateTime.Now; private readonly Dictionary)> moduleDictionary = new Dictionary)>(); private readonly ConcurrentQueue> messageQueue = new ConcurrentQueue>(); #endregion + #region Properties + public bool HeartbeatEnabled { get; set; } + public TimeSpan HeartbeatFrequency { get; set; } + #endregion #region Public Methods + public WebsocketRoutingHandler WithHeartbeatFrequency(TimeSpan heartbeatFrequency) + { + this.HeartbeatFrequency = heartbeatFrequency; + return this; + } + public WebsocketRoutingHandler WithHeartbeatEnabled(bool heartbeatEnabled) + { + this.HeartbeatEnabled = heartbeatEnabled; + return this; + } public WebsocketRoutingHandler AddRoute(string uri) where T : WebsocketRouteBase { @@ -212,6 +229,11 @@ namespace MTSC.ServerSide.Handlers this.QueueMessage(client, new byte[0] ,WebsocketMessage.Opcodes.Pong); return true; } + else if (receivedMessage.Opcode == WebsocketMessage.Opcodes.Pong) + { + // According to https://developer.mozilla.org/en-US/docs/Web/API/WebSockets_API/Writing_WebSocket_servers, ignore unrequested pings. + return true; + } else { try @@ -244,11 +266,18 @@ namespace MTSC.ServerSide.Handlers void IHandler.Tick(Server server) { + this.ellapsedTime = DateTime.Now; foreach(var client in server.Clients) { if (client.Resources.TryGetResource(out var route)) { route.Tick(); + if (this.HeartbeatEnabled is true && + (this.ellapsedTime - this.previousHeartbeatProc) > this.HeartbeatFrequency) + { + this.previousHeartbeatProc = DateTime.Now; + this.QueueMessage(client, emptyData, WebsocketMessage.Opcodes.Ping); + } } } diff --git a/MTSC/ServerSide/Server.cs b/MTSC/ServerSide/Server.cs index c2ec50c..05a867d 100644 --- a/MTSC/ServerSide/Server.cs +++ b/MTSC/ServerSide/Server.cs @@ -28,14 +28,14 @@ namespace MTSC.ServerSide 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 ProducerConsumerQueue addQueue = new(); + private readonly List clients = new(); + private readonly List toRemove = new(); + private readonly List handlers = new(); + private readonly List loggers = new(); + private readonly List exceptionHandlers = new(); + private readonly List serverUsageMonitors = new(); + private readonly ProducerConsumerQueue<(ClientData, byte[])> messageOutQueue = new(); private readonly IServiceManager serviceManager = new ServiceManager(); #endregion #region Private Properties @@ -540,13 +540,7 @@ namespace MTSC.ServerSide } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) - { - if (exceptionHandler.HandleException(e)) - { - break; - } - } + this.HandleException(e); } /* * Check and gather messages from clients and place them in their queues. @@ -567,13 +561,7 @@ namespace MTSC.ServerSide } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) - { - if (exceptionHandler.HandleException(e)) - { - break; - } - } + this.HandleException(e); } /* * Add all accepted clients to the list @@ -626,13 +614,7 @@ namespace MTSC.ServerSide } catch(Exception e) { - foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) - { - if (exceptionHandler.HandleException(e)) - { - break; - } - } + this.HandleException(e); } } } @@ -700,13 +682,7 @@ namespace MTSC.ServerSide } catch(Exception e) { - foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) - { - if (exceptionHandler.HandleException(e)) - { - break; - } - } + this.HandleException(e); } } } @@ -732,13 +708,7 @@ namespace MTSC.ServerSide } catch(Exception e) { - foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) - { - if (exceptionHandler.HandleException(e)) - { - break; - } - } + this.HandleException(e); } this.clients.Remove(client); } @@ -752,13 +722,7 @@ namespace MTSC.ServerSide } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) - { - if (exceptionHandler.HandleException(e)) - { - break; - } - } + this.HandleException(e); } } private void CheckAndGatherMessages() @@ -819,13 +783,7 @@ namespace MTSC.ServerSide } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) - { - if (exceptionHandler.HandleException(e)) - { - break; - } - } + this.HandleException(e); } try { @@ -833,13 +791,7 @@ namespace MTSC.ServerSide } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) - { - if (exceptionHandler.HandleException(e)) - { - break; - } - } + this.HandleException(e); } } private void HandleClientMessage(ClientData client, Message message) @@ -855,13 +807,7 @@ namespace MTSC.ServerSide } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) - { - if (exceptionHandler.HandleException(e)) - { - break; - } - } + this.HandleException(e); } } foreach (IHandler handler in this.handlers) @@ -875,13 +821,7 @@ namespace MTSC.ServerSide } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) - { - if (exceptionHandler.HandleException(e)) - { - break; - } - } + this.HandleException(e); } } } @@ -891,7 +831,7 @@ namespace MTSC.ServerSide { if (this.certificate != null) { - SslStream sslStream = new SslStream(client.SafeNetworkStream, + SslStream sslStream = new(client.SafeNetworkStream, false, this.RemoteCertificateValidationCallback, this.LocalCertificateSelectionCallback, @@ -917,16 +857,20 @@ namespace MTSC.ServerSide } catch (Exception e) { - foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) - { - if (exceptionHandler.HandleException(e)) - { - break; - } - } + this.HandleException(e); client.Dispose(); } } + private void HandleException(Exception exception) + { + foreach (IExceptionHandler exceptionHandler in this.exceptionHandlers) + { + if (exceptionHandler.HandleException(exception)) + { + break; + } + } + } #endregion } }