From a9589d2d7d0d14e803fa4d44ff609fb82479d49c Mon Sep 17 00:00:00 2001 From: ddidderr Date: Sun, 19 Jul 2026 09:43:03 +0200 Subject: [PATCH] feat(compress): move stream initialization policy into Rust Move the transparent ZSTD_CCtx_init_compressStream2 initialization policy into Rust. Rust now owns the ordered local-dictionary, prefix, parameter-resolution, pledged-size, worker-selection, and ordinary-buffering decisions through a scalar projection and explicit callbacks. Keep the private CCtx and parameter layouts, allocator and trace state, MT context lifecycle, codec operations, reset behavior, and mutation details in C. The C shim therefore remains the ABI and private-state boundary while the high-level stream setup flow is testable in Rust without duplicating those layouts. Test Plan: - cargo test --manifest-path rust/Cargo.toml --all-targets -- --test-threads=1 - cargo test --manifest-path rust/cli/Cargo.toml --all-targets -- --test-threads=1 - run the legacy Rust feature matrix and all six library/CLI clippy gates with -D warnings - run lib and program native rebuilds plus test-cli-tests, test-rust-lib-smoke, and test-zstd with make -j1 - run fuzzer, zstream, and decode-corpus stress gates serially with ulimit -v 41943040 Commit is intentionally unsigned because GPG pinentry hangs in this non-interactive environment. --- lib/compress/zstd_compress.c | 516 ++++++++++++++++++++----- rust/src/zstd_compress.rs | 705 +++++++++++++++++++++++++++++++++++ 2 files changed, 1134 insertions(+), 87 deletions(-) diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index 868b7f55e..6caf369c1 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -272,6 +272,131 @@ typedef char ZSTD_rust_compress_stream_state_layout[ && sizeof(ZSTD_rust_compressStreamState) == 15 * sizeof(void*) + 2 * sizeof(int) + 2 * sizeof(size_t)) ? 1 : -1]; + +/* Rust owns only the high-level transparent initialization policy. The + * parameter object, dictionary/context layouts, allocator, trace state, and + * codec entry points stay behind these scalar/opaque callbacks. */ +typedef struct { + const void* prefixDict; + size_t prefixDictSize; + int prefixDictContentType; + const void* cdict; + int cdictIsLocal; + int cdictCompressionLevel; + size_t cdictDictContentSize; +} ZSTD_rust_compressStreamInitDictionaryState; +typedef size_t (*ZSTD_rust_compressStreamInitLocalDict_f)(void* context); +typedef void (*ZSTD_rust_compressStreamInitRefreshCDict_f)( + void* context, ZSTD_rust_compressStreamInitDictionaryState* state); +typedef void (*ZSTD_rust_compressStreamInitClearPrefix_f)(void* context); +typedef void (*ZSTD_rust_compressStreamInitAssertDictionaries_f)( + void* context, const void* prefixDict); +typedef void (*ZSTD_rust_compressStreamInitSetLevel_f)(void* params, int level); +typedef void (*ZSTD_rust_compressStreamInitDebug_f)(void* context); +typedef U64 (*ZSTD_rust_compressStreamInitGetPledged_f)(void* context); +typedef void (*ZSTD_rust_compressStreamInitSetPledged_f)(void* context, size_t inSize); +typedef int (*ZSTD_rust_compressStreamInitGetCParamMode_f)( + void* params, const void* cdict, U64 pledgedSrcSize); +typedef void (*ZSTD_rust_compressStreamInitBuildCParams_f)( + void* params, U64 pledgedSrcSize, size_t dictSize, int mode); +typedef void (*ZSTD_rust_compressStreamInitResolveParams_f)( + void* params, int operation); +typedef unsigned (*ZSTD_rust_compressStreamInitGetNbWorkers_f)(void* params); +typedef void (*ZSTD_rust_compressStreamInitSetNbWorkers_f)( + void* params, unsigned nbWorkers); +typedef int (*ZSTD_rust_compressStreamInitHasExtSeqProd_f)(void* params); +typedef void (*ZSTD_rust_compressStreamInitTrace_f)(void* context); +typedef void* (*ZSTD_rust_compressStreamInitGetMTContext_f)(void* context); +typedef size_t (*ZSTD_rust_compressStreamInitCreateMTContext_f)( + void* context, unsigned nbWorkers); +typedef size_t (*ZSTD_rust_compressStreamInitMT_f)( + void* context, void* mtctx, + const ZSTD_rust_compressStreamInitDictionaryState* dictionaries, + void* params, U64 pledgedSrcSize); +typedef void (*ZSTD_rust_compressStreamInitCommitMT_f)( + void* context, + const ZSTD_rust_compressStreamInitDictionaryState* dictionaries, + void* params); +typedef void (*ZSTD_rust_compressStreamInitCheckCParams_f)(void* params); +typedef size_t (*ZSTD_rust_compressStreamInitBegin_f)( + void* context, + const ZSTD_rust_compressStreamInitDictionaryState* dictionaries, + void* params, U64 pledgedSrcSize); +typedef void (*ZSTD_rust_compressStreamInitAssertOrdinary_f)(void* context); +typedef int (*ZSTD_rust_compressStreamInitGetBufferMode_f)(void* context); +typedef size_t (*ZSTD_rust_compressStreamInitGetBlockSize_f)(void* context); +typedef void (*ZSTD_rust_compressStreamInitCommitOrdinary_f)( + void* context, size_t inBuffTarget); +typedef struct { + void* callbackContext; + void* params; + ZSTD_rust_compressStreamInitDictionaryState* dictionaries; + int endOp; + size_t inSize; + int multithreaded; + size_t mtJobSizeMin; + ZSTD_rust_compressStreamInitLocalDict_f initLocalDict; + ZSTD_rust_compressStreamInitRefreshCDict_f refreshCDict; + ZSTD_rust_compressStreamInitClearPrefix_f clearPrefix; + ZSTD_rust_compressStreamInitAssertDictionaries_f assertDictionaries; + ZSTD_rust_compressStreamInitSetLevel_f setCompressionLevel; + ZSTD_rust_compressStreamInitDebug_f debugInit; + ZSTD_rust_compressStreamInitGetPledged_f getPledgedSrcSizePlusOne; + ZSTD_rust_compressStreamInitSetPledged_f setPledgedSrcSize; + ZSTD_rust_compressStreamInitGetCParamMode_f getCParamMode; + ZSTD_rust_compressStreamInitBuildCParams_f buildCParams; + ZSTD_rust_compressStreamInitResolveParams_f resolveParams; + ZSTD_rust_compressStreamInitGetNbWorkers_f getNbWorkers; + ZSTD_rust_compressStreamInitSetNbWorkers_f setNbWorkers; + ZSTD_rust_compressStreamInitHasExtSeqProd_f hasExtSeqProd; + ZSTD_rust_compressStreamInitTrace_f traceBegin; + ZSTD_rust_compressStreamInitGetMTContext_f getMTContext; + ZSTD_rust_compressStreamInitCreateMTContext_f createMTContext; + ZSTD_rust_compressStreamInitMT_f initMT; + ZSTD_rust_compressStreamInitCommitMT_f commitMT; + ZSTD_rust_compressStreamInitCheckCParams_f checkCParams; + ZSTD_rust_compressStreamInitBegin_f compressBegin; + ZSTD_rust_compressStreamInitAssertOrdinary_f assertOrdinary; + ZSTD_rust_compressStreamInitGetBufferMode_f getBufferMode; + ZSTD_rust_compressStreamInitGetBlockSize_f getBlockSize; + ZSTD_rust_compressStreamInitCommitOrdinary_f commitOrdinary; +} ZSTD_rust_compressStreamInitState; +typedef char ZSTD_rust_compress_stream_init_dictionary_layout[ + (offsetof(ZSTD_rust_compressStreamInitDictionaryState, prefixDict) == 0 + && offsetof(ZSTD_rust_compressStreamInitDictionaryState, prefixDictSize) + == sizeof(void*) + && offsetof(ZSTD_rust_compressStreamInitDictionaryState, prefixDictContentType) + == 2 * sizeof(void*) + && offsetof(ZSTD_rust_compressStreamInitDictionaryState, cdict) + == 3 * sizeof(void*) + && offsetof(ZSTD_rust_compressStreamInitDictionaryState, cdictIsLocal) + == 4 * sizeof(void*) + && offsetof(ZSTD_rust_compressStreamInitDictionaryState, cdictCompressionLevel) + == 4 * sizeof(void*) + sizeof(int) + && offsetof(ZSTD_rust_compressStreamInitDictionaryState, cdictDictContentSize) + == 4 * sizeof(void*) + 2 * sizeof(int) + && sizeof(ZSTD_rust_compressStreamInitDictionaryState) + == (sizeof(void*) == 8 ? 48 : 28)) + ? 1 : -1]; +typedef char ZSTD_rust_compress_stream_init_state_layout[ + (offsetof(ZSTD_rust_compressStreamInitState, callbackContext) == 0 + && offsetof(ZSTD_rust_compressStreamInitState, params) == sizeof(void*) + && offsetof(ZSTD_rust_compressStreamInitState, dictionaries) + == 2 * sizeof(void*) + && offsetof(ZSTD_rust_compressStreamInitState, endOp) + == 3 * sizeof(void*) + && offsetof(ZSTD_rust_compressStreamInitState, inSize) + == 4 * sizeof(void*) + && offsetof(ZSTD_rust_compressStreamInitState, multithreaded) + == 4 * sizeof(void*) + sizeof(size_t) + && offsetof(ZSTD_rust_compressStreamInitState, mtJobSizeMin) + == 5 * sizeof(void*) + sizeof(size_t) + && sizeof(ZSTD_rust_compressStreamInitState) + == (sizeof(void*) == 8 ? 256 : 128)) + ? 1 : -1]; +size_t ZSTD_rust_compressStreamInit( + const ZSTD_rust_compressStreamInitState* state); + /* The target-sized block body only needs this narrow projection of ZSTD_CCtx. * Matchfinder/window state, sequence-store construction, and outer repeat-mode * cleanup remain in C. */ @@ -4740,100 +4865,317 @@ static size_t ZSTD_checkBufferStability(ZSTD_CCtx const* cctx, * Otherwise, it's ignored. * @return: 0 on success, or a ZSTD_error code otherwise. */ +enum { + ZSTD_RUST_INIT_RESOLVE_BLOCK_SPLITTER = 0, + ZSTD_RUST_INIT_RESOLVE_LDM = 1, + ZSTD_RUST_INIT_RESOLVE_ROW_MATCH_FINDER = 2, + ZSTD_RUST_INIT_RESOLVE_VALIDATE_SEQUENCES = 3, + ZSTD_RUST_INIT_RESOLVE_MAX_BLOCK_SIZE = 4, + ZSTD_RUST_INIT_RESOLVE_EXTERNAL_REPCODE_SEARCH = 5 +}; + +static size_t ZSTD_rust_compressStreamInit_localDict(void* context) +{ + return ZSTD_initLocalDict((ZSTD_CCtx*)context); +} + +static void ZSTD_rust_compressStreamInit_refreshCDict( + void* context, ZSTD_rust_compressStreamInitDictionaryState* state) +{ + ZSTD_CCtx const* const cctx = (ZSTD_CCtx const*)context; + state->cdict = cctx->cdict; + state->cdictIsLocal = cctx->localDict.cdict != NULL; + state->cdictCompressionLevel = cctx->cdict ? cctx->cdict->compressionLevel : 0; + state->cdictDictContentSize = cctx->cdict ? cctx->cdict->dictContentSize : 0; +} + +static void ZSTD_rust_compressStreamInit_clearPrefix(void* context) +{ + ZSTD_CCtx* const cctx = (ZSTD_CCtx*)context; + ZSTD_memset(&cctx->prefixDict, 0, sizeof(cctx->prefixDict)); +} + +static void ZSTD_rust_compressStreamInit_assertDictionaries( + void* context, const void* prefixDict) +{ + ZSTD_CCtx const* const cctx = (ZSTD_CCtx const*)context; + (void)cctx; + (void)prefixDict; + assert(prefixDict == NULL || cctx->cdict == NULL); +} + +static void ZSTD_rust_compressStreamInit_setLevel(void* params, int level) +{ + ((ZSTD_CCtx_params*)params)->compressionLevel = level; +} + +static void ZSTD_rust_compressStreamInit_debug(void* context) +{ + (void)context; + DEBUGLOG(4, "ZSTD_CCtx_init_compressStream2 : transparent init stage"); +} + +static U64 ZSTD_rust_compressStreamInit_getPledged(void* context) +{ + return ((ZSTD_CCtx const*)context)->pledgedSrcSizePlusOne; +} + +static void ZSTD_rust_compressStreamInit_setPledged(void* context, size_t inSize) +{ + ((ZSTD_CCtx*)context)->pledgedSrcSizePlusOne = inSize + 1; +} + +static int ZSTD_rust_compressStreamInit_getCParamMode( + void* params, const void* cdict, U64 pledgedSrcSize) +{ + return (int)ZSTD_getCParamMode( + (const ZSTD_CDict*)cdict, (const ZSTD_CCtx_params*)params, pledgedSrcSize); +} + +static void ZSTD_rust_compressStreamInit_buildCParams( + void* params, U64 pledgedSrcSize, size_t dictSize, int mode) +{ + ZSTD_CCtx_params* const cctxParams = (ZSTD_CCtx_params*)params; + cctxParams->cParams = ZSTD_getCParamsFromCCtxParams( + cctxParams, pledgedSrcSize, dictSize, (ZSTD_CParamMode_e)mode); +} + +static void ZSTD_rust_compressStreamInit_resolveParams(void* params, int operation) +{ + ZSTD_CCtx_params* const cctxParams = (ZSTD_CCtx_params*)params; + switch (operation) { + case ZSTD_RUST_INIT_RESOLVE_BLOCK_SPLITTER: + cctxParams->postBlockSplitter = ZSTD_resolveBlockSplitterMode( + cctxParams->postBlockSplitter, &cctxParams->cParams); + break; + case ZSTD_RUST_INIT_RESOLVE_LDM: + cctxParams->ldmParams.enableLdm = ZSTD_resolveEnableLdm( + cctxParams->ldmParams.enableLdm, &cctxParams->cParams); + break; + case ZSTD_RUST_INIT_RESOLVE_ROW_MATCH_FINDER: + cctxParams->useRowMatchFinder = ZSTD_resolveRowMatchFinderMode( + cctxParams->useRowMatchFinder, &cctxParams->cParams); + break; + case ZSTD_RUST_INIT_RESOLVE_VALIDATE_SEQUENCES: + cctxParams->validateSequences = ZSTD_resolveExternalSequenceValidation( + cctxParams->validateSequences); + break; + case ZSTD_RUST_INIT_RESOLVE_MAX_BLOCK_SIZE: + cctxParams->maxBlockSize = ZSTD_resolveMaxBlockSize(cctxParams->maxBlockSize); + break; + case ZSTD_RUST_INIT_RESOLVE_EXTERNAL_REPCODE_SEARCH: + cctxParams->searchForExternalRepcodes = ZSTD_resolveExternalRepcodeSearch( + cctxParams->searchForExternalRepcodes, cctxParams->compressionLevel); + break; + default: + assert(0); + break; + } +} + +static unsigned ZSTD_rust_compressStreamInit_getNbWorkers(void* params) +{ + return ((ZSTD_CCtx_params const*)params)->nbWorkers; +} + +static void ZSTD_rust_compressStreamInit_setNbWorkers(void* params, unsigned nbWorkers) +{ + ((ZSTD_CCtx_params*)params)->nbWorkers = nbWorkers; +} + +static int ZSTD_rust_compressStreamInit_hasExtSeqProd(void* params) +{ + return ZSTD_hasExtSeqProd((ZSTD_CCtx_params const*)params); +} + +static void ZSTD_rust_compressStreamInit_trace(void* context) +{ +#if ZSTD_TRACE + ZSTD_CCtx* const cctx = (ZSTD_CCtx*)context; + cctx->traceCtx = (ZSTD_trace_compress_begin != NULL) + ? ZSTD_trace_compress_begin(cctx) : 0; +#else + (void)context; +#endif +} + +static void* ZSTD_rust_compressStreamInit_getMTContext(void* context) +{ +#ifdef ZSTD_MULTITHREAD + return ((ZSTD_CCtx*)context)->mtctx; +#else + (void)context; + return NULL; +#endif +} + +static size_t ZSTD_rust_compressStreamInit_createMTContext( + void* context, unsigned nbWorkers) +{ +#ifdef ZSTD_MULTITHREAD + ZSTD_CCtx* const cctx = (ZSTD_CCtx*)context; + DEBUGLOG(4, "ZSTD_compressStream2: creating new mtctx for nbWorkers=%u", nbWorkers); + cctx->mtctx = ZSTDMT_createCCtx_advanced( + (U32)nbWorkers, cctx->customMem, cctx->pool); + RETURN_ERROR_IF(cctx->mtctx == NULL, memory_allocation, "NULL pointer!"); + return 0; +#else + (void)context; + (void)nbWorkers; + return ERROR(memory_allocation); +#endif +} + +static size_t ZSTD_rust_compressStreamInit_mt( + void* context, void* mtctx, + const ZSTD_rust_compressStreamInitDictionaryState* dictionaries, + void* params, U64 pledgedSrcSize) +{ +#ifdef ZSTD_MULTITHREAD + (void)context; + DEBUGLOG(4, "call ZSTDMT_initCStream_internal as nbWorkers=%u", + ((ZSTD_CCtx_params const*)params)->nbWorkers); + return ZSTDMT_initCStream_internal( + (ZSTDMT_CCtx*)mtctx, + dictionaries->prefixDict, dictionaries->prefixDictSize, + (ZSTD_dictContentType_e)dictionaries->prefixDictContentType, + (const ZSTD_CDict*)dictionaries->cdict, + *(const ZSTD_CCtx_params*)params, pledgedSrcSize); +#else + (void)context; + (void)mtctx; + (void)dictionaries; + (void)params; + (void)pledgedSrcSize; + return ERROR(memory_allocation); +#endif +} + +static void ZSTD_rust_compressStreamInit_commitMT( + void* context, + const ZSTD_rust_compressStreamInitDictionaryState* dictionaries, + void* params) +{ +#ifdef ZSTD_MULTITHREAD + ZSTD_CCtx* const cctx = (ZSTD_CCtx*)context; + cctx->dictID = cctx->cdict ? cctx->cdict->dictID : 0; + cctx->dictContentSize = cctx->cdict + ? cctx->cdict->dictContentSize : dictionaries->prefixDictSize; + cctx->consumedSrcSize = 0; + cctx->producedCSize = 0; + cctx->streamStage = zcss_load; + cctx->appliedParams = *(const ZSTD_CCtx_params*)params; +#else + (void)context; + (void)dictionaries; + (void)params; +#endif +} + +static void ZSTD_rust_compressStreamInit_checkCParams(void* params) +{ + (void)params; + assert(!ZSTD_isError(ZSTD_checkCParams( + ((ZSTD_CCtx_params const*)params)->cParams))); +} + +static size_t ZSTD_rust_compressStreamInit_begin( + void* context, + const ZSTD_rust_compressStreamInitDictionaryState* dictionaries, + void* params, U64 pledgedSrcSize) +{ + ZSTD_CCtx* const cctx = (ZSTD_CCtx*)context; + return ZSTD_compressBegin_internal( + cctx, + dictionaries->prefixDict, dictionaries->prefixDictSize, + (ZSTD_dictContentType_e)dictionaries->prefixDictContentType, ZSTD_dtlm_fast, + (const ZSTD_CDict*)dictionaries->cdict, + (const ZSTD_CCtx_params*)params, pledgedSrcSize, ZSTDb_buffered); +} + +static void ZSTD_rust_compressStreamInit_assertOrdinary(void* context) +{ + (void)context; + assert(((ZSTD_CCtx const*)context)->appliedParams.nbWorkers == 0); +} + +static int ZSTD_rust_compressStreamInit_getBufferMode(void* context) +{ + return (int)((ZSTD_CCtx const*)context)->appliedParams.inBufferMode; +} + +static size_t ZSTD_rust_compressStreamInit_getBlockSize(void* context) +{ + return ((ZSTD_CCtx const*)context)->blockSizeMax; +} + +static void ZSTD_rust_compressStreamInit_commitOrdinary( + void* context, size_t inBuffTarget) +{ + ZSTD_CCtx* const cctx = (ZSTD_CCtx*)context; + cctx->inToCompress = 0; + cctx->inBuffPos = 0; + cctx->inBuffTarget = inBuffTarget; + cctx->outBuffContentSize = cctx->outBuffFlushedSize = 0; + cctx->streamStage = zcss_load; + cctx->frameEnded = 0; +} + static size_t ZSTD_CCtx_init_compressStream2(ZSTD_CCtx* cctx, ZSTD_EndDirective endOp, size_t inSize) { ZSTD_CCtx_params params = cctx->requestedParams; ZSTD_prefixDict const prefixDict = cctx->prefixDict; - FORWARD_IF_ERROR( ZSTD_initLocalDict(cctx) , ""); /* Init the local dict if present. */ - ZSTD_memset(&cctx->prefixDict, 0, sizeof(cctx->prefixDict)); /* single usage */ - assert(prefixDict.dict==NULL || cctx->cdict==NULL); /* only one can be set */ - if (cctx->cdict && !cctx->localDict.cdict) { - /* Let the cdict's compression level take priority over the requested params. - * But do not take the cdict's compression level if the "cdict" is actually a localDict - * generated from ZSTD_initLocalDict(). - */ - params.compressionLevel = cctx->cdict->compressionLevel; - } - DEBUGLOG(4, "ZSTD_CCtx_init_compressStream2 : transparent init stage"); - if (endOp == ZSTD_e_end) cctx->pledgedSrcSizePlusOne = inSize + 1; /* auto-determine pledgedSrcSize */ - - { size_t const dictSize = prefixDict.dict - ? prefixDict.dictSize - : (cctx->cdict ? cctx->cdict->dictContentSize : 0); - ZSTD_CParamMode_e const mode = ZSTD_getCParamMode(cctx->cdict, ¶ms, cctx->pledgedSrcSizePlusOne - 1); - params.cParams = ZSTD_getCParamsFromCCtxParams( - ¶ms, cctx->pledgedSrcSizePlusOne-1, - dictSize, mode); - } - - params.postBlockSplitter = ZSTD_resolveBlockSplitterMode(params.postBlockSplitter, ¶ms.cParams); - params.ldmParams.enableLdm = ZSTD_resolveEnableLdm(params.ldmParams.enableLdm, ¶ms.cParams); - params.useRowMatchFinder = ZSTD_resolveRowMatchFinderMode(params.useRowMatchFinder, ¶ms.cParams); - params.validateSequences = ZSTD_resolveExternalSequenceValidation(params.validateSequences); - params.maxBlockSize = ZSTD_resolveMaxBlockSize(params.maxBlockSize); - params.searchForExternalRepcodes = ZSTD_resolveExternalRepcodeSearch(params.searchForExternalRepcodes, params.compressionLevel); - + ZSTD_rust_compressStreamInitDictionaryState dictionaries = { + prefixDict.dict, + prefixDict.dictSize, + (int)prefixDict.dictContentType, + NULL, + 0, + 0, + 0 + }; + ZSTD_rust_compressStreamInitState state = { + cctx, + ¶ms, + &dictionaries, + (int)endOp, + inSize, #ifdef ZSTD_MULTITHREAD - /* If external matchfinder is enabled, make sure to fail before checking job size (for consistency) */ - RETURN_ERROR_IF( - ZSTD_hasExtSeqProd(¶ms) && params.nbWorkers >= 1, - parameter_combination_unsupported, - "External sequence producer isn't supported with nbWorkers >= 1" - ); - - if ((cctx->pledgedSrcSizePlusOne-1) <= ZSTDMT_JOBSIZE_MIN) { - params.nbWorkers = 0; /* do not invoke multi-threading when src size is too small */ - } - if (params.nbWorkers > 0) { -# if ZSTD_TRACE - cctx->traceCtx = (ZSTD_trace_compress_begin != NULL) ? ZSTD_trace_compress_begin(cctx) : 0; -# endif - /* mt context creation */ - if (cctx->mtctx == NULL) { - DEBUGLOG(4, "ZSTD_compressStream2: creating new mtctx for nbWorkers=%u", - params.nbWorkers); - cctx->mtctx = ZSTDMT_createCCtx_advanced((U32)params.nbWorkers, cctx->customMem, cctx->pool); - RETURN_ERROR_IF(cctx->mtctx == NULL, memory_allocation, "NULL pointer!"); - } - /* mt compression */ - DEBUGLOG(4, "call ZSTDMT_initCStream_internal as nbWorkers=%u", params.nbWorkers); - FORWARD_IF_ERROR( ZSTDMT_initCStream_internal( - cctx->mtctx, - prefixDict.dict, prefixDict.dictSize, prefixDict.dictContentType, - cctx->cdict, params, cctx->pledgedSrcSizePlusOne-1) , ""); - cctx->dictID = cctx->cdict ? cctx->cdict->dictID : 0; - cctx->dictContentSize = cctx->cdict ? cctx->cdict->dictContentSize : prefixDict.dictSize; - cctx->consumedSrcSize = 0; - cctx->producedCSize = 0; - cctx->streamStage = zcss_load; - cctx->appliedParams = params; - } else -#endif /* ZSTD_MULTITHREAD */ - { U64 const pledgedSrcSize = cctx->pledgedSrcSizePlusOne - 1; - assert(!ZSTD_isError(ZSTD_checkCParams(params.cParams))); - FORWARD_IF_ERROR( ZSTD_compressBegin_internal(cctx, - prefixDict.dict, prefixDict.dictSize, prefixDict.dictContentType, ZSTD_dtlm_fast, - cctx->cdict, - ¶ms, pledgedSrcSize, - ZSTDb_buffered) , ""); - assert(cctx->appliedParams.nbWorkers == 0); - cctx->inToCompress = 0; - cctx->inBuffPos = 0; - if (cctx->appliedParams.inBufferMode == ZSTD_bm_buffered) { - /* for small input: avoid automatic flush on reaching end of block, since - * it would require to add a 3-bytes null block to end frame - */ - cctx->inBuffTarget = cctx->blockSizeMax + (cctx->blockSizeMax == pledgedSrcSize); - } else { - cctx->inBuffTarget = 0; - } - cctx->outBuffContentSize = cctx->outBuffFlushedSize = 0; - cctx->streamStage = zcss_load; - cctx->frameEnded = 0; - } - return 0; + 1, + ZSTDMT_JOBSIZE_MIN, +#else + 0, + 0, +#endif + ZSTD_rust_compressStreamInit_localDict, + ZSTD_rust_compressStreamInit_refreshCDict, + ZSTD_rust_compressStreamInit_clearPrefix, + ZSTD_rust_compressStreamInit_assertDictionaries, + ZSTD_rust_compressStreamInit_setLevel, + ZSTD_rust_compressStreamInit_debug, + ZSTD_rust_compressStreamInit_getPledged, + ZSTD_rust_compressStreamInit_setPledged, + ZSTD_rust_compressStreamInit_getCParamMode, + ZSTD_rust_compressStreamInit_buildCParams, + ZSTD_rust_compressStreamInit_resolveParams, + ZSTD_rust_compressStreamInit_getNbWorkers, + ZSTD_rust_compressStreamInit_setNbWorkers, + ZSTD_rust_compressStreamInit_hasExtSeqProd, + ZSTD_rust_compressStreamInit_trace, + ZSTD_rust_compressStreamInit_getMTContext, + ZSTD_rust_compressStreamInit_createMTContext, + ZSTD_rust_compressStreamInit_mt, + ZSTD_rust_compressStreamInit_commitMT, + ZSTD_rust_compressStreamInit_checkCParams, + ZSTD_rust_compressStreamInit_begin, + ZSTD_rust_compressStreamInit_assertOrdinary, + ZSTD_rust_compressStreamInit_getBufferMode, + ZSTD_rust_compressStreamInit_getBlockSize, + ZSTD_rust_compressStreamInit_commitOrdinary + }; + return ZSTD_rust_compressStreamInit(&state); } /* @return provides a minimum amount of data remaining to be flushed from internal buffers diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index b3a09dbd7..a2900af0c 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -696,6 +696,291 @@ const _: () = { ); }; +const ZSTD_RUST_INIT_RESOLVE_BLOCK_SPLITTER: c_int = 0; +const ZSTD_RUST_INIT_RESOLVE_LDM: c_int = 1; +const ZSTD_RUST_INIT_RESOLVE_ROW_MATCH_FINDER: c_int = 2; +const ZSTD_RUST_INIT_RESOLVE_VALIDATE_SEQUENCES: c_int = 3; +const ZSTD_RUST_INIT_RESOLVE_MAX_BLOCK_SIZE: c_int = 4; +const ZSTD_RUST_INIT_RESOLVE_EXTERNAL_REPCODE_SEARCH: c_int = 5; + +type CompressStreamInitLocalDictFn = unsafe extern "C" fn(*mut c_void) -> usize; +type CompressStreamInitRefreshCDictFn = + unsafe extern "C" fn(*mut c_void, *mut ZSTD_rust_compressStreamInitDictionaryState); +type CompressStreamInitClearPrefixFn = unsafe extern "C" fn(*mut c_void); +type CompressStreamInitAssertDictionariesFn = unsafe extern "C" fn(*mut c_void, *const c_void); +type CompressStreamInitSetLevelFn = unsafe extern "C" fn(*mut c_void, c_int); +type CompressStreamInitDebugFn = unsafe extern "C" fn(*mut c_void); +type CompressStreamInitGetPledgedFn = unsafe extern "C" fn(*mut c_void) -> u64; +type CompressStreamInitSetPledgedFn = unsafe extern "C" fn(*mut c_void, usize); +type CompressStreamInitGetCParamModeFn = + unsafe extern "C" fn(*mut c_void, *const c_void, u64) -> c_int; +type CompressStreamInitBuildCParamsFn = unsafe extern "C" fn(*mut c_void, u64, usize, c_int); +type CompressStreamInitResolveParamsFn = unsafe extern "C" fn(*mut c_void, c_int); +type CompressStreamInitGetNbWorkersFn = unsafe extern "C" fn(*mut c_void) -> c_uint; +type CompressStreamInitSetNbWorkersFn = unsafe extern "C" fn(*mut c_void, c_uint); +type CompressStreamInitHasExtSeqProdFn = unsafe extern "C" fn(*mut c_void) -> c_int; +type CompressStreamInitTraceFn = unsafe extern "C" fn(*mut c_void); +type CompressStreamInitGetMTContextFn = unsafe extern "C" fn(*mut c_void) -> *mut c_void; +type CompressStreamInitCreateMTContextFn = unsafe extern "C" fn(*mut c_void, c_uint) -> usize; +type CompressStreamInitMTFn = unsafe extern "C" fn( + *mut c_void, + *mut c_void, + *const ZSTD_rust_compressStreamInitDictionaryState, + *mut c_void, + u64, +) -> usize; +type CompressStreamInitCommitMTFn = unsafe extern "C" fn( + *mut c_void, + *const ZSTD_rust_compressStreamInitDictionaryState, + *mut c_void, +); +type CompressStreamInitCheckCParamsFn = unsafe extern "C" fn(*mut c_void); +type CompressStreamInitBeginFn = unsafe extern "C" fn( + *mut c_void, + *const ZSTD_rust_compressStreamInitDictionaryState, + *mut c_void, + u64, +) -> usize; +type CompressStreamInitAssertOrdinaryFn = unsafe extern "C" fn(*mut c_void); +type CompressStreamInitGetBufferModeFn = unsafe extern "C" fn(*mut c_void) -> c_int; +type CompressStreamInitGetBlockSizeFn = unsafe extern "C" fn(*mut c_void) -> usize; +type CompressStreamInitCommitOrdinaryFn = unsafe extern "C" fn(*mut c_void, usize); + +/// Scalar dictionary snapshot used by transparent stream initialization. +/// +/// The prefix is populated by C before local-dictionary initialization, which +/// preserves the original single-use snapshot order. C refreshes only the +/// CDict fields after local-dictionary initialization because that operation +/// may create the local CDict. +#[repr(C)] +pub struct ZSTD_rust_compressStreamInitDictionaryState { + prefix_dict: *const c_void, + prefix_dict_size: usize, + prefix_dict_content_type: c_int, + cdict: *const c_void, + cdict_is_local: c_int, + cdict_compression_level: c_int, + cdict_dict_content_size: usize, +} + +const _: () = { + assert!(offset_of!(ZSTD_rust_compressStreamInitDictionaryState, prefix_dict) == 0); + assert!( + offset_of!( + ZSTD_rust_compressStreamInitDictionaryState, + prefix_dict_size + ) == size_of::() + ); + assert!( + offset_of!( + ZSTD_rust_compressStreamInitDictionaryState, + prefix_dict_content_type + ) == 2 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compressStreamInitDictionaryState, cdict) == 3 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compressStreamInitDictionaryState, cdict_is_local) + == 4 * size_of::() + ); + assert!( + offset_of!( + ZSTD_rust_compressStreamInitDictionaryState, + cdict_compression_level + ) == 4 * size_of::() + size_of::() + ); + assert!( + offset_of!( + ZSTD_rust_compressStreamInitDictionaryState, + cdict_dict_content_size + ) == 4 * size_of::() + 2 * size_of::() + ); + assert!( + size_of::() + == if size_of::() == 8 { 48 } else { 28 } + ); +}; + +/// Projection for `ZSTD_CCtx_init_compressStream2()`. +/// +/// Rust owns the ordering and branch policy. The C callback slots retain +/// private parameter/context layouts, local dictionary storage, allocators, +/// trace setup, MT construction/init, and codec/reset operations. +#[repr(C)] +pub struct ZSTD_rust_compressStreamInitState { + callback_context: *mut c_void, + params: *mut c_void, + dictionaries: *mut ZSTD_rust_compressStreamInitDictionaryState, + end_op: c_int, + in_size: usize, + multithreaded: c_int, + mt_job_size_min: usize, + init_local_dict: CompressStreamInitLocalDictFn, + refresh_cdict: CompressStreamInitRefreshCDictFn, + clear_prefix: CompressStreamInitClearPrefixFn, + assert_dictionaries: CompressStreamInitAssertDictionariesFn, + set_compression_level: CompressStreamInitSetLevelFn, + debug_init: CompressStreamInitDebugFn, + get_pledged_src_size_plus_one: CompressStreamInitGetPledgedFn, + set_pledged_src_size: CompressStreamInitSetPledgedFn, + get_cparam_mode: CompressStreamInitGetCParamModeFn, + build_cparams: CompressStreamInitBuildCParamsFn, + resolve_params: CompressStreamInitResolveParamsFn, + get_nb_workers: CompressStreamInitGetNbWorkersFn, + set_nb_workers: CompressStreamInitSetNbWorkersFn, + has_ext_seq_prod: CompressStreamInitHasExtSeqProdFn, + trace_begin: CompressStreamInitTraceFn, + get_mt_context: CompressStreamInitGetMTContextFn, + create_mt_context: CompressStreamInitCreateMTContextFn, + init_mt: CompressStreamInitMTFn, + commit_mt: CompressStreamInitCommitMTFn, + check_cparams: CompressStreamInitCheckCParamsFn, + compress_begin: CompressStreamInitBeginFn, + assert_ordinary: CompressStreamInitAssertOrdinaryFn, + get_buffer_mode: CompressStreamInitGetBufferModeFn, + get_block_size: CompressStreamInitGetBlockSizeFn, + commit_ordinary: CompressStreamInitCommitOrdinaryFn, +} + +const _: () = { + assert!(offset_of!(ZSTD_rust_compressStreamInitState, callback_context) == 0); + assert!(offset_of!(ZSTD_rust_compressStreamInitState, params) == size_of::()); + assert!(offset_of!(ZSTD_rust_compressStreamInitState, dictionaries) == 2 * size_of::()); + assert!(offset_of!(ZSTD_rust_compressStreamInitState, end_op) == 3 * size_of::()); + assert!(offset_of!(ZSTD_rust_compressStreamInitState, in_size) == 4 * size_of::()); + assert!( + offset_of!(ZSTD_rust_compressStreamInitState, multithreaded) + == 4 * size_of::() + size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compressStreamInitState, mt_job_size_min) + == 5 * size_of::() + size_of::() + ); + assert!( + size_of::() + == if size_of::() == 8 { 256 } else { 128 } + ); +}; + +#[inline] +unsafe fn compress_stream_init_body_with(state: &ZSTD_rust_compressStreamInitState) -> usize { + if state.params.is_null() || state.dictionaries.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + + let result = unsafe { (state.init_local_dict)(state.callback_context) }; + if ERR_isError(result) { + return result; + } + + unsafe { (state.refresh_cdict)(state.callback_context, state.dictionaries) }; + let dictionaries = unsafe { &*state.dictionaries }; + unsafe { (state.clear_prefix)(state.callback_context) }; + unsafe { (state.assert_dictionaries)(state.callback_context, dictionaries.prefix_dict) }; + if !dictionaries.cdict.is_null() && dictionaries.cdict_is_local == 0 { + unsafe { + (state.set_compression_level)(state.params, dictionaries.cdict_compression_level) + }; + } + unsafe { (state.debug_init)(state.callback_context) }; + + if state.end_op == ZSTD_E_END { + unsafe { (state.set_pledged_src_size)(state.callback_context, state.in_size) }; + } + let pledged_src_size_plus_one = + unsafe { (state.get_pledged_src_size_plus_one)(state.callback_context) }; + let pledged_src_size = pledged_src_size_plus_one.wrapping_sub(1); + let dict_size = if !dictionaries.prefix_dict.is_null() { + dictionaries.prefix_dict_size + } else if !dictionaries.cdict.is_null() { + dictionaries.cdict_dict_content_size + } else { + 0 + }; + let mode = + unsafe { (state.get_cparam_mode)(state.params, dictionaries.cdict, pledged_src_size) }; + unsafe { + (state.build_cparams)(state.params, pledged_src_size, dict_size, mode); + (state.resolve_params)(state.params, ZSTD_RUST_INIT_RESOLVE_BLOCK_SPLITTER); + (state.resolve_params)(state.params, ZSTD_RUST_INIT_RESOLVE_LDM); + (state.resolve_params)(state.params, ZSTD_RUST_INIT_RESOLVE_ROW_MATCH_FINDER); + (state.resolve_params)(state.params, ZSTD_RUST_INIT_RESOLVE_VALIDATE_SEQUENCES); + (state.resolve_params)(state.params, ZSTD_RUST_INIT_RESOLVE_MAX_BLOCK_SIZE); + (state.resolve_params)(state.params, ZSTD_RUST_INIT_RESOLVE_EXTERNAL_REPCODE_SEARCH); + } + + if state.multithreaded != 0 { + let has_ext_seq_prod = unsafe { (state.has_ext_seq_prod)(state.params) }; + let mut nb_workers = unsafe { (state.get_nb_workers)(state.params) }; + if has_ext_seq_prod != 0 && nb_workers >= 1 { + return ERROR(ZstdErrorCode::ParameterCombinationUnsupported); + } + if pledged_src_size <= state.mt_job_size_min as u64 { + nb_workers = 0; + unsafe { (state.set_nb_workers)(state.params, 0) }; + } + if nb_workers > 0 { + unsafe { (state.trace_begin)(state.callback_context) }; + if unsafe { (state.get_mt_context)(state.callback_context) }.is_null() { + let result = + unsafe { (state.create_mt_context)(state.callback_context, nb_workers) }; + if ERR_isError(result) { + return result; + } + } + let result = unsafe { + (state.init_mt)( + state.callback_context, + (state.get_mt_context)(state.callback_context), + dictionaries, + state.params, + pledged_src_size, + ) + }; + if ERR_isError(result) { + return result; + } + unsafe { (state.commit_mt)(state.callback_context, dictionaries, state.params) }; + return 0; + } + } + + unsafe { (state.check_cparams)(state.params) }; + let result = unsafe { + (state.compress_begin)( + state.callback_context, + dictionaries, + state.params, + pledged_src_size, + ) + }; + if ERR_isError(result) { + return result; + } + unsafe { (state.assert_ordinary)(state.callback_context) }; + let buffer_mode = unsafe { (state.get_buffer_mode)(state.callback_context) }; + let block_size = unsafe { (state.get_block_size)(state.callback_context) }; + let in_buff_target = if buffer_mode == ZSTD_BM_BUFFERED { + block_size.wrapping_add(usize::from(block_size as u64 == pledged_src_size)) + } else { + 0 + }; + unsafe { (state.commit_ordinary)(state.callback_context, in_buff_target) }; + 0 +} + +/// Rust-owned high-level policy for transparent stream initialization. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_compressStreamInit( + state: *const ZSTD_rust_compressStreamInitState, +) -> usize { + if state.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + unsafe { compress_stream_init_body_with(&*state) } +} + #[inline] unsafe fn stream_limit_copy( dst: *mut c_void, @@ -4270,6 +4555,426 @@ mod tests { const ZSTD_BTOPT: c_int = 7; const ZSTD_BTULTRA2: c_int = 9; + #[derive(Default)] + struct CompressStreamInitTestContext { + events: Vec<&'static str>, + local_result: usize, + create_result: usize, + mt_result: usize, + begin_result: usize, + pledged: u64, + cdict: *const c_void, + cdict_is_local: c_int, + cdict_compression_level: c_int, + cdict_dict_content_size: usize, + expected_prefix: *const c_void, + nb_workers: c_uint, + has_ext_seq_prod: c_int, + mt_context: *mut c_void, + buffer_mode: c_int, + block_size: usize, + ordinary_target: usize, + compression_level: c_int, + } + + unsafe fn compress_stream_init_test_context( + context: *mut c_void, + ) -> &'static mut CompressStreamInitTestContext { + unsafe { &mut *context.cast::() } + } + + unsafe extern "C" fn compress_stream_init_test_local_dict(context: *mut c_void) -> usize { + let context = unsafe { compress_stream_init_test_context(context) }; + context.events.push("local-dict"); + context.local_result + } + + unsafe extern "C" fn compress_stream_init_test_refresh_cdict( + context: *mut c_void, + dictionaries: *mut ZSTD_rust_compressStreamInitDictionaryState, + ) { + let context = unsafe { compress_stream_init_test_context(context) }; + context.events.push("refresh-cdict"); + let dictionaries = unsafe { &mut *dictionaries }; + dictionaries.cdict = context.cdict; + dictionaries.cdict_is_local = context.cdict_is_local; + dictionaries.cdict_compression_level = context.cdict_compression_level; + dictionaries.cdict_dict_content_size = context.cdict_dict_content_size; + } + + unsafe extern "C" fn compress_stream_init_test_clear_prefix(context: *mut c_void) { + unsafe { compress_stream_init_test_context(context) } + .events + .push("clear-prefix"); + } + + unsafe extern "C" fn compress_stream_init_test_assert_dictionaries( + context: *mut c_void, + prefix: *const c_void, + ) { + let context = unsafe { compress_stream_init_test_context(context) }; + context.events.push("assert-dictionaries"); + assert_eq!(prefix, context.expected_prefix); + } + + unsafe extern "C" fn compress_stream_init_test_set_level(params: *mut c_void, level: c_int) { + let context = unsafe { &mut *params.cast::() }; + context.events.push("set-level"); + context.compression_level = level; + } + + unsafe extern "C" fn compress_stream_init_test_debug(context: *mut c_void) { + unsafe { compress_stream_init_test_context(context) } + .events + .push("debug"); + } + + unsafe extern "C" fn compress_stream_init_test_get_pledged(context: *mut c_void) -> u64 { + let context = unsafe { compress_stream_init_test_context(context) }; + context.events.push("get-pledged"); + context.pledged + } + + unsafe extern "C" fn compress_stream_init_test_set_pledged( + context: *mut c_void, + in_size: usize, + ) { + let context = unsafe { compress_stream_init_test_context(context) }; + context.events.push("set-pledged"); + context.pledged = (in_size as u64).wrapping_add(1); + } + + unsafe extern "C" fn compress_stream_init_test_get_cparam_mode( + params: *mut c_void, + _cdict: *const c_void, + _pledged: u64, + ) -> c_int { + unsafe { compress_stream_init_test_context(params) } + .events + .push("get-cparam-mode"); + 0 + } + + unsafe extern "C" fn compress_stream_init_test_build_cparams( + params: *mut c_void, + _pledged: u64, + _dict_size: usize, + _mode: c_int, + ) { + unsafe { compress_stream_init_test_context(params) } + .events + .push("build-cparams"); + } + + unsafe extern "C" fn compress_stream_init_test_resolve_params( + params: *mut c_void, + operation: c_int, + ) { + let context = unsafe { compress_stream_init_test_context(params) }; + context.events.push(match operation { + ZSTD_RUST_INIT_RESOLVE_BLOCK_SPLITTER => "resolve:block-splitter", + ZSTD_RUST_INIT_RESOLVE_LDM => "resolve:ldm", + ZSTD_RUST_INIT_RESOLVE_ROW_MATCH_FINDER => "resolve:row-match-finder", + ZSTD_RUST_INIT_RESOLVE_VALIDATE_SEQUENCES => "resolve:validate-sequences", + ZSTD_RUST_INIT_RESOLVE_MAX_BLOCK_SIZE => "resolve:max-block-size", + ZSTD_RUST_INIT_RESOLVE_EXTERNAL_REPCODE_SEARCH => "resolve:external-repcodes", + _ => "resolve:unknown", + }); + } + + unsafe extern "C" fn compress_stream_init_test_get_nb_workers(params: *mut c_void) -> c_uint { + let context = unsafe { compress_stream_init_test_context(params) }; + context.events.push("get-workers"); + context.nb_workers + } + + unsafe extern "C" fn compress_stream_init_test_set_nb_workers( + params: *mut c_void, + nb_workers: c_uint, + ) { + let context = unsafe { compress_stream_init_test_context(params) }; + context.events.push("set-workers"); + context.nb_workers = nb_workers; + } + + unsafe extern "C" fn compress_stream_init_test_has_ext_seq_prod(params: *mut c_void) -> c_int { + let context = unsafe { compress_stream_init_test_context(params) }; + context.events.push("has-ext-seq-prod"); + context.has_ext_seq_prod + } + + unsafe extern "C" fn compress_stream_init_test_trace(context: *mut c_void) { + unsafe { compress_stream_init_test_context(context) } + .events + .push("trace"); + } + + unsafe extern "C" fn compress_stream_init_test_get_mt_context( + context: *mut c_void, + ) -> *mut c_void { + let context = unsafe { compress_stream_init_test_context(context) }; + context.events.push("get-mt-context"); + context.mt_context + } + + unsafe extern "C" fn compress_stream_init_test_create_mt_context( + context: *mut c_void, + _nb_workers: c_uint, + ) -> usize { + let context = unsafe { compress_stream_init_test_context(context) }; + context.events.push("create-mt-context"); + if !ERR_isError(context.create_result) { + context.mt_context = ptr::dangling_mut::(); + } + context.create_result + } + + unsafe extern "C" fn compress_stream_init_test_mt( + context: *mut c_void, + _mt_context: *mut c_void, + _dictionaries: *const ZSTD_rust_compressStreamInitDictionaryState, + _params: *mut c_void, + _pledged: u64, + ) -> usize { + let context = unsafe { compress_stream_init_test_context(context) }; + context.events.push("init-mt"); + context.mt_result + } + + unsafe extern "C" fn compress_stream_init_test_commit_mt( + context: *mut c_void, + _dictionaries: *const ZSTD_rust_compressStreamInitDictionaryState, + _params: *mut c_void, + ) { + unsafe { compress_stream_init_test_context(context) } + .events + .push("commit-mt"); + } + + unsafe extern "C" fn compress_stream_init_test_check_cparams(params: *mut c_void) { + unsafe { compress_stream_init_test_context(params) } + .events + .push("check-cparams"); + } + + unsafe extern "C" fn compress_stream_init_test_begin( + context: *mut c_void, + _dictionaries: *const ZSTD_rust_compressStreamInitDictionaryState, + _params: *mut c_void, + _pledged: u64, + ) -> usize { + let context = unsafe { compress_stream_init_test_context(context) }; + context.events.push("compress-begin"); + context.begin_result + } + + unsafe extern "C" fn compress_stream_init_test_assert_ordinary(context: *mut c_void) { + unsafe { compress_stream_init_test_context(context) } + .events + .push("assert-ordinary"); + } + + unsafe extern "C" fn compress_stream_init_test_get_buffer_mode(context: *mut c_void) -> c_int { + let context = unsafe { compress_stream_init_test_context(context) }; + context.events.push("get-buffer-mode"); + context.buffer_mode + } + + unsafe extern "C" fn compress_stream_init_test_get_block_size(context: *mut c_void) -> usize { + let context = unsafe { compress_stream_init_test_context(context) }; + context.events.push("get-block-size"); + context.block_size + } + + unsafe extern "C" fn compress_stream_init_test_commit_ordinary( + context: *mut c_void, + in_buff_target: usize, + ) { + let context = unsafe { compress_stream_init_test_context(context) }; + context.events.push("commit-ordinary"); + context.ordinary_target = in_buff_target; + } + + fn compress_stream_init_test_state( + context: &mut CompressStreamInitTestContext, + dictionaries: &mut ZSTD_rust_compressStreamInitDictionaryState, + end_op: c_int, + in_size: usize, + multithreaded: c_int, + ) -> ZSTD_rust_compressStreamInitState { + ZSTD_rust_compressStreamInitState { + callback_context: (context as *mut CompressStreamInitTestContext).cast(), + params: (context as *mut CompressStreamInitTestContext).cast(), + dictionaries, + end_op, + in_size, + multithreaded, + mt_job_size_min: 10, + init_local_dict: compress_stream_init_test_local_dict, + refresh_cdict: compress_stream_init_test_refresh_cdict, + clear_prefix: compress_stream_init_test_clear_prefix, + assert_dictionaries: compress_stream_init_test_assert_dictionaries, + set_compression_level: compress_stream_init_test_set_level, + debug_init: compress_stream_init_test_debug, + get_pledged_src_size_plus_one: compress_stream_init_test_get_pledged, + set_pledged_src_size: compress_stream_init_test_set_pledged, + get_cparam_mode: compress_stream_init_test_get_cparam_mode, + build_cparams: compress_stream_init_test_build_cparams, + resolve_params: compress_stream_init_test_resolve_params, + get_nb_workers: compress_stream_init_test_get_nb_workers, + set_nb_workers: compress_stream_init_test_set_nb_workers, + has_ext_seq_prod: compress_stream_init_test_has_ext_seq_prod, + trace_begin: compress_stream_init_test_trace, + get_mt_context: compress_stream_init_test_get_mt_context, + create_mt_context: compress_stream_init_test_create_mt_context, + init_mt: compress_stream_init_test_mt, + commit_mt: compress_stream_init_test_commit_mt, + check_cparams: compress_stream_init_test_check_cparams, + compress_begin: compress_stream_init_test_begin, + assert_ordinary: compress_stream_init_test_assert_ordinary, + get_buffer_mode: compress_stream_init_test_get_buffer_mode, + get_block_size: compress_stream_init_test_get_block_size, + commit_ordinary: compress_stream_init_test_commit_ordinary, + } + } + + #[test] + fn compress_stream_init_preserves_ordinary_callback_order_and_policy() { + let mut context = CompressStreamInitTestContext { + pledged: 0, + cdict: ptr::dangling::(), + cdict_compression_level: 17, + cdict_dict_content_size: 3, + nb_workers: 0, + buffer_mode: ZSTD_BM_BUFFERED, + block_size: 4, + ..CompressStreamInitTestContext::default() + }; + let mut dictionaries = ZSTD_rust_compressStreamInitDictionaryState { + prefix_dict: ptr::null(), + prefix_dict_size: 0, + prefix_dict_content_type: 0, + cdict: ptr::null(), + cdict_is_local: 0, + cdict_compression_level: 0, + cdict_dict_content_size: 0, + }; + let state = + compress_stream_init_test_state(&mut context, &mut dictionaries, ZSTD_E_END, 4, 0); + + let result = unsafe { ZSTD_rust_compressStreamInit(&state) }; + + assert_eq!(result, 0); + assert_eq!(context.compression_level, 17); + assert_eq!(context.pledged, 5); + assert_eq!(context.ordinary_target, 5); + assert_eq!( + context.events, + [ + "local-dict", + "refresh-cdict", + "clear-prefix", + "assert-dictionaries", + "set-level", + "debug", + "set-pledged", + "get-pledged", + "get-cparam-mode", + "build-cparams", + "resolve:block-splitter", + "resolve:ldm", + "resolve:row-match-finder", + "resolve:validate-sequences", + "resolve:max-block-size", + "resolve:external-repcodes", + "check-cparams", + "compress-begin", + "assert-ordinary", + "get-buffer-mode", + "get-block-size", + "commit-ordinary", + ] + ); + } + + #[test] + fn compress_stream_init_stops_before_later_callbacks_on_error() { + let mut context = CompressStreamInitTestContext { + local_result: ERROR(ZstdErrorCode::MemoryAllocation), + ..CompressStreamInitTestContext::default() + }; + let mut dictionaries = ZSTD_rust_compressStreamInitDictionaryState { + prefix_dict: ptr::null(), + prefix_dict_size: 0, + prefix_dict_content_type: 0, + cdict: ptr::null(), + cdict_is_local: 0, + cdict_compression_level: 0, + cdict_dict_content_size: 0, + }; + let state = + compress_stream_init_test_state(&mut context, &mut dictionaries, ZSTD_E_CONTINUE, 0, 0); + + let result = unsafe { ZSTD_rust_compressStreamInit(&state) }; + + assert_eq!(result, ERROR(ZstdErrorCode::MemoryAllocation)); + assert_eq!(context.events, ["local-dict"]); + } + + #[test] + fn compress_stream_init_preserves_mt_branch_and_init_error_order() { + let mut context = CompressStreamInitTestContext { + cdict: ptr::dangling::(), + cdict_dict_content_size: 3, + nb_workers: 2, + mt_result: ERROR(ZstdErrorCode::MemoryAllocation), + ..CompressStreamInitTestContext::default() + }; + let mut dictionaries = ZSTD_rust_compressStreamInitDictionaryState { + prefix_dict: ptr::null(), + prefix_dict_size: 0, + prefix_dict_content_type: 0, + cdict: ptr::null(), + cdict_is_local: 0, + cdict_compression_level: 0, + cdict_dict_content_size: 0, + }; + let state = + compress_stream_init_test_state(&mut context, &mut dictionaries, ZSTD_E_END, 99, 1); + + let result = unsafe { ZSTD_rust_compressStreamInit(&state) }; + + assert_eq!(result, ERROR(ZstdErrorCode::MemoryAllocation)); + assert_eq!( + context.events, + [ + "local-dict", + "refresh-cdict", + "clear-prefix", + "assert-dictionaries", + "set-level", + "debug", + "set-pledged", + "get-pledged", + "get-cparam-mode", + "build-cparams", + "resolve:block-splitter", + "resolve:ldm", + "resolve:row-match-finder", + "resolve:validate-sequences", + "resolve:max-block-size", + "resolve:external-repcodes", + "has-ext-seq-prod", + "get-workers", + "trace", + "get-mt-context", + "create-mt-context", + "get-mt-context", + "init-mt", + ] + ); + } + fn system_round_trip(compressed: &[u8]) -> Option> { let mut child = Command::new("zstd") .args(["-q", "-d", "-c"])