#include "deflate.h"
#include "../storages/common/inflate_fast.h"
#include <zlib-ng.h>
namespace whiteout {
namespace {
constexpr int kDefaultLevel = 6;
std::vector<u8> inflateStream(std::span<const u8> data, size_t expectedSize, std::string* err) {
auto setErr = [&](const char* m) {
if (err)
*err = m;
};
if (data.size() < 6) {
setErr("Zlib data too short");
return {};
}
zng_stream zs{};
if (zng_inflateInit(&zs) != Z_OK) {
setErr("inflateInit failed");
return {};
}
size_t cap = expectedSize > 0 ? expectedSize : data.size() * 4;
if (cap < 4096)
cap = 4096;
std::vector<u8> out(cap);
zs.next_in = data.data();
zs.avail_in = static_cast<uint32_t>(data.size());
zs.next_out = out.data();
zs.avail_out = static_cast<uint32_t>(out.size());
for (;;) {
int const rc = zng_inflate(&zs, Z_NO_FLUSH);
if (rc == Z_STREAM_END)
break;
if (rc != Z_OK && rc != Z_BUF_ERROR) {
const char* msg = zs.msg;
zng_inflateEnd(&zs);
setErr(msg ? msg : "inflate failed");
return {};
}
if (zs.avail_out == 0) {
size_t const used = out.size();
out.resize(out.size() * 2);
zs.next_out = out.data() + used;
zs.avail_out = static_cast<uint32_t>(out.size() - used);
} else if (rc == Z_BUF_ERROR) {
zng_inflateEnd(&zs);
setErr("Truncated or incomplete zlib stream");
return {};
}
}
out.resize(zs.total_out);
zng_inflateEnd(&zs);
return out;
}
}
std::vector<u8> zlib_decompress(std::span<const u8> data, std::string* out_error,
size_t expectedSize) {
return inflateStream(data, expectedSize, out_error);
}
std::vector<u8> zlib_compress(std::span<const u8> data, std::string* out_error) {
size_t const bound = zng_compressBound(data.size());
std::vector<u8> out(bound);
size_t destLen = bound;
int const rc = zng_compress2(out.data(), &destLen, data.data(), data.size(), kDefaultLevel);
if (rc != Z_OK) {
if (out_error)
*out_error = "zlib compression failed";
return {};
}
out.resize(destLen);
return out;
}
}
namespace whiteout::storages::common {
std::vector<u8> zlibInflateFast(std::span<const u8> src, size_t expectedSize) {
if (expectedSize == 0)
return ::whiteout::zlib_decompress(src, nullptr, 0);
if (src.size() < 6)
return {};
std::vector<u8> out(expectedSize);
size_t destLen = expectedSize;
int const rc = zng_uncompress(out.data(), &destLen, src.data(), src.size());
if (rc != Z_OK) {
return ::whiteout::zlib_decompress(src, nullptr, expectedSize);
}
out.resize(destLen);
return out;
}
}