// SPDX-FileCopyrightText: 2024 Filip Leonarski, Paul Scherrer Institute // SPDX-License-Identifier: GPL-3.0-only #pragma once #include #include #include #include #include #include #include "BitShuffleBlock.h" #include "../compression/CompressionAlgorithmEnum.h" #include "../common/JFJochException.h" #include "../common/CompressedImage.h" extern "C" { uint64_t bshuf_read_uint64_BE(const void* buf); }; inline size_t JFJochDecompressHperfPtr(uint8_t *output, CompressionAlgorithm algorithm, const uint8_t *source, size_t source_size, size_t nelements, size_t elem_size, size_t block_size) { if ((algorithm != CompressionAlgorithm::BSHUF_LZ4) && (algorithm != CompressionAlgorithm::BSHUF_ZSTD) && (algorithm != CompressionAlgorithm::BSHUF_ZSTD_RLE) && (algorithm != CompressionAlgorithm::BSHUF_ZSTD_RLE_HUFF)) throw JFJochException(JFJochExceptionCategory::Compression, "Algorithm not supported by hperf decompressor"); if ((block_size == 0) || ((block_size % BSHUF_BLOCKED_MULT) != 0)) throw JFJochException(JFJochExceptionCategory::Compression, "Invalid block size"); std::vector decompressed_block(block_size * elem_size); std::vector scratch(block_size * elem_size); const uint8_t *src_ptr = source; const uint8_t *const source_end = source + source_size; uint8_t *dst_ptr = output; const size_t num_full_blocks = nelements / block_size; const size_t reminder_size = nelements - num_full_blocks * block_size; const size_t last_block_size = reminder_size - reminder_size % BSHUF_BLOCKED_MULT; auto decode_block = [&](size_t current_nelements) { // Both the block length and the block itself come from the stream, which for the CBOR path // and for the XDS plugin is data we did not produce. if (source_end - src_ptr < 4) throw JFJochException(JFJochExceptionCategory::Compression, "Truncated compressed block header"); const auto compressed_size = static_cast(bshuf_read_uint32_BE(src_ptr)); src_ptr += 4; if (static_cast(source_end - src_ptr) < compressed_size) throw JFJochException(JFJochExceptionCategory::Compression, "Compressed block extends past the input buffer"); const size_t expected_size = current_nelements * elem_size; size_t decompressed_size = 0; switch (algorithm) { case CompressionAlgorithm::BSHUF_LZ4: { const int ret = LZ4_decompress_safe(reinterpret_cast(src_ptr), decompressed_block.data(), static_cast(compressed_size), static_cast(expected_size)); if (ret < 0 || static_cast(ret) != expected_size) throw JFJochException(JFJochExceptionCategory::Compression, "LZ4 decompression error"); decompressed_size = static_cast(ret); break; } case CompressionAlgorithm::BSHUF_ZSTD: case CompressionAlgorithm::BSHUF_ZSTD_RLE: case CompressionAlgorithm::BSHUF_ZSTD_RLE_HUFF: { const size_t ret = ZSTD_decompress(decompressed_block.data(), expected_size, src_ptr, compressed_size); if (ZSTD_isError(ret) || ret != expected_size) throw JFJochException(JFJochExceptionCategory::Compression, "ZSTD decompression error"); decompressed_size = ret; break; } default: throw JFJochException(JFJochExceptionCategory::Compression, "Algorithm not supported"); } if (JFJochBitUnshuffleBlock(reinterpret_cast(dst_ptr), decompressed_block.data(), scratch.data(), current_nelements, elem_size) < 0) throw JFJochException(JFJochExceptionCategory::Compression, "bitshuffle block decode error"); src_ptr += compressed_size; dst_ptr += decompressed_size; }; for (size_t i = 0; i < num_full_blocks; ++i) decode_block(block_size); if (last_block_size > 0) decode_block(last_block_size); const size_t leftover_bytes = (reminder_size % BSHUF_BLOCKED_MULT) * elem_size; if (leftover_bytes > 0) { if (static_cast(source_end - src_ptr) < leftover_bytes) throw JFJochException(JFJochExceptionCategory::Compression, "Truncated trailing bytes"); memcpy(dst_ptr, src_ptr, leftover_bytes); src_ptr += leftover_bytes; } return static_cast(src_ptr - source); } inline void JFJochDecompressPtr(uint8_t *output, CompressionAlgorithm algorithm, const uint8_t *source, size_t source_size, size_t nelements, size_t elem_size, bool use_hperf = true) { size_t block_size = 0; if (algorithm != CompressionAlgorithm::NO_COMPRESSION) { // The 12-byte bitshuffle header must be there before it can be read, and before // source_size - 12 is handed to the decompressors below. if (source_size < 12) throw JFJochException(JFJochExceptionCategory::Compression, "Buffer too short for the bitshuffle header"); if (bshuf_read_uint64_BE(const_cast(source)) != nelements * elem_size) throw JFJochException(JFJochExceptionCategory::Compression, "Mismatch in size"); auto tmp = bshuf_read_uint32_BE(source + 8); block_size = tmp / elem_size; } switch (algorithm) { case CompressionAlgorithm::NO_COMPRESSION: if (source_size != nelements * elem_size) throw JFJochException(JFJochExceptionCategory::Compression, "Mismatch in size"); memcpy(output, source, source_size); break; case CompressionAlgorithm::BSHUF_LZ4: if (use_hperf) { if (JFJochDecompressHperfPtr(output, algorithm, source + 12, source_size - 12, nelements, elem_size, block_size) != source_size - 12) throw JFJochException(JFJochExceptionCategory::Compression, "Decompression error"); } else { if (bshuf_decompress_lz4(source + 12, output, nelements, elem_size, block_size) != source_size - 12) throw JFJochException(JFJochExceptionCategory::Compression, "Decompression error"); } break; case CompressionAlgorithm::BSHUF_ZSTD_RLE: case CompressionAlgorithm::BSHUF_ZSTD_RLE_HUFF: case CompressionAlgorithm::BSHUF_ZSTD: if (use_hperf) { if (JFJochDecompressHperfPtr(output, algorithm, source + 12, source_size - 12, nelements, elem_size, block_size) != source_size - 12) throw JFJochException(JFJochExceptionCategory::Compression, "Decompression error"); } else { if (bshuf_decompress_zstd(source + 12, output, nelements, elem_size, block_size) != source_size - 12) throw JFJochException(JFJochExceptionCategory::Compression, "Decompression error"); } break; default: throw JFJochException(JFJochExceptionCategory::Compression, "Not implemented algorithm"); } } template void JFJochDecompress(std::vector &output, CompressionAlgorithm algorithm, const Ts *source_v, size_t source_size, size_t nelements, bool use_hperf = true) { output.resize(nelements); JFJochDecompressPtr((uint8_t *) output.data(), algorithm, (uint8_t *) source_v, source_size, nelements, sizeof(Td), use_hperf); } template void JFJochDecompress(std::vector &output, CompressionAlgorithm algorithm, const std::vector source_v, size_t nelements, bool use_hperf = true) { JFJochDecompress(output, algorithm, source_v.data(), source_v.size() * sizeof(Ts), nelements, use_hperf); }