diff --git a/MTSC.UnitTests/E2ETests.cs b/MTSC.UnitTests/E2ETests.cs index 776f761..4773e43 100644 --- a/MTSC.UnitTests/E2ETests.cs +++ b/MTSC.UnitTests/E2ETests.cs @@ -59,7 +59,7 @@ namespace MTSC.UnitTests .AddLogger(new ConsoleLogger()) .AddLogger(new DebugConsoleLogger()) .AddExceptionHandler(new ExceptionConsoleLogger()) - .SetScheduler(new TaskAwaiterScheduler()) + .SetScheduler(new ParallelScheduler()) .WithSslAuthenticationTimeout(TimeSpan.FromMilliseconds(100)); Server.RunAsync(); } @@ -204,7 +204,7 @@ namespace MTSC.UnitTests } } var response = HttpResponse.FromBytes(receivedMessage); - Assert.AreEqual(response.StatusCode, HttpMessage.StatusCodes.BadRequest); + Assert.AreNotEqual(response.StatusCode, HttpMessage.StatusCodes.OK); } [TestMethod] @@ -212,6 +212,7 @@ namespace MTSC.UnitTests { HttpClient httpClient = new HttpClient(); httpClient.BaseAddress = new Uri("https://localhost:800"); + httpClient.DefaultRequestHeaders.ExpectContinue = true; string s = string.Empty; for(int i = 0; i < 200000; i++) { diff --git a/MTSC/Common/Http/PartialHttpRequest.cs b/MTSC/Common/Http/PartialHttpRequest.cs index 0441af3..e17d505 100644 --- a/MTSC/Common/Http/PartialHttpRequest.cs +++ b/MTSC/Common/Http/PartialHttpRequest.cs @@ -17,6 +17,7 @@ namespace MTSC.Common.Http /// public List Cookies { get; } = new List(); public bool Complete { get; private set; } = false; + public int HeaderByteCount { get; private set; } = 0; public Form Form { get; } = new Form(); public HttpMethods Method { get; set; } public string RequestURI { get; set; } @@ -70,6 +71,10 @@ namespace MTSC.Common.Http Array.Copy(bytesToBeAdded, 0, newBody, Body.Length, bytesToBeAdded.Length); } Body = newBody; + if (this.Headers.ContainsHeader(EntityHeaders.ContentLength) && int.Parse(Headers[EntityHeaders.ContentLength]) == Body.Length) + { + this.Complete = true; + } } private HttpMethods GetMethod(string methodString) @@ -384,6 +389,7 @@ namespace MTSC.Common.Http throw new IncompleteRequestException($"Incomplete request.", new HttpRequestParsingException("Exception during parsing of http request. Buffer: " + UTF8Encoding.UTF8.GetString(ms.ToArray()))); } + this.HeaderByteCount = (int)ms.Position; if (Headers.ContainsHeader(EntityHeaders.ContentLength)) { int remainingBytes = int.Parse(Headers[EntityHeaders.ContentLength]); diff --git a/MTSC/MTSC.csproj b/MTSC/MTSC.csproj index 359828d..123c936 100644 --- a/MTSC/MTSC.csproj +++ b/MTSC/MTSC.csproj @@ -5,12 +5,12 @@ netcoreapp2.1;net48;netstandard2.0;netcoreapp3.0;netcoreapp3.1 - 2.5.1 + 2.6 Alexandru-Victor Macocian MTSC Modular TCP Server and Client - 0.2.5.1 - 0.2.5.1 + 0.2.6.0 + 0.2.6.0 true AnyCPU;x64 https://github.com/AlexMacocian/MTSC diff --git a/MTSC/ServerSide/ClientData.cs b/MTSC/ServerSide/ClientData.cs index a57060b..93859fb 100644 --- a/MTSC/ServerSide/ClientData.cs +++ b/MTSC/ServerSide/ClientData.cs @@ -1,4 +1,5 @@ using MTSC.Common; +using MTSC.ServerSide.Handlers; using System; using System.Net.Security; using System.Net.Sockets; @@ -21,6 +22,10 @@ namespace MTSC.ServerSide /// Latest datetime when a message has been received or sent to the client /// public DateTime LastActivityTime { get; private set; } = DateTime.Now; + /// + /// Sets the affinity of the client to a specific handler, ignoring all other handlers + /// + public IHandler Affinity { get; private set; } IConsumerQueue IQueueHolder.ConsumerQueue { get => messageQueue; } @@ -34,6 +39,33 @@ namespace MTSC.ServerSide this.TcpClient = client; this.SafeNetworkStream = new SafeNetworkStream(this.TcpClient); } + /// + /// Sets the affinity of the client. + /// + /// Handler to bind to. + public void SetAffinity(IHandler handler) + { + this.Affinity = handler; + } + /// + /// Resets the affinity of the client. + /// + public void ResetAffinity() + { + this.Affinity = null; + } + /// + /// Resets the affinity if the handler is the one binded. + /// + /// The handler requesting reset. + public void ResetAffinityIfMe(IHandler handler) + { + if(this.Affinity == handler) + { + this.Affinity = null; + } + } + #region IActiveClient Implementation void IActiveClient.UpdateLastReceivedMessage() { diff --git a/MTSC/ServerSide/Handlers/HttpHandler.cs b/MTSC/ServerSide/Handlers/HttpHandler.cs index 512fc87..d1b411d 100644 --- a/MTSC/ServerSide/Handlers/HttpHandler.cs +++ b/MTSC/ServerSide/Handlers/HttpHandler.cs @@ -126,9 +126,13 @@ namespace MTSC.ServerSide.Handlers messageBytes = messageBytes.TrimTrailingNullBytes(); var partialRequest = PartialHttpRequest.FromBytes(messageBytes); if (partialRequest.Complete) + { request = partialRequest.ToRequest(); + client.ResetAffinityIfMe(this); + } else { + client.SetAffinity(this); HandleIncompleteRequest(client, server, messageBytes, partialRequest); return true; } @@ -146,10 +150,12 @@ namespace MTSC.ServerSide.Handlers { server.LogDebug("Malformed request, not saving!"); server.LogDebug(ex.Message + "\n" + ex.StackTrace); + client.ResetAffinityIfMe(this); return false; } catch (Exception e) { + client.ResetAffinityIfMe(this); throw e; } diff --git a/MTSC/ServerSide/Handlers/HttpRoutingHandler.cs b/MTSC/ServerSide/Handlers/HttpRoutingHandler.cs index 54a0929..949b505 100644 --- a/MTSC/ServerSide/Handlers/HttpRoutingHandler.cs +++ b/MTSC/ServerSide/Handlers/HttpRoutingHandler.cs @@ -87,107 +87,90 @@ namespace MTSC.ServerSide.Handlers bool IHandler.HandleReceivedMessage(Server server, ClientData client, Message message) { - // Parse the request. If the message is incomplete, return 100 and queue the message to be parsed later. - HttpRequest request = null; - byte[] messageBytes = null; - try + /* + * If a fragmented request exists, add the new messages to the body. + * 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; + if (client.Resources.TryGetResource(out var fragmentedMessage)) { - var trimmedMessageBytes = message.MessageBytes.TrimTrailingNullBytes(); - if (client.Resources.TryGetResource(out var fragmentedMessage)) + var bytesToBeAdded = message.MessageBytes.TrimTrailingNullBytes(); + if(fragmentedMessage.PartialRequest.HeaderByteCount + + fragmentedMessage.PartialRequest.Body.Length + + message.MessageLength > this.MaximumRequestSize) { - byte[] previousBytes = fragmentedMessage.Message; - if (previousBytes.Length + trimmedMessageBytes.Length > MaximumRequestSize) - { - // Discard the message if it is too big - server.LogDebug($"Discarded message. Message size [{previousBytes.Length + trimmedMessageBytes.Length}] > [{MaximumRequestSize}]"); - client.Resources.RemoveResource(); - QueueResponse(client, new HttpResponse { StatusCode = StatusCodes.BadRequest, BodyString = $"Request disallowed because it exceeds [{MaximumRequestSize}] bytes!" }); - return true; - } - byte[] repackagingBuffer = new byte[previousBytes.Length + trimmedMessageBytes.Length]; - Array.Copy(previousBytes, 0, repackagingBuffer, 0, previousBytes.Length); - Array.Copy(trimmedMessageBytes, 0, repackagingBuffer, previousBytes.Length, trimmedMessageBytes.Length); - messageBytes = repackagingBuffer; - } - else - { - if (trimmedMessageBytes.Length > MaximumRequestSize) - { - // Discard the message if it is too big - server.LogDebug($"Discarded message. Message size [{trimmedMessageBytes.Length}] > [{MaximumRequestSize}]"); - QueueResponse(client, new HttpResponse { StatusCode = StatusCodes.BadRequest, BodyString = $"Request disallowed because it exceeds [{MaximumRequestSize}] bytes!" }); - return true; - } - messageBytes = trimmedMessageBytes; - } - var partialRequest = PartialHttpRequest.FromBytes(messageBytes); - if (partialRequest.Complete) - request = partialRequest.ToRequest(); - else - { - HandleIncompleteRequest(client, server, messageBytes, partialRequest); - return true; + QueueResponse(client, new HttpResponse { StatusCode = StatusCodes.BadRequest, BodyString = $"Request exceeded [{MaximumRequestSize}] bytes!" }); + client.ResetAffinityIfMe(this); + client.Resources.RemoveResource(); + client.Resources.RemoveResourceIfExists(); } + fragmentedMessage.AddToMessage(bytesToBeAdded); + request = fragmentedMessage.PartialRequest; } - catch (Exception ex) when ( - ex is IncompleteHeaderKeyException || - ex is IncompleteHeaderValueException || - ex is IncompleteHttpVersionException || - ex is IncompleteMethodException || - ex is IncompleteRequestBodyException || - ex is IncompleteRequestQueryException || - ex is IncompleteRequestURIException || - ex is IncompleteRequestException || - ex is InvalidPostFormException) + else { - server.LogDebug("Malformed request, not saving!"); - server.LogDebug(ex.Message + "\n" + ex.StackTrace); - return false; - } - catch (Exception e) - { - throw e; - } + if (message.MessageLength > this.MaximumRequestSize) + { + QueueResponse(client, new HttpResponse { StatusCode = StatusCodes.BadRequest, BodyString = $"Request exceeded [{MaximumRequestSize}] bytes!" }); + return false; + } - // The message has been parsed. If there was a cache for the current message, remove it. - if (client.Resources.Contains()) - { - client.Resources.RemoveResource(); + request = PartialHttpRequest.FromBytes(message.MessageBytes.TrimTrailingNullBytes()); + if (!request.Complete) + { + client.Resources.SetResource(new FragmentedMessage { LastReceived = DateTime.Now, PartialRequest = request }); + client.SetAffinity(this); + } } /* - * Now find if a routing module exists. If not let other handlers try and handle the message. + * Once headers are loaded, check if mapping exists. If mapping exists between request and module, + * verify that the request is complete and send it to module. + * Otherwise, if the request is complete, send it to the mapped module. If the request is not complete, + * handle it and return. */ - if (moduleDictionary[request.Method].ContainsKey(request.RequestURI)) + + if (client.Resources.TryGetResource(out var mapping)) { - (var module, var routeEnabler) = moduleDictionary[request.Method][request.RequestURI]; - var routeEnablerResponse = routeEnabler.Invoke(server, request, client); - if (routeEnablerResponse is RouteEnablerResponse.RouteEnablerResponseAccept) + if (request.Complete) { - try - { - module.CallHandleRequest(request, client, server).ContinueWith((task) => { QueueResponse(client, task.Result); }); - //response = module.HandleRequest(requestTemplate.Invoke(request), client, server); - } - catch (Exception e) - { - server.LogDebug("Exception: " + e.Message); - server.LogDebug("Stacktrace: " + e.StackTrace); - QueueResponse(client, new HttpResponse() { StatusCode = StatusCodes.InternalServerError }); - } + client.Resources.RemoveResource(); + client.Resources.RemoveResource(); + client.ResetAffinityIfMe(this); + HandleCompleteRequest(client, server, request.ToRequest(), mapping.MappedModule, mapping.RouteEnabler); return true; } - else if (routeEnablerResponse is RouteEnablerResponse.RouteEnablerResponseIgnore) + else { - return false; - } - else if (routeEnablerResponse is RouteEnablerResponse.RouteEnablerResponseError) - { - QueueResponse(client, (routeEnablerResponse as RouteEnablerResponse.RouteEnablerResponseError).Response); + HandleIncompleteRequest(client, server, client.Resources.GetResource()); return true; } } - return false; + else + { + if (moduleDictionary[request.Method].ContainsKey(request.RequestURI)) + { + (var module, var routeEnabler) = moduleDictionary[request.Method][request.RequestURI]; + if (request.Complete) + { + var httpRequest = request.ToRequest(); + client.ResetAffinityIfMe(this); + return HandleCompleteRequest(client, server, httpRequest, module, routeEnabler); + } + else + { + client.Resources.SetResource(new RequestMapping { MappedModule = module, RouteEnabler = routeEnabler }); + return true; + } + } + else + { + client.ResetAffinityIfMe(this); + client.Resources.RemoveResource(); + return false; + } + } } bool IHandler.HandleSendMessage(Server server, ClientData client, ref Message message) => false; @@ -216,12 +199,12 @@ namespace MTSC.ServerSide.Handlers } } - private void HandleIncompleteRequest(ClientData client, Server server, byte[] messageBytes, PartialHttpRequest partialRequest = null) + private void HandleIncompleteRequest(ClientData client, Server server, FragmentedMessage fragmentedMessage) { - client.Resources.SetResource(new FragmentedMessage() { Message = messageBytes, LastReceived = DateTime.Now }); + fragmentedMessage.LastReceived = DateTime.Now; server.LogDebug("Incomplete request received!"); - if (partialRequest != null && partialRequest.Headers.ContainsHeader(HttpMessage.RequestHeaders.Expect) && - partialRequest.Headers[HttpMessage.RequestHeaders.Expect].Equals("100-continue", StringComparison.OrdinalIgnoreCase)) + if (fragmentedMessage.PartialRequest.Headers.ContainsHeader(HttpMessage.RequestHeaders.Expect) && + fragmentedMessage.PartialRequest.Headers[HttpMessage.RequestHeaders.Expect].Equals("100-continue", StringComparison.OrdinalIgnoreCase)) { server.LogDebug("Returning 100-Continue"); var contResponse = new HttpResponse { StatusCode = HttpMessage.StatusCodes.Continue }; @@ -230,18 +213,59 @@ namespace MTSC.ServerSide.Handlers } } + private bool HandleCompleteRequest( + ClientData client, + Server server, + HttpRequest request, + HttpRouteBase module, + Func routeEnabler) + { + var routeEnablerResponse = routeEnabler.Invoke(server, request, client); + if (routeEnablerResponse is RouteEnablerResponse.RouteEnablerResponseAccept) + { + try + { + module.CallHandleRequest(request, client, server).ContinueWith((task) => { QueueResponse(client, task.Result); }); + } + catch (Exception e) + { + server.LogDebug("Exception: " + e.Message); + server.LogDebug("Stacktrace: " + e.StackTrace); + QueueResponse(client, new HttpResponse() { StatusCode = StatusCodes.InternalServerError }); + } + return true; + } + else if (routeEnablerResponse is RouteEnablerResponse.RouteEnablerResponseIgnore) + { + return false; + } + else if (routeEnablerResponse is RouteEnablerResponse.RouteEnablerResponseError) + { + QueueResponse(client, (routeEnablerResponse as RouteEnablerResponse.RouteEnablerResponseError).Response); + return true; + } + else + { + throw new InvalidOperationException($"RouteEnablerResponse should be one of the types {typeof(RouteEnablerResponse.RouteEnablerResponseAccept)}, {typeof(RouteEnablerResponse.RouteEnablerResponseError)} or {typeof(RouteEnablerResponse.RouteEnablerResponseIgnore)}!"); + } + } + private class FragmentedMessage { - public byte[] Message { get; set; } + public PartialHttpRequest PartialRequest { get; set; } public DateTime LastReceived { get; set; } = DateTime.Now; public void AddToMessage(byte[] bytes) { - byte[] newMessage = new byte[Message.Length + bytes.Length]; - Array.Copy(Message, newMessage, Message.Length); - Array.Copy(bytes, 0, newMessage, Message.Length, bytes.Length); + this.PartialRequest.AddToBody(bytes); } } + + private class RequestMapping + { + public HttpRouteBase MappedModule { get; set; } + public Func RouteEnabler { get; set; } + } } } diff --git a/MTSC/ServerSide/Handlers/WebsocketHandler.cs b/MTSC/ServerSide/Handlers/WebsocketHandler.cs index 9d66c19..86f0b48 100644 --- a/MTSC/ServerSide/Handlers/WebsocketHandler.cs +++ b/MTSC/ServerSide/Handlers/WebsocketHandler.cs @@ -98,6 +98,7 @@ namespace MTSC.ServerSide.Handlers request.Headers[HttpMessage.GeneralHeaders.Connection].ToLower() == "upgrade" && request.Headers.ContainsHeader(WebsocketProtocolVersionKey) && request.Headers[WebsocketProtocolVersionKey] == "13") { + client.SetAffinity(this); /* * Prepare the handshake string. */ @@ -169,6 +170,7 @@ namespace MTSC.ServerSide.Handlers server.QueueMessage(tuple.Item1, tuple.Item2.GetMessageBytes()); if (tuple.Item2.Opcode == WebsocketMessage.Opcodes.Close) { + tuple.Item1.ResetAffinityIfMe(this); foreach (IWebsocketModule websocketModule in websocketModules) { websocketModule.ConnectionClosed(server, this, tuple.Item1); diff --git a/MTSC/ServerSide/ResourceDictionary.cs b/MTSC/ServerSide/ResourceDictionary.cs index 6393b9a..e800736 100644 --- a/MTSC/ServerSide/ResourceDictionary.cs +++ b/MTSC/ServerSide/ResourceDictionary.cs @@ -7,6 +7,14 @@ namespace MTSC.ServerSide { private Dictionary Resources = new Dictionary(); + public void RemoveResourceIfExists() + { + if (Resources.ContainsKey(typeof(TValue))) + { + Resources.Remove((typeof(TValue))); + } + } + public void RemoveResource() { Resources.Remove(typeof(TValue)); diff --git a/MTSC/ServerSide/Schedulers/ParallelScheduler.cs b/MTSC/ServerSide/Schedulers/ParallelScheduler.cs index c627b46..53aa4a8 100644 --- a/MTSC/ServerSide/Schedulers/ParallelScheduler.cs +++ b/MTSC/ServerSide/Schedulers/ParallelScheduler.cs @@ -16,6 +16,8 @@ namespace MTSC.ServerSide.Schedulers (var client, var messageQueue) = tuple; actionList.Add(new Action(() => messageHandlingProcedure.Invoke(client, messageQueue))); } + + Parallel.Invoke(actionList.ToArray()); } } } diff --git a/MTSC/ServerSide/Server.cs b/MTSC/ServerSide/Server.cs index 24d4aaf..a26c992 100644 --- a/MTSC/ServerSide/Server.cs +++ b/MTSC/ServerSide/Server.cs @@ -524,6 +524,13 @@ namespace MTSC.ServerSide while (this._ConsumerMessageOutQueue.TryDequeue(out var tuple)) { (var client, var bytes) = tuple; + if (client.TcpClient.Available > 0) + { + /* + * Don't send while client is still sending + */ + continue; + } try { Message sendMessage = CommunicationPrimitives.BuildMessage(bytes); @@ -639,7 +646,47 @@ namespace MTSC.ServerSide { while(messages.TryDequeue(out var message)) { - HandleClientMessage(client, message); + if (client.Affinity is null) + { + HandleClientMessage(client, message); + } + else + { + AffinityHandleClientMessage(client, message); + } + } + } + + private void AffinityHandleClientMessage(ClientData client, Message message) + { + var handler = client.Affinity; + try + { + handler.PreHandleReceivedMessage(this, client, ref message); + } + catch (Exception e) + { + foreach (IExceptionHandler exceptionHandler in exceptionHandlers) + { + if (exceptionHandler.HandleException(e)) + { + break; + } + } + } + try + { + handler.HandleReceivedMessage(this, client, message); + } + catch (Exception e) + { + foreach (IExceptionHandler exceptionHandler in exceptionHandlers) + { + if (exceptionHandler.HandleException(e)) + { + break; + } + } } } @@ -676,8 +723,6 @@ namespace MTSC.ServerSide } catch (Exception e) { - LogDebug("Exception: " + e.Message); - LogDebug("Stacktrace: " + e.StackTrace); foreach (IExceptionHandler exceptionHandler in exceptionHandlers) { if (exceptionHandler.HandleException(e))