diff --git a/src/node_zlib.cc b/src/node_zlib.cc index 95201624cfa2fa..edc69f4582b357 100644 --- a/src/node_zlib.cc +++ b/src/node_zlib.cc @@ -47,10 +47,11 @@ #include +#include #include #include #include -#include +#include namespace node { @@ -349,6 +350,7 @@ class ZstdCompressContext final : public ZstdContext { DeleteFnPtr cctx_; uint64_t pledged_src_size_ = ZSTD_CONTENTSIZE_UNKNOWN; + std::optional consumed_src_size_; }; class ZstdDecompressContext final : public ZstdContext { @@ -1653,6 +1655,11 @@ void ZstdCompressContext::Close() { CompressionError ZstdCompressContext::Init(uint64_t pledged_src_size, std::string_view dictionary) { pledged_src_size_ = pledged_src_size; + if (pledged_src_size == ZSTD_CONTENTSIZE_UNKNOWN) { + consumed_src_size_.reset(); + } else { + consumed_src_size_ = 0; + } #ifdef NODE_BUNDLED_ZSTD ZSTD_customMem custom_mem = { CompressionStreamMemoryOwner::AllocForBrotli, @@ -1692,12 +1699,26 @@ CompressionError ZstdCompressContext::ResetStream() { } void ZstdCompressContext::DoThreadPoolWork() { + // Zstd overrides a configured pledge when the first call uses ZSTD_e_end. + size_t const input_pos = input_.pos; size_t const remaining = ZSTD_compressStream2(cctx_.get(), &output_, &input_, flush_); + if (consumed_src_size_.has_value()) { + *consumed_src_size_ += input_.pos - input_pos; + } if (ZSTD_isError(remaining)) { error_ = ZSTD_getErrorCode(remaining); error_code_string_ = ZstdStrerror(error_); error_string_ = ZSTD_getErrorString(error_); + } else if (remaining == 0 && flush_ == ZSTD_e_end && + consumed_src_size_.has_value()) { + uint64_t const consumed_src_size = *consumed_src_size_; + consumed_src_size_.reset(); + if (consumed_src_size != pledged_src_size_) { + error_ = ZSTD_error_srcSize_wrong; + error_code_string_ = ZstdStrerror(error_); + error_string_ = ZSTD_getErrorString(error_); + } } } diff --git a/test/parallel/test-zlib-zstd-pledged-src-size.js b/test/parallel/test-zlib-zstd-pledged-src-size.js index b1e32e14ae732a..a1222b8e809a10 100644 --- a/test/parallel/test-zlib-zstd-pledged-src-size.js +++ b/test/parallel/test-zlib-zstd-pledged-src-size.js @@ -3,6 +3,11 @@ const common = require('../common'); const assert = require('assert'); const zlib = require('zlib'); +const pledgedSrcSizeError = { + code: 'ZSTD_error_srcSize_wrong', + errno: zlib.constants.ZSTD_error_srcSize_wrong, +}; + function compressWithPledgedSrcSize({ pledgedSrcSize, actualSrcSize }) { return new Promise((resolve, reject) => { const compressor = zlib.createZstdCompress({ pledgedSrcSize }); @@ -18,20 +23,56 @@ function compressWithPledgedSrcSize({ pledgedSrcSize, actualSrcSize }) { // Compression should only succeed if sizes match assert.strictEqual(pledgedSrcSize, actualSrcSize); }, (error) => { - assert.strictEqual(error.code, 'ZSTD_error_srcSize_wrong'); + assert.strictEqual(error.code, pledgedSrcSizeError.code); + assert.strictEqual(error.errno, pledgedSrcSizeError.errno); // Size error should only happen when sizes do not match assert.notStrictEqual(pledgedSrcSize, actualSrcSize); }).then(common.mustCall()); } -compressWithPledgedSrcSize({ pledgedSrcSize: 0, actualSrcSize: 0 }); +function compressSyncWithPledgedSrcSize({ pledgedSrcSize, actualSrcSize }) { + const compress = () => zlib.zstdCompressSync( + 'x'.repeat(actualSrcSize), + { pledgedSrcSize }, + ); -compressWithPledgedSrcSize({ pledgedSrcSize: 0, actualSrcSize: 42 }); + if (pledgedSrcSize === actualSrcSize) { + compress(); + } else { + assert.throws(compress, pledgedSrcSizeError); + } +} -compressWithPledgedSrcSize({ pledgedSrcSize: 13, actualSrcSize: 42 }); +const testCases = [ + { pledgedSrcSize: 0, actualSrcSize: 0 }, + { pledgedSrcSize: 0, actualSrcSize: 42 }, + { pledgedSrcSize: 1, actualSrcSize: 42 }, + { pledgedSrcSize: 13, actualSrcSize: 42 }, + { pledgedSrcSize: 42, actualSrcSize: 0 }, + { pledgedSrcSize: 42, actualSrcSize: 13 }, + { pledgedSrcSize: 42, actualSrcSize: 42 }, +]; -compressWithPledgedSrcSize({ pledgedSrcSize: 42, actualSrcSize: 0 }); +for (const testCase of testCases) { + compressWithPledgedSrcSize(testCase); + compressSyncWithPledgedSrcSize(testCase); +} + +const retryInput = Buffer.allocUnsafe(256 * 1024); +let randomState = 0x12345678; +for (let i = 0; i < retryInput.length; i++) { + randomState = (Math.imul(randomState, 1664525) + 1013904223) | 0; + retryInput[i] = randomState >>> 24; +} -compressWithPledgedSrcSize({ pledgedSrcSize: 42, actualSrcSize: 13 }); +const compressed = zlib.zstdCompressSync(retryInput, { + pledgedSrcSize: retryInput.length, + chunkSize: 64, +}); +assert.ok(compressed.length > 64); +assert.deepStrictEqual(zlib.zstdDecompressSync(compressed), retryInput); -compressWithPledgedSrcSize({ pledgedSrcSize: 42, actualSrcSize: 42 }); +assert.throws(() => zlib.zstdCompressSync(retryInput, { + pledgedSrcSize: retryInput.length - 1, + chunkSize: 64, +}), pledgedSrcSizeError);