From fc36acc494bbfdece3fe27c4317c51f762bb2d83 Mon Sep 17 00:00:00 2001 From: Alexandru Macocian Date: Wed, 11 Mar 2020 22:36:13 +0100 Subject: [PATCH] Added SendingTemplate to HttpResponse Implemented WebsocketRouteBase --- MTSC.UnitTests/E2ETests.cs | 15 +- MTSC.UnitTests/EchoWebsocketModule.cs | 24 ++ .../Http/RoutingModules/HttpRouteBase.cs | 37 ++- .../RoutingModules/WebsocketRouteBase.cs | 98 +++++++ MTSC/MTSC.csproj | 6 +- MTSC/ServerSide/Handlers/WebsocketHandler.cs | 2 +- .../Handlers/WebsocketRoutingHandler.cs | 241 ++++++++++++++++++ 7 files changed, 415 insertions(+), 8 deletions(-) create mode 100644 MTSC.UnitTests/EchoWebsocketModule.cs create mode 100644 MTSC/Common/WebSockets/RoutingModules/WebsocketRouteBase.cs create mode 100644 MTSC/ServerSide/Handlers/WebsocketRoutingHandler.cs diff --git a/MTSC.UnitTests/E2ETests.cs b/MTSC.UnitTests/E2ETests.cs index 931e501..7475fc9 100644 --- a/MTSC.UnitTests/E2ETests.cs +++ b/MTSC.UnitTests/E2ETests.cs @@ -2,6 +2,7 @@ using Microsoft.VisualStudio.TestTools.UnitTesting; using MTSC.Common.Http; using MTSC.Common.Http.RoutingModules; using MTSC.Common.Http.ServerModules; +using MTSC.Common.WebSockets; using MTSC.Exceptions; using MTSC.Logging; using MTSC.ServerSide.Handlers; @@ -29,8 +30,16 @@ namespace MTSC.UnitTests public static void InitializeServer(TestContext testContext) { Server = new ServerSide.Server(800) - .AddHandler(new WebsocketHandler() - .AddWebsocketHandler(new Common.WebSockets.ServerModules.EchoModule())) + .AddHandler(new WebsocketRoutingHandler() + .AddRoute("echo", new EchoWebsocketModule() + .WithReceiveTemplateProvider((message) => UTF8Encoding.UTF8.GetString(message.Data)) + .WithSendTemplateProvider((s) => + { + WebsocketMessage websocketMessage = new WebsocketMessage(); + websocketMessage.Data = UTF8Encoding.UTF8.GetBytes(s); + websocketMessage.Opcode = WebsocketMessage.Opcodes.Text; + return websocketMessage; + }))) .AddHandler(new HttpHandler() .AddHttpModule(new HttpRoutingModule() .AddRoute(HttpMessage.HttpMethods.Get, "", new Http200Module()) @@ -184,7 +193,7 @@ namespace MTSC.UnitTests { byte[] bytes = new byte[100]; ClientWebSocket client = new ClientWebSocket(); - client.ConnectAsync(new Uri("ws://localhost:800"), CancellationToken.None).Wait(); + client.ConnectAsync(new Uri("ws://localhost:800/echo"), CancellationToken.None).Wait(); client.SendAsync(ASCIIEncoding.ASCII.GetBytes("Hello world!"), WebSocketMessageType.Text, true, CancellationToken.None).Wait(); client.ReceiveAsync(bytes, CancellationToken.None).Wait(); var resultString = ASCIIEncoding.ASCII.GetString(bytes, 0, 12); diff --git a/MTSC.UnitTests/EchoWebsocketModule.cs b/MTSC.UnitTests/EchoWebsocketModule.cs new file mode 100644 index 0000000..4b3e412 --- /dev/null +++ b/MTSC.UnitTests/EchoWebsocketModule.cs @@ -0,0 +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 ConnectionInitialized(Server server, WebsocketRoutingHandler handler, ClientData client) + { + + } + + public override void HandleReceivedMessage(Server server, WebsocketRoutingHandler handler, ClientData client, string message) + { + this.SendMessage(message, client, handler); + } + } +} diff --git a/MTSC/Common/Http/RoutingModules/HttpRouteBase.cs b/MTSC/Common/Http/RoutingModules/HttpRouteBase.cs index 0ed1f5d..7de0941 100644 --- a/MTSC/Common/Http/RoutingModules/HttpRouteBase.cs +++ b/MTSC/Common/Http/RoutingModules/HttpRouteBase.cs @@ -26,7 +26,7 @@ namespace MTSC.Common.Http.RoutingModules } - public HttpRouteBase WithTemplateProvider(Func templateProvider) + public HttpRouteBase WithTemplateProvider(Func templateProvider) { this.template = templateProvider; return this; @@ -39,4 +39,39 @@ namespace MTSC.Common.Http.RoutingModules public abstract HttpResponse HandleRequest(T request, ClientData client, ServerSide.Server server); } + public abstract class HttpRouteBase : HttpRouteBase + { + private Func receiveTemplate; + private Func sendTemplate; + + public HttpRouteBase(Func receiveTemplate, Func sendTemplate) + { + this.receiveTemplate = receiveTemplate; + this.sendTemplate = sendTemplate; + } + + public HttpRouteBase() + { + + } + + public HttpRouteBase WithReceiveTemplateProvider(Func templateProvider) + { + this.receiveTemplate = templateProvider; + return this; + } + + public HttpRouteBase WithSendTemplateProvider(Func templateProvider) + { + this.sendTemplate = templateProvider; + return this; + } + + public override HttpResponse HandleRequest(HttpRequest request, ClientData client, ServerSide.Server server) + { + return sendTemplate.Invoke(HandleRequest(receiveTemplate.Invoke(request), client, server)); + } + + public abstract TSend HandleRequest(TReceive request, ClientData client, ServerSide.Server server); + } } diff --git a/MTSC/Common/WebSockets/RoutingModules/WebsocketRouteBase.cs b/MTSC/Common/WebSockets/RoutingModules/WebsocketRouteBase.cs new file mode 100644 index 0000000..f5890cc --- /dev/null +++ b/MTSC/Common/WebSockets/RoutingModules/WebsocketRouteBase.cs @@ -0,0 +1,98 @@ +using MTSC.ServerSide; +using MTSC.ServerSide.Handlers; +using System; + +namespace MTSC.Common.WebSockets.RoutingModules +{ + public abstract class WebsocketRouteBase + { + public void CallConnectionInitialized(Server server, WebsocketRoutingHandler handler, ClientData client) + { + ConnectionInitialized(server, handler, client); + } + public void CallHandleReceivedMessage(Server server, WebsocketRoutingHandler handler, ClientData client, WebsocketMessage receivedMessage) + { + HandleReceivedMessage(server, handler, client, receivedMessage); + } + public void CallConnectionClosed(Server server, WebsocketRoutingHandler handler, ClientData client) + { + ConnectionClosed(server, handler, client); + } + + public void SendMessage(WebsocketMessage message, ClientData client, WebsocketRoutingHandler handler) + { + handler.QueueMessage(client, 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 class WebsocketRouteBase : WebsocketRouteBase + { + private Func receiveTemplate; + + public WebsocketRouteBase(Func receiveTemplate) + { + this.receiveTemplate = receiveTemplate; + } + + public WebsocketRouteBase() + { + + } + + public WebsocketRouteBase WithReceiveTemplateProvider(Func templateProvider) + { + this.receiveTemplate = templateProvider; + return this; + } + + public override void HandleReceivedMessage(Server server, WebsocketRoutingHandler handler, ClientData client, WebsocketMessage receivedMessage) + { + HandleReceivedMessage(server, handler, client, receiveTemplate.Invoke(receivedMessage)); + } + + public abstract void HandleReceivedMessage(Server server, WebsocketRoutingHandler handler, ClientData client, TReceive message); + } + public abstract class WebsocketRouteBase : WebsocketRouteBase + { + private Func receiveTemplate; + private Func sendTemplate; + + public WebsocketRouteBase(Func receiveTemplate, Func sendTemplate) + { + this.receiveTemplate = receiveTemplate; + this.sendTemplate = sendTemplate; + } + + public WebsocketRouteBase() + { + + } + + public WebsocketRouteBase WithReceiveTemplateProvider(Func templateProvider) + { + this.receiveTemplate = templateProvider; + return this; + } + + public WebsocketRouteBase WithSendTemplateProvider(Func templateProvider) + { + this.sendTemplate = templateProvider; + return this; + } + + public void SendMessage(TSend message, ClientData client, WebsocketRoutingHandler handler) + { + base.SendMessage(sendTemplate.Invoke(message), client, handler); + } + + public override void HandleReceivedMessage(Server server, WebsocketRoutingHandler handler, ClientData client, WebsocketMessage receivedMessage) + { + HandleReceivedMessage(server, handler, client, receiveTemplate.Invoke(receivedMessage)); + } + + public abstract void HandleReceivedMessage(Server server, WebsocketRoutingHandler handler, ClientData client, TReceive message); + } +} diff --git a/MTSC/MTSC.csproj b/MTSC/MTSC.csproj index 436d135..5ae7ee0 100644 --- a/MTSC/MTSC.csproj +++ b/MTSC/MTSC.csproj @@ -5,12 +5,12 @@ netcoreapp2.1;net48;netstandard2.0;netcoreapp3.0 - 1.9.2 + 2.0.0 Alexandru-Victor Macocian MTSC Modular TCP Server and Client - 0.1.9.2 - 0.1.9.2 + 0.2.0.0 + 0.2.0.0 true AnyCPU;x64 https://github.com/AlexMacocian/MTSC diff --git a/MTSC/ServerSide/Handlers/WebsocketHandler.cs b/MTSC/ServerSide/Handlers/WebsocketHandler.cs index 6d52a8c..53b8af1 100644 --- a/MTSC/ServerSide/Handlers/WebsocketHandler.cs +++ b/MTSC/ServerSide/Handlers/WebsocketHandler.cs @@ -97,7 +97,7 @@ namespace MTSC.ServerSide.Handlers { if (webSockets[client] == SocketState.Initial) { - HttpRequest request = new HttpRequest(message.MessageBytes); + PartialHttpRequest request = new PartialHttpRequest(message.MessageBytes); if(request.Method == HttpMessage.HttpMethods.Get && request.Headers.ContainsHeader(HttpMessage.GeneralHeaders.Connection) && request.Headers[HttpMessage.GeneralHeaders.Connection].ToLower() == "upgrade" && request.Headers.ContainsHeader(WebsocketProtocolVersionKey) && request.Headers[WebsocketProtocolVersionKey] == "13") diff --git a/MTSC/ServerSide/Handlers/WebsocketRoutingHandler.cs b/MTSC/ServerSide/Handlers/WebsocketRoutingHandler.cs new file mode 100644 index 0000000..64b905d --- /dev/null +++ b/MTSC/ServerSide/Handlers/WebsocketRoutingHandler.cs @@ -0,0 +1,241 @@ +using MTSC.Common.Http; +using MTSC.Common.WebSockets; +using MTSC.Common.WebSockets.RoutingModules; +using System; +using System.Collections.Concurrent; +using System.Collections.Generic; +using System.Security.Cryptography; +using System.Text; + +namespace MTSC.ServerSide.Handlers +{ + public class WebsocketRoutingHandler : IHandler + { + 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(); + public enum SocketState + { + Initial, + Handshaking, + Established, + Closed + } + #region Fields + public ConcurrentDictionary webSockets = new ConcurrentDictionary(); + Dictionary)> moduleDictionary = + new Dictionary)>(); + ConcurrentDictionary routingTable = new ConcurrentDictionary(); + ConcurrentQueue> messageQueue = new ConcurrentQueue>(); + #endregion + #region Public Methods + public WebsocketRoutingHandler AddRoute(string uri, WebsocketRouteBase module) + { + this.moduleDictionary.Add(uri, (module, alwaysEnabled)); + return this; + } + + public WebsocketRoutingHandler AddRoute( + string uri, + WebsocketRouteBase module, + Func routeEnabler) + { + this.moduleDictionary.Add(uri, (module, routeEnabler)); + return this; + } + + public WebsocketRoutingHandler RemoveRoute(string uri) + { + 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; + rng.GetBytes(sendMessage.Mask); + sendMessage.Opcode = opcode; + messageQueue.Enqueue(new Tuple(client, sendMessage)); + } + + public void QueueMessage(ClientData client, WebsocketMessage message) + { + 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); + } + #endregion + #region Handler Implementation + void IHandler.ClientRemoved(Server server, ClientData client) + { + SocketState state = SocketState.Initial; + while (webSockets.ContainsKey(client)) + { + webSockets.TryRemove(client, out state); + } + routingTable[client].CallConnectionClosed(server, this, client); + while (routingTable.ContainsKey(client)) + { + routingTable.TryRemove(client, out _); + } + } + + bool IHandler.HandleClient(Server server, ClientData client) + { + webSockets[client] = SocketState.Initial; + return false; + } + + bool IHandler.HandleReceivedMessage(Server server, ClientData client, Message message) + { + if (webSockets[client] == SocketState.Initial) + { + PartialHttpRequest request; + try + { + request = PartialHttpRequest.FromBytes(message.MessageBytes); + } + catch(Exception e) + { + server.LogDebug(e.Message + "\n" + e.StackTrace); + return false; + } + if (request.Method == HttpMessage.HttpMethods.Get && request.Headers.ContainsHeader(HttpMessage.GeneralHeaders.Connection) && + request.Headers[HttpMessage.GeneralHeaders.Connection].ToLower() == "upgrade" && request.Headers.ContainsHeader(WebsocketProtocolVersionKey) && + request.Headers[WebsocketProtocolVersionKey] == "13") + { + if (!moduleDictionary.ContainsKey(request.RequestURI)) + { + QueueMessage(client, new HttpResponse { StatusCode = HttpMessage.StatusCodes.NotFound, BodyString = "URI not found" }.GetPackedResponse(true)); + } + (var module, var routeEnabler) = moduleDictionary[request.RequestURI]; + var routeEnablerResponse = routeEnabler.Invoke(server, request.ToRequest(), client); + if(routeEnablerResponse is RouteEnablerResponse.RouteEnablerResponseIgnore) + { + return false; + } + else if(routeEnablerResponse is RouteEnablerResponse.RouteEnablerResponseError) + { + 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]; + base64Key = base64Key.Trim(); + string handshakeKey = base64Key + GlobalUniqueIdentifier; + string returnBase64Key = Convert.ToBase64String(sha1Provider.ComputeHash(Encoding.UTF8.GetBytes(handshakeKey))); + + /* + * Prepare the response. + */ + HttpResponse response = new HttpResponse(); + response.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(true)); + webSockets[client] = SocketState.Established; + server.LogDebug("Websocket initialized " + client.TcpClient.Client.RemoteEndPoint.ToString()); + routingTable[client] = module; + module.CallConnectionInitialized(server, this, client); + return true; + } + else + { + return false; + } + } + else if (webSockets[client] == SocketState.Established) + { + WebsocketMessage receivedMessage = null; + try + { + receivedMessage = new WebsocketMessage(message.MessageBytes); + } + catch(Exception e) + { + server.LogDebug(e.Message + "\n" + e.StackTrace); + return false; + } + + if (receivedMessage.Opcode == WebsocketMessage.Opcodes.Close) + { + client.ToBeRemoved = true; + while (webSockets.ContainsKey(client)) + { + webSockets.TryRemove(client, out SocketState _); + } + WebsocketMessage closeFrame = new WebsocketMessage(); + closeFrame.Opcode = WebsocketMessage.Opcodes.Close; + QueueMessage(client, closeFrame); + return true; + } + else + { + try + { + routingTable[client].CallHandleReceivedMessage(server, this, client, receivedMessage); + return true; + } + catch(Exception e) + { + server.LogDebug(e.Message + "\n" + e.StackTrace); + return false; + } + } + } + else + { + return false; + } + } + + bool IHandler.HandleSendMessage(Server server, ClientData client, ref Message message) + { + return false; + } + + bool IHandler.PreHandleReceivedMessage(Server server, ClientData client, ref Message message) + { + return false; + } + + void IHandler.Tick(Server server) + { + while (messageQueue.Count > 0) + { + if (messageQueue.TryDequeue(out Tuple tuple)) + { + server.QueueMessage(tuple.Item1, tuple.Item2.GetMessageBytes()); + if (tuple.Item2.Opcode == WebsocketMessage.Opcodes.Close) + { + if (routingTable.ContainsKey(tuple.Item1)) + { + routingTable[tuple.Item1].CallConnectionClosed(server, this, tuple.Item1); + } + tuple.Item1.ToBeRemoved = true; + } + } + } + } + #endregion + } +}