updater: Add Zstandard for compressed downloads

Using zstd reduces the download size for updates by about 2/3.
This commit is contained in:
derrod
2023-03-07 15:34:27 -05:00
committed by Ryan Foster
parent f4e90a5bd7
commit 787c5f67a8
4 changed files with 140 additions and 6 deletions
+11 -2
View File
@@ -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")
+71
View File
@@ -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<uint8_t> 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<uint8_t> 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;
}
+44 -4
View File
@@ -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)
+14
View File
@@ -37,6 +37,7 @@
#include <zlib.h>
#include <ctype.h>
#include <blake2.h>
#include <zstd.h>
#include <string>
@@ -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<HINTERNET, FreeWinHttpHandle>;
/* ------------------------------------------------------------------------ */
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; }
};