#include "../../common/byte_order.h"
#include "../../common/inflate_fast.h"
#include "../../common/md5.h"
#include "../../common/zlib.h"
#include "blte.h"
#include "crypto.h"
#include <whiteout/interfaces.h>
#include <whiteout/utils/job_group.h>
#include <algorithm>
#include <atomic>
#include <cstring>
#include <mutex>
namespace whiteout::storages::casc {
using storages::common::readBE32;
using storages::common::writeBE32;
static constexpr u32 kBlteMagic = 0x424C5445;
static constexpr u32 kDefaultBlteFrameSize = 0x10000;
static constexpr size_t kBlteMinHeaderSize = 8;
static constexpr size_t kBlteFrameTableEntrySize = 24;
static constexpr u8 kBlteMultiFrameFlags = 0x0F;
namespace BlteFrameMode {
static constexpr u8 kRaw = 'N'; static constexpr u8 kZlib = 'Z'; static constexpr u8 kEncrypted = 'E'; static constexpr u8 kRecursive = 'F'; }
struct FrameDecodeResult {
std::vector<u8> data;
bool success = false;
std::string error;
};
static FrameDecodeResult decodeFramePayload(std::span<const u8> payload,
u32 expectedUncompressedSize, const KeyRing* keys) {
FrameDecodeResult result;
if (payload.empty()) {
result.error = "empty frame payload";
return result;
}
u8 const mode = payload[0];
auto inner = payload.subspan(1);
switch (mode) {
case BlteFrameMode::kRaw: {
result.data.assign(inner.begin(), inner.end());
result.success = true;
break;
}
case BlteFrameMode::kZlib: {
if (expectedUncompressedSize > 0)
result.data = storages::common::zlibInflateFast(inner, expectedUncompressedSize);
else
result.data = storages::common::zlibDecompress(inner);
if (result.data.empty() && expectedUncompressedSize > 0) {
result.error = "zlib decompression failed";
} else {
result.success = true;
}
break;
}
case BlteFrameMode::kEncrypted: {
if (inner.size() < 9) {
result.error = "encrypted frame too small for key header";
return result;
}
u64 keyName = 0;
std::memcpy(&keyName, inner.data(), 8);
u8 const ivSize = inner[8];
if (inner.size() < 9u + ivSize) {
result.error = "encrypted frame truncated at IV";
return result;
}
if (!keys) {
result.error = "encrypted frame but no KeyRing provided";
return result;
}
const auto* key = keys->findKey(keyName);
if (!key) {
result.error = "missing encryption key 0x" + std::to_string(keyName); return result;
}
auto encStart = inner.subspan(9 + ivSize);
std::vector<u8> decrypted(encStart.begin(), encStart.end());
if (ivSize >= 8) {
std::array<u8, 8> iv{};
std::memcpy(iv.data(), inner.data() + 9, 8);
salsa20Decrypt(decrypted, *key, std::span<const u8, 8>(iv));
}
auto innerResult = decodeFramePayload(decrypted, expectedUncompressedSize, keys);
result = std::move(innerResult);
break;
}
case BlteFrameMode::kRecursive: {
auto innerResult = blteDecode(inner, keys, nullptr);
result.data = std::move(innerResult.data);
result.success = innerResult.success;
result.error = std::move(innerResult.error);
break;
}
default:
result.error = "unknown BLTE encoding mode: 0x";
result.error += "0123456789ABCDEF"[(mode >> 4) & 0xF];
result.error += "0123456789ABCDEF"[mode & 0xF];
break;
}
return result;
}
BlteDecodeResult blteDecode(std::span<const u8> blteData, const KeyRing* keys,
interfaces::WorkerPool* pool) {
BlteDecodeResult result;
if (blteData.size() < kBlteMinHeaderSize) {
result.error = "data too small for BLTE header";
return result;
}
u32 const magic = readBE32(blteData.data());
if (magic != kBlteMagic) {
result.error = "invalid BLTE magic";
return result;
}
u32 const headerSize = readBE32(blteData.data() + 4);
struct FrameInfo {
u32 compressedSize;
u32 uncompressedSize;
};
std::vector<FrameInfo> frames;
size_t dataOffset = 0;
if (headerSize == 0) {
FrameInfo fi;
fi.compressedSize = u32(blteData.size() - kBlteMinHeaderSize);
fi.uncompressedSize = 0; frames.push_back(fi);
dataOffset = kBlteMinHeaderSize;
} else {
size_t const hdrStart = kBlteMinHeaderSize;
if (blteData.size() < hdrStart + 4) {
result.error = "truncated BLTE frame table header";
return result;
}
u32 const frameCount = (u32(blteData[hdrStart + 1]) << 16) |
(u32(blteData[hdrStart + 2]) << 8) | u32(blteData[hdrStart + 3]);
size_t const tableStart = hdrStart + 4;
size_t const tableSize = frameCount * kBlteFrameTableEntrySize;
if (blteData.size() < tableStart + tableSize) {
result.error = "truncated BLTE frame table";
return result;
}
frames.reserve(frameCount);
for (u32 i = 0; i < frameCount; ++i) {
const u8* entry = blteData.data() + tableStart + i * kBlteFrameTableEntrySize;
FrameInfo fi;
fi.compressedSize = readBE32(entry);
fi.uncompressedSize = readBE32(entry + 4);
frames.push_back(fi);
}
dataOffset = tableStart + tableSize;
}
std::vector<size_t> offsets(frames.size());
{
size_t off = dataOffset;
for (size_t i = 0; i < frames.size(); ++i) {
offsets[i] = off;
off += frames[i].compressedSize;
}
if (off > blteData.size()) {
result.error = "frame data extends past end of BLTE blob";
return result;
}
}
if (pool && frames.size() > 1) {
utils::JobGroup jobGroup;
std::vector<FrameDecodeResult> frameResults(frames.size());
std::mutex errMutex;
std::string firstError;
jobGroup.add(frames.size());
for (size_t i = 0; i < frames.size(); ++i) {
interfaces::WorkerTask task;
task.fn = [&, i]() {
auto payload = blteData.subspan(offsets[i], frames[i].compressedSize);
frameResults[i] = decodeFramePayload(payload, frames[i].uncompressedSize, keys);
if (!frameResults[i].success) {
std::lock_guard<std::mutex> const lock(errMutex);
if (firstError.empty())
firstError = frameResults[i].error;
}
jobGroup.done();
};
pool->submit(task);
}
jobGroup.wait();
if (!firstError.empty()) {
result.error = std::move(firstError);
return result;
}
size_t totalSize = 0;
for (auto& fr : frameResults)
totalSize += fr.data.size();
result.data.reserve(totalSize);
for (auto& fr : frameResults) {
result.data.insert(result.data.end(), std::make_move_iterator(fr.data.begin()),
std::make_move_iterator(fr.data.end()));
fr.data.clear();
fr.data.shrink_to_fit();
}
} else {
if (frames.size() == 1) {
auto payload = blteData.subspan(offsets[0], frames[0].compressedSize);
auto fr = decodeFramePayload(payload, frames[0].uncompressedSize, keys);
if (!fr.success) {
result.error = "frame 0: " + fr.error;
return result;
}
result.data = std::move(fr.data);
} else {
size_t totalUncompressed = 0;
for (auto& f : frames)
totalUncompressed += f.uncompressedSize;
if (totalUncompressed > 0)
result.data.reserve(totalUncompressed);
for (size_t i = 0; i < frames.size(); ++i) {
auto payload = blteData.subspan(offsets[i], frames[i].compressedSize);
auto fr = decodeFramePayload(payload, frames[i].uncompressedSize, keys);
if (!fr.success) {
result.error = "frame " + std::to_string(i) + ": " + fr.error;
return result;
}
result.data.insert(result.data.end(), fr.data.begin(), fr.data.end());
}
}
}
result.success = true;
return result;
}
BlteFrameLayout blteParseFrameLayout(std::span<const u8> blob) {
BlteFrameLayout layout;
if (blob.size() < kBlteMinHeaderSize) {
layout.error = "data too small for BLTE header";
return layout;
}
u32 const magic = readBE32(blob.data());
if (magic != kBlteMagic) {
layout.error = "invalid BLTE magic";
return layout;
}
u32 const headerSize = readBE32(blob.data() + 4);
if (headerSize == 0) {
BlteFrameLayout::Frame f;
f.compressedSize = u32(blob.size() - kBlteMinHeaderSize);
f.uncompressedSize = 0;
layout.frames.push_back(f);
layout.offsets.push_back(kBlteMinHeaderSize);
layout.valid = true;
return layout;
}
size_t const hdrStart = kBlteMinHeaderSize;
if (blob.size() < hdrStart + 4) {
layout.error = "truncated BLTE frame table header";
return layout;
}
u32 const frameCount =
(u32(blob[hdrStart + 1]) << 16) | (u32(blob[hdrStart + 2]) << 8) | u32(blob[hdrStart + 3]);
size_t const tableStart = hdrStart + 4;
size_t const tableSize = frameCount * kBlteFrameTableEntrySize;
if (blob.size() < tableStart + tableSize) {
layout.error = "truncated BLTE frame table";
return layout;
}
layout.frames.reserve(frameCount);
layout.offsets.reserve(frameCount);
size_t off = tableStart + tableSize;
for (u32 j = 0; j < frameCount; ++j) {
const u8* entry = blob.data() + tableStart + j * kBlteFrameTableEntrySize;
BlteFrameLayout::Frame f;
f.compressedSize = readBE32(entry);
f.uncompressedSize = readBE32(entry + 4);
layout.frames.push_back(f);
layout.offsets.push_back(off);
off += f.compressedSize;
}
if (off > blob.size()) {
layout.error = "frame data extends past end of BLTE blob";
return layout;
}
layout.valid = true;
return layout;
}
static BlteBatchResult decodeSingleFrameBlte(std::span<const u8> blob,
const BlteFrameLayout& layout, const KeyRing* keys) {
BlteBatchResult result;
auto payload = blob.subspan(layout.offsets[0], layout.frames[0].compressedSize);
auto fr = decodeFramePayload(payload, layout.frames[0].uncompressedSize, keys);
result.data = std::move(fr.data);
result.success = fr.success;
result.error = std::move(fr.error);
return result;
}
BlteBatchResult blteDecodeFrame(std::span<const u8> blteData, const BlteFrameLayout& layout,
size_t frameIdx, const KeyRing* keys) {
BlteBatchResult result;
if (!layout.valid || frameIdx >= layout.frames.size()) {
result.error = "invalid frame index";
return result;
}
auto& frame = layout.frames[frameIdx];
size_t const off = layout.offsets[frameIdx];
if (off + frame.compressedSize > blteData.size()) {
result.error = "frame extends past blob end";
return result;
}
auto payload = blteData.subspan(off, frame.compressedSize);
auto fr = decodeFramePayload(payload, frame.uncompressedSize, keys);
result.data = std::move(fr.data);
result.success = fr.success;
result.error = std::move(fr.error);
return result;
}
std::vector<BlteBatchResult> blteDecodeBatch(std::span<const BlteBatchEntry> entries,
const KeyRing* keys, interfaces::WorkerPool* pool) {
std::vector<BlteBatchResult> results(entries.size());
if (entries.empty())
return results;
if (!pool || pool->threadCount() == 0) {
for (size_t i = 0; i < entries.size(); ++i) {
auto decoded = blteDecode(entries[i].blteData, keys, nullptr);
results[i].data = std::move(decoded.data);
results[i].success = decoded.success;
results[i].error = std::move(decoded.error);
}
return results;
}
auto testSem = pool->createTimelineSemaphore();
if (!testSem) {
for (size_t i = 0; i < entries.size(); ++i) {
auto decoded = blteDecode(entries[i].blteData, keys, pool);
results[i].data = std::move(decoded.data);
results[i].success = decoded.success;
results[i].error = std::move(decoded.error);
}
return results;
}
testSem.reset();
std::vector<BlteFrameLayout> layouts(entries.size());
std::vector<std::unique_ptr<interfaces::TimelineSemaphore>> sems(entries.size());
std::vector<std::vector<FrameDecodeResult>> perFileFrameResults(entries.size());
std::vector<std::atomic<bool>> fileFailed(entries.size());
for (auto& f : fileFailed)
f.store(false, std::memory_order_relaxed);
std::vector<std::shared_ptr<utils::JobGroup>> decodeGroups(entries.size());
for (size_t i = 0; i < entries.size(); ++i) {
auto& blob = entries[i].blteData;
auto& layout = layouts[i];
layout = blteParseFrameLayout(blob);
if (!layout.valid) {
results[i].error = std::move(layout.error);
continue;
}
if (layout.frames.size() <= 1) {
results[i] = decodeSingleFrameBlte(blob, layout, keys);
continue;
}
u32 const frameCount = u32(layout.frames.size());
sems[i] = pool->createTimelineSemaphore();
perFileFrameResults[i].resize(frameCount);
decodeGroups[i] = std::make_shared<utils::JobGroup>();
decodeGroups[i]->add(frameCount);
auto framesDone = sems[i]->next();
auto assemblyDone = sems[i]->next();
decodeGroups[i]->signalOnComplete(sems[i].get(), framesDone);
for (u32 j = 0; j < frameCount; ++j) {
interfaces::WorkerTask task;
task.fn = [&entries, &perFileFrameResults, &fileFailed, &layouts, &decodeGroups, keys,
i, j]() {
if (fileFailed[i].load(std::memory_order_relaxed)) {
decodeGroups[i]->done();
return;
}
auto& blob = entries[i].blteData;
auto& layout = layouts[i];
auto payload = blob.subspan(layout.offsets[j], layout.frames[j].compressedSize);
perFileFrameResults[i][j] =
decodeFramePayload(payload, layout.frames[j].uncompressedSize, keys);
if (!perFileFrameResults[i][j].success)
fileFailed[i].store(true, std::memory_order_relaxed);
decodeGroups[i]->done();
};
pool->submit(task);
}
interfaces::WorkerTask assemblyTask;
assemblyTask.fn = [&results, &perFileFrameResults, &fileFailed, i]() {
if (fileFailed[i].load(std::memory_order_relaxed)) {
for (auto& fr : perFileFrameResults[i]) {
if (!fr.error.empty()) {
results[i].error = std::move(fr.error);
break;
}
}
if (results[i].error.empty())
results[i].error = "frame decode failed";
return;
}
size_t totalSize = 0;
for (auto& fr : perFileFrameResults[i])
totalSize += fr.data.size();
results[i].data.reserve(totalSize);
for (auto& fr : perFileFrameResults[i]) {
results[i].data.insert(results[i].data.end(),
std::make_move_iterator(fr.data.begin()),
std::make_move_iterator(fr.data.end()));
fr.data.clear();
fr.data.shrink_to_fit();
}
results[i].success = true;
};
assemblyTask.waitSemaphore = sems[i].get();
assemblyTask.waitValue = framesDone;
assemblyTask.signalSemaphore = sems[i].get();
assemblyTask.signalValue = assemblyDone;
pool->submit(assemblyTask);
}
for (size_t i = 0; i < entries.size(); ++i) {
if (sems[i]) {
sems[i]->wait(2);
}
}
return results;
}
std::vector<u8> blteEncode(std::span<const u8> rawData, const BlteEncodeOptions& opts,
interfaces::WorkerPool* pool) {
u32 frameSize = opts.frameSize;
if (frameSize == 0)
frameSize = kDefaultBlteFrameSize;
size_t const frameCount = rawData.empty() ? 1 : (rawData.size() + frameSize - 1) / frameSize;
struct EncodedFrame {
std::vector<u8> data; u32 uncompressedSize;
std::array<u8, 16> hash;
};
std::vector<EncodedFrame> encodedFrames(frameCount);
auto encodeOneFrame = [&](size_t i) {
size_t const start = i * frameSize;
size_t const end = std::min(start + (size_t)frameSize, rawData.size());
auto chunk = rawData.subspan(start, end - start);
EncodedFrame& ef = encodedFrames[i];
ef.uncompressedSize = u32(chunk.size());
if (opts.compress && !chunk.empty()) {
auto compressed = storages::common::zlibCompress(chunk);
if (!compressed.empty() && compressed.size() + 1 < chunk.size() + 1) {
ef.data.resize(1 + compressed.size());
ef.data[0] = BlteFrameMode::kZlib;
std::memcpy(ef.data.data() + 1, compressed.data(), compressed.size());
} else {
ef.data.resize(1 + chunk.size());
ef.data[0] = BlteFrameMode::kRaw;
std::memcpy(ef.data.data() + 1, chunk.data(), chunk.size());
}
} else {
ef.data.resize(1 + chunk.size());
ef.data[0] = BlteFrameMode::kRaw;
if (!chunk.empty())
std::memcpy(ef.data.data() + 1, chunk.data(), chunk.size());
}
ef.hash = storages::common::md5Hash(ef.data);
};
if (rawData.empty()) {
encodedFrames[0].uncompressedSize = 0;
encodedFrames[0].data = {BlteFrameMode::kRaw}; encodedFrames[0].hash = storages::common::md5Hash(encodedFrames[0].data);
} else if (pool && frameCount > 1) {
utils::JobGroup jobGroup;
jobGroup.add(frameCount);
for (size_t i = 0; i < frameCount; ++i) {
interfaces::WorkerTask task;
task.fn = [&, i]() {
encodeOneFrame(i);
jobGroup.done();
};
pool->submit(task);
}
jobGroup.wait();
} else {
for (size_t i = 0; i < frameCount; ++i)
encodeOneFrame(i);
}
std::vector<u8> output;
if (frameCount <= 1) {
output.resize(kBlteMinHeaderSize + encodedFrames[0].data.size());
writeBE32(output.data(), kBlteMagic);
writeBE32(output.data() + 4, 0); std::memcpy(output.data() + kBlteMinHeaderSize, encodedFrames[0].data.data(),
encodedFrames[0].data.size());
} else {
u32 const tableSize = u32(frameCount) * kBlteFrameTableEntrySize;
u32 const headerSize = 0x0C + tableSize;
size_t totalFrameData = 0;
for (auto& ef : encodedFrames)
totalFrameData += ef.data.size();
output.resize(headerSize + totalFrameData);
writeBE32(output.data(), kBlteMagic);
writeBE32(output.data() + 4, headerSize);
u8* hdr = output.data() + kBlteMinHeaderSize;
hdr[0] = kBlteMultiFrameFlags;
hdr[1] = u8((frameCount >> 16) & 0xFF);
hdr[2] = u8((frameCount >> 8) & 0xFF);
hdr[3] = u8(frameCount & 0xFF);
u8* table = hdr + 4;
for (size_t i = 0; i < frameCount; ++i) {
u8* entry = table + i * kBlteFrameTableEntrySize;
writeBE32(entry, u32(encodedFrames[i].data.size()));
writeBE32(entry + 4, encodedFrames[i].uncompressedSize);
std::memcpy(entry + 8, encodedFrames[i].hash.data(), 16);
}
u8* dest = output.data() + headerSize;
for (auto& ef : encodedFrames) {
std::memcpy(dest, ef.data.data(), ef.data.size());
dest += ef.data.size();
}
}
return output;
}
}