From 038111fb98b8dea3482414fe6569565297f2149f Mon Sep 17 00:00:00 2001 From: Alex Hultman Date: Sat, 22 Dec 2018 05:57:22 +0100 Subject: [PATCH] Add SHARED and DEDICATED compressor options --- misc/main.cpp | 2 +- src/App.h | 45 ++++++++++++---- src/LoopData.h | 3 +- src/PerMessageDeflate.h | 108 ++++++++++++++++++-------------------- src/WebSocket.h | 12 +++-- src/WebSocketContext.h | 4 +- src/WebSocketData.h | 12 ++++- src/WebSocketExtensions.h | 19 +++---- 8 files changed, 114 insertions(+), 91 deletions(-) diff --git a/misc/main.cpp b/misc/main.cpp index 2d4482b..fddbca6 100644 --- a/misc/main.cpp +++ b/misc/main.cpp @@ -22,7 +22,7 @@ int main(int argc, char **argv) { res->end("Hello HTTP!"); }).ws("/*", { /* Settings */ - .compression = true, + .compression = uWS::DEDICATED_COMPRESSOR, .maxPayloadLength = 16 * 1024, /* Handlers */ .open = [](auto *ws, auto *req) { diff --git a/src/App.h b/src/App.h index 21b74f4..9760e1f 100644 --- a/src/App.h +++ b/src/App.h @@ -30,6 +30,16 @@ namespace uWS { +/* Compress options (really more like PerMessageDeflateOptions) */ +enum CompressOptions { + /* Compression disabled */ + DISABLED = 0, + /* We compress using a shared non-sliding window. No added memory usage, worse compression. */ + SHARED_COMPRESSOR = 1, + /* We compress using a dedicated sliding window. Major memory usage added, better compression of similarly repeated messages. */ + DEDICATED_COMPRESSOR = 2 +}; + template struct TemplatedApp : StaticDispatch { private: @@ -53,7 +63,7 @@ public: } struct WebSocketBehavior { - bool compression = false; + CompressOptions compression = DISABLED; int maxPayloadLength = 16 * 1024; std::function *, HttpRequest *)> open = nullptr; std::function *, std::string_view, uWS::OpCode)> message = nullptr; @@ -70,14 +80,12 @@ public: /* If we are the first one to use compression, initialize it */ if (behavior.compression) { - LoopData *loopData = (LoopData *) us_loop_ext(static_dispatch(us_ssl_socket_context_loop, us_socket_context_loop)(webSocketContext->getSocketContext())); - if (!loopData->inflationStream) { + /* Initialize loop's deflate inflate streams */ + if (!loopData->zlibContext) { + loopData->zlibContext = new ZlibContext; loopData->inflationStream = new InflationStream; - } - - if (!loopData->deflationStream) { loopData->deflationStream = new DeflationStream; } } @@ -102,14 +110,24 @@ public: /* Negotiate compression */ bool perMessageDeflate = false; - if (behavior.compression) { + bool slidingDeflateWindow = false; + if (behavior.compression != DISABLED) { std::string_view extensions = req->getHeader("sec-websocket-extensions"); if (extensions.length()) { - // basically: parse(options, extensions) - ExtensionsNegotiator extensionsNegotiator(PERMESSAGE_DEFLATE | CLIENT_NO_CONTEXT_TAKEOVER); // take options + /* We never support client context takeover (the client cannot compress with a sliding window). */ + int wantedOptions = PERMESSAGE_DEFLATE | CLIENT_NO_CONTEXT_TAKEOVER; + + /* Shared compressor is the default */ + if (behavior.compression == SHARED_COMPRESSOR) { + /* Disable per-socket compressor */ + wantedOptions |= SERVER_NO_CONTEXT_TAKEOVER; + } + + /* isServer = true */ + ExtensionsNegotiator extensionsNegotiator(wantedOptions); extensionsNegotiator.readOffer(extensions); - //std::cout << extensions << " => " << extensionsNegotiator.generateOffer() << std::endl; + std::cout << extensions << " => " << extensionsNegotiator.generateOffer() << std::endl; /* Todo: remove these mid string copies */ res->writeHeader("Sec-WebSocket-Extensions", extensionsNegotiator.generateOffer()); @@ -118,6 +136,11 @@ public: if (extensionsNegotiator.getNegotiatedOptions() & PERMESSAGE_DEFLATE) { perMessageDeflate = true; } + + /* Is the server allowed to compress with a sliding window? */ + if (!(extensionsNegotiator.getNegotiatedOptions() & SERVER_NO_CONTEXT_TAKEOVER)) { + slidingDeflateWindow = true; + } } } @@ -132,7 +155,7 @@ public: webSocket->cork(); httpContext->upgradeToWebSocket( - webSocket->init(perMessageDeflate) + webSocket->init(perMessageDeflate, slidingDeflateWindow) ); /* Emit open event */ diff --git a/src/LoopData.h b/src/LoopData.h index d1988cb..f87b150 100644 --- a/src/LoopData.h +++ b/src/LoopData.h @@ -46,7 +46,8 @@ public: int corkOffset = 0; void *corkedSocket = nullptr; - /* Compression data */ + /* Per message deflate data */ + ZlibContext *zlibContext = nullptr; InflationStream *inflationStream = nullptr; DeflationStream *deflationStream = nullptr; }; diff --git a/src/PerMessageDeflate.h b/src/PerMessageDeflate.h index 1f70737..890cc27 100644 --- a/src/PerMessageDeflate.h +++ b/src/PerMessageDeflate.h @@ -14,7 +14,7 @@ * limitations under the License. */ -// inflationStream? Ciompression +/* This standalone module implements deflate / inflate streams */ #ifndef PERMESSAGEDEFLATE_H #define PERMESSAGEDEFLATE_H @@ -28,73 +28,73 @@ #define LARGE_BUFFER_SIZE 16000 // fix this -// we also need DeflationStream +struct ZlibContext { + /* Any returned data is valid until next same-class call. + * We need to have two classes to allow inflation followed + * by many deflations without modifying the inflation */ + std::string dynamicDeflationBuffer; + std::string dynamicInflationBuffer; + char *deflationBuffer; + char *inflationBuffer; + + ZlibContext() { + deflationBuffer = (char *) malloc(LARGE_BUFFER_SIZE); + inflationBuffer = (char *) malloc(LARGE_BUFFER_SIZE); + } + + ~ZlibContext() { + free(deflationBuffer); + free(inflationBuffer); + } +}; struct DeflationStream { - - // share this under the Loop - std::string dynamicZlibBuffer; z_stream deflationStream = {}; - char *zlibBuffer; DeflationStream() { - std::cout << "Constructing DeflationStream" << std::endl; - zlibBuffer = (char *) malloc(LARGE_BUFFER_SIZE); - deflateInit2(&deflationStream, 1, Z_DEFLATED, -15, 8, Z_DEFAULT_STRATEGY); } - std::string_view deflate(std::string_view raw) { + /* Deflate and optionally reset */ + std::string_view deflate(ZlibContext *zlibContext, std::string_view raw, bool reset) { + /* Odd place to clear this one, fix */ + zlibContext->dynamicDeflationBuffer.clear(); - // slidingDeflateWindow är input, length är in/ut + deflationStream.next_in = (Bytef *) raw.data(); + deflationStream.avail_in = (unsigned int) raw.length(); - z_stream *slidingDeflateWindow = nullptr; - - dynamicZlibBuffer.clear(); - - z_stream *compressor = slidingDeflateWindow ? slidingDeflateWindow : &deflationStream; - - compressor->next_in = (Bytef *) raw.data(); - compressor->avail_in = (unsigned int) raw.length(); - - // note: zlib requires more than 6 bytes with Z_SYNC_FLUSH + /* This buffer size has to be at least 6 bytes for Z_SYNC_FLUSH to work */ const int DEFLATE_OUTPUT_CHUNK = LARGE_BUFFER_SIZE; int err; do { - compressor->next_out = (Bytef *) zlibBuffer; - compressor->avail_out = DEFLATE_OUTPUT_CHUNK; + deflationStream.next_out = (Bytef *) zlibContext->deflationBuffer; + deflationStream.avail_out = DEFLATE_OUTPUT_CHUNK; - err = ::deflate(compressor, Z_SYNC_FLUSH); - if (Z_OK == err && compressor->avail_out == 0) { - dynamicZlibBuffer.append(zlibBuffer, DEFLATE_OUTPUT_CHUNK - compressor->avail_out); + err = ::deflate(&deflationStream, Z_SYNC_FLUSH); + if (Z_OK == err && deflationStream.avail_out == 0) { + zlibContext->dynamicDeflationBuffer.append(zlibContext->deflationBuffer, DEFLATE_OUTPUT_CHUNK - deflationStream.avail_out); continue; } else { break; } } while (true); - // note: should not change avail_out - if (!slidingDeflateWindow) { - deflateReset(compressor); + /* This must not change avail_out */ + if (reset) { + deflateReset(&deflationStream); } - if (dynamicZlibBuffer.length()) { - dynamicZlibBuffer.append(zlibBuffer, DEFLATE_OUTPUT_CHUNK - compressor->avail_out); + if (zlibContext->dynamicDeflationBuffer.length()) { + zlibContext->dynamicDeflationBuffer.append(zlibContext->deflationBuffer, DEFLATE_OUTPUT_CHUNK - deflationStream.avail_out); - return {(char *) dynamicZlibBuffer.data(), dynamicZlibBuffer.length() - 4}; - - //length = dynamicZlibBuffer.length() - 4; - //return (char *) dynamicZlibBuffer.data(); + return {(char *) zlibContext->dynamicDeflationBuffer.data(), zlibContext->dynamicDeflationBuffer.length() - 4}; } return { - zlibBuffer, - DEFLATE_OUTPUT_CHUNK - compressor->avail_out - 4 + zlibContext->deflationBuffer, + DEFLATE_OUTPUT_CHUNK - deflationStream.avail_out - 4 }; - - //length = DEFLATE_OUTPUT_CHUNK - compressor->avail_out - 4; - //return zlibBuffer; } ~DeflationStream() { @@ -102,54 +102,46 @@ struct DeflationStream { } }; -// the loop holds one of these struct InflationStream { - - // share this under the Loop - std::string dynamicZlibBuffer; z_stream inflationStream = {}; - char *zlibBuffer; InflationStream() { - std::cout << "Initliazing shared compression" << std::endl; - zlibBuffer = (char *) malloc(LARGE_BUFFER_SIZE); - inflateInit2(&inflationStream, -15); } - std::string_view inflate(std::string_view compressed) { + std::string_view inflate(ZlibContext *zlibContext, std::string_view compressed) { int maxPayload = 160000; // todo: fix this - dynamicZlibBuffer.clear(); + zlibContext->dynamicInflationBuffer.clear(); inflationStream.next_in = (Bytef *) compressed.data(); inflationStream.avail_in = (unsigned int) compressed.length(); int err; do { - inflationStream.next_out = (Bytef *) zlibBuffer; + inflationStream.next_out = (Bytef *) zlibContext->inflationBuffer; inflationStream.avail_out = LARGE_BUFFER_SIZE; err = ::inflate(&inflationStream, Z_FINISH); if (!inflationStream.avail_in) { break; } - dynamicZlibBuffer.append(zlibBuffer, LARGE_BUFFER_SIZE - inflationStream.avail_out); - } while (err == Z_BUF_ERROR && dynamicZlibBuffer.length() <= maxPayload); + zlibContext->dynamicInflationBuffer.append(zlibContext->inflationBuffer, LARGE_BUFFER_SIZE - inflationStream.avail_out); + } while (err == Z_BUF_ERROR && zlibContext->dynamicInflationBuffer.length() <= maxPayload); inflateReset(&inflationStream); - if ((err != Z_BUF_ERROR && err != Z_OK) || dynamicZlibBuffer.length() > maxPayload) { + if ((err != Z_BUF_ERROR && err != Z_OK) || zlibContext->dynamicInflationBuffer.length() > maxPayload) { return {nullptr, 0}; } - if (dynamicZlibBuffer.length()) { - dynamicZlibBuffer.append(zlibBuffer, LARGE_BUFFER_SIZE - inflationStream.avail_out); - return {dynamicZlibBuffer.data(), dynamicZlibBuffer.length()}; + if (zlibContext->dynamicInflationBuffer.length()) { + zlibContext->dynamicInflationBuffer.append(zlibContext->inflationBuffer, LARGE_BUFFER_SIZE - inflationStream.avail_out); + return {zlibContext->dynamicInflationBuffer.data(), zlibContext->dynamicInflationBuffer.length()}; } - return {zlibBuffer, LARGE_BUFFER_SIZE - inflationStream.avail_out}; + return {zlibContext->inflationBuffer, LARGE_BUFFER_SIZE - inflationStream.avail_out}; } }; diff --git a/src/WebSocket.h b/src/WebSocket.h index b7709de..14e1f2d 100644 --- a/src/WebSocket.h +++ b/src/WebSocket.h @@ -33,8 +33,8 @@ private: using SOCKET_TYPE = typename StaticDispatch::SOCKET_TYPE; using StaticDispatch::static_dispatch; - void *init(bool perMessageDeflate) { - new (static_dispatch(us_ssl_socket_ext, us_socket_ext)((SOCKET_TYPE *) this)) WebSocketData(perMessageDeflate); + void *init(bool perMessageDeflate, bool slidingCompression) { + new (static_dispatch(us_ssl_socket_ext, us_socket_ext)((SOCKET_TYPE *) this)) WebSocketData(perMessageDeflate, slidingCompression); return this; } public: @@ -58,9 +58,13 @@ public: /* Check and correct the compress hint */ if (opCode < 3 && webSocketData->compressionStatus == WebSocketData::ENABLED) { - // todo: shared deflate window LoopData *loopData = Super::getLoopData(); - message = loopData->deflationStream->deflate(message); + /* Compress using either shared or dedicated deflationStream */ + if (webSocketData->deflationStream) { + message = webSocketData->deflationStream->deflate(loopData->zlibContext, message, false); + } else { + message = loopData->deflationStream->deflate(loopData->zlibContext, message, true); + } } else { compress = false; } diff --git a/src/WebSocketContext.h b/src/WebSocketContext.h index 65a4bc2..ba186e1 100644 --- a/src/WebSocketContext.h +++ b/src/WebSocketContext.h @@ -101,7 +101,7 @@ private: ); - std::string_view inflatedFrame = loopData->inflationStream->inflate({data, length}); + std::string_view inflatedFrame = loopData->inflationStream->inflate(loopData->zlibContext, {data, length}); if (!inflatedFrame.length()) { forceClose(webSocketState, s); return true; @@ -156,7 +156,7 @@ private: ); - std::string_view inflatedFrame = loopData->inflationStream->inflate({webSocketData->fragmentBuffer.data(), webSocketData->fragmentBuffer.length() - 4}); + std::string_view inflatedFrame = loopData->inflationStream->inflate(loopData->zlibContext, {webSocketData->fragmentBuffer.data(), webSocketData->fragmentBuffer.length() - 4}); if (!inflatedFrame.length()) { forceClose(webSocketState, s); return true; diff --git a/src/WebSocketData.h b/src/WebSocketData.h index b64feb7..1eb08cd 100644 --- a/src/WebSocketData.h +++ b/src/WebSocketData.h @@ -19,6 +19,7 @@ #include "WebSocketProtocol.h" #include "AsyncSocketData.h" +#include "PerMessageDeflate.h" #include @@ -37,10 +38,17 @@ private: ENABLED, COMPRESSED_FRAME } compressionStatus; + + /* We might have a dedicated compressor */ + DeflationStream *deflationStream = nullptr; public: - WebSocketData(bool perMessageDeflate) : WebSocketState() { - //std::cout << "perMessageDeflate: " << perMessageDeflate << std::endl; + WebSocketData(bool perMessageDeflate, bool slidingCompression) : WebSocketState() { compressionStatus = perMessageDeflate ? ENABLED : DISABLED; + + /* Initialize the dedicated sliding window */ + if (perMessageDeflate && slidingCompression) { + deflationStream = new DeflationStream; + } } }; diff --git a/src/WebSocketExtensions.h b/src/WebSocketExtensions.h index 43afe35..111cbb5 100644 --- a/src/WebSocketExtensions.h +++ b/src/WebSocketExtensions.h @@ -129,16 +129,10 @@ std::string ExtensionsNegotiator::generateOffer() { extensionsOffer += "; client_no_context_takeover"; } - // we do not support accepting this yet - // todo: if we agree on this, do not allocate a compressor - // per socket! - - // It is RECOMMENDED that a server supports the - // "server_no_context_takeover" extension parameter in an extension - // negotiation offer. - if (options & Options::SERVER_NO_CONTEXT_TAKEOVER) { - //extensionsOffer += "; server_no_context_takeover"; - } + /* It is questionable sending this improves anything */ + /*if (options & Options::SERVER_NO_CONTEXT_TAKEOVER) { + extensionsOffer += "; server_no_context_takeover"; + }*/ } return extensionsOffer; @@ -153,11 +147,12 @@ void ExtensionsNegotiator::readOffer(std::string_view offer) { options |= CLIENT_NO_CONTEXT_TAKEOVER; } + /* We leave this option for us to read even if the client did not send it */ if (extensionsParser.serverNoContextTakeover) { options |= SERVER_NO_CONTEXT_TAKEOVER; - } else { + }/* else { options &= ~SERVER_NO_CONTEXT_TAKEOVER; - } + }*/ } else { options &= ~PERMESSAGE_DEFLATE; }