From e70eb78ae3bc233bb292dfdc82437df8d404495b Mon Sep 17 00:00:00 2001 From: Alex Hultman Date: Sat, 29 Dec 2018 13:47:15 +0100 Subject: [PATCH] Hook up maxPayloadLength for Zlib too --- src/PerMessageDeflate.h | 11 +++++------ src/WebSocketContext.h | 4 ++-- 2 files changed, 7 insertions(+), 8 deletions(-) diff --git a/src/PerMessageDeflate.h b/src/PerMessageDeflate.h index 890cc27..fc60ee8 100644 --- a/src/PerMessageDeflate.h +++ b/src/PerMessageDeflate.h @@ -26,7 +26,7 @@ #include #include -#define LARGE_BUFFER_SIZE 16000 // fix this +#define LARGE_BUFFER_SIZE 1024 * 16 // fix this struct ZlibContext { /* Any returned data is valid until next same-class call. @@ -109,10 +109,9 @@ struct InflationStream { inflateInit2(&inflationStream, -15); } - std::string_view inflate(ZlibContext *zlibContext, std::string_view compressed) { - - int maxPayload = 160000; // todo: fix this + std::string_view inflate(ZlibContext *zlibContext, std::string_view compressed, size_t maxPayloadLength) { + /* We clear this one here, could be done better */ zlibContext->dynamicInflationBuffer.clear(); inflationStream.next_in = (Bytef *) compressed.data(); @@ -128,11 +127,11 @@ struct InflationStream { } zlibContext->dynamicInflationBuffer.append(zlibContext->inflationBuffer, LARGE_BUFFER_SIZE - inflationStream.avail_out); - } while (err == Z_BUF_ERROR && zlibContext->dynamicInflationBuffer.length() <= maxPayload); + } while (err == Z_BUF_ERROR && zlibContext->dynamicInflationBuffer.length() <= maxPayloadLength); inflateReset(&inflationStream); - if ((err != Z_BUF_ERROR && err != Z_OK) || zlibContext->dynamicInflationBuffer.length() > maxPayload) { + if ((err != Z_BUF_ERROR && err != Z_OK) || zlibContext->dynamicInflationBuffer.length() > maxPayloadLength) { return {nullptr, 0}; } diff --git a/src/WebSocketContext.h b/src/WebSocketContext.h index 37232ad..d034aae 100644 --- a/src/WebSocketContext.h +++ b/src/WebSocketContext.h @@ -82,7 +82,7 @@ private: ) ); - std::string_view inflatedFrame = loopData->inflationStream->inflate(loopData->zlibContext, {data, length}); + std::string_view inflatedFrame = loopData->inflationStream->inflate(loopData->zlibContext, {data, length}, webSocketContextData->maxPayloadLength); if (!inflatedFrame.length()) { forceClose(webSocketState, s); return true; @@ -129,7 +129,7 @@ private: ) ); - std::string_view inflatedFrame = loopData->inflationStream->inflate(loopData->zlibContext, {webSocketData->fragmentBuffer.data(), webSocketData->fragmentBuffer.length() - 4}); + std::string_view inflatedFrame = loopData->inflationStream->inflate(loopData->zlibContext, {webSocketData->fragmentBuffer.data(), webSocketData->fragmentBuffer.length() - 4}, webSocketContextData->maxPayloadLength); if (!inflatedFrame.length()) { forceClose(webSocketState, s); return true;