Use std::optional in permessage-deflate

This commit is contained in:
Alex Hultman
2020-10-23 02:16:49 +02:00
parent 59a38b3a13
commit 21cb4846df
2 changed files with 18 additions and 17 deletions
+10 -9
View File
@@ -47,6 +47,7 @@ namespace uWS {
#endif
#include <string>
#include <optional>
namespace uWS {
@@ -54,9 +55,9 @@ namespace uWS {
#ifdef UWS_NO_ZLIB
struct ZlibContext {};
struct InflationStream {
std::pair<std::string_view, bool> inflate(ZlibContext *zlibContext, std::string_view compressed, size_t maxPayloadLength) {
std::optional<std::string_view> inflate(ZlibContext *zlibContext, std::string_view compressed, size_t maxPayloadLength) {
/* Anything here goes, it is never going to be called */
return {compressed, false};
return std::nullopt;
}
};
struct DeflationStream {
@@ -164,7 +165,7 @@ struct DeflationStream {
if (zlibContext->dynamicDeflationBuffer.length()) {
zlibContext->dynamicDeflationBuffer.append(zlibContext->deflationBuffer, DEFLATE_OUTPUT_CHUNK - deflationStream.avail_out);
return {(char *) zlibContext->dynamicDeflationBuffer.data(), zlibContext->dynamicDeflationBuffer.length() - 4};
return std::string_view((char *) zlibContext->dynamicDeflationBuffer.data(), zlibContext->dynamicDeflationBuffer.length() - 4);
}
/* Note: We will get an interger overflow resulting in heap buffer overflow if Z_BUF_ERROR is returned
@@ -192,7 +193,7 @@ struct InflationStream {
}
/* Zero length inflates are possible and valid */
std::pair<std::string_view, bool> inflate(ZlibContext *zlibContext, std::string_view compressed, size_t maxPayloadLength) {
std::optional<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();
@@ -218,7 +219,7 @@ struct InflationStream {
inflateReset(&inflationStream);
if ((err != Z_BUF_ERROR && err != Z_OK) || zlibContext->dynamicInflationBuffer.length() > maxPayloadLength) {
return {{nullptr, 0}, false};
return std::nullopt;
}
if (zlibContext->dynamicInflationBuffer.length()) {
@@ -226,18 +227,18 @@ struct InflationStream {
/* Let's be strict about the max size */
if (zlibContext->dynamicInflationBuffer.length() > maxPayloadLength) {
return {{nullptr, 0}, false};
return std::nullopt;
}
return {zlibContext->dynamicInflationBuffer.data(), zlibContext->dynamicInflationBuffer.length()};
return std::string_view(zlibContext->dynamicInflationBuffer.data(), zlibContext->dynamicInflationBuffer.length());
}
/* Let's be strict about the max size */
if ((LARGE_BUFFER_SIZE - inflationStream.avail_out) > maxPayloadLength) {
return {{nullptr, 0}, false};
return std::nullopt;
}
return {{zlibContext->inflationBuffer, LARGE_BUFFER_SIZE - inflationStream.avail_out}, true};
return std::string_view(zlibContext->inflationBuffer, LARGE_BUFFER_SIZE - inflationStream.avail_out);
}
};
+8 -8
View File
@@ -72,13 +72,13 @@ private:
webSocketData->compressionStatus = WebSocketData::CompressionStatus::ENABLED;
LoopData *loopData = (LoopData *) us_loop_ext(us_socket_context_loop(SSL, us_socket_context(SSL, (us_socket_t *) s)));
auto [inflatedFrame, valid] = loopData->inflationStream->inflate(loopData->zlibContext, {data, length}, webSocketContextData->maxPayloadLength);
if (!valid) {
auto inflatedFrame = loopData->inflationStream->inflate(loopData->zlibContext, {data, length}, webSocketContextData->maxPayloadLength);
if (!inflatedFrame.has_value()) {
forceClose(webSocketState, s, ERR_TOO_BIG_MESSAGE_INFLATION);
return true;
} else {
data = (char *) inflatedFrame.data();
length = inflatedFrame.length();
data = (char *) inflatedFrame->data();
length = inflatedFrame->length();
}
}
@@ -124,13 +124,13 @@ private:
)
);
auto [inflatedFrame, valid] = loopData->inflationStream->inflate(loopData->zlibContext, {webSocketData->fragmentBuffer.data(), webSocketData->fragmentBuffer.length() - 4}, webSocketContextData->maxPayloadLength);
if (!valid) {
auto inflatedFrame = loopData->inflationStream->inflate(loopData->zlibContext, {webSocketData->fragmentBuffer.data(), webSocketData->fragmentBuffer.length() - 4}, webSocketContextData->maxPayloadLength);
if (!inflatedFrame.has_value()) {
forceClose(webSocketState, s, ERR_TOO_BIG_MESSAGE_INFLATION);
return true;
} else {
data = (char *) inflatedFrame.data();
length = inflatedFrame.length();
data = (char *) inflatedFrame->data();
length = inflatedFrame->length();
}