From 787c5f67a88ca3043112c6239732c13418c86ae2 Mon Sep 17 00:00:00 2001 From: derrod Date: Fri, 3 Mar 2023 12:34:54 +0100 Subject: [PATCH] updater: Add Zstandard for compressed downloads Using zstd reduces the download size for updates by about 2/3. --- UI/win-update/updater/CMakeLists.txt | 13 ++++- UI/win-update/updater/patch.cpp | 71 ++++++++++++++++++++++++++++ UI/win-update/updater/updater.cpp | 48 +++++++++++++++++-- UI/win-update/updater/updater.hpp | 14 ++++++ 4 files changed, 140 insertions(+), 6 deletions(-) diff --git a/UI/win-update/updater/CMakeLists.txt b/UI/win-update/updater/CMakeLists.txt index 64672fae3..cad418b00 100644 --- a/UI/win-update/updater/CMakeLists.txt +++ b/UI/win-update/updater/CMakeLists.txt @@ -12,6 +12,8 @@ if(NOT DEFINED STATIC_ZLIB_PATH OR "${STATIC_ZLIB_PATH}" STREQUAL "") return() endif() +find_package(zstd) + add_executable(updater WIN32) target_sources( @@ -40,7 +42,14 @@ if(MSVC) target_compile_options(updater PRIVATE "/utf-8") endif() -target_link_libraries(updater PRIVATE OBS::blake2 OBS::lzma ${STATIC_ZLIB_PATH} - comctl32 shell32 winhttp) +target_link_libraries( + updater + PRIVATE OBS::blake2 + OBS::lzma + zstd::libzstd_static + ${STATIC_ZLIB_PATH} + comctl32 + shell32 + winhttp) set_target_properties(updater PROPERTIES FOLDER "frontend") diff --git a/UI/win-update/updater/patch.cpp b/UI/win-update/updater/patch.cpp index b42c95f26..7c034e152 100644 --- a/UI/win-update/updater/patch.cpp +++ b/UI/win-update/updater/patch.cpp @@ -318,3 +318,74 @@ try { } catch (int code) { return code; } + +int DecompressFile(ZSTD_DCtx *ctx, const wchar_t *tempFile, size_t newSize) +try { + WinHandle hTemp; + + hTemp = CreateFile(tempFile, GENERIC_READ, 0, nullptr, OPEN_EXISTING, 0, + nullptr); + if (!hTemp.Valid()) + throw int(GetLastError()); + + /* --------------------------------- * + * read compressed data */ + + DWORD read; + DWORD compressedFileSize; + + compressedFileSize = GetFileSize(hTemp, nullptr); + if (compressedFileSize == INVALID_FILE_SIZE) + throw int(GetLastError()); + + vector oldData; + try { + oldData.resize(compressedFileSize); + } catch (...) { + throw int(-1); + } + + if (!ReadFile(hTemp, &oldData[0], compressedFileSize, &read, nullptr)) + throw int(GetLastError()); + if (read != compressedFileSize) + throw int(-1); + + /* --------------------------------- * + * decompress data */ + + vector newData; + try { + newData.resize((size_t)newSize); + } catch (...) { + throw int(-1); + } + + size_t result = ZSTD_decompressDCtx(ctx, &newData[0], newData.size(), + oldData.data(), oldData.size()); + + if (result != newSize) + throw int(-9); + if (ZSTD_isError(result)) + throw int(-10); + + /* --------------------------------- * + * overwrite temp file with new data */ + + hTemp = nullptr; + hTemp = CreateFile(tempFile, GENERIC_WRITE, 0, nullptr, CREATE_ALWAYS, + 0, nullptr); + if (!hTemp.Valid()) + throw int(GetLastError()); + + bool success; + DWORD written; + + success = !!WriteFile(hTemp, newData.data(), (DWORD)newSize, &written, + nullptr); + if (!success || written != newSize) + throw int(GetLastError()); + + return 0; +} catch (int code) { + return code; +} diff --git a/UI/win-update/updater/updater.cpp b/UI/win-update/updater/updater.cpp index 7c1d477bf..d735b2691 100644 --- a/UI/win-update/updater/updater.cpp +++ b/UI/win-update/updater/updater.cpp @@ -262,6 +262,7 @@ struct update_t { state_t state = STATE_INVALID; bool has_hash = false; bool patchable = false; + bool compressed = false; inline update_t() {} inline update_t(const update_t &from) @@ -274,7 +275,8 @@ struct update_t { fileSize(from.fileSize), state(from.state), has_hash(from.has_hash), - patchable(from.patchable) + patchable(from.patchable), + compressed(from.compressed) { memcpy(hash, from.hash, sizeof(hash)); memcpy(downloadhash, from.downloadhash, sizeof(downloadhash)); @@ -291,7 +293,8 @@ struct update_t { fileSize(from.fileSize), state(from.state), has_hash(from.has_hash), - patchable(from.patchable) + patchable(from.patchable), + compressed(from.compressed) { from.state = STATE_INVALID; @@ -330,6 +333,7 @@ struct update_t { state = from.state; has_hash = from.has_hash; patchable = from.patchable; + compressed = from.compressed; memcpy(hash, from.hash, sizeof(hash)); memcpy(downloadhash, from.downloadhash, sizeof(downloadhash)); @@ -399,6 +403,8 @@ bool DownloadWorkerThread() return false; } + ZSTDDCtx zCtx; + for (;;) { bool foundWork = false; @@ -470,6 +476,20 @@ bool DownloadWorkerThread() return 1; } + if (update.compressed && !update.patchable) { + int res = DecompressFile( + zCtx, update.tempPath.c_str(), + update.fileSize); + if (res) { + downloadThreadFailure = true; + DeleteFile(update.tempPath.c_str()); + Status(L"Update failed: Decompression " + L"failed on %s (error code %d)", + update.outputPath.c_str(), res); + return 1; + } + } + ulock.lock(); update.state = STATE_DOWNLOADED; @@ -712,6 +732,7 @@ static bool AddPackageUpdateFiles(const Json &root, size_t idx, const Json &file = files[j]; const Json &fileName = file["name"]; const Json &hash = file["hash"]; + const Json &dlHash = file["compressed_hash"]; const Json &size = file["size"]; if (!fileName.is_string()) @@ -723,22 +744,33 @@ static bool AddPackageUpdateFiles(const Json &root, size_t idx, const string &fileUTF8 = fileName.string_value(); const string &hashUTF8 = hash.string_value(); + const string &dlHashUTF8 = dlHash.string_value(); int fileSize = size.int_value(); if (hashUTF8.size() != BLAKE2_HASH_LENGTH * 2) continue; + /* The download hash may not exist if a file is uncompressed */ + + bool compressed = false; + if (dlHashUTF8.size() == BLAKE2_HASH_LENGTH * 2) + compressed = true; + /* convert strings to wide */ wchar_t sourceURL[1024]; wchar_t updateFileName[MAX_PATH]; wchar_t updateHashStr[BLAKE2_HASH_STR_LENGTH]; + wchar_t downloadHashStr[BLAKE2_HASH_STR_LENGTH]; wchar_t tempFilePath[MAX_PATH]; if (!UTF8ToWideBuf(updateFileName, fileUTF8.c_str())) continue; if (!UTF8ToWideBuf(updateHashStr, hashUTF8.c_str())) continue; + if (compressed && + !UTF8ToWideBuf(downloadHashStr, dlHashUTF8.c_str())) + continue; /* make sure paths are safe */ @@ -783,9 +815,17 @@ static bool AddPackageUpdateFiles(const Json &root, size_t idx, update.packageName = packageName; update.state = STATE_PENDING_DOWNLOAD; update.patchable = false; + update.compressed = compressed; - StringToHash(updateHashStr, update.downloadhash); - memcpy(update.hash, update.downloadhash, sizeof(update.hash)); + StringToHash(updateHashStr, update.hash); + + if (compressed) { + update.sourceURL += L".zst"; + StringToHash(downloadHashStr, update.downloadhash); + } else { + memcpy(update.downloadhash, update.hash, + sizeof(update.downloadhash)); + } update.has_hash = has_hash; if (has_hash) diff --git a/UI/win-update/updater/updater.hpp b/UI/win-update/updater/updater.hpp index 8d1adaec1..5458d6dd0 100644 --- a/UI/win-update/updater/updater.hpp +++ b/UI/win-update/updater/updater.hpp @@ -37,6 +37,7 @@ #include #include #include +#include #include @@ -91,6 +92,7 @@ void StringToHash(const wchar_t *in, BYTE *out); bool CalculateFileHash(const wchar_t *path, BYTE *hash); int ApplyPatch(LPCTSTR patchFile, LPCTSTR targetFile); +int DecompressFile(ZSTD_DCtx *ctx, LPCTSTR tempFile, size_t newSize); extern HWND hwndMain; extern HCRYPTPROV hProvider; @@ -109,3 +111,15 @@ typedef struct { void FreeWinHttpHandle(HINTERNET handle); using HttpHandle = CustomHandle; + +/* ------------------------------------------------------------------------ */ + +class ZSTDDCtx { + ZSTD_DCtx *ctx = nullptr; + +public: + inline ZSTDDCtx() { ctx = ZSTD_createDCtx(); } + inline ~ZSTDDCtx() { ZSTD_freeDCtx(ctx); } + inline operator ZSTD_DCtx *() const { return ctx; } + inline ZSTD_DCtx *get() const { return ctx; } +};