From 533d8f8fc70b0edb2f471adb73ec2e04fe46dbf1 Mon Sep 17 00:00:00 2001 From: Alexandru Macocian Date: Tue, 3 Mar 2020 19:41:40 +0100 Subject: [PATCH] Introduced automatic 100-continue response in case user provided correct headers. --- MTSC/Common/Http/PartialHttpRequest.cs | 563 +++++++++++++++++++++++++ MTSC/MTSC.csproj | 6 +- MTSC/Server/Handlers/HttpHandler.cs | 44 +- 3 files changed, 589 insertions(+), 24 deletions(-) create mode 100644 MTSC/Common/Http/PartialHttpRequest.cs diff --git a/MTSC/Common/Http/PartialHttpRequest.cs b/MTSC/Common/Http/PartialHttpRequest.cs new file mode 100644 index 0000000..b4504f6 --- /dev/null +++ b/MTSC/Common/Http/PartialHttpRequest.cs @@ -0,0 +1,563 @@ +using MTSC.Exceptions; +using System; +using System.Collections.Generic; +using System.IO; +using System.Text; +using static MTSC.Common.Http.HttpMessage; + +namespace MTSC.Common.Http +{ + class PartialHttpRequest + { + public HttpRequestHeaderDictionary Headers { get; } = new HttpRequestHeaderDictionary(); + + /// + /// List of cookies. + /// + public List Cookies { get; } = new List(); + public bool Complete { get; private set; } = false; + public Dictionary Form { get; } = new Dictionary(); + public HttpMethods Method { get; set; } + public string RequestURI { get; set; } + public string RequestQuery { get; set; } + public byte[] Body { get; set; } = new byte[0]; + public string BodyString { get => ASCIIEncoding.ASCII.GetString(Body).Trim('\0'); set => Body = ASCIIEncoding.ASCII.GetBytes(value); } + + public PartialHttpRequest() + { + + } + + public PartialHttpRequest(byte[] requestBytes) + { + ParseRequest(requestBytes); + if (this.Method == HttpMethods.Post) + { + Form = GetPostForm(); + } + } + + public static PartialHttpRequest FromBytes(byte[] requestBytes) + { + return new PartialHttpRequest(requestBytes); + } + + public HttpRequest ToRequest() + { + HttpRequest httpRequest = new HttpRequest(); + foreach(var header in this.Headers) + { + httpRequest.Headers[header.Key] = header.Value; + } + httpRequest.Method = this.Method; + httpRequest.RequestQuery = this.RequestQuery; + httpRequest.RequestURI = this.RequestURI; + httpRequest.Body = this.Body; + foreach(var cookie in this.Cookies) + { + httpRequest.Cookies.Add(cookie); + } + return httpRequest; + } + + public byte[] GetPackedRequest() + { + return BuildRequest(); + } + + public void AddToBody(byte[] bytesToBeAdded) + { + var newBody = new byte[Body.Length + bytesToBeAdded.Length]; + if (Body.Length > 0) + { + Array.Copy(Body, 0, newBody, 0, Body.Length); + } + if (bytesToBeAdded.Length > 0) + { + Array.Copy(bytesToBeAdded, 0, newBody, Body.Length, bytesToBeAdded.Length); + } + Body = newBody; + } + + private HttpMethods GetMethod(string methodString) + { + return (HttpMethods)Enum.Parse(typeof(HttpMethods), methodString.ToUpper(), true); + } + + private HttpMethods ParseMethod(MemoryStream ms) + { + /* + * Get each character one by one. When meeting a SP character, parse the method, clear the buffer + * and continue with parsing the next step. + */ + StringBuilder parseBuffer = new StringBuilder(); + while (ms.Position < ms.Length) + { + try + { + char c = (char)ms.ReadByte(); + if (c == HttpHeaders.SP) + { + string methodString = parseBuffer.ToString(); + return GetMethod(methodString); + } + else + { + parseBuffer.Append(c); + } + } + catch (Exception e) + { + throw new InvalidMethodException("Invalid request method. Buffer: " + parseBuffer.ToString(), + new HttpRequestParsingException("Exception during parsing of http request. Buffer: " + UTF8Encoding.UTF8.GetString(ms.ToArray()), e)); + } + } + throw new IncompleteMethodException("Incomplete request method. Buffer: " + parseBuffer.ToString(), + new HttpRequestParsingException("Exception during parsing of http request. Buffer: " + UTF8Encoding.UTF8.GetString(ms.ToArray()))); + } + + private string ParseRequestURI(MemoryStream ms) + { + /* + * Get each character one by one. When meeting a SP character, parse the URI and clear the buffer. + */ + StringBuilder parseBuffer = new StringBuilder(); + ms.ReadByte(); //Ignore the first '/' + while (ms.Position < ms.Length) + { + try + { + char c = (char)ms.ReadByte(); + if (c == HttpHeaders.SP) + { + return parseBuffer.ToString(); + } + if (c == '?') + { + return parseBuffer.ToString(); + } + else + { + parseBuffer.Append(c); + } + } + catch (Exception e) + { + throw new InvalidRequestURIException("Invalid request URI. Buffer: " + parseBuffer.ToString(), + new HttpRequestParsingException("Exception during parsing of http request. Buffer: " + UTF8Encoding.UTF8.GetString(ms.ToArray()), e)); + } + } + throw new IncompleteRequestURIException("Incomplete request URI. Buffer: " + parseBuffer.ToString(), + new HttpRequestParsingException("Exception during parsing of http request. Buffer: " + UTF8Encoding.UTF8.GetString(ms.ToArray()))); + } + + private string ParseRequestQuery(MemoryStream ms) + { + /* + * Get each character one by one. When meeting a SP character, parse the URI and clear the buffer. + */ + StringBuilder parseBuffer = new StringBuilder(); + while (ms.Position < ms.Length) + { + try + { + char c = (char)ms.ReadByte(); + if (c == (byte)HttpHeaders.SP) + { + return parseBuffer.ToString(); + } + else + { + parseBuffer.Append(c); + } + } + catch (Exception e) + { + throw new InvalidRequestURIException("Invalid request query. Buffer: " + parseBuffer.ToString(), + new HttpRequestParsingException("Exception during parsing of http request. Buffer: " + UTF8Encoding.UTF8.GetString(ms.ToArray()), e)); + } + } + throw new IncompleteRequestQueryException("Incomplete request query. Buffer: " + parseBuffer.ToString(), + new HttpRequestParsingException("Exception during parsing of http request. Buffer: " + UTF8Encoding.UTF8.GetString(ms.ToArray()))); + } + + private void ParseHTTPVer(MemoryStream ms) + { + /* + * Get each character one by one. When meeting a LF character, parse the HTTPVer. + * Check if the HTTPVer matches the implementation version. + * If not, throw an exception. + */ + StringBuilder parseBuffer = new StringBuilder(); + while (ms.Position < ms.Length) + { + try + { + char c = (char)ms.ReadByte(); + if (c == HttpHeaders.CRLF[1] || c == HttpHeaders.SP) + { + string httpVer = parseBuffer.ToString(); + if (httpVer != HttpHeaders.HTTPVER) + { + throw new InvalidHttpVersionException("Invalid HTTP version. Buffer: " + parseBuffer.ToString()); + } + return; + } + else if (c == HttpHeaders.CRLF[0]) + { + /* + * If a termination character is detected, ignore it and wait for the full terminator. + */ + continue; + } + else + { + parseBuffer.Append(c); + } + } + catch (Exception e) + { + throw new InvalidHttpVersionException("Invalid HTTP version. Buffer: " + parseBuffer.ToString(), + new HttpRequestParsingException("Exception during parsing of http request. Buffer: " + UTF8Encoding.UTF8.GetString(ms.ToArray()), e)); + } + } + // If code reaches here, it means the message is incomplete. + throw new IncompleteHttpVersionException("Incomplete HTTP version. Buffer: " + parseBuffer.ToString(), + new HttpRequestParsingException("Exception during parsing of http request. Buffer: " + UTF8Encoding.UTF8.GetString(ms.ToArray()))); + } + + private string ParseHeaderKey(MemoryStream ms) + { + /* + * Get each character one by one. When meeting a ':' character, parse the header key. + */ + StringBuilder parseBuffer = new StringBuilder(); + while (ms.Position < ms.Length) + { + try + { + char c = (char)ms.ReadByte(); + if (c == ':') + { + return parseBuffer.ToString(); + } + else + { + parseBuffer.Append(c); + } + } + catch (Exception e) + { + throw new InvalidHeaderException("Invalid Header key. Buffer: " + parseBuffer.ToString(), + new HttpRequestParsingException("Exception during parsing of http request. Buffer: " + UTF8Encoding.UTF8.GetString(ms.ToArray()), e)); + } + } + throw new IncompleteHeaderKeyException("Incomplete Header key. Buffer: " + parseBuffer.ToString(), + new HttpRequestParsingException("Exception during parsing of http request. Buffer: " + UTF8Encoding.UTF8.GetString(ms.ToArray()))); + } + + private string ParseHeaderValue(MemoryStream ms) + { + /* + * Get each character one by one. When meeting a LF character, parse the value. + */ + StringBuilder parseBuffer = new StringBuilder(); + while (ms.Position < ms.Length) + { + try + { + char c = (char)ms.ReadByte(); + if (c == HttpHeaders.CRLF[1]) + { + return parseBuffer.ToString().Trim(); + } + else if (c == HttpHeaders.CRLF[0]) + { + /* + * If a termination character is detected, ignore it and wait for the full terminator. + */ + continue; + } + else + { + parseBuffer.Append(c); + } + } + catch (Exception e) + { + throw new InvalidHeaderException("Invalid header value. Buffer: " + parseBuffer.ToString(), + new HttpRequestParsingException("Exception during parsing of http request. Buffer: " + UTF8Encoding.UTF8.GetString(ms.ToArray()), e)); + } + } + throw new IncompleteHeaderValueException("Incomplete header value. Buffer: " + parseBuffer.ToString(), + new HttpRequestParsingException("Exception during parsing of http request. Buffer: " + UTF8Encoding.UTF8.GetString(ms.ToArray()))); + } + + /// + /// Build the request bytes based on the message contents. + /// + /// Array of bytes. + private byte[] BuildRequest() + { + StringBuilder requestString = new StringBuilder(); + requestString.Append(HttpHeaders.Methods[(int)Method]).Append(HttpHeaders.SP).Append(RequestURI); + if (!string.IsNullOrWhiteSpace(RequestQuery)) + { + requestString.Append('?').Append(RequestQuery); + } + requestString.Append(HttpHeaders.SP).Append(HttpHeaders.HTTPVER).Append(HttpHeaders.CRLF); + foreach (KeyValuePair header in Headers) + { + requestString.Append(header.Key).Append(':').Append(HttpHeaders.SP).Append(header.Value).Append(HttpHeaders.CRLF); + } + requestString.Append(HttpHeaders.CRLF); + if (Cookies.Count > 0) + { + requestString.Append(HttpHeaders.RequestCookieHeader).Append(':').Append(HttpHeaders.SP); + for (int i = 0; i < Cookies.Count; i++) + { + Cookie cookie = Cookies[i]; + requestString.Append(cookie.BuildCookieString()); + if (i < Cookies.Count - 1) + { + requestString.Append(';'); + } + } + } + byte[] request = new byte[requestString.Length + (Body == null ? 0 : Body.Length)]; + byte[] requestBytes = ASCIIEncoding.ASCII.GetBytes(requestString.ToString()); + Array.Copy(requestBytes, 0, request, 0, requestBytes.Length); + if (Body != null) + { + Array.Copy(Body, 0, request, requestBytes.Length, Body.Length); + } + return request; + } + /// + /// Parse the received bytes and populate the message contents. + /// + /// Message bytes to be parsed. + private void ParseRequest(byte[] requestBytes) + { + /* + * Parse the bytes one by one, respecting the reference manual. + */ + MemoryStream ms = new MemoryStream(requestBytes); + /* + * Keep the index of the byte array, to identify the message body. + * Step value indicates at what point the parsing algorithm currently is. + * Step 0 - Method, 1 - URI, 2 - Query, 3 - HTTPVer, 4 - Header, 5 - Value + */ + int step = 0; + string headerKey = string.Empty; + string headerValue = string.Empty; + while (ms.Position < ms.Length) + { + if (step == 0) + { + Method = ParseMethod(ms); + step++; + } + else if (step == 1) + { + RequestURI = ParseRequestURI(ms); + ms.Seek(-1, SeekOrigin.Current); + if (ms.ReadByte() == '?') + { + step++; + } + else + { + step += 2; + } + } + else if (step == 2) + { + RequestQuery = ParseRequestQuery(ms); + step++; + } + else if (step == 3) + { + ParseHTTPVer(ms); + step++; + } + else if (step == 4) + { + char c = Convert.ToChar(ms.ReadByte()); + if (c == HttpHeaders.CRLF[0]) + { + continue; + } + else if (c == HttpHeaders.CRLF[1]) + { + break; + } + else + { + ms.Seek(-1, SeekOrigin.Current); + headerKey = ParseHeaderKey(ms); + step++; + } + } + else if (step == 5) + { + char c = Convert.ToChar(ms.ReadByte()); + if (c == HttpHeaders.CRLF[0]) + { + continue; + } + else if (c == HttpHeaders.CRLF[1]) + { + break; + } + else + { + ms.Seek(-1, SeekOrigin.Current); + headerValue = ParseHeaderValue(ms); + if (headerKey == HttpHeaders.RequestCookieHeader) + { + Cookies.Add(new Cookie(headerValue)); + } + else + { + Headers[headerKey] = headerValue; + } + step--; + } + } + } + if (step < 4) + { + throw new IncompleteRequestException($"Incomplete request.", + new HttpRequestParsingException("Exception during parsing of http request. Buffer: " + UTF8Encoding.UTF8.GetString(ms.ToArray()))); + } + if (Headers.ContainsHeader(EntityHeaders.ContentLength)) + { + int remainingBytes = int.Parse(Headers[EntityHeaders.ContentLength]); + if (remainingBytes <= ms.Length - ms.Position) + { + this.Body = ms.ReadRemainingBytes(); + } + else + { + return; + } + } + else + { + if (ms.Length - ms.Position > 1) + { + /* + * If the message contains a body, copy it into a different array + * and save it into the HTTP message; + */ + this.Body = ms.ReadRemainingBytes(); + } + } + /* + * Trim all trailing null characters left over from SSL encryption. + */ + this.BodyString = this.BodyString.Trim('\0'); + Complete = true; + return; + } + /// + /// Parse the body into a posted from respecting the reference manual. + /// + /// Dictionary with posted from. + private Dictionary GetPostForm() + { + if (Headers.ContainsHeader("Content-Type") && Body != null) + { + if (Headers["Content-Type"] == "application/x-www-form-urlencoded") + { + Dictionary returnDictionary = new Dictionary(); + /* + * Walk through the buffer and get the form contents. + * Step 0 - key, 1 - value. + */ + string formKey = string.Empty; + int step = 0; + for (int i = 0; i < Body.Length; i++) + { + if (step == 0) + { + formKey = GetField(Body, ref i); + step++; + } + else + { + returnDictionary[formKey] = GetValue(Body, ref i); + step--; + } + } + return returnDictionary; + } + else if (Headers["Content-Type"].Contains("multipart/form-data")) + { + throw new NotImplementedException("Multipart posting not implemented"); + } + else + { + return null; + } + } + else + { + return null; + } + } + private string GetField(byte[] buffer, ref int index) + { + /* + * Get each character one by one. When meeting a LF character, parse the value. + */ + StringBuilder parseBuffer = new StringBuilder(); + for (; index < buffer.Length; index++) + { + try + { + if (buffer[index] == '=') + { + return parseBuffer.ToString().Trim(); + } + else + { + parseBuffer.Append((char)buffer[index]); + } + } + catch (Exception e) + { + throw new InvalidPostFormException("Invalid form field. Buffer: " + parseBuffer.ToString(), e); + } + } + throw new InvalidHeaderException("Invalid form field. Buffer: " + parseBuffer.ToString()); + } + private string GetValue(byte[] buffer, ref int index) + { + /* + * Get each character one by one. When meeting a LF character, parse the value. + */ + StringBuilder parseBuffer = new StringBuilder(); + for (; index < buffer.Length; index++) + { + try + { + if (buffer[index] == '&') + { + return parseBuffer.ToString().Trim(); + } + else + { + parseBuffer.Append((char)buffer[index]); + } + } + catch (Exception e) + { + throw new InvalidPostFormException("Invalid form field. Buffer: " + parseBuffer.ToString(), e); + } + } + return parseBuffer.ToString().Trim(); + } + } +} diff --git a/MTSC/MTSC.csproj b/MTSC/MTSC.csproj index 4b07783..28504f8 100644 --- a/MTSC/MTSC.csproj +++ b/MTSC/MTSC.csproj @@ -5,12 +5,12 @@ netcoreapp2.1;net48;netstandard2.0;netcoreapp3.0 - 1.6.5 + 1.6.6 Alexandru-Victor Macocian MTSC Modular TCP Server and Client - 0.1.6.5 - 0.1.6.5 + 0.1.6.6 + 0.1.6.6 true AnyCPU;x64 https://github.com/AlexMacocian/MTSC diff --git a/MTSC/Server/Handlers/HttpHandler.cs b/MTSC/Server/Handlers/HttpHandler.cs index 0e830e7..a2cc85c 100644 --- a/MTSC/Server/Handlers/HttpHandler.cs +++ b/MTSC/Server/Handlers/HttpHandler.cs @@ -23,7 +23,6 @@ namespace MTSC.Server.Handlers #region Public Properties public TimeSpan FragmentsExpirationTime { get; set; } = TimeSpan.FromSeconds(15); public double MaximumRequestSize { get; set; } = 15000; - public bool Return100Continue { get; set; } = true; #endregion #region Constructors public HttpHandler() @@ -48,16 +47,6 @@ namespace MTSC.Server.Handlers return this; } /// - /// Sets Return100Continue property - /// - /// - /// This object. - public HttpHandler WithContinueResponse(bool response) - { - this.Return100Continue = response; - return this; - } - /// /// Add a http module onto the server. /// /// Module to be added. @@ -133,7 +122,14 @@ namespace MTSC.Server.Handlers messageBytes = message.MessageBytes; } messageBytes = messageBytes.TrimTrailingNullBytes(); - request = HttpRequest.FromBytes(messageBytes); + var partialRequest = PartialHttpRequest.FromBytes(messageBytes); + if (partialRequest.Complete) + request = partialRequest.ToRequest(); + else + { + HandleIncompleteRequest(client, server, messageBytes, partialRequest); + return true; + } } catch (Exception ex) when ( ex is IncompleteHeaderKeyException || @@ -145,17 +141,9 @@ namespace MTSC.Server.Handlers ex is IncompleteRequestURIException || ex is IncompleteRequestException) { - fragmentedMessages[client] = (messageBytes, DateTime.Now); server.LogDebug(ex.Message); server.LogDebug(ex.StackTrace); - - if (Return100Continue) - { - var contResponse = new HttpResponse { StatusCode = HttpMessage.StatusCodes.Continue }; - contResponse.Headers[HttpMessage.GeneralHeaders.Connection] = "keep-alive"; - QueueResponse(client, contResponse); - } - + HandleIncompleteRequest(client, server, messageBytes); return true; } catch (Exception e) @@ -240,5 +228,19 @@ namespace MTSC.Server.Handlers } } #endregion + + private void HandleIncompleteRequest(ClientData client, Server server, byte[] messageBytes, PartialHttpRequest partialRequest = null) + { + fragmentedMessages[client] = (messageBytes, 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)) + { + server.LogDebug("Returning 100-Continue"); + var contResponse = new HttpResponse { StatusCode = HttpMessage.StatusCodes.Continue }; + contResponse.Headers[HttpMessage.GeneralHeaders.Connection] = "keep-alive"; + QueueResponse(client, contResponse); + } + } } }